diff --git a/.devcontainer/init-setup.sh b/.devcontainer/init-setup.sh index 9a96131ee5..f84c58ecb9 100644 --- a/.devcontainer/init-setup.sh +++ b/.devcontainer/init-setup.sh @@ -20,8 +20,8 @@ set -e # - To build the project, use: # ./gradlew build # -# - For running pre-commit hooks (if configured), use: -# pre-commit run --all-files +# - To run the lint/format/secret checks, use: +# task pre-commit # # Make sure you are in the project root directory after this script executes. # ============================================================================= @@ -70,6 +70,6 @@ echo "" echo " To build the project: " echo -e "\e[34m gradle build\e[0m" echo "" -echo " To run pre-commit hooks (if configured):" -echo -e "\e[34m pre-commit run --all-files -c .pre-commit-config.yaml\e[0m" +echo " To run the lint/format/secret checks:" +echo -e "\e[34m task pre-commit\e[0m" echo "==================================================================" diff --git a/.github/aur/stirling-pdf-desktop/PKGBUILD b/.github/aur/stirling-pdf-desktop/PKGBUILD index 706bfaacc1..de60325133 100644 --- a/.github/aur/stirling-pdf-desktop/PKGBUILD +++ b/.github/aur/stirling-pdf-desktop/PKGBUILD @@ -1,6 +1,6 @@ # Maintainer: Stirling PDF Inc pkgname=stirling-pdf-desktop -pkgver=2.12.0 +pkgver=2.13.0 pkgrel=1 pkgdesc="Locally hosted, web-based PDF manipulation tool (Tauri desktop app, official Stirling PDF Inc build)" arch=('x86_64') diff --git a/.github/aur/stirling-pdf-server-bin/PKGBUILD b/.github/aur/stirling-pdf-server-bin/PKGBUILD index 0622fed1d3..3853bb6256 100644 --- a/.github/aur/stirling-pdf-server-bin/PKGBUILD +++ b/.github/aur/stirling-pdf-server-bin/PKGBUILD @@ -1,6 +1,6 @@ # Maintainer: Stirling PDF Inc pkgname=stirling-pdf-server-bin -pkgver=2.12.0 +pkgver=2.13.0 pkgrel=1 pkgdesc="Locally hosted, web-based PDF manipulation tool (server JAR, prebuilt)" arch=('any') diff --git a/.github/config/.files.yaml b/.github/config/.files.yaml index 5343246c25..2d617f8f20 100644 --- a/.github/config/.files.yaml +++ b/.github/config/.files.yaml @@ -1,6 +1,8 @@ build: &build - build.gradle - app/(common|core|proprietary)/build.gradle + - Taskfile.yml + - .taskfiles/backend.yml openapi: &openapi - *build @@ -38,6 +40,9 @@ project: &project - frontend/** - docker/** - scripts/RestartHelper.java + - Taskfile.yml + - .taskfiles/backend.yml + - .taskfiles/docker.yml - scripts/db-migration/** - .github/workflows/db-migration-test.yml @@ -55,6 +60,9 @@ frontend: &frontend - scripts/summarize_type3_signatures.py - scripts/type3_to_cff.py - scripts/update_type3_library.py + - Taskfile.yml + - .taskfiles/frontend.yml + - .taskfiles/e2e.yml # Files that affect the Tauri desktop bundle. Gate the multi-OS Tauri build # job on changes to any of these. @@ -66,6 +74,8 @@ tauri: &tauri - frontend/package-lock.json - frontend/editor/vite.config.ts - .github/workflows/tauri-build.yml + - Taskfile.yml + - .taskfiles/desktop.yml # Files that affect the AI engine (Python tool models, fixers, tests). Gate # the engine validation job on changes to engine sources or to the Java @@ -74,6 +84,8 @@ engine: &engine - engine/** - app/(common|core|proprietary)/src/main/java/** - .github/workflows/ai-engine.yml + - Taskfile.yml + - .taskfiles/engine.yml licenses-frontend: &licenses-frontend - ".github/workflows/frontend-backend-licenses-update.yml" @@ -102,4 +114,4 @@ proprietary: &proprietary - configs/settings.yml.template - build.gradle - app/proprietary/build.gradle - - .github/workflows/build-enterprise.yml \ No newline at end of file + - .github/workflows/build-enterprise.yml diff --git a/.github/scripts/check_language_toml.py b/.github/scripts/check_language_toml.py index b931bac52b..638951693e 100644 --- a/.github/scripts/check_language_toml.py +++ b/.github/scripts/check_language_toml.py @@ -13,7 +13,7 @@ Usage: """ # Sample for Windows: -# python .github/scripts/check_language_toml.py --reference-file frontend/editor/public/locales/en-GB/translation.toml --branch "" --files frontend/editor/public/locales/de-DE/translation.toml frontend/editor/public/locales/fr-FR/translation.toml +# python .github/scripts/check_language_toml.py --reference-file frontend/editor/public/locales/en-US/translation.toml --branch "" --files frontend/editor/public/locales/de-DE/translation.toml frontend/editor/public/locales/fr-FR/translation.toml import argparse import glob @@ -211,7 +211,7 @@ def check_for_differences(reference_file, file_list, branch, actor): ) continue - if basename_current_file == basename_reference_file and locale_dir == "en-GB": + if basename_current_file == basename_reference_file and locale_dir == "en-US": continue if ( @@ -308,7 +308,7 @@ def check_for_differences(reference_file, file_list, branch, actor): report.append("## ❌ Overall Check Status: **_Failed_**") report.append("") report.append( - f"@{actor} please check your translation if it conforms to the standard. Follow the format of [en-GB/translation.toml](https://github.com/Stirling-Tools/Stirling-PDF/blob/main/frontend/editor/public/locales/en-GB/translation.toml)" + f"@{actor} please check your translation if it conforms to the standard. Follow the format of [en-US/translation.toml](https://github.com/Stirling-Tools/Stirling-PDF/blob/main/frontend/editor/public/locales/en-US/translation.toml)" ) else: report.append("## ✅ Overall Check Status: **_Success_**") diff --git a/.github/scripts/requirements_pre_commit.in b/.github/scripts/requirements_pre_commit.in deleted file mode 100644 index 416634f528..0000000000 --- a/.github/scripts/requirements_pre_commit.in +++ /dev/null @@ -1 +0,0 @@ -pre-commit diff --git a/.github/scripts/requirements_pre_commit.txt b/.github/scripts/requirements_pre_commit.txt deleted file mode 100644 index a476a5268b..0000000000 --- a/.github/scripts/requirements_pre_commit.txt +++ /dev/null @@ -1,121 +0,0 @@ -# -# This file is autogenerated by pip-compile with Python 3.12 -# by the following command: -# -# pip-compile --generate-hashes --output-file='.github\scripts\requirements_pre_commit.txt' --strip-extras '.github\scripts\requirements_pre_commit.in' -# -cfgv==3.5.0 \ - --hash=sha256:a8dc6b26ad22ff227d2634a65cb388215ce6cc96bbcc5cfde7641ae87e8dacc0 \ - --hash=sha256:d5b1034354820651caa73ede66a6294d6e95c1b00acc5e9b098e917404669132 - # via pre-commit -distlib==0.4.0 \ - --hash=sha256:9659f7d87e46584a30b5780e43ac7a2143098441670ff0a49d5f9034c54a6c16 \ - --hash=sha256:feec40075be03a04501a973d81f633735b4b69f98b05450592310c0f401a4e0d - # via virtualenv -filelock==3.29.0 \ - --hash=sha256:69974355e960702e789734cb4871f884ea6fe50bd8404051a3530bc07809cf90 \ - --hash=sha256:96f5f6344709aa1572bbf631c640e4ebeeb519e08da902c39a001882f30ac258 - # via - # python-discovery - # virtualenv -identify==2.6.19 \ - --hash=sha256:20e6a87f786f768c092a721ad107fc9df0eb89347be9396cadf3f4abbd1fb78a \ - --hash=sha256:6be5020c38fcb07da56c53733538a3081ea5aa70d36a156f83044bfbf9173842 - # via pre-commit -nodeenv==1.10.0 \ - --hash=sha256:5bb13e3eed2923615535339b3c620e76779af4cb4c6a90deccc9e36b274d3827 \ - --hash=sha256:996c191ad80897d076bdfba80a41994c2b47c68e224c542b48feba42ba00f8bb - # via pre-commit -platformdirs==4.9.6 \ - --hash=sha256:3bfa75b0ad0db84096ae777218481852c0ebc6c727b3168c1b9e0118e458cf0a \ - --hash=sha256:e61adb1d5e5cb3441b4b7710bea7e4c12250ca49439228cc1021c00dcfac0917 - # via - # python-discovery - # virtualenv -pre-commit==4.6.0 \ - --hash=sha256:718d2208cef53fdc38206e40524a6d4d9576d103eb16f0fec11c875e7716e9d9 \ - --hash=sha256:e2cf246f7299edcabcf15f9b0571fdce06058527f0a06535068a86d38089f29b - # via -r .github/scripts/requirements_pre_commit.in -python-discovery==1.2.2 \ - --hash=sha256:876e9c57139eb757cb5878cbdd9ae5379e5d96266c99ef731119e04fffe533bb \ - --hash=sha256:e1ae95d9af875e78f15e19aed0c6137ab1bb49c200f21f5061786490c9585c7a - # via virtualenv -pyyaml==6.0.3 \ - --hash=sha256:00c4bdeba853cc34e7dd471f16b4114f4162dc03e6b7afcc2128711f0eca823c \ - --hash=sha256:0150219816b6a1fa26fb4699fb7daa9caf09eb1999f3b70fb6e786805e80375a \ - --hash=sha256:02893d100e99e03eda1c8fd5c441d8c60103fd175728e23e431db1b589cf5ab3 \ - --hash=sha256:02ea2dfa234451bbb8772601d7b8e426c2bfa197136796224e50e35a78777956 \ - --hash=sha256:0f29edc409a6392443abf94b9cf89ce99889a1dd5376d94316ae5145dfedd5d6 \ - --hash=sha256:10892704fc220243f5305762e276552a0395f7beb4dbf9b14ec8fd43b57f126c \ - --hash=sha256:16249ee61e95f858e83976573de0f5b2893b3677ba71c9dd36b9cf8be9ac6d65 \ - --hash=sha256:1d37d57ad971609cf3c53ba6a7e365e40660e3be0e5175fa9f2365a379d6095a \ - --hash=sha256:1ebe39cb5fc479422b83de611d14e2c0d3bb2a18bbcb01f229ab3cfbd8fee7a0 \ - --hash=sha256:214ed4befebe12df36bcc8bc2b64b396ca31be9304b8f59e25c11cf94a4c033b \ - --hash=sha256:2283a07e2c21a2aa78d9c4442724ec1eb15f5e42a723b99cb3d822d48f5f7ad1 \ - --hash=sha256:22ba7cfcad58ef3ecddc7ed1db3409af68d023b7f940da23c6c2a1890976eda6 \ - --hash=sha256:27c0abcb4a5dac13684a37f76e701e054692a9b2d3064b70f5e4eb54810553d7 \ - --hash=sha256:28c8d926f98f432f88adc23edf2e6d4921ac26fb084b028c733d01868d19007e \ - --hash=sha256:2e71d11abed7344e42a8849600193d15b6def118602c4c176f748e4583246007 \ - --hash=sha256:34d5fcd24b8445fadc33f9cf348c1047101756fd760b4dacb5c3e99755703310 \ - --hash=sha256:37503bfbfc9d2c40b344d06b2199cf0e96e97957ab1c1b546fd4f87e53e5d3e4 \ - --hash=sha256:3c5677e12444c15717b902a5798264fa7909e41153cdf9ef7ad571b704a63dd9 \ - --hash=sha256:3ff07ec89bae51176c0549bc4c63aa6202991da2d9a6129d7aef7f1407d3f295 \ - --hash=sha256:41715c910c881bc081f1e8872880d3c650acf13dfa8214bad49ed4cede7c34ea \ - --hash=sha256:418cf3f2111bc80e0933b2cd8cd04f286338bb88bdc7bc8e6dd775ebde60b5e0 \ - --hash=sha256:44edc647873928551a01e7a563d7452ccdebee747728c1080d881d68af7b997e \ - --hash=sha256:4a2e8cebe2ff6ab7d1050ecd59c25d4c8bd7e6f400f5f82b96557ac0abafd0ac \ - --hash=sha256:4ad1906908f2f5ae4e5a8ddfce73c320c2a1429ec52eafd27138b7f1cbe341c9 \ - --hash=sha256:501a031947e3a9025ed4405a168e6ef5ae3126c59f90ce0cd6f2bfc477be31b7 \ - --hash=sha256:5190d403f121660ce8d1d2c1bb2ef1bd05b5f68533fc5c2ea899bd15f4399b35 \ - --hash=sha256:5498cd1645aa724a7c71c8f378eb29ebe23da2fc0d7a08071d89469bf1d2defb \ - --hash=sha256:5cf4e27da7e3fbed4d6c3d8e797387aaad68102272f8f9752883bc32d61cb87b \ - --hash=sha256:5e0b74767e5f8c593e8c9b5912019159ed0533c70051e9cce3e8b6aa699fcd69 \ - --hash=sha256:5ed875a24292240029e4483f9d4a4b8a1ae08843b9c54f43fcc11e404532a8a5 \ - --hash=sha256:5fcd34e47f6e0b794d17de1b4ff496c00986e1c83f7ab2fb8fcfe9616ff7477b \ - --hash=sha256:5fdec68f91a0c6739b380c83b951e2c72ac0197ace422360e6d5a959d8d97b2c \ - --hash=sha256:6344df0d5755a2c9a276d4473ae6b90647e216ab4757f8426893b5dd2ac3f369 \ - --hash=sha256:64386e5e707d03a7e172c0701abfb7e10f0fb753ee1d773128192742712a98fd \ - --hash=sha256:652cb6edd41e718550aad172851962662ff2681490a8a711af6a4d288dd96824 \ - --hash=sha256:66291b10affd76d76f54fad28e22e51719ef9ba22b29e1d7d03d6777a9174198 \ - --hash=sha256:66e1674c3ef6f541c35191caae2d429b967b99e02040f5ba928632d9a7f0f065 \ - --hash=sha256:6adc77889b628398debc7b65c073bcb99c4a0237b248cacaf3fe8a557563ef6c \ - --hash=sha256:79005a0d97d5ddabfeeea4cf676af11e647e41d81c9a7722a193022accdb6b7c \ - --hash=sha256:7c6610def4f163542a622a73fb39f534f8c101d690126992300bf3207eab9764 \ - --hash=sha256:7f047e29dcae44602496db43be01ad42fc6f1cc0d8cd6c83d342306c32270196 \ - --hash=sha256:8098f252adfa6c80ab48096053f512f2321f0b998f98150cea9bd23d83e1467b \ - --hash=sha256:850774a7879607d3a6f50d36d04f00ee69e7fc816450e5f7e58d7f17f1ae5c00 \ - --hash=sha256:8d1fab6bb153a416f9aeb4b8763bc0f22a5586065f86f7664fc23339fc1c1fac \ - --hash=sha256:8da9669d359f02c0b91ccc01cac4a67f16afec0dac22c2ad09f46bee0697eba8 \ - --hash=sha256:8dc52c23056b9ddd46818a57b78404882310fb473d63f17b07d5c40421e47f8e \ - --hash=sha256:9149cad251584d5fb4981be1ecde53a1ca46c891a79788c0df828d2f166bda28 \ - --hash=sha256:93dda82c9c22deb0a405ea4dc5f2d0cda384168e466364dec6255b293923b2f3 \ - --hash=sha256:96b533f0e99f6579b3d4d4995707cf36df9100d67e0c8303a0c55b27b5f99bc5 \ - --hash=sha256:9c57bb8c96f6d1808c030b1687b9b5fb476abaa47f0db9c0101f5e9f394e97f4 \ - --hash=sha256:9c7708761fccb9397fe64bbc0395abcae8c4bf7b0eac081e12b809bf47700d0b \ - --hash=sha256:9f3bfb4965eb874431221a3ff3fdcddc7e74e3b07799e0e84ca4a0f867d449bf \ - --hash=sha256:a33284e20b78bd4a18c8c2282d549d10bc8408a2a7ff57653c0cf0b9be0afce5 \ - --hash=sha256:a80cb027f6b349846a3bf6d73b5e95e782175e52f22108cfa17876aaeff93702 \ - --hash=sha256:b30236e45cf30d2b8e7b3e85881719e98507abed1011bf463a8fa23e9c3e98a8 \ - --hash=sha256:b3bc83488de33889877a0f2543ade9f70c67d66d9ebb4ac959502e12de895788 \ - --hash=sha256:b865addae83924361678b652338317d1bd7e79b1f4596f96b96c77a5a34b34da \ - --hash=sha256:b8bb0864c5a28024fac8a632c443c87c5aa6f215c0b126c449ae1a150412f31d \ - --hash=sha256:ba1cc08a7ccde2d2ec775841541641e4548226580ab850948cbfda66a1befcdc \ - --hash=sha256:bdb2c67c6c1390b63c6ff89f210c8fd09d9a1217a465701eac7316313c915e4c \ - --hash=sha256:c1ff362665ae507275af2853520967820d9124984e0f7466736aea23d8611fba \ - --hash=sha256:c2514fceb77bc5e7a2f7adfaa1feb2fb311607c9cb518dbc378688ec73d8292f \ - --hash=sha256:c3355370a2c156cffb25e876646f149d5d68f5e0a3ce86a5084dd0b64a994917 \ - --hash=sha256:c458b6d084f9b935061bc36216e8a69a7e293a2f1e68bf956dcd9e6cbcd143f5 \ - --hash=sha256:d0eae10f8159e8fdad514efdc92d74fd8d682c933a6dd088030f3834bc8e6b26 \ - --hash=sha256:d76623373421df22fb4cf8817020cbb7ef15c725b9d5e45f17e189bfc384190f \ - --hash=sha256:ebc55a14a21cb14062aa4162f906cd962b28e2e9ea38f9b4391244cd8de4ae0b \ - --hash=sha256:eda16858a3cab07b80edaf74336ece1f986ba330fdb8ee0d6c0d68fe82bc96be \ - --hash=sha256:ee2922902c45ae8ccada2c5b501ab86c36525b883eff4255313a253a3160861c \ - --hash=sha256:efd7b85f94a6f21e4932043973a7ba2613b059c4a000551892ac9f1d11f5baf3 \ - --hash=sha256:f7057c9a337546edc7973c0d3ba84ddcdf0daa14533c2065749c9075001090e6 \ - --hash=sha256:fa160448684b4e94d80416c0fa4aac48967a969efe22931448d853ada8baf926 \ - --hash=sha256:fc09d0aa354569bc501d4e787133afc08552722d3ab34836a80547331bb5d4a0 - # via pre-commit -virtualenv==21.2.4 \ - --hash=sha256:29d21e941795206138d0f22f4e45ff7050e5da6c6472299fb7103318763861ac \ - --hash=sha256:b294ef68192638004d72524ce7ef303e9d0cf5a44c95ce2e54a7500a6381cada - # via pre-commit diff --git a/.github/scripts/verify-updater-signatures.py b/.github/scripts/verify-updater-signatures.py new file mode 100644 index 0000000000..a1624448e1 --- /dev/null +++ b/.github/scripts/verify-updater-signatures.py @@ -0,0 +1,83 @@ +#!/usr/bin/env python3 +"""Verify Tauri updater .sig files against plugins.updater.pubkey in tauri.conf.json. + +Usage: verify-updater-signatures.py [tauri.conf.json] +""" + +import binascii +import sys +import json +import base64 +import hashlib +from pathlib import Path +from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PublicKey +from cryptography.exceptions import InvalidSignature + +ART_ROOT = Path(sys.argv[1]) +CONF = Path( + sys.argv[2] if len(sys.argv) > 2 else "frontend/editor/src-tauri/tauri.conf.json" +) + + +def load_pubkey(): + # tauri pubkey = base64 of a minisign .pub file; last line is base64 of + # [2 algo][8 key-id][32 ed25519 public key]. + raw = json.loads(CONF.read_text())["plugins"]["updater"]["pubkey"] + blob = base64.b64decode(base64.b64decode(raw).decode().splitlines()[-1]) + return blob[2:10], Ed25519PublicKey.from_public_bytes(blob[10:]) + + +def hash_file(path: Path) -> bytes: + h = hashlib.blake2b(digest_size=64) + with path.open("rb") as f: + for chunk in iter(lambda: f.read(1 << 16), b""): + h.update(chunk) + return h.digest() + + +def verify(artifact: Path, sig_file: Path, keyid_pub, pub) -> str: + # tauri .sig = base64 of a minisign signature file (4 lines). + try: + lines = base64.b64decode(sig_file.read_text()).decode().splitlines() + sig_blob = base64.b64decode(lines[1]) + except (binascii.Error, IndexError, UnicodeDecodeError) as e: + return f"FAIL malformed sig ({type(e).__name__})" + algo, keyid, sig = sig_blob[:2], sig_blob[2:10], sig_blob[10:74] + if keyid != keyid_pub: + return f"FAIL key-id mismatch (sig {keyid.hex()} vs pub {keyid_pub.hex()})" + # 'ED' = prehashed (BLAKE2b-512), 'Ed' = legacy (raw message). + msg = hash_file(artifact) if algo == b"ED" else artifact.read_bytes() + try: + pub.verify(sig, msg) + except InvalidSignature: + return f"FAIL signature invalid (algo={algo.decode()})" + # Global signature covers sig + trusted_comment. + gc = "global-sig FAIL" + try: + tc = lines[2].split("trusted comment: ", 1)[1] + pub.verify(base64.b64decode(lines[3]), sig + tc.encode()) + gc = "global-sig OK" + except (InvalidSignature, IndexError, binascii.Error): + pass + return f"VALID (algo={algo.decode()}, keyid={keyid.hex()}, {gc})" + + +keyid_pub, pub = load_pubkey() +print(f"updater pubkey keyid={keyid_pub.hex()}\n") +sigs = sorted(ART_ROOT.rglob("*.sig")) +if not sigs: + print(f"WARN: no .sig files under {ART_ROOT} - nothing to verify") + sys.exit(0) +bad = 0 +for sig_file in sigs: + artifact = sig_file.with_suffix("") + if not artifact.exists(): + print(f" ? {sig_file.name}: artifact missing") + bad += 1 + continue + res = verify(artifact, sig_file, keyid_pub, pub) + print(f" {artifact.name}: {res}") + if not res.startswith("VALID") or "global-sig FAIL" in res: + bad += 1 +print(f"\n{'ALL SIGNATURES VALID' if bad == 0 else f'{bad} SIGNATURE(S) FAILED'}") +sys.exit(1 if bad else 0) diff --git a/.github/workflows/PR-Auto-Deploy-V2.yml b/.github/workflows/PR-Auto-Deploy-V2.yml index 71071262cc..b6ab8d34c8 100644 --- a/.github/workflows/PR-Auto-Deploy-V2.yml +++ b/.github/workflows/PR-Auto-Deploy-V2.yml @@ -239,7 +239,7 @@ jobs: - name: Build and push V2 image (Depot) if: env.USE_DEPOT == 'true' && steps.check-image.outputs.exists == 'false' - uses: depot/build-push-action@5f3b3c2e5a00f0093de47f657aeaefcedff27d18 # v1.16.0 + uses: depot/build-push-action@98e78adca7817480b8185f474a400b451d74e287 # v1.16.0 with: project: ${{ vars.DEPOT_PROJECT_ID }} context: . @@ -293,7 +293,7 @@ jobs: SECURITY_ENABLELOGIN: "true" SECURITY_INITIALLOGIN_USERNAME: "${{ secrets.TEST_LOGIN_USERNAME }}" SECURITY_INITIALLOGIN_PASSWORD: "${{ secrets.TEST_LOGIN_PASSWORD }}" - SYSTEM_DEFAULTLOCALE: en-GB + SYSTEM_DEFAULTLOCALE: en-US UI_APPNAME: "Stirling-PDF V2 PR#${{ needs.check-pr.outputs.pr_number }}" UI_HOMEDESCRIPTION: "V2 PR#${{ needs.check-pr.outputs.pr_number }} - Embedded Architecture" UI_APPNAMENAVBAR: "V2 PR#${{ needs.check-pr.outputs.pr_number }}" diff --git a/.github/workflows/PR-Demo-Comment-with-react.yml b/.github/workflows/PR-Demo-Comment-with-react.yml index b55650ec30..95afebfe4c 100644 --- a/.github/workflows/PR-Demo-Comment-with-react.yml +++ b/.github/workflows/PR-Demo-Comment-with-react.yml @@ -222,10 +222,10 @@ jobs: - name: Setup Gradle uses: gradle/actions/setup-gradle@50e97c2cd7a37755bbfafc9c5b7cafaece252f6e # v6.1.0 with: - gradle-version: 9.3.1 + gradle-version: 9.5.1 - name: Install Task - uses: go-task/setup-task@3be4020d41929789a01026e0e427a4321ce0ad44 # v2.0.0 + uses: go-task/setup-task@01a4adf9db2d14c1de7a560f09170b6e0df736aa # v2.1.0 - name: Run Gradle Command run: | if [ "${{ needs.check-comment.outputs.disable_security }}" == "true" ]; then @@ -256,7 +256,7 @@ jobs: - name: Build and push PR-specific image (Depot) if: env.USE_DEPOT == 'true' - uses: depot/build-push-action@5f3b3c2e5a00f0093de47f657aeaefcedff27d18 # v1.16.0 + uses: depot/build-push-action@98e78adca7817480b8185f474a400b451d74e287 # v1.16.0 with: project: ${{ vars.DEPOT_PROJECT_ID }} context: . @@ -285,7 +285,7 @@ jobs: - name: Build and push engine image (Depot) if: env.USE_DEPOT == 'true' && needs.check-comment.outputs.enable_prototypes == 'true' - uses: depot/build-push-action@5f3b3c2e5a00f0093de47f657aeaefcedff27d18 # v1.16.0 + uses: depot/build-push-action@98e78adca7817480b8185f474a400b451d74e287 # v1.16.0 with: project: ${{ vars.DEPOT_PROJECT_ID }} context: ./engine @@ -388,7 +388,7 @@ jobs: environment: DISABLE_ADDITIONAL_FEATURES: "${DISABLE_ADDITIONAL_FEATURES}" SECURITY_ENABLELOGIN: "${LOGIN_SECURITY}" - SYSTEM_DEFAULTLOCALE: en-GB + SYSTEM_DEFAULTLOCALE: en-US UI_APPNAME: "Stirling-PDF PR#${PR_NUMBER}" UI_HOMEDESCRIPTION: "PR#${PR_NUMBER} for Stirling-PDF Latest" UI_APPNAMENAVBAR: "PR#${PR_NUMBER}" diff --git a/.github/workflows/ai-engine.yml b/.github/workflows/ai-engine.yml index 9c701a5a61..33212b2330 100644 --- a/.github/workflows/ai-engine.yml +++ b/.github/workflows/ai-engine.yml @@ -43,10 +43,10 @@ jobs: - name: Setup Gradle uses: gradle/actions/setup-gradle@50e97c2cd7a37755bbfafc9c5b7cafaece252f6e # v6.1.0 with: - gradle-version: 9.3.1 + gradle-version: 9.5.1 - name: Install Task - uses: go-task/setup-task@3be4020d41929789a01026e0e427a4321ce0ad44 # v2.0.0 + uses: go-task/setup-task@01a4adf9db2d14c1de7a560f09170b6e0df736aa # v2.1.0 - name: Regenerate tool models run: task engine:tool-models diff --git a/.github/workflows/backend-build.yml b/.github/workflows/backend-build.yml index 8d3023a89b..0a252100a0 100644 --- a/.github/workflows/backend-build.yml +++ b/.github/workflows/backend-build.yml @@ -58,11 +58,11 @@ jobs: - name: Setup Gradle uses: gradle/actions/setup-gradle@50e97c2cd7a37755bbfafc9c5b7cafaece252f6e # v6.1.0 with: - gradle-version: 9.3.1 + gradle-version: 9.5.1 cache-disabled: true - name: Install Task - uses: go-task/setup-task@3be4020d41929789a01026e0e427a4321ce0ad44 # v2.0.0 + uses: go-task/setup-task@01a4adf9db2d14c1de7a560f09170b6e0df736aa # v2.1.0 - name: Check Java formatting (Spotless) # Runs once per matrix combination - pick the cheapest leg # (core - no proprietary, no saas) so we don't wait for the diff --git a/.github/workflows/build-enterprise.yml b/.github/workflows/build-enterprise.yml index 413129f0a2..d29002ac37 100644 --- a/.github/workflows/build-enterprise.yml +++ b/.github/workflows/build-enterprise.yml @@ -78,7 +78,7 @@ jobs: cache: "npm" cache-dependency-path: frontend/package-lock.json - name: Install Task - uses: go-task/setup-task@3be4020d41929789a01026e0e427a4321ce0ad44 # v2.0.0 + uses: go-task/setup-task@01a4adf9db2d14c1de7a560f09170b6e0df736aa # v2.1.0 - name: Install Playwright (chromium only) run: task e2e:install -- chromium - name: Build frontend (needed for playwright's vite preview webServer) diff --git a/.github/workflows/check-licence.yml b/.github/workflows/check-licence.yml index 400bdc94ac..1159c75dc6 100644 --- a/.github/workflows/check-licence.yml +++ b/.github/workflows/check-licence.yml @@ -40,11 +40,11 @@ jobs: - name: Setup Gradle uses: gradle/actions/setup-gradle@50e97c2cd7a37755bbfafc9c5b7cafaece252f6e # v6.1.0 with: - gradle-version: 9.3.1 + gradle-version: 9.5.1 cache-disabled: true - name: Install Task - uses: go-task/setup-task@3be4020d41929789a01026e0e427a4321ce0ad44 # v2.0.0 + uses: go-task/setup-task@01a4adf9db2d14c1de7a560f09170b6e0df736aa # v2.1.0 - name: Check licenses for compatibility run: task backend:licenses:check env: diff --git a/.github/workflows/check-openapi.yml b/.github/workflows/check-openapi.yml index c6d4731cd2..49c19433c3 100644 --- a/.github/workflows/check-openapi.yml +++ b/.github/workflows/check-openapi.yml @@ -45,11 +45,11 @@ jobs: - name: Setup Gradle uses: gradle/actions/setup-gradle@50e97c2cd7a37755bbfafc9c5b7cafaece252f6e # v6.1.0 with: - gradle-version: 9.3.1 + gradle-version: 9.5.1 cache-disabled: true - name: Install Task - uses: go-task/setup-task@3be4020d41929789a01026e0e427a4321ce0ad44 # v2.0.0 + uses: go-task/setup-task@01a4adf9db2d14c1de7a560f09170b6e0df736aa # v2.1.0 - name: Generate OpenAPI documentation run: task backend:swagger env: diff --git a/.github/workflows/check_toml.yml b/.github/workflows/check_toml.yml index 7e31863616..b7a277f873 100644 --- a/.github/workflows/check_toml.yml +++ b/.github/workflows/check_toml.yml @@ -166,16 +166,16 @@ jobs: // Determine reference file let referenceFilePath; - if (changedFiles.includes("frontend/editor/public/locales/en-GB/translation.toml")) { + if (changedFiles.includes("frontend/editor/public/locales/en-US/translation.toml")) { console.log("Using PR branch reference file."); const { data: fileContent } = await github.rest.repos.getContent({ owner: prRepoOwner, repo: prRepoName, - path: "frontend/editor/public/locales/en-GB/translation.toml", + path: "frontend/editor/public/locales/en-US/translation.toml", ref: branch, }); - referenceFilePath = "pr-branch-translation-en-GB.toml"; + referenceFilePath = "pr-branch-translation-en-US.toml"; const content = Buffer.from(fileContent.content, "base64").toString("utf-8"); fs.writeFileSync(referenceFilePath, content); } else { @@ -183,11 +183,11 @@ jobs: const { data: fileContent } = await github.rest.repos.getContent({ owner: repoOwner, repo: repoName, - path: "frontend/editor/public/locales/en-GB/translation.toml", + path: "frontend/editor/public/locales/en-US/translation.toml", ref: "main", }); - referenceFilePath = "main-branch-translation-en-GB.toml"; + referenceFilePath = "main-branch-translation-en-US.toml"; const content = Buffer.from(fileContent.content, "base64").toString("utf-8"); fs.writeFileSync(referenceFilePath, content); } @@ -293,6 +293,6 @@ jobs: run: | echo "Cleaning up temporary files..." rm -rf pr-branch - rm -f pr-branch-translation-en-GB.toml main-branch-translation-en-GB.toml changed_files.txt result.txt + rm -f pr-branch-translation-en-US.toml main-branch-translation-en-US.toml changed_files.txt result.txt echo "Cleanup complete." continue-on-error: true # Ensure cleanup runs even if previous steps fail diff --git a/.github/workflows/db-migration-test.yml b/.github/workflows/db-migration-test.yml index 2a386d1a5a..114855c4bc 100644 --- a/.github/workflows/db-migration-test.yml +++ b/.github/workflows/db-migration-test.yml @@ -48,7 +48,7 @@ jobs: - name: Setup Gradle uses: gradle/actions/setup-gradle@50e97c2cd7a37755bbfafc9c5b7cafaece252f6e # v6.1.0 with: - gradle-version: 9.3.1 + gradle-version: 9.5.1 cache-disabled: true # No `-PnoSpotless` here yet because the upstream cache layer matches the diff --git a/.github/workflows/deploy-on-v2-commit.yml b/.github/workflows/deploy-on-v2-commit.yml index e7af240c3d..850dcb12e4 100644 --- a/.github/workflows/deploy-on-v2-commit.yml +++ b/.github/workflows/deploy-on-v2-commit.yml @@ -107,7 +107,7 @@ jobs: - name: Build and push frontend image (Depot) if: env.USE_DEPOT == 'true' && steps.check-frontend.outputs.exists == 'false' - uses: depot/build-push-action@5f3b3c2e5a00f0093de47f657aeaefcedff27d18 # v1.16.0 + uses: depot/build-push-action@98e78adca7817480b8185f474a400b451d74e287 # v1.16.0 with: project: ${{ vars.DEPOT_PROJECT_ID }} context: . @@ -136,7 +136,7 @@ jobs: - name: Build and push backend image (Depot) if: env.USE_DEPOT == 'true' && steps.check-backend.outputs.exists == 'false' - uses: depot/build-push-action@5f3b3c2e5a00f0093de47f657aeaefcedff27d18 # v1.16.0 + uses: depot/build-push-action@98e78adca7817480b8185f474a400b451d74e287 # v1.16.0 with: project: ${{ vars.DEPOT_PROJECT_ID }} context: . @@ -188,7 +188,7 @@ jobs: environment: DISABLE_ADDITIONAL_FEATURES: "true" SECURITY_ENABLELOGIN: "false" - SYSTEM_DEFAULTLOCALE: en-GB + SYSTEM_DEFAULTLOCALE: en-US UI_APPNAME: "Stirling-PDF V2" UI_HOMEDESCRIPTION: "V2 Frontend/Backend Split" UI_APPNAMENAVBAR: "V2 Deployment" diff --git a/.github/workflows/docker-compose-tests.yml b/.github/workflows/docker-compose-tests.yml index 69031fed54..add65d6a8e 100644 --- a/.github/workflows/docker-compose-tests.yml +++ b/.github/workflows/docker-compose-tests.yml @@ -61,7 +61,7 @@ jobs: - name: Setup Gradle uses: gradle/actions/setup-gradle@50e97c2cd7a37755bbfafc9c5b7cafaece252f6e # v6.1.0 with: - gradle-version: 9.3.1 + gradle-version: 9.5.1 cache-disabled: true - name: Set up Docker Buildx diff --git a/.github/workflows/e2e-live.yml b/.github/workflows/e2e-live.yml index ec46a19f97..eb6d8d7be5 100644 --- a/.github/workflows/e2e-live.yml +++ b/.github/workflows/e2e-live.yml @@ -42,7 +42,7 @@ jobs: cache: "npm" cache-dependency-path: frontend/package-lock.json - name: Install Task - uses: go-task/setup-task@3be4020d41929789a01026e0e427a4321ce0ad44 # v2.0.0 + uses: go-task/setup-task@01a4adf9db2d14c1de7a560f09170b6e0df736aa # v2.1.0 - name: Install Playwright (chromium only) run: task e2e:install -- chromium - name: Build frontend (production bundle for vite preview) @@ -188,10 +188,30 @@ jobs: name: backend-log-live-${{ github.run_id }} path: .test-state/playwright/backend.log retention-days: 7 - - name: Upload Playwright report + - name: List Playwright output locations (debug) + if: always() + run: | + echo "::group::Playwright output dirs" + # Playwright anchors its default outputDir + HTML report to the + # nearest package.json, which is frontend/ (frontend/editor has + # none), so artifacts land under frontend/, not frontend/editor/. + ls -la frontend/playwright-report 2>/dev/null \ + || echo "no playwright-report at frontend/" + ls -la frontend/test-results 2>/dev/null \ + || echo "no test-results at frontend/" + find . -name node_modules -prune -o -name 'trace.zip' -print 2>/dev/null || true + echo "::endgroup::" + - name: Upload Playwright report + traces if: always() uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1 with: name: playwright-report-live-${{ github.run_id }} - path: frontend/editor/playwright-report/ + # test-results/ holds the per-test trace.zip (with browser console + # logs) + screenshots/video; playwright-report/ is the HTML report. + # Both live under frontend/ (Playwright anchors them to the nearest + # package.json, which is frontend/; frontend/editor has none). + path: | + frontend/playwright-report/ + frontend/test-results/ retention-days: 7 + if-no-files-found: warn diff --git a/.github/workflows/e2e-stubbed.yml b/.github/workflows/e2e-stubbed.yml index 99f7c11580..dd6bcc8ea7 100644 --- a/.github/workflows/e2e-stubbed.yml +++ b/.github/workflows/e2e-stubbed.yml @@ -36,7 +36,7 @@ jobs: cache: "npm" cache-dependency-path: frontend/package-lock.json - name: Install Task - uses: go-task/setup-task@3be4020d41929789a01026e0e427a4321ce0ad44 # v2.0.0 + uses: go-task/setup-task@01a4adf9db2d14c1de7a560f09170b6e0df736aa # v2.1.0 - name: Install Playwright (chromium only) run: task e2e:install -- chromium - name: Build frontend (production bundle for vite preview) diff --git a/.github/workflows/frontend-backend-licenses-update.yml b/.github/workflows/frontend-backend-licenses-update.yml index f551ca5ceb..8bf89bf210 100644 --- a/.github/workflows/frontend-backend-licenses-update.yml +++ b/.github/workflows/frontend-backend-licenses-update.yml @@ -97,7 +97,7 @@ jobs: run: npm ci --ignore-scripts --audit=false --fund=false - name: Install Task - uses: go-task/setup-task@3be4020d41929789a01026e0e427a4321ce0ad44 # v2.0.0 + uses: go-task/setup-task@01a4adf9db2d14c1de7a560f09170b6e0df736aa # v2.1.0 - name: Generate frontend license report (internal PR) if: github.event_name == 'pull_request' && github.event.pull_request.head.repo.fork == false env: @@ -110,8 +110,8 @@ jobs: NPM_CONFIG_IGNORE_SCRIPTS: "true" working-directory: frontend run: | - mkdir -p src/assets - npx --yes license-report --only=prod --output=json > src/assets/3rdPartyLicenses.json + mkdir -p editor/src/assets + npx --yes license-report --only=prod --output=json > editor/src/assets/3rdPartyLicenses.json - name: Postprocess with project script (BASE version) if: github.event_name == 'pull_request' && github.event.pull_request.head.repo.fork == true @@ -349,10 +349,10 @@ jobs: - name: Setup Gradle uses: gradle/actions/setup-gradle@50e97c2cd7a37755bbfafc9c5b7cafaece252f6e # v6.1.0 with: - gradle-version: 9.3.1 + gradle-version: 9.5.1 - name: Install Task - uses: go-task/setup-task@3be4020d41929789a01026e0e427a4321ce0ad44 # v2.0.0 + uses: go-task/setup-task@01a4adf9db2d14c1de7a560f09170b6e0df736aa # v2.1.0 - name: Check licenses and generate report id: license-check run: task backend:licenses:generate || echo "LICENSE_CHECK_FAILED=true" >> $GITHUB_ENV diff --git a/.github/workflows/frontend-validation.yml b/.github/workflows/frontend-validation.yml index 1e1068e222..70187df412 100644 --- a/.github/workflows/frontend-validation.yml +++ b/.github/workflows/frontend-validation.yml @@ -31,7 +31,7 @@ jobs: cache: "npm" cache-dependency-path: frontend/package-lock.json - name: Install Task - uses: go-task/setup-task@3be4020d41929789a01026e0e427a4321ce0ad44 # v2.0.0 + uses: go-task/setup-task@01a4adf9db2d14c1de7a560f09170b6e0df736aa # v2.1.0 - name: Quality-check frontend id: frontend-check run: task frontend:check:all diff --git a/.github/workflows/multiOSReleases.yml b/.github/workflows/multiOSReleases.yml index 438126f904..f29714d94e 100644 --- a/.github/workflows/multiOSReleases.yml +++ b/.github/workflows/multiOSReleases.yml @@ -73,10 +73,10 @@ jobs: - name: Setup Gradle uses: gradle/actions/setup-gradle@50e97c2cd7a37755bbfafc9c5b7cafaece252f6e # v6.1.0 with: - gradle-version: 9.3.1 + gradle-version: 9.5.1 - name: Install Task - uses: go-task/setup-task@3be4020d41929789a01026e0e427a4321ce0ad44 # v2.0.0 + uses: go-task/setup-task@01a4adf9db2d14c1de7a560f09170b6e0df736aa # v2.1.0 - name: Get version number id: versionNumber run: | @@ -148,7 +148,7 @@ jobs: - name: Setup Gradle uses: gradle/actions/setup-gradle@50e97c2cd7a37755bbfafc9c5b7cafaece252f6e # v6.1.0 with: - gradle-version: 9.3.1 + gradle-version: 9.5.1 - name: Setup Node.js if: matrix.variant.build_frontend == true @@ -159,7 +159,7 @@ jobs: cache-dependency-path: frontend/package-lock.json - name: Install Task - uses: go-task/setup-task@3be4020d41929789a01026e0e427a4321ce0ad44 # v2.0.0 + uses: go-task/setup-task@01a4adf9db2d14c1de7a560f09170b6e0df736aa # v2.1.0 - name: Build JAR run: ./gradlew build ${{ matrix.variant.build_frontend && '-PbuildWithFrontend=true' || '' }} -x spotlessApply -x spotlessCheck -x test -x sonarqube @@ -252,10 +252,10 @@ jobs: - name: Setup Gradle uses: gradle/actions/setup-gradle@50e97c2cd7a37755bbfafc9c5b7cafaece252f6e # v6.1.0 with: - gradle-version: 9.3.1 + gradle-version: 9.5.1 - name: Install Task - uses: go-task/setup-task@3be4020d41929789a01026e0e427a4321ce0ad44 # v2.0.0 + uses: go-task/setup-task@01a4adf9db2d14c1de7a560f09170b6e0df736aa # v2.1.0 # Build the universal JRE before desktop:prepare so the jlink:runtime # task short-circuits on its `test -d runtime/jre` status check. @@ -588,33 +588,34 @@ jobs: mkdir -p "$DIST" cd ./frontend/editor/src-tauri/target - # Find and rename artifacts based on platform + echo "=== tauri bundle artifacts ===" + find . -path "*/bundle/*" \( -name "*.msi" -o -name "*.deb" \ + -o -name "*.rpm" -o -name "*.AppImage" -o -name "*.dmg" \ + -o -name "*.app.tar.gz" -o -name "*.sig" \) 2>/dev/null | sort || true + echo "==============================" + + # createUpdaterArtifacts:true signs the native installers in place; + # each ships with a sibling .sig consumed by latest.json. if [ "${{ matrix.platform }}" = "windows-latest" ]; then - # Only ship the MSI installer on Windows. The loose exe and WiX toolset exes - # are not the user-facing installer - the MSI contains the signed inner exe. find . -name "*.msi" -exec cp {} "$DIST/Stirling-PDF-${{ matrix.name }}.msi" \; - find . -name "*.msi.zip" -exec cp {} "$DIST/Stirling-PDF-${{ matrix.name }}.msi.zip" \; - find . -name "*.msi.zip.sig" -exec cp {} "$DIST/Stirling-PDF-${{ matrix.name }}.msi.zip.sig" \; - find . -name "*.nsis.zip" -exec cp {} "$DIST/Stirling-PDF-${{ matrix.name }}.nsis.zip" \; - find . -name "*.nsis.zip.sig" -exec cp {} "$DIST/Stirling-PDF-${{ matrix.name }}.nsis.zip.sig" \; + find . -name "*.msi.sig" -exec cp {} "$DIST/Stirling-PDF-${{ matrix.name }}.msi.sig" \; elif [ "${{ matrix.platform }}" = "macos-15" ]; then + # DMG = manual install; .app.tar.gz (+ .sig) = updater payload. + # Raw .app is intentionally not shipped (hundreds of MB of uncompressed input). find . -name "*.dmg" -exec cp {} "$DIST/Stirling-PDF-${{ matrix.name }}.dmg" \; - find . -name "*.app" -exec cp -r {} "$DIST/Stirling-PDF-${{ matrix.name }}.app" \; find . -name "*.app.tar.gz" -exec cp {} "$DIST/Stirling-PDF-${{ matrix.name }}.app.tar.gz" \; find . -name "*.app.tar.gz.sig" -exec cp {} "$DIST/Stirling-PDF-${{ matrix.name }}.app.tar.gz.sig" \; else + # The raw .AppImage IS its updater payload (signed -> .AppImage.sig), + # not a .tar.gz wrapper - that's only produced under v1Compatible. find . -name "*.deb" -exec cp {} "$DIST/Stirling-PDF-${{ matrix.name }}.deb" \; + find . -name "*.deb.sig" -exec cp {} "$DIST/Stirling-PDF-${{ matrix.name }}.deb.sig" \; find . -name "*.rpm" -exec cp {} "$DIST/Stirling-PDF-${{ matrix.name }}.rpm" \; + find . -name "*.rpm.sig" -exec cp {} "$DIST/Stirling-PDF-${{ matrix.name }}.rpm.sig" \; find . -name "*.AppImage" -exec cp {} "$DIST/Stirling-PDF-${{ matrix.name }}.AppImage" \; - find . -name "*.AppImage.tar.gz" -exec cp {} "$DIST/Stirling-PDF-${{ matrix.name }}.AppImage.tar.gz" \; - find . -name "*.AppImage.tar.gz.sig" -exec cp {} "$DIST/Stirling-PDF-${{ matrix.name }}.AppImage.tar.gz.sig" \; + find . -name "*.AppImage.sig" -exec cp {} "$DIST/Stirling-PDF-${{ matrix.name }}.AppImage.sig" \; fi - # Copy updater latest.json if generated by tauri-action (rename to be platform-specific). - # Scope strictly to the tauri target tree - a broader find from repo root - # could pick up stray latest.json files and corrupt the release manifest. - find . -name "latest.json" -not -path "*/deps/*" 2>/dev/null | head -1 | xargs -I{} cp {} "$DIST/latest-${{ matrix.name }}.json" 2>/dev/null || true - - name: Upload build artifacts if: always() && steps.digicert-setup.conclusion != 'failure' uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1 @@ -634,6 +635,16 @@ jobs: with: egress-policy: audit + # Sparse-check out the verifier + pubkey before the artifact downloads + # so the checkout cannot clobber ./artifacts. + - name: Checkout updater verifier + uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2 + with: + sparse-checkout: | + .github/scripts/verify-updater-signatures.py + frontend/editor/src-tauri/tauri.conf.json + sparse-checkout-cone-mode: false + - name: Download all Tauri artifacts uses: actions/download-artifact@3e5f45b2cfb9172054b4087a40e8e0b5a5461e7c # v8.0.1 with: @@ -661,41 +672,106 @@ jobs: - name: Display structure of downloaded files run: ls -R ./artifacts - - name: Generate merged updater latest.json + # tauri-action only emits latest.json when it also publishes the release + # (tagName/releaseId set). We publish separately via action-gh-release, + # so build latest.json here from the per-platform .sig files. + - name: Generate updater latest.json + env: + VERSION: ${{ needs.determine-matrix.outputs.version }} + TAG: v${{ needs.determine-matrix.outputs.version }} + REPO: ${{ github.repository }} run: | python3 - << 'PYEOF' - import json, glob, os, sys + import json, os, sys + from pathlib import Path + from datetime import datetime, timezone - files = sorted(glob.glob('./artifacts/tauri/**/latest-*.json', recursive=True)) - if not files: - print("No latest-*.json files found — skipping latest.json generation") + VERSION = os.environ['VERSION'] + TAG = os.environ['TAG'] + REPO = os.environ['REPO'] + + ART = Path('./artifacts/tauri') + + # Tauri updater looks up {os}-{arch}-{installer} (e.g. linux-x86_64-deb) + # before bare {os}-{arch}, so per-format Linux keys let deb/rpm/appimage + # each self-update from their matching file. macOS universal serves both + # arches from the one .app.tar.gz. + PLATFORM_MAP = [ + { + 'bundles': ['Stirling-PDF-linux-x86_64.deb'], + 'targets': ['linux-x86_64-deb'], + }, + { + 'bundles': ['Stirling-PDF-linux-x86_64.rpm'], + 'targets': ['linux-x86_64-rpm'], + }, + { + 'bundles': ['Stirling-PDF-linux-x86_64.AppImage'], + 'targets': ['linux-x86_64-appimage'], + }, + { + 'bundles': ['Stirling-PDF-windows-x86_64.msi'], + 'targets': ['windows-x86_64-msi', 'windows-x86_64'], + }, + { + 'bundles': ['Stirling-PDF-macos-universal.app.tar.gz'], + 'targets': ['darwin-x86_64', 'darwin-aarch64'], + }, + ] + + # rglob() because download-artifact varies layout: one artifact -> flat, + # many -> nested under /. + def find_signed(name): + for bundle_path in sorted(ART.rglob(name)): + sig_path = bundle_path.with_name(bundle_path.name + '.sig') + if sig_path.exists(): + return bundle_path, sig_path + return None + + platforms = {} + skipped = [] + for entry in PLATFORM_MAP: + picked = None + for name in entry['bundles']: + picked = find_signed(name) + if picked: + break + if not picked: + skipped.append( + f"{entry['targets']} (no signed bundle among " + f"{entry['bundles']} - TAURI_SIGNING_PRIVATE_KEY unset " + f"or createUpdaterArtifacts disabled?)" + ) + continue + bundle_path, sig_path = picked + signature = sig_path.read_text(encoding='utf-8').strip() + url = f"https://github.com/{REPO}/releases/download/{TAG}/{bundle_path.name}" + for target in entry['targets']: + platforms[target] = {'signature': signature, 'url': url} + print(f"Added {entry['targets']} from {bundle_path.name}") + + if skipped: + print("Skipped platforms:") + for s in skipped: + print(f" - {s}") + + if not platforms: + print( + "WARN: no signed updater bundles found - " + "skipping latest.json generation" + ) sys.exit(0) - merged = None - for f in files: - print(f"Merging: {f}") - with open(f) as fh: - data = json.load(fh) - if merged is None: - merged = dict(data) - merged['platforms'] = {} - else: - # Guard against stale artifacts from a rerun bleeding into a new - # release's latest.json. All per-platform files must agree on version. - if merged.get('version') != data.get('version'): - sys.exit( - f"Version mismatch: {merged.get('version')} vs " - f"{data.get('version')} in {f}" - ) - for platform, payload in data.get('platforms', {}).items(): - if platform in merged['platforms']: - sys.exit(f"Duplicate platform entry for {platform} in {f}") - merged['platforms'][platform] = payload + manifest = { + 'version': VERSION, + 'notes': f"See https://github.com/{REPO}/releases/tag/{TAG}", + 'pub_date': datetime.now(timezone.utc).strftime('%Y-%m-%dT%H:%M:%SZ'), + 'platforms': platforms, + } - if merged and merged.get('platforms'): - with open('./artifacts/latest.json', 'w') as fh: - json.dump(merged, fh, indent=2) - print(f"Generated latest.json with platforms: {list(merged['platforms'].keys())}") + out = Path('./artifacts/latest.json') + out.write_text(json.dumps(manifest, indent=2) + '\n', encoding='utf-8') + print(f"Generated {out} with platforms: {sorted(platforms.keys())}") PYEOF - name: Upload merged artifacts for review @@ -705,23 +781,37 @@ jobs: path: ./artifacts/ retention-days: 7 + # Gate publish on valid updater sigs. Runs after the review upload (so + # artifacts survive for debugging) and before action-gh-release. + - name: Verify updater signatures + run: | + python3 -m pip install --quiet 'cryptography==44.0.0' + python3 .github/scripts/verify-updater-signatures.py \ + ./artifacts/tauri frontend/editor/src-tauri/tauri.conf.json + + # workflow_dispatch path requires platform=='all' so a single-platform + # dispatch can't overwrite an existing release's full latest.json with a + # partial one (action-gh-release defaults overwrite_files:true). + # release / V2-master always build the full matrix so no extra guard needed. + # fail_on_unmatched_files makes a missing latest.json or installer fail loudly + # instead of silently shipping a broken auto-update. - name: Upload binaries to Release - if: (github.event_name == 'workflow_dispatch' && github.event.inputs.test_mode != 'true') || github.event_name == 'release' || github.ref == 'refs/heads/V2-master' + if: (github.event_name == 'workflow_dispatch' && github.event.inputs.test_mode != 'true' && github.event.inputs.platform == 'all') || github.event_name == 'release' || github.ref == 'refs/heads/V2-master' uses: softprops/action-gh-release@b4309332981a82ec1c5618f44dd2e27cc8bfbfda # v3.0.0 with: tag_name: v${{ needs.determine-matrix.outputs.version }} generate_release_notes: true + fail_on_unmatched_files: true + # Installers + updater payloads + manifest. .sig contents are embedded + # in latest.json so the .sig files themselves are not uploaded. files: | ./artifacts/**/*.jar ./artifacts/**/*.msi - ./artifacts/**/*.msi.zip - ./artifacts/**/*.nsis.zip ./artifacts/**/*.dmg ./artifacts/**/*.app.tar.gz ./artifacts/**/*.deb ./artifacts/**/*.rpm ./artifacts/**/*.AppImage - ./artifacts/**/*.AppImage.tar.gz ./artifacts/latest.json draft: false prerelease: false diff --git a/.github/workflows/nightly.yml b/.github/workflows/nightly.yml index 691d20cc4b..8731eb8a84 100644 --- a/.github/workflows/nightly.yml +++ b/.github/workflows/nightly.yml @@ -37,7 +37,7 @@ jobs: cache-dependency-path: frontend/package-lock.json - name: Install Task - uses: go-task/setup-task@3be4020d41929789a01026e0e427a4321ce0ad44 # v2.0.0 + uses: go-task/setup-task@01a4adf9db2d14c1de7a560f09170b6e0df736aa # v2.1.0 - name: Install all Playwright browsers run: task e2e:install diff --git a/.github/workflows/pre_commit.yml b/.github/workflows/pre_commit.yml index fe2d1b6877..85d2efa996 100644 --- a/.github/workflows/pre_commit.yml +++ b/.github/workflows/pre_commit.yml @@ -1,8 +1,7 @@ name: Pre-commit -# Runs `pre-commit run` for ruff / codespell / gitleaks / EOF / trailing-ws. -# Called from build.yml on PRs and merge_group; also runnable on demand via -# workflow_dispatch for manual local-equivalent linting. +# Runs the repo-wide lint/format/secret checks via `task pre-commit`. +# Called from build.yml on PRs and merge_group; also runnable on demand via workflow_dispatch. on: workflow_call: workflow_dispatch: @@ -13,10 +12,6 @@ permissions: jobs: pre-commit: runs-on: ubuntu-latest - env: - # Prevents sdist builds → no tar extraction - PIP_ONLY_BINARY: ":all:" - PIP_DISABLE_PIP_VERSION_CHECK: "1" steps: - name: Harden Runner uses: step-security/harden-runner@ab7a9404c0f3da075243ca237b5fac12c98deaa5 # v2.19.3 @@ -29,23 +24,13 @@ jobs: fetch-depth: 0 persist-credentials: false - - name: Set up Python - uses: actions/setup-python@a309ff8b426b58ec0e2a45f0f869d46889d02405 # v6.2.0 + - name: Install uv + uses: astral-sh/setup-uv@08807647e7069bb48b6ef5acd8ec9567f424441b # v8.1.0 with: - python-version: 3.12 - cache: "pip" # caching pip dependencies - cache-dependency-path: ./.github/scripts/requirements_pre_commit.txt + enable-cache: true - - name: Run Pre-Commit Hooks - run: | - pip install --require-hashes --only-binary=:all: -r ./.github/scripts/requirements_pre_commit.txt + - name: Install Task + uses: go-task/setup-task@3be4020d41929789a01026e0e427a4321ce0ad44 # v2.0.0 - - name: Run Pre-Commit - run: | - pre-commit run ruff --all-files -c .pre-commit-config.yaml - pre-commit run ruff-format --all-files -c .pre-commit-config.yaml - pre-commit run codespell --all-files -c .pre-commit-config.yaml - pre-commit run gitleaks --all-files -c .pre-commit-config.yaml - pre-commit run end-of-file-fixer --all-files -c .pre-commit-config.yaml - pre-commit run trailing-whitespace --all-files -c .pre-commit-config.yaml - git diff --exit-code + - name: Run pre-commit checks + run: task pre-commit diff --git a/.github/workflows/push-docker-base.yml b/.github/workflows/push-docker-base.yml index 801a4328f9..57d0935d1d 100644 --- a/.github/workflows/push-docker-base.yml +++ b/.github/workflows/push-docker-base.yml @@ -75,7 +75,7 @@ jobs: - name: Generate tags for base image id: meta - uses: docker/metadata-action@030e881283bb7a6894de51c315a6bfe6a94e05cf # v6.0.0 + uses: docker/metadata-action@80c7e94dd9b9319bd5eb7a0e0fe9291e23a2a2e9 # v6.1.0 with: images: | ${{ secrets.DOCKER_HUB_ORG_USERNAME }}/stirling-pdf-base diff --git a/.github/workflows/push-docker.yml b/.github/workflows/push-docker.yml index c637b91580..ad7653eab3 100644 --- a/.github/workflows/push-docker.yml +++ b/.github/workflows/push-docker.yml @@ -78,14 +78,14 @@ jobs: - name: Setup Gradle uses: gradle/actions/setup-gradle@50e97c2cd7a37755bbfafc9c5b7cafaece252f6e # v6.1.0 with: - gradle-version: 9.3.1 + gradle-version: 9.5.1 - name: Set up Docker Buildx id: buildx uses: docker/setup-buildx-action@4d04d5d9486b7bd6fa91e7baf45bbb4f8b9deedd # v4.0.0 - name: Install Task - uses: go-task/setup-task@3be4020d41929789a01026e0e427a4321ce0ad44 # v2.0.0 + uses: go-task/setup-task@01a4adf9db2d14c1de7a560f09170b6e0df736aa # v2.1.0 - name: Get version number id: versionNumber run: echo "versionNumber=$(./gradlew printVersion --quiet | tail -1)" >> $GITHUB_OUTPUT @@ -129,7 +129,7 @@ jobs: - name: Generate tags for latest id: meta if: env.RUN_MAIN_APP == 'true' - uses: docker/metadata-action@030e881283bb7a6894de51c315a6bfe6a94e05cf # v6.0.0 + uses: docker/metadata-action@80c7e94dd9b9319bd5eb7a0e0fe9291e23a2a2e9 # v6.1.0 with: images: | ${{ secrets.DOCKER_HUB_USERNAME }}/s-pdf @@ -178,7 +178,7 @@ jobs: - name: Generate tags for latest-fat id: meta-fat - uses: docker/metadata-action@030e881283bb7a6894de51c315a6bfe6a94e05cf # v6.0.0 + uses: docker/metadata-action@80c7e94dd9b9319bd5eb7a0e0fe9291e23a2a2e9 # v6.1.0 if: env.RUN_MAIN_APP == 'true' && github.ref != 'refs/heads/main' && github.ref != 'refs/heads/testMain' with: images: | @@ -222,7 +222,7 @@ jobs: - name: Generate tags for ultra-lite id: meta-lite - uses: docker/metadata-action@030e881283bb7a6894de51c315a6bfe6a94e05cf # v6.0.0 + uses: docker/metadata-action@80c7e94dd9b9319bd5eb7a0e0fe9291e23a2a2e9 # v6.1.0 if: env.RUN_MAIN_APP == 'true' && github.ref != 'refs/heads/main' && github.ref != 'refs/heads/testMain' with: images: | diff --git a/.github/workflows/stale.yml b/.github/workflows/stale.yml index b4e9d9c91a..6972a1fc95 100644 --- a/.github/workflows/stale.yml +++ b/.github/workflows/stale.yml @@ -22,7 +22,7 @@ jobs: egress-policy: audit - name: 30 days stale issues - uses: actions/stale@b5d41d4e1d5dceea10e7104786b73624c18a190f # v10.2.0 + uses: actions/stale@eb5cf3af3ac0a1aa4c9c45633dd1ae542a27a899 # v10.3.0 with: repo-token: ${{ secrets.GITHUB_TOKEN }} days-before-stale: 30 diff --git a/.github/workflows/swagger.yml b/.github/workflows/swagger.yml index 295498e571..00c368c75f 100644 --- a/.github/workflows/swagger.yml +++ b/.github/workflows/swagger.yml @@ -48,7 +48,7 @@ jobs: - name: Setup Gradle uses: gradle/actions/setup-gradle@50e97c2cd7a37755bbfafc9c5b7cafaece252f6e # v6.1.0 with: - gradle-version: 9.3.1 + gradle-version: 9.5.1 - name: Generate Swagger documentation run: ./gradlew :stirling-pdf:generateOpenApiDocs @@ -63,7 +63,7 @@ jobs: SWAGGERHUB_USER: "Frooodle" - name: Install Task - uses: go-task/setup-task@3be4020d41929789a01026e0e427a4321ce0ad44 # v2.0.0 + uses: go-task/setup-task@01a4adf9db2d14c1de7a560f09170b6e0df736aa # v2.1.0 - name: Get version number id: versionNumber run: echo "versionNumber=$(./gradlew printVersion --quiet | tail -1)" >> $GITHUB_OUTPUT diff --git a/.github/workflows/sync_files_v2.yml b/.github/workflows/sync_files_v2.yml index 4b93c27d41..4647442d4f 100644 --- a/.github/workflows/sync_files_v2.yml +++ b/.github/workflows/sync_files_v2.yml @@ -58,15 +58,23 @@ jobs: - name: Install Python dependencies run: | - pip install --require-hashes --only-binary=:all: -r ./.github/scripts/requirements_sync_readme.txt -r ./.github/scripts/requirements_pre_commit.txt + pip install --require-hashes --only-binary=:all: -r ./.github/scripts/requirements_sync_readme.txt + + - name: Install uv + uses: astral-sh/setup-uv@08807647e7069bb48b6ef5acd8ec9567f424441b # v8.1.0 + with: + enable-cache: true + + - name: Install Task + uses: go-task/setup-task@3be4020d41929789a01026e0e427a4321ce0ad44 # v2.0.0 - name: Sync translation TOML files run: | - python .github/scripts/check_language_toml.py --reference-file "frontend/editor/public/locales/en-GB/translation.toml" --branch main + python .github/scripts/check_language_toml.py --reference-file "frontend/editor/public/locales/en-US/translation.toml" --branch main - - name: pre-commit run + - name: Sort translation TOML files run: | - pre-commit run toml-sort-fix --all-files + task pre-commit:toml-sort FIX=1 - name: Commit translation files run: | @@ -100,7 +108,7 @@ jobs: This Pull Request was automatically generated to synchronize updates to translation files and documentation. Below are the details of the changes made: #### **1. Synchronization of Translation Files** - - Updated translation files (`frontend/editor/public/locales/*/translation.toml`) to reflect changes in the reference file `en-GB/translation.toml`. + - Updated translation files (`frontend/editor/public/locales/*/translation.toml`) to reflect changes in the reference file `en-US/translation.toml`. - Ensured consistency and synchronization across all supported language files. - Highlighted any missing or incomplete translations. - **Format**: TOML diff --git a/.github/workflows/tauri-build.yml b/.github/workflows/tauri-build.yml index 82dc1b63fa..d8bece6803 100644 --- a/.github/workflows/tauri-build.yml +++ b/.github/workflows/tauri-build.yml @@ -136,10 +136,10 @@ jobs: - name: Setup Gradle uses: gradle/actions/setup-gradle@50e97c2cd7a37755bbfafc9c5b7cafaece252f6e # v6.1.0 with: - gradle-version: 9.3.1 + gradle-version: 9.5.1 - name: Setup Task - uses: go-task/setup-task@3be4020d41929789a01026e0e427a4321ce0ad44 # v2.0.0 + uses: go-task/setup-task@01a4adf9db2d14c1de7a560f09170b6e0df736aa # v2.1.0 - name: Build universal macOS JRE if: matrix.platform == 'macos-15' diff --git a/.github/workflows/test-build-docker.yml b/.github/workflows/test-build-docker.yml index 9d9c8f2bfa..0b692a008a 100644 --- a/.github/workflows/test-build-docker.yml +++ b/.github/workflows/test-build-docker.yml @@ -106,11 +106,11 @@ jobs: - name: Setup Gradle uses: gradle/actions/setup-gradle@50e97c2cd7a37755bbfafc9c5b7cafaece252f6e # v6.1.0 with: - gradle-version: 9.3.1 + gradle-version: 9.5.1 cache-disabled: true - name: Install Task - uses: go-task/setup-task@3be4020d41929789a01026e0e427a4321ce0ad44 # v2.0.0 + uses: go-task/setup-task@01a4adf9db2d14c1de7a560f09170b6e0df736aa # v2.1.0 - name: Build application run: task backend:build env: @@ -157,7 +157,7 @@ jobs: - name: Build ${{ matrix.docker-rev }} (Depot) if: env.USE_DEPOT == 'true' - uses: depot/build-push-action@5f3b3c2e5a00f0093de47f657aeaefcedff27d18 # v1.16.0 + uses: depot/build-push-action@98e78adca7817480b8185f474a400b451d74e287 # v1.16.0 with: project: ${{ vars.DEPOT_PROJECT_ID }} context: . @@ -230,7 +230,7 @@ jobs: - name: Build docker/unoserver/Dockerfile (Depot) if: env.USE_DEPOT == 'true' - uses: depot/build-push-action@5f3b3c2e5a00f0093de47f657aeaefcedff27d18 # v1.16.0 + uses: depot/build-push-action@98e78adca7817480b8185f474a400b451d74e287 # v1.16.0 with: project: ${{ vars.DEPOT_PROJECT_ID }} context: . diff --git a/.github/workflows/testdriver.yml b/.github/workflows/testdriver.yml index 9a395120aa..05fce0199c 100644 --- a/.github/workflows/testdriver.yml +++ b/.github/workflows/testdriver.yml @@ -51,7 +51,7 @@ jobs: - name: Setup Gradle uses: gradle/actions/setup-gradle@50e97c2cd7a37755bbfafc9c5b7cafaece252f6e # v6.1.0 with: - gradle-version: 9.3.1 + gradle-version: 9.5.1 - name: Build with Gradle run: ./gradlew build @@ -83,7 +83,7 @@ jobs: - name: Build and push test image (Depot) if: env.USE_DEPOT == 'true' - uses: depot/build-push-action@5f3b3c2e5a00f0093de47f657aeaefcedff27d18 # v1.16.0 + uses: depot/build-push-action@98e78adca7817480b8185f474a400b451d74e287 # v1.16.0 with: project: ${{ vars.DEPOT_PROJECT_ID }} context: . @@ -129,7 +129,7 @@ jobs: environment: DISABLE_ADDITIONAL_FEATURES: "true" SECURITY_ENABLELOGIN: "false" - SYSTEM_DEFAULTLOCALE: en-GB + SYSTEM_DEFAULTLOCALE: en-US UI_APPNAME: "Stirling-PDF Test" UI_HOMEDESCRIPTION: "Test Deployment" UI_APPNAMENAVBAR: "Test" diff --git a/.gitignore b/.gitignore index 9e221c7eb7..a379cf1db0 100644 --- a/.gitignore +++ b/.gitignore @@ -46,6 +46,11 @@ app/core/storage/ # These are generated by npm build and should not be committed app/core/src/main/resources/static/assets/ app/core/src/main/resources/static/index.html +# Prerendered per-route SPA pages (OG/social-preview), e.g. compress.html. api-landing.html is source. +app/core/src/main/resources/static/*.html +!app/core/src/main/resources/static/api-landing.html +# Prerendered nested-route pages (e.g. settings/people.html) +app/core/src/main/resources/static/settings/ app/core/src/main/resources/static/locales/ app/core/src/main/resources/static/Login/ app/core/src/main/resources/static/classic-logo/ @@ -53,10 +58,14 @@ app/core/src/main/resources/static/modern-logo/ app/core/src/main/resources/static/og_images/ app/core/src/main/resources/static/samples/ app/core/src/main/resources/static/manifest-classic.json +app/core/src/main/resources/static/og-metadata.json +app/core/src/main/resources/static/sw-folder-retry.js app/core/src/main/resources/static/robots.txt app/core/src/main/resources/static/pdfium/ app/core/src/main/resources/static/pdfjs/ app/core/src/main/resources/static/vendor/ +app/core/src/main/resources/static/**/*.gz +app/core/src/main/resources/static/**/*.br # Note: Keep backend-managed files like fonts/, css/, js/, pdfjs/, etc. # Gradle @@ -279,4 +288,4 @@ docs/type3/signatures/ *.playwright-mcp.png # Local screenshot artifacts from *-screenshots.spec.ts -frontend/screenshots/ +frontend/editor/screenshots/ diff --git a/.gitleaksignore b/.gitleaksignore index a5f4b02f91..43f9db5957 100644 --- a/.gitleaksignore +++ b/.gitleaksignore @@ -1,5 +1,21 @@ -# PostHog project-level key — phc_ prefix keys are public/client-side by design +# PostHog project-level key - phc_ prefix keys are public/client-side by design # (PostHog client-side tracking embeds them in the browser bundle). Committed # intentionally in #6150 so engine/.env has a working default, with real # credentials overridden via engine/.env.local. engine/.env:generic-api-key:41 + +# MCP test fixtures / harness - no real secrets: +# - test-only API key constant in an integration test +# - JDBC URL + throwaway Keycloak creds in the local test compose +# - placeholder / shell-variable Bearer headers in curl-based validation scripts +app/proprietary/src/test/java/stirling/software/proprietary/mcp/security/McpApiKeyIntegrationTest.java:generic-api-key:40 +testing/compose/docker-compose-keycloak-mcp.yml:generic-api-key:25 +testing/compose/validate-mcp-apikey.sh:curl-auth-header:73 +testing/compose/validate-mcp-test.sh:curl-auth-header:92 +testing/compose/validate-mcp-test.sh:curl-auth-header:116 + +# Storybook example showing curl with a fake Bearer token placeholder (sk_live_a3f8...). +frontend/shared/components/CodeBlock.stories.tsx:curl-auth-header:4 + +# Truncated placeholder API key in portal docs example (sk_live_8f2c...e10) - not a real secret. +frontend/portal/src/components/docs/GettingStartedSection.tsx:generic-api-key:31 diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 2490b4ae6e..83753a9263 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -1,52 +1,14 @@ +# The actual checks live in .taskfiles/pre-commit.yml (with helper scripts under +# scripts/pre-commit/) and are driven by Task. This hook just delegates to `task +# pre-commit` so the git pre-commit hook, CI and a manual `task pre-commit` all +# run the exact same thing. Requires `task` and `uv` on PATH. To auto-fix instead +# of only checking, run `task pre-commit:fix`. repos: - - repo: https://github.com/astral-sh/ruff-pre-commit - rev: v0.15.14 + - repo: local hooks: - - id: ruff - args: - - --fix - - --line-length=127 - files: ^((\.github/scripts|scripts|app/core/src/main/resources/static/python)/.+)?[^/]+\.py$ - exclude: (split_photos.py) - - id: ruff-format - files: ^((\.github/scripts|scripts|app/core/src/main/resources/static/python)/.+)?[^/]+\.py$ - exclude: (split_photos.py) - - repo: https://github.com/codespell-project/codespell - rev: v2.4.2 - hooks: - - id: codespell - args: - - --ignore-words-list=thirdParty,tabEl,tabEls,Sie,ist,fulfilment - - --skip="./.*,*.csv,*.json,*.ambr" - - --quiet-level=2 - files: \.(html|css|js|py|md)$ - exclude: (.vscode|.devcontainer|app/core/src/main/resources|app/proprietary/src/main/resources|frontend/editor/public/vendor|Dockerfile|.*/pdfjs.*|.*/thirdParty.*|bootstrap.*|.*\.min\..*|.*diff\.js) - - repo: https://github.com/gitleaks/gitleaks - rev: v8.30.0 - hooks: - - id: gitleaks - - repo: https://github.com/pre-commit/pre-commit-hooks - rev: v6.0.0 - hooks: - - id: end-of-file-fixer - files: ^.*(\.js|\.java|\.py|\.yml)$ - exclude: ^(.*/pdfjs.*|.*/thirdParty.*|bootstrap.*|.*\.min\..*|.*diff\.js|\.github/workflows/.*$) - - id: trailing-whitespace - files: ^.*(\.js|\.java|\.py|\.yml)$ - exclude: ^(.*/pdfjs.*|.*/thirdParty.*|bootstrap.*|.*\.min\..*|.*diff\.js|\.github/workflows/.*$) - - repo: https://github.com/pappasam/toml-sort - rev: v0.24.4 - hooks: - - id: toml-sort-fix - files: frontend/editor/public/locales/.*\.toml$ - args: ['--in-place', '--all', '--ignore-case'] - # - repo: https://github.com/thibaudcolas/pre-commit-stylelint - # rev: v16.21.1 - # hooks: - # - id: stylelint - # additional_dependencies: - # - stylelint@16.21.1 - # - stylelint-config-standard@38.0.0 - # - "@stylistic/stylelint-plugin@3.1.3" - # files: \.(css)$ - # args: [--fix] + - id: task-pre-commit + name: task pre-commit + entry: task pre-commit + language: system + pass_filenames: false + always_run: true diff --git a/.taskfiles/backend.yml b/.taskfiles/backend.yml index 2a3adffeb7..8a7f6c622c 100644 --- a/.taskfiles/backend.yml +++ b/.taskfiles/backend.yml @@ -18,17 +18,28 @@ version: '3' tasks: dev: desc: "Start backend dev server" + cmds: + - task: dev:proprietary + vars: + PORT: '{{.PORT}}' + AIENGINE_URL: '{{.AIENGINE_URL}}' + AIENGINE_ENABLED: '{{.AIENGINE_ENABLED}}' + AIENGINE_TIMEOUTSECONDS: '{{.AIENGINE_TIMEOUTSECONDS}}' + + dev:proprietary: + desc: "Start backend dev server in proprietary mode" ignore_error: true vars: PORT: '{{.PORT | default "8080"}}' AIENGINE_URL: '{{.AIENGINE_URL | default ""}}' + AIENGINE_ENABLED: '{{.AIENGINE_ENABLED | default "false"}}' AIENGINE_TIMEOUTSECONDS: '{{.AIENGINE_TIMEOUTSECONDS | default "120"}}' env: SERVER_PORT: '{{.PORT}}' cmds: - - cmd: '{{if .AIENGINE_URL}}AIENGINE_URL={{.AIENGINE_URL}} AIENGINE_ENABLED=true AIENGINE_TIMEOUTSECONDS={{.AIENGINE_TIMEOUTSECONDS}} {{end}}cmd /c ".\gradlew.bat :stirling-pdf:bootRun"' + - cmd: '{{if .AIENGINE_URL}}AIENGINE_URL={{.AIENGINE_URL}} AIENGINE_ENABLED={{.AIENGINE_ENABLED}} AIENGINE_TIMEOUTSECONDS={{.AIENGINE_TIMEOUTSECONDS}} {{end}}cmd /c ".\gradlew.bat :stirling-pdf:bootRun"' platforms: [windows] - - cmd: '{{if .AIENGINE_URL}}AIENGINE_URL={{.AIENGINE_URL}} AIENGINE_ENABLED=true AIENGINE_TIMEOUTSECONDS={{.AIENGINE_TIMEOUTSECONDS}} {{end}}./gradlew :stirling-pdf:bootRun' + - cmd: '{{if .AIENGINE_URL}}AIENGINE_URL={{.AIENGINE_URL}} AIENGINE_ENABLED={{.AIENGINE_ENABLED}} AIENGINE_TIMEOUTSECONDS={{.AIENGINE_TIMEOUTSECONDS}} {{end}}./gradlew :stirling-pdf:bootRun' platforms: [linux, darwin] dev:bundled: @@ -50,9 +61,15 @@ tasks: PORT: '{{.PORT | default "8080"}}' # Override to "" to run the pure `saas` profile against your own SAAS_DB_*. PROFILES: '{{.PROFILES | default "dev"}}' + AIENGINE_URL: '{{.AIENGINE_URL | default ""}}' + AIENGINE_ENABLED: '{{.AIENGINE_ENABLED | default "false"}}' + AIENGINE_TIMEOUTSECONDS: '{{.AIENGINE_TIMEOUTSECONDS | default "120"}}' env: SERVER_PORT: '{{.PORT}}' STIRLING_FLAVOR: saas + AIENGINE_URL: '{{.AIENGINE_URL}}' + AIENGINE_ENABLED: '{{.AIENGINE_ENABLED}}' + AIENGINE_TIMEOUTSECONDS: '{{.AIENGINE_TIMEOUTSECONDS}}' cmds: - cmd: cmd /c ".\gradlew.bat :stirling-pdf:bootRun {{if .PROFILES}}--args=\"--spring.profiles.include={{.PROFILES}}\"{{end}}" platforms: [windows] diff --git a/.taskfiles/desktop.yml b/.taskfiles/desktop.yml index ac4ccbe806..0b37df9f6f 100644 --- a/.taskfiles/desktop.yml +++ b/.taskfiles/desktop.yml @@ -1,7 +1,9 @@ version: '3' vars: - JLINK_MODULES: "java.base,java.compiler,java.desktop,java.instrument,java.logging,java.management,java.naming,java.net.http,java.prefs,java.rmi,java.scripting,java.security.jgss,java.security.sasl,java.sql,java.transaction.xa,java.xml,java.xml.crypto,jdk.crypto.ec,jdk.crypto.cryptoki,jdk.unsupported" + # jdk.dynalink is required by VeraPDF (PDF/A validation); without it the bundled JRE throws + # NoClassDefFoundError: jdk/dynalink/Namespace at runtime in get-info-on-pdf and verify-pdf + JLINK_MODULES: "java.base,java.compiler,java.desktop,java.instrument,java.logging,java.management,java.naming,java.net.http,java.prefs,java.rmi,java.scripting,java.security.jgss,java.security.sasl,java.sql,java.transaction.xa,java.xml,java.xml.crypto,jdk.crypto.ec,jdk.crypto.cryptoki,jdk.unsupported,jdk.dynalink" # Override via JPDFIUM_PLATFORMS env (csv of platform keys, or 'all'). JPDFIUM_PLATFORMS: @@ -62,21 +64,21 @@ tasks: deps: [prepare] dir: editor cmds: - - npx tauri build --bundles app + - npx tauri build --bundles app --config '{"bundle":{"createUpdaterArtifacts":false}}' build:dev:windows: desc: "Build Tauri desktop NSIS installer (Windows)" deps: [prepare] dir: editor cmds: - - npx tauri build --bundles nsis + - npx tauri build --bundles nsis --config '{"bundle":{"createUpdaterArtifacts":false}}' build:dev:linux: desc: "Build Tauri desktop AppImage (Linux)" deps: [prepare] dir: editor cmds: - - npx tauri build --bundles appimage + - npx tauri build --bundles appimage --config '{"bundle":{"createUpdaterArtifacts":false}}' test: desc: "Run Tauri/Cargo tests" @@ -125,14 +127,15 @@ tasks: cmds: - rm -rf runtime/jre - mkdir -p runtime - - >- - jlink - --add-modules {{.JLINK_MODULES}} - --strip-debug - --compress=zip-6 - --no-header-files - --no-man-pages - --output runtime/jre + - | + JLINK_COMPRESS="$(jlink --help 2>&1 | grep -q 'zip-\[0-9\]' && echo zip-6 || echo 2)" + jlink \ + --add-modules {{.JLINK_MODULES}} \ + --strip-debug \ + --compress="$JLINK_COMPRESS" \ + --no-header-files \ + --no-man-pages \ + --output runtime/jre # jlink emits its files mode 444 (read-only). Tauri's build-script # resource copier preserves source permissions when staging # `runtime/jre/**/*` into `target//runtime/jre/...`, so the diff --git a/.taskfiles/docker.yml b/.taskfiles/docker.yml index abe885a13e..cba35493fa 100644 --- a/.taskfiles/docker.yml +++ b/.taskfiles/docker.yml @@ -20,6 +20,11 @@ tasks: cmds: - docker build -t stirling-pdf-ultra-lite -f {{.EMBEDDED_DIR}}/Dockerfile.ultra-lite . + build:backend: + desc: "Build backend-only Docker image (no embedded frontend)" + cmds: + - docker build -t stirling-pdf-backend -f docker/backend/Dockerfile . + build:frontend: desc: "Build frontend-only Docker image" cmds: diff --git a/.taskfiles/e2e.yml b/.taskfiles/e2e.yml index 825b7cb69f..618041807f 100644 --- a/.taskfiles/e2e.yml +++ b/.taskfiles/e2e.yml @@ -212,3 +212,55 @@ tasks: desc: "Stop the SAML keycloak test environment" cmds: - docker compose -f testing/compose/docker-compose-keycloak-saml.yml down -v + + mcp:up: + desc: "Start the MCP keycloak test environment (Stirling as OAuth resource server)" + summary: | + Brings up Keycloak (OAuth authorization server) + Stirling configured as an + MCP resource server, then you can exercise /mcp with real Keycloak tokens. + Set LICENSE_KEY= to skip the interactive license prompt: + task e2e:mcp:up LICENSE_KEY=abc123 + Pass extra flags via -- : + task e2e:mcp:up -- --validate --nobuild + ignore_error: true + cmds: + - bash testing/compose/start-mcp-test.sh {{if .LICENSE_KEY}}--license-key "{{.LICENSE_KEY}}"{{end}} {{.CLI_ARGS}} + + mcp:manual: + desc: "Start the MCP keycloak test env in manual mode (prints URLs + a live token for your client)" + summary: | + Brings the stack up and prints copy-paste URLs/commands plus a freshly minted + access token so you can drive your own MCP client (Inspector, curl, ...). + task e2e:mcp:manual LICENSE_KEY= + Add --nobuild if the images are already built: + task e2e:mcp:manual LICENSE_KEY= -- --nobuild + ignore_error: true + cmds: + - bash testing/compose/start-mcp-test.sh --manual {{if .LICENSE_KEY}}--license-key "{{.LICENSE_KEY}}"{{end}} {{.CLI_ARGS}} + + mcp:apikey: + desc: "Start the MCP test env in API-KEY manual mode (no OAuth/IdP): mints a key + prints client settings" + summary: | + Brings Stirling up in apikey auth mode and prints copy-paste client settings with a freshly + minted X-API-KEY - ideal for clients whose OAuth layer can't reach localhost. + task e2e:mcp:apikey LICENSE_KEY= + Add --nobuild if images are already built: + task e2e:mcp:apikey LICENSE_KEY= -- --nobuild + ignore_error: true + cmds: + - bash testing/compose/start-mcp-test.sh --apikey {{if .LICENSE_KEY}}--license-key "{{.LICENSE_KEY}}"{{end}} {{.CLI_ARGS}} + + mcp:validate: + desc: "Validate the running MCP keycloak test environment end-to-end (oauth mode + real MCP SDK client)" + cmds: + - bash testing/compose/validate-mcp-test.sh + + mcp:validate-apikey: + desc: "Validate the MCP server in API-KEY auth mode (mints a key + real MCP SDK client), then restore oauth" + cmds: + - bash testing/compose/validate-mcp-apikey.sh + + mcp:down: + desc: "Stop the MCP keycloak test environment" + cmds: + - docker compose -f testing/compose/docker-compose-keycloak-mcp.yml down -v diff --git a/.taskfiles/engine.yml b/.taskfiles/engine.yml index 5c9d6a8e65..cfe0790241 100644 --- a/.taskfiles/engine.yml +++ b/.taskfiles/engine.yml @@ -33,7 +33,7 @@ tasks: env: PYTHONUNBUFFERED: "1" cmds: - - uv run uvicorn stirling.api.app:app --host 0.0.0.0 --port {{.PORT}} + - uv run uvicorn stirling.api.app:app --host 0.0.0.0 --port {{.PORT}} --workers "${STIRLING_ENGINE_WORKERS:-4}" dev: desc: "Start engine dev server with hot reload" diff --git a/.taskfiles/frontend.yml b/.taskfiles/frontend.yml index b2fc9d45d1..472e38729f 100644 --- a/.taskfiles/frontend.yml +++ b/.taskfiles/frontend.yml @@ -40,6 +40,21 @@ tasks: cmds: - node editor/scripts/generate-icons.js + prepare:og: + internal: true + run: when_changed + desc: "Regenerate OG/social-preview metadata from the tool registry" + cmds: + - node editor/scripts/generate-og-metadata.mjs + sources: + - editor/src/core/types/toolId.ts + - editor/src/core/utils/urlMapping.ts + - editor/src/core/data/useTranslatedToolRegistry.tsx + - editor/public/og_images/*.png + generates: + - editor/src/core/data/ogImageMap.json + - editor/public/og-metadata.json + prepare: desc: "Set up dev environment" run: when_changed @@ -49,6 +64,7 @@ tasks: - task: prepare:env vars: { MODE: '{{.MODE}}' } - prepare:icons + - prepare:og # ============================================================ # Development @@ -187,12 +203,23 @@ tasks: lint: desc: "Run linting" deps: [install] + cmds: + - task: lint:eslint + - task: lint:dpdm + + lint:eslint: + desc: "Run ESLint linting" + deps: [install] cmds: - npx eslint --max-warnings=0 - # Globs (not a bare dir) so dpdm walks the whole tree — `editor/src` - # alone matched only 2 files. dpdm expands the braces itself, so this is + + lint:dpdm: + desc: "Run circular import linting" + deps: [install] + cmds: + # Globs so dpdm walks the whole tree. dpdm expands the braces itself, so this is # shell-agnostic. Covers editor, portal, and the shared design system. - - 'npx dpdm "editor/src/**/*.{ts,tsx}" "portal/src/**/*.{ts,tsx}" "shared/**/*.{ts,tsx}" --circular --no-warning --no-tree --exit-code circular:1' + - npx dpdm "editor/src/**/*.{ts,tsx}" "portal/src/**/*.{ts,tsx}" "shared/**/*.{ts,tsx}" --circular --no-warning --no-tree --exit-code circular:1 lint:fix: desc: "Auto-fix lint issues" @@ -251,6 +278,12 @@ tasks: cmds: - npx tsc --noEmit --project editor/src/desktop/tsconfig.json + typecheck:cloud: + desc: "Typecheck cloud shared layer (standalone)" + deps: [prepare] + cmds: + - npx tsc --noEmit --project editor/src/cloud/tsconfig.json + typecheck:scripts: desc: "Typecheck scripts" deps: [prepare] @@ -282,6 +315,7 @@ tasks: - task: typecheck:proprietary - task: typecheck:saas - task: typecheck:desktop + - task: typecheck:cloud - task: typecheck:scripts - task: typecheck:prototypes - task: typecheck:portal @@ -299,9 +333,17 @@ tasks: - task: format:check - task: test + og:check: + desc: "Fail if committed OG/social-preview metadata is out of date" + cmds: + - node editor/scripts/generate-og-metadata.mjs --check + check:all: desc: "Full CI quality gate" cmds: + # Runs first, before prepare regenerates: guards the committed og-metadata.json / + # ogImageMap.json that the Cloudflare Pages (plain `vite build`) deploy relies on. + - task: og:check - task: typecheck:all - task: lint - task: format:check @@ -316,19 +358,19 @@ tasks: test: desc: "Run tests" - deps: [install] + deps: [prepare] cmds: - npx vitest run --root editor test:watch: desc: "Run tests in watch mode" - deps: [install] + deps: [prepare] cmds: - npx vitest --watch --root editor test:coverage: desc: "Run tests with coverage (one-shot; CI-friendly)." - deps: [install] + deps: [prepare] cmds: # `vitest run` makes this CI-safe (the bare `vitest` form enters watch # mode). Explicit reporter list because v8 + json-summary is what the @@ -356,3 +398,15 @@ tasks: deps: [install] cmds: - node editor/scripts/generate-licenses.js + + # ============================================================ + # Clean + # ============================================================ + + clean: + desc: "Clean build artifacts and caches" + cmds: + - cmd: powershell rm -Recurse -Force -ErrorAction SilentlyContinue node_modules/.vite, editor/dist, dist, dist-portal + platforms: [windows] + - cmd: rm -rf node_modules/.vite editor/dist dist dist-portal + platforms: [linux, darwin] diff --git a/.taskfiles/pre-commit.yml b/.taskfiles/pre-commit.yml new file mode 100644 index 0000000000..63fdd9ea01 --- /dev/null +++ b/.taskfiles/pre-commit.yml @@ -0,0 +1,159 @@ +version: '3' + +# Repo-wide lint/format/secret checks - the single source of truth that the git +# pre-commit hook (.pre-commit-config.yaml) and CI (pre_commit.yml) both call. + +vars: + GITLEAKS: '8.30.0' + + # File selections as git pathspecs: git does the include/exclude matching, so + # there is no grep/xargs and it behaves identically on every platform. + PY_FILES: >- + 'scripts/*.py' + '.github/scripts/*.py' + 'app/core/src/main/resources/static/python/*.py' + ':(exclude)*split_photos.py' + SPELL_FILES: >- + '*.html' + '*.css' + '*.js' + '*.py' + '*.md' + ':(exclude).vscode/*' + ':(exclude).devcontainer/*' + ':(exclude)app/core/src/main/resources/*' + ':(exclude)app/proprietary/src/main/resources/*' + ':(exclude)frontend/editor/public/vendor/*' + ':(exclude)*Dockerfile*' + ':(exclude)*pdfjs*' + ':(exclude)*thirdParty*' + ':(exclude)*bootstrap*' + ':(exclude)*.min.*' + ':(exclude)*diff.js' + WS_FILES: >- + '*.js' + '*.java' + '*.py' + '*.yml' + ':(exclude)*pdfjs*' + ':(exclude)*thirdParty*' + ':(exclude)*bootstrap*' + ':(exclude)*.min.*' + ':(exclude)*diff.js' + ':(exclude).github/workflows/*' + LOCALE_TOML: 'frontend/editor/public/locales/*/translation.toml' + + GITLEAKS_BIN: '.task/bin/gitleaks-{{.GITLEAKS}}{{if eq OS "windows"}}.exe{{end}}' + +tasks: + default: + desc: "Check formatting, spelling, and secrets across the repo" + cmds: + - task: ruff + - task: ruff-format + - task: codespell + - task: gitleaks + - task: whitespace + - task: toml-sort + + fix: + desc: "Auto-fix formatting, spelling, and secrets issues across the repo" + cmds: + # Auto-fixers first, then the report-only tools (codespell, gitleaks) so a + # finding there does not stop the fixers from running. + - task: ruff + vars: { FIX: '1' } + - task: ruff-format + vars: { FIX: '1' } + - task: whitespace + vars: { FIX: '1' } + - task: toml-sort + vars: { FIX: '1' } + - task: codespell + - task: gitleaks + + install: + desc: "Install the pinned pre-commit Python tools (ruff, codespell, toml-sort)" + run: once + cmds: + - uv sync --project scripts/pre-commit --locked + sources: + - scripts/pre-commit/uv.lock + - scripts/pre-commit/pyproject.toml + status: + - test -d scripts/pre-commit/.venv + + clean: + desc: "Remove the cache/build artifacts" + cmds: + - task: '{{if eq OS "windows"}}clean-windows{{else}}clean-unix{{end}}' + + clean-unix: + internal: true + cmds: + - rm -rf scripts/pre-commit/.venv .task/bin/gitleaks-* + + # On Windows, use PowerShell so it matches the same paths and tolerates absent + # files without erroring. + clean-windows: + internal: true + ignore_error: true + cmds: + - powershell -NoProfile -Command "Remove-Item -Recurse -Force -ErrorAction SilentlyContinue scripts/pre-commit/.venv, .task/bin/gitleaks-*" + + # Individual checks (hidden from `task --list`, but callable, e.g. + # `task pre-commit:toml-sort FIX=1`). Pass FIX=1 to auto-fix where supported. + ruff: + deps: [install] + cmds: + - uv run --project scripts/pre-commit --no-sync ruff check --line-length=127 {{if .FIX}}--fix {{end}}$(git ls-files {{.PY_FILES}}) + + ruff-format: + deps: [install] + cmds: + - uv run --project scripts/pre-commit --no-sync ruff format {{if .FIX}}{{else}}--check {{end}}$(git ls-files {{.PY_FILES}}) + + codespell: + deps: [install] + cmds: + - uv run --project scripts/pre-commit --no-sync codespell --ignore-words-list=thirdParty,tabEl,tabEls,Sie,ist,fulfilment --quiet-level=2 $(git ls-files {{.SPELL_FILES}}) + + toml-sort: + deps: [install] + cmds: + - uv run --project scripts/pre-commit --no-sync toml-sort --all --ignore-case {{if .FIX}}--in-place{{else}}--check{{end}} {{.LOCALE_TOML}} + + whitespace: + cmds: + - uv run --no-project python scripts/pre-commit/whitespace.py {{if .FIX}}--fix {{end}}$(git ls-files {{.WS_FILES}}) + + gitleaks: + deps: [gitleaks-bin] + # Scan staged changes only, matching the old hook: the git-mode fingerprints + # in .gitleaksignore (file:rule:line) still apply, and with nothing staged + # this is a no-op. Secrets are never auto-fixed, so FIX has no effect. + cmds: + - "{{.GITLEAKS_BIN}} git --pre-commit --redact --staged --verbose" + + gitleaks-bin: + internal: true + desc: "Ensure the pinned gitleaks binary is cached in .task/bin" + status: + - test -f {{.GITLEAKS_BIN}} + vars: + GL_ARCH: '{{if eq ARCH "amd64"}}x64{{else if eq ARCH "arm64"}}arm64{{else if eq ARCH "386"}}x32{{else}}{{ARCH}}{{end}}' + GL_PLATFORM: '{{OS}}_{{.GL_ARCH}}' + GL_URL: 'https://github.com/gitleaks/gitleaks/releases/download/v{{.GITLEAKS}}/gitleaks_{{.GITLEAKS}}_{{.GL_PLATFORM}}' + # SHA-256 of each release asset, from gitleaks_{{.GITLEAKS}}_checksums.txt. + GL_SHA: >- + {{if eq .GL_PLATFORM "linux_x64"}}79a3ab579b53f71efd634f3aaf7e04a0fa0cf206b7ed434638d1547a2470a66e + {{- else if eq .GL_PLATFORM "linux_arm64"}}b4cbbb6ddf7d1b2a603088cd03a4e3f7ce48ee7fd449b51f7de6ee2906f5fa2f + {{- else if eq .GL_PLATFORM "darwin_x64"}}ca221d012d247080c2f6f61f4b7a83bffa2453806b0c195c795bbe9a8c775ed5 + {{- else if eq .GL_PLATFORM "darwin_arm64"}}b251ab2bcd4cd8ba9e56ff37698c033ebf38582b477d21ebd86586d927cf87e7 + {{- else if eq .GL_PLATFORM "windows_x64"}}54fe94f644b832dd08e8c3a5915efb3bfa862386d59fb27ca0792cb687a83573 + {{- end}} + cmds: + - cmd: bash scripts/pre-commit/install-gitleaks.sh "{{.GL_URL}}.tar.gz" "{{.GL_SHA}}" "{{.GITLEAKS_BIN}}" + platforms: [linux, darwin] + - cmd: powershell -NoProfile -File scripts/pre-commit/install-gitleaks.ps1 -Url "{{.GL_URL}}.zip" -Sha "{{.GL_SHA}}" -Dest "{{.GITLEAKS_BIN}}" + platforms: [windows] diff --git a/ADDING_TOOLS.md b/ADDING_TOOLS.md index 579a9647a5..9c22246a7c 100644 --- a/ADDING_TOOLS.md +++ b/ADDING_TOOLS.md @@ -200,9 +200,9 @@ const [ToolName] = (props: BaseToolProps) => { ``` ## 5. Add Translations -Update translation files. **Important: Only update `en-GB` files** - other languages are handled separately. +Update translation files. **Important: Only update `en-US` files** - other languages are handled separately. -**File to update:** `frontend/editor/public/locales/en-GB/translation.toml` +**File to update:** `frontend/editor/public/locales/en-US/translation.toml` **Required Translation Keys**: ```toml @@ -251,7 +251,7 @@ Update translation files. **Important: Only update `en-GB` files** - other langu ``` **Translation Notes:** -- **Only update `en-GB/translation.toml`** - other locale files are managed separately +- **Only update `en-US/translation.toml`** - other locale files are managed separately - Use descriptive keys that match your component's `t()` calls - Include tooltip translations if you created tooltip hooks - Add `options.*` keys if your tool has settings with descriptions diff --git a/AGENTS.md b/AGENTS.md index 7c89413c49..ae8eb60316 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -169,7 +169,31 @@ import { useFileContext } from "@proprietary/contexts/FileContext"; - Building layer-specific override that wraps a lower layer's component - Example: `import { AppProviders as CoreAppProviders } from "@core/components/AppProviders"` when creating proprietary/AppProviders.tsx that extends the core version -The `@app/*` alias automatically resolves to the correct layer based on build target (core/proprietary/desktop) and handles the fallback cascade. +The `@app/*` alias automatically resolves to the correct layer based on build target (core/proprietary/saas/desktop/cloud) and handles the fallback cascade — see "Frontend `cloud/` Layer" below for the full per-flavor order. + +#### Frontend `cloud/` Layer + +`@app/*` resolves through a per-flavor cascade — first existing file wins (shadow/override): + +- **core** → core +- **proprietary** → proprietary → core +- **saas** → saas → cloud → proprietary → core +- **desktop** → desktop → cloud → proprietary → core +- **cloud** → cloud → proprietary → core + +What goes where: + +- **core** — OSS base. +- **proprietary** — licensed / offline features. +- **cloud** — the SHARED hosted/SaaS experience used by BOTH saas + desktop: PAYG, wallet, plan, billing, usage meters, cloud config/team/onboarding. +- **saas** — web-only: Supabase web auth, AuthCallback, avatar canvas, `window.location`. +- **desktop** — Tauri-only: keyring authService, tauriHttpClient, native files/windows, backend routing. + +`cloud/` MUST NOT import `@supabase/*`, `@tauri-apps/*`, raw `fetch`, `window.location`, `localStorage`, `sessionStorage`, or `import.meta.env.VITE_*` (enforced by ESLint). It reaches platform-specific things only via `@app/*` seams: `services/apiClient`, `auth/session.getAccessToken`, `auth/supabase`, `platform/openExternal`, `services/billing`, `hooks/useSaaSMode` — each provided per-platform in `saas/` and `desktop/`. + +Rule of thumb — **move, don't copy**: share via `cloud/`, override by shadowing the same `@app/*` path in a leaf (`saas/` or `desktop/`). + +**Cloud feature flags on desktop.** The local `AppConfigContext` reads `/api/v1/config/app-config` from the LOCAL bundled backend, so cloud-only flags (`aiEngineEnabled`, `premiumEnabled`, …) are never seen on desktop. To read the cloud's view, use `useSaasAppConfig()` (`desktop/hooks/useSaasAppConfig.ts`, backed by the general `saasAppConfigService` — SaaS-mode-only, public endpoint, native HTTP, 5-min cache). It returns `null` outside SaaS mode, so cloud features stay off in local/self-hosted and the server keeps the on/off switch (no desktop release needed to flip a flag). Gate a feature behind a per-platform seam — e.g. `useAiEngineEnabled()` (core reads `useAppConfig()`, desktop reads `useSaasAppConfig()`) — rather than hardcoding the flag on. #### Component Override Pattern (Stub/Shadow) Use this pattern for desktop-specific or proprietary-specific features WITHOUT runtime checks or conditionals. @@ -426,7 +450,7 @@ The frontend is organized with a clear separation of concerns: ## Translation Rules -- **CRITICAL**: Always update translations in `en-GB` only, never `en-US` +- **CRITICAL**: Always update translations in `en-US` only - all other languages (including `en-GB`) are handled separately - Translation files are located in `frontend/editor/public/locales/` ## Important Notes diff --git a/DeveloperGuide.md b/DeveloperGuide.md index f6f87e8243..23b2652ce1 100644 --- a/DeveloperGuide.md +++ b/DeveloperGuide.md @@ -52,6 +52,17 @@ This guide focuses on developing for Stirling 2.0, including both the React fron - Rust and Cargo (required for Tauri desktop app development) - Tauri CLI (install with `cargo install tauri-cli`) +### Optional System Dependencies + +These are not required to run the app but enable specific features. The app detects them at startup and disables the relevant features if they are missing. + +| Dependency | Feature | Install | +|---|---|---| +| LibreOffice | File-to-PDF conversions | `brew install libreoffice` / `apt install libreoffice` | +| Tesseract | OCR | `brew install tesseract` / `apt install tesseract-ocr` | +| WeasyPrint | AI document creation | `brew install weasyprint` / `apt install weasyprint` | +| qpdf | PDF optimisation | `brew install qpdf` / `apt install qpdf` | + ### Setup Steps 1. Clone the repository: @@ -576,7 +587,7 @@ When adding a new feature or modifying existing ones in Stirling-PDF, you'll nee Find the existing `messages.properties` files in the `stirling-pdf/src/main/resources` directory. You'll see files like: - `messages.properties` (default, usually English) -- `messages_en_GB.properties` +- `messages_en_US.properties` - `messages_fr_FR.properties` - `messages_de_DE.properties` - etc. diff --git a/LICENSE b/LICENSE index efc31d3a4d..2b2f7fc085 100644 --- a/LICENSE +++ b/LICENSE @@ -16,6 +16,8 @@ if that directory exists, is licensed under the license defined in "frontend/edi if that directory exists, is licensed under the license defined in "frontend/editor/src/desktop/LICENSE". * All content that resides under the "frontend/editor/src/saas/" directory of this repository, if that directory exists, is licensed under the license defined in "frontend/editor/src/saas/LICENSE". +* All content that resides under the "frontend/editor/src/cloud/" directory of this repository, +if that directory exists, is licensed under the license defined in "frontend/editor/src/cloud/LICENSE". * All content that resides under the "frontend/editor/src/prototypes/" directory of this repository, if that directory exists, is licensed under the license defined in "frontend/editor/src/prototypes/LICENSE". * All content that resides under the "frontend/portal/" directory of this repository, diff --git a/README.md b/README.md index 9329b20eed..c1d96eca1d 100644 --- a/README.md +++ b/README.md @@ -60,7 +60,7 @@ For full installation options (including desktop and Kubernetes), see our [Docum We welcome contributions! Please see [CONTRIBUTING.md](CONTRIBUTING.md) for guidelines. -This project uses [Task](https://taskfile.dev/) as a unified command runner for all build, dev, and test commands. Run `task install` to get started, or see the [Developer Guide](DeveloperGuide.md) for full details. +This project uses [Task](https://taskfile.dev/) as a unified command runner for all build, dev, and test commands. Run `task dev` to get started running the editor, run `task` to see the most common commands, or see the [Developer Guide](DeveloperGuide.md) for full details. For adding translations, see the [Translation Guide](devGuide/HowToAddNewLanguage.md). diff --git a/Taskfile.yml b/Taskfile.yml index dc4d5130bd..ac4160182a 100644 --- a/Taskfile.yml +++ b/Taskfile.yml @@ -25,8 +25,28 @@ includes: e2e: taskfile: .taskfiles/e2e.yml dir: . + pre-commit: + taskfile: .taskfiles/pre-commit.yml + dir: . tasks: + # ============================================================ + # Help (shown when you run `task` with no arguments) + # ============================================================ + + default: + desc: "List the most common commands" + silent: true + cmds: + - | + echo "Common commands (run 'task --list' to see all):" + echo "" + echo " task dev Start backend & frontend on free ports" + echo " task backend:dev Start backend on default port" + echo " task frontend:dev Start frontend on default port" + echo " task desktop:dev Start desktop app" + echo " task check Quality gate (lint, typecheck, test, etc.)" + # ============================================================ # Setup & Prerequisites # ============================================================ @@ -60,24 +80,20 @@ tasks: dev:saas: desc: "Start SaaS backend + frontend concurrently on free ports" - vars: - PORTS: - sh: '{{if eq OS "windows"}}{{.FIND_FREE_PORT_PS}} 8080 5173{{else}}{{.FIND_FREE_PORT_SH}} 8080 5173{{end}}' - BACKEND_PORT: '{{index (splitList "\n" .PORTS) 0}}' - FRONTEND_PORT: '{{index (splitList "\n" .PORTS) 1}}' - deps: - - task: backend:dev:saas - vars: - PORT: '{{.BACKEND_PORT}}' - - task: frontend:dev:saas - vars: - PORT: '{{.FRONTEND_PORT}}' - BACKEND_URL: 'http://localhost:{{.BACKEND_PORT}}' - OPEN: "true" + cmds: + - task: dev:_all + vars: { FRONTEND: saas, BACKEND: saas } dev:all: desc: "Start backend + frontend + engine concurrently on free ports" + cmds: + - task: dev:_all + + dev:_all: + internal: true vars: + FRONTEND: '{{.FRONTEND | default "proprietary"}}' + BACKEND: '{{.BACKEND | default "proprietary"}}' PORTS: sh: '{{if eq OS "windows"}}{{.FIND_FREE_PORT_PS}} 8080 5173 5001{{else}}{{.FIND_FREE_PORT_SH}} 8080 5173 5001{{end}}' BACKEND_PORT: '{{index (splitList "\n" .PORTS) 0}}' @@ -87,11 +103,12 @@ tasks: - task: engine:dev vars: PORT: '{{.ENGINE_PORT}}' - - task: backend:dev + - task: 'backend:dev:{{.BACKEND}}' vars: PORT: '{{.BACKEND_PORT}}' AIENGINE_URL: 'http://localhost:{{.ENGINE_PORT}}' - - task: frontend:dev + AIENGINE_ENABLED: "true" + - task: 'frontend:dev:{{.FRONTEND}}' vars: PORT: '{{.FRONTEND_PORT}}' BACKEND_URL: 'http://localhost:{{.BACKEND_PORT}}' @@ -175,4 +192,6 @@ tasks: desc: "Clean all build artifacts" cmds: - task: backend:clean + - task: frontend:clean - task: engine:clean + - task: pre-commit:clean diff --git a/app/common/src/main/java/stirling/software/SPDF/pdf/parser/CompositeTableParser.java b/app/common/src/main/java/stirling/software/SPDF/pdf/parser/CompositeTableParser.java deleted file mode 100644 index 429f180f3e..0000000000 --- a/app/common/src/main/java/stirling/software/SPDF/pdf/parser/CompositeTableParser.java +++ /dev/null @@ -1,73 +0,0 @@ -package stirling.software.SPDF.pdf.parser; - -import static stirling.software.SPDF.pdf.parser.PdfModels.*; - -import java.io.IOException; -import java.util.List; - -import org.apache.pdfbox.pdmodel.PDDocument; -import org.springframework.context.annotation.Primary; -import org.springframework.stereotype.Service; - -import lombok.RequiredArgsConstructor; -import lombok.extern.slf4j.Slf4j; - -/** - * Chains table parsers in priority order: Tabula lattice → Tabula stream → {@link - * LineAlignmentTableParser}. The first parser returning a result above {@link - * #TABULA_CONFIDENCE_THRESHOLD} wins; results from different parsers are never mixed on one page. - */ -@Service -@Primary -@RequiredArgsConstructor -@Slf4j -public class CompositeTableParser implements TableParser { - - /** Min Tabula confidence to accept results; below this LineAlignment is tried instead. */ - static final float TABULA_CONFIDENCE_THRESHOLD = 0.5f; - - private final TabulaTableParser tabulaParser; - private final LineAlignmentTableParser lineAlignmentParser; - - @Override - public List parse(PDDocument document, RawPage rawPage) throws IOException { - // Step 1: Tabula lattice mode (ruled/bordered tables). - List latticeResults = filterConfident(tabulaParser.parse(document, rawPage)); - if (!latticeResults.isEmpty()) { - log.debug( - "Page {}: using Tabula lattice ({} table(s))", - rawPage.pageNumber(), - latticeResults.size()); - return latticeResults; - } - - // Step 2: Tabula stream mode (borderless/whitespace-delimited tables). - // parseStream is not on the TableParser interface — this intentionally couples to the - // concrete TabulaTableParser since stream mode is a Tabula-specific concept. - List streamResults = - filterConfident(tabulaParser.parseStream(document, rawPage)); - if (!streamResults.isEmpty()) { - log.debug( - "Page {}: using Tabula stream ({} table(s))", - rawPage.pageNumber(), - streamResults.size()); - return streamResults; - } - - // Step 3: Geometry-based line-alignment fallback. - List lineResults = lineAlignmentParser.parse(document, rawPage); - if (!lineResults.isEmpty()) { - log.debug( - "Page {}: using LineAlignment ({} table(s))", - rawPage.pageNumber(), - lineResults.size()); - return lineResults; - } - - return List.of(); - } - - private List filterConfident(List tables) { - return tables.stream().filter(t -> t.confidence() >= TABULA_CONFIDENCE_THRESHOLD).toList(); - } -} diff --git a/app/common/src/main/java/stirling/software/SPDF/pdf/parser/LineAlignmentTableParser.java b/app/common/src/main/java/stirling/software/SPDF/pdf/parser/LineAlignmentTableParser.java deleted file mode 100644 index b2d8de5167..0000000000 --- a/app/common/src/main/java/stirling/software/SPDF/pdf/parser/LineAlignmentTableParser.java +++ /dev/null @@ -1,528 +0,0 @@ -package stirling.software.SPDF.pdf.parser; - -import static stirling.software.SPDF.pdf.parser.PdfModels.*; - -import java.io.IOException; -import java.util.ArrayList; -import java.util.Arrays; -import java.util.Collections; -import java.util.Comparator; -import java.util.HashMap; -import java.util.List; -import java.util.Map; -import java.util.Optional; -import java.util.TreeMap; -import java.util.regex.Pattern; - -import org.apache.pdfbox.pdmodel.PDDocument; -import org.springframework.stereotype.Service; - -import lombok.extern.slf4j.Slf4j; - -/** - * Fallback {@link TableParser} for borderless financial tables using text geometry. - * - *

Identifies "anchor lines" (≥2 numeric tokens), builds a column grid from their right-edge - * positions, groups vertically proximate anchor lines into table candidates, then scores each group - * on column consistency and anchor density (confidence ceiling 0.85). - */ -@Service -@Slf4j -public class LineAlignmentTableParser implements TableParser { - - /** Width in points of each column position bucket. */ - static final float COLUMN_BUCKET_PT = 5f; - - /** Tolerance in buckets when matching a token's right-edge to a confirmed column position. */ - private static final int COLUMN_MATCH_BUCKETS = 2; - - /** Maximum gap (as a multiple of modal line spacing) before splitting a group. */ - private static final float MAX_GAP_FACTOR = 2.5f; - - /** Minimum anchor rows (numeric-heavy) to form a valid table. */ - static final int MIN_TABLE_ROWS = 3; - - /** Minimum confirmed column positions to form a valid table. */ - static final int MIN_COLUMNS = 2; - - /** - * Min fraction of anchor lines a column must appear on to be confirmed (permissive for N/A - * rows). - */ - private static final double COLUMN_MIN_FREQUENCY = 0.40; - - /** - * Matches financial numeric tokens: integers, decimals, parenthetical negatives, currency, - * percent, nil dashes. - */ - private static final Pattern NUMERIC = - Pattern.compile("^[\\(\\-\\$£€¥]?\\d[\\d,\\.]*[\\)%]?$|^[-–—]$"); - - /** - * Lines within this y-distance are merged into one row (restores rows split by LineBuilder's - * column-gap logic). - */ - static final float ROW_MERGE_TOLERANCE_PT = 2f; - - // ── public API ─────────────────────────────────────────────────────────────────────────────── - - @Override - public List parse(PDDocument document, RawPage rawPage) throws IOException { - List lines = rawPage.lines(); - if (lines.size() < MIN_TABLE_ROWS) return List.of(); - - float modalSpacing = computeModalSpacing(lines); - List tokenized = - mergeCoincidentLines(lines.stream().map(this::tokenize).toList()); - - List anchors = tokenized.stream().filter(TokenizedLine::isAnchor).toList(); - - if (anchors.size() < MIN_TABLE_ROWS) return List.of(); - - List columnGrid = buildColumnGrid(anchors); - if (columnGrid.size() < MIN_COLUMNS) { - log.debug( - "Page {}: LineAlignment — fewer than {} confirmed columns, skipping", - rawPage.pageNumber(), - MIN_COLUMNS); - return List.of(); - } - - List> groups = groupRows(tokenized, columnGrid, modalSpacing); - - List results = new ArrayList<>(); - for (int i = 0; i < groups.size(); i++) { - buildFragment(groups.get(i), columnGrid, rawPage.pageNumber(), i) - .ifPresent(results::add); - } - - log.debug( - "Page {}: LineAlignment detected {} table(s) ({} anchor lines, {} columns)", - rawPage.pageNumber(), - results.size(), - anchors.size(), - columnGrid.size()); - return results; - } - - // ── coincident-line merging ────────────────────────────────────────────────────────────────── - - /** - * Merges tokenised lines sharing the same y-position into one row, rejoining label/value halves - * split by LineBuilder. - */ - List mergeCoincidentLines(List tokenized) { - if (tokenized.size() < 2) return tokenized; - - List result = new ArrayList<>(); - int i = 0; - - while (i < tokenized.size()) { - float baseY = tokenized.get(i).line().bounds().y(); - int j = i + 1; - while (j < tokenized.size() - && Math.abs(tokenized.get(j).line().bounds().y() - baseY) - <= ROW_MERGE_TOLERANCE_PT) { - j++; - } - - if (j == i + 1) { - result.add(tokenized.get(i)); - } else { - result.add(mergeGroup(tokenized.subList(i, j))); - } - i = j; - } - - return result; - } - - private TokenizedLine mergeGroup(List group) { - List mergedFragments = - group.stream() - .flatMap(tl -> tl.line().fragments().stream()) - .sorted(Comparator.comparingDouble(f -> f.bounds().x())) - .toList(); - - Bounds mergedBounds = - group.stream() - .map(tl -> tl.line().bounds()) - .reduce(Bounds::merge) - .orElse(group.get(0).line().bounds()); - - RawLine mergedLine = - new RawLine( - group.get(0).line().lineId(), - mergedFragments, - mergedBounds, - group.get(0).line().pageNumber()); - - return tokenize(mergedLine); - } - - // ── tokenisation ───────────────────────────────────────────────────────────────────────────── - - /** - * Splits fragments into word-level tokens; x-positions are estimated linearly within each - * fragment. - */ - TokenizedLine tokenize(RawLine line) { - List tokens = new ArrayList<>(); - for (TextFragment frag : line.fragments()) { - tokens.addAll(tokensFromFragment(frag)); - } - List numeric = tokens.stream().filter(LineToken::numeric).toList(); - return new TokenizedLine(line, tokens, numeric); - } - - private List tokensFromFragment(TextFragment frag) { - String raw = frag.text(); - if (raw == null || raw.isBlank()) return List.of(); - - float fragX = frag.bounds().x(); - float fragWidth = frag.bounds().width(); - int rawLen = raw.length(); - - List result = new ArrayList<>(); - int offset = 0; - for (String part : raw.split("\\s+")) { - if (part.isEmpty()) { - offset++; - continue; - } - int idx = raw.indexOf(part, offset); - if (idx < 0) idx = offset; - - float tokenX = rawLen > 0 ? fragX + ((float) idx / rawLen) * fragWidth : fragX; - float tokenRight = - rawLen > 0 - ? fragX + ((float) (idx + part.length()) / rawLen) * fragWidth - : fragX + fragWidth; - - result.add(new LineToken(part, tokenX, tokenRight, NUMERIC.matcher(part).matches())); - offset = idx + part.length(); - } - return result; - } - - // ── column grid ────────────────────────────────────────────────────────────────────────────── - - /** - * Returns confirmed column right-edge positions — those appearing on ≥ {@value - * #COLUMN_MIN_FREQUENCY} × N anchor lines. - */ - private List buildColumnGrid(List anchors) { - // bucket → set of line indices that contributed a numeric token to that bucket - Map> bucketLines = new HashMap<>(); - for (int i = 0; i < anchors.size(); i++) { - for (LineToken t : anchors.get(i).numeric()) { - int bucket = bucket(t.right()); - bucketLines.computeIfAbsent(bucket, k -> new ArrayList<>()).add(i); - } - } - - int minHits = - Math.max(MIN_TABLE_ROWS, (int) Math.ceil(anchors.size() * COLUMN_MIN_FREQUENCY)); - - // Confirmed buckets → average right-edge for that bucket - TreeMap confirmed = new TreeMap<>(); - for (Map.Entry> entry : bucketLines.entrySet()) { - // Count distinct lines - long distinctLines = entry.getValue().stream().distinct().count(); - if (distinctLines >= minHits) { - double avg = - entry.getValue().stream() - .distinct() // weight each line equally regardless of token count - .mapToDouble( - lineIdx -> - avgRightEdgeForBucket( - anchors, lineIdx, entry.getKey())) - .average() - .orElse(entry.getKey() * (double) COLUMN_BUCKET_PT); - confirmed.put(entry.getKey(), (float) avg); - } - } - - return new ArrayList<>(confirmed.values()); // already sorted by bucket (left to right) - } - - /** - * Returns the average right-edge position of tokens in {@code line} whose bucket matches {@code - * targetBucket}, falling back to the bucket's nominal centre when no tokens match. - */ - private double avgRightEdgeForBucket( - List anchors, int lineIdx, int targetBucket) { - return anchors.get(lineIdx).numeric().stream() - .filter(t -> bucket(t.right()) == targetBucket) - .mapToDouble(LineToken::right) - .average() - .orElse(targetBucket * (double) COLUMN_BUCKET_PT); - } - - // ── grouping ───────────────────────────────────────────────────────────────────────────────── - - /** - * Groups anchor lines into table candidates, including adjacent label rows; a gap > - * MAX_GAP_FACTOR × modal spacing splits groups. - */ - private List> groupRows( - List all, List columnGrid, float modalSpacing) { - float maxGap = modalSpacing > 0 ? modalSpacing * MAX_GAP_FACTOR : 30f; - - List> groups = new ArrayList<>(); - List current = new ArrayList<>(); - - for (int i = 0; i < all.size(); i++) { - TokenizedLine tl = all.get(i); - boolean fits = tl.isAnchor() && matchesGrid(tl, columnGrid); - - if (current.isEmpty()) { - if (fits) current.add(tl); - continue; - } - - float gap = - tl.line().bounds().y() - - current.get(current.size() - 1).line().bounds().bottom(); - - if (gap > maxGap) { - groups.add(current); - current = new ArrayList<>(); - if (fits) current.add(tl); - continue; - } - - if (fits) { - current.add(tl); - } else if (!tl.line().text().isBlank()) { - // Include non-anchor lines (labels) only if they have text and are within - // proximity. - current.add(tl); - } - } - - if (!current.isEmpty()) groups.add(current); - - return groups.stream().filter(g -> hasEnoughAnchorRows(g, columnGrid)).toList(); - } - - private boolean hasEnoughAnchorRows(List group, List columnGrid) { - return group.stream().filter(r -> r.isAnchor() && matchesGrid(r, columnGrid)).count() - >= MIN_TABLE_ROWS; - } - - /** A line "matches" the grid when ≥ 60 % of its numeric tokens land in confirmed columns. */ - private boolean matchesGrid(TokenizedLine tl, List columnGrid) { - if (tl.numeric().isEmpty()) return false; - long matches = - tl.numeric().stream() - .filter(t -> nearestColumnIndex(t.right(), columnGrid) >= 0) - .count(); - return (double) matches / tl.numeric().size() >= 0.60; - } - - private boolean hasInconsistentColumnMatch(TokenizedLine tl, List columnGrid) { - if (tl.numeric().isEmpty()) return false; - long hits = - tl.numeric().stream() - .filter(t -> nearestColumnIndex(t.right(), columnGrid) >= 0) - .count(); - return (double) hits / tl.numeric().size() < 0.60; - } - - // ── fragment assembly ──────────────────────────────────────────────────────────────────────── - - private Optional buildFragment( - List group, List columnGrid, int pageNumber, int tableIndex) { - - long anchorCount = - group.stream().filter(r -> r.isAnchor() && matchesGrid(r, columnGrid)).count(); - if (anchorCount < MIN_TABLE_ROWS) return Optional.empty(); - - List warnings = new ArrayList<>(); - List> rawRows = new ArrayList<>(); - List rows = new ArrayList<>(); - - for (int rowIdx = 0; rowIdx < group.size(); rowIdx++) { - TokenizedLine tl = group.get(rowIdx); - List rawRow = buildRawRow(tl, columnGrid); - rawRows.add(Collections.unmodifiableList(rawRow)); - rows.add(buildTableRow(rowIdx, tl, rawRow, columnGrid)); - } - - // Column count = 1 label column + confirmed numeric columns - int colCount = columnGrid.size() + 1; - Bounds bounds = computeGroupBounds(group); - float confidence = computeConfidence(group, columnGrid, warnings); - - return Optional.of( - new TableFragment( - "tbl-la-p" + pageNumber + "-" + tableIndex, - pageNumber, - bounds, - List.of(), - Collections.unmodifiableList(rows), - Collections.unmodifiableList(rawRows), - colCount, - confidence, - Collections.unmodifiableList(warnings), - null)); - } - - /** - * Builds a raw row as a list of strings: index 0 = label text, indices 1..N = column values. - */ - private List buildRawRow(TokenizedLine tl, List columnGrid) { - String[] cells = new String[columnGrid.size() + 1]; - Arrays.fill(cells, ""); - - // Separate label tokens (those not landing in any confirmed column) from column tokens. - List labelParts = new ArrayList<>(); - for (LineToken token : tl.all()) { - int col = nearestColumnIndex(token.right(), columnGrid); - if (col >= 0 && token.numeric()) { - int cellIdx = col + 1; - cells[cellIdx] = - cells[cellIdx].isEmpty() - ? token.text() - : cells[cellIdx] + " " + token.text(); - } else { - labelParts.add(token.text()); - } - } - cells[0] = String.join(" ", labelParts).trim(); - return Arrays.asList(cells); - } - - private TableRow buildTableRow( - int rowIdx, TokenizedLine tl, List rawRow, List columnGrid) { - List cells = new ArrayList<>(rawRow.size()); - - // Label cell: use the line's full bounds as an approximation. - cells.add(TableCell.of(0, rawRow.get(0), tl.line().bounds())); - - for (int col = 0; col < columnGrid.size(); col++) { - String text = col + 1 < rawRow.size() ? rawRow.get(col + 1) : ""; - float right = columnGrid.get(col); - float left = col > 0 ? columnGrid.get(col - 1) : right - 50f; - Bounds cellBounds = - new Bounds( - left, - tl.line().bounds().y(), - right - left, - tl.line().bounds().height()); - cells.add(TableCell.of(col + 1, text, cellBounds)); - } - return new TableRow(rowIdx, Collections.unmodifiableList(cells)); - } - - // ── confidence scoring ─────────────────────────────────────────────────────────────────────── - - /** - * Heuristic score in [0.0, 0.85] (ceiling keeps results below Tabula lattice which starts at - * 1.0). Base 0.70; +0.05/col beyond 2 (max +0.10); +0.05 at ≥5 anchors, +0.05 at ≥8; −0.15 if - * >30 % of anchors have inconsistent columns; −0.10 if non-anchors outnumber anchors. - */ - private float computeConfidence( - List group, List columnGrid, List warnings) { - float score = 0.70f; - - long anchorCount = - group.stream().filter(r -> r.isAnchor() && matchesGrid(r, columnGrid)).count(); - long totalRows = group.size(); - - // More columns - int extraCols = Math.min(columnGrid.size() - MIN_COLUMNS, 2); - score += extraCols * 0.05f; - - // More anchor rows - if (anchorCount >= 5) score += 0.05f; - if (anchorCount >= 8) score += 0.05f; - - // Inconsistent column matching - long inconsistent = - group.stream() - .filter(TokenizedLine::isAnchor) - .filter(tl -> hasInconsistentColumnMatch(tl, columnGrid)) - .count(); - if (inconsistent > anchorCount * 0.30) { - score -= 0.15f; - warnings.add( - "Column match inconsistent on " - + inconsistent - + "/" - + anchorCount - + " anchor rows"); - } - - // Label-heavy - long nonAnchor = totalRows - anchorCount; - if (nonAnchor > anchorCount) { - score -= 0.10f; - warnings.add( - "Non-anchor rows (" - + nonAnchor - + ") outnumber anchor rows (" - + anchorCount - + ")"); - } - - return Math.max(0f, Math.min(0.85f, score)); - } - - // ── utility ────────────────────────────────────────────────────────────────────────────────── - - /** - * Returns the grid index nearest to {@code rightEdge}, or -1 if none is within {@value - * #COLUMN_MATCH_BUCKETS} buckets. - */ - private int nearestColumnIndex(float rightEdge, List grid) { - int nearest = -1; - float minDist = COLUMN_MATCH_BUCKETS * COLUMN_BUCKET_PT + 1f; - for (int i = 0; i < grid.size(); i++) { - float dist = Math.abs(rightEdge - grid.get(i)); - if (dist < minDist) { - minDist = dist; - nearest = i; - } - } - return nearest; - } - - private Bounds computeGroupBounds(List group) { - return group.stream() - .map(tl -> tl.line().bounds()) - .reduce(Bounds::merge) - .orElse(new Bounds(0, 0, 0, 0)); - } - - /** Modal gap between consecutive line edges, used to calibrate the group-split threshold. */ - private float computeModalSpacing(List lines) { - if (lines.size() < 2) return 0f; - Map freq = new HashMap<>(); - for (int i = 1; i < lines.size(); i++) { - float gap = lines.get(i).bounds().y() - lines.get(i - 1).bounds().bottom(); - if (gap > 0) freq.merge(Math.round(gap / 2f) * 2f, 1L, Long::sum); - } - return freq.entrySet().stream() - .max(Map.Entry.comparingByValue()) - .map(Map.Entry::getKey) - .orElse(0f); - } - - private static int bucket(float x) { - return Math.round(x / COLUMN_BUCKET_PT); - } - - // ── private data types ─────────────────────────────────────────────────────────────────────── - - /** A word-level token with an approximate right-edge x-position. */ - record LineToken(String text, float x, float right, boolean numeric) {} - - /** A {@link RawLine} with tokens pre-computed; an "anchor" has ≥ 2 numeric tokens. */ - record TokenizedLine(RawLine line, List all, List numeric) { - boolean isAnchor() { - return numeric.size() >= 2; - } - } -} diff --git a/app/common/src/main/java/stirling/software/SPDF/pdf/parser/LineBuilder.java b/app/common/src/main/java/stirling/software/SPDF/pdf/parser/LineBuilder.java deleted file mode 100644 index 6831f6d734..0000000000 --- a/app/common/src/main/java/stirling/software/SPDF/pdf/parser/LineBuilder.java +++ /dev/null @@ -1,139 +0,0 @@ -package stirling.software.SPDF.pdf.parser; - -import static stirling.software.SPDF.pdf.parser.PdfModels.*; - -import java.util.ArrayList; -import java.util.Comparator; -import java.util.List; - -import org.springframework.stereotype.Service; - -import lombok.extern.slf4j.Slf4j; - -/** - * Groups {@link TextFragment} objects into visual {@link RawLine}s using baseline proximity. - * - *

Fragments are on the same line when their baselines are within a font-size-derived tolerance. - * A new line starts whenever the horizontal gap exceeds an adaptive column-gap threshold ({@code - * max(effectiveWidth * COLUMN_GAP_RATIO, COLUMN_GAP_MIN_PT)}), splitting two-column text. - */ -@Service -@Slf4j -public class LineBuilder { - - /** Baseline tolerance as a fraction of font size; 0.5 keeps mixed-size text on one line. */ - private static final float BASELINE_TOLERANCE_FACTOR = 0.5f; - - /** Absolute minimum tolerance so tiny font sizes don't collapse multi-line content. */ - private static final float MIN_BASELINE_TOLERANCE = 2f; - - /** - * Column-gap threshold as a fraction of page width; 0.10 clears tab stops but stays below - * two-column gutters. - */ - static final float COLUMN_GAP_RATIO = 0.10f; - - /** Floor for the column-gap threshold so narrow pages don't over-split lines. */ - static final float COLUMN_GAP_MIN_PT = 40f; - - public List build(List fragments, int pageNumber) { - if (fragments.isEmpty()) return List.of(); - - float effectiveWidth = inferEffectiveWidth(fragments); - float columnGapThreshold = Math.max(effectiveWidth * COLUMN_GAP_RATIO, COLUMN_GAP_MIN_PT); - log.debug( - "LineBuilder page {}: effectiveWidth={:.1f}pt, columnGapThreshold={:.1f}pt", - pageNumber, - effectiveWidth, - columnGapThreshold); - - // Sort top-to-bottom first, then left-to-right within the same baseline band. - List sorted = - fragments.stream() - .sorted( - Comparator.comparingDouble(TextFragment::baseline) - .thenComparingDouble(f -> f.bounds().x())) - .toList(); - - List> groups = groupByBaseline(sorted, columnGapThreshold); - - List lines = new ArrayList<>(groups.size()); - for (int i = 0; i < groups.size(); i++) { - List group = - groups.get(i).stream() - .sorted(Comparator.comparingDouble(f -> f.bounds().x())) - .toList(); - - Bounds lineBounds = - group.stream() - .map(TextFragment::bounds) - .reduce(Bounds::merge) - .orElse(new Bounds(0, 0, 0, 0)); - - lines.add(new RawLine("ln-p" + pageNumber + "-" + i, group, lineBounds, pageNumber)); - } - return lines; - } - - private List> groupByBaseline( - List sorted, float columnGapThreshold) { - List> groups = new ArrayList<>(); - List current = new ArrayList<>(); - float currentBaseline = Float.NaN; - - for (TextFragment fragment : sorted) { - if (current.isEmpty()) { - current.add(fragment); - currentBaseline = fragment.baseline(); - continue; - } - - float maxFontSize = - Math.max( - fragment.fontSize(), - (float) - current.stream() - .mapToDouble(TextFragment::fontSize) - .max() - .orElse(0)); - float tolerance = - Math.max(maxFontSize * BASELINE_TOLERANCE_FACTOR, MIN_BASELINE_TOLERANCE); - - boolean sameBaseline = Math.abs(fragment.baseline() - currentBaseline) <= tolerance; - boolean columnGap = sameBaseline && hasColumnGap(fragment, current, columnGapThreshold); - - if (sameBaseline && !columnGap) { - current.add(fragment); - // Anchor to the weighted mean baseline so long lines stay stable. - currentBaseline = - (currentBaseline * (current.size() - 1) + fragment.baseline()) - / current.size(); - } else { - groups.add(current); - current = new ArrayList<>(); - current.add(fragment); - currentBaseline = fragment.baseline(); - } - } - - if (!current.isEmpty()) groups.add(current); - return groups; - } - - /** - * True when the gap from the rightmost fragment in {@code group} to {@code next} exceeds {@code - * threshold}. - */ - private static boolean hasColumnGap( - TextFragment next, List group, float threshold) { - float lastRight = group.get(group.size() - 1).bounds().right(); - return next.bounds().x() - lastRight > threshold; - } - - /** Infers effective page width from the rightmost fragment right-edge plus a 10 % margin. */ - private static float inferEffectiveWidth(List fragments) { - double maxRight = - fragments.stream().mapToDouble(f -> f.bounds().right()).max().orElse(500.0); - return (float) maxRight * 1.10f; - } -} diff --git a/app/common/src/main/java/stirling/software/SPDF/pdf/parser/PdfIngester.java b/app/common/src/main/java/stirling/software/SPDF/pdf/parser/PdfIngester.java deleted file mode 100644 index a7dc9c282b..0000000000 --- a/app/common/src/main/java/stirling/software/SPDF/pdf/parser/PdfIngester.java +++ /dev/null @@ -1,79 +0,0 @@ -package stirling.software.SPDF.pdf.parser; - -import static stirling.software.SPDF.pdf.parser.PdfModels.*; - -import java.io.IOException; -import java.util.ArrayList; -import java.util.List; - -import org.apache.pdfbox.pdmodel.PDDocument; -import org.apache.pdfbox.pdmodel.PDPage; -import org.apache.pdfbox.pdmodel.common.PDRectangle; -import org.springframework.stereotype.Service; - -import lombok.RequiredArgsConstructor; -import lombok.extern.slf4j.Slf4j; - -/** - * Runs the per-page ingestion pipeline: {@link WordExtractingStripper} → {@link LineBuilder} → - * {@link TableParser}, producing a {@link PdfModels.ParsedPage} per page. The caller owns the - * {@link PDDocument} lifecycle. - */ -@Service -@RequiredArgsConstructor -@Slf4j -public class PdfIngester { - - private final LineBuilder lineBuilder; - private final TableParser tableParser; - - public List parse(PDDocument document) throws IOException { - return parse(document, document.getNumberOfPages()); - } - - public List parse(PDDocument document, int maxPages) throws IOException { - int pageCount = Math.min(document.getNumberOfPages(), maxPages); - List pages = new ArrayList<>(pageCount); - long fragmentsMs = 0; - long tablesMs = 0; - long t0 = System.currentTimeMillis(); - - for (int p = 1; p <= pageCount; p++) { - long ft = System.currentTimeMillis(); - List fragments = extractFragments(document, p); - fragmentsMs += System.currentTimeMillis() - ft; - - PDPage page = document.getPage(p - 1); - PDRectangle mediaBox = page.getMediaBox(); - List lines = lineBuilder.build(fragments, p); - RawPage rawPage = new RawPage(p, mediaBox.getWidth(), mediaBox.getHeight(), lines); - - long tt = System.currentTimeMillis(); - List tables = tableParser.parse(document, rawPage); - tablesMs += System.currentTimeMillis() - tt; - - log.debug( - "Page {}: {} fragments → {} lines, {} table(s)", - p, - fragments.size(), - lines.size(), - tables.size()); - pages.add(new ParsedPage(p, mediaBox.getWidth(), mediaBox.getHeight(), tables, lines)); - } - - log.info( - "[timing] parse pages={} total={}ms fragments={}ms tables={}ms", - pageCount, - System.currentTimeMillis() - t0, - fragmentsMs, - tablesMs); - return pages; - } - - private List extractFragments(PDDocument document, int pageNumber) - throws IOException { - WordExtractingStripper stripper = new WordExtractingStripper(pageNumber); - stripper.getText(document); - return stripper.getFragments(); - } -} diff --git a/app/common/src/main/java/stirling/software/SPDF/pdf/parser/WordExtractingStripper.java b/app/common/src/main/java/stirling/software/SPDF/pdf/parser/WordExtractingStripper.java deleted file mode 100644 index 52ab9d9a18..0000000000 --- a/app/common/src/main/java/stirling/software/SPDF/pdf/parser/WordExtractingStripper.java +++ /dev/null @@ -1,113 +0,0 @@ -package stirling.software.SPDF.pdf.parser; - -import static stirling.software.SPDF.pdf.parser.PdfModels.*; - -import java.io.IOException; -import java.util.ArrayList; -import java.util.Collections; -import java.util.List; - -import org.apache.pdfbox.pdmodel.PDPage; -import org.apache.pdfbox.pdmodel.font.PDFont; -import org.apache.pdfbox.text.PDFTextStripper; -import org.apache.pdfbox.text.TextPosition; - -/** - * Extends {@link PDFTextStripper} to capture per-fragment geometry and font metadata. - * - *

Overrides {@link #writeString} to split each content-stream string into word-level {@link - * TextFragment}s with bounding boxes, baseline, font name, and bold flag. Coordinates are in - * PDFTextStripper space: (0,0) top-left, Y increases downward, {@code getY()} is the baseline. - */ -class WordExtractingStripper extends PDFTextStripper { - - private final int targetPage; - private final List fragments = new ArrayList<>(); - private int fragmentIndex = 0; - - WordExtractingStripper(int pageNumber) throws IOException { - this.targetPage = pageNumber; - setStartPage(pageNumber); - setEndPage(pageNumber); - setSortByPosition(true); - } - - @Override - protected void startPage(PDPage page) throws IOException { - super.startPage(page); - fragments.clear(); - fragmentIndex = 0; - } - - @Override - protected void writeString(String text, List textPositions) throws IOException { - if (text == null || text.isBlank()) return; - - // Fast path: no whitespace → emit one fragment (most financial PDFs have each - // number as its own string operation, so this is the common case). - if (text.indexOf(' ') < 0) { - emitFragment(text, textPositions); - return; - } - - // Per-word splitting requires 1:1 text-char to TextPosition correspondence. - // Fall back to one fragment when sizes differ (ligatures, encoding edge cases). - if (textPositions.size() != text.length()) { - emitFragment(text, textPositions); - return; - } - - // Emit one TextFragment per whitespace-delimited word with accurate per-word bounds. - int start = 0; - for (int i = 0; i <= text.length(); i++) { - if (i == text.length() || text.charAt(i) == ' ') { - if (start < i) { - emitFragment(text.substring(start, i), textPositions.subList(start, i)); - } - start = i + 1; - } - } - } - - private void emitFragment(String text, List positions) { - if (positions.isEmpty()) return; - - float minX = Float.MAX_VALUE; - float minY = Float.MAX_VALUE; - float maxRight = -Float.MAX_VALUE; - float maxBaseline = -Float.MAX_VALUE; - TextPosition first = null; - - for (TextPosition tp : positions) { - if (tp == null) continue; - if (first == null) first = tp; - - float x = tp.getX(); - // getY() is the baseline; top of character = getY() - getHeight(). - float top = tp.getY() - tp.getHeight(); - float right = x + tp.getWidth(); - float baseline = tp.getY(); - - minX = Math.min(minX, x); - minY = Math.min(minY, top); - maxRight = Math.max(maxRight, right); - maxBaseline = Math.max(maxBaseline, baseline); - } - - if (first == null) return; - - PDFont font = first.getFont(); - String fontName = font != null ? font.getName() : ""; - boolean bold = fontName != null && fontName.toLowerCase().contains("bold"); - // getHeight() gives the rendered glyph height, which is the most reliable visual size. - float fontSize = first.getHeight(); - - Bounds bounds = new Bounds(minX, minY, maxRight - minX, maxBaseline - minY); - String id = "tf-p" + targetPage + "-" + fragmentIndex++; - fragments.add(new TextFragment(id, text, bounds, maxBaseline, fontSize, fontName, bold)); - } - - List getFragments() { - return Collections.unmodifiableList(fragments); - } -} diff --git a/app/common/src/main/java/stirling/software/common/cluster/FileStore.java b/app/common/src/main/java/stirling/software/common/cluster/FileStore.java index f576e1e974..33e48eb40a 100644 --- a/app/common/src/main/java/stirling/software/common/cluster/FileStore.java +++ b/app/common/src/main/java/stirling/software/common/cluster/FileStore.java @@ -11,23 +11,39 @@ public interface FileStore { /** Stored file record. */ record Stored(String fileId, long size) {} - /** Store the given stream and return a generated file id and total bytes written. */ - Stored store(InputStream in, String originalName) throws IOException; + /** + * Store the given stream and return a generated file id and total bytes written. {@code owner} + * may be null to indicate the file has no associated user (anonymous / desktop / async job with + * no propagated security context); a non-null value is persisted alongside the data so {@link + * #getOwner(String)} can return it later for authorization checks. + */ + Stored store(InputStream in, String originalName, String owner) throws IOException; + + /** Store with no owner. Equivalent to {@link #store(InputStream, String, String)} with null. */ + default Stored store(InputStream in, String originalName) throws IOException { + return store(in, originalName, null); + } /** * Store the file at {@code source} and return a generated file id and total bytes written. * *

Default implementation opens {@code source} as a stream and delegates to {@link - * #store(InputStream, String)}. Local-disk implementations should override to use a direct - * file-to-file copy ({@code Files.copy(source, dest)} can use {@code sendfile(2)} on Linux), - * which avoids the two-memory-copy hit of streaming a disk-backed upload through the JVM heap. + * #store(InputStream, String, String)}. Local-disk implementations should override to use a + * direct file-to-file copy ({@code Files.copy(source, dest)} can use {@code sendfile(2)} on + * Linux), which avoids the two-memory-copy hit of streaming a disk-backed upload through the + * JVM heap. */ - default Stored store(Path source, String originalName) throws IOException { + default Stored store(Path source, String originalName, String owner) throws IOException { try (InputStream in = Files.newInputStream(source)) { - return store(in, originalName); + return store(in, originalName, owner); } } + /** Store with no owner. Equivalent to {@link #store(Path, String, String)} with null. */ + default Stored store(Path source, String originalName) throws IOException { + return store(source, originalName, null); + } + /** Open the stored file for streaming reads. Caller closes. */ InputStream retrieve(String fileId) throws IOException; @@ -42,4 +58,12 @@ public interface FileStore { /** Whether the file id exists in the store. */ boolean exists(String fileId); + + /** + * Returns the owner identifier recorded at store time, or {@code null} if the file does not + * exist or was stored without an owner. Implementations must not throw when the file is missing + * or when the owner record is absent; they should return null so callers can treat "no owner" + * as a non-authoritative case. + */ + String getOwner(String fileId) throws IOException; } diff --git a/app/common/src/main/java/stirling/software/common/cluster/inprocess/LocalDiskFileStore.java b/app/common/src/main/java/stirling/software/common/cluster/inprocess/LocalDiskFileStore.java index 6d8d0d91fa..67ba48b4af 100644 --- a/app/common/src/main/java/stirling/software/common/cluster/inprocess/LocalDiskFileStore.java +++ b/app/common/src/main/java/stirling/software/common/cluster/inprocess/LocalDiskFileStore.java @@ -3,9 +3,12 @@ package stirling.software.common.cluster.inprocess; import java.io.BufferedInputStream; import java.io.IOException; import java.io.InputStream; +import java.nio.charset.StandardCharsets; import java.nio.file.Files; import java.nio.file.Path; import java.util.UUID; +import java.util.concurrent.locks.ReentrantLock; +import java.util.regex.Pattern; import lombok.extern.slf4j.Slf4j; @@ -15,33 +18,47 @@ import stirling.software.common.cluster.FileStore; @Slf4j public class LocalDiskFileStore implements FileStore { + private static final String OWNER_SUFFIX = ".owner"; + + // File ids are generated as random UUIDs; reject anything else so a tainted id can never reach + // Files.* APIs (defence in depth on top of the resolve() prefix check, and silences CodeQL's + // path-injection finding on the resolveOwner sidecar lookup). + private static final Pattern UUID_PATTERN = + Pattern.compile( + "^[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}$"); + private final String baseDirPath; + // Fixed-size lock stripes so concurrent store/delete on the same (or colliding) fileId + // serialise the data-file + owner-sidecar pair as one critical section. Striped (not + // per-id) so the map never has to be cleaned up; collisions across unrelated ids are + // harmless contention. + private static final int LOCK_STRIPES = 64; + private final ReentrantLock[] stripes = new ReentrantLock[LOCK_STRIPES]; public LocalDiskFileStore(String baseDirPath) { this.baseDirPath = baseDirPath; + for (int i = 0; i < LOCK_STRIPES; i++) { + stripes[i] = new ReentrantLock(); + } } @Override - public Stored store(InputStream in, String originalName) throws IOException { + public Stored store(InputStream in, String originalName, String owner) throws IOException { String fileId = UUID.randomUUID().toString(); Path filePath = resolve(fileId); Files.createDirectories(filePath.getParent()); + ReentrantLock lock = acquire(fileId); boolean success = false; try { long size = Files.copy(in, filePath); + writeOwner(fileId, owner); success = true; return new Stored(fileId, size); } finally { if (!success) { - try { - Files.deleteIfExists(filePath); - } catch (IOException cleanupEx) { - log.warn( - "Failed to clean up partial file {} after store failure", - filePath, - cleanupEx); - } + cleanupAfterFailedStore(fileId, filePath); } + release(fileId, lock); } } @@ -52,27 +69,44 @@ public class LocalDiskFileStore implements FileStore { * the source size before copying so the post-copy stat is unnecessary. */ @Override - public Stored store(Path source, String originalName) throws IOException { + public Stored store(Path source, String originalName, String owner) throws IOException { String fileId = UUID.randomUUID().toString(); Path filePath = resolve(fileId); Files.createDirectories(filePath.getParent()); long size = Files.size(source); + ReentrantLock lock = acquire(fileId); boolean success = false; try { Files.copy(source, filePath); + writeOwner(fileId, owner); success = true; return new Stored(fileId, size); } finally { if (!success) { - try { - Files.deleteIfExists(filePath); - } catch (IOException cleanupEx) { - log.warn( - "Failed to clean up partial file {} after store failure", - filePath, - cleanupEx); - } + cleanupAfterFailedStore(fileId, filePath); } + release(fileId, lock); + } + } + + private void writeOwner(String fileId, String owner) throws IOException { + if (owner == null || owner.isBlank()) { + return; + } + Path ownerPath = resolveOwner(fileId); + Files.write(ownerPath, owner.getBytes(StandardCharsets.UTF_8)); + } + + private void cleanupAfterFailedStore(String fileId, Path filePath) { + try { + Files.deleteIfExists(filePath); + } catch (IOException cleanupEx) { + log.warn("Failed to clean up partial file {} after store failure", filePath, cleanupEx); + } + try { + Files.deleteIfExists(resolveOwner(fileId)); + } catch (IOException cleanupEx) { + log.warn("Failed to clean up owner sidecar for {} after store failure", fileId); } } @@ -101,11 +135,26 @@ public class LocalDiskFileStore implements FileStore { @Override public boolean delete(String fileId) { + ReentrantLock lock = acquire(fileId); try { - return Files.deleteIfExists(resolve(fileId)); - } catch (IOException e) { - log.error("Error deleting file with ID: {}", fileId, e); - return false; + // Data first, owner second: a concurrent retrieve that observes the transient + // (data-gone, owner-still-present) window simply fails with IOException; the inverse + // order would briefly look like an unowned file and could grant cross-user access. + boolean removed; + try { + removed = Files.deleteIfExists(resolve(fileId)); + } catch (IOException e) { + log.error("Error deleting file with ID: {}", fileId, e); + return false; + } + try { + Files.deleteIfExists(resolveOwner(fileId)); + } catch (IOException e) { + log.warn("Error deleting owner sidecar for file ID: {}", fileId, e); + } + return removed; + } finally { + release(fileId, lock); } } @@ -114,8 +163,21 @@ public class LocalDiskFileStore implements FileStore { return Files.exists(resolve(fileId)); } + @Override + public String getOwner(String fileId) throws IOException { + Path ownerPath = resolveOwner(fileId); + if (!Files.exists(ownerPath)) { + return null; + } + byte[] bytes = Files.readAllBytes(ownerPath); + if (bytes.length == 0) { + return null; + } + return new String(bytes, StandardCharsets.UTF_8); + } + public Path resolve(String fileId) { - if (fileId.contains("..") || fileId.contains("/") || fileId.contains("\\")) { + if (fileId == null || !UUID_PATTERN.matcher(fileId).matches()) { throw new IllegalArgumentException("Invalid file ID"); } Path basePath = Path.of(baseDirPath).normalize().toAbsolutePath(); @@ -125,4 +187,19 @@ public class LocalDiskFileStore implements FileStore { } return resolvedPath; } + + private Path resolveOwner(String fileId) { + Path data = resolve(fileId); + return data.resolveSibling(data.getFileName().toString() + OWNER_SUFFIX); + } + + private ReentrantLock acquire(String fileId) { + ReentrantLock lock = stripes[(fileId.hashCode() & Integer.MAX_VALUE) % LOCK_STRIPES]; + lock.lock(); + return lock; + } + + private void release(String fileId, ReentrantLock lock) { + lock.unlock(); + } } diff --git a/app/common/src/main/java/stirling/software/common/configuration/AppConfig.java b/app/common/src/main/java/stirling/software/common/configuration/AppConfig.java index 6d5e2507e5..ba896b4e34 100644 --- a/app/common/src/main/java/stirling/software/common/configuration/AppConfig.java +++ b/app/common/src/main/java/stirling/software/common/configuration/AppConfig.java @@ -3,7 +3,6 @@ package stirling.software.common.configuration; import java.io.IOException; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import java.util.List; import java.util.Locale; import java.util.Properties; @@ -122,17 +121,17 @@ public class AppConfig { @Bean(name = "RunningInDocker") public boolean runningInDocker() { - return Files.exists(Paths.get("/.dockerenv")); + return Files.exists(Path.of("/.dockerenv")); } @Bean(name = "configDirMounted") public boolean isRunningInDockerWithConfig() { - Path dockerEnv = Paths.get("/.dockerenv"); + Path dockerEnv = Path.of("/.dockerenv"); // default to true if not docker if (!Files.exists(dockerEnv)) { return true; } - Path mountInfo = Paths.get("/proc/1/mountinfo"); + Path mountInfo = Path.of("/proc/1/mountinfo"); // this should always exist, if not some unknown usecase if (!Files.exists(mountInfo)) { return true; diff --git a/app/common/src/main/java/stirling/software/common/configuration/ConfigInitializer.java b/app/common/src/main/java/stirling/software/common/configuration/ConfigInitializer.java index 4b3c237ec8..6d3f366b8e 100644 --- a/app/common/src/main/java/stirling/software/common/configuration/ConfigInitializer.java +++ b/app/common/src/main/java/stirling/software/common/configuration/ConfigInitializer.java @@ -7,7 +7,6 @@ import java.net.URISyntaxException; import java.net.URL; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import java.nio.file.StandardCopyOption; import java.util.List; @@ -27,7 +26,7 @@ public class ConfigInitializer { public void ensureConfigExists() throws IOException, URISyntaxException { // 1) If settings file doesn't exist, create from template - Path destPath = Paths.get(InstallationPathConfig.getSettingsPath()); + Path destPath = Path.of(InstallationPathConfig.getSettingsPath()); boolean settingsFileExists = Files.exists(destPath); @@ -39,7 +38,7 @@ public class ConfigInitializer { if (settingsFileExists) { // move settings.yml to settings.yml.{timestamp}.bak Path backupPath = - Paths.get( + Path.of( InstallationPathConfig.getSettingsPath() + "." + System.currentTimeMillis() @@ -96,7 +95,7 @@ public class ConfigInitializer { } // 3) Ensure custom settings file exists - Path customSettingsPath = Paths.get(InstallationPathConfig.getCustomSettingsPath()); + Path customSettingsPath = Path.of(InstallationPathConfig.getCustomSettingsPath()); if (Files.notExists(customSettingsPath)) { Files.createFile(customSettingsPath); log.info("Created custom_settings file: {}", customSettingsPath); diff --git a/app/common/src/main/java/stirling/software/common/configuration/RuntimePathConfig.java b/app/common/src/main/java/stirling/software/common/configuration/RuntimePathConfig.java index b1ecc1020e..82921c0ba2 100644 --- a/app/common/src/main/java/stirling/software/common/configuration/RuntimePathConfig.java +++ b/app/common/src/main/java/stirling/software/common/configuration/RuntimePathConfig.java @@ -3,7 +3,6 @@ package stirling.software.common.configuration; import java.nio.file.Files; import java.nio.file.InvalidPathException; import java.nio.file.Path; -import java.nio.file.Paths; import java.util.ArrayList; import java.util.Collections; import java.util.LinkedHashSet; @@ -201,7 +200,7 @@ public class RuntimePathConfig { try { // Normalize to absolute path - Path path = Paths.get(pathStr.trim()).toAbsolutePath().normalize(); + Path path = Path.of(pathStr.trim()).toAbsolutePath().normalize(); String normalizedPath = path.toString(); // Check for duplicates @@ -224,9 +223,9 @@ public class RuntimePathConfig { private void detectOverlappingPaths(List paths) { for (int i = 0; i < paths.size(); i++) { - Path path1 = Paths.get(paths.get(i)); + Path path1 = Path.of(paths.get(i)); for (int j = i + 1; j < paths.size(); j++) { - Path path2 = Paths.get(paths.get(j)); + Path path2 = Path.of(paths.get(j)); // Check if one path is a parent of the other if (path1.startsWith(path2)) { @@ -246,10 +245,10 @@ public class RuntimePathConfig { private void validatePipelinePaths() { try { - Path finishedPath = Paths.get(pipelineFinishedFoldersPath).toAbsolutePath().normalize(); + Path finishedPath = Path.of(pipelineFinishedFoldersPath).toAbsolutePath().normalize(); for (String watchedPathStr : pipelineWatchedFoldersPaths) { - Path watchedPath = Paths.get(watchedPathStr).toAbsolutePath().normalize(); + Path watchedPath = Path.of(watchedPathStr).toAbsolutePath().normalize(); // Check if watched folder is same as finished folder if (watchedPath.equals(finishedPath)) { diff --git a/app/common/src/main/java/stirling/software/common/model/ApplicationProperties.java b/app/common/src/main/java/stirling/software/common/model/ApplicationProperties.java index bb2eb81678..b955a2fe32 100644 --- a/app/common/src/main/java/stirling/software/common/model/ApplicationProperties.java +++ b/app/common/src/main/java/stirling/software/common/model/ApplicationProperties.java @@ -77,8 +77,10 @@ public class ApplicationProperties { private ProcessExecutor processExecutor = new ProcessExecutor(); private PdfEditor pdfEditor = new PdfEditor(); private AiEngine aiEngine = new AiEngine(); + private Mcp mcp = new Mcp(); private InternalApi internalApi = new InternalApi(); private Cluster cluster = new Cluster(); + private Policies policies = new Policies(); @Bean public PropertySource dynamicYamlPropertySource(ConfigurableEnvironment environment) @@ -202,6 +204,45 @@ public class ApplicationProperties { } } + @Data + public static class Policies { + /** + * Absolute directories that policy folder input sources and output sinks may read from or + * write to. Empty (the default) disables folder access entirely, so a policy can never be + * pointed at an arbitrary server path. Stirling's own config directory is always + * off-limits, and folder access is always disabled in SaaS mode regardless of this list. + */ + private List allowedFolderRoots = new java.util.ArrayList<>(); + + /** How often (seconds) the schedule trigger checks for policies whose schedule is due. */ + private long scheduleSweepSeconds = 60; + + /** + * How often (seconds) the folder-watch trigger reconciles its watch registrations and + * re-runs every folder-watch policy as a safety net for filesystem events that were missed + * (NFS, bind mounts, inotify-queue overflow). + */ + private long watchReconcileSeconds = 300; + + /** + * How long (milliseconds) the folder-watch trigger keeps draining filesystem events after + * the first, so a burst from a single file copy coalesces into one run instead of many. + */ + private long watchQuietPeriodMs = 500; + + /** + * SSE emitter timeout (milliseconds) for streamed runs; generous for long multi-step runs. + */ + private long streamTimeoutMs = 1800000; + + /** + * How long (minutes) a finished run's in-memory state is retained before eviction, + * mirroring the job-result expiry so rich run state does not outlive the process. Active + * and paused runs are kept regardless of age. + */ + private int runExpiryMinutes = 30; + } + @Data public static class PdfEditor { private Cache cache = new Cache(); @@ -256,6 +297,103 @@ public class ApplicationProperties { private int longRunningTimeoutSeconds = 600; } + /** + * Model Context Protocol (MCP) server configuration. All keys live under the top-level {@code + * mcp.*} prefix. {@link #enabled} defaults to {@code false}: when off, no MCP beans are wired, + * no /mcp endpoint exists, and no protected-resource metadata is published. + */ + @Data + public static class Mcp { + + /** Master switch. When {@code false} (default), no MCP beans are wired. */ + private boolean enabled = false; + + /** + * When {@code true} (default), invocations require an OAuth scope: {@code mcp.tools.read} + * for read-style operations and {@code mcp.tools.write} for write/destructive ones. When + * {@code false}, scope checks are skipped (use only if your IdP issues a single coarse + * scope). + */ + private boolean scopesEnabled = true; + + /** How often to refresh the AI capabilities manifest from the engine. */ + private int engineCapabilityRefreshMinutes = 5; + + /** + * Tool allow-list (operation ids, e.g. {@code compress-pdf}). When non-empty, ONLY these + * operations are exposed over MCP; everything else is hidden, undescribable, and + * uninvocable - on top of the global endpoint enable/disable config. Empty = allow all. + */ + private List allowedOperations = new ArrayList<>(); + + /** + * Tool deny-list (operation ids). Any operation listed here is removed from MCP even if it + * would otherwise be allowed. Applied after {@link #allowedOperations}. + */ + private List blockedOperations = new ArrayList<>(); + + /** Max MCP request body size in bytes; inline file uploads ride in the JSON-RPC body. */ + private long maxRequestBytes = 10L * 1024 * 1024; + + /** Results up to this size return inline as base64; larger ones return a fileId only. */ + private long maxInlineResponseBytes = 10L * 1024 * 1024; + + private Auth auth = new Auth(); + + @Data + public static class Auth { + /** + * Authentication mode for the MCP endpoint. {@code oauth} (default) runs a full OAuth2 + * resource server (JWT, RFC 8707 audience, RFC 9728 metadata). {@code apikey} accepts a + * Stirling per-user API key via the {@code X-API-KEY} header (or {@code Authorization: + * Bearer }) and binds the request to that user - the low-friction self-host path, + * no external IdP required. + */ + private String mode = "oauth"; + + /** OAuth2 issuer URI, e.g. {@code http://localhost:9000}. Required when MCP is on. */ + private String issuerUri = ""; + + /** + * JWKS URI. When blank, derived from the issuer's {@code + * /.well-known/openid-configuration} document. + */ + private String jwksUri = ""; + + /** + * RFC 8707 resource identifier of THIS MCP server, e.g. {@code + * http://localhost:8080/mcp}. Tokens that do not list this id in their {@code aud} + * claim are rejected with HTTP 401. + */ + private String resourceId = ""; + + /** + * Additional JWT audiences accepted at the MCP endpoint, on top of {@link #resourceId}. + * Empty (default) keeps strict RFC 8707 binding. Some IdPs cannot mint + * resource-specific audiences - e.g. Supabase's OAuth server always issues {@code + * aud=authenticated} - so operators list the audience their IdP actually emits here + * (env: {@code MCP_AUTH_ACCEPTEDAUDIENCES}, comma-separated). + */ + private List acceptedAudiences = new ArrayList<>(); + + /** + * JWT claim whose value is matched against a provisioned Stirling username. Defaults to + * {@code sub}; set to {@code email} or {@code preferred_username} to match how your IdP + * maps users to Stirling accounts. + */ + private String usernameClaim = "sub"; + + /** + * When {@code true} (default), a validated token is accepted only if its {@link + * #usernameClaim} value resolves to an existing, enabled Stirling user account. Tokens + * whose subject has no Stirling account (or a disabled one) are rejected with HTTP 403. + * Set to {@code false} only if you intentionally want any IdP-valid token to use MCP + * without a local account. + */ + private boolean requireExistingAccount = true; + } + } + /** * Cluster backplane configuration. All keys live under the top-level {@code cluster.*} prefix * (e.g. env var {@code CLUSTER_ENABLED}). The master switch is {@link #enabled} and defaults to @@ -867,6 +1005,10 @@ public class ApplicationProperties { @Data public static class Signing { private boolean enabled = false; + + // Signing user-picker scope: 'org' (default) = whole instance, anything else = + // caller's team only (fail-closed). The saas profile pins 'team'. + private String userListScope = "org"; } } diff --git a/app/common/src/main/java/stirling/software/common/model/FileInfo.java b/app/common/src/main/java/stirling/software/common/model/FileInfo.java index e894202969..0d2777c43e 100644 --- a/app/common/src/main/java/stirling/software/common/model/FileInfo.java +++ b/app/common/src/main/java/stirling/software/common/model/FileInfo.java @@ -1,7 +1,6 @@ package stirling.software.common.model; import java.nio.file.Path; -import java.nio.file.Paths; import java.time.LocalDateTime; import java.time.format.DateTimeFormatter; import java.util.Locale; @@ -24,7 +23,7 @@ public class FileInfo { // Converts the file path string to a Path object. public Path getFilePathAsPath() { - return Paths.get(filePath); + return Path.of(filePath); } // Formats the file size into a human-readable string. diff --git a/app/common/src/main/java/stirling/software/common/pdf/HeadingDetector.java b/app/common/src/main/java/stirling/software/common/pdf/HeadingDetector.java new file mode 100644 index 0000000000..0937cef646 --- /dev/null +++ b/app/common/src/main/java/stirling/software/common/pdf/HeadingDetector.java @@ -0,0 +1,191 @@ +package stirling.software.common.pdf; + +import java.util.ArrayList; +import java.util.Collections; +import java.util.HashMap; +import java.util.List; +import java.util.Map; + +import stirling.software.jpdfium.text.PageText; +import stirling.software.jpdfium.text.TextChar; +import stirling.software.jpdfium.text.TextLine; +import stirling.software.jpdfium.text.TextWord; + +final class HeadingDetector { + + private HeadingDetector() {} + + /** A heading is at most this many words; longer lines are treated as body text. */ + private static final int MAX_HEADING_WORDS = 12; + + /** + * Returns the Markdown heading prefix for a line. The decision combines several signals, never + * text matching, so a plain line that merely shares text with a heading is never promoted: + * + *

    + *
  • Size — dominant glyph font size vs. the document body median (primary signal). + * Some PDFs encode visual size in the text matrix, so every glyph reports ~1.0; for those + * the line height is used as the proxy instead. + *
  • Brevity — headings are short labels; a line over {@value #MAX_HEADING_WORDS} + * words is body text regardless of size. + *
  • Not a sentence — a line ending in {@code . ! ?} reads as prose, not a heading. + *
+ * + *

Boldness is deliberately not a heading signal — a bold-but-not-larger line is + * emphasis, not a heading (see {@link #isBoldLabel}); promoting it to {@code #}/{@code ##} is + * the main source of false-positive headings. + * + *

    + *
  • size > baseline * 1.4 → {@code "# "} + *
  • size > baseline * 1.2 → {@code "## "} + *
  • otherwise → {@code ""} + *
+ */ + static String headingPrefix(TextLine line, float medianBodySize, float medianBodyHeight) { + String text = line.text().strip(); + if (text.isEmpty() || wordCount(text) > MAX_HEADING_WORDS || endsLikeSentence(text)) { + return ""; + } + + float dominant = dominantFontSize(line); + float value; + float baseline; + if (dominant > 2f && medianBodySize > 2f) { + value = dominant; + baseline = medianBodySize; + } else { + value = line.height(); + baseline = medianBodyHeight; + } + if (baseline <= 0f) { + return ""; + } + + float ratio = value / baseline; + if (ratio > 1.4f) { + return "# "; + } + if (ratio > 1.2f) { + return "## "; + } + return ""; + } + + /** + * True when a line should be emphasised as bold (rendered {@code **like this**}) rather than + * promoted to a heading: it is bold, short, and not a full sentence. Used for bold labels that + * are not large enough to be headings. + */ + static boolean isBoldLabel(TextLine line) { + String text = line.text().strip(); + if (text.isEmpty() || wordCount(text) > MAX_HEADING_WORDS || endsLikeSentence(text)) { + return false; + } + return isBold(line); + } + + private static int wordCount(String text) { + return text.split("\\s+").length; + } + + private static boolean endsLikeSentence(String text) { + char last = text.charAt(text.length() - 1); + return last == '.' || last == '!' || last == '?'; + } + + /** True when the line's dominant font is bold, inferred from PostScript font names. */ + private static boolean isBold(TextLine line) { + Map counts = new HashMap<>(); + for (TextWord word : line.words()) { + for (TextChar ch : word.chars()) { + if (ch.isWhitespace() || ch.isNewline()) { + continue; + } + String name = ch.fontName(); + if (name != null && !name.isBlank()) { + counts.merge(name, 1, Integer::sum); + } + } + } + String dominantFont = ""; + int max = -1; + for (Map.Entry e : counts.entrySet()) { + if (e.getValue() > max) { + max = e.getValue(); + dominantFont = e.getKey(); + } + } + String lower = dominantFont.toLowerCase(java.util.Locale.ROOT); + return lower.contains("bold") + || lower.contains("black") + || lower.contains("heavy") + || lower.contains("semibold"); + } + + /** Computes the median glyph font size across all pages. */ + static float medianFontSize(List allPages) { + List sizes = new ArrayList<>(); + for (PageText page : allPages) { + for (TextChar ch : page.chars()) { + if (!ch.isWhitespace() && !ch.isNewline() && ch.fontSize() > 0f) { + sizes.add(ch.fontSize()); + } + } + } + return median(sizes, 12f); + } + + /** Computes the median TextLine height across all pages. Used when font size is degenerate. */ + static float medianLineHeight(List allPages) { + List heights = new ArrayList<>(); + for (PageText page : allPages) { + for (TextLine line : page.lines()) { + if (line.height() > 0f && !line.text().isBlank()) { + heights.add(line.height()); + } + } + } + return median(heights, 12f); + } + + private static float median(List values, float fallback) { + if (values.isEmpty()) { + return fallback; + } + Collections.sort(values); + int mid = values.size() / 2; + if (values.size() % 2 == 0) { + return (values.get(mid - 1) + values.get(mid)) / 2f; + } + return values.get(mid); + } + + /** + * Returns the font size that appears most often (by character count) in the given line. Ties + * are broken in favour of the larger size. + */ + private static float dominantFontSize(TextLine line) { + Map counts = new HashMap<>(); + for (TextWord word : line.words()) { + for (TextChar ch : word.chars()) { + if (!ch.isWhitespace() && !ch.isNewline() && ch.fontSize() > 0f) { + counts.merge(ch.fontSize(), 1, Integer::sum); + } + } + } + if (counts.isEmpty()) { + return 0f; + } + float dominant = 0f; + int maxCount = -1; + for (Map.Entry entry : counts.entrySet()) { + int count = entry.getValue(); + float size = entry.getKey(); + if (count > maxCount || (count == maxCount && size > dominant)) { + maxCount = count; + dominant = size; + } + } + return dominant; + } +} diff --git a/app/common/src/main/java/stirling/software/common/pdf/PdfMarkdownConverter.java b/app/common/src/main/java/stirling/software/common/pdf/PdfMarkdownConverter.java new file mode 100644 index 0000000000..c19468b5ed --- /dev/null +++ b/app/common/src/main/java/stirling/software/common/pdf/PdfMarkdownConverter.java @@ -0,0 +1,1043 @@ +package stirling.software.common.pdf; + +import java.io.IOException; +import java.util.ArrayList; +import java.util.Comparator; +import java.util.HashSet; +import java.util.List; +import java.util.Set; +import java.util.regex.Pattern; +import java.util.stream.Collectors; + +import stirling.software.jpdfium.PdfDocument; +import stirling.software.jpdfium.PdfPage; +import stirling.software.jpdfium.doc.ExtractedImage; +import stirling.software.jpdfium.doc.PdfImageExtractor; +import stirling.software.jpdfium.model.Rect; +import stirling.software.jpdfium.text.PageText; +import stirling.software.jpdfium.text.PdfTableExtractor; +import stirling.software.jpdfium.text.PdfTextExtractor; +import stirling.software.jpdfium.text.Table; +import stirling.software.jpdfium.text.TextLine; +import stirling.software.jpdfium.text.TextWord; + +/** + * Converts a PDF to Markdown using a TextLine-driven body pipeline. + * + *

Body text is rebuilt from {@link PdfTextExtractor} {@link TextLine}s. TextLines group words + * faithfully and keep paragraph order, so the only pre-processing needed is stitching narrow + * standalone glyph fragments (apostrophes, quotes, asterisks, superscript footnote markers, + * bullets) back into the line they belong to. Column layout and tables are derived from line/word + * geometry directly. + */ +public class PdfMarkdownConverter { + + private static final Pattern SOFT_HYPHEN = Pattern.compile("(\\w+)-\\n([a-z])"); + + /** Width below which a TextLine is treated as a stray glyph fragment to be stitched. */ + private static final float GLYPH_WIDTH = 7.5f; + + public String convert(PdfDocument doc) throws IOException { + List allPageText = PdfTextExtractor.extractAll(doc); + float medianSize = HeadingDetector.medianFontSize(allPageText); + float medianHeight = HeadingDetector.medianLineHeight(allPageText); + + int pageCount = doc.pageCount(); + // Elements are either rendered text (String) or a structured TableBlock. Tables stay + // structured until after the page loop so a table split across a page break can be stitched + // back together before rendering. + List output = new ArrayList<>(); + // Header text of a table that ended the previous page, used to spot a continuation whose + // header repeats at the top of the current page. Null when the previous page did not end in + // a table. + String prevPageTrailingTableHeader = null; + + for (int pageIndex = 0; pageIndex < pageCount; pageIndex++) { + List rawLines = + pageIndex < allPageText.size() ? allPageText.get(pageIndex).lines() : List.of(); + + // Stitch stray glyph fragments (apostrophes, asterisks, superscripts, bullets) into + // their host lines so paragraph assembly sees faithful, complete lines. + List lines = stitchGlyphs(rawLines); + if (lines.isEmpty()) { + emitImages(doc, pageIndex, output); + prevPageTrailingTableHeader = null; + continue; + } + + // Sort top-to-bottom (PDF y=0 is the bottom of the page). + lines.sort(Comparator.comparingDouble((Line l) -> l.y).reversed()); + + // Multi-column guard: only genuine two-column prose should be split. A table's column + // gutters must NOT be mistaken for a page-layout gutter, so this looks at whether row + // lines span the gutter (table) or stay within one side (two-column prose). + // A table that ran to the bottom of the previous page and repeats its header at the top + // of this page is a continuation, not a new two-column layout. Detecting the repeated + // header keeps this page out of the two-column path so the continuation is rebuilt as a + // table and stitched back onto the previous block. + final String continuationHeader = prevPageTrailingTableHeader; + boolean tableContinuation = + continuationHeader != null + && lines.stream() + .anyMatch( + l -> normaliseSpace(l.text).equals(continuationHeader)); + + boolean twoColumn = !tableContinuation && detectsTwoColumns(lines); + + // Tables are detected from text/word geometry (the word-grid detector), which handles + // both ruled and borderless tables and places cells by column alignment. The native + // ruled-line extractor is not used: it both mis-renders cells and double-emits rows. + Set tableRowTexts = new HashSet<>(); + List blocks = twoColumn ? List.of() : findTableBlocks(lines); + Set tableLines = new HashSet<>(); + for (TableBlock b : blocks) { + for (List row : b.rows()) { + for (Line l : row) { + tableLines.add(l); + tableRowTexts.add(repairHyphens(l.text).strip()); + } + } + } + + List pageItems = new ArrayList<>(); + if (twoColumn) { + for (List col : splitIntoColumns(lines)) { + List paras = new ArrayList<>(); + assembleParagraphs(col, medianSize, medianHeight, paras, tableRowTexts); + pageItems.addAll(paras); + } + } else { + // Interleave tables with surrounding text by vertical position. Each block sits in + // its own slot; non-table lines fall into the slot for their y (text above a block, + // between blocks, or below the last). This keeps multiple tables on one page + // separate and in reading order. + List> segments = new ArrayList<>(); + for (int s = 0; s <= blocks.size(); s++) { + segments.add(new ArrayList<>()); + } + for (Line l : lines) { + if (tableLines.contains(l)) { + continue; + } + int slot = 0; + for (TableBlock b : blocks) { + if (b.bottom() > l.y) { + slot++; + } + } + segments.get(slot).add(l); + } + for (int s = 0; s <= blocks.size(); s++) { + List paras = new ArrayList<>(); + assembleParagraphs( + segments.get(s), medianSize, medianHeight, paras, tableRowTexts); + pageItems.addAll(paras); + if (s < blocks.size()) { + pageItems.add(blocks.get(s)); + } + } + } + + emitImages(doc, pageIndex, pageItems); + + if (pageItems.isEmpty()) { + continue; + } + + mergeAcrossPageBoundary(output, pageItems); + output.addAll(pageItems); + prevPageTrailingTableHeader = trailingTableHeader(pageItems); + } + + // Stitch tables split across page breaks, then render every element to Markdown. + List stitched = stitchTables(output); + List rendered = new ArrayList<>(); + for (Object e : stitched) { + rendered.add(e instanceof TableBlock tb ? tb.render() : (String) e); + } + return String.join("\n\n", rendered); + } + + // --- Glyph stitching --------------------------------------------------- + + /** A mutable assembled line: text plus geometry used for ordering and heading detection. */ + private static final class Line { + String text; + float x; + float y; + float width; + float height; + final TextLine source; + + Line(TextLine src) { + this.source = src; + this.text = src.text(); + this.x = src.x(); + this.y = src.y(); + this.width = src.width(); + this.height = src.height(); + } + } + + /** + * Merges narrow glyph fragments (width < {@link #GLYPH_WIDTH}) into the line they belong to. + * + *
    + *
  • A glyph between a left fragment that ends near it and a right fragment that starts near + * it (both on the same baseline) is inserted inline: {@code aren} + {@code '} + {@code t} + * → {@code aren't}. + *
  • A glyph immediately right of a line's end is appended (e.g. superscript footnote marker + * after a number). + *
  • A glyph immediately left of a line's start is prepended (e.g. footnote marker before + * its text). + *
+ */ + private static List stitchGlyphs(List raw) { + List hosts = new ArrayList<>(); + List glyphs = new ArrayList<>(); + for (TextLine l : raw) { + String t = l.text().strip(); + if (t.isEmpty()) { + continue; + } + if (l.width() < GLYPH_WIDTH && t.length() <= 2) { + glyphs.add(l); + } else { + hosts.add(l); + } + } + + List lines = hosts.stream().map(Line::new).collect(Collectors.toList()); + + for (TextLine g : glyphs) { + String gt = g.text().strip(); + if (isBulletGlyph(gt)) { + attachBullet(g, gt, lines); + } else { + attachInlineGlyph(g, gt, lines); + } + } + return lines; + } + + private static boolean isBulletGlyph(String gt) { + return "•".equals(gt) || "▪".equals(gt) || "◦".equals(gt); + } + + /** + * Attaches a bullet glyph to the body line it introduces: the closest line that begins to the + * right of the bullet at roughly the same height or just below it. + */ + private static void attachBullet(TextLine g, String gt, List lines) { + Line best = null; + float bestScore = Float.MAX_VALUE; + for (Line h : lines) { + if (h.x < g.x() - 2f) { + continue; + } + float dy = g.y() - h.y; + if (dy < -4f || dy > 28f) { + continue; + } + float score = Math.abs(dy) + (h.x - g.x()) * 0.2f; + if (score < bestScore) { + bestScore = score; + best = h; + } + } + if (best != null && !best.text.startsWith("•")) { + best.text = "• " + best.text; + best.x = g.x(); + } else { + lines.add(new Line(g)); + } + } + + /** + * Stitches a narrow inline glyph (apostrophe, quote, asterisk, superscript marker) into the + * line it belongs to: inline between two same-baseline fragments, appended to the line that + * ends at it, or prepended to the line that starts at it. + */ + private static void attachInlineGlyph(TextLine g, String gt, List lines) { + Line left = null; + Line right = null; + float lb = 7f; + float rb = 7f; + for (Line h : lines) { + boolean sameBaseline = g.y() >= h.y - 4f && g.y() <= h.y + h.height + 5f; + if (!sameBaseline) { + continue; + } + float rightEdge = h.x + h.width; + float dxLeft = Math.abs(rightEdge - g.x()); + if (dxLeft < lb) { + lb = dxLeft; + left = h; + } + float dxRight = Math.abs(h.x - g.x()); + if (dxRight < rb) { + rb = dxRight; + right = h; + } + } + + if (left != null && right != null && left != right && Math.abs(left.y - right.y) < 6f) { + left.text = left.text + gt + right.text; + left.width = (right.x + right.width) - left.x; + lines.remove(right); + } else if (left != null) { + left.text = left.text + gt; + left.width = Math.max(left.width, g.x() + g.width() - left.x); + } else if (right != null) { + right.text = gt + right.text; + right.x = g.x(); + } else { + lines.add(new Line(g)); + } + } + + // --- Column detection (guard only) ------------------------------------- + + /** + * Returns true when the page is a genuine two-column layout. Uses line/word geometry: body + * blocks (ignoring narrow glyph blocks) and requires a wide horizontal gutter populated on both + * sides, so single apostrophe glyphs cannot create a false second column. + */ + private static boolean detectsTwoColumns(List lines) { + if (lines.size() < 8) { + return false; + } + float minX = Float.MAX_VALUE; + float maxX = -Float.MAX_VALUE; + for (Line l : lines) { + minX = Math.min(minX, l.x); + maxX = Math.max(maxX, l.x + l.width); + } + if (maxX - minX < 200f) { + return false; + } + + // Scan candidate gutter positions across the central band (35%-65% of width) and pick the + // one crossed by the fewest lines. Two-column prose has a gutter that only a handful of + // full-width lines (title, section headings) cross; a table's rows all span the full width, + // so every candidate gutter is crossed by most lines. + float centreLo = minX + (maxX - minX) * 0.35f; + float centreHi = minX + (maxX - minX) * 0.65f; + int bestCrossing = Integer.MAX_VALUE; + int bestLeft = 0; + int bestRight = 0; + for (float gutter = centreLo; gutter <= centreHi; gutter += 2f) { + int crossing = 0; + int leftOnly = 0; + int rightOnly = 0; + for (Line l : lines) { + float lx = l.x; + float rx = l.x + l.width; + if (lx < gutter - 5f && rx > gutter + 5f) { + crossing++; + } else if (rx <= gutter) { + leftOnly++; + } else { + rightOnly++; + } + } + if (crossing < bestCrossing) { + bestCrossing = crossing; + bestLeft = leftOnly; + bestRight = rightOnly; + } + } + + return bestLeft >= 4 && bestRight >= 4 && bestCrossing <= (int) (lines.size() * 0.25f); + } + + private static List> splitIntoColumns(List lines) { + List xs = + lines.stream() + .filter(l -> l.width >= 40f) + .map(l -> l.x) + .sorted() + .collect(Collectors.toList()); + if (xs.isEmpty()) { + return List.of(lines); + } + float minX = xs.get(0); + float maxX = xs.get(xs.size() - 1); + float splitAt = (minX + maxX) / 2f; + float biggestGap = 0; + for (int i = 1; i < xs.size(); i++) { + float gap = xs.get(i) - xs.get(i - 1); + if (gap > biggestGap) { + biggestGap = gap; + splitAt = (xs.get(i - 1) + xs.get(i)) / 2f; + } + } + List left = new ArrayList<>(); + List right = new ArrayList<>(); + for (Line l : lines) { + if (l.x < splitAt) { + left.add(l); + } else { + right.add(l); + } + } + if (left.isEmpty()) { + return List.of(right); + } + if (right.isEmpty()) { + return List.of(left); + } + return List.of(left, right); + } + + // --- Paragraph assembly ------------------------------------------------ + + private static void assembleParagraphs( + List lines, + float medianSize, + float medianHeight, + List out, + Set tableRowTexts) { + StringBuilder para = new StringBuilder(); + float prevBottomY = Float.MAX_VALUE; + float prevHeight = 0f; + + for (Line line : lines) { + String text = repairHyphens(line.text).strip(); + if (text.isEmpty()) { + continue; + } + if (tableRowTexts.contains(text)) { + continue; + } + + float blockTop = line.y + line.height; + float gap = prevBottomY - blockTop; + boolean paragraphBreak = prevHeight > 0f && gap > prevHeight * 0.8f; + + String prefix = HeadingDetector.headingPrefix(line.source, medianSize, medianHeight); + boolean isHeading = !prefix.isEmpty(); + boolean isBullet = startsWithBullet(text); + + if (isHeading) { + flushParagraph(para, out); + out.add(prefix + escapeMarkdown(text)); + } else if (isBullet) { + flushParagraph(para, out); + out.add(escapeMarkdown(text)); + } else if (HeadingDetector.isBoldLabel(line.source)) { + // Bold but not large enough to be a heading → emphasise as bold, don't promote. + flushParagraph(para, out); + out.add("**" + escapeMarkdown(text) + "**"); + } else if (paragraphBreak) { + flushParagraph(para, out); + para.append(text); + } else { + if (!para.isEmpty()) { + char fc = text.charAt(0); + boolean noSpace = fc == '\'' || fc == '’' || fc == '‘' || fc == '"'; + if (!noSpace) { + para.append(' '); + } + } + para.append(text); + } + + prevBottomY = line.y; + prevHeight = line.height; + } + flushParagraph(para, out); + } + + private static boolean startsWithBullet(String text) { + return text.startsWith("•") || text.startsWith("▪") || text.startsWith("◦"); + } + + // --- Word-grid table detection ----------------------------------------- + + /** + * A detected table. Each row is a list of source lines: usually one, but more when a cell wraps + * onto extra lines (those continuation lines are absorbed into the row they belong to). + */ + private record TableBlock(List> rows, float top, float bottom) { + String render() { + return buildTableFromRows(rows); + } + } + + /** + * Detects table blocks on a page. Anchor rows (lines with table-like column gaps) are grouped + * into vertically-contiguous runs separated by large vertical gaps, so multiple separate tables + * on one page stay separate. Non-anchor lines that fall within a run's vertical span are + * treated as wrapped-cell continuations and absorbed into the nearest anchor row above them. + */ + private static List findTableBlocks(List lines) { + List cands = + lines.stream() + .filter(l -> isTableCandidate(l.source)) + .sorted(Comparator.comparingDouble((Line l) -> l.y).reversed()) + .collect(Collectors.toList()); + if (cands.size() < 2) { + return List.of(); + } + + List gaps = new ArrayList<>(); + for (int i = 1; i < cands.size(); i++) { + gaps.add(cands.get(i - 1).y - cands.get(i).y); + } + List sorted = new ArrayList<>(gaps); + sorted.sort(Comparator.naturalOrder()); + float medianGap = sorted.get(sorted.size() / 2); + float splitThreshold = Math.max(medianGap * 2.5f, medianGap + 6f); + + List> anchorGroups = new ArrayList<>(); + List current = new ArrayList<>(); + current.add(cands.get(0)); + for (int i = 1; i < cands.size(); i++) { + float gap = cands.get(i - 1).y - cands.get(i).y; + if (gap > splitThreshold) { + anchorGroups.add(current); + current = new ArrayList<>(); + } + current.add(cands.get(i)); + } + anchorGroups.add(current); + + List nonCandidates = + lines.stream() + .filter(l -> !isTableCandidate(l.source)) + .collect(Collectors.toList()); + + List blocks = new ArrayList<>(); + for (List anchors : anchorGroups) { + if (anchors.size() < 2) { + continue; + } + float top = anchors.get(0).y; + float bottom = anchors.get(anchors.size() - 1).y; + + // Each anchor seeds a row; absorb wrapped continuation lines (non-anchors within the + // run's vertical span, with a little slack below the last row) into the anchor above. + List> rows = new ArrayList<>(); + for (Line a : anchors) { + List row = new ArrayList<>(); + row.add(a); + rows.add(row); + } + for (Line nc : nonCandidates) { + if (nc.y > top || nc.y < bottom - medianGap) { + continue; + } + int owner = 0; + float bestDelta = Float.MAX_VALUE; + for (int i = 0; i < anchors.size(); i++) { + float delta = anchors.get(i).y - nc.y; // positive when anchor is above nc + if (delta >= -1f && delta < bestDelta) { + bestDelta = delta; + owner = i; + } + } + rows.get(owner).add(nc); + } + + if (buildTableFromRows(rows).isBlank()) { + continue; + } + blocks.add(new TableBlock(rows, top, bottom)); + } + return blocks; + } + + private static String buildTableFromRows(List> rowGroups) { + // Detect columns by vertical-whitespace projection across all lines, rather than a 1-D gap + // threshold on pooled word x's. Pooled-gap detection is fragile when numbers are + // right-aligned (a 10-digit value starts well left of a 7-digit one) or when sparse cells + // sit in their own x-band. Projection asks "which x-bands are occupied across many rows", + // which is stable under those conditions. + List flat = rowGroups.stream().flatMap(List::stream).collect(Collectors.toList()); + List columns = findColumnRanges(flat); + if (columns.size() < 2 || columns.size() > 15) { + return ""; + } + + float[] centers = new float[columns.size()]; + for (int i = 0; i < columns.size(); i++) { + centers[i] = (columns.get(i)[0] + columns.get(i)[1]) / 2f; + } + + int cols = centers.length; + List rows = new ArrayList<>(); + for (List rowLines : rowGroups) { + String[] row = new String[cols]; + for (int i = 0; i < cols; i++) { + row[i] = ""; + } + // Top line first so a wrapped cell's words stay in reading order within the cell. + rowLines.sort(Comparator.comparingDouble((Line l) -> l.y).reversed()); + for (Line line : rowLines) { + for (TextWord word : line.source.words()) { + String wt = word.text().strip(); + if (wt.isEmpty()) { + continue; + } + int col = nearestColumn(word.x() + word.width() / 2f, centers); + row[col] = row[col].isEmpty() ? wt : row[col] + " " + wt; + } + } + rows.add(row); + } + + // Guard against false positives while tolerating uneven rows (sparse cells, merged/spanning + // headers). The columns already come from cross-row whitespace alignment, so a stable grid + // exists. Additionally require: at least one "anchor" row that nearly fills the grid (so + // the + // column count is real, not an artefact), and that most rows are genuinely multi-column. + int anchorWidth = Math.max(2, Math.round(cols * 0.6f)); + long anchorRows = rows.stream().filter(r -> filledCells(r) >= anchorWidth).count(); + long multiColumnRows = rows.stream().filter(r -> filledCells(r) >= 2).count(); + if (anchorRows < 1 || multiColumnRows < 2 || multiColumnRows < rows.size() * 0.5) { + return ""; + } + return renderGfm(rows, cols); + } + + /** + * Visible for testing: column detection depends only on word geometry, so tests can drive it + * from synthetic {@link TextLine}s to exercise degenerate-coordinate handling (the crash path + * an extreme text matrix can produce) without needing a binary PDF fixture. + */ + static List findColumnRangesFromLines(List rows) { + return findColumnRanges(rows.stream().map(Line::new).collect(Collectors.toList())); + } + + /** + * Finds column x-ranges by vertical-whitespace projection. Each row contributes coverage for + * the x-bands its words occupy; a column is a contiguous band covered by a sufficient fraction + * of rows, and the gaps between such bands are the gutters. + */ + private static List findColumnRanges(List rows) { + float minX = Float.MAX_VALUE; + float maxX = -Float.MAX_VALUE; + for (Line l : rows) { + for (TextWord w : l.source.words()) { + minX = Math.min(minX, w.x()); + maxX = Math.max(maxX, w.x() + w.width()); + } + } + // Real pages are under ~2000pt wide; anything larger is a malformed/crafted coordinate + // that would allocate a multi-GB array or produce a negative span on overflow. + if (maxX <= minX || (maxX - minX) > 2000f) { + return List.of(); + } + + int lo = (int) Math.floor(minX); + int span = Math.min((int) Math.ceil(maxX) - lo + 1, 2001); + int[] coverage = new int[span]; + for (Line l : rows) { + boolean[] covered = new boolean[span]; + for (TextWord w : l.source.words()) { + int a = Math.max(0, (int) Math.floor(w.x()) - lo); + int b = Math.min(span, (int) Math.ceil(w.x() + w.width()) - lo); + for (int x = a; x < b; x++) { + covered[x] = true; + } + } + for (int x = 0; x < span; x++) { + if (covered[x]) { + coverage[x]++; + } + } + } + + // A column band must be occupied by at least this many rows; below it is gutter. + int support = Math.max(2, Math.round(rows.size() * 0.35f)); + List columns = new ArrayList<>(); + int start = -1; + for (int x = 0; x < span; x++) { + boolean isColumn = coverage[x] >= support; + if (isColumn && start < 0) { + start = x; + } else if (!isColumn && start >= 0) { + columns.add(new float[] {lo + start, lo + x}); + start = -1; + } + } + if (start >= 0) { + columns.add(new float[] {(float) (lo + start), (float) (lo + span)}); + } + + // Merge bands separated by only a narrow gutter. A real column separator is several + // characters wide; the gaps *inside* a multi-word cell (ordinary word spacing) are about + // one character. Without this, a cell like "January 20th, 2026" — whose words align + // vertically across every row — would be split into three spurious columns. + float charWidth = averageCharWidth(rows); + float minGutter = Math.max(10f, charWidth * 2.5f); + List merged = new ArrayList<>(); + for (float[] band : columns) { + if (!merged.isEmpty() && band[0] - merged.get(merged.size() - 1)[1] < minGutter) { + merged.get(merged.size() - 1)[1] = band[1]; + } else { + merged.add(new float[] {band[0], band[1]}); + } + } + return merged; + } + + private static float averageCharWidth(List rows) { + double totalWidth = 0; + int totalChars = 0; + for (Line l : rows) { + for (TextWord w : l.source.words()) { + totalWidth += w.width(); + totalChars += Math.max(1, w.text().strip().length()); + } + } + return totalChars == 0 ? 6f : (float) (totalWidth / totalChars); + } + + private static int nearestColumn(float x, float[] centers) { + int best = 0; + float bestDist = Float.MAX_VALUE; + for (int i = 0; i < centers.length; i++) { + float d = Math.abs(x - centers[i]); + if (d < bestDist) { + bestDist = d; + best = i; + } + } + return best; + } + + private static int filledCells(String[] row) { + int count = 0; + for (String cell : row) { + if (!cell.isEmpty()) { + count++; + } + } + return count; + } + + private static String renderGfm(List rows, int cols) { + if (rows.isEmpty()) { + return ""; + } + int[] widths = new int[cols]; + for (int c = 0; c < cols; c++) { + widths[c] = 3; + } + for (String[] row : rows) { + for (int c = 0; c < cols; c++) { + if (c < row.length) { + widths[c] = Math.max(widths[c], escapeCell(row[c]).length()); + } + } + } + StringBuilder sb = new StringBuilder(); + sb.append(buildGfmRow(rows.get(0), widths, cols)).append('\n'); + sb.append('|'); + for (int c = 0; c < cols; c++) { + sb.append('-').append("-".repeat(widths[c])).append('-').append('|'); + } + for (int r = 1; r < rows.size(); r++) { + sb.append('\n').append(buildGfmRow(rows.get(r), widths, cols)); + } + return sb.toString(); + } + + /** + * A line looks like a table row if it has at least two words separated by a gap far wider than + * normal inter-word spacing. The threshold is derived from the line's own character width + * rather than a document font size, because some PDFs report a unit (matrix-scaled) font size + * that makes absolute thresholds meaningless. (Two-word rows are allowed so two-column tables + * are detected; spurious matches are filtered later by block contiguity and column + * consistency.) + */ + private static boolean isTableCandidate(TextLine line) { + List words = line.words(); + if (words.size() < 2) { + return false; + } + double totalWidth = 0; + int totalChars = 0; + for (TextWord w : words) { + totalWidth += w.width(); + totalChars += Math.max(1, w.text().strip().length()); + } + float charWidth = (float) (totalWidth / Math.max(1, totalChars)); + // A deliberate cell gap is several blank characters wide; ordinary word spaces are ~a third + // of a character. Floor at 8pt so tiny fonts still need a real gap. + float cellGap = Math.max(8f, charWidth * 3f); + for (int i = 1; i < words.size(); i++) { + TextWord prev = words.get(i - 1); + float gap = words.get(i).x() - (prev.x() + prev.width()); + if (gap >= cellGap) { + return true; + } + } + return false; + } + + private static String buildGfmRow(String[] row, int[] widths, int cols) { + StringBuilder sb = new StringBuilder().append('|'); + for (int c = 0; c < cols; c++) { + String cell = c < row.length ? escapeCell(row[c]) : ""; + sb.append(' ').append(padRight(cell, widths[c])).append(' ').append('|'); + } + return sb.toString(); + } + + private static String escapeCell(String cell) { + // Cell content is inline context: escape inline markdown (including the column delimiter) + // but not leading block markers, which have no meaning inside a table cell. + return escapeMarkdownInline(cell); + } + + /** + * Escapes Markdown control characters in body text extracted from the PDF so that literal + * characters (e.g. a line that reads {@code # Heading} or {@code [label](url)}, or an embedded + * {@code }) are emitted as text rather than being reinterpreted as structure or raw HTML. + * Applied to all body text — headings, paragraphs, bold labels, bullets — before emission. + * + *

The generated Markdown should still be treated as untrusted content by any downstream + * renderer: this hardens fidelity and is defence-in-depth, not a substitute for safe rendering. + */ + private static String escapeMarkdown(String text) { + if (text.isEmpty()) { + return text; + } + String inline = escapeMarkdownInline(text); + return escapeLeadingBlockMarker(inline, text); + } + + /** Escapes inline-significant Markdown characters anywhere in the string. */ + private static String escapeMarkdownInline(String text) { + StringBuilder sb = new StringBuilder(text.length() + 8); + for (int i = 0; i < text.length(); i++) { + char c = text.charAt(i); + switch (c) { + case '\\', '`', '*', '_', '[', ']', '<', '>', '|', '~' -> sb.append('\\').append(c); + default -> sb.append(c); + } + } + return sb.toString(); + } + + /** + * Escapes block-level markers that are only significant at the start of a line: ATX headings + * ({@code #}), unordered list / thematic break markers ({@code -}, {@code +}), and ordered list + * markers ({@code 1.} / {@code 1)}). {@code original} carries the unescaped leading characters, + * none of which are altered by inline escaping, so positions line up with {@code escaped}. + */ + private static String escapeLeadingBlockMarker(String escaped, String original) { + char c0 = original.charAt(0); + if (c0 == '#' || c0 == '-' || c0 == '+') { + return "\\" + escaped; + } + int i = 0; + while (i < original.length() && Character.isDigit(original.charAt(i))) { + i++; + } + if (i > 0 && i < original.length()) { + char delim = original.charAt(i); + if (delim == '.' || delim == ')') { + return escaped.substring(0, i) + "\\" + escaped.substring(i); + } + } + return escaped; + } + + private static String padRight(String s, int width) { + return s.length() >= width ? s : s + " ".repeat(width - s.length()); + } + + // --- Page-level emission helpers --------------------------------------- + + private static void emitImages(PdfDocument doc, int pageIndex, List pageItems) + throws IOException { + try (PdfPage page = doc.page(pageIndex)) { + List images = + PdfImageExtractor.extract(page.rawDocHandle(), page.rawHandle(), pageIndex); + for (ExtractedImage img : images) { + pageItems.add(describeImage(img)); + } + } + } + + /** + * Builds an image placeholder annotated with whatever metadata JPDFium exposes: pixel + * dimensions, on-page placement (points), effective DPI, encoded format, colour space and bit + * depth. Missing fields are simply omitted so the line stays valid for any image. + */ + private static String describeImage(ExtractedImage img) { + List parts = new ArrayList<>(); + if (img.width() > 0 && img.height() > 0) { + parts.add(img.width() + "x" + img.height() + "px"); + } + Rect b = img.bounds(); + if (b != null && b.width() > 0 && b.height() > 0) { + parts.add(String.format("%.0fx%.0fpt", b.width(), b.height())); + if (img.width() > 0) { + float dpiX = img.width() / (b.width() / 72f); + float dpiY = img.height() / (b.height() / 72f); + if (Float.isFinite(dpiX) && dpiX > 0) { + parts.add(String.format("~%.0fdpi", (dpiX + dpiY) / 2f)); + } + } + } + String ext = img.suggestedExtension(); + if (ext != null && !ext.isBlank()) { + parts.add(ext.replaceFirst("^\\.", "").toUpperCase(java.util.Locale.ROOT)); + } + if (img.colorSpace() != null) { + parts.add(img.colorSpace().toString()); + } + if (img.bitsPerPixel() > 0) { + parts.add(img.bitsPerPixel() + "bpp"); + } + + StringBuilder sb = new StringBuilder("'); + return sb.toString(); + } + + private static void mergeAcrossPageBoundary(List output, List pageItems) { + if (output.isEmpty() || pageItems.isEmpty()) { + return; + } + // Only merge a sentence continuation between two text paragraphs, never into/out of a + // table. + if (!(output.get(output.size() - 1) instanceof String last) + || !(pageItems.get(0) instanceof String first)) { + return; + } + if (!first.isEmpty() + && Character.isLowerCase(first.charAt(0)) + && !endsWithSentencePunctuation(last)) { + output.set(output.size() - 1, last + " " + first); + pageItems.remove(0); + } + } + + /** + * Joins tables split across a page break. Two consecutive {@link TableBlock}s (no text between + * them — i.e. one ended a page and the next began the following page) are merged when their + * column layouts match; a repeated header row on the continuation is dropped. + */ + private static List stitchTables(List elements) { + List out = new ArrayList<>(); + for (Object e : elements) { + if (e instanceof TableBlock tb + && !out.isEmpty() + && out.get(out.size() - 1) instanceof TableBlock prev + && columnsMatch(flatten(prev.rows()), flatten(tb.rows()))) { + List> merged = new ArrayList<>(prev.rows()); + List> tail = tb.rows(); + if (!tail.isEmpty() + && !prev.rows().isEmpty() + && rowText(tail.get(0)).equals(rowText(prev.rows().get(0)))) { + tail = tail.subList(1, tail.size()); + } + merged.addAll(tail); + out.set(out.size() - 1, new TableBlock(merged, prev.top(), tb.bottom())); + } else { + out.add(e); + } + } + return out; + } + + private static String normaliseSpace(String s) { + return s.strip().replaceAll("\\s+", " "); + } + + private static List flatten(List> rows) { + return rows.stream().flatMap(List::stream).collect(Collectors.toList()); + } + + /** Whitespace-normalised text of a row's lines (top to bottom), for header de-duplication. */ + /** + * Header text of a table at the very bottom of a page, or null if the page does not end in one. + * Trailing image placeholders are skipped; any other text after a table means it did not run to + * the page bottom and so is not a continuation candidate. + */ + private static String trailingTableHeader(List pageItems) { + for (int i = pageItems.size() - 1; i >= 0; i--) { + Object e = pageItems.get(i); + if (e instanceof String s && s.strip().startsWith(" row) { + List ordered = new ArrayList<>(row); + ordered.sort(Comparator.comparingDouble((Line l) -> l.y).reversed()); + StringBuilder sb = new StringBuilder(); + for (Line l : ordered) { + if (sb.length() > 0) { + sb.append(' '); + } + sb.append(l.text); + } + return normaliseSpace(sb.toString()); + } + + /** True when two table blocks have the same number of columns at near-identical x-centres. */ + private static boolean columnsMatch(List a, List b) { + List ca = findColumnRanges(a); + List cb = findColumnRanges(b); + if (ca.size() < 2 || ca.size() != cb.size()) { + return false; + } + for (int i = 0; i < ca.size(); i++) { + float centreA = (ca.get(i)[0] + ca.get(i)[1]) / 2f; + float centreB = (cb.get(i)[0] + cb.get(i)[1]) / 2f; + if (Math.abs(centreA - centreB) > 15f) { + return false; + } + } + return true; + } + + private static void flushParagraph(StringBuilder para, List out) { + if (!para.isEmpty()) { + out.add(escapeMarkdown(para.toString())); + para.setLength(0); + } + } + + private static String repairHyphens(String text) { + return SOFT_HYPHEN.matcher(text).replaceAll("$1$2"); + } + + private static boolean endsWithSentencePunctuation(String s) { + if (s.isEmpty()) { + return false; + } + char last = s.charAt(s.length() - 1); + return last == '.' || last == '?' || last == '!' || last == ':'; + } + + // --- Methods used by other components / tests -------------------------- + + List extractAllPageText(PdfDocument doc) throws IOException { + return PdfTextExtractor.extractAll(doc); + } + + List extractTables(PdfDocument doc, int pageIndex) throws IOException { + return PdfTableExtractor.extract(doc, pageIndex); + } + + List renderTables(List
tables) { + return tables.stream().map(TableRenderer::render).toList(); + } +} diff --git a/app/common/src/main/java/stirling/software/common/pdf/TableRenderer.java b/app/common/src/main/java/stirling/software/common/pdf/TableRenderer.java new file mode 100644 index 0000000000..3f468699fb --- /dev/null +++ b/app/common/src/main/java/stirling/software/common/pdf/TableRenderer.java @@ -0,0 +1,82 @@ +package stirling.software.common.pdf; + +import stirling.software.jpdfium.text.Table; + +final class TableRenderer { + private TableRenderer() {} + + /** Renders a Table as a GitHub-Flavoured Markdown table string. */ + static String render(Table table) { + if (table.rowCount() == 0) { + return ""; + } + + String[][] grid = table.asGrid(); + + if (table.rowCount() < 2) { + // No separator row possible — return plain lines + StringBuilder sb = new StringBuilder(); + for (int c = 0; c < grid[0].length; c++) { + if (c > 0) sb.append('\n'); + sb.append(escape(grid[0][c].trim())); + } + return sb.toString(); + } + + int cols = grid[0].length; + + // Compute column widths: max(3, max content length across all rows) + int[] widths = new int[cols]; + for (int c = 0; c < cols; c++) { + widths[c] = 3; + } + for (String[] row : grid) { + for (int c = 0; c < cols; c++) { + String cell = c < row.length ? row[c].trim() : ""; + widths[c] = Math.max(widths[c], escape(cell).length()); + } + } + + StringBuilder sb = new StringBuilder(); + + // Header row + sb.append(buildRow(grid[0], widths, cols)); + sb.append('\n'); + + // Separator row + sb.append('|'); + for (int c = 0; c < cols; c++) { + sb.append('-').append("-".repeat(widths[c])).append('-').append('|'); + } + sb.append('\n'); + + // Data rows + for (int r = 1; r < grid.length; r++) { + sb.append(buildRow(grid[r], widths, cols)); + if (r < grid.length - 1) { + sb.append('\n'); + } + } + + return sb.toString(); + } + + private static String buildRow(String[] row, int[] widths, int cols) { + StringBuilder sb = new StringBuilder(); + sb.append('|'); + for (int c = 0; c < cols; c++) { + String cell = c < row.length ? escape(row[c].trim()) : ""; + sb.append(' ').append(padRight(cell, widths[c])).append(' ').append('|'); + } + return sb.toString(); + } + + private static String escape(String cell) { + return cell.replace("|", "\\|"); + } + + private static String padRight(String s, int width) { + if (s.length() >= width) return s; + return s + " ".repeat(width - s.length()); + } +} diff --git a/app/common/src/main/java/stirling/software/common/service/FileStorage.java b/app/common/src/main/java/stirling/software/common/service/FileStorage.java index e77e6d6fc6..c5c1ddb5db 100644 --- a/app/common/src/main/java/stirling/software/common/service/FileStorage.java +++ b/app/common/src/main/java/stirling/software/common/service/FileStorage.java @@ -5,6 +5,7 @@ import java.io.IOException; import java.io.InputStream; import java.io.PipedInputStream; import java.io.PipedOutputStream; +import java.util.Optional; import java.util.concurrent.Executors; import java.util.concurrent.atomic.AtomicReference; @@ -17,6 +18,7 @@ import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; import stirling.software.common.cluster.FileStore; +import stirling.software.common.util.JobContext; /** * Service for storing and retrieving files with unique file IDs. Used by the AutoJobPostMapping @@ -32,8 +34,10 @@ public class FileStorage { private final FileOrUploadService fileOrUploadService; private final FileStore fileStore; + private final Optional jobOwnershipService; public String storeFile(MultipartFile file) throws IOException { + String owner = resolveOwner(); // Fast path: when Spring buffered the multipart to disk (typical for large uploads), the // backing Resource exposes a real File. Hand the Path to the FileStore so it can do a // file-to-file copy (Linux sendfile, no copy through Java heap) rather than streaming @@ -48,7 +52,7 @@ public class FileStorage { if (res != null && res.isFile()) { try { FileStore.Stored stored = - fileStore.store(res.getFile().toPath(), file.getOriginalFilename()); + fileStore.store(res.getFile().toPath(), file.getOriginalFilename(), owner); log.debug("Stored file with ID: {} (fast path)", stored.fileId()); return stored.fileId(); } catch (IOException ex) { @@ -57,40 +61,45 @@ public class FileStorage { } } try (InputStream in = file.getInputStream()) { - FileStore.Stored stored = fileStore.store(in, file.getOriginalFilename()); + FileStore.Stored stored = fileStore.store(in, file.getOriginalFilename(), owner); log.debug("Stored file with ID: {}", stored.fileId()); return stored.fileId(); } } public String storeBytes(byte[] bytes, String originalName) throws IOException { - FileStore.Stored stored = fileStore.store(new ByteArrayInputStream(bytes), originalName); + FileStore.Stored stored = + fileStore.store(new ByteArrayInputStream(bytes), originalName, resolveOwner()); log.debug("Stored byte array with ID: {}", stored.fileId()); return stored.fileId(); } public MultipartFile retrieveFile(String fileId) throws IOException { + enforceOwnership(fileId); byte[] fileData = fileStore.retrieveBytes(fileId); return fileOrUploadService.toMockMultipartFile(fileId, fileData); } public byte[] retrieveBytes(String fileId) throws IOException { + enforceOwnership(fileId); return fileStore.retrieveBytes(fileId); } public InputStream retrieveInputStream(String fileId) throws IOException { + enforceOwnership(fileId); return fileStore.retrieve(fileId); } public StoredFile storeInputStream(InputStream inputStream, String originalName) throws IOException { - FileStore.Stored stored = fileStore.store(inputStream, originalName); + FileStore.Stored stored = fileStore.store(inputStream, originalName, resolveOwner()); log.debug("Stored input stream with ID: {}", stored.fileId()); return new StoredFile(stored.fileId(), stored.size()); } public String storeFromStreamingBody(StreamingResponseBody body, String originalName) throws IOException { + String owner = resolveOwner(); // Hold Throwable not IOException: an unchecked failure (NPE, IllegalState, OOM, etc.) // from the body writer would otherwise close the pipe with EOF and the consumer would // return a truncated file with no error surfaced to the caller. @@ -115,7 +124,7 @@ public class FileStorage { } } }); - FileStore.Stored stored = fileStore.store(in, originalName); + FileStore.Stored stored = fileStore.store(in, originalName, owner); Throwable writerErr = bodyError.get(); if (writerErr != null) { // Body failed mid-write: the FileStore persisted a truncated entry. @@ -159,21 +168,62 @@ public class FileStorage { public String storeFromResource(Resource resource, String originalName) throws IOException { try (InputStream in = resource.getInputStream()) { - FileStore.Stored stored = fileStore.store(in, originalName); + FileStore.Stored stored = fileStore.store(in, originalName, resolveOwner()); log.debug("Stored Resource with ID: {}", stored.fileId()); return stored.fileId(); } } public boolean deleteFile(String fileId) { + enforceOwnership(fileId); return fileStore.delete(fileId); } public boolean fileExists(String fileId) { + enforceOwnership(fileId); return fileStore.exists(fileId); } public long getFileSize(String fileId) throws IOException { + enforceOwnership(fileId); return fileStore.size(fileId); } + + private String resolveOwner() { + String propagated = JobContext.getOwner(); + if (propagated != null) { + return propagated; + } + return jobOwnershipService.flatMap(JobOwnershipService::getCurrentUserId).orElse(null); + } + + private void enforceOwnership(String fileId) { + if (jobOwnershipService.isEmpty()) { + return; + } + Optional currentUser = jobOwnershipService.get().getCurrentUserId(); + if (currentUser.isEmpty()) { + return; + } + String owner; + try { + owner = fileStore.getOwner(fileId); + } catch (IOException e) { + log.warn("Failed to read owner for file {}: {}", fileId, e.getMessage()); + throw new SecurityException( + "Access denied: could not verify ownership of the requested file"); + } + if (owner == null) { + return; + } + if (!owner.equals(currentUser.get())) { + log.warn( + "Access denied: user {} attempted to access file {} owned by {}", + currentUser.get(), + fileId, + owner); + throw new SecurityException( + "Access denied: you do not have permission to access this file"); + } + } } diff --git a/app/common/src/main/java/stirling/software/common/service/InternalApiClient.java b/app/common/src/main/java/stirling/software/common/service/InternalApiClient.java index f72eca14d5..8df2d5e410 100644 --- a/app/common/src/main/java/stirling/software/common/service/InternalApiClient.java +++ b/app/common/src/main/java/stirling/software/common/service/InternalApiClient.java @@ -50,6 +50,16 @@ public class InternalApiClient { "^/api/v1/(general|misc|security|convert|filter)(/[A-Za-z0-9_-]+)+$" + "|^/api/v1/ai/tools(/[A-Za-z0-9_-]+)+$"); + /** + * Marker propagated on every internal sub-step dispatch so the saas PAYG interceptor classifies + * the call as {@code BillingCategory.AUTOMATION}. By construction every {@link + * InternalApiClient#post} caller is an automation surface (pipeline executor, AI workflow, + * policy runner) running a child tool inside a parent automation flow — see the saas {@code + * PaygChargeInterceptor.determineCategory} precedence chain, where this header dominates any + * per-tool {@code @RequiresFeature} annotation. + */ + public static final String AUTOMATION_HEADER = "X-Stirling-Automation"; + private final ServletContext servletContext; private final UserServiceInterface userService; private final TempFileManager tempFileManager; @@ -96,7 +106,23 @@ public class InternalApiClient { if (apiKey != null && !apiKey.isEmpty()) { headers.add("X-API-KEY", apiKey); } + // Tag the sub-step as automation so PAYG bills it under AUTOMATION regardless of which + // tool-level @RequiresFeature annotation the dispatched controller carries (e.g. an AI-OCR + // step inside a policy run must bill as AUTOMATION, not AI). Set unconditionally because + // every caller of this dispatcher is an automation surface by design. + headers.add(AUTOMATION_HEADER, "true"); + // A no-file ai/tools call (e.g. create-pdf-from-html-agent) sends only string params, so + // without this RestTemplate would use urlencoded instead of the multipart the controller + // expects. File-bearing calls get the right multipart content-type from RestTemplate. + boolean isAiTool = endpointPath.startsWith("/api/v1/ai/tools/"); + boolean hasFilePart = + body.values().stream() + .flatMap(java.util.List::stream) + .anyMatch(v -> v instanceof Resource); + if (isAiTool && !hasFilePart) { + headers.setContentType(MediaType.MULTIPART_FORM_DATA); + } HttpEntity> entity = new HttpEntity<>(body, headers); RequestCallback requestCallback = restTemplate.httpEntityCallback(entity, Resource.class); diff --git a/app/common/src/main/java/stirling/software/common/service/JobExecutorService.java b/app/common/src/main/java/stirling/software/common/service/JobExecutorService.java index 629d28ba64..23a23e868b 100644 --- a/app/common/src/main/java/stirling/software/common/service/JobExecutorService.java +++ b/app/common/src/main/java/stirling/software/common/service/JobExecutorService.java @@ -89,6 +89,11 @@ public class JobExecutorService { String jobId = scopedJobKey; + final String jobOwner = + jobOwnershipService != null + ? jobOwnershipService.getCurrentUserId().orElse(null) + : null; + long timeoutToUse = customTimeoutMs > 0 ? customTimeoutMs : effectiveTimeoutMs; log.debug( @@ -119,6 +124,7 @@ public class JobExecutorService { try { stirling.software.common.util.JobContext.setJobId( capturedJobIdForQueue); + stirling.software.common.util.JobContext.setOwner(jobOwner); Object result = work.get(); processJobResult(capturedJobIdForQueue, result); return result; @@ -153,6 +159,7 @@ public class JobExecutorService { timeoutToUse); stirling.software.common.util.JobContext.setJobId(capturedJobId); + stirling.software.common.util.JobContext.setOwner(jobOwner); Object result = executeWithTimeout(() -> work.get(), timeoutToUse); processJobResult(capturedJobId, result); } catch (TimeoutException te) { diff --git a/app/common/src/main/java/stirling/software/common/service/MobileScannerService.java b/app/common/src/main/java/stirling/software/common/service/MobileScannerService.java index 49512af708..7c544b6242 100644 --- a/app/common/src/main/java/stirling/software/common/service/MobileScannerService.java +++ b/app/common/src/main/java/stirling/software/common/service/MobileScannerService.java @@ -3,7 +3,6 @@ package stirling.software.common.service; import java.io.IOException; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import java.util.ArrayList; import java.util.HashMap; import java.util.List; @@ -35,7 +34,7 @@ public class MobileScannerService { public MobileScannerService() throws IOException { // Create temp directory for mobile scanner uploads this.tempDirectory = - Paths.get(System.getProperty("java.io.tmpdir"), "stirling-mobile-scanner"); + Path.of(System.getProperty("java.io.tmpdir"), "stirling-mobile-scanner"); Files.createDirectories(tempDirectory); log.info("Mobile scanner temp directory: {}", tempDirectory); } diff --git a/app/common/src/main/java/stirling/software/common/service/PostHogService.java b/app/common/src/main/java/stirling/software/common/service/PostHogService.java index 92093762fc..f6a339ab9a 100644 --- a/app/common/src/main/java/stirling/software/common/service/PostHogService.java +++ b/app/common/src/main/java/stirling/software/common/service/PostHogService.java @@ -8,7 +8,7 @@ import java.lang.management.OperatingSystemMXBean; import java.lang.management.RuntimeMXBean; import java.lang.management.ThreadMXBean; import java.nio.file.Files; -import java.nio.file.Paths; +import java.nio.file.Path; import java.util.HashMap; import java.util.Locale; import java.util.Map; @@ -160,7 +160,7 @@ public class PostHogService { } private boolean isRunningInDocker() { - return Files.exists(Paths.get("/.dockerenv")); + return Files.exists(Path.of("/.dockerenv")); } private Map getDockerMetrics() { diff --git a/app/common/src/main/java/stirling/software/common/service/ToolMetadataService.java b/app/common/src/main/java/stirling/software/common/service/ToolMetadataService.java index 662878b741..fb7a928d76 100644 --- a/app/common/src/main/java/stirling/software/common/service/ToolMetadataService.java +++ b/app/common/src/main/java/stirling/software/common/service/ToolMetadataService.java @@ -1,11 +1,21 @@ package stirling.software.common.service; +import java.util.List; + /** Provides metadata about tool endpoints for internal dispatch. */ public interface ToolMetadataService { /** Returns true if the given operation path accepts multiple input files. */ boolean isMultiInput(String operationPath); + /** + * Returns the file extensions (lowercase, no leading dot, e.g. {@code "pdf"}) that the + * operation accepts as input ({@code output=false}) or produces as output ({@code + * output=true}), derived from the endpoint's declared type. Returns {@code null} when the + * endpoint declares no specific type, which callers should treat as "any type accepted". + */ + List getExtensionTypes(boolean output, String operationPath); + /** * Returns true when the endpoint's ZIP response is a transport for multiple typed results and * should be unpacked: multi-output endpoints (Type:SIMO / Type:MIMO) and wrapper declarations diff --git a/app/common/src/main/java/stirling/software/common/util/GeneralUtils.java b/app/common/src/main/java/stirling/software/common/util/GeneralUtils.java index 61ba8670fd..fbcf30fdff 100644 --- a/app/common/src/main/java/stirling/software/common/util/GeneralUtils.java +++ b/app/common/src/main/java/stirling/software/common/util/GeneralUtils.java @@ -255,7 +255,7 @@ public class GeneralUtils { String pattern = locationPattern; if (pattern.startsWith("file:")) { String rawPath = pattern.substring(5).replace("\\*", "").replace("/*", ""); - Path normalizePath = Paths.get(rawPath).normalize(); + Path normalizePath = Path.of(rawPath).normalize(); pattern = "file:" + normalizePath.toString().replace("\\", "/") + "/*"; } return ResourcePatternUtils.getResourcePatternResolver(resourceLoader) @@ -837,7 +837,7 @@ public class GeneralUtils { } public boolean createDir(String path) { - Path folder = Paths.get(path); + Path folder = Path.of(path); if (!Files.exists(folder)) { try { Files.createDirectories(folder); @@ -867,7 +867,7 @@ public class GeneralUtils { public void saveKeyToSettings(String key, Object newValue) throws IOException { String[] keyArray = key.split("\\."); - Path settingsPath = Paths.get(InstallationPathConfig.getSettingsPath()); + Path settingsPath = Path.of(InstallationPathConfig.getSettingsPath()); YamlHelper settingsYaml = new YamlHelper(settingsPath); settingsYaml.updateValue(Arrays.asList(keyArray), newValue); settingsYaml.saveOverride(settingsPath); @@ -888,7 +888,7 @@ public class GeneralUtils { return; } - Path settingsPath = Paths.get(InstallationPathConfig.getSettingsPath()); + Path settingsPath = Path.of(InstallationPathConfig.getSettingsPath()); YamlHelper settingsYaml = new YamlHelper(settingsPath); // Apply all updates to the same YamlHelper instance @@ -974,11 +974,11 @@ public class GeneralUtils { */ public void extractPipeline() throws IOException { Path pipelineDir = - Paths.get(InstallationPathConfig.getPipelinePath(), DEFAULT_WEBUI_CONFIGS_DIR); + Path.of(InstallationPathConfig.getPipelinePath(), DEFAULT_WEBUI_CONFIGS_DIR); Files.createDirectories(pipelineDir); for (String name : DEFAULT_VALID_PIPELINE) { - if (!Paths.get(name).getFileName().toString().equals(name)) { + if (!Path.of(name).getFileName().toString().equals(name)) { log.error("Invalid pipeline file name: {}", name); throw new IllegalArgumentException("Invalid pipeline file name: " + name); } @@ -1014,7 +1014,7 @@ public class GeneralUtils { throw new IllegalArgumentException( "scriptName must not contain path traversal characters"); } - if (!Paths.get(scriptName).getFileName().toString().equals(scriptName)) { + if (!Path.of(scriptName).getFileName().toString().equals(scriptName)) { throw new IllegalArgumentException( "scriptName must not contain path traversal characters"); } @@ -1024,7 +1024,7 @@ public class GeneralUtils { "scriptName must be either 'png_to_webp.py' or 'split_photos.py'"); } - Path scriptsDir = Paths.get(InstallationPathConfig.getScriptsPath(), PYTHON_SCRIPTS_DIR); + Path scriptsDir = Path.of(InstallationPathConfig.getScriptsPath(), PYTHON_SCRIPTS_DIR); Files.createDirectories(scriptsDir); Path target = scriptsDir.resolve(scriptName); diff --git a/app/common/src/main/java/stirling/software/common/util/JarPathUtil.java b/app/common/src/main/java/stirling/software/common/util/JarPathUtil.java index e5c6488eb2..bb0e39491f 100644 --- a/app/common/src/main/java/stirling/software/common/util/JarPathUtil.java +++ b/app/common/src/main/java/stirling/software/common/util/JarPathUtil.java @@ -4,7 +4,6 @@ import java.io.File; import java.net.URISyntaxException; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import lombok.extern.slf4j.Slf4j; @@ -20,7 +19,7 @@ public class JarPathUtil { public static Path currentJar() { try { Path jar = - Paths.get( + Path.of( JarPathUtil.class .getProtectionDomain() .getCodeSource() @@ -61,14 +60,14 @@ public class JarPathUtil { } // Location 2: ./build/libs/ (development build) - possibleLocations[1] = Paths.get("build", "libs", "restart-helper.jar").toAbsolutePath(); + possibleLocations[1] = Path.of("build", "libs", "restart-helper.jar").toAbsolutePath(); // Location 3: app/common/build/libs/ (multi-module build) possibleLocations[2] = - Paths.get("app", "common", "build", "libs", "restart-helper.jar").toAbsolutePath(); + Path.of("app", "common", "build", "libs", "restart-helper.jar").toAbsolutePath(); // Location 4: Current working directory - possibleLocations[3] = Paths.get("restart-helper.jar").toAbsolutePath(); + possibleLocations[3] = Path.of("restart-helper.jar").toAbsolutePath(); // Check each location for (Path location : possibleLocations) { diff --git a/app/common/src/main/java/stirling/software/common/util/JobContext.java b/app/common/src/main/java/stirling/software/common/util/JobContext.java index a413949147..f7016b8f71 100644 --- a/app/common/src/main/java/stirling/software/common/util/JobContext.java +++ b/app/common/src/main/java/stirling/software/common/util/JobContext.java @@ -1,8 +1,9 @@ package stirling.software.common.util; -/** Thread-local context for passing job ID across async boundaries */ +/** Thread-local context for passing job ID and owner across async boundaries */ public class JobContext { private static final ThreadLocal CURRENT_JOB_ID = new ThreadLocal<>(); + private static final ThreadLocal CURRENT_OWNER = new ThreadLocal<>(); public static void setJobId(String jobId) { CURRENT_JOB_ID.set(jobId); @@ -12,7 +13,16 @@ public class JobContext { return CURRENT_JOB_ID.get(); } + public static void setOwner(String owner) { + CURRENT_OWNER.set(owner); + } + + public static String getOwner() { + return CURRENT_OWNER.get(); + } + public static void clear() { CURRENT_JOB_ID.remove(); + CURRENT_OWNER.remove(); } } diff --git a/app/common/src/main/java/stirling/software/common/util/OfficeDocumentSanitizer.java b/app/common/src/main/java/stirling/software/common/util/OfficeDocumentSanitizer.java new file mode 100644 index 0000000000..9cdf8cf53b --- /dev/null +++ b/app/common/src/main/java/stirling/software/common/util/OfficeDocumentSanitizer.java @@ -0,0 +1,310 @@ +package stirling.software.common.util; + +import java.io.ByteArrayInputStream; +import java.io.ByteArrayOutputStream; +import java.io.IOException; +import java.util.ArrayList; +import java.util.List; +import java.util.Locale; +import java.util.Set; +import java.util.zip.ZipEntry; +import java.util.zip.ZipInputStream; +import java.util.zip.ZipOutputStream; + +import javax.xml.XMLConstants; +import javax.xml.parsers.DocumentBuilder; +import javax.xml.parsers.DocumentBuilderFactory; +import javax.xml.parsers.ParserConfigurationException; +import javax.xml.transform.OutputKeys; +import javax.xml.transform.Transformer; +import javax.xml.transform.TransformerException; +import javax.xml.transform.TransformerFactory; +import javax.xml.transform.dom.DOMSource; +import javax.xml.transform.stream.StreamResult; + +import org.springframework.stereotype.Component; +import org.w3c.dom.Document; +import org.w3c.dom.Element; +import org.w3c.dom.NamedNodeMap; +import org.w3c.dom.Node; +import org.w3c.dom.NodeList; +import org.xml.sax.SAXException; + +import io.github.pixee.security.ZipSecurity; + +import lombok.extern.slf4j.Slf4j; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.common.service.SsrfProtectionService; + +// Strips external refs from OOXML/ODF uploads so LibreOffice can't be made to fetch them. +@Component +@Slf4j +public class OfficeDocumentSanitizer { + + private static final Set OOXML_EXTENSIONS = + Set.of( + "docx", "docm", "dotx", "dotm", "xlsx", "xlsm", "xltx", "xltm", "pptx", "pptm", + "potx", "potm", "ppsx", "ppsm"); + + private static final Set ODF_EXTENSIONS = + Set.of( + "odt", "ott", "ods", "ots", "odp", "otp", "odg", "otg", "odf", "odc", "odi", + "odm"); + + private static final Set ODF_XML_PARTS = + Set.of("content.xml", "styles.xml", "meta.xml", "settings.xml"); + + private final SsrfProtectionService ssrfProtectionService; + private final ApplicationProperties applicationProperties; + + public OfficeDocumentSanitizer( + SsrfProtectionService ssrfProtectionService, + ApplicationProperties applicationProperties) { + this.ssrfProtectionService = ssrfProtectionService; + this.applicationProperties = applicationProperties; + } + + public boolean isSanitizableExtension(String extension) { + if (extension == null) { + return false; + } + String lower = extension.toLowerCase(Locale.ROOT); + return OOXML_EXTENSIONS.contains(lower) || ODF_EXTENSIONS.contains(lower); + } + + public byte[] sanitize(byte[] documentBytes, String extension) throws IOException { + if (documentBytes == null || documentBytes.length == 0) { + throw new IOException("Office document input is empty or null"); + } + if (applicationProperties.getSystem().isDisableSanitize()) { + log.debug("Office document sanitization disabled by configuration"); + return documentBytes; + } + if (!isSanitizableExtension(extension)) { + return documentBytes; + } + + ByteArrayOutputStream out = new ByteArrayOutputStream(documentBytes.length); + try (ZipInputStream zipIn = + ZipSecurity.createHardenedInputStream( + new ByteArrayInputStream(documentBytes)); + ZipOutputStream zipOut = new ZipOutputStream(out)) { + + ZipEntry entry; + while ((entry = zipIn.getNextEntry()) != null) { + String name = entry.getName(); + byte[] bytes = entry.isDirectory() ? new byte[0] : zipIn.readAllBytes(); + + if (!entry.isDirectory()) { + bytes = sanitizeEntry(name, bytes); + } + + ZipEntry outEntry = new ZipEntry(name); + if (entry.getComment() != null) { + outEntry.setComment(entry.getComment()); + } + if (entry.getExtra() != null) { + outEntry.setExtra(entry.getExtra()); + } + zipOut.putNextEntry(outEntry); + if (!entry.isDirectory()) { + zipOut.write(bytes); + } + zipOut.closeEntry(); + } + } + return out.toByteArray(); + } + + private byte[] sanitizeEntry(String entryName, byte[] entryBytes) { + String lower = entryName.toLowerCase(Locale.ROOT); + try { + if (lower.endsWith(".rels")) { + return sanitizeOoxmlRels(entryBytes); + } + if (isOdfXmlPart(lower)) { + return sanitizeOdfXml(entryBytes); + } + } catch (ParserConfigurationException + | SAXException + | IOException + | TransformerException e) { + log.warn( + "Failed to parse XML part '{}' for sanitization, leaving as-is: {}", + entryName, + e.getMessage()); + } + return entryBytes; + } + + private boolean isOdfXmlPart(String lowerName) { + int slash = lowerName.lastIndexOf('/'); + String base = slash >= 0 ? lowerName.substring(slash + 1) : lowerName; + return ODF_XML_PARTS.contains(base); + } + + private byte[] sanitizeOoxmlRels(byte[] xmlBytes) + throws IOException, ParserConfigurationException, SAXException, TransformerException { + Document doc = parseSecurely(xmlBytes); + Element root = doc.getDocumentElement(); + if (root == null) { + return xmlBytes; + } + NodeList relationships = root.getElementsByTagNameNS("*", "Relationship"); + List toRemove = new ArrayList<>(); + for (int i = 0; i < relationships.getLength(); i++) { + Node node = relationships.item(i); + NamedNodeMap attrs = node.getAttributes(); + if (attrs == null) { + continue; + } + Node targetMode = attrs.getNamedItem("TargetMode"); + if (targetMode == null || !"external".equalsIgnoreCase(targetMode.getNodeValue())) { + continue; + } + Node target = attrs.getNamedItem("Target"); + String targetValue = target == null ? "" : target.getNodeValue(); + if (isAdminAllowed(targetValue)) { + continue; + } + log.warn( + "Stripping OOXML external relationship target: {}", + truncateForLog(targetValue)); + toRemove.add(node); + } + if (toRemove.isEmpty()) { + return xmlBytes; + } + for (Node n : toRemove) { + n.getParentNode().removeChild(n); + } + return serializeDocument(doc); + } + + private byte[] sanitizeOdfXml(byte[] xmlBytes) + throws IOException, ParserConfigurationException, SAXException, TransformerException { + Document doc = parseSecurely(xmlBytes); + Element root = doc.getDocumentElement(); + if (root == null) { + return xmlBytes; + } + boolean modified = stripExternalHrefs(root); + if (!modified) { + return xmlBytes; + } + return serializeDocument(doc); + } + + private boolean stripExternalHrefs(Node node) { + boolean modified = false; + if (node.getNodeType() == Node.ELEMENT_NODE) { + NamedNodeMap attrs = node.getAttributes(); + List hrefAttrsToRemove = new ArrayList<>(); + for (int i = 0; i < attrs.getLength(); i++) { + Node attr = attrs.item(i); + String name = attr.getNodeName(); + if (name == null) { + continue; + } + String lower = name.toLowerCase(Locale.ROOT); + if (!(lower.equals("xlink:href") + || lower.endsWith(":href") + || lower.equals("href"))) { + continue; + } + String value = attr.getNodeValue(); + if (!isExternalUrl(value)) { + continue; + } + if (isAdminAllowed(value)) { + continue; + } + log.warn( + "Stripping ODF external href attribute ({}): {}", + name, + truncateForLog(value)); + hrefAttrsToRemove.add(name); + } + Element element = (Element) node; + for (String attrName : hrefAttrsToRemove) { + element.removeAttribute(attrName); + modified = true; + } + } + NodeList children = node.getChildNodes(); + for (int i = 0; i < children.getLength(); i++) { + if (stripExternalHrefs(children.item(i))) { + modified = true; + } + } + return modified; + } + + private boolean isExternalUrl(String url) { + if (url == null) { + return false; + } + String trimmed = url.trim().toLowerCase(Locale.ROOT); + if (trimmed.isEmpty() || trimmed.startsWith("#") || trimmed.startsWith("../")) { + return false; + } + return trimmed.startsWith("http://") + || trimmed.startsWith("https://") + || trimmed.startsWith("ftp://") + || trimmed.startsWith("ftps://") + || trimmed.startsWith("file:") + || trimmed.startsWith("smb:") + || trimmed.startsWith("\\\\") + || trimmed.startsWith("//"); + } + + // Preserved only with an explicit allowedDomains entry; MEDIUM default would admit public URLs. + private boolean isAdminAllowed(String url) { + if (ssrfProtectionService == null || url == null || url.isBlank()) { + return false; + } + ApplicationProperties.Html.UrlSecurity config = + applicationProperties.getSystem().getHtml().getUrlSecurity(); + if (config == null + || config.getAllowedDomains() == null + || config.getAllowedDomains().isEmpty()) { + return false; + } + return ssrfProtectionService.isUrlAllowed(url); + } + + private Document parseSecurely(byte[] xmlBytes) + throws ParserConfigurationException, SAXException, IOException { + DocumentBuilderFactory factory = DocumentBuilderFactory.newInstance(); + factory.setFeature(XMLConstants.FEATURE_SECURE_PROCESSING, true); + factory.setFeature("http://apache.org/xml/features/disallow-doctype-decl", true); + factory.setFeature("http://xml.org/sax/features/external-general-entities", false); + factory.setFeature("http://xml.org/sax/features/external-parameter-entities", false); + factory.setFeature("http://apache.org/xml/features/nonvalidating/load-external-dtd", false); + factory.setXIncludeAware(false); + factory.setExpandEntityReferences(false); + factory.setNamespaceAware(true); + DocumentBuilder builder = factory.newDocumentBuilder(); + return builder.parse(new ByteArrayInputStream(xmlBytes)); + } + + private byte[] serializeDocument(Document doc) throws TransformerException { + TransformerFactory tf = TransformerFactory.newInstance(); + tf.setFeature(XMLConstants.FEATURE_SECURE_PROCESSING, true); + Transformer transformer = tf.newTransformer(); + transformer.setOutputProperty(OutputKeys.ENCODING, "UTF-8"); + transformer.setOutputProperty(OutputKeys.INDENT, "no"); + transformer.setOutputProperty(OutputKeys.OMIT_XML_DECLARATION, "no"); + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + transformer.transform(new DOMSource(doc), new StreamResult(baos)); + return baos.toByteArray(); + } + + private String truncateForLog(String value) { + if (value == null) { + return "null"; + } + return value.length() > 80 ? value.substring(0, 80) + "..." : value; + } +} diff --git a/app/common/src/test/java/stirling/software/SPDF/config/EndpointConfigurationGapTest.java b/app/common/src/test/java/stirling/software/SPDF/config/EndpointConfigurationGapTest.java new file mode 100644 index 0000000000..6275b49343 --- /dev/null +++ b/app/common/src/test/java/stirling/software/SPDF/config/EndpointConfigurationGapTest.java @@ -0,0 +1,481 @@ +package stirling.software.SPDF.config; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.util.List; +import java.util.Set; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; + +import stirling.software.SPDF.config.EndpointConfiguration.DisableReason; +import stirling.software.SPDF.config.EndpointConfiguration.EndpointAvailability; +import stirling.software.common.model.ApplicationProperties; + +/** + * Unit tests for {@link EndpointConfiguration}. The class wires up its endpoint/group registry in + * {@code init()} during construction and then applies environment overrides. We build it with a + * real {@link ApplicationProperties} (whose System/Endpoints sub-objects are non-null by default) + * so the constructor runs cleanly without any mocking. + */ +class EndpointConfigurationGapTest { + + private ApplicationProperties applicationProperties; + + /** + * Construct an EndpointConfiguration with the given pro flag and current applicationProperties. + */ + private EndpointConfiguration build(boolean runningProOrHigher) { + return new EndpointConfiguration(applicationProperties, runningProOrHigher); + } + + /** Default config: not pro, no removals, url-to-pdf disabled (default System flag is false). */ + private EndpointConfiguration buildDefault() { + return build(false); + } + + @BeforeEach + void setUp() { + applicationProperties = new ApplicationProperties(); + } + + @Nested + @DisplayName("endpointKeyForUri (static)") + class EndpointKeyForUriTests { + + @Test + @DisplayName("returns null for null uri") + void nullUri() { + assertNull(EndpointConfiguration.endpointKeyForUri(null)); + } + + @Test + @DisplayName("returns null when uri does not contain /api/v1") + void notApiPath() { + assertNull(EndpointConfiguration.endpointKeyForUri("/foo/bar/baz")); + assertNull(EndpointConfiguration.endpointKeyForUri("https://example.com/home")); + } + + @Test + @DisplayName("returns null when uri has too few path segments") + void tooFewSegments() { + // "/api/v1/general" splits to ["", "api", "v1", "general"] -> length 4, not > 4 + assertNull(EndpointConfiguration.endpointKeyForUri("/api/v1/general")); + } + + @Test + @DisplayName("extracts plain endpoint key from a standard /api/v1// uri") + void plainEndpoint() { + assertEquals( + "remove-pages", + EndpointConfiguration.endpointKeyForUri("/api/v1/general/remove-pages")); + } + + @Test + @DisplayName("builds a -to- key for convert endpoints") + void convertEndpoint() { + assertEquals( + "pdf-to-img", + EndpointConfiguration.endpointKeyForUri("/api/v1/convert/pdf/img")); + } + + @Test + @DisplayName("convert path without a target segment falls back to the segment after group") + void convertWithoutTarget() { + // "/api/v1/convert/pdf" -> length 5, the convert branch needs length > 5 + assertEquals("pdf", EndpointConfiguration.endpointKeyForUri("/api/v1/convert/pdf")); + } + } + + @Nested + @DisplayName("enable / disable endpoint") + class EnableDisableEndpointTests { + + @Test + @DisplayName("a freshly registered endpoint is enabled by default") + void enabledByDefault() { + EndpointConfiguration config = buildDefault(); + assertTrue(config.isEndpointEnabled("merge-pdfs")); + } + + @Test + @DisplayName("disableEndpoint marks the endpoint disabled") + void disableEndpoint() { + EndpointConfiguration config = buildDefault(); + config.disableEndpoint("merge-pdfs"); + assertFalse(config.isEndpointEnabled("merge-pdfs")); + } + + @Test + @DisplayName("enableEndpoint re-enables a previously disabled endpoint") + void reEnableEndpoint() { + EndpointConfiguration config = buildDefault(); + config.disableEndpoint("merge-pdfs"); + assertFalse(config.isEndpointEnabled("merge-pdfs")); + config.enableEndpoint("merge-pdfs"); + assertTrue(config.isEndpointEnabled("merge-pdfs")); + } + + @Test + @DisplayName("leading slash is normalized away on disable") + void leadingSlashNormalizedOnDisable() { + EndpointConfiguration config = buildDefault(); + config.disableEndpoint("/merge-pdfs"); + // both forms resolve to the same key + assertFalse(config.isEndpointEnabled("merge-pdfs")); + assertFalse(config.isEndpointEnabled("/merge-pdfs")); + } + + @Test + @DisplayName("isEndpointEnabled tolerates a leading slash on the query") + void leadingSlashOnQuery() { + EndpointConfiguration config = buildDefault(); + assertTrue(config.isEndpointEnabled("/merge-pdfs")); + } + + @Test + @DisplayName("disabling clears with enable, removing the disable reason") + void enableClearsReason() { + EndpointConfiguration config = buildDefault(); + config.disableEndpoint("split-pages", DisableReason.DEPENDENCY); + assertEquals( + DisableReason.DEPENDENCY, + config.getEndpointAvailability("split-pages").getReason()); + config.enableEndpoint("split-pages"); + EndpointAvailability availability = config.getEndpointAvailability("split-pages"); + assertTrue(availability.isEnabled()); + assertNull(availability.getReason()); + } + } + + @Nested + @DisplayName("isEndpointEnabledForUri") + class IsEndpointEnabledForUriTests { + + @Test + @DisplayName("translates a /api/v1 uri to a key and reports its status") + void translatesUri() { + EndpointConfiguration config = buildDefault(); + assertTrue(config.isEndpointEnabledForUri("/api/v1/general/merge-pdfs")); + config.disableEndpoint("merge-pdfs"); + assertFalse(config.isEndpointEnabledForUri("/api/v1/general/merge-pdfs")); + } + + @Test + @DisplayName("falls back to treating a non-api uri as a raw key") + void fallsBackToRawKey() { + EndpointConfiguration config = buildDefault(); + config.disableEndpoint("merge-pdfs"); + // non-api path: key resolution returns null, so the uri itself is used as the key + assertFalse(config.isEndpointEnabledForUri("merge-pdfs")); + } + } + + @Nested + @DisplayName("group enable / disable") + class GroupTests { + + @Test + @DisplayName("a functional group with all endpoints enabled reports enabled") + void functionalGroupEnabled() { + EndpointConfiguration config = buildDefault(); + assertTrue(config.isGroupEnabled("PageOps")); + } + + @Test + @DisplayName("disabling a functional group cascades to all its endpoints") + void disableFunctionalGroupCascades() { + EndpointConfiguration config = buildDefault(); + config.disableGroup("PageOps"); + assertFalse(config.isGroupEnabled("PageOps")); + assertFalse(config.isEndpointEnabled("remove-pages")); + assertFalse(config.isEndpointEnabled("split-pages")); + } + + @Test + @DisplayName("re-enabling a functional group re-enables its endpoints") + void enableFunctionalGroupRestores() { + EndpointConfiguration config = buildDefault(); + config.disableGroup("PageOps"); + assertFalse(config.isEndpointEnabled("remove-pages")); + config.enableGroup("PageOps"); + assertTrue(config.isEndpointEnabled("remove-pages")); + assertTrue(config.isGroupEnabled("PageOps")); + } + + @Test + @DisplayName("a functional group with one disabled endpoint is not enabled") + void functionalGroupWithDisabledEndpoint() { + EndpointConfiguration config = buildDefault(); + config.disableEndpoint("remove-pages"); + assertFalse(config.isGroupEnabled("PageOps")); + } + + @Test + @DisplayName("disabledGroups reflects disabled groups and getDisabledGroups returns a copy") + void getDisabledGroupsReturnsCopy() { + EndpointConfiguration config = buildDefault(); + config.disableGroup("PageOps"); + Set disabled = config.getDisabledGroups(); + assertTrue(disabled.contains("PageOps")); + // mutating the returned set must not affect internal state + disabled.clear(); + assertTrue(config.getDisabledGroups().contains("PageOps")); + } + + @Test + @DisplayName("an unknown group with no endpoints is not enabled") + void unknownGroupNotEnabled() { + EndpointConfiguration config = buildDefault(); + assertFalse(config.isGroupEnabled("NoSuchGroupXyz")); + } + } + + @Nested + @DisplayName("tool group semantics") + class ToolGroupTests { + + @Test + @DisplayName("a tool group is enabled until explicitly disabled") + void toolGroupEnabledUntilDisabled() { + EndpointConfiguration config = buildDefault(); + assertTrue(config.isGroupEnabled("qpdf")); + config.disableGroup("qpdf"); + assertFalse(config.isGroupEnabled("qpdf")); + } + + @Test + @DisplayName("disabling a tool group does NOT cascade to its endpoints directly") + void toolGroupNoCascade() { + EndpointConfiguration config = buildDefault(); + // repair has alternatives (qpdf, Ghostscript); disabling only qpdf keeps it enabled + config.disableGroup("qpdf"); + assertTrue(config.isEndpointEnabled("repair")); + } + + @Test + @DisplayName("endpoint with alternatives is disabled only when all tool groups are gone") + void allAlternativesDisabled() { + EndpointConfiguration config = buildDefault(); + config.disableGroup("qpdf"); + config.disableGroup("Ghostscript"); + // repair's only alternatives are qpdf and Ghostscript + assertFalse(config.isEndpointEnabled("repair")); + } + + @Test + @DisplayName("endpoint with a still-enabled alternative stays enabled") + void oneAlternativeRemains() { + EndpointConfiguration config = buildDefault(); + // compress-pdf alternatives: qpdf, Ghostscript, Java + config.disableGroup("qpdf"); + config.disableGroup("Ghostscript"); + assertTrue(config.isEndpointEnabled("compress-pdf")); + config.disableGroup("Java"); + assertFalse(config.isEndpointEnabled("compress-pdf")); + } + + @Test + @DisplayName("single-dependency endpoint (no alternatives) disabled when its tool group is") + void singleDependencyDisabled() { + EndpointConfiguration config = buildDefault(); + // pdf-to-epub depends on Calibre, no alternatives registered + assertTrue(config.isEndpointEnabled("pdf-to-epub")); + config.disableGroup("Calibre"); + assertFalse(config.isEndpointEnabled("pdf-to-epub")); + } + } + + @Nested + @DisplayName("addEndpointToGroup / addEndpointAlternative") + class RegistrationTests { + + @Test + @DisplayName("addEndpointToGroup makes the endpoint part of the group") + void addEndpointToGroup() { + EndpointConfiguration config = buildDefault(); + config.addEndpointToGroup("CustomGroup", "custom-endpoint"); + Set endpoints = config.getEndpointsForGroup("CustomGroup"); + assertTrue(endpoints.contains("custom-endpoint")); + } + + @Test + @DisplayName("disabling a custom functional group disables its added endpoint") + void customFunctionalGroupCascades() { + EndpointConfiguration config = buildDefault(); + config.addEndpointToGroup("CustomGroup", "custom-endpoint"); + assertTrue(config.isEndpointEnabled("custom-endpoint")); + config.disableGroup("CustomGroup"); + assertFalse(config.isEndpointEnabled("custom-endpoint")); + } + + @Test + @DisplayName("getEndpointsForGroup returns an empty set for unknown groups") + void unknownGroupEmptySet() { + EndpointConfiguration config = buildDefault(); + Set endpoints = config.getEndpointsForGroup("NoSuchGroupXyz"); + assertNotNull(endpoints); + assertTrue(endpoints.isEmpty()); + } + } + + @Nested + @DisplayName("getEndpointAvailability / determineDisableReason") + class AvailabilityTests { + + @Test + @DisplayName("an enabled endpoint has a null disable reason") + void enabledHasNullReason() { + EndpointConfiguration config = buildDefault(); + EndpointAvailability availability = config.getEndpointAvailability("merge-pdfs"); + assertTrue(availability.isEnabled()); + assertNull(availability.getReason()); + } + + @Test + @DisplayName("explicit disable preserves the supplied reason") + void explicitDisableReason() { + EndpointConfiguration config = buildDefault(); + config.disableEndpoint("merge-pdfs", DisableReason.DEPENDENCY); + EndpointAvailability availability = config.getEndpointAvailability("merge-pdfs"); + assertFalse(availability.isEnabled()); + assertEquals(DisableReason.DEPENDENCY, availability.getReason()); + } + + @Test + @DisplayName("default disableEndpoint reason is CONFIG") + void defaultDisableReasonIsConfig() { + EndpointConfiguration config = buildDefault(); + config.disableEndpoint("merge-pdfs"); + assertEquals( + DisableReason.CONFIG, config.getEndpointAvailability("merge-pdfs").getReason()); + } + + @Test + @DisplayName("endpoint disabled via functional group reports the group's reason") + void functionalGroupReason() { + EndpointConfiguration config = buildDefault(); + config.disableGroup("PageOps", DisableReason.DEPENDENCY); + EndpointAvailability availability = config.getEndpointAvailability("crop"); + assertFalse(availability.isEnabled()); + // crop is disabled both via group cascade and group membership; reason is DEPENDENCY + assertEquals(DisableReason.DEPENDENCY, availability.getReason()); + } + } + + @Nested + @DisplayName("getAllEndpoints") + class GetAllEndpointsTests { + + @Test + @DisplayName("aggregates endpoints across all groups") + void aggregatesAcrossGroups() { + EndpointConfiguration config = buildDefault(); + Set all = config.getAllEndpoints(); + assertTrue(all.contains("merge-pdfs")); + assertTrue(all.contains("compress-pdf")); + assertTrue(all.contains("ocr-pdf")); + assertFalse(all.isEmpty()); + } + + @Test + @DisplayName("custom endpoints registered after init appear in getAllEndpoints") + void includesCustomEndpoints() { + EndpointConfiguration config = buildDefault(); + config.addEndpointToGroup("CustomGroup", "brand-new-endpoint"); + assertTrue(config.getAllEndpoints().contains("brand-new-endpoint")); + } + } + + @Nested + @DisplayName("environment / constructor driven configuration") + class EnvironmentConfigTests { + + @Test + @DisplayName("url-to-pdf is disabled when enableUrlToPDF is false (default)") + void urlToPdfDisabledByDefault() { + EndpointConfiguration config = buildDefault(); + assertFalse(config.isEndpointEnabled("url-to-pdf")); + } + + @Test + @DisplayName("url-to-pdf stays enabled when enableUrlToPDF is true") + void urlToPdfEnabledWhenFlagSet() { + applicationProperties.getSystem().setEnableUrlToPDF(true); + EndpointConfiguration config = build(false); + assertTrue(config.isEndpointEnabled("url-to-pdf")); + } + + @Test + @DisplayName("endpoints.toRemove disables the listed endpoints at construction") + void endpointsToRemove() { + applicationProperties + .getEndpoints() + .setToRemove(List.of(" merge-pdfs ", "split-pages")); + EndpointConfiguration config = build(false); + // values are trimmed before disabling + assertFalse(config.isEndpointEnabled("merge-pdfs")); + assertFalse(config.isEndpointEnabled("split-pages")); + } + + @Test + @DisplayName("endpoints.groupsToRemove disables the listed groups at construction") + void groupsToRemove() { + applicationProperties.getEndpoints().setGroupsToRemove(List.of(" PageOps ")); + EndpointConfiguration config = build(false); + assertTrue(config.getDisabledGroups().contains("PageOps")); + assertFalse(config.isEndpointEnabled("remove-pages")); + } + + @Test + @DisplayName("non-pro build disables the enterprise group") + void nonProDisablesEnterprise() { + EndpointConfiguration config = build(false); + assertTrue(config.getDisabledGroups().contains("enterprise")); + } + + @Test + @DisplayName("pro build does not disable the enterprise group") + void proDoesNotDisableEnterprise() { + EndpointConfiguration config = build(true); + assertFalse(config.getDisabledGroups().contains("enterprise")); + } + } + + @Nested + @DisplayName("getEndpointStatuses (Lombok getter) and logging summary") + class MiscTests { + + @Test + @DisplayName("getEndpointStatuses reflects explicit disable state") + void endpointStatusesReflectDisable() { + EndpointConfiguration config = buildDefault(); + config.disableEndpoint("merge-pdfs"); + assertEquals(Boolean.FALSE, config.getEndpointStatuses().get("merge-pdfs")); + } + + @Test + @DisplayName("logDisabledEndpointsSummary runs without throwing") + void logSummaryDoesNotThrow() { + EndpointConfiguration config = buildDefault(); + config.disableGroup("PageOps"); + config.disableGroup("qpdf"); + // purely a smoke test of the logging branch coverage + config.logDisabledEndpointsSummary(); + } + + @Test + @DisplayName("logDisabledEndpointsSummary runs when nothing is disabled") + void logSummaryNothingDisabled() { + applicationProperties.getSystem().setEnableUrlToPDF(true); + EndpointConfiguration config = build(true); + config.logDisabledEndpointsSummary(); + } + } +} diff --git a/app/common/src/test/java/stirling/software/SPDF/pdf/parser/LineAlignmentTableParserTest.java b/app/common/src/test/java/stirling/software/SPDF/pdf/parser/LineAlignmentTableParserTest.java deleted file mode 100644 index fbbf5af9cf..0000000000 --- a/app/common/src/test/java/stirling/software/SPDF/pdf/parser/LineAlignmentTableParserTest.java +++ /dev/null @@ -1,153 +0,0 @@ -package stirling.software.SPDF.pdf.parser; - -import static org.assertj.core.api.Assertions.assertThat; -import static stirling.software.SPDF.pdf.parser.PdfModels.*; - -import java.util.List; - -import org.junit.jupiter.api.Test; - -/** - * Unit tests for {@link LineAlignmentTableParser}, focused on the coincident-line merge logic and - * column-grid construction. - */ -class LineAlignmentTableParserTest { - - private final LineAlignmentTableParser parser = new LineAlignmentTableParser(); - - // ── mergeCoincidentLines ───────────────────────────────────────────────────────────────────── - - @Test - void mergeCoincidentLines_singleLine_unchanged() { - var lines = List.of(tokenized(rawLine(10f, 100f, "Revenue"))); - assertThat(parser.mergeCoincidentLines(lines)).hasSize(1); - } - - @Test - void mergeCoincidentLines_distinctYLines_unchanged() { - // Two lines at different y positions — must NOT be merged. - var lines = - List.of( - tokenized(rawLine(10f, 100f, "Revenue")), - tokenized(rawLine(10f, 115f, "Cost"))); - assertThat(parser.mergeCoincidentLines(lines)).hasSize(2); - } - - @Test - void mergeCoincidentLines_sameY_merged() { - // Simulates a financial-table row split by LineBuilder at the column gap: - // label fragment at x=72 → "Revenue" - // value fragment at x=350 → "1,234" - // Both have y=100. After merge they should form one TokenizedLine. - var label = rawLine(72f, 100f, "Revenue"); - var value = rawLine(350f, 100f, "1,234"); - - var merged = parser.mergeCoincidentLines(List.of(tokenized(label), tokenized(value))); - - assertThat(merged).hasSize(1); - // The merged line should contain tokens from both halves. - var tokens = merged.get(0).all(); - assertThat(tokens.stream().map(t -> t.text()).toList()) - .containsExactlyInAnyOrder("Revenue", "1,234"); - } - - @Test - void mergeCoincidentLines_sameY_mergedLineHasCorrectBounds() { - var label = rawLine(72f, 100f, "Revenue"); // 7 chars × 6pt = 42pt wide → right = 114 - var value = rawLine(350f, 100f, "1,234"); // 5 chars × 6pt = 30pt wide → right = 380 - - var merged = parser.mergeCoincidentLines(List.of(tokenized(label), tokenized(value))); - - var bounds = merged.get(0).line().bounds(); - assertThat(bounds.x()).isEqualTo(72f); - assertThat(bounds.right()).isEqualTo(380f); - } - - @Test - void mergeCoincidentLines_withinTolerance_merged() { - // Lines 1.5pt apart (within ROW_MERGE_TOLERANCE_PT = 2pt) should merge. - var a = rawLine(10f, 100.0f, "Alpha"); - var b = rawLine(200f, 101.5f, "99"); - - var merged = parser.mergeCoincidentLines(List.of(tokenized(a), tokenized(b))); - assertThat(merged).hasSize(1); - } - - @Test - void mergeCoincidentLines_beyondTolerance_notMerged() { - // Lines 3pt apart (beyond ROW_MERGE_TOLERANCE_PT = 2pt) should NOT merge. - var a = rawLine(10f, 100.0f, "Alpha"); - var b = rawLine(200f, 103.0f, "99"); - - var merged = parser.mergeCoincidentLines(List.of(tokenized(a), tokenized(b))); - assertThat(merged).hasSize(2); - } - - @Test - void mergeCoincidentLines_threeCoincident_allMerged() { - // Three fragments at the same y (e.g. wide financial table with two value columns). - var a = rawLine(72f, 100f, "Revenue"); - var b = rawLine(300f, 100f, "1,234"); - var c = rawLine(400f, 100f, "5,678"); - - var merged = parser.mergeCoincidentLines(List.of(tokenized(a), tokenized(b), tokenized(c))); - assertThat(merged).hasSize(1); - assertThat(merged.get(0).all()).hasSize(3); - } - - @Test - void mergeCoincidentLines_coincidentPairFollowedByDistinctLine_twoGroups() { - var a = rawLine(72f, 100f, "Revenue"); - var b = rawLine(350f, 100f, "1,234"); // same y as a → merges with a - var c = rawLine(10f, 115f, "Expenses"); // different y → stays separate - - var merged = parser.mergeCoincidentLines(List.of(tokenized(a), tokenized(b), tokenized(c))); - assertThat(merged).hasSize(2); - } - - @Test - void mergeCoincidentLines_numericAnchorStatus_correctAfterMerge() { - // After merging, the combined line should be an anchor (≥2 numeric tokens). - // "Revenue" alone → not an anchor. "1,234 567" alone → anchor. - // Merged → anchor with at least 2 numerics. - var label = rawLine(72f, 100f, "Revenue"); - var values = rawLineMultiWord(350f, 100f, "1,234", 30f, "567", 30f); - - var merged = parser.mergeCoincidentLines(List.of(tokenized(label), tokenized(values))); - - assertThat(merged).hasSize(1); - assertThat(merged.get(0).isAnchor()).isTrue(); - } - - // ── helpers ────────────────────────────────────────────────────────────────────────────────── - - /** Creates a RawLine with a single TextFragment of the given text at the given position. */ - private static RawLine rawLine(float x, float y, String text) { - float width = text.length() * 6f; // ~6pt per char — rough but consistent - float height = 12f; - Bounds bounds = new Bounds(x, y, width, height); - TextFragment fragment = - new TextFragment("tf-test", text, bounds, y + height, 11f, "Helvetica", false); - return new RawLine("ln-test", List.of(fragment), bounds, 1); - } - - /** - * Creates a RawLine with two TextFragments representing two words separated by a small gap. - * Used to simulate a values-only line with multiple numeric tokens. - */ - private static RawLine rawLineMultiWord( - float x, float y, String word1, float w1, String word2, float w2) { - float height = 12f; - Bounds b1 = new Bounds(x, y, w1, height); - Bounds b2 = new Bounds(x + w1 + 5f, y, w2, height); - TextFragment f1 = new TextFragment("tf-1", word1, b1, y + height, 11f, "Helvetica", false); - TextFragment f2 = new TextFragment("tf-2", word2, b2, y + height, 11f, "Helvetica", false); - Bounds lineBounds = new Bounds(x, y, x + w1 + 5f + w2 - x, height); - return new RawLine("ln-test", List.of(f1, f2), lineBounds, 1); - } - - /** Tokenises a RawLine via the parser's own tokenise logic (package-private access). */ - private LineAlignmentTableParser.TokenizedLine tokenized(RawLine line) { - return parser.tokenize(line); - } -} diff --git a/app/common/src/test/java/stirling/software/SPDF/pdf/parser/TabulaTableParserGapTest.java b/app/common/src/test/java/stirling/software/SPDF/pdf/parser/TabulaTableParserGapTest.java new file mode 100644 index 0000000000..8e92c2fe5c --- /dev/null +++ b/app/common/src/test/java/stirling/software/SPDF/pdf/parser/TabulaTableParserGapTest.java @@ -0,0 +1,345 @@ +package stirling.software.SPDF.pdf.parser; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertSame; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static stirling.software.SPDF.pdf.parser.PdfModels.RawPage; +import static stirling.software.SPDF.pdf.parser.PdfModels.TableCell; +import static stirling.software.SPDF.pdf.parser.PdfModels.TableFragment; +import static stirling.software.SPDF.pdf.parser.PdfModels.TableRow; + +import java.awt.Color; +import java.io.ByteArrayOutputStream; +import java.util.List; + +import org.apache.pdfbox.Loader; +import org.apache.pdfbox.pdmodel.PDDocument; +import org.apache.pdfbox.pdmodel.PDPage; +import org.apache.pdfbox.pdmodel.PDPageContentStream; +import org.apache.pdfbox.pdmodel.common.PDRectangle; +import org.apache.pdfbox.pdmodel.font.PDType1Font; +import org.apache.pdfbox.pdmodel.font.Standard14Fonts; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; + +/** + * Unit tests for {@link TabulaTableParser}. Tables are built in-memory with PDFBox so the tests are + * deterministic and need no fixtures, network, or external processes. + */ +class TabulaTableParserGapTest { + + private final TabulaTableParser parser = new TabulaTableParser(); + + // ── error / empty branches ─────────────────────────────────────────────── + + @Nested + @DisplayName("Empty and error branches") + class EmptyAndErrorBranches { + + @Test + @DisplayName("page number 0 is out of Tabula's 1-based range -> empty list, no throw") + void pageNumberZeroReturnsEmpty() throws Exception { + byte[] pdf = pdfWithText(new String[] {"hello"}); + try (PDDocument doc = Loader.loadPDF(pdf)) { + List result = parser.parse(doc, 0); + assertNotNull(result); + assertTrue(result.isEmpty()); + } + } + + @Test + @DisplayName("page number beyond the document -> empty list, exception swallowed") + void pageNumberOutOfRangeReturnsEmpty() throws Exception { + byte[] pdf = pdfWithText(new String[] {"hello"}); + try (PDDocument doc = Loader.loadPDF(pdf)) { + List result = parser.parse(doc, 99); + assertNotNull(result); + assertTrue(result.isEmpty()); + } + } + + @Test + @DisplayName("negative page number -> empty list") + void negativePageNumberReturnsEmpty() throws Exception { + byte[] pdf = pdfWithText(new String[] {"hello"}); + try (PDDocument doc = Loader.loadPDF(pdf)) { + assertTrue(parser.parse(doc, -5).isEmpty()); + } + } + + @Test + @DisplayName("lattice mode on a page with no ruled lines -> no tables") + void latticeWithNoRulingsReturnsEmpty() throws Exception { + byte[] pdf = pdfWithText(new String[] {"just some prose", "no table here"}); + try (PDDocument doc = Loader.loadPDF(pdf)) { + List result = parser.parse(doc, new RawPage(1, 0f, 0f, List.of())); + assertNotNull(result); + assertTrue( + result.isEmpty(), "borderless text must not be detected in lattice mode"); + } + } + + @Test + @DisplayName("blank page in lattice mode -> empty list") + void blankPageLatticeReturnsEmpty() throws Exception { + byte[] pdf = blankPdf(); + try (PDDocument doc = Loader.loadPDF(pdf)) { + assertTrue(parser.parse(doc, new RawPage(1, 0f, 0f, List.of())).isEmpty()); + } + } + } + + // ── stream mode (BasicExtractionAlgorithm) ─────────────────────────────── + + @Nested + @DisplayName("Stream mode") + class StreamMode { + + @Test + @DisplayName("page with text yields at least one well-formed fragment") + void streamOnTextProducesFragment() throws Exception { + byte[] pdf = + pdfWithText(new String[] {"Name Age City", "Alice 30 Paris", "Bob 25 Rome"}); + try (PDDocument doc = Loader.loadPDF(pdf)) { + List fragments = + parser.parseStream(doc, new RawPage(1, 0f, 0f, List.of())); + assertNotNull(fragments); + assertFalse(fragments.isEmpty(), "stream mode always builds a table from text"); + assertFragmentWellFormed(fragments.get(0), 1, 0); + } + } + + @Test + @DisplayName("fragment ids encode page and index") + void streamFragmentIdFormat() throws Exception { + byte[] pdf = pdfWithText(new String[] {"col1 col2", "a b"}); + try (PDDocument doc = Loader.loadPDF(pdf)) { + List fragments = + parser.parseStream(doc, new RawPage(1, 0f, 0f, List.of())); + assertFalse(fragments.isEmpty()); + assertEquals("tbl-p1-0", fragments.get(0).tableId()); + assertEquals(1, fragments.get(0).pageNumber()); + } + } + + @Test + @DisplayName("rawRows and the parsed rows stay in lockstep") + void streamRowsMatchRawRows() throws Exception { + byte[] pdf = pdfWithText(new String[] {"x y", "1 2", "3 4"}); + try (PDDocument doc = Loader.loadPDF(pdf)) { + List fragments = + parser.parseStream(doc, new RawPage(1, 0f, 0f, List.of())); + assertFalse(fragments.isEmpty()); + TableFragment f = fragments.get(0); + assertEquals(f.rawRows().size(), f.rows().size()); + } + } + } + + // ── lattice mode with a real bordered grid ─────────────────────────────── + + @Nested + @DisplayName("Lattice mode") + class LatticeMode { + + @Test + @DisplayName("bordered grid is detected and produces well-formed fragments") + void latticeDetectsBorderedTable() throws Exception { + byte[] pdf = pdfWithGrid(); + try (PDDocument doc = Loader.loadPDF(pdf)) { + List fragments = + parser.parse(doc, new RawPage(1, 0f, 0f, List.of())); + assertNotNull(fragments); + assertFalse( + fragments.isEmpty(), "a clean ruled grid must be detected in lattice mode"); + TableFragment f = fragments.get(0); + assertFragmentWellFormed(f, 1, 0); + assertTrue(f.columnCount() >= 1, "a detected grid must have at least one column"); + assertFalse(f.rawRows().isEmpty(), "a detected grid must have rows"); + } + } + + @Test + @DisplayName("convenience overload with page number routes to lattice mode") + void parseByPageNumberDetectsGrid() throws Exception { + byte[] pdf = pdfWithGrid(); + try (PDDocument doc = Loader.loadPDF(pdf)) { + List fragments = parser.parse(doc, 1); + assertNotNull(fragments); + assertFalse(fragments.isEmpty()); + assertEquals(1, fragments.get(0).pageNumber()); + } + } + + @Test + @DisplayName("cell text is normalised (trimmed, newlines collapsed)") + void latticeCellTextIsNormalised() throws Exception { + byte[] pdf = pdfWithGrid(); + try (PDDocument doc = Loader.loadPDF(pdf)) { + List fragments = + parser.parse(doc, new RawPage(1, 0f, 0f, List.of())); + assertFalse(fragments.isEmpty()); + for (List row : fragments.get(0).rawRows()) { + for (String cell : row) { + assertNotNull(cell); + assertFalse(cell.contains("\n"), "newlines must be collapsed"); + assertFalse(cell.contains("\r"), "carriage returns must be collapsed"); + assertEquals(cell.trim(), cell, "cell text must be trimmed"); + } + } + } + } + } + + // ── contract invariants ────────────────────────────────────────────────── + + @Nested + @DisplayName("Contract invariants") + class ContractInvariants { + + @Test + @DisplayName("parse never returns null") + void parseNeverReturnsNull() throws Exception { + byte[] pdf = pdfWithText(new String[] {"abc"}); + try (PDDocument doc = Loader.loadPDF(pdf)) { + assertNotNull(parser.parse(doc, new RawPage(1, 0f, 0f, List.of()))); + assertNotNull(parser.parse(doc, 1)); + assertNotNull(parser.parseStream(doc, new RawPage(1, 0f, 0f, List.of()))); + } + } + + @Test + @DisplayName("the document is not closed by the parser") + void documentRemainsOpenAfterParse() throws Exception { + byte[] pdf = pdfWithText(new String[] {"keep me open"}); + try (PDDocument doc = Loader.loadPDF(pdf)) { + parser.parse(doc, new RawPage(1, 0f, 0f, List.of())); + parser.parseStream(doc, new RawPage(1, 0f, 0f, List.of())); + // ObjectExtractor.close() would close the underlying COSDocument; the parser must + // not. + assertFalse( + doc.getDocument().isClosed(), + "parser must not close the caller's document"); + assertEquals(1, doc.getNumberOfPages()); + } + } + } + + // ── helpers ────────────────────────────────────────────────────────────── + + /** Asserts every field of a fragment satisfies the documented contract. */ + private static void assertFragmentWellFormed( + TableFragment f, int expectedPage, int expectedIndex) { + assertNotNull(f); + assertEquals(expectedPage, f.pageNumber()); + assertEquals("tbl-p" + expectedPage + "-" + expectedIndex, f.tableId()); + assertNotNull(f.bounds()); + assertNotNull(f.headers()); + assertTrue(f.headers().isEmpty(), "headers are deferred to v2 and must be empty"); + assertNotNull(f.rows()); + assertNotNull(f.rawRows()); + assertNotNull(f.warnings()); + assertSame(null, f.continuedFromPage(), "continuedFromPage is deferred to v2"); + assertTrue(f.columnCount() >= 0); + assertTrue(f.confidence() >= 0f && f.confidence() <= 1f, "confidence must be within [0,1]"); + assertEquals(f.rawRows().size(), f.rows().size()); + + for (TableRow row : f.rows()) { + assertNotNull(row.cells()); + for (TableCell cell : row.cells()) { + assertNotNull(cell.text()); + assertNotNull(cell.bounds()); + assertEquals(1, cell.colSpan(), "colSpan is always 1 in v1"); + assertEquals(1, cell.rowSpan(), "rowSpan is always 1 in v1"); + } + } + } + + private static byte[] pdfWithText(String[] lines) throws Exception { + try (PDDocument doc = new PDDocument()) { + PDPage page = new PDPage(PDRectangle.A4); + doc.addPage(page); + try (PDPageContentStream cs = new PDPageContentStream(doc, page)) { + cs.setFont(new PDType1Font(Standard14Fonts.FontName.HELVETICA), 12); + cs.setNonStrokingColor(Color.BLACK); + float y = 720f; + for (String line : lines) { + cs.beginText(); + cs.newLineAtOffset(72f, y); + cs.showText(line); + cs.endText(); + y -= 20f; + } + } + return save(doc); + } + } + + private static byte[] blankPdf() throws Exception { + try (PDDocument doc = new PDDocument()) { + doc.addPage(new PDPage(PDRectangle.A4)); + return save(doc); + } + } + + /** + * Builds a small 3-row x 3-column ruled grid with text in each cell. The ruled lines make the + * table detectable by lattice mode. + */ + private static byte[] pdfWithGrid() throws Exception { + try (PDDocument doc = new PDDocument()) { + PDPage page = new PDPage(PDRectangle.A4); + doc.addPage(page); + + float left = 100f; + float right = 400f; + float top = 700f; + float bottom = 550f; + int cols = 3; + int rows = 3; + float colStep = (right - left) / cols; + float rowStep = (top - bottom) / rows; + + try (PDPageContentStream cs = new PDPageContentStream(doc, page)) { + cs.setStrokingColor(Color.BLACK); + cs.setLineWidth(1f); + + // vertical lines + for (int c = 0; c <= cols; c++) { + float x = left + c * colStep; + cs.moveTo(x, bottom); + cs.lineTo(x, top); + } + // horizontal lines + for (int r = 0; r <= rows; r++) { + float yLine = bottom + r * rowStep; + cs.moveTo(left, yLine); + cs.lineTo(right, yLine); + } + cs.stroke(); + + // cell text + cs.setFont(new PDType1Font(Standard14Fonts.FontName.HELVETICA), 10); + cs.setNonStrokingColor(Color.BLACK); + for (int r = 0; r < rows; r++) { + for (int c = 0; c < cols; c++) { + cs.beginText(); + cs.newLineAtOffset(left + c * colStep + 5f, top - (r + 1) * rowStep + 6f); + cs.showText("R" + r + "C" + c); + cs.endText(); + } + } + } + return save(doc); + } + } + + private static byte[] save(PDDocument doc) throws Exception { + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + doc.save(baos); + return baos.toByteArray(); + } +} diff --git a/app/common/src/test/java/stirling/software/common/cluster/inprocess/LocalDiskFileStoreTest.java b/app/common/src/test/java/stirling/software/common/cluster/inprocess/LocalDiskFileStoreTest.java index 296d012c8f..df8516851d 100644 --- a/app/common/src/test/java/stirling/software/common/cluster/inprocess/LocalDiskFileStoreTest.java +++ b/app/common/src/test/java/stirling/software/common/cluster/inprocess/LocalDiskFileStoreTest.java @@ -3,11 +3,13 @@ package stirling.software.common.cluster.inprocess; import static org.junit.jupiter.api.Assertions.assertArrayEquals; import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNull; import static org.junit.jupiter.api.Assertions.assertThrows; import static org.junit.jupiter.api.Assertions.assertTrue; import java.io.ByteArrayInputStream; import java.io.IOException; +import java.nio.file.Files; import java.nio.file.Path; import org.junit.jupiter.api.Test; @@ -39,4 +41,46 @@ class LocalDiskFileStoreTest { assertThrows(IllegalArgumentException.class, () -> store.resolve("a/b")); assertThrows(IllegalArgumentException.class, () -> store.resolve("a\\b")); } + + @Test + void ownerSidecarCannotBeReadAsFileId(@TempDir Path dir) throws IOException { + LocalDiskFileStore store = new LocalDiskFileStore(dir.toString()); + FileStore.Stored stored = + store.store(new ByteArrayInputStream("hi".getBytes()), "f.bin", "alice"); + String sidecarId = stored.fileId() + ".owner"; + assertThrows(IllegalArgumentException.class, () -> store.resolve(sidecarId)); + assertThrows(IllegalArgumentException.class, () -> store.retrieveBytes(sidecarId)); + } + + @Test + void ownerIsPersistedAndReturnedByGetOwner(@TempDir Path dir) throws IOException { + LocalDiskFileStore store = new LocalDiskFileStore(dir.toString()); + FileStore.Stored stored = + store.store(new ByteArrayInputStream("hi".getBytes()), "f.bin", "alice"); + assertEquals("alice", store.getOwner(stored.fileId())); + } + + @Test + void getOwnerReturnsNullWhenNoOwnerWasRecorded(@TempDir Path dir) throws IOException { + LocalDiskFileStore store = new LocalDiskFileStore(dir.toString()); + FileStore.Stored stored = + store.store(new ByteArrayInputStream("hi".getBytes()), "f.bin", null); + assertNull(store.getOwner(stored.fileId())); + } + + @Test + void getOwnerReturnsNullForUnknownFileId(@TempDir Path dir) throws IOException { + LocalDiskFileStore store = new LocalDiskFileStore(dir.toString()); + assertNull(store.getOwner("00000000-0000-0000-0000-000000000000")); + } + + @Test + void deleteRemovesOwnerSidecar(@TempDir Path dir) throws IOException { + LocalDiskFileStore store = new LocalDiskFileStore(dir.toString()); + FileStore.Stored stored = + store.store(new ByteArrayInputStream("hi".getBytes()), "f.bin", "alice"); + assertTrue(store.delete(stored.fileId())); + assertFalse(Files.exists(dir.resolve(stored.fileId() + ".owner"))); + assertNull(store.getOwner(stored.fileId())); + } } diff --git a/app/common/src/test/java/stirling/software/common/configuration/RuntimePathConfigTest.java b/app/common/src/test/java/stirling/software/common/configuration/RuntimePathConfigTest.java new file mode 100644 index 0000000000..f3e0490b14 --- /dev/null +++ b/app/common/src/test/java/stirling/software/common/configuration/RuntimePathConfigTest.java @@ -0,0 +1,520 @@ +package stirling.software.common.configuration; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertSame; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.nio.file.Path; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.List; + +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.common.model.ApplicationProperties.CustomPaths.Operations; +import stirling.software.common.model.ApplicationProperties.CustomPaths.Pipeline; +import stirling.software.common.model.ApplicationProperties.ProcessExecutor.UnoServerEndpoint; + +/** + * Unit tests for {@link RuntimePathConfig}. All of the resolution logic lives in the constructor, + * so each test builds a real {@link ApplicationProperties} (a plain @Data POJO with sensible + * defaults), constructs the config, and asserts on the exposed getters. + */ +class RuntimePathConfigTest { + + /** The base path the production code derives from {@link InstallationPathConfig#getPath()}. */ + private static final String BASE_PATH = InstallationPathConfig.getPath(); + + private static ApplicationProperties newProperties() { + return new ApplicationProperties(); + } + + private static RuntimePathConfig build(ApplicationProperties properties) { + return new RuntimePathConfig(properties); + } + + @Nested + @DisplayName("Pipeline directory resolution") + class PipelinePaths { + + @Test + @DisplayName("Defaults to /pipeline and derived sub-folders") + void defaultPipelinePaths() { + RuntimePathConfig config = build(newProperties()); + + String expectedPipeline = Path.of(BASE_PATH, "pipeline").toString(); + assertEquals(expectedPipeline, config.getPipelinePath()); + // Watched folders are resolved to an absolute, normalized path by the production code. + assertEquals( + Path.of(expectedPipeline, "watchedFolders") + .toAbsolutePath() + .normalize() + .toString(), + config.getPipelineWatchedFoldersPath()); + assertEquals( + Path.of(expectedPipeline, "finishedFolders").toString(), + config.getPipelineFinishedFoldersPath()); + assertEquals( + Path.of(expectedPipeline, "defaultWebUIConfigs").toString(), + config.getPipelineDefaultWebUiConfigs()); + } + + @Test + @DisplayName("Custom pipelineDir overrides the default pipeline path") + void customPipelineDir() { + ApplicationProperties properties = newProperties(); + Pipeline pipeline = properties.getSystem().getCustomPaths().getPipeline(); + pipeline.setPipelineDir("/custom/pipeline"); + + RuntimePathConfig config = build(properties); + + assertEquals("/custom/pipeline", config.getPipelinePath()); + // Sub-folders are derived from the (already-resolved) custom pipeline path. + assertEquals( + Path.of("/custom/pipeline", "finishedFolders").toString(), + config.getPipelineFinishedFoldersPath()); + assertEquals( + Path.of("/custom/pipeline", "defaultWebUIConfigs").toString(), + config.getPipelineDefaultWebUiConfigs()); + } + + @Test + @DisplayName("Blank pipelineDir falls back to the default") + void blankPipelineDirFallsBackToDefault() { + ApplicationProperties properties = newProperties(); + properties.getSystem().getCustomPaths().getPipeline().setPipelineDir(" "); + + RuntimePathConfig config = build(properties); + + assertEquals(Path.of(BASE_PATH, "pipeline").toString(), config.getPipelinePath()); + } + + @Test + @DisplayName("Custom finished and webUI configs dirs override defaults") + void customFinishedAndWebUiDirs() { + ApplicationProperties properties = newProperties(); + Pipeline pipeline = properties.getSystem().getCustomPaths().getPipeline(); + pipeline.setFinishedFoldersDir("/custom/finished"); + pipeline.setWebUIConfigsDir("/custom/webui"); + + RuntimePathConfig config = build(properties); + + assertEquals("/custom/finished", config.getPipelineFinishedFoldersPath()); + assertEquals("/custom/webui", config.getPipelineDefaultWebUiConfigs()); + } + } + + @Nested + @DisplayName("Watched folder resolution") + class WatchedFolders { + + @Test + @DisplayName("Default watched folder is /watchedFolders and list has one entry") + void defaultWatchedFolder() { + RuntimePathConfig config = build(newProperties()); + + // Watched folders are resolved to an absolute, normalized path by the production code. + String expected = + Path.of(Path.of(BASE_PATH, "pipeline").toString(), "watchedFolders") + .toAbsolutePath() + .normalize() + .toString(); + assertEquals(expected, config.getPipelineWatchedFoldersPath()); + assertEquals(1, config.getPipelineWatchedFoldersPaths().size()); + assertEquals(expected, config.getPipelineWatchedFoldersPaths().get(0)); + } + + @Test + @DisplayName("Legacy single watchedFoldersDir is used when no list is provided") + void legacyWatchedFolder() { + ApplicationProperties properties = newProperties(); + properties + .getSystem() + .getCustomPaths() + .getPipeline() + .setWatchedFoldersDir("relativeWatched"); + + RuntimePathConfig config = build(properties); + + // Legacy paths are normalized to absolute. + String expected = Path.of("relativeWatched").toAbsolutePath().normalize().toString(); + assertEquals(1, config.getPipelineWatchedFoldersPaths().size()); + assertEquals(expected, config.getPipelineWatchedFoldersPath()); + } + + @Test + @DisplayName("New list config takes precedence over the legacy single dir") + void listTakesPrecedenceOverLegacy() { + ApplicationProperties properties = newProperties(); + Pipeline pipeline = properties.getSystem().getCustomPaths().getPipeline(); + pipeline.setWatchedFoldersDir("legacyDir"); + pipeline.setWatchedFoldersDirs(new ArrayList<>(Arrays.asList("listDirA", "listDirB"))); + + RuntimePathConfig config = build(properties); + + List paths = config.getPipelineWatchedFoldersPaths(); + assertEquals(2, paths.size()); + assertEquals(Path.of("listDirA").toAbsolutePath().normalize().toString(), paths.get(0)); + assertEquals(Path.of("listDirB").toAbsolutePath().normalize().toString(), paths.get(1)); + // The legacy value must NOT appear when the list is present. + assertFalse( + paths.contains(Path.of("legacyDir").toAbsolutePath().normalize().toString())); + } + + @Test + @DisplayName("Duplicate paths in the list are de-duplicated after normalization") + void duplicatePathsAreDeduplicated() { + ApplicationProperties properties = newProperties(); + properties + .getSystem() + .getCustomPaths() + .getPipeline() + .setWatchedFoldersDirs( + new ArrayList<>(Arrays.asList("dupDir", "dupDir", "otherDir"))); + + RuntimePathConfig config = build(properties); + + List paths = config.getPipelineWatchedFoldersPaths(); + assertEquals(2, paths.size()); + assertEquals(Path.of("dupDir").toAbsolutePath().normalize().toString(), paths.get(0)); + assertEquals(Path.of("otherDir").toAbsolutePath().normalize().toString(), paths.get(1)); + } + + @Test + @DisplayName("Blank and whitespace-only list entries are sanitized out") + void blankListEntriesAreFiltered() { + ApplicationProperties properties = newProperties(); + properties + .getSystem() + .getCustomPaths() + .getPipeline() + .setWatchedFoldersDirs( + new ArrayList<>(Arrays.asList(" ", "", "validDir", " "))); + + RuntimePathConfig config = build(properties); + + List paths = config.getPipelineWatchedFoldersPaths(); + assertEquals(1, paths.size()); + assertEquals(Path.of("validDir").toAbsolutePath().normalize().toString(), paths.get(0)); + } + + @Test + @DisplayName("List entries are trimmed before resolution") + void listEntriesAreTrimmed() { + ApplicationProperties properties = newProperties(); + properties + .getSystem() + .getCustomPaths() + .getPipeline() + .setWatchedFoldersDirs(new ArrayList<>(Arrays.asList(" spacedDir "))); + + RuntimePathConfig config = build(properties); + + assertEquals( + Path.of("spacedDir").toAbsolutePath().normalize().toString(), + config.getPipelineWatchedFoldersPath()); + } + + @Test + @DisplayName("An all-blank list falls back to the legacy dir, then default") + void allBlankListFallsBackToDefault() { + ApplicationProperties properties = newProperties(); + properties + .getSystem() + .getCustomPaths() + .getPipeline() + .setWatchedFoldersDirs(new ArrayList<>(Arrays.asList("", " "))); + + RuntimePathConfig config = build(properties); + + // sanitizePathList strips everything -> empty -> falls through to default watched + // folder. + // The default is also resolved to an absolute, normalized path by the production code. + String expectedDefault = + Path.of(Path.of(BASE_PATH, "pipeline").toString(), "watchedFolders") + .toAbsolutePath() + .normalize() + .toString(); + assertEquals(1, config.getPipelineWatchedFoldersPaths().size()); + assertEquals(expectedDefault, config.getPipelineWatchedFoldersPath()); + } + + @Test + @DisplayName("First watched folder path is always exposed via the singular getter") + void singularGetterReturnsFirstEntry() { + ApplicationProperties properties = newProperties(); + properties + .getSystem() + .getCustomPaths() + .getPipeline() + .setWatchedFoldersDirs(new ArrayList<>(Arrays.asList("firstDir", "secondDir"))); + + RuntimePathConfig config = build(properties); + + assertEquals( + config.getPipelineWatchedFoldersPaths().get(0), + config.getPipelineWatchedFoldersPath()); + assertEquals( + Path.of("firstDir").toAbsolutePath().normalize().toString(), + config.getPipelineWatchedFoldersPath()); + } + } + + @Nested + @DisplayName("Operation tool path resolution") + class OperationPaths { + + @Test + @DisplayName("Defaults to bare command names when not running in Docker") + void defaultOperationPaths() { + // The test host has no /.dockerenv, so the non-docker defaults apply. + RuntimePathConfig config = build(newProperties()); + + assertEquals("weasyprint", config.getWeasyPrintPath()); + assertEquals("unoconvert", config.getUnoConvertPath()); + assertEquals("ebook-convert", config.getCalibrePath()); + assertEquals("ocrmypdf", config.getOcrMyPdfPath()); + assertEquals("soffice", config.getSOfficePath()); + } + + @Test + @DisplayName("Custom operation paths override the defaults") + void customOperationPaths() { + ApplicationProperties properties = newProperties(); + Operations operations = properties.getSystem().getCustomPaths().getOperations(); + operations.setWeasyprint("/opt/custom/weasyprint"); + operations.setUnoconvert("/opt/custom/unoconvert"); + operations.setCalibre("/opt/custom/ebook-convert"); + operations.setOcrmypdf("/opt/custom/ocrmypdf"); + operations.setSoffice("/opt/custom/soffice"); + + RuntimePathConfig config = build(properties); + + assertEquals("/opt/custom/weasyprint", config.getWeasyPrintPath()); + assertEquals("/opt/custom/unoconvert", config.getUnoConvertPath()); + assertEquals("/opt/custom/ebook-convert", config.getCalibrePath()); + assertEquals("/opt/custom/ocrmypdf", config.getOcrMyPdfPath()); + assertEquals("/opt/custom/soffice", config.getSOfficePath()); + } + + @Test + @DisplayName("Blank custom operation path falls back to the default") + void blankOperationPathFallsBack() { + ApplicationProperties properties = newProperties(); + properties.getSystem().getCustomPaths().getOperations().setWeasyprint(" "); + + RuntimePathConfig config = build(properties); + + assertEquals("weasyprint", config.getWeasyPrintPath()); + } + + @Test + @DisplayName("A single custom path leaves the other operation paths at defaults") + void partialOperationOverride() { + ApplicationProperties properties = newProperties(); + properties + .getSystem() + .getCustomPaths() + .getOperations() + .setSoffice("/usr/local/soffice"); + + RuntimePathConfig config = build(properties); + + assertEquals("/usr/local/soffice", config.getSOfficePath()); + assertEquals("weasyprint", config.getWeasyPrintPath()); + assertEquals("unoconvert", config.getUnoConvertPath()); + } + } + + @Nested + @DisplayName("Tesseract data path resolution") + class TessdataPath { + + @Test + @DisplayName("Explicit tessdataDir config wins over env var and default") + void configuredTessdataDirWins() { + ApplicationProperties properties = newProperties(); + properties.getSystem().setTessdataDir("/my/tessdata"); + + RuntimePathConfig config = build(properties); + + // Config setting has the highest priority regardless of TESSDATA_PREFIX env state. + assertEquals("/my/tessdata", config.getTessDataPath()); + } + + @Test + @DisplayName("tessDataPath is never null even with no config") + void tessDataPathNeverNull() { + RuntimePathConfig config = build(newProperties()); + + // With no config setting, the value comes from TESSDATA_PREFIX or the hard default, + // either of which is non-null. + assertNotNull(config.getTessDataPath()); + assertFalse(config.getTessDataPath().isEmpty()); + } + } + + @Nested + @DisplayName("UNO server endpoint resolution") + class UnoServerEndpoints { + + @Test + @DisplayName("Auto mode builds one endpoint when session limit is unset (defaults to 1)") + void autoSingleEndpointByDefault() { + // Default ApplicationProperties: autoUnoServer = true, libreOfficeSessionLimit = 0 -> + // 1. + RuntimePathConfig config = build(newProperties()); + + List endpoints = config.getUnoServerEndpoints(); + assertEquals(1, endpoints.size()); + assertEquals("127.0.0.1", endpoints.get(0).getHost()); + assertEquals(2003, endpoints.get(0).getPort()); + } + + @Test + @DisplayName("Auto mode builds N endpoints on consecutive even ports") + void autoMultipleEndpoints() { + ApplicationProperties properties = newProperties(); + properties.getProcessExecutor().getSessionLimit().setLibreOfficeSessionLimit(3); + + RuntimePathConfig config = build(properties); + + List endpoints = config.getUnoServerEndpoints(); + assertEquals(3, endpoints.size()); + assertEquals(2003, endpoints.get(0).getPort()); + assertEquals(2005, endpoints.get(1).getPort()); + assertEquals(2007, endpoints.get(2).getPort()); + for (UnoServerEndpoint endpoint : endpoints) { + assertEquals("127.0.0.1", endpoint.getHost()); + } + } + + @Test + @DisplayName("Manual mode returns the configured (valid) endpoints") + void manualEndpointsAreUsed() { + ApplicationProperties properties = newProperties(); + ApplicationProperties.ProcessExecutor processExecutor = properties.getProcessExecutor(); + processExecutor.setAutoUnoServer(false); + + UnoServerEndpoint endpoint = new UnoServerEndpoint(); + endpoint.setHost("10.0.0.5"); + endpoint.setPort(4000); + processExecutor.setUnoServerEndpoints(new ArrayList<>(Arrays.asList(endpoint))); + + RuntimePathConfig config = build(properties); + + List endpoints = config.getUnoServerEndpoints(); + assertEquals(1, endpoints.size()); + assertEquals("10.0.0.5", endpoints.get(0).getHost()); + assertEquals(4000, endpoints.get(0).getPort()); + } + + @Test + @DisplayName("Manual mode filters out endpoints with blank host or non-positive port") + void manualEndpointsAreSanitized() { + ApplicationProperties properties = newProperties(); + ApplicationProperties.ProcessExecutor processExecutor = properties.getProcessExecutor(); + processExecutor.setAutoUnoServer(false); + + UnoServerEndpoint valid = new UnoServerEndpoint(); + valid.setHost("192.168.1.10"); + valid.setPort(5000); + + UnoServerEndpoint blankHost = new UnoServerEndpoint(); + blankHost.setHost(" "); + blankHost.setPort(5001); + + UnoServerEndpoint badPort = new UnoServerEndpoint(); + badPort.setHost("192.168.1.11"); + badPort.setPort(0); + + processExecutor.setUnoServerEndpoints( + new ArrayList<>(Arrays.asList(valid, blankHost, badPort))); + + RuntimePathConfig config = build(properties); + + List endpoints = config.getUnoServerEndpoints(); + assertEquals(1, endpoints.size()); + assertEquals("192.168.1.10", endpoints.get(0).getHost()); + assertEquals(5000, endpoints.get(0).getPort()); + } + + @Test + @DisplayName("Manual mode with no usable endpoints falls back to a single default endpoint") + void manualModeNoEndpointsFallsBackToDefault() { + ApplicationProperties properties = newProperties(); + ApplicationProperties.ProcessExecutor processExecutor = properties.getProcessExecutor(); + processExecutor.setAutoUnoServer(false); + processExecutor.setUnoServerEndpoints(new ArrayList<>()); + + RuntimePathConfig config = build(properties); + + List endpoints = config.getUnoServerEndpoints(); + assertEquals(1, endpoints.size()); + assertEquals("127.0.0.1", endpoints.get(0).getHost()); + assertEquals(2003, endpoints.get(0).getPort()); + } + + @Test + @DisplayName("Null processExecutor defaults to a single UNO endpoint") + void nullProcessExecutorDefaultsToSingleEndpoint() { + ApplicationProperties properties = newProperties(); + properties.setProcessExecutor(null); + + RuntimePathConfig config = build(properties); + + List endpoints = config.getUnoServerEndpoints(); + assertEquals(1, endpoints.size()); + assertEquals("127.0.0.1", endpoints.get(0).getHost()); + assertEquals(2003, endpoints.get(0).getPort()); + } + } + + @Nested + @DisplayName("General contract") + class GeneralContract { + + @Test + @DisplayName("getProperties returns the same instance passed to the constructor") + void propertiesAccessorReturnsSameInstance() { + ApplicationProperties properties = newProperties(); + + RuntimePathConfig config = build(properties); + + assertSame(properties, config.getProperties()); + } + + @Test + @DisplayName("basePath matches InstallationPathConfig.getPath()") + void basePathMatchesInstallationPath() { + RuntimePathConfig config = build(newProperties()); + + assertEquals(BASE_PATH, config.getBasePath()); + } + + @Test + @DisplayName("All resolved path getters are non-null") + void allPathsNonNull() { + RuntimePathConfig config = build(newProperties()); + + assertNotNull(config.getPipelinePath()); + assertNotNull(config.getPipelineWatchedFoldersPath()); + assertNotNull(config.getPipelineWatchedFoldersPaths()); + assertNotNull(config.getPipelineFinishedFoldersPath()); + assertNotNull(config.getPipelineDefaultWebUiConfigs()); + assertNotNull(config.getWeasyPrintPath()); + assertNotNull(config.getUnoConvertPath()); + assertNotNull(config.getCalibrePath()); + assertNotNull(config.getOcrMyPdfPath()); + assertNotNull(config.getSOfficePath()); + assertNotNull(config.getTessDataPath()); + assertNotNull(config.getUnoServerEndpoints()); + assertTrue(config.getUnoServerEndpoints().size() >= 1); + } + } +} diff --git a/app/common/src/test/java/stirling/software/common/model/ApplicationPropertiesLogicTest.java b/app/common/src/test/java/stirling/software/common/model/ApplicationPropertiesLogicTest.java index f075f15185..98e0e8ca21 100644 --- a/app/common/src/test/java/stirling/software/common/model/ApplicationPropertiesLogicTest.java +++ b/app/common/src/test/java/stirling/software/common/model/ApplicationPropertiesLogicTest.java @@ -2,7 +2,7 @@ package stirling.software.common.model; import static org.junit.jupiter.api.Assertions.*; -import java.nio.file.Paths; +import java.nio.file.Path; import java.util.ArrayList; import java.util.Collection; import java.util.List; @@ -31,18 +31,33 @@ class ApplicationPropertiesLogicTest { assertTrue(sys.isAnalyticsEnabled()); } + @Test + void storageSigning_userListScope_defaultsToOrg_andIsSettable() { + // Self-host backward-compat: scope must default to "org" (saas profile pins "team"). + ApplicationProperties.Storage.Signing signing = new ApplicationProperties.Storage.Signing(); + + assertFalse(signing.isEnabled()); + assertEquals("org", signing.getUserListScope()); + + signing.setUserListScope("team"); + assertEquals("team", signing.getUserListScope()); + + // Reachable from the full tree as storage.signing.userListScope. + assertEquals( + "org", new ApplicationProperties().getStorage().getSigning().getUserListScope()); + } + @Test void tempFileManagement_defaults_and_overrides() { - Function normalize = s -> Paths.get(s).normalize().toString(); + Function normalize = s -> Path.of(s).normalize().toString(); ApplicationProperties.TempFileManagement tfm = new ApplicationProperties.TempFileManagement(); String expectedBase = - Paths.get(java.lang.System.getProperty("java.io.tmpdir"), "stirling-pdf") - .toString(); + Path.of(java.lang.System.getProperty("java.io.tmpdir"), "stirling-pdf").toString(); assertEquals(expectedBase, tfm.getBaseTmpDir()); - String expectedLibre = Paths.get(expectedBase, "libreoffice").toString(); + String expectedLibre = Path.of(expectedBase, "libreoffice").toString(); assertEquals(expectedLibre, tfm.getLibreofficeDir()); tfm.setBaseTmpDir("/custom/base"); diff --git a/app/common/src/test/java/stirling/software/common/pdf/PdfMarkdownConverterTest.java b/app/common/src/test/java/stirling/software/common/pdf/PdfMarkdownConverterTest.java new file mode 100644 index 0000000000..b3c104da85 --- /dev/null +++ b/app/common/src/test/java/stirling/software/common/pdf/PdfMarkdownConverterTest.java @@ -0,0 +1,269 @@ +package stirling.software.common.pdf; + +import static org.junit.jupiter.api.Assertions.assertDoesNotThrow; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.junit.jupiter.api.Assertions.fail; + +import java.io.IOException; +import java.io.InputStream; +import java.nio.charset.StandardCharsets; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.ArrayList; +import java.util.List; +import java.util.stream.Stream; + +import org.junit.jupiter.api.Disabled; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.io.TempDir; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.Arguments; +import org.junit.jupiter.params.provider.MethodSource; + +import stirling.software.jpdfium.PdfDocument; +import stirling.software.jpdfium.text.TextLine; +import stirling.software.jpdfium.text.TextWord; + +/** + * Accuracy and robustness tests for {@link PdfMarkdownConverter}, comparing conversion output + * against hand-authored golden Markdown for a set of owned/synthetic fixtures. + * + *

The {@link #gatedFixtures()} set is enforced in CI: those fixtures currently convert within + * the accuracy threshold and guard against regressions. Fixtures still being iterated on live in + * {@link #wipFixtures()} under a {@link Disabled} test so the goldens stay in the tree without + * breaking the build. Enable the WIP test locally to see per-fixture scores while working on the + * converter. + */ +class PdfMarkdownConverterTest { + + /** Accuracy threshold: output must share at least this fraction of content with the golden. */ + private static final double THRESHOLD = 0.95; + + @TempDir Path tmp; + + /** Fixtures that meet the accuracy threshold today and therefore gate CI. */ + static Stream gatedFixtures() { + return Stream.of( + Arguments.of("multi-column-test_lorem.pdf", "multi-column-test_lorem.md"), + Arguments.of("bordered-table-test_widget.pdf", "bordered-table-test_widget.md"), + Arguments.of("many-tables-test_stress.pdf", "many-tables-test_stress.md")); + } + + /** Fixtures still below the threshold; tracked here, enable locally to iterate. */ + static Stream wipFixtures() { + return Stream.of( + Arguments.of( + "wrapped-cell-test_expense-report.pdf", + "wrapped-cell-test_expense-report.md")); + } + + @ParameterizedTest(name = "{0}") + @MethodSource("gatedFixtures") + void convertMatchesGoldenMarkdown(String pdfName, String mdName) throws IOException { + assertConversionMatchesGolden(pdfName, mdName); + } + + @Disabled("WIP fixtures below the accuracy threshold; enable locally to iterate") + @ParameterizedTest(name = "{0}") + @MethodSource("wipFixtures") + void convertMatchesGoldenMarkdownWip(String pdfName, String mdName) throws IOException { + assertConversionMatchesGolden(pdfName, mdName); + } + + /** + * Degenerate/extreme geometry must not crash the converter. A crafted or malformed PDF can + * position text anywhere via a text matrix, so a row's words can span from near the origin to a + * coordinate beyond {@link Integer#MAX_VALUE}. The old column-detection code sized an {@code + * int[]} straight from {@code (int) Math.ceil(maxX) - lo}, which either allocated a multi-GB + * array (OutOfMemoryError) or overflowed to a negative length (NegativeArraySizeException) — + * taking down the request thread. Detection must instead bail out and return no columns. + */ + @Test + void columnDetectionSurvivesDegenerateGeometry() { + // x ≈ 2.5e9 is past Integer.MAX_VALUE; combined with a near-origin word it yields an + // implausible span that the pre-fix code turned into a fatal array allocation. + List rows = new ArrayList<>(); + for (int r = 0; r < 4; r++) { + float y = 400f - r * 12f; + TextWord near = new TextWord(List.of(), 50f, y, 30f, 10f); + TextWord far = new TextWord(List.of(), 2_500_000_000f, y, 30f, 10f); + rows.add(new TextLine(List.of(near, far), 50f, y, 2_499_999_980f, 10f)); + } + + List columns = + assertDoesNotThrow(() -> PdfMarkdownConverter.findColumnRangesFromLines(rows)); + assertTrue( + columns.isEmpty(), + "implausible page span should disable column detection, not allocate from it"); + } + + private void assertConversionMatchesGolden(String pdfName, String mdName) throws IOException { + Path pdfPath = tmp.resolve(pdfName); + try (InputStream in = + getClass().getResourceAsStream("/pdf-ingestion-fixtures/" + pdfName)) { + if (in == null) { + fail("Fixture not found on classpath: /pdf-ingestion-fixtures/" + pdfName); + } + Files.copy(in, pdfPath); + } + + String actual; + try (PdfDocument doc = PdfDocument.open(pdfPath)) { + actual = new PdfMarkdownConverter().convert(doc); + } + + String expected; + try (InputStream in = getClass().getResourceAsStream("/pdf-ingestion-fixtures/" + mdName)) { + if (in == null) { + fail("Golden file not found on classpath: /pdf-ingestion-fixtures/" + mdName); + } + expected = new String(in.readAllBytes(), StandardCharsets.UTF_8); + } + + // Image placeholders are not scored: their body text is a TODO ("ideally, add the info + // available about the image...") rather than real content, so comparing it would penalise + // output for matching a placeholder we intend to replace. Drop those lines from both sides. + expected = stripImagePlaceholders(expected); + actual = stripImagePlaceholders(actual); + + double similarity = similarity(expected, actual); + if (similarity < THRESHOLD) { + fail( + String.format( + "Markdown output differs from golden file '%s' by %.1f%% (threshold %.0f%%):%n%s", + mdName, + (1.0 - similarity) * 100, + (1.0 - THRESHOLD) * 100, + unifiedDiff(expected, actual))); + } + } + + /** Substring identifying an image-placeholder line, which is excluded from scoring. */ + private static final String IMAGE_PLACEHOLDER_MARKER = "Image intentionally redacted"; + + /** + * Removes non-content lines from the comparison: image placeholders (TODO text we intend to + * replace) and GFM table separator rows (the {@code |---|---|} divider, whose exact dash count + * is cosmetic — any run of three or more dashes is valid Markdown). + */ + private static String stripImagePlaceholders(String md) { + StringBuilder sb = new StringBuilder(); + for (String line : md.split("\n", -1)) { + if (line.contains(IMAGE_PLACEHOLDER_MARKER) + || line.strip().startsWith(" 0) { + sb.append('\n'); + } + sb.append(line); + } + return sb.toString(); + } + + /** True for a GFM table separator row, e.g. {@code |---|:--:|---|} (only |, -, :, space). */ + private static boolean isTableSeparatorRow(String line) { + String t = line.strip(); + if (!t.contains("-")) { + return false; + } + return t.chars().allMatch(c -> c == '|' || c == '-' || c == ':' || c == ' '); + } + + /** + * Character-level similarity: proportion of expected characters that appear in the LCS. O(n*m) + * but golden files are small enough that this is fine. + */ + private static double similarity(String expected, String actual) { + if (expected.isEmpty() && actual.isEmpty()) return 1.0; + if (expected.isEmpty() || actual.isEmpty()) return 0.0; + // Strip all whitespace for a content-focused comparison + String e = expected.replaceAll("\\s+", " ").strip(); + String a = actual.replaceAll("\\s+", " ").strip(); + int lcs = lcsLength(e, a); + return (double) lcs / Math.max(e.length(), a.length()); + } + + private static int lcsLength(String a, String b) { + // Use two-row DP to keep memory reasonable + int m = a.length(), n = b.length(); + int[] prev = new int[n + 1]; + int[] curr = new int[n + 1]; + for (int i = 1; i <= m; i++) { + for (int j = 1; j <= n; j++) { + if (a.charAt(i - 1) == b.charAt(j - 1)) { + curr[j] = prev[j - 1] + 1; + } else { + curr[j] = Math.max(curr[j - 1], prev[j]); + } + } + int[] tmp = prev; + prev = curr; + curr = tmp; + java.util.Arrays.fill(curr, 0); + } + return prev[n]; + } + + private static String unifiedDiff(String expected, String actual) { + String[] expectedLines = expected.split("\n", -1); + String[] actualLines = actual.split("\n", -1); + + List diff = new ArrayList<>(); + diff.add("--- expected"); + diff.add("+++ actual"); + + int maxLines = Math.max(expectedLines.length, actualLines.length); + int context = 3; + boolean inHunk = false; + int hunkStart = -1; + List hunkLines = new ArrayList<>(); + + for (int i = 0; i < maxLines; i++) { + String exp = i < expectedLines.length ? expectedLines[i] : null; + String act = i < actualLines.length ? actualLines[i] : null; + + boolean changed = exp == null || act == null || !exp.equals(act); + if (changed) { + if (!inHunk) { + inHunk = true; + hunkStart = Math.max(0, i - context); + // add context lines before change + for (int c = hunkStart; c < i; c++) { + hunkLines.add(" " + (c < expectedLines.length ? expectedLines[c] : "")); + } + } + if (exp != null) hunkLines.add("-" + exp); + if (act != null) hunkLines.add("+" + act); + } else { + if (inHunk) { + hunkLines.add(" " + exp); + // check if we're far enough past the last change to close the hunk + boolean moreChanges = false; + for (int j = i + 1; j < Math.min(i + context, maxLines); j++) { + String e2 = j < expectedLines.length ? expectedLines[j] : null; + String a2 = j < actualLines.length ? actualLines[j] : null; + if (e2 == null || a2 == null || !e2.equals(a2)) { + moreChanges = true; + break; + } + } + if (!moreChanges && (i - hunkStart) >= context) { + diff.add("@@ -" + (hunkStart + 1) + " @@"); + diff.addAll(hunkLines); + hunkLines.clear(); + inHunk = false; + } + } + } + } + + if (inHunk && !hunkLines.isEmpty()) { + diff.add("@@ -" + (hunkStart + 1) + " @@"); + diff.addAll(hunkLines); + } + + return String.join("\n", diff); + } +} diff --git a/app/common/src/test/java/stirling/software/common/service/FileStorageDelegationTest.java b/app/common/src/test/java/stirling/software/common/service/FileStorageDelegationTest.java index a929a15af2..b013a26927 100644 --- a/app/common/src/test/java/stirling/software/common/service/FileStorageDelegationTest.java +++ b/app/common/src/test/java/stirling/software/common/service/FileStorageDelegationTest.java @@ -5,6 +5,7 @@ import static org.mockito.Mockito.mock; import java.io.IOException; import java.nio.file.Path; +import java.util.Optional; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.io.TempDir; @@ -19,7 +20,8 @@ class FileStorageDelegationTest { FileStorage fs = new FileStorage( mock(FileOrUploadService.class), - new LocalDiskFileStore(tempDir.toString())); + new LocalDiskFileStore(tempDir.toString()), + Optional.empty()); byte[] payload = "round-trip".getBytes(); String id = fs.storeBytes(payload, "x.bin"); assertArrayEquals(payload, fs.retrieveBytes(id)); diff --git a/app/common/src/test/java/stirling/software/common/service/FileStorageOwnershipTest.java b/app/common/src/test/java/stirling/software/common/service/FileStorageOwnershipTest.java new file mode 100644 index 0000000000..861efa8509 --- /dev/null +++ b/app/common/src/test/java/stirling/software/common/service/FileStorageOwnershipTest.java @@ -0,0 +1,107 @@ +package stirling.software.common.service; + +import static org.junit.jupiter.api.Assertions.assertArrayEquals; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; + +import java.io.IOException; +import java.nio.file.Path; +import java.util.Optional; +import java.util.concurrent.atomic.AtomicReference; + +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.io.TempDir; + +import stirling.software.common.cluster.inprocess.LocalDiskFileStore; +import stirling.software.common.util.JobContext; + +class FileStorageOwnershipTest { + + private FileStorage newStorageWithoutSecurity(Path tempDir) { + return new FileStorage( + mock(FileOrUploadService.class), + new LocalDiskFileStore(tempDir.toString()), + Optional.empty()); + } + + private FileStorage newStorageWithCurrentUser(Path tempDir, AtomicReference userRef) { + JobOwnershipService svc = mock(JobOwnershipService.class); + when(svc.getCurrentUserId()).thenAnswer(invocation -> Optional.ofNullable(userRef.get())); + return new FileStorage( + mock(FileOrUploadService.class), + new LocalDiskFileStore(tempDir.toString()), + Optional.of(svc)); + } + + @Test + void desktopMode_noOwnershipService_storesAndRetrievesWithoutChecks(@TempDir Path tempDir) + throws IOException { + FileStorage fs = newStorageWithoutSecurity(tempDir); + byte[] payload = "desktop".getBytes(); + String id = fs.storeBytes(payload, "x.bin"); + assertArrayEquals(payload, fs.retrieveBytes(id)); + } + + @Test + void sameUserStoresAndRetrieves_allowed(@TempDir Path tempDir) throws IOException { + AtomicReference user = new AtomicReference<>("alice"); + FileStorage fs = newStorageWithCurrentUser(tempDir, user); + byte[] payload = "alice's file".getBytes(); + String id = fs.storeBytes(payload, "x.bin"); + assertArrayEquals(payload, fs.retrieveBytes(id)); + } + + @Test + void differentUserRetrieves_throwsSecurityException(@TempDir Path tempDir) throws IOException { + AtomicReference user = new AtomicReference<>("alice"); + FileStorage fs = newStorageWithCurrentUser(tempDir, user); + String id = fs.storeBytes("alice's file".getBytes(), "x.bin"); + user.set("bob"); + assertThrows(SecurityException.class, () -> fs.retrieveBytes(id)); + assertThrows(SecurityException.class, () -> fs.retrieveInputStream(id)); + assertThrows(SecurityException.class, () -> fs.getFileSize(id)); + assertThrows(SecurityException.class, () -> fs.fileExists(id)); + assertThrows(SecurityException.class, () -> fs.deleteFile(id)); + } + + @Test + void anonymousRetrieveOfOwnedFile_allowed_noCurrentUserMeansNoCompare(@TempDir Path tempDir) + throws IOException { + AtomicReference user = new AtomicReference<>("alice"); + FileStorage fs = newStorageWithCurrentUser(tempDir, user); + byte[] payload = "alice's file".getBytes(); + String id = fs.storeBytes(payload, "x.bin"); + user.set(null); + assertArrayEquals(payload, fs.retrieveBytes(id)); + } + + @Test + void authedRetrieveOfAnonymousFile_allowed_noOwnerOnFile(@TempDir Path tempDir) + throws IOException { + AtomicReference user = new AtomicReference<>(null); + FileStorage fs = newStorageWithCurrentUser(tempDir, user); + byte[] payload = "no-owner".getBytes(); + String id = fs.storeBytes(payload, "x.bin"); + user.set("alice"); + assertArrayEquals(payload, fs.retrieveBytes(id)); + } + + @Test + void propagatedOwner_scopesAsyncWriteWithNoLiveUser(@TempDir Path tempDir) throws IOException { + AtomicReference user = new AtomicReference<>(null); + FileStorage fs = newStorageWithCurrentUser(tempDir, user); + byte[] payload = "alice's async result".getBytes(); + String id; + try { + JobContext.setOwner("alice"); + id = fs.storeBytes(payload, "x.bin"); + } finally { + JobContext.clear(); + } + user.set("alice"); + assertArrayEquals(payload, fs.retrieveBytes(id)); + user.set("bob"); + assertThrows(SecurityException.class, () -> fs.retrieveBytes(id)); + } +} diff --git a/app/common/src/test/java/stirling/software/common/service/FileStorageTest.java b/app/common/src/test/java/stirling/software/common/service/FileStorageTest.java index ace0dfa567..32a3cb65de 100644 --- a/app/common/src/test/java/stirling/software/common/service/FileStorageTest.java +++ b/app/common/src/test/java/stirling/software/common/service/FileStorageTest.java @@ -9,6 +9,8 @@ import java.io.InputStream; import java.nio.charset.StandardCharsets; import java.nio.file.Files; import java.nio.file.Path; +import java.util.Optional; +import java.util.UUID; import java.util.stream.Stream; import org.junit.jupiter.api.BeforeEach; @@ -37,7 +39,10 @@ class FileStorageTest { void setUp() throws IOException { MockitoAnnotations.openMocks(this); fileStorage = - new FileStorage(fileOrUploadService, new LocalDiskFileStore(tempDir.toString())); + new FileStorage( + fileOrUploadService, + new LocalDiskFileStore(tempDir.toString()), + Optional.empty()); // Create a mock MultipartFile mockFile = mock(MultipartFile.class); @@ -79,7 +84,7 @@ class FileStorageTest { void testRetrieveFile() throws IOException { // Arrange byte[] fileContent = "Test PDF content".getBytes(); - String fileId = "test-file-1"; + String fileId = UUID.randomUUID().toString(); Path filePath = tempDir.resolve(fileId); Files.write(filePath, fileContent); @@ -99,7 +104,7 @@ class FileStorageTest { void testRetrieveBytes() throws IOException { // Arrange byte[] fileContent = "Test PDF content".getBytes(); - String fileId = "test-file-2"; + String fileId = UUID.randomUUID().toString(); Path filePath = tempDir.resolve(fileId); Files.write(filePath, fileContent); @@ -113,7 +118,7 @@ class FileStorageTest { @Test void testRetrieveFile_FileNotFound() { // Arrange - String nonExistentFileId = "non-existent-file"; + String nonExistentFileId = UUID.randomUUID().toString(); // Act & Assert assertThrows(IOException.class, () -> fileStorage.retrieveFile(nonExistentFileId)); @@ -122,7 +127,7 @@ class FileStorageTest { @Test void testRetrieveBytes_FileNotFound() { // Arrange - String nonExistentFileId = "non-existent-file"; + String nonExistentFileId = UUID.randomUUID().toString(); // Act & Assert assertThrows(IOException.class, () -> fileStorage.retrieveBytes(nonExistentFileId)); @@ -132,7 +137,7 @@ class FileStorageTest { void testDeleteFile() throws IOException { // Arrange byte[] fileContent = "Test PDF content".getBytes(); - String fileId = "test-file-3"; + String fileId = UUID.randomUUID().toString(); Path filePath = tempDir.resolve(fileId); Files.write(filePath, fileContent); @@ -147,7 +152,7 @@ class FileStorageTest { @Test void testDeleteFile_FileNotFound() { // Arrange - String nonExistentFileId = "non-existent-file"; + String nonExistentFileId = UUID.randomUUID().toString(); // Act boolean result = fileStorage.deleteFile(nonExistentFileId); @@ -160,7 +165,7 @@ class FileStorageTest { void testFileExists() throws IOException { // Arrange byte[] fileContent = "Test PDF content".getBytes(); - String fileId = "test-file-4"; + String fileId = UUID.randomUUID().toString(); Path filePath = tempDir.resolve(fileId); Files.write(filePath, fileContent); @@ -174,7 +179,7 @@ class FileStorageTest { @Test void testFileExists_FileNotFound() { // Arrange - String nonExistentFileId = "non-existent-file"; + String nonExistentFileId = UUID.randomUUID().toString(); // Act boolean result = fileStorage.fileExists(nonExistentFileId); diff --git a/app/common/src/test/java/stirling/software/common/service/InternalApiClientTest.java b/app/common/src/test/java/stirling/software/common/service/InternalApiClientTest.java index 06543e47b0..7bab141888 100644 --- a/app/common/src/test/java/stirling/software/common/service/InternalApiClientTest.java +++ b/app/common/src/test/java/stirling/software/common/service/InternalApiClientTest.java @@ -59,6 +59,53 @@ class InternalApiClientTest { servletContext, userService, tempFileManager, environment, applicationProperties); } + @Test + void postTagsRequestAsAutomation() throws Exception { + // Every InternalApiClient.post() caller is a parent automation flow dispatching a child + // tool (pipeline executor, AI workflow, policy runner). Tagging the sub-step here means + // the saas PaygChargeInterceptor classifies it as BillingCategory.AUTOMATION regardless of + // the dispatched controller's @RequiresFeature — so an AI-OCR step inside a policy run + // bills as AUTOMATION, not AI. The header value is the literal string "true" because the + // interceptor compares case-insensitively-trimmed against that token. + MultiValueMap body = new LinkedMultiValueMap<>(); + body.add("fileInput", namedResource("input.pdf", "data")); + + Path tempPath = Files.createTempFile("internal-api-automation-test", ".tmp"); + TempFile tempFile = mock(TempFile.class); + when(tempFile.getPath()).thenReturn(tempPath); + when(tempFile.getFile()).thenReturn(tempPath.toFile()); + when(tempFileManager.createManagedTempFile("internal-api")).thenReturn(tempFile); + + HttpHeaders[] captured = {null}; + + try (var ignored = + mockConstruction( + RestTemplate.class, + (rt, ctx) -> { + when(rt.httpEntityCallback(any(), eq(Resource.class))) + .thenAnswer( + inv -> { + HttpEntity entity = inv.getArgument(0); + captured[0] = entity.getHeaders(); + return (RequestCallback) req -> {}; + }); + when(rt.execute(anyString(), eq(HttpMethod.POST), any(), any())) + .thenAnswer(inv -> fakeOkResponse(inv.getArgument(3))); + })) { + + InternalApiClient mockedClient = newClient(); + mockedClient.post("/api/v1/general/merge-pdfs", body); + + assertNotNull(captured[0]); + assertEquals( + "true", + captured[0].getFirst(InternalApiClient.AUTOMATION_HEADER), + "Sub-step dispatch must carry the automation marker header"); + } finally { + Files.deleteIfExists(tempPath); + } + } + @Test void postDoesNotForceContentType() throws Exception { MultiValueMap body = new LinkedMultiValueMap<>(); diff --git a/app/common/src/test/java/stirling/software/common/service/MobileScannerServiceTest.java b/app/common/src/test/java/stirling/software/common/service/MobileScannerServiceTest.java new file mode 100644 index 0000000000..1bfc166676 --- /dev/null +++ b/app/common/src/test/java/stirling/software/common/service/MobileScannerServiceTest.java @@ -0,0 +1,471 @@ +package stirling.software.common.service; + +import static org.junit.jupiter.api.Assertions.*; + +import java.io.IOException; +import java.nio.charset.StandardCharsets; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.List; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.io.TempDir; +import org.springframework.mock.web.MockMultipartFile; +import org.springframework.test.util.ReflectionTestUtils; +import org.springframework.web.multipart.MultipartFile; + +import stirling.software.common.service.MobileScannerService.FileMetadata; +import stirling.software.common.service.MobileScannerService.SessionInfo; + +/** + * Unit tests for {@link MobileScannerService}. The service stores uploaded files in a temp + * directory. To keep tests isolated and deterministic, the {@code tempDirectory} field is + * redirected to a JUnit {@link TempDir} via reflection after construction. + */ +class MobileScannerServiceTest { + + @TempDir Path tempDir; + + private MobileScannerService service; + + @BeforeEach + void setUp() throws IOException { + service = new MobileScannerService(); + // Redirect the service's temp directory to the isolated test temp dir. + ReflectionTestUtils.setField(service, "tempDirectory", tempDir); + } + + private MultipartFile file(String name, String content) { + return new MockMultipartFile( + "file", name, "text/plain", content.getBytes(StandardCharsets.UTF_8)); + } + + private MultipartFile emptyFile(String name) { + return new MockMultipartFile("file", name, "text/plain", new byte[0]); + } + + @Nested + @DisplayName("createSession") + class CreateSession { + + @Test + @DisplayName("creates a session and returns coherent SessionInfo") + void createsSession() { + SessionInfo info = service.createSession("abc-123"); + + assertNotNull(info); + assertEquals("abc-123", info.getSessionId()); + assertTrue(info.getCreatedAt() > 0); + assertEquals(10 * 60 * 1000L, info.getTimeoutMs()); + assertEquals(info.getCreatedAt() + info.getTimeoutMs(), info.getExpiresAt()); + } + + @Test + @DisplayName("session is retrievable via validateSession after creation") + void createdSessionIsValid() { + service.createSession("sess1"); + assertNotNull(service.validateSession("sess1")); + } + + @Test + @DisplayName("rejects null session ID") + void rejectsNull() { + assertThrows(IllegalArgumentException.class, () -> service.createSession(null)); + } + + @Test + @DisplayName("rejects blank session ID") + void rejectsBlank() { + assertThrows(IllegalArgumentException.class, () -> service.createSession(" ")); + } + + @Test + @DisplayName("rejects session ID with invalid characters") + void rejectsInvalidChars() { + assertThrows(IllegalArgumentException.class, () -> service.createSession("bad/id")); + assertThrows(IllegalArgumentException.class, () -> service.createSession("bad id")); + assertThrows(IllegalArgumentException.class, () -> service.createSession("bad_id")); + } + + @Test + @DisplayName("accepts alphanumeric and hyphen session IDs") + void acceptsValidChars() { + assertNotNull(service.createSession("ABC-def-123")); + } + } + + @Nested + @DisplayName("validateSession") + class ValidateSession { + + @Test + @DisplayName("returns null for unknown session") + void unknownReturnsNull() { + assertNull(service.validateSession("does-not-exist")); + } + + @Test + @DisplayName("returns SessionInfo for an existing session") + void existingReturnsInfo() { + service.createSession("s1"); + SessionInfo info = service.validateSession("s1"); + + assertNotNull(info); + assertEquals("s1", info.getSessionId()); + assertEquals(10 * 60 * 1000L, info.getTimeoutMs()); + } + + @Test + @DisplayName("expires and removes a session whose last access is in the past") + void expiredSessionRemoved() { + service.createSession("expired"); + + // Force the underlying session's last access far into the past. + forceLastAccess("expired", System.currentTimeMillis() - (20 * 60 * 1000L)); + + assertNull(service.validateSession("expired")); + // After expiry the session should be gone entirely. + assertNull(service.validateSession("expired")); + } + } + + @Nested + @DisplayName("uploadFiles") + class UploadFiles { + + @Test + @DisplayName("stores files and records metadata") + void storesFiles() throws IOException { + service.createSession("up1"); + service.uploadFiles("up1", List.of(file("scan.txt", "hello"))); + + List metas = service.getSessionFiles("up1"); + assertEquals(1, metas.size()); + FileMetadata meta = metas.get(0); + assertEquals("scan.txt", meta.getFilename()); + assertEquals(5, meta.getSize()); + assertEquals("text/plain", meta.getContentType()); + + // File physically exists on disk. + Path stored = tempDir.resolve("up1").resolve("scan.txt"); + assertTrue(Files.exists(stored)); + assertEquals("hello", Files.readString(stored)); + } + + @Test + @DisplayName("auto-creates a session when uploading to an unregistered session ID") + void autoCreatesSession() throws IOException { + service.uploadFiles("new-session", List.of(file("a.txt", "data"))); + + List metas = service.getSessionFiles("new-session"); + assertEquals(1, metas.size()); + } + + @Test + @DisplayName("skips empty files") + void skipsEmptyFiles() throws IOException { + service.createSession("up2"); + service.uploadFiles("up2", List.of(emptyFile("empty.txt"), file("real.txt", "x"))); + + List metas = service.getSessionFiles("up2"); + assertEquals(1, metas.size()); + assertEquals("real.txt", metas.get(0).getFilename()); + } + + @Test + @DisplayName("sanitizes dangerous filename characters") + void sanitizesFilename() throws IOException { + service.createSession("up3"); + service.uploadFiles("up3", List.of(file("we ird@na#me.txt", "x"))); + + List metas = service.getSessionFiles("up3"); + assertEquals(1, metas.size()); + String stored = metas.get(0).getFilename(); + // Disallowed chars replaced with underscores; allowed set is [a-zA-Z0-9._-]. + assertTrue(stored.matches("[a-zA-Z0-9._-]+"), "unexpected filename: " + stored); + assertTrue(Files.exists(tempDir.resolve("up3").resolve(stored))); + } + + @Test + @DisplayName("handles duplicate filenames by appending a counter") + void handlesDuplicateFilenames() throws IOException { + service.createSession("up4"); + service.uploadFiles("up4", List.of(file("dup.txt", "one"))); + service.uploadFiles("up4", List.of(file("dup.txt", "two"))); + + List metas = service.getSessionFiles("up4"); + assertEquals(2, metas.size()); + + Path original = tempDir.resolve("up4").resolve("dup.txt"); + Path renamed = tempDir.resolve("up4").resolve("dup-1.txt"); + assertTrue(Files.exists(original)); + assertTrue(Files.exists(renamed)); + assertEquals("one", Files.readString(original)); + assertEquals("two", Files.readString(renamed)); + } + + @Test + @DisplayName("falls back to a generated name when original filename is null") + void generatesNameWhenNull() throws IOException { + service.createSession("up5"); + MultipartFile noName = + new MockMultipartFile("file", null, "text/plain", "x".getBytes()); + service.uploadFiles("up5", List.of(noName)); + + List metas = service.getSessionFiles("up5"); + assertEquals(1, metas.size()); + assertTrue(metas.get(0).getFilename().startsWith("upload-")); + } + + @Test + @DisplayName("rejects invalid session ID before any storage") + void rejectsInvalidSessionId() { + assertThrows( + IllegalArgumentException.class, + () -> service.uploadFiles("bad/id", List.of(file("a.txt", "x")))); + } + + @Test + @DisplayName("uploading an empty list leaves no files") + void emptyListNoFiles() throws IOException { + service.createSession("up6"); + service.uploadFiles("up6", List.of()); + + assertTrue(service.getSessionFiles("up6").isEmpty()); + } + } + + @Nested + @DisplayName("getSessionFiles") + class GetSessionFiles { + + @Test + @DisplayName("returns empty list for unknown session") + void unknownReturnsEmpty() { + assertTrue(service.getSessionFiles("nope").isEmpty()); + } + + @Test + @DisplayName("returns a defensive copy of the metadata list") + void returnsDefensiveCopy() throws IOException { + service.createSession("g1"); + service.uploadFiles("g1", List.of(file("a.txt", "x"))); + + List first = service.getSessionFiles("g1"); + first.clear(); + + // Mutating the returned list must not affect the service's internal state. + assertEquals(1, service.getSessionFiles("g1").size()); + } + } + + @Nested + @DisplayName("getFile") + class GetFile { + + @Test + @DisplayName("returns the path of an uploaded file") + void returnsPath() throws IOException { + service.createSession("f1"); + service.uploadFiles("f1", List.of(file("doc.txt", "body"))); + + Path path = service.getFile("f1", "doc.txt"); + assertTrue(Files.exists(path)); + assertEquals("body", Files.readString(path)); + } + + @Test + @DisplayName("throws when the session does not exist") + void unknownSessionThrows() { + IOException ex = + assertThrows(IOException.class, () -> service.getFile("ghost", "doc.txt")); + assertTrue(ex.getMessage().contains("Session not found")); + } + + @Test + @DisplayName("throws when the file does not exist in an existing session") + void unknownFileThrows() throws IOException { + service.createSession("f2"); + service.uploadFiles("f2", List.of(file("present.txt", "x"))); + + IOException ex = + assertThrows(IOException.class, () -> service.getFile("f2", "missing.txt")); + assertTrue(ex.getMessage().contains("File not found")); + } + + @Test + @DisplayName("rejects filenames containing path separators") + void rejectsPathSeparators() throws IOException { + service.createSession("f3"); + service.uploadFiles("f3", List.of(file("ok.txt", "x"))); + + assertThrows(IOException.class, () -> service.getFile("f3", "../escape.txt")); + assertThrows(IOException.class, () -> service.getFile("f3", "sub/file.txt")); + assertThrows(IOException.class, () -> service.getFile("f3", "sub\\file.txt")); + } + + @Test + @DisplayName("rejects blank filename") + void rejectsBlankFilename() throws IOException { + service.createSession("f4"); + service.uploadFiles("f4", List.of(file("ok.txt", "x"))); + + assertThrows(IOException.class, () -> service.getFile("f4", " ")); + } + } + + @Nested + @DisplayName("deleteFileAfterDownload") + class DeleteFileAfterDownload { + + @Test + @DisplayName("deletes a single file but keeps the session if others remain") + void deletesOneFile() throws IOException { + service.createSession("d1"); + service.uploadFiles("d1", List.of(file("a.txt", "x"), file("b.txt", "y"))); + + service.deleteFileAfterDownload("d1", "a.txt"); + + assertFalse(Files.exists(tempDir.resolve("d1").resolve("a.txt"))); + // Session still present because not all files have been downloaded. + assertNotNull(service.validateSession("d1")); + } + + @Test + @DisplayName("deletes the entire session once all files are marked downloaded") + void deletesSessionWhenAllDownloaded() throws IOException { + service.createSession("d2"); + service.uploadFiles("d2", List.of(file("only.txt", "x"))); + + // Mark the file as downloaded via getFile, then delete it. + service.getFile("d2", "only.txt"); + service.deleteFileAfterDownload("d2", "only.txt"); + + assertNull(service.validateSession("d2")); + assertFalse(Files.exists(tempDir.resolve("d2"))); + } + + @Test + @DisplayName("does not throw for an unknown session") + void unknownSessionNoThrow() { + assertDoesNotThrow(() -> service.deleteFileAfterDownload("ghost", "a.txt")); + } + + @Test + @DisplayName("swallows invalid filename input without throwing") + void invalidFilenameNoThrow() throws IOException { + service.createSession("d3"); + service.uploadFiles("d3", List.of(file("a.txt", "x"))); + + assertDoesNotThrow(() -> service.deleteFileAfterDownload("d3", "../escape.txt")); + // Original file untouched. + assertTrue(Files.exists(tempDir.resolve("d3").resolve("a.txt"))); + } + } + + @Nested + @DisplayName("deleteSession") + class DeleteSession { + + @Test + @DisplayName("removes the session and all its files") + void removesSessionAndFiles() throws IOException { + service.createSession("x1"); + service.uploadFiles("x1", List.of(file("a.txt", "x"), file("b.txt", "y"))); + + assertTrue(Files.exists(tempDir.resolve("x1"))); + + service.deleteSession("x1"); + + assertNull(service.validateSession("x1")); + assertFalse(Files.exists(tempDir.resolve("x1"))); + } + + @Test + @DisplayName("is a no-op for an unknown session") + void unknownSessionNoOp() { + assertDoesNotThrow(() -> service.deleteSession("never-existed")); + } + } + + @Nested + @DisplayName("cleanupExpiredSessions") + class CleanupExpiredSessions { + + @Test + @DisplayName("removes sessions past the timeout") + void removesExpired() throws IOException { + service.createSession("old"); + service.uploadFiles("old", List.of(file("a.txt", "x"))); + forceLastAccess("old", System.currentTimeMillis() - (20 * 60 * 1000L)); + + service.cleanupExpiredSessions(); + + assertNull(service.validateSession("old")); + assertFalse(Files.exists(tempDir.resolve("old"))); + } + + @Test + @DisplayName("keeps sessions that are still fresh") + void keepsFresh() { + service.createSession("fresh"); + + service.cleanupExpiredSessions(); + + assertNotNull(service.validateSession("fresh")); + } + + @Test + @DisplayName("does not throw when there are no sessions") + void noSessionsNoThrow() { + assertDoesNotThrow(() -> service.cleanupExpiredSessions()); + } + } + + @Nested + @DisplayName("SessionInfo accessors") + class SessionInfoAccessors { + + @Test + @DisplayName("exposes all constructor values") + void exposesValues() { + SessionInfo info = new SessionInfo("id", 100L, 200L, 50L); + assertEquals("id", info.getSessionId()); + assertEquals(100L, info.getCreatedAt()); + assertEquals(200L, info.getExpiresAt()); + assertEquals(50L, info.getTimeoutMs()); + } + } + + @Nested + @DisplayName("FileMetadata accessors") + class FileMetadataAccessors { + + @Test + @DisplayName("exposes all constructor values") + void exposesValues() { + FileMetadata meta = new FileMetadata("name.pdf", 1234L, "application/pdf"); + assertEquals("name.pdf", meta.getFilename()); + assertEquals(1234L, meta.getSize()); + assertEquals("application/pdf", meta.getContentType()); + } + } + + /** + * Reaches into the internal SessionData for a given session and forces its lastAccessTime, used + * to deterministically simulate expiry without sleeping. + */ + @SuppressWarnings("unchecked") + private void forceLastAccess(String sessionId, long lastAccessTime) { + java.util.Map sessions = + (java.util.Map) + ReflectionTestUtils.getField(service, "activeSessions"); + assertNotNull(sessions); + Object sessionData = sessions.get(sessionId); + assertNotNull(sessionData, "session not found: " + sessionId); + ReflectionTestUtils.setField(sessionData, "lastAccessTime", lastAccessTime); + } +} diff --git a/app/common/src/test/java/stirling/software/common/service/PdfMetadataServiceTest.java b/app/common/src/test/java/stirling/software/common/service/PdfMetadataServiceTest.java new file mode 100644 index 0000000000..6519d667a0 --- /dev/null +++ b/app/common/src/test/java/stirling/software/common/service/PdfMetadataServiceTest.java @@ -0,0 +1,416 @@ +package stirling.software.common.service; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.mockito.Mockito.lenient; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; + +import java.time.LocalDateTime; +import java.time.ZoneId; +import java.time.ZonedDateTime; +import java.util.Calendar; + +import org.apache.pdfbox.pdmodel.PDDocument; +import org.apache.pdfbox.pdmodel.PDDocumentInformation; +import org.apache.pdfbox.pdmodel.PDPage; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.common.model.ApplicationProperties.Premium; +import stirling.software.common.model.ApplicationProperties.Premium.ProFeatures; +import stirling.software.common.model.ApplicationProperties.Premium.ProFeatures.CustomMetadata; +import stirling.software.common.model.PdfMetadata; + +class PdfMetadataServiceTest { + + private static final String LABEL = "Stirling-PDF v1.0.0"; + + /** + * Builds a service whose pro-features are disabled (real ApplicationProperties, all defaults). + */ + private PdfMetadataService nonProService(UserServiceInterface userService) { + return new PdfMetadataService(new ApplicationProperties(), LABEL, false, userService); + } + + @Nested + @DisplayName("toCalendar(ZonedDateTime)") + class ToCalendarTests { + + @Test + @DisplayName("returns null for null input") + void nullReturnsNull() { + assertNull(PdfMetadataService.toCalendar(null)); + } + + @Test + @DisplayName("converts ZonedDateTime preserving the instant") + void convertsInstant() { + ZonedDateTime zdt = ZonedDateTime.of(2021, 6, 15, 10, 30, 45, 0, ZoneId.of("UTC")); + Calendar cal = PdfMetadataService.toCalendar(zdt); + + assertNotNull(cal); + assertEquals(zdt.toInstant().toEpochMilli(), cal.getTimeInMillis()); + } + } + + @Nested + @DisplayName("parseToCalendar(String)") + class ParseToCalendarTests { + + @Test + @DisplayName("returns null for null input") + void nullReturnsNull() { + assertNull(PdfMetadataService.parseToCalendar(null)); + } + + @Test + @DisplayName("returns null for empty / blank input") + void blankReturnsNull() { + assertNull(PdfMetadataService.parseToCalendar("")); + assertNull(PdfMetadataService.parseToCalendar(" ")); + } + + @Test + @DisplayName("returns null for unparsable input") + void invalidReturnsNull() { + assertNull(PdfMetadataService.parseToCalendar("not a date")); + assertNull(PdfMetadataService.parseToCalendar("2021-06-15")); + assertNull(PdfMetadataService.parseToCalendar("2021/13/40 99:99:99")); + } + + @Test + @DisplayName("parses a valid 'yyyy/MM/dd HH:mm:ss' string") + void parsesValidDate() { + Calendar cal = PdfMetadataService.parseToCalendar("2021/06/15 10:30:45"); + assertNotNull(cal); + + // Build the expected instant the same way the implementation does so the + // assertion is independent of the JVM's default time zone. + long expectedMillis = + LocalDateTime.of(2021, 6, 15, 10, 30, 45) + .atZone(ZoneId.systemDefault()) + .toInstant() + .toEpochMilli(); + assertEquals(expectedMillis, cal.getTimeInMillis()); + } + } + + @Nested + @DisplayName("extractMetadataFromPdf(PDDocument)") + class ExtractMetadataTests { + + @Test + @DisplayName("returns all-null fields for a fresh empty document") + void emptyDocumentYieldsNulls() throws Exception { + PdfMetadataService service = nonProService(null); + try (PDDocument doc = new PDDocument()) { + PdfMetadata md = service.extractMetadataFromPdf(doc); + + assertNotNull(md); + assertNull(md.getAuthor()); + assertNull(md.getProducer()); + assertNull(md.getTitle()); + assertNull(md.getCreator()); + assertNull(md.getSubject()); + assertNull(md.getKeywords()); + assertNull(md.getCreationDate()); + assertNull(md.getModificationDate()); + } + } + + @Test + @DisplayName("reads back string and date fields set on the document") + void readsBackPopulatedFields() throws Exception { + PdfMetadataService service = nonProService(null); + try (PDDocument doc = new PDDocument()) { + PDDocumentInformation info = doc.getDocumentInformation(); + info.setAuthor("Alice"); + info.setProducer("ProducerX"); + info.setTitle("My Title"); + info.setCreator("CreatorY"); + info.setSubject("Subject Z"); + info.setKeywords("k1, k2"); + + Calendar creation = Calendar.getInstance(); + creation.setTimeInMillis(1_600_000_000_000L); + Calendar modification = Calendar.getInstance(); + modification.setTimeInMillis(1_700_000_000_000L); + info.setCreationDate(creation); + info.setModificationDate(modification); + + PdfMetadata md = service.extractMetadataFromPdf(doc); + + assertEquals("Alice", md.getAuthor()); + assertEquals("ProducerX", md.getProducer()); + assertEquals("My Title", md.getTitle()); + assertEquals("CreatorY", md.getCreator()); + assertEquals("Subject Z", md.getSubject()); + assertEquals("k1, k2", md.getKeywords()); + + assertNotNull(md.getCreationDate()); + assertNotNull(md.getModificationDate()); + assertEquals(1_600_000_000_000L, md.getCreationDate().toInstant().toEpochMilli()); + assertEquals( + 1_700_000_000_000L, md.getModificationDate().toInstant().toEpochMilli()); + } + } + } + + @Nested + @DisplayName("setMetadataToPdf / setDefaultMetadata (non-pro path)") + class SetMetadataNonProTests { + + @Test + @DisplayName("writes producer label, title, subject, keywords and author from metadata") + void writesCommonMetadata() throws Exception { + PdfMetadataService service = nonProService(null); + PdfMetadata md = + PdfMetadata.builder() + .author("Bob") + .title("Doc Title") + .subject("Doc Subject") + .keywords("a, b, c") + .creationDate( + ZonedDateTime.of(2020, 1, 1, 0, 0, 0, 0, ZoneId.of("UTC"))) + .modificationDate( + ZonedDateTime.of(2021, 1, 1, 0, 0, 0, 0, ZoneId.of("UTC"))) + .build(); + + try (PDDocument doc = new PDDocument()) { + doc.addPage(new PDPage()); + service.setMetadataToPdf(doc, md); + + PDDocumentInformation info = doc.getDocumentInformation(); + assertEquals(LABEL, info.getProducer()); + assertEquals("Doc Title", info.getTitle()); + assertEquals("Doc Subject", info.getSubject()); + assertEquals("a, b, c", info.getKeywords()); + // Non-pro: author is taken verbatim from the metadata. + assertEquals("Bob", info.getAuthor()); + assertNotNull(info.getModificationDate()); + } + } + + @Test + @DisplayName("existing creation date is left untouched when not newly created") + void keepsExistingCreationDate() throws Exception { + PdfMetadataService service = nonProService(null); + ZonedDateTime creation = ZonedDateTime.of(2019, 5, 20, 8, 15, 0, 0, ZoneId.of("UTC")); + PdfMetadata md = PdfMetadata.builder().title("T").creationDate(creation).build(); + + try (PDDocument doc = new PDDocument()) { + service.setMetadataToPdf(doc, md); + + Calendar creationCal = doc.getDocumentInformation().getCreationDate(); + // creationDate is non-null and newlyCreated=false, so setNewDocumentMetadata + // is skipped and no creation date is written. + assertNull(creationCal); + } + } + + @Test + @DisplayName("sets a fresh creation date when metadata has none") + void setsCreationDateWhenMissing() throws Exception { + PdfMetadataService service = nonProService(null); + PdfMetadata md = PdfMetadata.builder().title("T").build(); + + try (PDDocument doc = new PDDocument()) { + service.setMetadataToPdf(doc, md); + + Calendar creationCal = doc.getDocumentInformation().getCreationDate(); + assertNotNull(creationCal); + // Non-pro path writes the Stirling label as the creator. + assertEquals(LABEL, doc.getDocumentInformation().getCreator()); + } + } + + @Test + @DisplayName("newlyCreated=true forces a fresh creation date even if metadata has one") + void newlyCreatedForcesCreationDate() throws Exception { + PdfMetadataService service = nonProService(null); + ZonedDateTime creation = ZonedDateTime.of(2018, 3, 3, 3, 3, 3, 0, ZoneId.of("UTC")); + PdfMetadata md = PdfMetadata.builder().title("T").creationDate(creation).build(); + + try (PDDocument doc = new PDDocument()) { + service.setMetadataToPdf(doc, md, true); + + Calendar creationCal = doc.getDocumentInformation().getCreationDate(); + assertNotNull(creationCal); + // The supplied creation date must have been honoured (not "now"). + assertEquals(creation.toInstant().toEpochMilli(), creationCal.getTimeInMillis()); + assertEquals(LABEL, doc.getDocumentInformation().getCreator()); + } + } + + @Test + @DisplayName( + "setDefaultMetadata round-trips existing document info through the producer label") + void setDefaultMetadataRewritesProducer() throws Exception { + PdfMetadataService service = nonProService(null); + try (PDDocument doc = new PDDocument()) { + PDDocumentInformation info = doc.getDocumentInformation(); + info.setTitle("Original Title"); + info.setAuthor("Original Author"); + info.setProducer("Some Other Producer"); + + service.setDefaultMetadata(doc); + + // extract + re-apply keeps title/author but rewrites producer to the label. + assertEquals("Original Title", info.getTitle()); + assertEquals("Original Author", info.getAuthor()); + assertEquals(LABEL, info.getProducer()); + } + } + + @Test + @DisplayName("null string fields in metadata are written through without error") + void handlesNullStringFields() throws Exception { + PdfMetadataService service = nonProService(null); + PdfMetadata md = PdfMetadata.builder().build(); + + try (PDDocument doc = new PDDocument()) { + service.setMetadataToPdf(doc, md, true); + + PDDocumentInformation info = doc.getDocumentInformation(); + assertEquals(LABEL, info.getProducer()); + assertNull(info.getTitle()); + assertNull(info.getSubject()); + assertNull(info.getKeywords()); + assertNull(info.getAuthor()); + // newlyCreated=true always stamps a creation date. + assertNotNull(info.getCreationDate()); + assertNotNull(info.getModificationDate()); + } + } + } + + @Nested + @DisplayName("setMetadataToPdf (pro path with custom metadata)") + class SetMetadataProTests { + + private ApplicationProperties propsWithCustomMetadata( + boolean autoUpdate, String author, String creator) { + ApplicationProperties props = mock(ApplicationProperties.class); + Premium premium = mock(Premium.class); + ProFeatures proFeatures = mock(ProFeatures.class); + CustomMetadata customMetadata = mock(CustomMetadata.class); + + lenient().when(props.getPremium()).thenReturn(premium); + lenient().when(premium.getProFeatures()).thenReturn(proFeatures); + lenient().when(proFeatures.getCustomMetadata()).thenReturn(customMetadata); + lenient().when(customMetadata.isAutoUpdateMetadata()).thenReturn(autoUpdate); + lenient().when(customMetadata.getAuthor()).thenReturn(author); + lenient().when(customMetadata.getCreator()).thenReturn(creator); + return props; + } + + @Test + @DisplayName("uses custom author and creator when pro and auto-update enabled") + void appliesCustomAuthorAndCreator() throws Exception { + ApplicationProperties props = + propsWithCustomMetadata(true, "Custom Author", "Custom Creator"); + PdfMetadataService service = new PdfMetadataService(props, LABEL, true, null); + + PdfMetadata md = PdfMetadata.builder().author("Ignored").title("T").build(); + + try (PDDocument doc = new PDDocument()) { + service.setMetadataToPdf(doc, md, true); + + PDDocumentInformation info = doc.getDocumentInformation(); + assertEquals("Custom Author", info.getAuthor()); + assertEquals("Custom Creator", info.getCreator()); + // Producer is set to the label by both setNewDocumentMetadata and + // setCommonMetadata. + assertEquals(LABEL, info.getProducer()); + } + } + + @Test + @DisplayName("replaces 'username' token with the current user when userService present") + void replacesUsernameToken() throws Exception { + ApplicationProperties props = + propsWithCustomMetadata(true, "Report by username", "Creator"); + UserServiceInterface userService = mock(UserServiceInterface.class); + when(userService.getCurrentUsername()).thenReturn("alice"); + + PdfMetadataService service = new PdfMetadataService(props, LABEL, true, userService); + PdfMetadata md = PdfMetadata.builder().title("T").build(); + + try (PDDocument doc = new PDDocument()) { + service.setMetadataToPdf(doc, md, true); + + assertEquals("Report by alice", doc.getDocumentInformation().getAuthor()); + } + } + + @Test + @DisplayName("leaves 'username' token intact when current user is null") + void keepsTokenWhenUsernameNull() throws Exception { + ApplicationProperties props = + propsWithCustomMetadata(true, "Report by username", "Creator"); + UserServiceInterface userService = mock(UserServiceInterface.class); + when(userService.getCurrentUsername()).thenReturn(null); + + PdfMetadataService service = new PdfMetadataService(props, LABEL, true, userService); + PdfMetadata md = PdfMetadata.builder().title("T").build(); + + try (PDDocument doc = new PDDocument()) { + service.setMetadataToPdf(doc, md, true); + + assertEquals("Report by username", doc.getDocumentInformation().getAuthor()); + } + } + + @Test + @DisplayName("custom author applied even without a userService") + void appliesCustomAuthorWithoutUserService() throws Exception { + ApplicationProperties props = propsWithCustomMetadata(true, "Static Author", "Creator"); + PdfMetadataService service = new PdfMetadataService(props, LABEL, true, null); + PdfMetadata md = PdfMetadata.builder().title("T").build(); + + try (PDDocument doc = new PDDocument()) { + service.setMetadataToPdf(doc, md, true); + + assertEquals("Static Author", doc.getDocumentInformation().getAuthor()); + } + } + + @Test + @DisplayName("pro flag without auto-update keeps metadata author and label creator") + void proButAutoUpdateDisabledUsesMetadata() throws Exception { + ApplicationProperties props = + propsWithCustomMetadata(false, "Custom Author", "Custom Creator"); + PdfMetadataService service = new PdfMetadataService(props, LABEL, true, null); + PdfMetadata md = PdfMetadata.builder().author("Metadata Author").title("T").build(); + + try (PDDocument doc = new PDDocument()) { + service.setMetadataToPdf(doc, md, true); + + PDDocumentInformation info = doc.getDocumentInformation(); + assertEquals("Metadata Author", info.getAuthor()); + assertEquals(LABEL, info.getCreator()); + } + } + + @Test + @DisplayName("auto-update enabled but not pro keeps metadata author and label creator") + void autoUpdateButNotProUsesMetadata() throws Exception { + ApplicationProperties props = + propsWithCustomMetadata(true, "Custom Author", "Custom Creator"); + PdfMetadataService service = new PdfMetadataService(props, LABEL, false, null); + PdfMetadata md = PdfMetadata.builder().author("Metadata Author").title("T").build(); + + try (PDDocument doc = new PDDocument()) { + service.setMetadataToPdf(doc, md, true); + + PDDocumentInformation info = doc.getDocumentInformation(); + assertEquals("Metadata Author", info.getAuthor()); + assertEquals(LABEL, info.getCreator()); + } + } + } +} diff --git a/app/common/src/test/java/stirling/software/common/service/PostHogServiceTest.java b/app/common/src/test/java/stirling/software/common/service/PostHogServiceTest.java new file mode 100644 index 0000000000..025ef71a6e --- /dev/null +++ b/app/common/src/test/java/stirling/software/common/service/PostHogServiceTest.java @@ -0,0 +1,441 @@ +package stirling.software.common.service; + +import static org.junit.jupiter.api.Assertions.*; +import static org.mockito.ArgumentMatchers.*; +import static org.mockito.Mockito.*; + +import java.util.HashMap; +import java.util.Map; + +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.ArgumentCaptor; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.mock.env.MockEnvironment; + +import com.posthog.java.PostHog; + +import stirling.software.common.model.ApplicationProperties; + +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class PostHogServiceTest { + + private static final String UUID = "test-uuid-1234"; + private static final String APP_VERSION = "9.9.9"; + + @Mock PostHog postHog; + @Mock UserServiceInterface userService; + + /** Build an ApplicationProperties with analytics/posthog toggled. */ + private ApplicationProperties props(boolean analyticsEnabled) { + ApplicationProperties appProps = new ApplicationProperties(); + appProps.getSystem().setEnableAnalytics(analyticsEnabled); + return appProps; + } + + /** Construct the service under test. */ + private PostHogService newService( + ApplicationProperties appProps, + UserServiceInterface user, + boolean configDirMounted, + MockEnvironment env) { + return new PostHogService( + postHog, UUID, configDirMounted, APP_VERSION, appProps, user, env); + } + + private MockEnvironment env() { + return new MockEnvironment(); + } + + @Nested + @DisplayName("Constructor / captureSystemInfo") + class ConstructorBehavior { + + @Test + @DisplayName("constructor captures system_info when posthog is enabled") + void constructorCapturesWhenEnabled() { + ApplicationProperties appProps = props(true); + + newService(appProps, userService, false, env()); + + verify(postHog).capture(eq(UUID), eq("system_info_captured"), anyMap()); + } + + @Test + @DisplayName("constructor does not capture when analytics disabled") + void constructorNoCaptureWhenDisabled() { + ApplicationProperties appProps = props(false); + + newService(appProps, userService, false, env()); + + verify(postHog, never()).capture(anyString(), anyString(), anyMap()); + } + + @Test + @DisplayName("constructor does not capture when posthog explicitly disabled") + void constructorNoCaptureWhenPosthogOff() { + ApplicationProperties appProps = props(true); + appProps.getSystem().setEnablePosthog(false); + + newService(appProps, userService, false, env()); + + verify(postHog, never()).capture(anyString(), anyString(), anyMap()); + } + + @Test + @DisplayName("constructor swallows exceptions thrown by postHog.capture") + void constructorSwallowsCaptureException() { + ApplicationProperties appProps = props(true); + doThrow(new RuntimeException("boom")) + .when(postHog) + .capture(anyString(), anyString(), anyMap()); + + // Must not propagate; constructor wraps capture in try/catch. + assertDoesNotThrow(() -> newService(appProps, userService, false, env())); + } + + @Test + @DisplayName("constructor works with null userService (optional dependency)") + void constructorWithNullUserService() { + ApplicationProperties appProps = props(true); + + assertDoesNotThrow(() -> newService(appProps, null, false, env())); + verify(postHog).capture(eq(UUID), eq("system_info_captured"), anyMap()); + } + } + + @Nested + @DisplayName("captureEvent") + class CaptureEvent { + + @Test + @DisplayName("captureEvent forwards to postHog when enabled and injects app_version") + void captureEventWhenEnabled() { + ApplicationProperties appProps = props(true); + PostHogService service = newService(appProps, userService, false, env()); + // Reset the constructor's capture so we only assert on captureEvent. + clearInvocations(postHog); + + Map properties = new HashMap<>(); + properties.put("foo", "bar"); + service.captureEvent("my_event", properties); + + @SuppressWarnings("unchecked") + ArgumentCaptor> captor = ArgumentCaptor.forClass(Map.class); + verify(postHog).capture(eq(UUID), eq("my_event"), captor.capture()); + Map sent = captor.getValue(); + assertEquals("bar", sent.get("foo")); + assertEquals(APP_VERSION, sent.get("app_version")); + } + + @Test + @DisplayName("captureEvent is a no-op when analytics disabled") + void captureEventWhenDisabled() { + ApplicationProperties appProps = props(false); + PostHogService service = newService(appProps, userService, false, env()); + clearInvocations(postHog); + + Map properties = new HashMap<>(); + service.captureEvent("my_event", properties); + + verify(postHog, never()).capture(anyString(), anyString(), anyMap()); + // app_version must not be added when disabled (early return). + assertFalse(properties.containsKey("app_version")); + } + + @Test + @DisplayName("captureEvent adds app_version key to the provided map") + void captureEventMutatesMap() { + ApplicationProperties appProps = props(true); + PostHogService service = newService(appProps, userService, false, env()); + clearInvocations(postHog); + + Map properties = new HashMap<>(); + service.captureEvent("evt", properties); + + assertTrue(properties.containsKey("app_version")); + assertEquals(APP_VERSION, properties.get("app_version")); + } + } + + @Nested + @DisplayName("captureServerMetrics") + class CaptureServerMetrics { + + private PostHogService disabledService() { + // Keep posthog disabled so the constructor performs no capture; metrics + // methods are independent of the enabled flag. + return newService(props(false), userService, true, env()); + } + + @Test + @DisplayName("includes core application and system metrics") + void includesCoreMetrics() { + PostHogService service = disabledService(); + + Map metrics = service.captureServerMetrics(); + + assertEquals(APP_VERSION, metrics.get("app_version")); + assertEquals(true, metrics.get("mounted_config_dir")); + assertNotNull(metrics.get("os_name")); + assertNotNull(metrics.get("java_version")); + assertTrue(metrics.containsKey("cpu_cores")); + assertTrue(metrics.containsKey("total_memory")); + assertTrue(metrics.containsKey("free_memory")); + assertTrue(metrics.containsKey("process_id")); + assertTrue(metrics.containsKey("jvm_uptime_ms")); + assertTrue(metrics.containsKey("thread_count")); + } + + @Test + @DisplayName("deployment_type defaults to JAR when not docker/exe") + void deploymentTypeJar() { + PostHogService service = disabledService(); + + Map metrics = service.captureServerMetrics(); + + // In the unit-test environment there is no /.dockerenv and no BROWSER_OPEN. + assertEquals("JAR", metrics.get("deployment_type")); + } + + @Test + @DisplayName("deployment_type becomes EXE when BROWSER_OPEN=true") + void deploymentTypeExe() { + MockEnvironment environment = env(); + environment.setProperty("BROWSER_OPEN", "true"); + PostHogService service = newService(props(false), userService, false, environment); + + Map metrics = service.captureServerMetrics(); + + assertEquals("EXE", metrics.get("deployment_type")); + } + + @Test + @DisplayName("BROWSER_OPEN matching is case-insensitive") + void deploymentTypeExeCaseInsensitive() { + MockEnvironment environment = env(); + environment.setProperty("BROWSER_OPEN", "TRUE"); + PostHogService service = newService(props(false), userService, false, environment); + + Map metrics = service.captureServerMetrics(); + + assertEquals("EXE", metrics.get("deployment_type")); + } + + @Test + @DisplayName("mounted_config_dir reflects the configDirMounted flag") + void mountedConfigDirFalse() { + PostHogService service = newService(props(false), userService, false, env()); + + Map metrics = service.captureServerMetrics(); + + assertEquals(false, metrics.get("mounted_config_dir")); + } + + @Test + @DisplayName("includes total_users_created when userService present") + void includesUserCountWhenUserServicePresent() { + when(userService.getTotalUsersCount()).thenReturn(42L); + PostHogService service = newService(props(false), userService, false, env()); + + Map metrics = service.captureServerMetrics(); + + assertEquals(42L, metrics.get("total_users_created")); + } + + @Test + @DisplayName("omits total_users_created when userService is null") + void omitsUserCountWhenUserServiceNull() { + PostHogService service = newService(props(false), null, false, env()); + + Map metrics = service.captureServerMetrics(); + + assertFalse(metrics.containsKey("total_users_created")); + } + + @Test + @DisplayName("always embeds nested application_properties map") + void embedsApplicationProperties() { + PostHogService service = disabledService(); + + Map metrics = service.captureServerMetrics(); + + assertTrue(metrics.get("application_properties") instanceof Map); + } + } + + @Nested + @DisplayName("captureApplicationProperties") + class CaptureApplicationProperties { + + private PostHogService serviceWith(ApplicationProperties appProps) { + // Disable analytics to keep the constructor from capturing. + appProps.getSystem().setEnableAnalytics(false); + return newService(appProps, userService, false, env()); + } + + @Test + @DisplayName("includes blank-trimmed legal strings only when non-empty") + void legalPropertiesFiltered() { + ApplicationProperties appProps = new ApplicationProperties(); + appProps.getLegal().setTermsAndConditions(" https://terms "); + appProps.getLegal().setPrivacyPolicy(""); // blank -> skipped + PostHogService service = serviceWith(appProps); + + Map p = service.captureApplicationProperties(); + + // String values are trimmed by addIfNotEmpty. + assertEquals("https://terms", p.get("legal_termsAndConditions")); + assertFalse(p.containsKey("legal_privacyPolicy")); + assertFalse(p.containsKey("legal_accessibilityStatement")); + } + + @Test + @DisplayName("always reports csrfDisabled true and login booleans") + void securityProperties() { + ApplicationProperties appProps = new ApplicationProperties(); + appProps.getSecurity().setEnableLogin(true); + appProps.getSecurity().setLoginAttemptCount(5); + appProps.getSecurity().setLoginResetTimeMinutes(10); + PostHogService service = serviceWith(appProps); + + Map p = service.captureApplicationProperties(); + + assertEquals(true, p.get("security_csrfDisabled")); + assertEquals(true, p.get("security_enableLogin")); + assertEquals(5, p.get("security_loginAttemptCount")); + assertEquals(10L, p.get("security_loginResetTimeMinutes")); + assertEquals("all", p.get("security_loginMethod")); + } + + @Test + @DisplayName("oauth2 nested fields are omitted when oauth2 disabled") + void oauth2DisabledOmitsNested() { + ApplicationProperties appProps = new ApplicationProperties(); + // oauth2.enabled defaults to false. + PostHogService service = serviceWith(appProps); + + Map p = service.captureApplicationProperties(); + + assertEquals(false, p.get("security_oauth2_enabled")); + assertFalse(p.containsKey("security_oauth2_autoCreateUser")); + assertFalse(p.containsKey("security_oauth2_provider")); + } + + @Test + @DisplayName("oauth2 nested fields are included when oauth2 enabled") + void oauth2EnabledIncludesNested() { + ApplicationProperties appProps = new ApplicationProperties(); + appProps.getSecurity().getOauth2().setEnabled(true); + appProps.getSecurity().getOauth2().setAutoCreateUser(true); + appProps.getSecurity().getOauth2().setBlockRegistration(false); + appProps.getSecurity().getOauth2().setUseAsUsername("email"); + appProps.getSecurity().getOauth2().setProvider("google"); + PostHogService service = serviceWith(appProps); + + Map p = service.captureApplicationProperties(); + + assertEquals(true, p.get("security_oauth2_enabled")); + assertEquals(true, p.get("security_oauth2_autoCreateUser")); + assertEquals(false, p.get("security_oauth2_blockRegistration")); + assertEquals("email", p.get("security_oauth2_useAsUsername")); + assertEquals("google", p.get("security_oauth2_provider")); + } + + @Test + @DisplayName("system analytics/posthog/scarf booleans are reported") + void systemAnalyticsBooleans() { + ApplicationProperties appProps = new ApplicationProperties(); + appProps.getSystem().setEnableAnalytics(true); + appProps.getSystem().setEnablePosthog(true); + appProps.getSystem().setEnableScarf(false); + appProps.getSystem().setDefaultLocale("en-US"); + PostHogService service = newService(appProps, userService, false, env()); + // Constructor will capture once because analytics is enabled; that's fine. + clearInvocations(postHog); + + Map p = service.captureApplicationProperties(); + + assertEquals("en-US", p.get("system_defaultLocale")); + assertEquals(true, p.get("system_enableAnalytics")); + assertEquals(true, p.get("system_enablePosthog")); + // isScarfEnabled() is false because enableScarf is false. + assertEquals(false, p.get("system_enableScarf")); + } + + @Test + @DisplayName("metrics_enabled and autoPipeline output folder included appropriately") + void metricsAndAutoPipeline() { + ApplicationProperties appProps = new ApplicationProperties(); + appProps.getMetrics().setEnabled(true); + appProps.getAutoPipeline().setOutputFolder("/tmp/out"); + PostHogService service = serviceWith(appProps); + + Map p = service.captureApplicationProperties(); + + assertEquals(true, p.get("metrics_enabled")); + assertEquals("/tmp/out", p.get("autoPipeline_outputFolder")); + } + + @Test + @DisplayName("enterprise metadata flag omitted when premium disabled") + void premiumDisabledOmitsMetadata() { + ApplicationProperties appProps = new ApplicationProperties(); + // premium.enabled defaults to false. + PostHogService service = serviceWith(appProps); + + Map p = service.captureApplicationProperties(); + + assertEquals(false, p.get("enterpriseEdition_enabled")); + assertFalse(p.containsKey("enterpriseEdition_customMetadata_autoUpdateMetadata")); + } + + @Test + @DisplayName("enterprise metadata flag included when premium enabled") + void premiumEnabledIncludesMetadata() { + ApplicationProperties appProps = new ApplicationProperties(); + appProps.getPremium().setEnabled(true); + appProps.getPremium().getProFeatures().getCustomMetadata().setAutoUpdateMetadata(true); + PostHogService service = serviceWith(appProps); + + Map p = service.captureApplicationProperties(); + + assertEquals(true, p.get("enterpriseEdition_enabled")); + assertEquals(true, p.get("enterpriseEdition_customMetadata_autoUpdateMetadata")); + } + + @Test + @DisplayName("ui appNameNavbar omitted when blank, included when set") + void uiAppNameNavbar() { + ApplicationProperties blankProps = new ApplicationProperties(); + // appNameNavbar getter returns null for blank/empty values. + PostHogService blankService = serviceWith(blankProps); + Map blank = blankService.captureApplicationProperties(); + assertFalse(blank.containsKey("ui_appNameNavbar")); + + ApplicationProperties namedProps = new ApplicationProperties(); + namedProps.getUi().setAppNameNavbar("My App"); + PostHogService namedService = serviceWith(namedProps); + Map named = namedService.captureApplicationProperties(); + assertEquals("My App", named.get("ui_appNameNavbar")); + } + + @Test + @DisplayName("returns a non-null map for a fresh ApplicationProperties") + void defaultsProduceNonNullMap() { + PostHogService service = serviceWith(new ApplicationProperties()); + + Map p = service.captureApplicationProperties(); + + assertNotNull(p); + // csrfDisabled is always added regardless of config, so map is never empty. + assertTrue(p.containsKey("security_csrfDisabled")); + } + } +} diff --git a/app/common/src/test/java/stirling/software/common/util/ExceptionUtilsGapTest.java b/app/common/src/test/java/stirling/software/common/util/ExceptionUtilsGapTest.java new file mode 100644 index 0000000000..8268027484 --- /dev/null +++ b/app/common/src/test/java/stirling/software/common/util/ExceptionUtilsGapTest.java @@ -0,0 +1,738 @@ +package stirling.software.common.util; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertInstanceOf; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertSame; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.Mockito.mockStatic; + +import java.io.IOException; +import java.util.List; + +import org.apache.pdfbox.pdmodel.PDPage; +import org.apache.pdfbox.pdmodel.common.PDRectangle; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.mockito.MockedStatic; + +import stirling.software.common.util.ExceptionUtils.BaseAppException; +import stirling.software.common.util.ExceptionUtils.CbrFormatException; +import stirling.software.common.util.ExceptionUtils.CbzFormatException; +import stirling.software.common.util.ExceptionUtils.EmlFormatException; +import stirling.software.common.util.ExceptionUtils.ErrorCode; +import stirling.software.common.util.ExceptionUtils.FfmpegRequiredException; +import stirling.software.common.util.ExceptionUtils.GhostscriptException; +import stirling.software.common.util.ExceptionUtils.OutOfMemoryDpiException; +import stirling.software.common.util.ExceptionUtils.PdfCorruptedException; + +/** + * Additional gap-filling unit tests for {@link ExceptionUtils}, covering areas not exercised by + * {@code ExceptionUtilsTest}: CBR/CBZ/EML factories, error-code hint/action lookups, rendering + * dimension validation, OOM rendering wrappers, Ghostscript output analysis, and wrapException. + * + *

The {@code messages} ResourceBundle is not on the common module test classpath, so {@link + * ExceptionUtils} falls back to the default messages baked into {@link ErrorCode}. Assertions here + * rely only on those default messages and on deterministic structural behavior. + */ +class ExceptionUtilsGapTest { + + @Nested + @DisplayName("ErrorCode enum metadata") + class ErrorCodeMetadataTests { + + @Test + @DisplayName("each error code exposes code, message key and default message") + void allErrorCodesHaveMetadata() { + for (ErrorCode code : ErrorCode.values()) { + assertNotNull(code.getCode(), "code for " + code); + assertTrue(code.getCode().startsWith("E"), "code prefix for " + code); + assertNotNull(code.getMessageKey(), "messageKey for " + code); + assertNotNull(code.getDefaultMessage(), "defaultMessage for " + code); + assertFalse(code.getDefaultMessage().isEmpty(), "defaultMessage empty for " + code); + } + } + + @Test + @DisplayName("known error codes map to expected identifiers") + void knownErrorCodeIdentifiers() { + assertEquals("E001", ErrorCode.PDF_CORRUPTED.getCode()); + assertEquals("E081", ErrorCode.OUT_OF_MEMORY_DPI.getCode()); + assertEquals("error.pdfCorrupted", ErrorCode.PDF_CORRUPTED.getMessageKey()); + } + } + + @Nested + @DisplayName("Hints and action lookups via resource bundle") + class HintAndActionTests { + + @Test + @DisplayName("getHintsForErrorCode returns empty list for null code") + void hintsNullCode() { + assertEquals(List.of(), ExceptionUtils.getHintsForErrorCode(null)); + } + + @Test + @DisplayName("getHintsForErrorCode returns empty list when no hints exist in bundle") + void hintsMissingFromBundle() { + // Fallback empty bundle has no hint keys, so the result is an empty list. + List hints = ExceptionUtils.getHintsForErrorCode("E001"); + assertNotNull(hints); + assertTrue(hints.isEmpty()); + } + + @Test + @DisplayName("getActionRequiredForErrorCode returns null for null code") + void actionNullCode() { + assertNull(ExceptionUtils.getActionRequiredForErrorCode(null)); + } + + @Test + @DisplayName("getActionRequiredForErrorCode returns null when key absent from bundle") + void actionMissingFromBundle() { + assertNull(ExceptionUtils.getActionRequiredForErrorCode("E001")); + } + } + + @Nested + @DisplayName("CBR format exception factories") + class CbrFactoryTests { + + @Test + @DisplayName("invalid format uses provided message when non-null") + void cbrInvalidFormatWithMessage() { + CbrFormatException ex = + ExceptionUtils.createCbrInvalidFormatException("custom cbr msg"); + assertEquals("custom cbr msg", ex.getMessage()); + assertEquals(ErrorCode.CBR_INVALID_FORMAT.getCode(), ex.getErrorCode()); + } + + @Test + @DisplayName("invalid format falls back to default message when null") + void cbrInvalidFormatNullMessage() { + CbrFormatException ex = ExceptionUtils.createCbrInvalidFormatException(null); + assertTrue(ex.getMessage().contains("CBR/RAR archive")); + assertEquals("E010", ex.getErrorCode()); + } + + @Test + @DisplayName("encrypted CBR reuses invalid-format code") + void cbrEncrypted() { + CbrFormatException ex = ExceptionUtils.createCbrEncryptedException(); + assertEquals(ErrorCode.CBR_INVALID_FORMAT.getCode(), ex.getErrorCode()); + } + + @Test + @DisplayName("no images and corrupted images both map to CBR_NO_IMAGES") + void cbrNoImages() { + CbrFormatException noImages = ExceptionUtils.createCbrNoImagesException(); + CbrFormatException corrupted = ExceptionUtils.createCbrCorruptedImagesException(); + assertEquals(ErrorCode.CBR_NO_IMAGES.getCode(), noImages.getErrorCode()); + assertEquals(ErrorCode.CBR_NO_IMAGES.getCode(), corrupted.getErrorCode()); + assertTrue(noImages.getMessage().contains("No valid images")); + } + + @Test + @DisplayName("not-a-CBR file uses CBR_NOT_CBR code") + void notCbr() { + CbrFormatException ex = ExceptionUtils.createNotCbrFileException(); + assertEquals(ErrorCode.CBR_NOT_CBR.getCode(), ex.getErrorCode()); + assertTrue(ex.getMessage().contains("CBR or RAR")); + } + + @Test + @DisplayName( + "CbrFormatException is an IllegalArgumentException via BaseValidationException") + void cbrIsIllegalArgument() { + CbrFormatException ex = ExceptionUtils.createNotCbrFileException(); + assertInstanceOf(IllegalArgumentException.class, ex); + } + } + + @Nested + @DisplayName("CBZ format exception factories") + class CbzFactoryTests { + + @Test + @DisplayName("invalid format wraps cause and uses CBZ_INVALID_FORMAT code") + void cbzInvalidFormat() { + Exception cause = new Exception("zip boom"); + CbzFormatException ex = ExceptionUtils.createCbzInvalidFormatException(cause); + assertSame(cause, ex.getCause()); + assertEquals(ErrorCode.CBZ_INVALID_FORMAT.getCode(), ex.getErrorCode()); + assertTrue(ex.getMessage().contains("CBZ/ZIP archive")); + } + + @Test + @DisplayName("empty CBZ reuses invalid-format code") + void cbzEmpty() { + CbzFormatException ex = ExceptionUtils.createCbzEmptyException(); + assertEquals(ErrorCode.CBZ_INVALID_FORMAT.getCode(), ex.getErrorCode()); + } + + @Test + @DisplayName("no images and corrupted images both map to CBZ_NO_IMAGES") + void cbzNoImages() { + CbzFormatException noImages = ExceptionUtils.createCbzNoImagesException(); + CbzFormatException corrupted = ExceptionUtils.createCbzCorruptedImagesException(); + assertEquals(ErrorCode.CBZ_NO_IMAGES.getCode(), noImages.getErrorCode()); + assertEquals(ErrorCode.CBZ_NO_IMAGES.getCode(), corrupted.getErrorCode()); + } + + @Test + @DisplayName("not-a-CBZ file uses CBZ_NOT_CBZ code") + void notCbz() { + CbzFormatException ex = ExceptionUtils.createNotCbzFileException(); + assertEquals(ErrorCode.CBZ_NOT_CBZ.getCode(), ex.getErrorCode()); + assertTrue(ex.getMessage().contains("CBZ or ZIP")); + } + } + + @Nested + @DisplayName("EML format exception factories") + class EmlFactoryTests { + + @Test + @DisplayName("empty EML uses EML_EMPTY code") + void emlEmpty() { + EmlFormatException ex = ExceptionUtils.createEmlEmptyException(); + assertEquals(ErrorCode.EML_EMPTY.getCode(), ex.getErrorCode()); + assertTrue(ex.getMessage().contains("EML file is empty")); + } + + @Test + @DisplayName("invalid EML uses EML_INVALID_FORMAT code") + void emlInvalid() { + EmlFormatException ex = ExceptionUtils.createEmlInvalidFormatException(); + assertEquals(ErrorCode.EML_INVALID_FORMAT.getCode(), ex.getErrorCode()); + assertTrue(ex.getMessage().contains("Invalid EML")); + } + } + + @Nested + @DisplayName("Image, OCR and processing factories") + class ImageOcrProcessingTests { + + @Test + @DisplayName("image read exception embeds filename and has no cause") + void imageRead() { + IOException ex = ExceptionUtils.createImageReadException("photo.png"); + assertTrue(ex.getMessage().contains("photo.png")); + assertNull(ex.getCause()); + } + + @Test + @DisplayName("image read exception rejects null filename") + void imageReadNullFilename() { + assertThrows( + IllegalArgumentException.class, + () -> ExceptionUtils.createImageReadException(null)); + } + + @Test + @DisplayName("ocr invalid render type uses default message") + void ocrInvalidRenderType() { + IOException ex = ExceptionUtils.createOcrInvalidRenderTypeException(); + assertTrue(ex.getMessage().contains("hocr")); + } + + @Test + @DisplayName("ocr processing failed includes return code") + void ocrProcessingFailed() { + IOException ex = ExceptionUtils.createOcrProcessingFailedException(7); + assertTrue(ex.getMessage().contains("7")); + } + + @Test + @DisplayName("processing interrupted wraps the InterruptedException cause") + void processingInterrupted() { + InterruptedException cause = new InterruptedException("stop"); + IOException ex = + ExceptionUtils.createProcessingInterruptedException("compression", cause); + assertSame(cause, ex.getCause()); + assertTrue(ex.getMessage().contains("compression")); + } + + @Test + @DisplayName("processing interrupted rejects null arguments") + void processingInterruptedNullArgs() { + assertThrows( + IllegalArgumentException.class, + () -> + ExceptionUtils.createProcessingInterruptedException( + null, new InterruptedException())); + assertThrows( + IllegalArgumentException.class, + () -> ExceptionUtils.createProcessingInterruptedException("x", null)); + } + + @Test + @DisplayName("ghostscript conversion exception embeds output type") + void ghostscriptConversion() { + IOException ex = ExceptionUtils.createGhostscriptConversionException("png"); + assertNotNull(ex.getMessage()); + assertThrows( + IllegalArgumentException.class, + () -> ExceptionUtils.createGhostscriptConversionException(null)); + } + } + + @Nested + @DisplayName("Validation factories: page size, file, ffmpeg") + class ValidationFactoryTests { + + @Test + @DisplayName("invalid page size rejects null size") + void invalidPageSizeNull() { + assertThrows( + IllegalArgumentException.class, + () -> ExceptionUtils.createInvalidPageSizeException(null)); + } + + @Test + @DisplayName("file null-or-empty uses FILE_NULL_OR_EMPTY default message") + void fileNullOrEmpty() { + IllegalArgumentException ex = ExceptionUtils.createFileNullOrEmptyException(); + assertTrue(ex.getMessage().contains("null or empty")); + } + + @Test + @DisplayName("file no-name uses FILE_NO_NAME default message") + void fileNoName() { + IllegalArgumentException ex = ExceptionUtils.createFileNoNameException(); + assertTrue(ex.getMessage().contains("must have a name")); + } + + @Test + @DisplayName("pdf no-pages uses PDF_NO_PAGES default message") + void pdfNoPages() { + IllegalArgumentException ex = ExceptionUtils.createPdfNoPages(); + assertTrue(ex.getMessage().contains("no pages")); + } + + @Test + @DisplayName("ffmpeg required exception exposes FFMPEG_REQUIRED code and null cause") + void ffmpegRequired() { + FfmpegRequiredException ex = ExceptionUtils.createFfmpegRequiredException(); + assertEquals(ErrorCode.FFMPEG_REQUIRED.getCode(), ex.getErrorCode()); + assertNull(ex.getCause()); + assertTrue(ex.getMessage().contains("FFmpeg")); + } + } + + @Nested + @DisplayName("ErrorCode-based argument and IO factories") + class ErrorCodeArgFactoryTests { + + @Test + @DisplayName("createIllegalArgumentException(ErrorCode, args) formats default message") + void illegalArgumentFromErrorCode() { + IllegalArgumentException ex = + ExceptionUtils.createIllegalArgumentException( + ErrorCode.INVALID_PAGE_SIZE, "B7"); + assertTrue(ex.getMessage().contains("B7")); + } + + @Test + @DisplayName("createIllegalArgumentException rejects null ErrorCode") + void illegalArgumentFromNullErrorCode() { + ErrorCode nullCode = null; + assertThrows( + IllegalArgumentException.class, + () -> ExceptionUtils.createIllegalArgumentException(nullCode)); + } + + @Test + @DisplayName("createFileProcessingException rejects null operation and cause") + void fileProcessingNullArgs() { + assertThrows( + IllegalArgumentException.class, + () -> ExceptionUtils.createFileProcessingException(null, new Exception())); + assertThrows( + IllegalArgumentException.class, + () -> ExceptionUtils.createFileProcessingException("op", null)); + } + + @Test + @DisplayName("createInvalidArgumentException rejects null name or value") + void invalidArgumentNullArgs() { + assertThrows( + IllegalArgumentException.class, + () -> ExceptionUtils.createInvalidArgumentException(null, "v")); + assertThrows( + IllegalArgumentException.class, + () -> ExceptionUtils.createInvalidArgumentException("n", null)); + } + + @Test + @DisplayName("createNullArgumentException rejects null argument name") + void nullArgumentNullName() { + assertThrows( + IllegalArgumentException.class, + () -> ExceptionUtils.createNullArgumentException(null)); + } + + @Test + @DisplayName("createIOException without cause leaves cause null") + void ioExceptionWithoutCause() { + IOException ex = ExceptionUtils.createIOException("key", "msg {0}", null, "A"); + assertEquals("msg A", ex.getMessage()); + assertNull(ex.getCause()); + } + + @Test + @DisplayName("createRuntimeException without cause leaves cause null") + void runtimeExceptionWithoutCause() { + RuntimeException ex = + ExceptionUtils.createRuntimeException("key", "msg {0}", null, "B"); + assertEquals("msg B", ex.getMessage()); + assertNull(ex.getCause()); + } + } + + @Nested + @DisplayName("createPdfCorruptedException null-cause handling") + class PdfCorruptedCauseTests { + + @Test + @DisplayName("rejects null cause") + void nullCause() { + assertThrows( + IllegalArgumentException.class, + () -> ExceptionUtils.createPdfCorruptedException("ctx", null)); + } + + @Test + @DisplayName("empty context behaves like no context") + void emptyContext() { + PdfCorruptedException ex = + ExceptionUtils.createPdfCorruptedException("", new Exception("x")); + assertTrue(ex.getMessage().contains("PDF file appears to be corrupted")); + assertEquals(ErrorCode.PDF_CORRUPTED.getCode(), ex.getErrorCode()); + } + } + + @Nested + @DisplayName("validateRenderingDimensions") + class ValidateRenderingDimensionsTests { + + @Test + @DisplayName("null page is a no-op") + void nullPage() { + // Should simply return without throwing. + org.junit.jupiter.api.Assertions.assertDoesNotThrow( + () -> ExceptionUtils.validateRenderingDimensions(null, 1, 300)); + } + + @Test + @DisplayName("normal letter-size page at 150 DPI passes validation") + void normalPagePasses() { + PDPage page = new PDPage(PDRectangle.LETTER); + org.junit.jupiter.api.Assertions.assertDoesNotThrow( + () -> ExceptionUtils.validateRenderingDimensions(page, 1, 150)); + } + + @Test + @DisplayName("page with zero DPI yields zero pixels and passes") + void zeroDpiPasses() { + PDPage page = new PDPage(PDRectangle.A4); + org.junit.jupiter.api.Assertions.assertDoesNotThrow( + () -> ExceptionUtils.validateRenderingDimensions(page, 2, 0)); + } + } + + @Nested + @DisplayName("handleOomRendering wrappers") + class HandleOomRenderingTests { + + @Test + @DisplayName("returns operation result on success (with page number)") + void successWithPage() throws IOException { + String result = ExceptionUtils.handleOomRendering(3, 300, () -> "ok"); + assertEquals("ok", result); + } + + @Test + @DisplayName("returns operation result on success (no page number)") + void successNoPage() throws IOException { + String result = ExceptionUtils.handleOomRendering(300, () -> "fine"); + assertEquals("fine", result); + } + + @Test + @DisplayName("propagates IOException from the operation unchanged") + void propagatesIoException() { + IOException boom = new IOException("io boom"); + IOException thrown = + assertThrows( + IOException.class, + () -> + ExceptionUtils.handleOomRendering( + 1, + 300, + () -> { + throw boom; + })); + assertSame(boom, thrown); + } + + @Test + @DisplayName("converts OutOfMemoryError to OutOfMemoryDpiException (with page)") + void oomToDpiExceptionWithPage() { + OutOfMemoryDpiException thrown = + assertThrows( + OutOfMemoryDpiException.class, + () -> + ExceptionUtils.handleOomRendering( + 5, + 300, + () -> { + throw new OutOfMemoryError("heap"); + })); + assertEquals(ErrorCode.OUT_OF_MEMORY_DPI.getCode(), thrown.getErrorCode()); + assertInstanceOf(OutOfMemoryError.class, thrown.getCause()); + } + + @Test + @DisplayName("converts NegativeArraySizeException to OutOfMemoryDpiException (no page)") + void negativeArraySizeToDpiExceptionNoPage() { + OutOfMemoryDpiException thrown = + assertThrows( + OutOfMemoryDpiException.class, + () -> + ExceptionUtils.handleOomRendering( + 300, + () -> { + throw new NegativeArraySizeException("-1"); + })); + assertEquals(ErrorCode.OUT_OF_MEMORY_DPI.getCode(), thrown.getErrorCode()); + assertInstanceOf(NegativeArraySizeException.class, thrown.getCause()); + } + } + + @Nested + @DisplayName("createOutOfMemoryDpiException overloads") + class OutOfMemoryDpiFactoryTests { + + @Test + @DisplayName("page + dpi + Throwable wraps cause and sets code") + void pageDpiThrowable() { + Throwable cause = new IllegalStateException("too big"); + OutOfMemoryDpiException ex = + ExceptionUtils.createOutOfMemoryDpiException(4, 600, cause); + assertSame(cause, ex.getCause()); + assertEquals(ErrorCode.OUT_OF_MEMORY_DPI.getCode(), ex.getErrorCode()); + } + + @Test + @DisplayName("page + dpi + OutOfMemoryError overload wraps the error") + void pageDpiOomError() { + OutOfMemoryError cause = new OutOfMemoryError("oom"); + OutOfMemoryDpiException ex = + ExceptionUtils.createOutOfMemoryDpiException(2, 300, cause); + assertSame(cause, ex.getCause()); + } + + @Test + @DisplayName("dpi + Throwable overload wraps cause") + void dpiThrowable() { + Throwable cause = new RuntimeException("x"); + OutOfMemoryDpiException ex = ExceptionUtils.createOutOfMemoryDpiException(300, cause); + assertSame(cause, ex.getCause()); + assertEquals(ErrorCode.OUT_OF_MEMORY_DPI.getCode(), ex.getErrorCode()); + } + + @Test + @DisplayName("dpi + OutOfMemoryError overload wraps the error") + void dpiOomError() { + OutOfMemoryError cause = new OutOfMemoryError("oom"); + OutOfMemoryDpiException ex = ExceptionUtils.createOutOfMemoryDpiException(300, cause); + assertSame(cause, ex.getCause()); + } + + @Test + @DisplayName("rejects null cause") + void nullCause() { + assertThrows( + IllegalArgumentException.class, + () -> ExceptionUtils.createOutOfMemoryDpiException(1, 300, (Throwable) null)); + } + } + + @Nested + @DisplayName("Ghostscript output analysis") + class GhostscriptAnalysisTests { + + @Test + @DisplayName("null/blank output produces generic compression exception") + void blankOutput() { + GhostscriptException ex = ExceptionUtils.createGhostscriptCompressionException(" "); + assertEquals(ErrorCode.GHOSTSCRIPT_COMPRESSION.getCode(), ex.getErrorCode()); + } + + @Test + @DisplayName("recognized page drawing error yields page-drawing error code") + void pageDrawingError() { + String output = "Page 3\nERROR: page drawing error encountered while processing"; + GhostscriptException ex = ExceptionUtils.createGhostscriptCompressionException(output); + assertEquals(ErrorCode.GHOSTSCRIPT_PAGE_DRAWING.getCode(), ex.getErrorCode()); + } + + @Test + @DisplayName("non-page-drawing output falls back to compression error code") + void unrecognizedOutput() { + String output = "Some random ghostscript chatter that is not an error marker"; + GhostscriptException ex = ExceptionUtils.createGhostscriptCompressionException(output); + assertEquals(ErrorCode.GHOSTSCRIPT_COMPRESSION.getCode(), ex.getErrorCode()); + } + + @Test + @DisplayName("detectGhostscriptCriticalError returns exception only for critical output") + void detectCritical() { + GhostscriptException critical = + ExceptionUtils.detectGhostscriptCriticalError( + "Page 1\ncould not draw this page"); + assertNotNull(critical); + assertEquals(ErrorCode.GHOSTSCRIPT_PAGE_DRAWING.getCode(), critical.getErrorCode()); + } + + @Test + @DisplayName("detectGhostscriptCriticalError returns null for non-critical output") + void detectNonCritical() { + assertNull(ExceptionUtils.detectGhostscriptCriticalError("just informational output")); + assertNull(ExceptionUtils.detectGhostscriptCriticalError(null)); + } + + @Test + @DisplayName("compression exception derived from cause message") + void compressionFromCauseMessage() { + GhostscriptException ex = + ExceptionUtils.createGhostscriptCompressionException( + new Exception("Page 2\npage drawing error")); + assertEquals(ErrorCode.GHOSTSCRIPT_PAGE_DRAWING.getCode(), ex.getErrorCode()); + } + + @Test + @DisplayName("createGhostscriptCompressionException rejects null cause overload") + void compressionNullCause() { + assertThrows( + IllegalArgumentException.class, + () -> ExceptionUtils.createGhostscriptCompressionException((Exception) null)); + } + + @Test + @DisplayName("multiple affected pages are summarized in the message") + void multiplePages() { + String output = "Page 1\npage drawing error\nPage 2\ncould not draw this page"; + GhostscriptException ex = ExceptionUtils.createGhostscriptCompressionException(output); + assertEquals(ErrorCode.GHOSTSCRIPT_PAGE_DRAWING.getCode(), ex.getErrorCode()); + assertNotNull(ex.getMessage()); + } + } + + @Nested + @DisplayName("wrapException") + class WrapExceptionTests { + + @Test + @DisplayName("RuntimeException is returned unchanged") + void runtimePassthrough() { + RuntimeException original = new IllegalStateException("boom"); + RuntimeException wrapped = ExceptionUtils.wrapException(original, "merge"); + assertSame(original, wrapped); + } + + @Test + @DisplayName("BaseAppException (IOException subtype) is wrapped in a RuntimeException") + void baseAppExceptionWrapped() { + // A corrupted-pdf IOException triggers handlePdfException -> PdfCorruptedException. + IOException corrupted = new IOException("Invalid PDF"); + try (MockedStatic mock = mockStatic(PdfErrorUtils.class)) { + mock.when(() -> PdfErrorUtils.isCorruptedPdfError(corrupted)).thenReturn(true); + RuntimeException wrapped = ExceptionUtils.wrapException(corrupted, "merge"); + assertInstanceOf(BaseAppException.class, wrapped.getCause()); + } + } + + @Test + @DisplayName("plain IOException is wrapped via file-processing exception") + void plainIoExceptionWrapped() { + IOException io = new IOException("disk full"); + try (MockedStatic mock = mockStatic(PdfErrorUtils.class)) { + mock.when(() -> PdfErrorUtils.isCorruptedPdfError(io)).thenReturn(false); + RuntimeException wrapped = ExceptionUtils.wrapException(io, "split"); + assertInstanceOf(IOException.class, wrapped.getCause()); + assertFalse(wrapped.getCause() instanceof BaseAppException); + } + } + + @Test + @DisplayName("checked non-IO exception is wrapped with operation context") + void checkedExceptionWrapped() { + Exception checked = new Exception("oops"); + RuntimeException wrapped = ExceptionUtils.wrapException(checked, "convert"); + assertSame(checked, wrapped.getCause()); + assertTrue(wrapped.getMessage().contains("convert")); + assertTrue(wrapped.getMessage().contains("oops")); + } + + @Test + @DisplayName("rejects null exception or operation") + void wrapNullArgs() { + assertThrows( + IllegalArgumentException.class, () -> ExceptionUtils.wrapException(null, "op")); + assertThrows( + IllegalArgumentException.class, + () -> ExceptionUtils.wrapException(new Exception(), null)); + } + } + + @Nested + @DisplayName("logException return value and handlePdfException null guard") + class LogAndHandleTests { + + @Test + @DisplayName("logException returns the same exception instance for fluent throw") + void logExceptionReturnsSame() { + Exception e = new RuntimeException("unexpected"); + try (MockedStatic mock = mockStatic(PdfErrorUtils.class)) { + mock.when(() -> PdfErrorUtils.isCorruptedPdfError(e)).thenReturn(false); + Exception returned = ExceptionUtils.logException("op", e); + assertSame(e, returned); + } + } + + @Test + @DisplayName("logException rejects null operation or exception") + void logExceptionNullArgs() { + assertThrows( + IllegalArgumentException.class, + () -> ExceptionUtils.logException(null, new Exception())); + assertThrows( + IllegalArgumentException.class, () -> ExceptionUtils.logException("op", null)); + } + + @Test + @DisplayName("handlePdfException rejects null exception") + void handlePdfNull() { + assertThrows( + IllegalArgumentException.class, () -> ExceptionUtils.handlePdfException(null)); + } + + @Test + @DisplayName("handlePdfException with context wraps corrupted PDF and includes context") + void handlePdfWithContext() { + IOException original = new IOException("damaged"); + try (MockedStatic mock = mockStatic(PdfErrorUtils.class)) { + mock.when(() -> PdfErrorUtils.isCorruptedPdfError(original)).thenReturn(true); + IOException result = ExceptionUtils.handlePdfException(original, "during merge"); + assertInstanceOf(PdfCorruptedException.class, result); + assertTrue(result.getMessage().contains("during merge")); + } + } + } +} diff --git a/app/common/src/test/java/stirling/software/common/util/FormUtilsGapTest.java b/app/common/src/test/java/stirling/software/common/util/FormUtilsGapTest.java new file mode 100644 index 0000000000..dd828327b1 --- /dev/null +++ b/app/common/src/test/java/stirling/software/common/util/FormUtilsGapTest.java @@ -0,0 +1,847 @@ +package stirling.software.common.util; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.io.IOException; +import java.util.ArrayList; +import java.util.List; +import java.util.Map; + +import org.apache.pdfbox.cos.COSDictionary; +import org.apache.pdfbox.cos.COSName; +import org.apache.pdfbox.pdmodel.PDDocument; +import org.apache.pdfbox.pdmodel.PDPage; +import org.apache.pdfbox.pdmodel.PDResources; +import org.apache.pdfbox.pdmodel.common.PDRectangle; +import org.apache.pdfbox.pdmodel.font.PDType1Font; +import org.apache.pdfbox.pdmodel.font.Standard14Fonts; +import org.apache.pdfbox.pdmodel.interactive.annotation.PDAnnotationWidget; +import org.apache.pdfbox.pdmodel.interactive.form.PDAcroForm; +import org.apache.pdfbox.pdmodel.interactive.form.PDCheckBox; +import org.apache.pdfbox.pdmodel.interactive.form.PDComboBox; +import org.apache.pdfbox.pdmodel.interactive.form.PDListBox; +import org.apache.pdfbox.pdmodel.interactive.form.PDRadioButton; +import org.apache.pdfbox.pdmodel.interactive.form.PDSignatureField; +import org.apache.pdfbox.pdmodel.interactive.form.PDTerminalField; +import org.apache.pdfbox.pdmodel.interactive.form.PDTextField; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; + +/** + * Gap coverage for {@link FormUtils} methods not exercised by {@code FormUtilsTest} (disabled) or + * {@code FormUtilsAdditionalTest}. Focuses on coordinate extraction, the page-map / repair / prune + * / delete / modify lifecycle, and the package-private parsing helpers. + */ +class FormUtilsGapTest { + + private record SetupDocument(PDPage page, PDAcroForm acroForm) {} + + private static SetupDocument createBasicDocument(PDDocument document) { + PDPage page = new PDPage(); + document.addPage(page); + + PDAcroForm acroForm = new PDAcroForm(document); + // Register a Helvetica font in the default resources and set a default appearance so + // PDFBox can write text-field values without throwing "/DA is a required entry". + PDResources dr = new PDResources(); + dr.put(COSName.getPDFName("Helv"), new PDType1Font(Standard14Fonts.FontName.HELVETICA)); + acroForm.setDefaultResources(dr); + acroForm.setDefaultAppearance("/Helv 12 Tf 0 g"); + acroForm.setNeedAppearances(true); + document.getDocumentCatalog().setAcroForm(acroForm); + + return new SetupDocument(page, acroForm); + } + + private static void attachWidget( + SetupDocument setup, PDTerminalField field, PDRectangle rectangle) throws IOException { + PDAnnotationWidget widget = new PDAnnotationWidget(); + widget.setRectangle(rectangle); + widget.setPage(setup.page()); + // Start from an empty list: a fresh terminal field has no /Kids, so getWidgets() would + // return a synthetic widget wrapping the field dict itself. Re-adding that turns the field + // into a self-referential non-terminal field whose getWidgets() is empty. + List widgets = new ArrayList<>(); + widgets.add(widget); + field.setWidgets(widgets); + setup.acroForm().getFields().add(field); + setup.page().getAnnotations().add(widget); + } + + // ---------------------------------------------------------------------- + // Constants + // ---------------------------------------------------------------------- + + @Nested + @DisplayName("Field type constants") + class Constants { + + @Test + void typeConstantsHaveExpectedValues() { + assertEquals("text", FormUtils.FIELD_TYPE_TEXT); + assertEquals("checkbox", FormUtils.FIELD_TYPE_CHECKBOX); + assertEquals("combobox", FormUtils.FIELD_TYPE_COMBOBOX); + assertEquals("listbox", FormUtils.FIELD_TYPE_LISTBOX); + assertEquals("radio", FormUtils.FIELD_TYPE_RADIO); + assertEquals("button", FormUtils.FIELD_TYPE_BUTTON); + assertEquals("signature", FormUtils.FIELD_TYPE_SIGNATURE); + } + + @Test + void choiceFieldTypesContainsExpectedMembers() { + assertTrue(FormUtils.CHOICE_FIELD_TYPES.contains("combobox")); + assertTrue(FormUtils.CHOICE_FIELD_TYPES.contains("listbox")); + assertTrue(FormUtils.CHOICE_FIELD_TYPES.contains("radio")); + assertFalse(FormUtils.CHOICE_FIELD_TYPES.contains("text")); + assertEquals(3, FormUtils.CHOICE_FIELD_TYPES.size()); + } + } + + // ---------------------------------------------------------------------- + // detectFieldType (choice/radio/signature/button branches) + // ---------------------------------------------------------------------- + + @Nested + @DisplayName("detectFieldType") + class DetectFieldType { + + @Test + void comboBoxDetected() throws IOException { + try (PDDocument doc = new PDDocument()) { + SetupDocument setup = createBasicDocument(doc); + assertEquals( + "combobox", FormUtils.detectFieldType(new PDComboBox(setup.acroForm()))); + } + } + + @Test + void listBoxDetected() throws IOException { + try (PDDocument doc = new PDDocument()) { + SetupDocument setup = createBasicDocument(doc); + assertEquals("listbox", FormUtils.detectFieldType(new PDListBox(setup.acroForm()))); + } + } + + @Test + void radioButtonDetected() throws IOException { + try (PDDocument doc = new PDDocument()) { + SetupDocument setup = createBasicDocument(doc); + assertEquals( + "radio", FormUtils.detectFieldType(new PDRadioButton(setup.acroForm()))); + } + } + + @Test + void signatureDetected() throws IOException { + try (PDDocument doc = new PDDocument()) { + SetupDocument setup = createBasicDocument(doc); + assertEquals( + "signature", + FormUtils.detectFieldType(new PDSignatureField(setup.acroForm()))); + } + } + } + + // ---------------------------------------------------------------------- + // isChecked + // ---------------------------------------------------------------------- + + @Nested + @DisplayName("isChecked") + class IsChecked { + + @Test + void nullIsFalse() { + assertFalse(FormUtils.isChecked(null)); + } + + @Test + void truthyValuesAreChecked() { + assertTrue(FormUtils.isChecked("true")); + assertTrue(FormUtils.isChecked("1")); + assertTrue(FormUtils.isChecked("yes")); + assertTrue(FormUtils.isChecked("on")); + assertTrue(FormUtils.isChecked("checked")); + } + + @Test + void truthyValuesAreCaseInsensitiveAndTrimmed() { + assertTrue(FormUtils.isChecked(" TRUE ")); + assertTrue(FormUtils.isChecked("Yes")); + assertTrue(FormUtils.isChecked("ON")); + } + + @Test + void falsyValuesAreNotChecked() { + assertFalse(FormUtils.isChecked("false")); + assertFalse(FormUtils.isChecked("0")); + assertFalse(FormUtils.isChecked("off")); + assertFalse(FormUtils.isChecked("")); + assertFalse(FormUtils.isChecked("anything")); + } + } + + // ---------------------------------------------------------------------- + // safeValue + // ---------------------------------------------------------------------- + + @Test + void safeValueEmptyStringPassesThrough() { + assertEquals("", FormUtils.safeValue("")); + } + + // ---------------------------------------------------------------------- + // parseMultiChoiceSelections + // ---------------------------------------------------------------------- + + @Nested + @DisplayName("parseMultiChoiceSelections") + class ParseMultiChoiceSelections { + + @Test + void nullReturnsEmpty() { + assertTrue(FormUtils.parseMultiChoiceSelections(null).isEmpty()); + } + + @Test + void blankReturnsEmpty() { + assertTrue(FormUtils.parseMultiChoiceSelections(" ").isEmpty()); + } + + @Test + void splitsAndTrims() { + List result = FormUtils.parseMultiChoiceSelections(" a , b ,c "); + assertEquals(List.of("a", "b", "c"), result); + } + + @Test + void dropsEmptySegments() { + List result = FormUtils.parseMultiChoiceSelections("a,,b,"); + assertEquals(List.of("a", "b"), result); + } + } + + // ---------------------------------------------------------------------- + // filterChoiceSelections + // ---------------------------------------------------------------------- + + @Nested + @DisplayName("filterChoiceSelections") + class FilterChoiceSelections { + + @Test + void nullSelectionsReturnsEmpty() { + assertTrue(FormUtils.filterChoiceSelections(null, List.of("A"), "f").isEmpty()); + } + + @Test + void emptySelectionsReturnsEmpty() { + assertTrue(FormUtils.filterChoiceSelections(List.of(), List.of("A"), "f").isEmpty()); + } + + @Test + void selectionsOfOnlyBlanksReturnsEmpty() { + List selections = new ArrayList<>(); + selections.add(" "); + selections.add(null); + assertTrue(FormUtils.filterChoiceSelections(selections, List.of("A"), "f").isEmpty()); + } + + @Test + void matchingSelectionsAreKeptCaseInsensitively() { + List result = + FormUtils.filterChoiceSelections( + List.of("apple", "BANANA"), List.of("Apple", "Banana", "Cherry"), "f"); + // The resolved (canonical) allowed option is returned, not the input. + assertEquals(List.of("Apple", "Banana"), result); + } + + @Test + void unsupportedSelectionsAreDropped() { + List result = + FormUtils.filterChoiceSelections( + List.of("Apple", "Grape"), List.of("Apple", "Banana"), "f"); + assertEquals(List.of("Apple"), result); + } + + @Test + void missingAllowedOptionsThrows() { + org.junit.jupiter.api.Assertions.assertThrows( + IllegalArgumentException.class, + () -> FormUtils.filterChoiceSelections(List.of("Apple"), List.of(), "fieldX")); + } + + @Test + void nullAllowedOptionsThrows() { + org.junit.jupiter.api.Assertions.assertThrows( + IllegalArgumentException.class, + () -> FormUtils.filterChoiceSelections(List.of("Apple"), null, "fieldX")); + } + } + + // ---------------------------------------------------------------------- + // resolveOptions / resolveDisplayOptions / collectChoiceAllowedValues + // ---------------------------------------------------------------------- + + @Nested + @DisplayName("option resolution") + class OptionResolution { + + @Test + void resolveOptionsForComboBox() throws IOException { + try (PDDocument doc = new PDDocument()) { + SetupDocument setup = createBasicDocument(doc); + PDComboBox combo = new PDComboBox(setup.acroForm()); + combo.setOptions(List.of("Red", "Green", "Blue")); + List options = FormUtils.resolveOptions(combo); + assertTrue(options.contains("Red")); + assertTrue(options.contains("Green")); + assertTrue(options.contains("Blue")); + } + } + + @Test + void resolveOptionsForTextFieldIsEmpty() throws IOException { + try (PDDocument doc = new PDDocument()) { + SetupDocument setup = createBasicDocument(doc); + PDTextField text = new PDTextField(setup.acroForm()); + assertTrue(FormUtils.resolveOptions(text).isEmpty()); + } + } + + @Test + void resolveOptionsForCheckBoxUsesExportValues() throws IOException { + try (PDDocument doc = new PDDocument()) { + SetupDocument setup = createBasicDocument(doc); + PDCheckBox checkBox = new PDCheckBox(setup.acroForm()); + checkBox.setExportValues(List.of("Yes")); + assertEquals(List.of("Yes"), FormUtils.resolveOptions(checkBox)); + } + } + + @Test + void resolveDisplayOptionsEmptyForTextField() throws IOException { + try (PDDocument doc = new PDDocument()) { + SetupDocument setup = createBasicDocument(doc); + PDTextField text = new PDTextField(setup.acroForm()); + assertTrue(FormUtils.resolveDisplayOptions(text).isEmpty()); + } + } + + @Test + void collectChoiceAllowedValuesNullReturnsEmpty() { + assertTrue(FormUtils.collectChoiceAllowedValues(null).isEmpty()); + } + + @Test + void collectChoiceAllowedValuesReturnsOptions() throws IOException { + try (PDDocument doc = new PDDocument()) { + SetupDocument setup = createBasicDocument(doc); + PDComboBox combo = new PDComboBox(setup.acroForm()); + combo.setOptions(List.of("One", "Two")); + List allowed = FormUtils.collectChoiceAllowedValues(combo); + assertTrue(allowed.contains("One")); + assertTrue(allowed.contains("Two")); + } + } + } + + // ---------------------------------------------------------------------- + // setTextValue + // ---------------------------------------------------------------------- + + @Nested + @DisplayName("setTextValue") + class SetTextValue { + + @Test + void writesValue() throws IOException { + try (PDDocument doc = new PDDocument()) { + SetupDocument setup = createBasicDocument(doc); + PDTextField text = new PDTextField(setup.acroForm()); + text.setPartialName("note"); + text.setDefaultAppearance("/Helv 12 Tf 0 g"); + attachWidget(setup, text, new PDRectangle(20, 600, 200, 20)); + + FormUtils.setTextValue(text, "hello world"); + assertEquals("hello world", text.getValueAsString()); + } + } + + @Test + void nullValueWritesEmptyString() throws IOException { + try (PDDocument doc = new PDDocument()) { + SetupDocument setup = createBasicDocument(doc); + PDTextField text = new PDTextField(setup.acroForm()); + text.setPartialName("note"); + text.setDefaultAppearance("/Helv 12 Tf 0 g"); + attachWidget(setup, text, new PDRectangle(20, 600, 200, 20)); + + FormUtils.setTextValue(text, null); + assertEquals("", text.getValueAsString()); + } + } + } + + // ---------------------------------------------------------------------- + // buildFillTemplateRecord (choice branches not covered elsewhere) + // ---------------------------------------------------------------------- + + @Nested + @DisplayName("buildFillTemplateRecord") + class BuildFillTemplateRecord { + + @Test + void comboBoxUsesCurrentValue() { + FormUtils.FormFieldInfo info = + new FormUtils.FormFieldInfo( + "color", "Color", "combobox", "Red", null, false, 0, false, null, 0); + Map result = FormUtils.buildFillTemplateRecord(List.of(info)); + assertEquals("Red", result.get("color")); + } + + @Test + void singleSelectListBoxUsesValue() { + FormUtils.FormFieldInfo info = + new FormUtils.FormFieldInfo( + "list", "List", "listbox", "Item1", null, false, 0, false, null, 0); + Map result = FormUtils.buildFillTemplateRecord(List.of(info)); + assertEquals("Item1", result.get("list")); + } + + @Test + void multiSelectListBoxUsesEmptyArray() { + FormUtils.FormFieldInfo info = + new FormUtils.FormFieldInfo( + "list", "List", "listbox", "Item1", null, false, 0, true, null, 0); + Map result = FormUtils.buildFillTemplateRecord(List.of(info)); + Object value = result.get("list"); + assertTrue(value instanceof List); + assertTrue(((List) value).isEmpty()); + } + + @Test + void nullValueDefaultsToEmptyString() { + FormUtils.FormFieldInfo info = + new FormUtils.FormFieldInfo( + "name", "Name", "text", null, null, false, 0, false, null, 0); + Map result = FormUtils.buildFillTemplateRecord(List.of(info)); + assertEquals("", result.get("name")); + } + + @Test + void entriesWithBlankNamesAreSkipped() { + FormUtils.FormFieldInfo blank = + new FormUtils.FormFieldInfo( + " ", "Blank", "text", "x", null, false, 0, false, null, 0); + FormUtils.FormFieldInfo good = + new FormUtils.FormFieldInfo( + "kept", "Kept", "text", "x", null, false, 0, false, null, 0); + Map result = FormUtils.buildFillTemplateRecord(List.of(blank, good)); + assertEquals(1, result.size()); + assertTrue(result.containsKey("kept")); + } + } + + // ---------------------------------------------------------------------- + // buildAnnotationPageMap + // ---------------------------------------------------------------------- + + @Nested + @DisplayName("buildAnnotationPageMap") + class BuildAnnotationPageMap { + + @Test + void nullDocumentReturnsEmpty() { + assertTrue(FormUtils.buildAnnotationPageMap(null).isEmpty()); + } + + @Test + void emptyDocumentReturnsEmpty() throws IOException { + try (PDDocument doc = new PDDocument()) { + doc.addPage(new PDPage()); + assertTrue(FormUtils.buildAnnotationPageMap(doc).isEmpty()); + } + } + + @Test + void mapsWidgetToItsPageIndex() throws IOException { + try (PDDocument doc = new PDDocument()) { + SetupDocument setup = createBasicDocument(doc); + PDTextField text = new PDTextField(setup.acroForm()); + text.setPartialName("a"); + attachWidget(setup, text, new PDRectangle(10, 10, 100, 20)); + + Map map = FormUtils.buildAnnotationPageMap(doc); + assertEquals(1, map.size()); + assertTrue(map.containsValue(0)); + } + } + } + + // ---------------------------------------------------------------------- + // extractFormFieldsWithCoordinates + // ---------------------------------------------------------------------- + + @Nested + @DisplayName("extractFormFieldsWithCoordinates") + class ExtractFormFieldsWithCoordinates { + + @Test + void nullDocumentReturnsEmpty() { + assertTrue(FormUtils.extractFormFieldsWithCoordinates(null).isEmpty()); + } + + @Test + void noAcroFormReturnsEmpty() throws IOException { + try (PDDocument doc = new PDDocument()) { + doc.addPage(new PDPage()); + assertTrue(FormUtils.extractFormFieldsWithCoordinates(doc).isEmpty()); + } + } + + @Test + void textFieldProducesWidgetCoordinates() throws IOException { + try (PDDocument doc = new PDDocument()) { + SetupDocument setup = createBasicDocument(doc); + PDTextField text = new PDTextField(setup.acroForm()); + text.setPartialName("firstName"); + attachWidget(setup, text, new PDRectangle(50, 700, 200, 20)); + + List fields = + FormUtils.extractFormFieldsWithCoordinates(doc); + assertEquals(1, fields.size()); + stirling.software.common.model.FormFieldWithCoordinates field = fields.get(0); + assertEquals("firstName", field.getName()); + assertEquals("text", field.getType()); + assertNotNull(field.getWidgets()); + assertEquals(1, field.getWidgets().size()); + stirling.software.common.model.FormFieldWithCoordinates.WidgetCoordinates wc = + field.getWidgets().get(0); + assertEquals(0, wc.getPageIndex()); + // x is relative to crop-box origin (0 here), so it equals the lower-left x. + assertEquals(50f, wc.getX(), 0.01f); + assertEquals(200f, wc.getWidth(), 0.01f); + assertEquals(20f, wc.getHeight(), 0.01f); + } + } + + @Test + void multipleFieldsAreSortedTopToBottom() throws IOException { + try (PDDocument doc = new PDDocument()) { + SetupDocument setup = createBasicDocument(doc); + + PDTextField lower = new PDTextField(setup.acroForm()); + lower.setPartialName("lower"); + attachWidget(setup, lower, new PDRectangle(50, 100, 200, 20)); + + PDTextField upper = new PDTextField(setup.acroForm()); + upper.setPartialName("upper"); + attachWidget(setup, upper, new PDRectangle(50, 700, 200, 20)); + + List fields = + FormUtils.extractFormFieldsWithCoordinates(doc); + assertEquals(2, fields.size()); + // The widget higher on the page (smaller CSS-y after flip) sorts first. + assertEquals("upper", fields.get(0).getName()); + assertEquals("lower", fields.get(1).getName()); + } + } + } + + // ---------------------------------------------------------------------- + // repairMissingWidgetPageReferences + // ---------------------------------------------------------------------- + + @Nested + @DisplayName("repairMissingWidgetPageReferences") + class RepairMissingWidgetPageReferences { + + @Test + void nullDocumentDoesNotThrow() { + FormUtils.repairMissingWidgetPageReferences(null); + } + + @Test + void documentWithoutAcroFormDoesNotThrow() throws IOException { + try (PDDocument doc = new PDDocument()) { + doc.addPage(new PDPage()); + FormUtils.repairMissingWidgetPageReferences(doc); + } + } + + @Test + void setsPageReferenceForOrphanWidget() throws IOException { + try (PDDocument doc = new PDDocument()) { + SetupDocument setup = createBasicDocument(doc); + PDTextField text = new PDTextField(setup.acroForm()); + text.setPartialName("orphan"); + + // Build a widget that is on the page's annotation list but has no /P page ref. + PDAnnotationWidget widget = new PDAnnotationWidget(); + widget.setRectangle(new PDRectangle(10, 10, 100, 20)); + List widgets = new ArrayList<>(text.getWidgets()); + widgets.add(widget); + text.setWidgets(widgets); + setup.acroForm().getFields().add(text); + setup.page().getAnnotations().add(widget); + + assertNull(widget.getPage()); + FormUtils.repairMissingWidgetPageReferences(doc); + assertNotNull(widget.getPage()); + } + } + } + + // ---------------------------------------------------------------------- + // deleteFormFields + // ---------------------------------------------------------------------- + + @Nested + @DisplayName("deleteFormFields") + class DeleteFormFields { + + @Test + void nullDocumentIsNoOp() { + FormUtils.deleteFormFields(null, List.of("a")); + } + + @Test + void nullNamesIsNoOp() throws IOException { + try (PDDocument doc = new PDDocument()) { + createBasicDocument(doc); + FormUtils.deleteFormFields(doc, null); + } + } + + @Test + void emptyNamesIsNoOp() throws IOException { + try (PDDocument doc = new PDDocument()) { + createBasicDocument(doc); + FormUtils.deleteFormFields(doc, List.of()); + } + } + + @Test + void removesNamedField() throws IOException { + try (PDDocument doc = new PDDocument()) { + SetupDocument setup = createBasicDocument(doc); + PDTextField keep = new PDTextField(setup.acroForm()); + keep.setPartialName("keep"); + attachWidget(setup, keep, new PDRectangle(50, 700, 200, 20)); + + PDTextField remove = new PDTextField(setup.acroForm()); + remove.setPartialName("remove"); + attachWidget(setup, remove, new PDRectangle(50, 660, 200, 20)); + + FormUtils.deleteFormFields(doc, List.of("remove")); + + List remaining = FormUtils.extractFormFields(doc); + assertEquals(1, remaining.size()); + assertEquals("keep", remaining.get(0).name()); + } + } + + @Test + void unknownFieldNameIsIgnored() throws IOException { + try (PDDocument doc = new PDDocument()) { + SetupDocument setup = createBasicDocument(doc); + PDTextField keep = new PDTextField(setup.acroForm()); + keep.setPartialName("keep"); + attachWidget(setup, keep, new PDRectangle(50, 700, 200, 20)); + + FormUtils.deleteFormFields(doc, List.of("doesNotExist", " ", "keep")); + assertTrue(FormUtils.extractFormFields(doc).isEmpty()); + } + } + } + + // ---------------------------------------------------------------------- + // modifyFormFields + // ---------------------------------------------------------------------- + + @Nested + @DisplayName("modifyFormFields") + class ModifyFormFields { + + @Test + void nullDocumentIsNoOp() { + FormUtils.modifyFormFields(null, List.of()); + } + + @Test + void nullModificationsIsNoOp() throws IOException { + try (PDDocument doc = new PDDocument()) { + createBasicDocument(doc); + FormUtils.modifyFormFields(doc, null); + } + } + + @Test + void emptyModificationsIsNoOp() throws IOException { + try (PDDocument doc = new PDDocument()) { + createBasicDocument(doc); + FormUtils.modifyFormFields(doc, List.of()); + } + } + + @Test + void inPlaceRenameAndLabelUpdate() throws IOException { + try (PDDocument doc = new PDDocument()) { + SetupDocument setup = createBasicDocument(doc); + PDTextField text = new PDTextField(setup.acroForm()); + text.setPartialName("oldName"); + text.setDefaultAppearance("/Helv 12 Tf 0 g"); + attachWidget(setup, text, new PDRectangle(50, 700, 200, 20)); + + FormUtils.ModifyFormFieldDefinition mod = + new FormUtils.ModifyFormFieldDefinition( + "oldName", + "newName", + "New Label", + null, // keep type (text) -> in-place path + Boolean.TRUE, + null, + null, + null, + null); + + FormUtils.modifyFormFields(doc, List.of(mod)); + + List fields = FormUtils.extractFormFields(doc); + assertEquals(1, fields.size()); + assertEquals("newName", fields.get(0).name()); + assertEquals("New Label", fields.get(0).label()); + assertTrue(fields.get(0).required()); + } + } + + @Test + void unknownTargetIsSkipped() throws IOException { + try (PDDocument doc = new PDDocument()) { + SetupDocument setup = createBasicDocument(doc); + PDTextField text = new PDTextField(setup.acroForm()); + text.setPartialName("present"); + attachWidget(setup, text, new PDRectangle(50, 700, 200, 20)); + + FormUtils.ModifyFormFieldDefinition mod = + new FormUtils.ModifyFormFieldDefinition( + "missing", null, null, null, null, null, null, null, null); + + FormUtils.modifyFormFields(doc, List.of(mod)); + + // Untouched field remains. + List fields = FormUtils.extractFormFields(doc); + assertEquals(1, fields.size()); + assertEquals("present", fields.get(0).name()); + } + } + + @Test + void nullEntriesAndBlankTargetsAreSkipped() throws IOException { + try (PDDocument doc = new PDDocument()) { + SetupDocument setup = createBasicDocument(doc); + PDTextField text = new PDTextField(setup.acroForm()); + text.setPartialName("present"); + attachWidget(setup, text, new PDRectangle(50, 700, 200, 20)); + + List mods = new ArrayList<>(); + mods.add(null); + mods.add( + new FormUtils.ModifyFormFieldDefinition( + " ", null, null, null, null, null, null, null, null)); + + FormUtils.modifyFormFields(doc, mods); + assertEquals(1, FormUtils.extractFormFields(doc).size()); + } + } + } + + // ---------------------------------------------------------------------- + // pruneOrphanedFormFields + // ---------------------------------------------------------------------- + + @Nested + @DisplayName("pruneOrphanedFormFields") + class PruneOrphanedFormFields { + + @Test + void nullDocumentIsNoOp() { + FormUtils.pruneOrphanedFormFields(null); + } + + @Test + void documentWithoutAcroFormIsNoOp() throws IOException { + try (PDDocument doc = new PDDocument()) { + doc.addPage(new PDPage()); + FormUtils.pruneOrphanedFormFields(doc); + assertNull(doc.getDocumentCatalog().getAcroForm(null)); + } + } + + @Test + void keepsFieldsWithLiveWidgets() throws IOException { + try (PDDocument doc = new PDDocument()) { + SetupDocument setup = createBasicDocument(doc); + PDTextField text = new PDTextField(setup.acroForm()); + text.setPartialName("live"); + attachWidget(setup, text, new PDRectangle(50, 700, 200, 20)); + + FormUtils.pruneOrphanedFormFields(doc); + + PDAcroForm form = doc.getDocumentCatalog().getAcroForm(null); + assertNotNull(form); + assertEquals(1, form.getFields().size()); + } + } + + @Test + void dropsAcroFormWhenAllWidgetsOrphaned() throws IOException { + try (PDDocument doc = new PDDocument()) { + SetupDocument setup = createBasicDocument(doc); + PDTextField text = new PDTextField(setup.acroForm()); + text.setPartialName("orphan"); + attachWidget(setup, text, new PDRectangle(50, 700, 200, 20)); + + // Remove the widget from the page so it is no longer "live". + setup.page().getAnnotations().clear(); + + FormUtils.pruneOrphanedFormFields(doc); + + assertNull(doc.getDocumentCatalog().getAcroForm(null)); + } + } + } + + // ---------------------------------------------------------------------- + // hasAnyRotatedPage (rotated branch) + // ---------------------------------------------------------------------- + + @Nested + @DisplayName("hasAnyRotatedPage") + class HasAnyRotatedPage { + + @Test + void rotatedPageDetected() throws IOException { + try (PDDocument doc = new PDDocument()) { + PDPage page = new PDPage(); + page.setRotation(90); + doc.addPage(page); + assertTrue(FormUtils.hasAnyRotatedPage(doc)); + } + } + + @Test + void unrotatedPageNotDetected() throws IOException { + try (PDDocument doc = new PDDocument()) { + doc.addPage(new PDPage()); + assertFalse(FormUtils.hasAnyRotatedPage(doc)); + } + } + } +} diff --git a/app/common/src/test/java/stirling/software/common/util/GeneralUtilsGapTest.java b/app/common/src/test/java/stirling/software/common/util/GeneralUtilsGapTest.java new file mode 100644 index 0000000000..51a08ee9ab --- /dev/null +++ b/app/common/src/test/java/stirling/software/common/util/GeneralUtilsGapTest.java @@ -0,0 +1,333 @@ +package stirling.software.common.util; + +import static org.junit.jupiter.api.Assertions.*; + +import java.io.File; +import java.io.IOException; +import java.nio.charset.StandardCharsets; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.List; +import java.util.Map; + +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.io.TempDir; +import org.springframework.core.io.DefaultResourceLoader; +import org.springframework.core.io.Resource; +import org.springframework.mock.web.MockMultipartFile; +import org.springframework.web.multipart.MultipartFile; + +/** + * Gap-coverage tests for {@link GeneralUtils}. Targets the public methods NOT already exercised by + * {@code GeneralUtilsAdditionalTest} (size/url/version/uuid) or {@code GeneralUtilsTest} (filename + * helpers, parsePageList basics, saveKeyToSettings): namely {@code generateFilename}, {@code + * convertToFileName}, {@code evaluateNFunc}, the n-function {@code parsePageList} path, {@code + * createDir}/{@code deleteDirectory}, multipart conversion, the {@code updateSettingsTransactional} + * early-return guards, {@code getResourcesFromLocationPattern}, and the environment helpers {@code + * generateMachineFingerprint}/{@code getLocalNetworkIp}. + */ +class GeneralUtilsGapTest { + + @Nested + @DisplayName("generateFilename") + class GenerateFilenameTests { + + @Test + @DisplayName("removes extension then appends suffix") + void removesAndAppends() { + assertEquals( + "report_out.pdf", GeneralUtils.generateFilename("report.docx", "_out.pdf")); + } + + @Test + @DisplayName("null filename uses default base") + void nullFilename() { + assertEquals("default_out.pdf", GeneralUtils.generateFilename(null, "_out.pdf")); + } + + @Test + @DisplayName("filename without extension is preserved") + void noExtension() { + assertEquals("README_x", GeneralUtils.generateFilename("README", "_x")); + } + } + + @Nested + @DisplayName("convertToFileName") + class ConvertToFileNameTests { + + @Test + @DisplayName("null returns underscore") + void nullReturnsUnderscore() { + assertEquals("_", GeneralUtils.convertToFileName(null)); + } + + @Test + @DisplayName("keeps letters and digits, replaces others with underscore") + void replacesUnsafeChars() { + assertEquals("my_file_2024_", GeneralUtils.convertToFileName("my file/2024!")); + } + + @Test + @DisplayName("alphanumeric input is unchanged") + void alphanumericUnchanged() { + assertEquals("File123", GeneralUtils.convertToFileName("File123")); + } + + @Test + @DisplayName("truncates to 50 characters") + void truncatesToFifty() { + String input = "a".repeat(100); + assertEquals(50, GeneralUtils.convertToFileName(input).length()); + } + } + + @Nested + @DisplayName("evaluateNFunc") + class EvaluateNFuncTests { + + @Test + @DisplayName("null expression throws") + void nullThrows() { + assertThrows( + IllegalArgumentException.class, () -> GeneralUtils.evaluateNFunc(null, 10)); + } + + @Test + @DisplayName("blank expression throws") + void blankThrows() { + assertThrows( + IllegalArgumentException.class, () -> GeneralUtils.evaluateNFunc(" ", 10)); + } + + @Test + @DisplayName("maxValue below 1 throws") + void maxValueTooLow() { + assertThrows(IllegalArgumentException.class, () -> GeneralUtils.evaluateNFunc("n", 0)); + } + + @Test + @DisplayName("maxValue above 10000 throws") + void maxValueTooHigh() { + assertThrows( + IllegalArgumentException.class, () -> GeneralUtils.evaluateNFunc("n", 10001)); + } + + @Test + @DisplayName("invalid characters throw") + void invalidCharsThrow() { + assertThrows( + IllegalArgumentException.class, () -> GeneralUtils.evaluateNFunc("n$", 10)); + } + + @Test + @DisplayName("identity 'n' yields all pages up to maxValue") + void identity() { + assertEquals(List.of(1, 2, 3, 4, 5), GeneralUtils.evaluateNFunc("n", 5)); + } + + @Test + @DisplayName("2n yields even values within bounds") + void doubling() { + assertEquals(List.of(2, 4, 6), GeneralUtils.evaluateNFunc("2n", 6)); + } + + @Test + @DisplayName("implicit multiplication 'n(n-1)' is handled") + void implicitMultiplication() { + // n*(n-1): n=1->0(excluded), n=2->2, n=3->6 ; capped at maxValue 6 + assertEquals(List.of(2, 6), GeneralUtils.evaluateNFunc("n(n-1)", 6)); + } + + @Test + @DisplayName("results outside (0, maxValue] are excluded") + void boundsExcluded() { + // n+10 always exceeds maxValue 5 -> empty + assertTrue(GeneralUtils.evaluateNFunc("n+10", 5).isEmpty()); + } + } + + @Nested + @DisplayName("parsePageList n-function path") + class ParsePageListNFuncTests { + + @Test + @DisplayName("n-function token expands to matching one-based pages") + void nFunctionOneBased() { + // 2n for total 6 (one-based) -> values 2,4,6 mapped to (v-1+1)=v + assertEquals(List.of(2, 4, 6), GeneralUtils.parsePageList("2n", 6, true)); + } + + @Test + @DisplayName("n-function token zero-based subtracts one") + void nFunctionZeroBased() { + // 2n for total 6 zero-based -> values 2,4,6 mapped to (v-1+0)=v-1 + assertEquals(List.of(1, 3, 5), GeneralUtils.parsePageList("2n", 6, false)); + } + } + + @Nested + @DisplayName("createDir and deleteDirectory") + class DirectoryTests { + + @Test + @DisplayName("createDir makes a nested directory and returns true") + void createNested(@TempDir Path tempDir) { + Path nested = tempDir.resolve("a").resolve("b").resolve("c"); + assertTrue(GeneralUtils.createDir(nested.toString())); + assertTrue(Files.isDirectory(nested)); + } + + @Test + @DisplayName("createDir returns true when directory already exists") + void createExisting(@TempDir Path tempDir) { + assertTrue(GeneralUtils.createDir(tempDir.toString())); + } + + @Test + @DisplayName("deleteDirectory removes a populated tree without touching siblings") + void deletePopulatedTree(@TempDir Path tempDir) throws IOException { + Path sibling = tempDir.resolve("sibling"); + Files.createDirectories(sibling); + Files.writeString(sibling.resolve("keep.txt"), "data"); + + Path root = tempDir.resolve("root"); + Files.createDirectories(root.resolve("nested")); + Files.writeString(root.resolve("a.txt"), "x"); + Files.writeString(root.resolve("nested").resolve("b.txt"), "y"); + + GeneralUtils.deleteDirectory(root); + + assertFalse(Files.exists(root)); + assertTrue(Files.exists(sibling.resolve("keep.txt"))); + } + } + + @Nested + @DisplayName("multipart conversion") + class MultipartTests { + + @Test + @DisplayName("convertMultipartFileToFile writes content to a temp file") + void convertWritesContent() throws IOException { + byte[] content = "hello world".getBytes(StandardCharsets.UTF_8); + MultipartFile mf = + new MockMultipartFile("file", "input.bin", "application/octet-stream", content); + + File out = GeneralUtils.convertMultipartFileToFile(mf); + try { + assertTrue(out.exists()); + assertArrayEquals(content, Files.readAllBytes(out.toPath())); + } finally { + Files.deleteIfExists(out.toPath()); + } + } + + @Test + @DisplayName("convertMultipartFileToFile handles empty input") + void convertEmpty() throws IOException { + MultipartFile mf = + new MockMultipartFile( + "file", "empty.bin", "application/octet-stream", new byte[0]); + + File out = GeneralUtils.convertMultipartFileToFile(mf); + try { + assertTrue(out.exists()); + assertEquals(0, out.length()); + } finally { + Files.deleteIfExists(out.toPath()); + } + } + + @Test + @DisplayName("multipartToFile writes content to a .pdf temp file") + void multipartToFileWritesContent() throws IOException { + byte[] content = "%PDF-1.7 minimal".getBytes(StandardCharsets.UTF_8); + MultipartFile mf = new MockMultipartFile("file", "doc.pdf", "application/pdf", content); + + File out = GeneralUtils.multipartToFile(mf); + try { + assertTrue(out.exists()); + assertTrue(out.getName().endsWith(".pdf")); + assertArrayEquals(content, Files.readAllBytes(out.toPath())); + } finally { + Files.deleteIfExists(out.toPath()); + } + } + } + + @Nested + @DisplayName("getResourcesFromLocationPattern") + class ResourcePatternTests { + + @Test + @DisplayName("file: pattern resolves matching files in a directory") + void filePatternResolves(@TempDir Path tempDir) throws Exception { + Files.writeString(tempDir.resolve("one.txt"), "1"); + Files.writeString(tempDir.resolve("two.txt"), "2"); + + String pattern = "file:" + tempDir.toString().replace("\\", "/") + "/*"; + Resource[] resources = + GeneralUtils.getResourcesFromLocationPattern( + pattern, new DefaultResourceLoader()); + + assertNotNull(resources); + assertEquals(2, resources.length); + } + + @Test + @DisplayName("classpath pattern with no matches returns an empty array") + void classpathNoMatches() throws Exception { + Resource[] resources = + GeneralUtils.getResourcesFromLocationPattern( + "classpath*:this/path/does/not/exist/**/*.nope", + new DefaultResourceLoader()); + + assertNotNull(resources); + assertEquals(0, resources.length); + } + } + + @Nested + @DisplayName("updateSettingsTransactional early-return guards") + class SettingsGuardTests { + + @Test + @DisplayName("null map returns without throwing") + void nullMapNoOp() { + assertDoesNotThrow(() -> GeneralUtils.updateSettingsTransactional(null)); + } + + @Test + @DisplayName("empty map returns without throwing") + void emptyMapNoOp() { + assertDoesNotThrow(() -> GeneralUtils.updateSettingsTransactional(Map.of())); + } + } + + @Nested + @DisplayName("environment-dependent helpers") + class EnvironmentHelperTests { + + @Test + @DisplayName("generateMachineFingerprint returns a non-blank, deterministic value") + void fingerprintStable() { + String first = GeneralUtils.generateMachineFingerprint(); + assertNotNull(first); + assertFalse(first.isBlank()); + // Deterministic within the same JVM/host + assertEquals(first, GeneralUtils.generateMachineFingerprint()); + } + + @Test + @DisplayName("getLocalNetworkIp returns null or a dotted IPv4 string") + void localIpFormat() { + String ip = GeneralUtils.getLocalNetworkIp(); + if (ip != null) { + assertTrue(ip.matches("\\d{1,3}(\\.\\d{1,3}){3}"), "unexpected IP form: " + ip); + } + } + } +} diff --git a/app/common/src/test/java/stirling/software/common/util/OfficeDocumentSanitizerTest.java b/app/common/src/test/java/stirling/software/common/util/OfficeDocumentSanitizerTest.java new file mode 100644 index 0000000000..eb22919068 --- /dev/null +++ b/app/common/src/test/java/stirling/software/common/util/OfficeDocumentSanitizerTest.java @@ -0,0 +1,370 @@ +package stirling.software.common.util; + +import static org.junit.jupiter.api.Assertions.assertArrayEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.lenient; +import static org.mockito.Mockito.mock; + +import java.io.ByteArrayInputStream; +import java.io.ByteArrayOutputStream; +import java.io.IOException; +import java.nio.charset.StandardCharsets; +import java.util.HashMap; +import java.util.LinkedHashMap; +import java.util.Map; +import java.util.zip.ZipEntry; +import java.util.zip.ZipInputStream; +import java.util.zip.ZipOutputStream; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.common.service.SsrfProtectionService; + +class OfficeDocumentSanitizerTest { + + private static final String EXTERNAL_URL = "https://webhook.site/ssrf-callback"; + private static final String INTERNAL_TARGET = "media/image1.png"; + + private static final String DOCX_RELS = + "" + + "" + + "" + + "" + + ""; + + private static final String DOCX_DOCUMENT = + "" + + "" + + ""; + + private static final String ODF_CONTENT_EXTERNAL = + "" + + "" + + "" + + "" + + "" + + ""; + + private SsrfProtectionService ssrfProtectionService; + private ApplicationProperties applicationProperties; + private OfficeDocumentSanitizer sanitizer; + + @BeforeEach + void setUp() { + applicationProperties = new ApplicationProperties(); + ssrfProtectionService = mock(SsrfProtectionService.class); + sanitizer = new OfficeDocumentSanitizer(ssrfProtectionService, applicationProperties); + } + + @Test + void isSanitizableExtension_recognizesOoxmlAndOdf() { + assertTrue(sanitizer.isSanitizableExtension("docx")); + assertTrue(sanitizer.isSanitizableExtension("DOCX")); + assertTrue(sanitizer.isSanitizableExtension("xlsx")); + assertTrue(sanitizer.isSanitizableExtension("pptx")); + assertTrue(sanitizer.isSanitizableExtension("odt")); + assertTrue(sanitizer.isSanitizableExtension("ods")); + assertTrue(sanitizer.isSanitizableExtension("odp")); + assertFalse(sanitizer.isSanitizableExtension("pdf")); + assertFalse(sanitizer.isSanitizableExtension("html")); + assertFalse(sanitizer.isSanitizableExtension("")); + assertFalse(sanitizer.isSanitizableExtension(null)); + } + + @Test + void sanitize_stripsOoxmlExternalRelationship() throws IOException { + Map entries = new LinkedHashMap<>(); + entries.put("word/_rels/document.xml.rels", DOCX_RELS.getBytes(StandardCharsets.UTF_8)); + entries.put("word/document.xml", DOCX_DOCUMENT.getBytes(StandardCharsets.UTF_8)); + byte[] docx = zip(entries); + + byte[] cleaned = sanitizer.sanitize(docx, "docx"); + + Map result = unzip(cleaned); + String rels = + new String(result.get("word/_rels/document.xml.rels"), StandardCharsets.UTF_8); + assertFalse(rels.contains(EXTERNAL_URL), "External URL should be stripped from .rels"); + assertFalse( + rels.toLowerCase().contains("targetmode=\"external\""), + "TargetMode=External relationship should be removed"); + assertTrue(rels.contains(INTERNAL_TARGET), "Internal image target should be preserved"); + assertArrayEquals( + DOCX_DOCUMENT.getBytes(StandardCharsets.UTF_8), + result.get("word/document.xml"), + "Non-rels entries must be untouched"); + } + + @Test + void sanitize_pptxExternalImageRelStripped() throws IOException { + String pptxRels = + "" + + "" + + "" + + ""; + Map entries = new LinkedHashMap<>(); + entries.put("ppt/slides/_rels/slide1.xml.rels", pptxRels.getBytes(StandardCharsets.UTF_8)); + byte[] pptx = zip(entries); + + byte[] cleaned = sanitizer.sanitize(pptx, "pptx"); + + Map result = unzip(cleaned); + String rels = + new String(result.get("ppt/slides/_rels/slide1.xml.rels"), StandardCharsets.UTF_8); + assertFalse(rels.contains(EXTERNAL_URL)); + } + + @Test + void sanitize_xlsxExternalImageRelStripped() throws IOException { + String xlsxRels = + "" + + "" + + "" + + ""; + Map entries = new LinkedHashMap<>(); + entries.put( + "xl/drawings/_rels/drawing1.xml.rels", xlsxRels.getBytes(StandardCharsets.UTF_8)); + byte[] xlsx = zip(entries); + + byte[] cleaned = sanitizer.sanitize(xlsx, "xlsx"); + + Map result = unzip(cleaned); + String rels = + new String( + result.get("xl/drawings/_rels/drawing1.xml.rels"), StandardCharsets.UTF_8); + assertFalse(rels.contains(EXTERNAL_URL)); + } + + @Test + void sanitize_odtStripsExternalXlinkHrefButKeepsInternal() throws IOException { + Map entries = new LinkedHashMap<>(); + entries.put("content.xml", ODF_CONTENT_EXTERNAL.getBytes(StandardCharsets.UTF_8)); + String manifestXml = + ""; + entries.put("META-INF/manifest.xml", manifestXml.getBytes(StandardCharsets.UTF_8)); + byte[] odt = zip(entries); + + byte[] cleaned = sanitizer.sanitize(odt, "odt"); + + Map result = unzip(cleaned); + String content = new String(result.get("content.xml"), StandardCharsets.UTF_8); + assertFalse(content.contains(EXTERNAL_URL), "External xlink:href should be stripped"); + assertTrue(content.contains("Pictures/image1.png"), "Internal href should be preserved"); + } + + @Test + void sanitize_odsStripsExternalXlinkHref() throws IOException { + Map entries = new LinkedHashMap<>(); + entries.put("content.xml", ODF_CONTENT_EXTERNAL.getBytes(StandardCharsets.UTF_8)); + byte[] ods = zip(entries); + + byte[] cleaned = sanitizer.sanitize(ods, "ods"); + + Map result = unzip(cleaned); + String content = new String(result.get("content.xml"), StandardCharsets.UTF_8); + assertFalse(content.contains(EXTERNAL_URL)); + } + + @Test + void sanitize_odpStripsExternalXlinkHrefInStylesXml() throws IOException { + String stylesXml = + "" + + "" + + ""; + Map entries = new LinkedHashMap<>(); + entries.put("styles.xml", stylesXml.getBytes(StandardCharsets.UTF_8)); + byte[] odp = zip(entries); + + byte[] cleaned = sanitizer.sanitize(odp, "odp"); + + Map result = unzip(cleaned); + String content = new String(result.get("styles.xml"), StandardCharsets.UTF_8); + assertFalse(content.contains(EXTERNAL_URL)); + } + + @Test + void sanitize_disabledByConfigReturnsOriginal() throws IOException { + applicationProperties.getSystem().setDisableSanitize(true); + Map entries = new LinkedHashMap<>(); + entries.put("word/_rels/document.xml.rels", DOCX_RELS.getBytes(StandardCharsets.UTF_8)); + byte[] docx = zip(entries); + + byte[] result = sanitizer.sanitize(docx, "docx"); + assertArrayEquals(docx, result); + } + + @Test + void sanitize_unrecognizedExtensionReturnsOriginal() throws IOException { + byte[] original = "irrelevant".getBytes(StandardCharsets.UTF_8); + byte[] result = sanitizer.sanitize(original, "pdf"); + assertArrayEquals(original, result); + } + + @Test + void sanitize_emptyInputThrows() { + assertThrows(IOException.class, () -> sanitizer.sanitize(new byte[0], "docx")); + } + + @Test + void sanitize_nullInputThrows() { + assertThrows(IOException.class, () -> sanitizer.sanitize(null, "docx")); + } + + @Test + void sanitize_preservesEntryWithExternalRefWhenAdminAllowsDomain() throws IOException { + applicationProperties + .getSystem() + .getHtml() + .getUrlSecurity() + .getAllowedDomains() + .add("webhook.site"); + lenient().when(ssrfProtectionService.isUrlAllowed(eq(EXTERNAL_URL))).thenReturn(true); + + Map entries = new LinkedHashMap<>(); + entries.put("word/_rels/document.xml.rels", DOCX_RELS.getBytes(StandardCharsets.UTF_8)); + byte[] docx = zip(entries); + + byte[] cleaned = sanitizer.sanitize(docx, "docx"); + + Map result = unzip(cleaned); + String rels = + new String(result.get("word/_rels/document.xml.rels"), StandardCharsets.UTF_8); + assertTrue(rels.contains(EXTERNAL_URL), "Allow-listed external URL should be preserved"); + } + + @Test + void sanitize_doesNotConsultSsrfServiceWhenAllowedDomainsEmpty() throws IOException { + // Even if mock would say allowed, we should not invoke it when there is no allow-list, + // because MEDIUM default would let public URLs through and re-introduce the vulnerability. + lenient().when(ssrfProtectionService.isUrlAllowed(eq(EXTERNAL_URL))).thenReturn(true); + + Map entries = new LinkedHashMap<>(); + entries.put("word/_rels/document.xml.rels", DOCX_RELS.getBytes(StandardCharsets.UTF_8)); + byte[] docx = zip(entries); + + byte[] cleaned = sanitizer.sanitize(docx, "docx"); + + Map result = unzip(cleaned); + String rels = + new String(result.get("word/_rels/document.xml.rels"), StandardCharsets.UTF_8); + assertFalse(rels.contains(EXTERNAL_URL)); + } + + @Test + void sanitize_handlesNonXmlEntriesSafely() throws IOException { + Map entries = new LinkedHashMap<>(); + byte[] imageBytes = new byte[] {(byte) 0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a}; + entries.put("word/media/image1.png", imageBytes); + entries.put("word/_rels/document.xml.rels", DOCX_RELS.getBytes(StandardCharsets.UTF_8)); + byte[] docx = zip(entries); + + byte[] cleaned = sanitizer.sanitize(docx, "docx"); + + Map result = unzip(cleaned); + assertArrayEquals(imageBytes, result.get("word/media/image1.png")); + } + + @Test + void sanitize_internalLinksKeptWhenNoExternalPresent() throws IOException { + String internalOnlyRels = + "" + + "" + + "" + + ""; + Map entries = new LinkedHashMap<>(); + entries.put( + "word/_rels/document.xml.rels", internalOnlyRels.getBytes(StandardCharsets.UTF_8)); + byte[] docx = zip(entries); + + byte[] cleaned = sanitizer.sanitize(docx, "docx"); + + Map result = unzip(cleaned); + String rels = + new String(result.get("word/_rels/document.xml.rels"), StandardCharsets.UTF_8); + assertTrue(rels.contains("media/image1.png")); + } + + @Test + void sanitize_corruptZipProducesSafeOutput() throws IOException { + byte[] garbage = "this is not a zip file".getBytes(StandardCharsets.UTF_8); + byte[] result = sanitizer.sanitize(garbage, "docx"); + Map entries = unzip(result); + assertTrue(entries.isEmpty(), "Garbage input must not yield exploitable entries"); + } + + @Test + void sanitize_relativeOdfPathsArePreserved() throws IOException { + String content = + "" + + "" + + "" + + "" + + ""; + Map entries = new LinkedHashMap<>(); + entries.put("content.xml", content.getBytes(StandardCharsets.UTF_8)); + byte[] odt = zip(entries); + + byte[] cleaned = sanitizer.sanitize(odt, "odt"); + + Map result = unzip(cleaned); + String out = new String(result.get("content.xml"), StandardCharsets.UTF_8); + assertTrue(out.contains("../Pictures/image1.png")); + assertTrue(out.contains("#anchor")); + } + + private static byte[] zip(Map entries) throws IOException { + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + try (ZipOutputStream zos = new ZipOutputStream(baos)) { + for (Map.Entry e : entries.entrySet()) { + ZipEntry entry = new ZipEntry(e.getKey()); + zos.putNextEntry(entry); + zos.write(e.getValue()); + zos.closeEntry(); + } + } + return baos.toByteArray(); + } + + private static Map unzip(byte[] data) throws IOException { + Map entries = new HashMap<>(); + try (ZipInputStream zis = new ZipInputStream(new ByteArrayInputStream(data))) { + ZipEntry e; + while ((e = zis.getNextEntry()) != null) { + entries.put(e.getName(), zis.readAllBytes()); + zis.closeEntry(); + } + } + return entries; + } +} diff --git a/app/common/src/test/java/stirling/software/common/util/PdfAttachmentHandlerGapTest.java b/app/common/src/test/java/stirling/software/common/util/PdfAttachmentHandlerGapTest.java new file mode 100644 index 0000000000..01d9128f26 --- /dev/null +++ b/app/common/src/test/java/stirling/software/common/util/PdfAttachmentHandlerGapTest.java @@ -0,0 +1,446 @@ +package stirling.software.common.util; + +import static org.junit.jupiter.api.Assertions.*; +import static org.mockito.Mockito.*; + +import java.io.ByteArrayOutputStream; +import java.nio.charset.StandardCharsets; +import java.time.ZoneId; +import java.time.ZonedDateTime; +import java.util.ArrayList; +import java.util.Base64; +import java.util.Date; +import java.util.List; +import java.util.Map; + +import org.apache.pdfbox.Loader; +import org.apache.pdfbox.pdmodel.PDDocument; +import org.apache.pdfbox.pdmodel.PDDocumentNameDictionary; +import org.apache.pdfbox.pdmodel.PDEmbeddedFilesNameTreeNode; +import org.apache.pdfbox.pdmodel.PDPage; +import org.apache.pdfbox.pdmodel.PDPageContentStream; +import org.apache.pdfbox.pdmodel.common.PDRectangle; +import org.apache.pdfbox.pdmodel.common.filespecification.PDComplexFileSpecification; +import org.apache.pdfbox.pdmodel.font.PDType1Font; +import org.apache.pdfbox.pdmodel.font.Standard14Fonts; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; + +import stirling.software.common.service.CustomPDFDocumentFactory; + +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class PdfAttachmentHandlerGapTest { + + @Mock private CustomPDFDocumentFactory pdfDocumentFactory; + + // ----- helpers ------------------------------------------------------- + + /** Builds a tiny one-page PDF whose page renders the given text lines, each on its own line. */ + private static byte[] pdfWithLines(String... lines) throws Exception { + try (PDDocument doc = new PDDocument()) { + PDPage page = new PDPage(PDRectangle.A4); + doc.addPage(page); + try (PDPageContentStream cs = new PDPageContentStream(doc, page)) { + cs.setFont(new PDType1Font(Standard14Fonts.FontName.HELVETICA), 12); + float y = 720f; + for (String line : lines) { + cs.beginText(); + cs.newLineAtOffset(72f, y); + cs.showText(line); + cs.endText(); + y -= 20f; + } + } + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + doc.save(baos); + return baos.toByteArray(); + } + } + + /** Builds a blank one-page PDF with no text. */ + private static byte[] blankPdf() throws Exception { + try (PDDocument doc = new PDDocument()) { + doc.addPage(new PDPage(PDRectangle.A4)); + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + doc.save(baos); + return baos.toByteArray(); + } + } + + private static EmlParser.EmailAttachment attachment(String filename, byte[] data) { + EmlParser.EmailAttachment a = new EmlParser.EmailAttachment(); + a.setFilename(filename); + a.setData(data); + a.setContentType("application/pdf"); + return a; + } + + private static List embeddedFileNames(byte[] pdfBytes) throws Exception { + List names = new ArrayList<>(); + try (PDDocument doc = Loader.loadPDF(pdfBytes)) { + PDDocumentNameDictionary docNames = doc.getDocumentCatalog().getNames(); + if (docNames == null) { + return names; + } + PDEmbeddedFilesNameTreeNode tree = docNames.getEmbeddedFiles(); + if (tree == null) { + return names; + } + Map map = tree.getNames(); + if (map != null) { + names.addAll(map.keySet()); + } + } + return names; + } + + // ----- attachFilesToPdf: short-circuit branches ---------------------- + + @Nested + @DisplayName("attachFilesToPdf short-circuit handling") + class ShortCircuitTests { + + @Test + @DisplayName("null attachment list returns the original bytes untouched") + void nullAttachments_returnsOriginalBytes() throws Exception { + byte[] original = {1, 2, 3, 4}; + byte[] result = + PdfAttachmentHandler.attachFilesToPdf(original, null, pdfDocumentFactory); + assertSame(original, result); + verifyNoInteractions(pdfDocumentFactory); + } + + @Test + @DisplayName("empty attachment list returns the original bytes untouched") + void emptyAttachments_returnsOriginalBytes() throws Exception { + byte[] original = {9, 8, 7}; + byte[] result = + PdfAttachmentHandler.attachFilesToPdf( + original, new ArrayList<>(), pdfDocumentFactory); + assertSame(original, result); + verifyNoInteractions(pdfDocumentFactory); + } + + @Test + @DisplayName("attachments with no usable data are skipped and a clean PDF is returned") + void attachmentsWithoutData_produceNoEmbeddedFiles() throws Exception { + byte[] pdfBytes = blankPdf(); + when(pdfDocumentFactory.load(pdfBytes)).thenReturn(Loader.loadPDF(pdfBytes)); + + List attachments = new ArrayList<>(); + attachments.add(attachment("empty.pdf", new byte[0])); + attachments.add(attachment("alsoEmpty.pdf", null)); + + byte[] result = + PdfAttachmentHandler.attachFilesToPdf( + pdfBytes, attachments, pdfDocumentFactory); + + assertNotNull(result); + assertTrue(result.length > 0); + assertTrue(embeddedFileNames(result).isEmpty()); + } + } + + // ----- attachFilesToPdf: embedding happy paths ----------------------- + + @Nested + @DisplayName("attachFilesToPdf embedding behaviour") + class EmbeddingTests { + + @Test + @DisplayName("embeds attachment data even when no '@' marker exists in the PDF text") + void embedsAttachment_withoutMarker() throws Exception { + byte[] pdfBytes = pdfWithLines("Just a plain document with no attachment markers"); + when(pdfDocumentFactory.load(pdfBytes)).thenReturn(Loader.loadPDF(pdfBytes)); + + List attachments = new ArrayList<>(); + attachments.add(attachment("report.pdf", "hello".getBytes(StandardCharsets.UTF_8))); + + byte[] result = + PdfAttachmentHandler.attachFilesToPdf( + pdfBytes, attachments, pdfDocumentFactory); + + List embedded = embeddedFileNames(result); + assertEquals(1, embedded.size()); + assertTrue(embedded.contains("report.pdf")); + } + + @Test + @DisplayName("embeds attachment and adds an annotation when an '@' marker matches") + void embedsAttachment_withMatchingMarker() throws Exception { + byte[] pdfBytes = + pdfWithLines("Email body text here", "Attachments (1)", "@report.pdf (5 KB)"); + when(pdfDocumentFactory.load(pdfBytes)).thenReturn(Loader.loadPDF(pdfBytes)); + + List attachments = new ArrayList<>(); + attachments.add(attachment("report.pdf", "PDFDATA".getBytes(StandardCharsets.UTF_8))); + + byte[] result = + PdfAttachmentHandler.attachFilesToPdf( + pdfBytes, attachments, pdfDocumentFactory); + + List embedded = embeddedFileNames(result); + assertTrue(embedded.contains("report.pdf")); + + // The annotation pass should have run and produced at least one annotation on the + // page that contains the marker (a blank source page has none). + try (PDDocument doc = Loader.loadPDF(result)) { + assertFalse(doc.getPage(0).getAnnotations().isEmpty()); + } + } + + @Test + @DisplayName("attachment without a filename falls back to a generated embedded name") + void embedsAttachment_withGeneratedName() throws Exception { + byte[] pdfBytes = blankPdf(); + when(pdfDocumentFactory.load(pdfBytes)).thenReturn(Loader.loadPDF(pdfBytes)); + + EmlParser.EmailAttachment a = new EmlParser.EmailAttachment(); + a.setFilename(null); + a.setData("x".getBytes(StandardCharsets.UTF_8)); + List attachments = new ArrayList<>(); + attachments.add(a); + + byte[] result = + PdfAttachmentHandler.attachFilesToPdf( + pdfBytes, attachments, pdfDocumentFactory); + + // A single embedded file should exist with a non-blank generated name. + List embedded = embeddedFileNames(result); + assertEquals(1, embedded.size()); + assertFalse(embedded.get(0).isBlank()); + } + + @Test + @DisplayName("duplicate attachment filenames produce uniquely named embedded files") + void embedsAttachments_withDuplicateNames() throws Exception { + byte[] pdfBytes = blankPdf(); + when(pdfDocumentFactory.load(pdfBytes)).thenReturn(Loader.loadPDF(pdfBytes)); + + List attachments = new ArrayList<>(); + attachments.add(attachment("dup.pdf", "a".getBytes(StandardCharsets.UTF_8))); + attachments.add(attachment("dup.pdf", "b".getBytes(StandardCharsets.UTF_8))); + + byte[] result = + PdfAttachmentHandler.attachFilesToPdf( + pdfBytes, attachments, pdfDocumentFactory); + + List embedded = embeddedFileNames(result); + assertEquals(2, embedded.size()); + assertTrue(embedded.contains("dup.pdf")); + // The second one must have been disambiguated, not overwritten. + assertTrue(embedded.stream().anyMatch(n -> !"dup.pdf".equals(n))); + } + } + + // ----- attachFilesToPdf: error wrapping ------------------------------ + + @Nested + @DisplayName("attachFilesToPdf error handling") + class ErrorHandlingTests { + + @Test + @DisplayName("IOException from the factory load propagates to the caller") + void factoryIOException_propagates() throws Exception { + byte[] pdfBytes = {0x25, 0x50, 0x44, 0x46}; // "%PDF" + when(pdfDocumentFactory.load(pdfBytes)) + .thenThrow(new java.io.IOException("boom from factory")); + + List attachments = new ArrayList<>(); + attachments.add(attachment("a.pdf", "data".getBytes(StandardCharsets.UTF_8))); + + java.io.IOException ex = + assertThrows( + java.io.IOException.class, + () -> + PdfAttachmentHandler.attachFilesToPdf( + pdfBytes, attachments, pdfDocumentFactory)); + assertTrue(ex.getMessage().contains("boom from factory")); + } + } + + // ----- AttachmentMarkerPositionFinder -------------------------------- + + @Nested + @DisplayName("AttachmentMarkerPositionFinder") + class MarkerFinderTests { + + @Test + @DisplayName("finds marker positions inside an attachments section") + void findsMarkerPositions() throws Exception { + byte[] pdfBytes = + pdfWithLines( + "Some intro text", + "Attachments (2)", + "@invoice.pdf (10 KB)", + "@photo.png (4 KB)"); + + try (PDDocument doc = Loader.loadPDF(pdfBytes)) { + PdfAttachmentHandler.AttachmentMarkerPositionFinder finder = + new PdfAttachmentHandler.AttachmentMarkerPositionFinder(); + finder.setSortByPosition(false); + String returned = finder.getText(doc); + + // getText is overridden to return an empty string (positions are the payload). + assertEquals("", returned); + + List positions = finder.getPositions(); + assertEquals(2, positions.size()); + + List filenames = + positions.stream() + .map(PdfAttachmentHandler.MarkerPosition::getFilename) + .toList(); + assertTrue(filenames.contains("invoice.pdf")); + assertTrue(filenames.contains("photo.png")); + + for (PdfAttachmentHandler.MarkerPosition p : positions) { + assertEquals("@", p.getCharacter()); + assertEquals(0, p.getPageIndex()); + } + } + } + + @Test + @DisplayName("collects no positions when there is no attachments section") + void noAttachmentSection_noPositions() throws Exception { + byte[] pdfBytes = + pdfWithLines("Plain email body", "Contact us @ support address", "Goodbye"); + + try (PDDocument doc = Loader.loadPDF(pdfBytes)) { + PdfAttachmentHandler.AttachmentMarkerPositionFinder finder = + new PdfAttachmentHandler.AttachmentMarkerPositionFinder(); + finder.getText(doc); + assertTrue(finder.getPositions().isEmpty()); + } + } + + @Test + @DisplayName("sortByPosition reorders collected positions deterministically") + void sortByPosition_sortsPositions() throws Exception { + byte[] pdfBytes = + pdfWithLines("Attachments (2)", "@first.pdf (1 KB)", "@second.pdf (2 KB)"); + + try (PDDocument doc = Loader.loadPDF(pdfBytes)) { + PdfAttachmentHandler.AttachmentMarkerPositionFinder finder = + new PdfAttachmentHandler.AttachmentMarkerPositionFinder(); + finder.setSortByPosition(true); + finder.getText(doc); + + List positions = finder.getPositions(); + assertEquals(2, positions.size()); + // With descending-Y sorting and same page, the higher-on-page marker comes first. + assertTrue(positions.get(0).getY() >= positions.get(1).getY()); + } + } + } + + // ----- processInlineImages ------------------------------------------- + + @Nested + @DisplayName("processInlineImages") + class ProcessInlineImagesTests { + + @Test + @DisplayName("replaces a cid: reference with an inline base64 data URI") + void replacesCidWithDataUri() { + byte[] imageData = {(byte) 0x89, 'P', 'N', 'G'}; + EmlParser.EmailAttachment img = new EmlParser.EmailAttachment(); + img.setEmbedded(true); + img.setContentId("img001"); + img.setFilename("pic.png"); + img.setContentType("image/png"); + img.setData(imageData); + + EmlParser.EmailContent content = new EmlParser.EmailContent(); + List list = new ArrayList<>(); + list.add(img); + content.setAttachments(list); + + String html = ""; + String result = PdfAttachmentHandler.processInlineImages(html, content); + + String expectedB64 = Base64.getEncoder().encodeToString(imageData); + assertTrue(result.contains("data:image/png;base64," + expectedB64)); + assertFalse(result.contains("cid:img001")); + } + + @Test + @DisplayName("leaves a cid: reference untouched when no attachment matches it") + void unmatchedCid_isUnchanged() { + EmlParser.EmailAttachment img = new EmlParser.EmailAttachment(); + img.setEmbedded(true); + img.setContentId("known"); + img.setFilename("known.png"); + img.setContentType("image/png"); + img.setData(new byte[] {1, 2, 3}); + + EmlParser.EmailContent content = new EmlParser.EmailContent(); + List list = new ArrayList<>(); + list.add(img); + content.setAttachments(list); + + String html = ""; + String result = PdfAttachmentHandler.processInlineImages(html, content); + + // The unknown cid reference is preserved verbatim. + assertTrue(result.contains("cid:unknown")); + } + + @Test + @DisplayName("returns original html when there are no embedded images to map") + void noEmbeddedImages_returnsOriginal() { + EmlParser.EmailAttachment nonEmbedded = new EmlParser.EmailAttachment(); + nonEmbedded.setEmbedded(false); + nonEmbedded.setContentId("x"); + nonEmbedded.setData(new byte[] {1}); + + EmlParser.EmailContent content = new EmlParser.EmailContent(); + List list = new ArrayList<>(); + list.add(nonEmbedded); + content.setAttachments(list); + + String html = ""; + assertEquals(html, PdfAttachmentHandler.processInlineImages(html, content)); + } + } + + // ----- formatEmailDate (deterministic UTC) --------------------------- + + @Nested + @DisplayName("formatEmailDate determinism") + class FormatEmailDateTests { + + @Test + @DisplayName("a known instant formats to a stable UTC string regardless of input zone") + void zonedDateTime_formatsToUtc() { + // 2024-06-15 12:00 in Tokyo is 03:00 UTC the same day. + ZonedDateTime tokyo = + ZonedDateTime.of(2024, 6, 15, 12, 0, 0, 0, ZoneId.of("Asia/Tokyo")); + String result = PdfAttachmentHandler.formatEmailDate(tokyo); + assertEquals("Sat, Jun 15, 2024 at 3:00 AM UTC", result); + } + + @Test + @DisplayName("Date overload converts a fixed epoch instant to the expected UTC string") + void date_formatsToUtc() { + // Epoch milli 0 == 1970-01-01T00:00:00Z. + String result = PdfAttachmentHandler.formatEmailDate(new Date(0L)); + assertEquals("Thu, Jan 1, 1970 at 12:00 AM UTC", result); + } + + @Test + @DisplayName("null inputs yield an empty string for both overloads") + void nullInputs_returnEmpty() { + assertEquals("", PdfAttachmentHandler.formatEmailDate((Date) null)); + assertEquals("", PdfAttachmentHandler.formatEmailDate((ZonedDateTime) null)); + } + } +} diff --git a/app/common/src/test/java/stirling/software/common/util/PdfUtilsGapTest.java b/app/common/src/test/java/stirling/software/common/util/PdfUtilsGapTest.java new file mode 100644 index 0000000000..b0f5c750ac --- /dev/null +++ b/app/common/src/test/java/stirling/software/common/util/PdfUtilsGapTest.java @@ -0,0 +1,551 @@ +package stirling.software.common.util; + +import static org.junit.jupiter.api.Assertions.*; +import static org.mockito.Mockito.*; + +import java.awt.Color; +import java.awt.Graphics2D; +import java.awt.image.BufferedImage; +import java.io.ByteArrayOutputStream; +import java.io.IOException; + +import javax.imageio.ImageIO; + +import org.apache.pdfbox.pdmodel.PDDocument; +import org.apache.pdfbox.pdmodel.PDPage; +import org.apache.pdfbox.pdmodel.PDPageContentStream; +import org.apache.pdfbox.pdmodel.common.PDRectangle; +import org.apache.pdfbox.pdmodel.font.PDType1Font; +import org.apache.pdfbox.pdmodel.font.Standard14Fonts; +import org.apache.pdfbox.pdmodel.graphics.image.PDImageXObject; +import org.apache.pdfbox.rendering.ImageType; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.springframework.http.MediaType; +import org.springframework.mock.web.MockMultipartFile; +import org.springframework.web.multipart.MultipartFile; + +import stirling.software.common.service.CustomPDFDocumentFactory; + +/** + * Additional unit tests for {@link PdfUtils} targeting methods not exercised by {@code + * PdfUtilsTest}: convertFromPdf, convertPdfToPdfImage, imageToPdf, addImageToDocument, + * overlayImage, containsTextInFile and the error branch of pageSize. + */ +@ExtendWith(MockitoExtension.class) +class PdfUtilsGapTest { + + @Mock private CustomPDFDocumentFactory pdfDocumentFactory; + + // ---- helpers ------------------------------------------------------------ + + /** Builds a tiny single-page PDF and returns it serialized to bytes. */ + private static byte[] simplePdfBytes() throws IOException { + try (PDDocument doc = new PDDocument()) { + doc.addPage(new PDPage(PDRectangle.A4)); + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + doc.save(baos); + return baos.toByteArray(); + } + } + + /** Builds a PDF with the given number of empty A4 pages. */ + private static PDDocument docWithPages(int pages) { + PDDocument doc = new PDDocument(); + for (int i = 0; i < pages; i++) { + doc.addPage(new PDPage(PDRectangle.A4)); + } + return doc; + } + + /** + * Builds a PDF with the given number of tiny pages. convertPdfToPdfImage rasterises every page + * at 300 DPI, so page area drives the cost; tiny pages keep render work minimal while still + * exercising the per-page loop. Page size is non-square so dimension preservation stays + * verifiable. + */ + private static PDDocument docWithTinyPages(int pages, float width, float height) { + PDDocument doc = new PDDocument(); + for (int i = 0; i < pages; i++) { + doc.addPage(new PDPage(new PDRectangle(width, height))); + } + return doc; + } + + /** Builds a PDF whose pages each contain the given text phrase. */ + private static PDDocument docWithText(String... pageTexts) throws IOException { + PDDocument doc = new PDDocument(); + for (String text : pageTexts) { + PDPage page = new PDPage(PDRectangle.A4); + doc.addPage(page); + try (PDPageContentStream cs = new PDPageContentStream(doc, page)) { + cs.beginText(); + cs.setFont(new PDType1Font(Standard14Fonts.FontName.HELVETICA), 12); + cs.newLineAtOffset(100, 700); + cs.showText(text); + cs.endText(); + } + } + return doc; + } + + /** Encodes a small solid-color image to bytes in the requested format. */ + private static byte[] imageBytes(String format, Color color) throws IOException { + BufferedImage img = new BufferedImage(20, 20, BufferedImage.TYPE_INT_RGB); + Graphics2D g = img.createGraphics(); + g.setColor(color); + g.fillRect(0, 0, 20, 20); + g.dispose(); + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + ImageIO.write(img, format, baos); + return baos.toByteArray(); + } + + // ---- convertFromPdf ----------------------------------------------------- + + @Nested + @DisplayName("convertFromPdf") + class ConvertFromPdf { + + @Test + @DisplayName("single PNG image is produced from a one-page PDF") + void singlePng() throws Exception { + byte[] bytes = simplePdfBytes(); + when(pdfDocumentFactory.load(bytes)).thenReturn(docWithPages(1)); + + byte[] out = + PdfUtils.convertFromPdf( + pdfDocumentFactory, bytes, "png", ImageType.RGB, true, 72, "doc", true); + + assertNotNull(out); + assertTrue(out.length > 0); + // A valid PNG starts with the 8-byte PNG signature. + assertEquals((byte) 0x89, out[0]); + assertEquals('P', out[1]); + assertEquals('N', out[2]); + assertEquals('G', out[3]); + } + + @Test + @DisplayName("single combined JPEG image is produced for multi-page PDF") + void singleJpegMultiPage() throws Exception { + byte[] bytes = simplePdfBytes(); + when(pdfDocumentFactory.load(bytes)).thenReturn(docWithPages(2)); + + byte[] out = + PdfUtils.convertFromPdf( + pdfDocumentFactory, + bytes, + "jpg", + ImageType.RGB, + true, + 72, + "doc", + false); + + assertNotNull(out); + assertTrue(out.length > 0); + // JPEG magic bytes. + assertEquals((byte) 0xFF, out[0]); + assertEquals((byte) 0xD8, out[1]); + } + + @Test + @DisplayName("single TIFF image sequence is produced for multi-page PDF") + void singleTiffMultiPage() throws Exception { + byte[] bytes = simplePdfBytes(); + when(pdfDocumentFactory.load(bytes)).thenReturn(docWithPages(2)); + + byte[] out = + PdfUtils.convertFromPdf( + pdfDocumentFactory, + bytes, + "tiff", + ImageType.RGB, + true, + 72, + "doc", + true); + + assertNotNull(out); + assertTrue(out.length > 0); + } + + @Test + @DisplayName("non-single image mode returns a non-empty zip of per-page images") + void zipOfImages() throws Exception { + byte[] bytes = simplePdfBytes(); + when(pdfDocumentFactory.load(bytes)).thenReturn(docWithPages(2)); + + byte[] out = + PdfUtils.convertFromPdf( + pdfDocumentFactory, + bytes, + "png", + ImageType.RGB, + false, + 72, + "myfile", + true); + + assertNotNull(out); + assertTrue(out.length > 0); + // ZIP local-file-header magic "PK\003\004". + assertEquals('P', out[0]); + assertEquals('K', out[1]); + } + + @Test + @DisplayName("DPI above the safe limit throws IllegalArgumentException") + void dpiTooHighThrows() { + byte[] bytes = new byte[] {1, 2, 3}; + // The DPI check happens before the document is loaded. + assertThrows( + IllegalArgumentException.class, + () -> + PdfUtils.convertFromPdf( + pdfDocumentFactory, + bytes, + "png", + ImageType.RGB, + true, + 9999, + "doc", + true)); + } + + @Test + @DisplayName("annotations excluded path still renders successfully") + void withoutAnnotations() throws Exception { + byte[] bytes = simplePdfBytes(); + when(pdfDocumentFactory.load(bytes)).thenReturn(docWithPages(1)); + + byte[] out = + PdfUtils.convertFromPdf( + pdfDocumentFactory, + bytes, + "png", + ImageType.RGB, + true, + 72, + "doc", + false); + + assertNotNull(out); + assertTrue(out.length > 0); + } + } + + // ---- convertPdfToPdfImage ----------------------------------------------- + + @Nested + @DisplayName("convertPdfToPdfImage") + class ConvertPdfToPdfImage { + + @Test + @DisplayName("returns a new document with the same page count") + void preservesPageCount() throws IOException { + // Page size is irrelevant to the count assertion; tiny pages avoid a 300 DPI A4 raster. + try (PDDocument source = docWithTinyPages(2, 6f, 9f); + PDDocument result = PdfUtils.convertPdfToPdfImage(source)) { + assertNotNull(result); + assertEquals(2, result.getNumberOfPages()); + } + } + + @Test + @DisplayName("preserves page dimensions of the source") + void preservesPageSize() throws IOException { + // A small non-square page still proves width/height are carried through (and not + // swapped) without rastering a full LETTER page at 300 DPI. + float width = 60f; + float height = 90f; + try (PDDocument source = new PDDocument()) { + source.addPage(new PDPage(new PDRectangle(width, height))); + try (PDDocument result = PdfUtils.convertPdfToPdfImage(source)) { + PDRectangle box = result.getPage(0).getMediaBox(); + assertEquals(width, box.getWidth(), 0.5f); + assertEquals(height, box.getHeight(), 0.5f); + } + } + } + + @Test + @DisplayName("empty document yields an empty document") + void emptyDocument() throws IOException { + try (PDDocument source = new PDDocument(); + PDDocument result = PdfUtils.convertPdfToPdfImage(source)) { + assertEquals(0, result.getNumberOfPages()); + } + } + } + + // ---- imageToPdf --------------------------------------------------------- + + @Nested + @DisplayName("imageToPdf") + class ImageToPdf { + + private byte[] runImageToPdf(MultipartFile[] files, String fitOption, boolean autoRotate) + throws IOException { + when(pdfDocumentFactory.createNewDocument()).thenReturn(new PDDocument()); + return PdfUtils.imageToPdf(files, fitOption, autoRotate, "color", pdfDocumentFactory); + } + + @Test + @DisplayName("single PNG image becomes a one-page PDF") + void singlePngImage() throws IOException { + MockMultipartFile file = + new MockMultipartFile( + "file", + "image.png", + MediaType.IMAGE_PNG_VALUE, + imageBytes("png", Color.RED)); + + byte[] pdfOut = runImageToPdf(new MultipartFile[] {file}, "fillPage", false); + + assertNotNull(pdfOut); + try (PDDocument doc = org.apache.pdfbox.Loader.loadPDF(pdfOut)) { + assertEquals(1, doc.getNumberOfPages()); + } + } + + @Test + @DisplayName("JPEG image uses the lossy factory path and produces a PDF") + void jpegImage() throws IOException { + MockMultipartFile file = + new MockMultipartFile( + "file", + "image.jpg", + MediaType.IMAGE_JPEG_VALUE, + imageBytes("jpg", Color.BLUE)); + + byte[] pdfOut = runImageToPdf(new MultipartFile[] {file}, "maintainAspectRatio", false); + + assertNotNull(pdfOut); + try (PDDocument doc = org.apache.pdfbox.Loader.loadPDF(pdfOut)) { + assertEquals(1, doc.getNumberOfPages()); + } + } + + @Test + @DisplayName("fitDocumentToImage sizes the page to the image") + void fitDocumentToImage() throws IOException { + MockMultipartFile file = + new MockMultipartFile( + "file", + "image.png", + MediaType.IMAGE_PNG_VALUE, + imageBytes("png", Color.GREEN)); + + byte[] pdfOut = runImageToPdf(new MultipartFile[] {file}, "fitDocumentToImage", false); + + try (PDDocument doc = org.apache.pdfbox.Loader.loadPDF(pdfOut)) { + PDRectangle box = doc.getPage(0).getMediaBox(); + assertEquals(20f, box.getWidth(), 0.5f); + assertEquals(20f, box.getHeight(), 0.5f); + } + } + + @Test + @DisplayName("multiple images become multiple pages") + void multipleImages() throws IOException { + MockMultipartFile a = + new MockMultipartFile( + "file", + "a.png", + MediaType.IMAGE_PNG_VALUE, + imageBytes("png", Color.RED)); + MockMultipartFile b = + new MockMultipartFile( + "file", + "b.png", + MediaType.IMAGE_PNG_VALUE, + imageBytes("png", Color.BLUE)); + + byte[] pdfOut = runImageToPdf(new MultipartFile[] {a, b}, "fillPage", true); + + try (PDDocument doc = org.apache.pdfbox.Loader.loadPDF(pdfOut)) { + assertEquals(2, doc.getNumberOfPages()); + } + } + } + + // ---- addImageToDocument ------------------------------------------------- + + @Nested + @DisplayName("addImageToDocument") + class AddImageToDocument { + + private PDImageXObject portraitImage(PDDocument doc) throws IOException { + BufferedImage img = new BufferedImage(40, 80, BufferedImage.TYPE_INT_RGB); + return org.apache.pdfbox.pdmodel.graphics.image.LosslessFactory.createFromImage( + doc, img); + } + + private PDImageXObject landscapeImage(PDDocument doc) throws IOException { + BufferedImage img = new BufferedImage(80, 40, BufferedImage.TYPE_INT_RGB); + return org.apache.pdfbox.pdmodel.graphics.image.LosslessFactory.createFromImage( + doc, img); + } + + @Test + @DisplayName("fillPage adds an A4 page") + void fillPage() throws IOException { + try (PDDocument doc = new PDDocument()) { + PdfUtils.addImageToDocument(doc, portraitImage(doc), "fillPage", false); + assertEquals(1, doc.getNumberOfPages()); + PDRectangle box = doc.getPage(0).getMediaBox(); + assertEquals(PDRectangle.A4.getWidth(), box.getWidth(), 0.5f); + } + } + + @Test + @DisplayName("maintainAspectRatio adds an A4 page and centers the image") + void maintainAspectRatio() throws IOException { + try (PDDocument doc = new PDDocument()) { + PdfUtils.addImageToDocument(doc, portraitImage(doc), "maintainAspectRatio", false); + assertEquals(1, doc.getNumberOfPages()); + } + } + + @Test + @DisplayName("fitDocumentToImage sizes the page to the image bounds") + void fitDocumentToImage() throws IOException { + try (PDDocument doc = new PDDocument()) { + PdfUtils.addImageToDocument(doc, portraitImage(doc), "fitDocumentToImage", false); + PDRectangle box = doc.getPage(0).getMediaBox(); + assertEquals(40f, box.getWidth(), 0.5f); + assertEquals(80f, box.getHeight(), 0.5f); + } + } + + @Test + @DisplayName("autoRotate with a landscape image swaps to landscape A4") + void autoRotateLandscape() throws IOException { + try (PDDocument doc = new PDDocument()) { + PdfUtils.addImageToDocument(doc, landscapeImage(doc), "maintainAspectRatio", true); + PDRectangle box = doc.getPage(0).getMediaBox(); + // Landscape: width should now exceed height. + assertTrue(box.getWidth() > box.getHeight()); + } + } + + @Test + @DisplayName("unknown fit option still adds a page without drawing") + void unknownFitOption() throws IOException { + try (PDDocument doc = new PDDocument()) { + PdfUtils.addImageToDocument(doc, portraitImage(doc), "unknownOption", false); + assertEquals(1, doc.getNumberOfPages()); + } + } + } + + // ---- overlayImage ------------------------------------------------------- + + @Nested + @DisplayName("overlayImage") + class OverlayImage { + + @Test + @DisplayName("overlays only the first page when everyPage is false") + void firstPageOnly() throws IOException { + byte[] pdf = simplePdfBytes(); + when(pdfDocumentFactory.load(pdf)).thenReturn(docWithPages(3)); + byte[] image = imageBytes("png", Color.RED); + + byte[] out = PdfUtils.overlayImage(pdfDocumentFactory, pdf, image, 10f, 10f, false); + + assertNotNull(out); + try (PDDocument doc = org.apache.pdfbox.Loader.loadPDF(out)) { + assertEquals(3, doc.getNumberOfPages()); + } + } + + @Test + @DisplayName("overlays every page when everyPage is true") + void everyPage() throws IOException { + byte[] pdf = simplePdfBytes(); + when(pdfDocumentFactory.load(pdf)).thenReturn(docWithPages(2)); + byte[] image = imageBytes("png", Color.BLUE); + + byte[] out = PdfUtils.overlayImage(pdfDocumentFactory, pdf, image, 0f, 0f, true); + + assertNotNull(out); + assertTrue(out.length > 0); + try (PDDocument doc = org.apache.pdfbox.Loader.loadPDF(out)) { + assertEquals(2, doc.getNumberOfPages()); + } + } + } + + // ---- containsTextInFile ------------------------------------------------- + + @Nested + @DisplayName("containsTextInFile") + class ContainsTextInFile { + + @Test + @DisplayName("finds text when searching all pages") + void allPagesMatch() throws IOException { + PDDocument doc = docWithText("HelloWorld"); + assertTrue(PdfUtils.containsTextInFile(doc, "HelloWorld", "all")); + } + + @Test + @DisplayName("null pagesToCheck is treated as all pages") + void nullPagesTreatedAsAll() throws IOException { + PDDocument doc = docWithText("FindThis"); + assertTrue(PdfUtils.containsTextInFile(doc, "FindThis", null)); + } + + @Test + @DisplayName("returns false when text is absent") + void noMatch() throws IOException { + PDDocument doc = docWithText("SomeText"); + assertFalse(PdfUtils.containsTextInFile(doc, "Missing", "all")); + } + + @Test + @DisplayName("matches text on an individual page number") + void individualPage() throws IOException { + PDDocument doc = docWithText("PageOne", "PageTwo"); + assertTrue(PdfUtils.containsTextInFile(doc, "PageTwo", "2")); + } + + @Test + @DisplayName("matches text within a page range") + void pageRange() throws IOException { + PDDocument doc = docWithText("Alpha", "Beta", "Gamma"); + assertTrue(PdfUtils.containsTextInFile(doc, "Gamma", "1-3")); + } + + @Test + @DisplayName("whitespace in the page spec is stripped before parsing") + void whitespaceStripped() throws IOException { + PDDocument doc = docWithText("One", "Two"); + assertTrue(PdfUtils.containsTextInFile(doc, "Two", " 1 , 2 ")); + } + } + + // ---- pageSize error branch --------------------------------------------- + + @Nested + @DisplayName("pageSize parsing") + class PageSizeParsing { + + @Test + @DisplayName("non-numeric expected size throws NumberFormatException") + void nonNumericThrows() throws IOException { + try (PDDocument doc = new PDDocument()) { + doc.addPage(new PDPage(PDRectangle.A4)); + assertThrows( + NumberFormatException.class, () -> PdfUtils.pageSize(doc, "widthxheight")); + } + } + } +} diff --git a/app/common/src/test/java/stirling/software/common/util/ProcessExecutorGapTest.java b/app/common/src/test/java/stirling/software/common/util/ProcessExecutorGapTest.java new file mode 100644 index 0000000000..961ae0b6cb --- /dev/null +++ b/app/common/src/test/java/stirling/software/common/util/ProcessExecutorGapTest.java @@ -0,0 +1,634 @@ +package stirling.software.common.util; + +import static org.junit.jupiter.api.Assertions.*; + +import java.io.IOException; +import java.lang.reflect.InvocationTargetException; +import java.lang.reflect.Method; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.ArrayList; +import java.util.List; + +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.io.TempDir; + +import stirling.software.common.model.ApplicationProperties; + +/** + * Gap-filling unit tests for {@link ProcessExecutor}. Focused on the pure logic that can be + * exercised without launching any real OS process: command validation branches, the unoserver + * endpoint helper methods (via reflection), the {@link ProcessExecutor.Processes} enum, the + * singleton/getInstance behaviour, the static unoserver pool setter, and the nested {@link + * ProcessExecutor.ProcessExecutorResult} value type. + */ +class ProcessExecutorGapTest { + + // ----- reflection helpers ------------------------------------------------- + + private void invokeValidateCommand(ProcessExecutor executor, List command) + throws Exception { + Method method = ProcessExecutor.class.getDeclaredMethod("validateCommand", List.class); + method.setAccessible(true); + try { + method.invoke(executor, command); + } catch (InvocationTargetException e) { + throw (Exception) e.getCause(); + } + } + + @SuppressWarnings("unchecked") + private List invokeStripUnoEndpointArgs(ProcessExecutor executor, List command) + throws Exception { + Method method = ProcessExecutor.class.getDeclaredMethod("stripUnoEndpointArgs", List.class); + method.setAccessible(true); + return (List) method.invoke(executor, command); + } + + @SuppressWarnings("unchecked") + private List invokeApplyUnoServerEndpoint( + ProcessExecutor executor, + List command, + ApplicationProperties.ProcessExecutor.UnoServerEndpoint endpoint) + throws Exception { + Method method = + ProcessExecutor.class.getDeclaredMethod( + "applyUnoServerEndpoint", + List.class, + ApplicationProperties.ProcessExecutor.UnoServerEndpoint.class); + method.setAccessible(true); + return (List) method.invoke(executor, command, endpoint); + } + + private boolean invokeShouldUseUnoServerPool(ProcessExecutor executor, List command) + throws Exception { + Method method = + ProcessExecutor.class.getDeclaredMethod("shouldUseUnoServerPool", List.class); + method.setAccessible(true); + return (boolean) method.invoke(executor, command); + } + + private ProcessExecutor qpdfExecutor() { + return ProcessExecutor.getInstance(ProcessExecutor.Processes.QPDF); + } + + private ProcessExecutor libreOfficeExecutor() { + return ProcessExecutor.getInstance(ProcessExecutor.Processes.LIBRE_OFFICE); + } + + /** The static unoserver pool is global state; clear it after every test that touches it. */ + @AfterEach + void resetUnoServerPool() { + ProcessExecutor.setUnoServerPool(null); + } + + // ----- validateCommand deeper branches ----------------------------------- + + @Nested + @DisplayName("validateCommand path/executable branches") + class ValidateCommandPathTests { + + @Test + @DisplayName("absolute path executable that does not exist is rejected") + void absolutePathExecutableMissing() { + String bogus = + System.getProperty("os.name").toLowerCase().contains("win") + ? "C:\\definitely\\does\\not\\exist\\tool.exe" + : "/definitely/does/not/exist/tool"; + IllegalArgumentException ex = + assertThrows( + IllegalArgumentException.class, + () -> invokeValidateCommand(qpdfExecutor(), List.of(bogus))); + assertTrue(ex.getMessage().contains("does not exist")); + } + + @Test + @DisplayName("path that exists but is a directory is rejected (not a regular file)") + void directoryPathExecutableRejected(@TempDir Path tempDir) { + String dirPath = tempDir.toString(); + // Ensure the path contains a separator so the path-based validation branch is taken. + assertTrue(dirPath.contains("/") || dirPath.contains("\\")); + IllegalArgumentException ex = + assertThrows( + IllegalArgumentException.class, + () -> invokeValidateCommand(qpdfExecutor(), List.of(dirPath))); + assertTrue(ex.getMessage().contains("not a regular file")); + } + + @Test + @DisplayName("absolute path to an existing regular file passes validation") + void existingRegularFileExecutablePasses(@TempDir Path tempDir) throws Exception { + Path file = tempDir.resolve("fakebinary"); + Files.writeString(file, "#!/bin/sh\n"); + assertTrue(Files.exists(file)); + // Should not throw. + invokeValidateCommand(qpdfExecutor(), List.of(file.toString(), "--version")); + } + + @Test + @DisplayName("path traversal anywhere in the executable is rejected") + void pathTraversalInExecutable() { + IllegalArgumentException ex = + assertThrows( + IllegalArgumentException.class, + () -> + invokeValidateCommand( + qpdfExecutor(), List.of("/usr/bin/../bin/tool"))); + assertTrue(ex.getMessage().contains("path traversal")); + } + + @Test + @DisplayName("null byte / newline checks run across every argument, not just the first") + void invalidCharactersInLaterArgument() { + assertThrows( + IllegalArgumentException.class, + () -> invokeValidateCommand(qpdfExecutor(), List.of("qpdf", "ok", "bad\0arg"))); + assertThrows( + IllegalArgumentException.class, + () -> invokeValidateCommand(qpdfExecutor(), List.of("qpdf", "ok", "bad\narg"))); + assertThrows( + IllegalArgumentException.class, + () -> invokeValidateCommand(qpdfExecutor(), List.of("qpdf", "ok", "bad\rarg"))); + } + + @Test + @DisplayName("relative simple command (no separators) is trusted and passes") + void relativeSimpleCommandPasses() throws Exception { + invokeValidateCommand(qpdfExecutor(), List.of("qpdf", "--help")); + } + + @Test + @DisplayName("null first-argument executable is rejected") + void nullExecutableRejected() { + List command = new ArrayList<>(); + command.add(null); + // null arg is caught by the per-arg null check before the executable check. + assertThrows( + IllegalArgumentException.class, + () -> invokeValidateCommand(qpdfExecutor(), command)); + } + } + + // ----- stripUnoEndpointArgs ---------------------------------------------- + + @Nested + @DisplayName("stripUnoEndpointArgs") + class StripUnoEndpointArgsTests { + + @Test + @DisplayName("removes space-separated --host/--port/--host-location/--protocol pairs") + void stripsSpaceSeparatedArgs() throws Exception { + List input = + List.of( + "unoconvert", + "--host", + "1.2.3.4", + "--port", + "9999", + "--host-location", + "remote", + "--protocol", + "https", + "in.docx", + "out.pdf"); + List result = invokeStripUnoEndpointArgs(qpdfExecutor(), input); + assertEquals(List.of("unoconvert", "in.docx", "out.pdf"), result); + } + + @Test + @DisplayName("removes equals-form --host=.../--port=... arguments") + void stripsEqualsFormArgs() throws Exception { + List input = + List.of( + "unoconvert", + "--host=5.6.7.8", + "--port=4002", + "--host-location=local", + "--protocol=http", + "doc.odt"); + List result = invokeStripUnoEndpointArgs(qpdfExecutor(), input); + assertEquals(List.of("unoconvert", "doc.odt"), result); + } + + @Test + @DisplayName("leaves a command without endpoint args unchanged") + void leavesPlainCommandUntouched() throws Exception { + List input = List.of("unoconvert", "in.docx", "out.pdf"); + List result = invokeStripUnoEndpointArgs(qpdfExecutor(), input); + assertEquals(input, result); + } + + @Test + @DisplayName("returns a fresh list, not the same instance") + void returnsNewList() throws Exception { + List input = new ArrayList<>(List.of("unoconvert", "a", "b")); + List result = invokeStripUnoEndpointArgs(qpdfExecutor(), input); + assertNotSame(input, result); + } + } + + // ----- applyUnoServerEndpoint -------------------------------------------- + + private ApplicationProperties.ProcessExecutor.UnoServerEndpoint endpoint( + String host, int port, String hostLocation, String protocol) { + ApplicationProperties.ProcessExecutor.UnoServerEndpoint ep = + new ApplicationProperties.ProcessExecutor.UnoServerEndpoint(); + ep.setHost(host); + ep.setPort(port); + ep.setHostLocation(hostLocation); + ep.setProtocol(protocol); + return ep; + } + + @Nested + @DisplayName("applyUnoServerEndpoint") + class ApplyUnoServerEndpointTests { + + @Test + @DisplayName( + "injects --host/--port after the executable, defaults omit host-location and protocol") + void injectsHostAndPortWithDefaults() throws Exception { + List command = List.of("unoconvert", "in.docx", "out.pdf"); + ApplicationProperties.ProcessExecutor.UnoServerEndpoint ep = + endpoint("9.9.9.9", 7777, "auto", "http"); + List result = invokeApplyUnoServerEndpoint(qpdfExecutor(), command, ep); + assertEquals( + List.of( + "unoconvert", + "--host", + "9.9.9.9", + "--port", + "7777", + "in.docx", + "out.pdf"), + result); + } + + @Test + @DisplayName("non-default host-location and protocol are injected") + void injectsHostLocationAndProtocolWhenNonDefault() throws Exception { + List command = List.of("unoconvert", "in.docx"); + ApplicationProperties.ProcessExecutor.UnoServerEndpoint ep = + endpoint("10.0.0.5", 2200, "remote", "https"); + List result = invokeApplyUnoServerEndpoint(qpdfExecutor(), command, ep); + assertEquals( + List.of( + "unoconvert", + "--host", + "10.0.0.5", + "--port", + "2200", + "--host-location", + "remote", + "--protocol", + "https", + "in.docx"), + result); + } + + @Test + @DisplayName("blank host falls back to 127.0.0.1 and non-positive port falls back to 2003") + void appliesHostAndPortFallbacks() throws Exception { + List command = List.of("unoconvert", "in.docx"); + ApplicationProperties.ProcessExecutor.UnoServerEndpoint ep = + endpoint(" ", 0, "auto", "http"); + List result = invokeApplyUnoServerEndpoint(qpdfExecutor(), command, ep); + assertEquals( + List.of("unoconvert", "--host", "127.0.0.1", "--port", "2003", "in.docx"), + result); + } + + @Test + @DisplayName("invalid host-location and protocol values are normalised to defaults") + void invalidHostLocationAndProtocolNormalised() throws Exception { + List command = List.of("unoconvert", "in.docx"); + ApplicationProperties.ProcessExecutor.UnoServerEndpoint ep = + endpoint("1.1.1.1", 3000, "sideways", "gopher"); + List result = invokeApplyUnoServerEndpoint(qpdfExecutor(), command, ep); + // Both invalid -> normalised to defaults (auto/http) -> neither injected. + assertEquals( + List.of("unoconvert", "--host", "1.1.1.1", "--port", "3000", "in.docx"), + result); + } + + @Test + @DisplayName("host-location and protocol matching is case-insensitive and trimmed") + void hostLocationAndProtocolCaseInsensitive() throws Exception { + List command = List.of("unoconvert", "in.docx"); + ApplicationProperties.ProcessExecutor.UnoServerEndpoint ep = + endpoint("1.1.1.1", 3000, " REMOTE ", " HTTPS "); + List result = invokeApplyUnoServerEndpoint(qpdfExecutor(), command, ep); + assertEquals( + List.of( + "unoconvert", + "--host", + "1.1.1.1", + "--port", + "3000", + "--host-location", + "remote", + "--protocol", + "https", + "in.docx"), + result); + } + + @Test + @DisplayName("existing endpoint args are stripped before re-injection") + void stripsExistingEndpointArgsBeforeInjecting() throws Exception { + List command = List.of("unoconvert", "--host", "old", "--port", "1", "in.docx"); + ApplicationProperties.ProcessExecutor.UnoServerEndpoint ep = + endpoint("2.2.2.2", 2222, "auto", "http"); + List result = invokeApplyUnoServerEndpoint(qpdfExecutor(), command, ep); + assertEquals( + List.of("unoconvert", "--host", "2.2.2.2", "--port", "2222", "in.docx"), + result); + } + + @Test + @DisplayName("null endpoint returns the command unchanged") + void nullEndpointReturnsCommandUnchanged() throws Exception { + List command = List.of("unoconvert", "in.docx"); + List result = invokeApplyUnoServerEndpoint(qpdfExecutor(), command, null); + assertEquals(command, result); + } + + @Test + @DisplayName("empty command returns the command unchanged") + void emptyCommandReturnedUnchanged() throws Exception { + List command = List.of(); + ApplicationProperties.ProcessExecutor.UnoServerEndpoint ep = + endpoint("1.1.1.1", 2003, "auto", "http"); + List result = invokeApplyUnoServerEndpoint(qpdfExecutor(), command, ep); + assertEquals(command, result); + } + } + + // ----- shouldUseUnoServerPool -------------------------------------------- + + @Nested + @DisplayName("shouldUseUnoServerPool") + class ShouldUseUnoServerPoolTests { + + private UnoServerPool nonEmptyPool() { + ApplicationProperties.ProcessExecutor.UnoServerEndpoint ep = + new ApplicationProperties.ProcessExecutor.UnoServerEndpoint(); + return new UnoServerPool(List.of(ep)); + } + + @Test + @DisplayName( + "false for non-LIBRE_OFFICE process type even with a pool and unoconvert command") + void falseForNonLibreOfficeProcessType() throws Exception { + ProcessExecutor.setUnoServerPool(nonEmptyPool()); + assertFalse( + invokeShouldUseUnoServerPool(qpdfExecutor(), List.of("unoconvert", "in.docx"))); + } + + @Test + @DisplayName("false when no pool is configured") + void falseWhenPoolNull() throws Exception { + ProcessExecutor.setUnoServerPool(null); + assertFalse( + invokeShouldUseUnoServerPool( + libreOfficeExecutor(), List.of("unoconvert", "in.docx"))); + } + + @Test + @DisplayName("false when the configured pool is empty") + void falseWhenPoolEmpty() throws Exception { + ProcessExecutor.setUnoServerPool(new UnoServerPool(List.of())); + assertFalse( + invokeShouldUseUnoServerPool( + libreOfficeExecutor(), List.of("unoconvert", "in.docx"))); + } + + @Test + @DisplayName("false for null or empty command") + void falseForNullOrEmptyCommand() throws Exception { + ProcessExecutor.setUnoServerPool(nonEmptyPool()); + assertFalse(invokeShouldUseUnoServerPool(libreOfficeExecutor(), null)); + assertFalse(invokeShouldUseUnoServerPool(libreOfficeExecutor(), List.of())); + } + + @Test + @DisplayName("true for a plain unoconvert command with a non-empty pool") + void trueForUnoconvertCommand() throws Exception { + ProcessExecutor.setUnoServerPool(nonEmptyPool()); + assertTrue( + invokeShouldUseUnoServerPool( + libreOfficeExecutor(), List.of("unoconvert", "in.docx", "out.pdf"))); + } + + @Test + @DisplayName("true for a unoconvert path with directories and a .exe extension") + void trueForUnoconvertWithPathAndExeExtension() throws Exception { + ProcessExecutor.setUnoServerPool(nonEmptyPool()); + assertTrue( + invokeShouldUseUnoServerPool( + libreOfficeExecutor(), + List.of("C:\\tools\\bin\\unoconvert.exe", "in.docx"))); + assertTrue( + invokeShouldUseUnoServerPool( + libreOfficeExecutor(), + List.of("/usr/local/bin/unoconvert", "in.docx"))); + } + + @Test + @DisplayName("true for the legacy 'unoconv' executable name") + void trueForLegacyUnoconv() throws Exception { + ProcessExecutor.setUnoServerPool(nonEmptyPool()); + assertTrue( + invokeShouldUseUnoServerPool( + libreOfficeExecutor(), List.of("unoconv", "in.docx"))); + } + + @Test + @DisplayName("false for soffice, which must not be routed through the pool") + void falseForSoffice() throws Exception { + ProcessExecutor.setUnoServerPool(nonEmptyPool()); + assertFalse( + invokeShouldUseUnoServerPool( + libreOfficeExecutor(), + List.of("/usr/bin/soffice", "--headless", "in.docx"))); + } + } + + // ----- Processes enum ----------------------------------------------------- + + @Nested + @DisplayName("Processes enum") + class ProcessesEnumTests { + + @Test + @DisplayName("contains all expected process types") + void containsExpectedValues() { + ProcessExecutor.Processes[] values = ProcessExecutor.Processes.values(); + assertEquals(13, values.length); + assertEquals( + ProcessExecutor.Processes.LIBRE_OFFICE, + ProcessExecutor.Processes.valueOf("LIBRE_OFFICE")); + assertEquals( + ProcessExecutor.Processes.CFF_CONVERTER, + ProcessExecutor.Processes.valueOf("CFF_CONVERTER")); + assertEquals( + ProcessExecutor.Processes.FFMPEG, ProcessExecutor.Processes.valueOf("FFMPEG")); + } + + @Test + @DisplayName("valueOf rejects an unknown name") + void valueOfRejectsUnknown() { + assertThrows( + IllegalArgumentException.class, + () -> ProcessExecutor.Processes.valueOf("NOT_A_PROCESS")); + } + + @Test + @DisplayName("getInstance resolves a non-null singleton for every enum value") + void getInstanceForEveryProcessType() { + for (ProcessExecutor.Processes p : ProcessExecutor.Processes.values()) { + ProcessExecutor instance = ProcessExecutor.getInstance(p); + assertNotNull(instance, "instance should not be null for " + p); + // Same key returns the cached singleton. + assertSame(instance, ProcessExecutor.getInstance(p)); + } + } + } + + // ----- getInstance / liveUpdates ----------------------------------------- + + @Nested + @DisplayName("getInstance behaviour") + class GetInstanceTests { + + @Test + @DisplayName("single-arg getInstance delegates to liveUpdates=true and is cached") + void singleArgDelegatesAndCaches() { + ProcessExecutor a = ProcessExecutor.getInstance(ProcessExecutor.Processes.GHOSTSCRIPT); + ProcessExecutor b = + ProcessExecutor.getInstance(ProcessExecutor.Processes.GHOSTSCRIPT, true); + assertSame(a, b); + } + + @Test + @DisplayName("the liveUpdates flag of the first call wins because the instance is cached") + void firstCallWinsForCachedInstance() { + // First resolution for this type fixes its configuration. + ProcessExecutor first = + ProcessExecutor.getInstance(ProcessExecutor.Processes.OCR_MY_PDF, false); + ProcessExecutor second = + ProcessExecutor.getInstance(ProcessExecutor.Processes.OCR_MY_PDF, true); + assertSame(first, second); + } + } + + // ----- setUnoServerPool --------------------------------------------------- + + @Nested + @DisplayName("setUnoServerPool") + class SetUnoServerPoolTests { + + @Test + @DisplayName("setting then clearing the pool flips shouldUseUnoServerPool") + void poolSetterAffectsRouting() throws Exception { + ProcessExecutor exec = libreOfficeExecutor(); + List command = List.of("unoconvert", "in.docx"); + + ProcessExecutor.setUnoServerPool(null); + assertFalse(invokeShouldUseUnoServerPool(exec, command)); + + ApplicationProperties.ProcessExecutor.UnoServerEndpoint ep = + new ApplicationProperties.ProcessExecutor.UnoServerEndpoint(); + ProcessExecutor.setUnoServerPool(new UnoServerPool(List.of(ep))); + assertTrue(invokeShouldUseUnoServerPool(exec, command)); + + ProcessExecutor.setUnoServerPool(null); + assertFalse(invokeShouldUseUnoServerPool(exec, command)); + } + } + + // ----- ProcessExecutorResult --------------------------------------------- + + @Nested + @DisplayName("ProcessExecutorResult value type") + class ProcessExecutorResultTests { + + @Test + @DisplayName("constructor stores rc and messages; setters mutate them") + void constructorAndSetters() { + ProcessExecutor exec = qpdfExecutor(); + ProcessExecutor.ProcessExecutorResult result = exec.new ProcessExecutorResult(0, "ok"); + assertEquals(0, result.getRc()); + assertEquals("ok", result.getMessages()); + + result.setRc(42); + result.setMessages("boom"); + assertEquals(42, result.getRc()); + assertEquals("boom", result.getMessages()); + } + + @Test + @DisplayName("messages may be null") + void allowsNullMessages() { + ProcessExecutor exec = qpdfExecutor(); + ProcessExecutor.ProcessExecutorResult result = exec.new ProcessExecutorResult(3, null); + assertEquals(3, result.getRc()); + assertNull(result.getMessages()); + } + } + + // ----- runCommandWithOutputHandling validation entry point --------------- + + @Nested + @DisplayName("runCommandWithOutputHandling validation (no process launched)") + class RunCommandValidationTests { + + @Test + @DisplayName("empty command is rejected before any process is started") + void emptyCommandRejected() { + ProcessExecutor exec = qpdfExecutor(); + assertThrows( + IllegalArgumentException.class, + () -> exec.runCommandWithOutputHandling(List.of())); + } + + @Test + @DisplayName("command containing a null byte is rejected before any process is started") + void nullByteCommandRejected() { + ProcessExecutor exec = qpdfExecutor(); + assertThrows( + IllegalArgumentException.class, + () -> exec.runCommandWithOutputHandling(List.of("qpdf", "bad\0arg"))); + } + + @Test + @DisplayName("absolute non-existent executable is rejected before any process is started") + void missingAbsoluteExecutableRejected() { + ProcessExecutor exec = qpdfExecutor(); + String bogus = + System.getProperty("os.name").toLowerCase().contains("win") + ? "C:\\no\\such\\tool.exe" + : "/no/such/tool"; + IllegalArgumentException ex = + assertThrows( + IllegalArgumentException.class, + () -> exec.runCommandWithOutputHandling(List.of(bogus))); + assertTrue(ex.getMessage().contains("does not exist")); + } + + @Test + @DisplayName("validation exception type is not an IOException for bad input") + void validationThrowsIllegalArgumentNotIOException() { + ProcessExecutor exec = qpdfExecutor(); + Exception thrown = + assertThrows( + Exception.class, + () -> exec.runCommandWithOutputHandling(List.of("qpdf", "x\ny"))); + assertInstanceOf(IllegalArgumentException.class, thrown); + assertFalse(thrown instanceof IOException); + } + } +} diff --git a/app/common/src/test/resources/pdf-ingestion-fixtures/bordered-table-test_widget.md b/app/common/src/test/resources/pdf-ingestion-fixtures/bordered-table-test_widget.md new file mode 100644 index 0000000000..4b590e4b63 --- /dev/null +++ b/app/common/src/test/resources/pdf-ingestion-fixtures/bordered-table-test_widget.md @@ -0,0 +1,10 @@ +# Widget Inventory Report + +This report lists current stock levels for each warehouse. + +| Region | Units | Status | +|---|---|---| +| North | 1200 | OK | +| South | 950 | Low | +| East | 1430 | OK | +| West | 875 | Low | diff --git a/app/common/src/test/resources/pdf-ingestion-fixtures/bordered-table-test_widget.pdf b/app/common/src/test/resources/pdf-ingestion-fixtures/bordered-table-test_widget.pdf new file mode 100644 index 0000000000..8da041e28d --- /dev/null +++ b/app/common/src/test/resources/pdf-ingestion-fixtures/bordered-table-test_widget.pdf @@ -0,0 +1,74 @@ +%PDF-1.4 +%“Œ‹ž ReportLab Generated PDF document (opensource) +1 0 obj +<< +/F1 2 0 R /F2 3 0 R +>> +endobj +2 0 obj +<< +/BaseFont /Helvetica /Encoding /WinAnsiEncoding /Name /F1 /Subtype /Type1 /Type /Font +>> +endobj +3 0 obj +<< +/BaseFont /Helvetica-Bold /Encoding /WinAnsiEncoding /Name /F2 /Subtype /Type1 /Type /Font +>> +endobj +4 0 obj +<< +/Contents 8 0 R /MediaBox [ 0 0 612 792 ] /Parent 7 0 R /Resources << +/Font 1 0 R /ProcSet [ /PDF /Text /ImageB /ImageC /ImageI ] +>> /Rotate 0 /Trans << + +>> + /Type /Page +>> +endobj +5 0 obj +<< +/PageMode /UseNone /Pages 7 0 R /Type /Catalog +>> +endobj +6 0 obj +<< +/Author (\(anonymous\)) /CreationDate (D:20260603003133+01'00') /Creator (\(unspecified\)) /Keywords () /ModDate (D:20260603003133+01'00') /Producer (ReportLab PDF Library - \(opensource\)) + /Subject (\(unspecified\)) /Title (\(anonymous\)) /Trapped /False +>> +endobj +7 0 obj +<< +/Count 1 /Kids [ 4 0 R ] /Type /Pages +>> +endobj +8 0 obj +<< +/Filter [ /ASCII85Decode /FlateDecode ] /Length 500 +>> +stream +Gas1[9i&Y\%))C:pc(6t3;pQ8%0T@tD+-Nf,MP8j(>R4tMCHL4NI*\1+UNiI'V9NC%VeJKn/YI0J];XQt&X83?=ihrg<*Mcn1n!1nWcDaQPe\P"9gnJuHl(jf]JQgZ[,&^uobI4QF',k"*^S)3c;)GMWC(T'=",ErnS#U=YCUN0&q4+*KmK1Zd*NI\GQDiZUG7;PTja8lulb"\PWWO#WcfI[ZB:6s*3g$be%?JH(n`oaEJ[XE'%QW=HE04M<,;ERm[MS=uYF=nN3jG'f@#?O48Ia,6Y-3m&tTWVq1?DeiBkp.Ug*;lVZX`Z=P.eklHhNV;!R_?QOuoeJ<0%7idG7GM8boU$^>N.N,2^;25]0Z8M<<]XMCct>noC'Qfb?`*[Mo+,F9#>t~>endstream +endobj +xref +0 9 +0000000000 65535 f +0000000061 00000 n +0000000102 00000 n +0000000209 00000 n +0000000321 00000 n +0000000514 00000 n +0000000582 00000 n +0000000862 00000 n +0000000921 00000 n +trailer +<< +/ID +[] +% ReportLab generated PDF document -- digest (opensource) + +/Info 6 0 R +/Root 5 0 R +/Size 9 +>> +startxref +1511 +%%EOF diff --git a/app/common/src/test/resources/pdf-ingestion-fixtures/many-tables-test_stress.md b/app/common/src/test/resources/pdf-ingestion-fixtures/many-tables-test_stress.md new file mode 100644 index 0000000000..4b456165df --- /dev/null +++ b/app/common/src/test/resources/pdf-ingestion-fixtures/many-tables-test_stress.md @@ -0,0 +1,222 @@ +Intro paragraph for section 1. + +| Name | Qty | +|---|---| +| alpha | 101 | +| delta | 201 | + +# Section 2 Heading + +| Name | Qty | Price | +|---|---|---| +| alpha | 101 | charlie | +| delta | 201 | foxtrot | +| golf | 301 | india | + +## Section 3 Heading + +| Name | Qty | Price | Region | +|---|---|---|---| +| alpha | 101 | charlie | 3 | +| delta | 201 | foxtrot | 13 | +| golf | 301 | india | 23 | +| juliet | 401 | lima | 33 | + +Intro paragraph for section 4. + +| Name | Qty | Price | Region | Status | +|---|---|---|---|---| +| alpha | 101 | charlie | 3 | echo | +| delta | 201 | foxtrot | 13 | hotel | +| golf | 301 | india | 23 | kilo | +| juliet | 401 | lima | 33 | november | +| mike | 501 | oscar | 43 | alpha | + +# Section 5 Heading + +| Name | Qty | +|---|---| +| alpha | 101 | +| delta | 201 | +| golf | 301 | +| juliet | 401 | +| mike | 501 | +| papa | 601 | + +| Name | Qty | Price | +|---|---|---| +| alpha | 101 | charlie | +| delta | 201 | foxtrot | + +# Section 7 Heading + +Intro paragraph for section 7. + +| Name | Qty | Price | Region | +|---|---|---|---| +| alpha | 101 | charlie | 3 | +| delta | 201 | foxtrot | 13 | +| golf | 301 | india | 23 | + +## Section 8 Heading + +| Name | Qty | Price | Region | Status | +|---|---|---|---|---| +| alpha | 101 | charlie | 3 | echo | +| delta | 201 | foxtrot | 13 | hotel | +| golf | 301 | india | 23 | kilo | +| juliet | 401 | lima | 33 | november | + +| Name | Qty | +|---|---| +| alpha | 101 | +| delta | 201 | +| golf | 301 | +| juliet | 401 | +| mike | 501 | + +# Section 10 Heading + +Intro paragraph for section 10. + +| Name | Qty | Price | +|---|---|---| +| alpha | 101 | charlie | +| delta | 201 | foxtrot | +| golf | 301 | india | +| juliet | 401 | lima | +| mike | 501 | oscar | +| papa | 601 | bravo | + +| Name | Qty | Price | Region | +|---|---|---|---| +| alpha | 101 | charlie | 3 | +| delta | 201 | foxtrot | 13 | + +# Section 12 Heading + +| Name | Qty | Price | Region | Status | +|---|---|---|---|---| +| alpha | 101 | charlie | 3 | echo | +| delta | 201 | foxtrot | 13 | hotel | +| golf | 301 | india | 23 | kilo | + +## Section 13 Heading + +Intro paragraph for section 13. + +| Name | Qty | +|---|---| +| alpha | 101 | +| delta | 201 | +| golf | 301 | +| juliet | 401 | + +| Name | Qty | Price | +|---|---|---| +| alpha | 101 | charlie | +| delta | 201 | foxtrot | +| golf | 301 | india | +| juliet | 401 | lima | +| mike | 501 | oscar | + +# Section 15 Heading + +| Name | Qty | Price | Region | +|---|---|---|---| +| alpha | 101 | charlie | 3 | +| delta | 201 | foxtrot | 13 | +| golf | 301 | india | 23 | +| juliet | 401 | lima | 33 | +| mike | 501 | oscar | 43 | +| papa | 601 | bravo | 53 | + +Intro paragraph for section 16. + +| Name | Qty | Price | Region | Status | +|---|---|---|---|---| +| alpha | 101 | charlie | 3 | echo | +| delta | 201 | foxtrot | 13 | hotel | + +# Section 17 Heading + +| Name | Qty | +|---|---| +| alpha | 101 | +| delta | 201 | +| golf | 301 | + +## Section 18 Heading + +| Name | Qty | Price | +|---|---|---| +| alpha | 101 | charlie | +| delta | 201 | foxtrot | +| golf | 301 | india | +| juliet | 401 | lima | + +Intro paragraph for section 19. + +| Name | Qty | Price | Region | +|---|---|---|---| +| alpha | 101 | charlie | 3 | +| delta | 201 | foxtrot | 13 | +| golf | 301 | india | 23 | +| juliet | 401 | lima | 33 | +| mike | 501 | oscar | 43 | + +# Section 20 Heading + +| Name | Qty | Price | Region | Status | +|---|---|---|---|---| +| alpha | 101 | charlie | 3 | echo | +| delta | 201 | foxtrot | 13 | hotel | +| golf | 301 | india | 23 | kilo | +| juliet | 401 | lima | 33 | november | +| mike | 501 | oscar | 43 | alpha | +| papa | 601 | bravo | 53 | delta | + +| Name | Qty | +|---|---| +| alpha | 101 | +| delta | 201 | + +# Section 22 Heading + +Intro paragraph for section 22. + +| Name | Qty | Price | +|---|---|---| +| alpha | 101 | charlie | +| delta | 201 | foxtrot | +| golf | 301 | india | + +## Section 23 Heading + +| Name | Qty | Price | Region | +|---|---|---|---| +| alpha | 101 | charlie | 3 | +| delta | 201 | foxtrot | 13 | +| golf | 301 | india | 23 | +| juliet | 401 | lima | 33 | + +| Name | Qty | Price | Region | Status | +|---|---|---|---|---| +| alpha | 101 | charlie | 3 | echo | +| delta | 201 | foxtrot | 13 | hotel | +| golf | 301 | india | 23 | kilo | +| juliet | 401 | lima | 33 | november | +| mike | 501 | oscar | 43 | alpha | + +# Section 25 Heading + +Intro paragraph for section 25. + +| Name | Qty | +|---|---| +| alpha | 101 | +| delta | 201 | +| golf | 301 | +| juliet | 401 | +| mike | 501 | +| papa | 601 | diff --git a/app/common/src/test/resources/pdf-ingestion-fixtures/many-tables-test_stress.pdf b/app/common/src/test/resources/pdf-ingestion-fixtures/many-tables-test_stress.pdf new file mode 100644 index 0000000000..f12925cda3 --- /dev/null +++ b/app/common/src/test/resources/pdf-ingestion-fixtures/many-tables-test_stress.pdf @@ -0,0 +1,169 @@ +%PDF-1.4 +%“Œ‹ž ReportLab Generated PDF document (opensource) +1 0 obj +<< +/F1 2 0 R /F2 3 0 R +>> +endobj +2 0 obj +<< +/BaseFont /Helvetica /Encoding /WinAnsiEncoding /Name /F1 /Subtype /Type1 /Type /Font +>> +endobj +3 0 obj +<< +/BaseFont /Helvetica-Bold /Encoding /WinAnsiEncoding /Name /F2 /Subtype /Type1 /Type /Font +>> +endobj +4 0 obj +<< +/Contents 13 0 R /MediaBox [ 0 0 612 792 ] /Parent 12 0 R /Resources << +/Font 1 0 R /ProcSet [ /PDF /Text /ImageB /ImageC /ImageI ] +>> /Rotate 0 /Trans << + +>> + /Type /Page +>> +endobj +5 0 obj +<< +/Contents 14 0 R /MediaBox [ 0 0 612 792 ] /Parent 12 0 R /Resources << +/Font 1 0 R /ProcSet [ /PDF /Text /ImageB /ImageC /ImageI ] +>> /Rotate 0 /Trans << + +>> + /Type /Page +>> +endobj +6 0 obj +<< +/Contents 15 0 R /MediaBox [ 0 0 612 792 ] /Parent 12 0 R /Resources << +/Font 1 0 R /ProcSet [ /PDF /Text /ImageB /ImageC /ImageI ] +>> /Rotate 0 /Trans << + +>> + /Type /Page +>> +endobj +7 0 obj +<< +/Contents 16 0 R /MediaBox [ 0 0 612 792 ] /Parent 12 0 R /Resources << +/Font 1 0 R /ProcSet [ /PDF /Text /ImageB /ImageC /ImageI ] +>> /Rotate 0 /Trans << + +>> + /Type /Page +>> +endobj +8 0 obj +<< +/Contents 17 0 R /MediaBox [ 0 0 612 792 ] /Parent 12 0 R /Resources << +/Font 1 0 R /ProcSet [ /PDF /Text /ImageB /ImageC /ImageI ] +>> /Rotate 0 /Trans << + +>> + /Type /Page +>> +endobj +9 0 obj +<< +/Contents 18 0 R /MediaBox [ 0 0 612 792 ] /Parent 12 0 R /Resources << +/Font 1 0 R /ProcSet [ /PDF /Text /ImageB /ImageC /ImageI ] +>> /Rotate 0 /Trans << + +>> + /Type /Page +>> +endobj +10 0 obj +<< +/PageMode /UseNone /Pages 12 0 R /Type /Catalog +>> +endobj +11 0 obj +<< +/Author (\(anonymous\)) /CreationDate (D:20260603005358+01'00') /Creator (\(unspecified\)) /Keywords () /ModDate (D:20260603005358+01'00') /Producer (ReportLab PDF Library - \(opensource\)) + /Subject (\(unspecified\)) /Title (\(anonymous\)) /Trapped /False +>> +endobj +12 0 obj +<< +/Count 6 /Kids [ 4 0 R 5 0 R 6 0 R 7 0 R 8 0 R 9 0 R ] /Type /Pages +>> +endobj +13 0 obj +<< +/Filter [ /ASCII85Decode /FlateDecode ] /Length 1007 +>> +stream +GauHKgN)$k&:O:SnDZmeMUK]*(BdhV"2[X,j0;*\%_E*0HXJRc`q%4S#/VM`O[Rh4i5T.Phi,$]aT$05!q:PBuCQU/.p2]@1koN*Tb*K9k5G\ht+Dr\K=+8\NZ"alMaOEo**@OK:8-.O1X3-?Gg`@m3%,ti3'">T-&c=M&Wuu?cDbGDp.gOF0!r6&&CLC$tM?fIR"M37["/*k9@YkpKKSS@Np"1/4#R]`I^(g*1Sc,9L3Qt6N(T@F`A>oBGgL-!r#Y4)R\G`D,iVtFeJUE9u4iuUQ?D%C&SB4kkp.>D5tii>nDKJ"Y07jANhOb$R(_$=U7Stjs)-/KZ(6IBm`6u<3;i.Bh4+MFJ"H:.XWUQX6%LU(sg4Tt$_":5`p.2,gZkpUfdg,Hd(qR1)9\ltF2^8b5,,XbY%VfhD52O7A.c.u]dhTc5t%0<0L5E38`Bq+i;"%J7kcc#@i)@okdN-qiaA"33DdPgPrUs7;j%+47_cnGN@&ug9/dqGOAHA4f.N*&guHspf?E;GIc-Wt>:<(m1AmcS_Zc2VlEI'_S_>@#!MF^m1$$endstream +endobj +14 0 obj +<< +/Filter [ /ASCII85Decode /FlateDecode ] /Length 987 +>> +stream +GatU3gN&c;&:O:SkV;H,eI!J\f?X"LSPIZ'"4Y.F#oAAY?N.Y?8Iu:s)eV;,aNG1Lea>G$!MSc\rU6m7jF'AOI09V1,^TSTi$A+f4sZVQ%2^>uj$40\AI>WbZo9q!NL^CS9P3E9X.:XEMb$a/X"qH9hN6fY,]D!OEOVT,722%RJnVqK4f'[4d3(/hbJWJRs,>>28jk5A:nCl/%FlC#LntsAC5cE$q`QO_@Pk%Lp>7V23")NNrFkcXfuoI=SJJ-`g)_+K\@,SQ@8ORGg(_(B.[u[WXD%Q8@I";9d(]M62K_)Q+-<\kHY5,)o99%BB2bcQVkd:+MZVH=$Z`S@BhH]-XfPANk@qBN.:Jc'?.Kn\p7(aPIQOI(du7U]WDrNNPW%d"8P_j^Y8d%G,'JJLXrV.UD!ah19&9b#*fgR2o730<9)QX)X0"EK6u;mBJs3/fl\d&K$B3bXI5R>F*E,^\%+b?0o8s$o.CuX4uDAT_?G0(p4N3La,8qfi'j?e`e88KrZ8*NDgc:62'icCb],4B]Z"SO&e/BRC*C$2PmbSKAjc\$FLD8?&qHVH3G4p834/77n[%()-p/JDRBapcO2b\,3dW%#U9u!ilSkd%2'rr&g[u".1O]W%0J<@#(ptUD;%0:^.^ls:.#.V6NgR[`2\RUW!k$-*GZIf3DOBiB_Mu"=LVGlO\1Z]0pF(@UaiVQ=dh?h5kfQ_gj>$jHW$8TCt:C.jVU98ne.XoiMtoV;q8XhnlOE;3a]=dGJ8KL&+Khj"&P4Q=P-:nHsDprY"B4+#A3dF,;b&gf(sT\Ts*iH)f!KhfHV;\mMSHej.9.XrBG41_<:b)MG;-Y(:&nCKS0F]r&CGl>EfMc0td\0GHScl3Lr;?OUoo6619Qdr)S8S);]7L5rmWW&(5t4Dil"USDhhY/Y=!=!X:85**$:~>endstream +endobj +15 0 obj +<< +/Filter [ /ASCII85Decode /FlateDecode ] /Length 976 +>> +stream +GatU3gN&c;'R]XVkgAc"70atdYFXp#3h7VV#7/=-%N$"G?N.YO&iX#cVsS_l_5i8)eum:1-)(7`mXJ7LnhqY0hGZal,We=q^e"$M]Lu;_=4AoIN$QOjOF.@i*Sg!_rLL=-'?kk;m;M&R`$D>'HCL7]$A,-S`H1I/1T]4\B3=gQ9:I4j9f]+af*Z7*GH5H+;j73Z-T#DAcgpOW(/'b?XF9+X+'-Dqi"+1U7?6=ZRJNBSJ6ojDTa+]o4RF^26C+i0s(9G2o;$@RNbs\(f75],L(SVpI0M!'7JVhZ"$t/%kb1+)F\p54/@jpVJF2+)7^F=gVcp8(h*m%F#'p9c%9Ep9?b7g^F5>(>I3t^d_"f'aue%^E\X4kN:g\P[;sC=g3h,nD/n9*`94uoM?sBCL4;XilfMPT4]&FC.U%DI(c6uf3*XPJP2MRdq@Ag3uYF().nqLFAE2Nih'(^p\5FnSNa>n+cI?U65A61_N1<9ZXH4Xr)Fgq(6md3H9[MQ3Q^%aq>E?XM:O:RKGjc+QV2MB!.I.17QpSt]X[C*6ap.`#'o;8F(ebZ.VMM6k`BkZB\I,2MNoX]J"k]Qd";,@($aWX&29q"6k;>Let)@]QHm:i/lf[DB_&U@ROq-0MLQp;@j+]?$"E)brXVun52AVQo$"[a7@!?!LT*>$&:5X_.R;9)'$l-S.>WjLQIVTo.JU%(H\,-*pl_kP9~>endstream +endobj +16 0 obj +<< +/Filter [ /ASCII85Decode /FlateDecode ] /Length 1024 +>> +stream +GauHKflEQ9'Rf^WgrHc4<#\/Qm7c-rFIIq+TFSD%Z#L'6o(Nl&^_!k0ds=.%-s&pt#i0QBKE91*<^/MJJCcfoHA_e;aL?\&_BAj]Dt;HW$6(8Ql\%O7-5%uT80ZIL#\l)(+[ABOTNKmq!2F,S:.F93Z6QIu[nf\Q[qCE`Y-=.8k$67iCR!7=El+50a@f&b^6'Yg'mqlcRQ-5j:%[K@'5!A_m0L)&VlW1T50^X)"3Ma3,U1P8Di$>uPZ)Zj_igBc0lm]WE0#e>*M53(Yh6\9jHYh6B+3bM\b[Q9j'O_9:!:W0$Ijm,#(bKn_D?UYMGJ7LLqC0\[$l<8U*O5^GnN;!N6YB'?d[d8DAPt^>+E:r0<2$'VL/TtuNg;C@5iM#%KdS-GU->0]`/20eLXVW#qnjsU.P:Xj&=enVE[mpqDCn2!qDRuQhI/#cZ-#m6`U9'SI4"9\"r:4p.\W>IK_$#C,EkU#uO1g6*HoUK':XiJjtQfI'Sdc)B3J1eqgGJ+l?"$//&C$^f'TA!2:K!0pd;^Zif)c[:f!C/8hNKc7+'+1WbOFKr9T7c_G$s*O634O2Zs:"iN>j,pWU+:kEQ>5YbF=$27*59iJEA//7(KsL\pUG.G[+JB-r8t;Sih+(W.A'8Q*LZK*I9WjCLjEt>lT%2RFJZg0Ps:5c,HdK#tSNI:#Rp!]O$Ic=T%K0[._E_Fcjn>OcM,iSAf9iF7,,+Qe:@o'Y4o_:54t*l+TAgendstream +endobj +17 0 obj +<< +/Filter [ /ASCII85Decode /FlateDecode ] /Length 1075 +>> +stream +GauHKgN&c;&:O:SkV;H,6&QZ[5)3-pL]0bpl':D91fR,o"F49.1/bfmFt32ljt6eO[j7!K]\nZ:^QQASu[l@3no(6*HL'@3I!Qfi7%,.N\.>A91O)D\n5*o@JWSOI@ng5p/"_I`g*Y$j0+5PWaso'25ltcglFd,QIaO'dPO0Yr\h5/,=m\PiP:P+Gg0&N;_5iLN$]BD.#VL+h,tSe\fL&,;+2fg&%\pAJiV#eVhq@ro%%/`[(&8n29YiYd>bDF^Bg[AVf/n>[FE(au"ZAe&>%Q[7/n3-QhsPXuPbp#<$d"UR8u-nBnr!Q$Wjd1;9WN/,KG8])Ca5rrDrI'Ue'*I0C%pn+f`Hd"4sB_&2*)S$pX=F/6"DC"TU`bMtd/m7u]?]=J5L[SNBuE:6Ri,NmW7.*qUp)97Tt+a84b\os`Ti8/uRDlhrl*Xrn(a0\,%4%*k^rL\k3%6ako`&A3MoN^CXV77D!Qg%=j9Ge+0B@dZDHELb2fDU(,iiXXRl^*U,U#Fld!R72UIpKGWu#-DXi653a?U$fqCs18gt(5/#U%!u,gkMK;>OT_/='S:kn,iD@u]<8-rlTRu'+\7L$U7-#!5-5\WQn(n:q@-p24u!~>endstream +endobj +18 0 obj +<< +/Filter [ /ASCII85Decode /FlateDecode ] /Length 590 +>> +stream +GasbXgMWc?&;KY!MRgt!"n`^Bf\9$c*`Q-Sb6tfd8P(q-dT1eng=QV-!OG*\kiWo2AGBdKLqI^%h,XNB)-gDk+GFVBa?g*a%:!J\97Vc8-jH8>8P>0S@Wr)o*"Hu<6qPp`EF:M3'l4@S\c2]`&[J%dO@5p$<(([H6:SlB;FS8L`&pO%6Y\V"/O7F]E%L@d.%>L@"4bK_XXs&";n_.W^Nr847HH,X@$.pC"4bYC?9RH,dN?h8a+8/!?Q=D\F41Nc@'B`';4XgMeh!;U23C+>AbMQD-o--lJ/":m#(Yt,X5(`1KL,GTMBLlr6=-1mi#k.-Ou\\T4$Pun2dEU(\?$&;1@T4dm^t!KuOD-U@N_AMr"&uVpsGm,+8I7B*f!%9.o4cC1a[CZ12(hd>1*0bU`k2-MXo1[Gor#kmXGIM'#R49X#NSAOpdf'0ilUH4:M(^Snc3;m/+>endstream +endobj +xref +0 19 +0000000000 65535 f +0000000061 00000 n +0000000102 00000 n +0000000209 00000 n +0000000321 00000 n +0000000516 00000 n +0000000711 00000 n +0000000906 00000 n +0000001101 00000 n +0000001296 00000 n +0000001491 00000 n +0000001561 00000 n +0000001842 00000 n +0000001932 00000 n +0000003031 00000 n +0000004109 00000 n +0000005176 00000 n +0000006292 00000 n +0000007459 00000 n +trailer +<< +/ID +[] +% ReportLab generated PDF document -- digest (opensource) + +/Info 11 0 R +/Root 10 0 R +/Size 19 +>> +startxref +8140 +%%EOF diff --git a/app/common/src/test/resources/pdf-ingestion-fixtures/multi-column-test_lorem.md b/app/common/src/test/resources/pdf-ingestion-fixtures/multi-column-test_lorem.md new file mode 100644 index 0000000000..5c35de111f --- /dev/null +++ b/app/common/src/test/resources/pdf-ingestion-fixtures/multi-column-test_lorem.md @@ -0,0 +1,25 @@ +# Lorem Ipsum in Two Columns + +## 1. Origins + +Lorem ipsum dolor sit amet consectetur adipiscing elit. Sed do eiusmod tempor incididunt ut labore et dolore magna aliqua. + +## 2. Structure + +Ut enim ad minim veniam quis nostrud exercitation ullamco laboris nisi ut aliquip ex ea commodo consequat. + +## 3. Usage + +Duis aute irure dolor in reprehenderit in voluptate velit esse cillum dolore eu fugiat nulla pariatur. + +## 4. Variations + +Excepteur sint occaecat cupidatat non proident sunt in culpa qui officia deserunt mollit anim id est laborum. + +## 5. Typography + +Curabitur pretium tincidunt lacus. Nulla gravida orci a odio. Nullam various turpis et commodo pharetra est. + +## 6. Conclusion + +Nunc nonummy metus. Vestibulum volutpat pretium libero. Cras id dui. Aenean ut eros et nisl sagittis vestibulum. diff --git a/app/common/src/test/resources/pdf-ingestion-fixtures/multi-column-test_lorem.pdf b/app/common/src/test/resources/pdf-ingestion-fixtures/multi-column-test_lorem.pdf new file mode 100644 index 0000000000..36dc3a1a65 --- /dev/null +++ b/app/common/src/test/resources/pdf-ingestion-fixtures/multi-column-test_lorem.pdf @@ -0,0 +1,74 @@ +%PDF-1.4 +%“Œ‹ž ReportLab Generated PDF document (opensource) +1 0 obj +<< +/F1 2 0 R /F2 3 0 R +>> +endobj +2 0 obj +<< +/BaseFont /Helvetica /Encoding /WinAnsiEncoding /Name /F1 /Subtype /Type1 /Type /Font +>> +endobj +3 0 obj +<< +/BaseFont /Helvetica-Bold /Encoding /WinAnsiEncoding /Name /F2 /Subtype /Type1 /Type /Font +>> +endobj +4 0 obj +<< +/Contents 8 0 R /MediaBox [ 0 0 612 792 ] /Parent 7 0 R /Resources << +/Font 1 0 R /ProcSet [ /PDF /Text /ImageB /ImageC /ImageI ] +>> /Rotate 0 /Trans << + +>> + /Type /Page +>> +endobj +5 0 obj +<< +/PageMode /UseNone /Pages 7 0 R /Type /Catalog +>> +endobj +6 0 obj +<< +/Author (\(anonymous\)) /CreationDate (D:20260603021636+01'00') /Creator (\(unspecified\)) /Keywords () /ModDate (D:20260603021636+01'00') /Producer (ReportLab PDF Library - \(opensource\)) + /Subject (\(unspecified\)) /Title (\(anonymous\)) /Trapped /False +>> +endobj +7 0 obj +<< +/Count 1 /Kids [ 4 0 R ] /Type /Pages +>> +endobj +8 0 obj +<< +/Filter [ /ASCII85Decode /FlateDecode ] /Length 861 +>> +stream +Gat=i>Ak00&;B$5/*8Q/Z)fmNp(^2^OK'MC[dT6#Oq.&b^4c4;1No6>+D%TC.n+fice+l9R0[L%'(osQ?dQP@AUa.A:T=#%ZIIl18hB7ZfI(hVU+?`]UK%;aabAH:q>9NUGQq^.=u])Wt-ETI))C'a[FE@elj!RSrf2Q'F>URAO.C,!DneTPqrj#4e2kb9%1"4qfZ)"#&^0j9nHQ?9nF!j7mVPP5\*Uq'_jMVS]9%`kQB\8*AF_bpr/hGj;HCUOSQU-%5:6S79Ud\b!*tPbr_'pCr$Ea#(FYP31NFhSX.-("1M:$cgH#hX8L(2]R3Q>'BYHCS%pI!;=WdJp,'ii[`QPZ_9mcd\baZ2U(_;c\-p,8EoIEpQ*lstL>]LE;C#\dLnT2R:)BM-fTc['3_He[U,k'!Bo".uERd>SkhRj^J+koSIrZ_dEf_5L'/1h.`+DTK(R:P,WH)h5\se=SZ"L/5b8b..,e/E\o+4YQ+*im^C>AERG/TieEK\)#>U@HXnJ,H0A9-MqhhkDp8%.6Lr,OrK*lih;B<-opZ8%EU?,$r^jmCDAQ`-0/-8/`[p]7Fm%0f:E&S*FV)DX2>#q\bRqA=^_`43#8EA$u%8r6F5`rc9>K@q>E4q^*~>endstream +endobj +xref +0 9 +0000000000 65535 f +0000000061 00000 n +0000000102 00000 n +0000000209 00000 n +0000000321 00000 n +0000000514 00000 n +0000000582 00000 n +0000000862 00000 n +0000000921 00000 n +trailer +<< +/ID +[<21a9fbd0a0991a91b6e6e2db0856056e><21a9fbd0a0991a91b6e6e2db0856056e>] +% ReportLab generated PDF document -- digest (opensource) + +/Info 6 0 R +/Root 5 0 R +/Size 9 +>> +startxref +1872 +%%EOF diff --git a/app/common/src/test/resources/pdf-ingestion-fixtures/wrapped-cell-test_expense-report.md b/app/common/src/test/resources/pdf-ingestion-fixtures/wrapped-cell-test_expense-report.md new file mode 100644 index 0000000000..8008a3b970 --- /dev/null +++ b/app/common/src/test/resources/pdf-ingestion-fixtures/wrapped-cell-test_expense-report.md @@ -0,0 +1,62 @@ +# Employee Expense Report + +Reimbursement Request + +EMP-1047 + +**Report Header** + +| Employee Name | Michael Tran | +|---|---| +| Employee ID | EMP-1047 | +| Department | Client Services | +| Report Date | January 20th, 2026 | +| Reporting Period | January 5th–16th, 2026 | +| Manager Approver | Laura Simmons | + +**Company Information** + +| Company | Summit Consulting Partners | +|---|---| +| Company Address | 88 Riverside Plaza, Suite 1400, New York, NY 10069 | +| Accounting Department Email | expenses@example.com | + +**Trip Purpose** + +The trip was undertaken for client onsite meetings with Atlantic Energy Solutions in Boston, MA. + +**Expense Details** + +| Description | Amount | Date | Category | +|---|---|---|---| +| Flight (NYC to Boston roundtrip) | $325.40 | January 5th, 2026 | Airline ticket | +| Hotel (3 nights at Harborview Hotel) | $822.75 | January 5th–8th, 2026 | Lodging | +| Taxi from airport to hotel | $48.00 | January 5th, 2026 | Ground transportation | +| Client dinner (3 attendees) | $186.20 | January 6th, 2026 | Meals | +| Parking at JFK Airport | $72.00 | January 5th–8th, 2026 | Parking | +| Breakfast (per diem not used) | $18.50 | January 7th, 2026 | Meals | + +| Description | Amount | Date | Category | +|---|---|---|---| +| Uber to client office | $22.10 | January 7th, 2026 | Ground transportation | +| Printing + presentation materials | $46.90 | January 8th, 2026 | Materials | +| Lunch with client | $39.75 | January 8th, 2026 | Meals | +| Office supplies (notebooks, pens) | $27.60 | January 10th, 2026 | Supplies | +| Mileage reimbursement (client visit in NJ, 42 miles @ $0.67/mile) | $28.14 | January 14th, 2026 | Mileage | +| Team lunch meeting (internal) | $64.30 | January 15th, 2026 | Meals | + +Total Expenses $1,701.64 + +Reimbursement Method + +Reimbursement method Direct deposit + +Notes + +All receipts are attached. Expenses are business-related and comply with company travel policy. + +**Approval** + +Michael Tran, Employee + +Laura Simmons, Manager diff --git a/app/common/src/test/resources/pdf-ingestion-fixtures/wrapped-cell-test_expense-report.pdf b/app/common/src/test/resources/pdf-ingestion-fixtures/wrapped-cell-test_expense-report.pdf new file mode 100644 index 0000000000..95a0b2e07a Binary files /dev/null and b/app/common/src/test/resources/pdf-ingestion-fixtures/wrapped-cell-test_expense-report.pdf differ diff --git a/app/core/build.gradle b/app/core/build.gradle index 54bdd7bd76..e505ec9838 100644 --- a/app/core/build.gradle +++ b/app/core/build.gradle @@ -207,13 +207,29 @@ def resourcesStaticDir = file('src/main/resources/static') def generatedFrontendPaths = [ 'assets', 'index.html', + 'index.html.gz', + 'index.html.br', + 'sw.js', + 'sw.js.gz', + 'sw.js.br', + 'manifest.json.gz', + 'manifest.json.br', + 'site.webmanifest.gz', + 'site.webmanifest.br', + 'browserconfig.xml.gz', + 'browserconfig.xml.br', + 'manifest-classic.json', + 'manifest-classic.json.gz', + 'manifest-classic.json.br', 'locales', 'Login', 'classic-logo', 'modern-logo', 'og_images', 'samples', - 'manifest-classic.json' + 'pdfium', + 'vendor', + 'pdfjs' ] tasks.register('npmInstall', Exec) { @@ -314,6 +330,12 @@ tasks.register('cleanFrontendAssets', Delete) { group = 'frontend' description = 'Remove previously generated frontend assets from static resources' delete generatedFrontendPaths.collect { new File(resourcesStaticDir, it) } + // Prerendered per-route SPA pages (e.g. compress.html) carry per-tool OG tags and are + // copied from the frontend build. Remove stale ones so renamed/removed tools don't linger. + // api-landing.html is a real backend source file, not a generated artifact. + delete fileTree(dir: resourcesStaticDir, includes: ['*.html'], excludes: ['api-landing.html']) + // Nested prerendered route pages (e.g. settings/people.html) + delete new File(resourcesStaticDir, 'settings') } tasks.register('copyApiLandingPage', Copy) { diff --git a/app/core/src/main/java/stirling/software/SPDF/SPDFApplication.java b/app/core/src/main/java/stirling/software/SPDF/SPDFApplication.java index 177d2443a6..3cc3100755 100644 --- a/app/core/src/main/java/stirling/software/SPDF/SPDFApplication.java +++ b/app/core/src/main/java/stirling/software/SPDF/SPDFApplication.java @@ -4,7 +4,6 @@ import java.io.IOException; import java.net.URISyntaxException; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import java.util.Collections; import java.util.HashMap; import java.util.Map; @@ -75,7 +74,7 @@ public class SPDFApplication { Map propertyFiles = new HashMap<>(); // External config files - Path settingsPath = Paths.get(InstallationPathConfig.getSettingsPath()); + Path settingsPath = Path.of(InstallationPathConfig.getSettingsPath()); log.info("Settings file: {}", settingsPath.toString()); if (Files.exists(settingsPath)) { propertyFiles.put( @@ -84,7 +83,7 @@ public class SPDFApplication { log.warn("External configuration file '{}' does not exist.", settingsPath.toString()); } - Path customSettingsPath = Paths.get(InstallationPathConfig.getCustomSettingsPath()); + Path customSettingsPath = Path.of(InstallationPathConfig.getCustomSettingsPath()); log.info("Custom settings file: {}", customSettingsPath.toString()); if (Files.exists(customSettingsPath)) { String existingLocation = diff --git a/app/core/src/main/java/stirling/software/SPDF/config/ExternalAppDepConfig.java b/app/core/src/main/java/stirling/software/SPDF/config/ExternalAppDepConfig.java index c08e53e1ab..8755dfe2ef 100644 --- a/app/core/src/main/java/stirling/software/SPDF/config/ExternalAppDepConfig.java +++ b/app/core/src/main/java/stirling/software/SPDF/config/ExternalAppDepConfig.java @@ -95,7 +95,7 @@ public class ExternalAppDepConfig { checkDependencyAndDisableGroup(cmd); return null; }) - .collect(Collectors.toList()); + .toList(); invokeAllWithTimeout(tasks, DEFAULT_TIMEOUT.plusSeconds(3)); // Python / OpenCV special handling diff --git a/app/core/src/main/java/stirling/software/SPDF/config/LocaleConfiguration.java b/app/core/src/main/java/stirling/software/SPDF/config/LocaleConfiguration.java index 97fbb4d219..7d57c1efe8 100644 --- a/app/core/src/main/java/stirling/software/SPDF/config/LocaleConfiguration.java +++ b/app/core/src/main/java/stirling/software/SPDF/config/LocaleConfiguration.java @@ -37,8 +37,8 @@ public class LocaleConfiguration implements WebMvcConfigurer { public LocaleResolver localeResolver() { SessionLocaleResolver slr = new SessionLocaleResolver(); String appLocaleEnv = applicationProperties.getSystem().getDefaultLocale(); - Locale defaultLocale = // Fallback to UK locale if environment variable is not set - Locale.UK; + Locale defaultLocale = // Fallback to US locale if environment variable is not set + Locale.US; if (appLocaleEnv != null && !appLocaleEnv.isEmpty()) { Locale tempLocale = Locale.forLanguageTag(appLocaleEnv); String tempLanguageTag = tempLocale.toLanguageTag(); @@ -51,7 +51,7 @@ public class LocaleConfiguration implements WebMvcConfigurer { defaultLocale = tempLocale; } else { System.err.println( - "Invalid SYSTEM_DEFAULTLOCALE environment variable value. Falling back to default en-GB."); + "Invalid SYSTEM_DEFAULTLOCALE environment variable value. Falling back to default en-US."); } } } diff --git a/app/core/src/main/java/stirling/software/SPDF/config/OpenApiConfig.java b/app/core/src/main/java/stirling/software/SPDF/config/OpenApiConfig.java index 205a0d5734..7c7fe115d0 100644 --- a/app/core/src/main/java/stirling/software/SPDF/config/OpenApiConfig.java +++ b/app/core/src/main/java/stirling/software/SPDF/config/OpenApiConfig.java @@ -18,6 +18,7 @@ import io.swagger.v3.oas.models.media.StringSchema; import io.swagger.v3.oas.models.security.SecurityRequirement; import io.swagger.v3.oas.models.security.SecurityScheme; import io.swagger.v3.oas.models.servers.Server; +import io.swagger.v3.oas.models.tags.Tag; import lombok.RequiredArgsConstructor; @@ -60,6 +61,15 @@ public class OpenApiConfig { OpenAPI openAPI = new OpenAPI().info(info).openapi("3.0.3"); + // Register a single global "AI" tag so every AI endpoint groups under it in the docs. + // The AI controllers are currently @Hidden, so they don't emit this tag themselves yet; + // defining it here keeps the grouping ready for when those endpoints are unhidden. + openAPI.addTagsItem( + new Tag() + .name("AI") + .description( + "AI-powered document creation, editing, and assistant endpoints.")); + // Add server configuration from environment variable String swaggerServerUrl = System.getenv("SWAGGER_SERVER_URL"); Server server; diff --git a/app/core/src/main/java/stirling/software/SPDF/config/WebMvcConfig.java b/app/core/src/main/java/stirling/software/SPDF/config/WebMvcConfig.java index 48fecf001e..367c875744 100644 --- a/app/core/src/main/java/stirling/software/SPDF/config/WebMvcConfig.java +++ b/app/core/src/main/java/stirling/software/SPDF/config/WebMvcConfig.java @@ -13,6 +13,7 @@ import org.springframework.web.servlet.config.annotation.CorsRegistry; import org.springframework.web.servlet.config.annotation.InterceptorRegistry; import org.springframework.web.servlet.config.annotation.ResourceHandlerRegistry; import org.springframework.web.servlet.config.annotation.WebMvcConfigurer; +import org.springframework.web.servlet.resource.EncodedResourceResolver; import lombok.RequiredArgsConstructor; @@ -49,14 +50,16 @@ public class WebMvcConfig implements WebMvcConfigurer { "/sw.js", "/manifest.json", "/site.webmanifest", "/browserconfig.xml") .addResourceLocations(staticPath, "classpath:/static/") .setCacheControl(CacheControl.noStore()) - .resourceChain(true); + .resourceChain(true) + .addResolver(new EncodedResourceResolver()); // 2. Vite fingerprinted assets (immutable) // These already have content hashes in filenames (e.g. index-ChAS4tCC.js) registry.addResourceHandler("/assets/**") .addResourceLocations(staticPath + "assets/", "classpath:/static/assets/") .setCacheControl(IMMUTABLE_ONE_YEAR) - .resourceChain(true); + .resourceChain(true) + .addResolver(new EncodedResourceResolver()); // 3. Media and fonts (immutable) registry.addResourceHandler("/images/**", "/fonts/**") @@ -66,7 +69,8 @@ public class WebMvcConfig implements WebMvcConfigurer { staticPath + "fonts/", "classpath:/static/fonts/") .setCacheControl(IMMUTABLE_ONE_YEAR) - .resourceChain(true); + .resourceChain(true) + .addResolver(new EncodedResourceResolver()); // 4. Branding and stable non-fingerprinted assets (1 day + SWR) // Use stale-while-revalidate to improve perceived performance. @@ -114,19 +118,27 @@ public class WebMvcConfig implements WebMvcConfigurer { staticPath + "og_images/", "classpath:/static/og_images/", staticPath + "Login/", - "classpath:/static/Login/") + "classpath:/static/Login/", + staticPath + "icons/", + "classpath:/static/icons/", + staticPath + "modern-logo/", + "classpath:/static/modern-logo/", + staticPath + "classic-logo/", + "classpath:/static/classic-logo/") .setCacheControl( CacheControl.maxAge(Duration.ofDays(1)) .cachePublic() .staleWhileRevalidate(Duration.ofDays(7))) - .resourceChain(true); + .resourceChain(true) + .addResolver(new EncodedResourceResolver()); // 5. Catch-all (SPA fallback) // Must check with server to ensure index.html is always fresh. registry.addResourceHandler("/**") .addResourceLocations(staticPath, "classpath:/static/") .setCacheControl(NO_CACHE) - .resourceChain(true); + .resourceChain(true) + .addResolver(new EncodedResourceResolver()); } @Override diff --git a/app/core/src/main/java/stirling/software/SPDF/controller/api/AdditionalLanguageJsController.java b/app/core/src/main/java/stirling/software/SPDF/controller/api/AdditionalLanguageJsController.java index e5d3ba8840..4b53ea4eae 100644 --- a/app/core/src/main/java/stirling/software/SPDF/controller/api/AdditionalLanguageJsController.java +++ b/app/core/src/main/java/stirling/software/SPDF/controller/api/AdditionalLanguageJsController.java @@ -48,7 +48,7 @@ public class AdditionalLanguageJsController { } } // Fallback - return "en_GB"; + return "en_US"; } """); writer.flush(); diff --git a/app/core/src/main/java/stirling/software/SPDF/controller/api/PdfOverlayController.java b/app/core/src/main/java/stirling/software/SPDF/controller/api/PdfOverlayController.java index 50b42d89dc..4aab03ff55 100644 --- a/app/core/src/main/java/stirling/software/SPDF/controller/api/PdfOverlayController.java +++ b/app/core/src/main/java/stirling/software/SPDF/controller/api/PdfOverlayController.java @@ -10,6 +10,7 @@ import java.util.Map; import org.apache.pdfbox.Loader; import org.apache.pdfbox.multipdf.Overlay; +import org.apache.pdfbox.pdfwriter.compress.CompressParameters; import org.apache.pdfbox.pdmodel.PDDocument; import org.springframework.core.io.Resource; import org.springframework.http.MediaType; @@ -157,7 +158,10 @@ public class PdfOverlayController { PDDocument singlePageDocument = new PDDocument()) { singlePageDocument.addPage(overlayPdf.getPage(pageCountInCurrentOverlay)); File tempFile = Files.createTempFile("overlay-page-", ".pdf").toFile(); - singlePageDocument.save(tempFile); + // NO_COMPRESSION: this single-page doc holds a page copied from overlayPdf. + // PDFBox 3.0.7's compressed writer (PDFBOX-6203) drops shared resources imported + // across documents, corrupting overlay fonts. Revert once on 3.0.8. + singlePageDocument.save(tempFile, CompressParameters.NO_COMPRESSION); overlayGuide.put(basePageIndex, tempFile.getAbsolutePath()); tempFiles.add(tempFile); // Keep track of the temporary file for cleanup diff --git a/app/core/src/main/java/stirling/software/SPDF/controller/api/RearrangePagesPDFController.java b/app/core/src/main/java/stirling/software/SPDF/controller/api/RearrangePagesPDFController.java index 1e87a15d57..6dd7aacd79 100644 --- a/app/core/src/main/java/stirling/software/SPDF/controller/api/RearrangePagesPDFController.java +++ b/app/core/src/main/java/stirling/software/SPDF/controller/api/RearrangePagesPDFController.java @@ -6,11 +6,9 @@ import java.util.Collections; import java.util.List; import java.util.Locale; -import org.apache.pdfbox.cos.COSName; import org.apache.pdfbox.pdmodel.PDDocument; -import org.apache.pdfbox.pdmodel.PDDocumentCatalog; import org.apache.pdfbox.pdmodel.PDPage; -import org.apache.pdfbox.pdmodel.interactive.form.PDAcroForm; +import org.apache.pdfbox.pdmodel.PDPageTree; import org.springframework.core.io.Resource; import org.springframework.http.MediaType; import org.springframework.http.ResponseEntity; @@ -262,38 +260,31 @@ public class RearrangePagesPDFController { } log.info("newPageOrder = {}", newPageOrder); log.info("totalPages = {}", totalPages); - // Create a new list to hold the pages in the new order - List newPages = new ArrayList<>(); - for (int i = 0; i < newPageOrder.size(); i++) { - newPages.add(document.getPage(newPageOrder.get(i))); + + // Snapshot the desired pages before mutating the source document's page tree. + List newPages = new ArrayList<>(newPageOrder.size()); + for (Integer idx : newPageOrder) { + newPages.add(document.getPage(idx)); } - // Create a new document based on the original one - try (PDDocument rearrangedDocument = - pdfDocumentFactory.createNewDocumentBasedOnOldDocument(document)) { - - // Add the pages in the new order - for (PDPage page : newPages) { - rearrangedDocument.addPage(page); - } - - PDDocumentCatalog sourceCatalog = document.getDocumentCatalog(); - if (sourceCatalog != null) { - PDAcroForm sourceForm = sourceCatalog.getAcroForm(null); - if (sourceForm != null) { - rearrangedDocument - .getDocumentCatalog() - .getCOSObject() - .setItem(COSName.ACRO_FORM, sourceForm.getCOSObject()); - } - } - - return WebResponseUtils.pdfDocToWebResponse( - rearrangedDocument, - GeneralUtils.generateFilename( - pdfFile.getOriginalFilename(), "_rearranged.pdf"), - tempFileManager); + // Rearrange in-place on the source document rather than copying pages into a + // freshly-created PDDocument. Copying pages across documents triggers a PDFBox + // 3.0.7 compressed-save regression (PDFBOX-6203, fixed for 3.0.8) where shared + // resource objects (fonts, etc.) imported from the source can be silently + // dropped from the output, producing pages with "font not found" errors. + PDPageTree pages = document.getPages(); + for (int i = totalPages - 1; i >= 0; i--) { + pages.remove(i); } + for (PDPage page : newPages) { + pages.add(page); + } + + return WebResponseUtils.pdfDocToWebResponse( + document, + GeneralUtils.generateFilename( + pdfFile.getOriginalFilename(), "_rearranged.pdf"), + tempFileManager); } } catch (IOException e) { ExceptionUtils.logException("document rearrangement", e); diff --git a/app/core/src/main/java/stirling/software/SPDF/controller/api/UIDataController.java b/app/core/src/main/java/stirling/software/SPDF/controller/api/UIDataController.java index c1afe8c406..de391c7c32 100644 --- a/app/core/src/main/java/stirling/software/SPDF/controller/api/UIDataController.java +++ b/app/core/src/main/java/stirling/software/SPDF/controller/api/UIDataController.java @@ -5,7 +5,6 @@ import java.io.InputStream; import java.nio.charset.StandardCharsets; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import java.util.*; import java.util.stream.Stream; @@ -115,7 +114,7 @@ public class UIDataController { if (new java.io.File(runtimePathConfig.getPipelineDefaultWebUiConfigs()).exists()) { try (Stream paths = - Files.walk(Paths.get(runtimePathConfig.getPipelineDefaultWebUiConfigs()))) { + Files.walk(Path.of(runtimePathConfig.getPipelineDefaultWebUiConfigs()))) { List jsonFiles = paths.filter(Files::isRegularFile) .filter(p -> p.toString().endsWith(".json")) diff --git a/app/core/src/main/java/stirling/software/SPDF/controller/api/converters/ConvertOfficeController.java b/app/core/src/main/java/stirling/software/SPDF/controller/api/converters/ConvertOfficeController.java index 8bf53c79e1..4fc6669bc7 100644 --- a/app/core/src/main/java/stirling/software/SPDF/controller/api/converters/ConvertOfficeController.java +++ b/app/core/src/main/java/stirling/software/SPDF/controller/api/converters/ConvertOfficeController.java @@ -35,6 +35,7 @@ import stirling.software.common.service.CustomPDFDocumentFactory; import stirling.software.common.util.CustomHtmlSanitizer; import stirling.software.common.util.ExceptionUtils; import stirling.software.common.util.GeneralUtils; +import stirling.software.common.util.OfficeDocumentSanitizer; import stirling.software.common.util.ProcessExecutor; import stirling.software.common.util.ProcessExecutor.ProcessExecutorResult; import stirling.software.common.util.RegexPatternUtils; @@ -50,6 +51,7 @@ public class ConvertOfficeController { private final CustomPDFDocumentFactory pdfDocumentFactory; private final RuntimePathConfig runtimePathConfig; private final CustomHtmlSanitizer customHtmlSanitizer; + private final OfficeDocumentSanitizer officeDocumentSanitizer; private final EndpointConfiguration endpointConfiguration; private final TempFileManager tempFileManager; @@ -83,14 +85,16 @@ public class ConvertOfficeController { Path inputPath = workDir.resolve(baseName + "." + extensionLower); Path outputPath = workDir.resolve(baseName + ".pdf"); - // Check if the file is HTML and apply sanitization if needed + // Sanitize input before LibreOffice sees it so embedded URLs can't trigger SSRF. if ("html".equals(extensionLower) || "htm".equals(extensionLower)) { - // Read and sanitize HTML content String htmlContent = new String(inputFile.getBytes(), StandardCharsets.UTF_8); String sanitizedHtml = customHtmlSanitizer.sanitize(htmlContent); Files.writeString(inputPath, sanitizedHtml, StandardCharsets.UTF_8); + } else if (officeDocumentSanitizer.isSanitizableExtension(extensionLower)) { + byte[] sanitized = + officeDocumentSanitizer.sanitize(inputFile.getBytes(), extensionLower); + Files.write(inputPath, sanitized); } else { - // copy file content Files.copy(inputFile.getInputStream(), inputPath, StandardCopyOption.REPLACE_EXISTING); } diff --git a/app/core/src/main/java/stirling/software/SPDF/controller/api/misc/AutoSplitPdfController.java b/app/core/src/main/java/stirling/software/SPDF/controller/api/misc/AutoSplitPdfController.java index d9714f0618..ad5776b3cd 100644 --- a/app/core/src/main/java/stirling/software/SPDF/controller/api/misc/AutoSplitPdfController.java +++ b/app/core/src/main/java/stirling/software/SPDF/controller/api/misc/AutoSplitPdfController.java @@ -14,6 +14,7 @@ import java.util.zip.ZipEntry; import java.util.zip.ZipOutputStream; import org.apache.pdfbox.cos.COSName; +import org.apache.pdfbox.pdfwriter.compress.CompressParameters; import org.apache.pdfbox.pdmodel.PDDocument; import org.apache.pdfbox.pdmodel.PDPage; import org.apache.pdfbox.pdmodel.graphics.image.PDImageXObject; @@ -357,7 +358,10 @@ public class AutoSplitPdfController { for (int i = 0; i < splitDocuments.size(); i++) { String fileName = filename + "_" + (i + 1) + ".pdf"; zipOut.putNextEntry(new ZipEntry(fileName)); - splitDocuments.get(i).save(zipOut); + // NO_COMPRESSION: split docs are built by addPage()-ing pages copied from the + // source document. PDFBox 3.0.7's compressed writer (PDFBOX-6203) drops shared + // resources imported across documents, corrupting fonts. Revert once on 3.0.8. + splitDocuments.get(i).save(zipOut, CompressParameters.NO_COMPRESSION); zipOut.closeEntry(); } } diff --git a/app/core/src/main/java/stirling/software/SPDF/controller/api/misc/OCRController.java b/app/core/src/main/java/stirling/software/SPDF/controller/api/misc/OCRController.java index 64ed5b79ea..4660d5f128 100644 --- a/app/core/src/main/java/stirling/software/SPDF/controller/api/misc/OCRController.java +++ b/app/core/src/main/java/stirling/software/SPDF/controller/api/misc/OCRController.java @@ -17,6 +17,7 @@ import javax.imageio.ImageIO; import org.apache.pdfbox.io.IOUtils; import org.apache.pdfbox.multipdf.PDFMergerUtility; +import org.apache.pdfbox.pdfwriter.compress.CompressParameters; import org.apache.pdfbox.pdmodel.PDDocument; import org.apache.pdfbox.pdmodel.PDPage; import org.apache.pdfbox.rendering.PDFRenderer; @@ -427,7 +428,10 @@ public class OCRController { // Save original page without OCR as fallback try (PDDocument pageDoc = new PDDocument()) { pageDoc.addPage(page); - pageDoc.save(pageOutputPath); + // NO_COMPRESSION: page is copied from another document; + // PDFBox 3.0.7 compressed writer (PDFBOX-6203) drops shared + // resources, corrupting fonts. Revert once on 3.0.8. + pageDoc.save(pageOutputPath, CompressParameters.NO_COMPRESSION); } } @@ -437,7 +441,10 @@ public class OCRController { // Save original page without OCR try (PDDocument pageDoc = new PDDocument()) { pageDoc.addPage(page); - pageDoc.save(pageOutputPath); + // NO_COMPRESSION: page is copied from another document; PDFBox 3.0.7 + // compressed writer (PDFBOX-6203) drops shared resources, corrupting + // fonts on retained text pages. Revert once on 3.0.8. + pageDoc.save(pageOutputPath, CompressParameters.NO_COMPRESSION); merger.addSource(pageOutputPath); } } diff --git a/app/core/src/main/java/stirling/software/SPDF/controller/api/misc/OverlayImageController.java b/app/core/src/main/java/stirling/software/SPDF/controller/api/misc/OverlayImageController.java index c55597b47d..557c77c76e 100644 --- a/app/core/src/main/java/stirling/software/SPDF/controller/api/misc/OverlayImageController.java +++ b/app/core/src/main/java/stirling/software/SPDF/controller/api/misc/OverlayImageController.java @@ -25,6 +25,7 @@ import stirling.software.common.annotations.api.MiscApi; import stirling.software.common.enumeration.ResourceWeight; import stirling.software.common.service.CustomPDFDocumentFactory; import stirling.software.common.util.GeneralUtils; +import stirling.software.common.util.SvgSanitizer; import stirling.software.common.util.TempFile; import stirling.software.common.util.TempFileManager; import stirling.software.common.util.WebResponseUtils; @@ -36,6 +37,7 @@ public class OverlayImageController { private final CustomPDFDocumentFactory pdfDocumentFactory; private final TempFileManager tempFileManager; + private final SvgSanitizer svgSanitizer; @AutoJobPostMapping( consumes = MediaType.MULTIPART_FORM_DATA_VALUE, @@ -61,6 +63,9 @@ public class OverlayImageController { byte[] imageBytes = imageFile.getBytes(); boolean isSvg = SvgOverlayUtil.isSvgImage(imageBytes); + if (isSvg) { + imageBytes = svgSanitizer.sanitize(imageBytes); + } try (PDDocument document = pdfDocumentFactory.load(pdfBytes)) { int pages = document.getNumberOfPages(); diff --git a/app/core/src/main/java/stirling/software/SPDF/controller/api/misc/PrintFileController.java b/app/core/src/main/java/stirling/software/SPDF/controller/api/misc/PrintFileController.java index 48a9b5586f..e369bbf102 100644 --- a/app/core/src/main/java/stirling/software/SPDF/controller/api/misc/PrintFileController.java +++ b/app/core/src/main/java/stirling/software/SPDF/controller/api/misc/PrintFileController.java @@ -9,7 +9,6 @@ import java.awt.print.PrinterJob; import java.io.IOException; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import java.nio.file.StandardCopyOption; import java.util.Arrays; import java.util.Locale; @@ -54,7 +53,7 @@ public class PrintFileController { MultipartFile file = request.getFileInput(); String originalFilename = file.getOriginalFilename(); if (originalFilename != null - && (originalFilename.contains("..") || Paths.get(originalFilename).isAbsolute())) { + && (originalFilename.contains("..") || Path.of(originalFilename).isAbsolute())) { throw ExceptionUtils.createIllegalArgumentException( "error.invalid.filepath", "Invalid file path detected: " + originalFilename); } diff --git a/app/core/src/main/java/stirling/software/SPDF/controller/api/pipeline/PipelineDirectoryProcessor.java b/app/core/src/main/java/stirling/software/SPDF/controller/api/pipeline/PipelineDirectoryProcessor.java index ef0afce94b..79223ddd38 100644 --- a/app/core/src/main/java/stirling/software/SPDF/controller/api/pipeline/PipelineDirectoryProcessor.java +++ b/app/core/src/main/java/stirling/software/SPDF/controller/api/pipeline/PipelineDirectoryProcessor.java @@ -7,7 +7,6 @@ import java.nio.file.FileVisitOption; import java.nio.file.FileVisitResult; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import java.nio.file.SimpleFileVisitor; import java.nio.file.StandardCopyOption; import java.nio.file.attribute.BasicFileAttributes; @@ -82,7 +81,7 @@ public class PipelineDirectoryProcessor { try { for (String watchedFoldersDir : watchedFoldersDirs) { - scanWatchedFolder(Paths.get(watchedFoldersDir).toAbsolutePath()); + scanWatchedFolder(Path.of(watchedFoldersDir).toAbsolutePath()); } } finally { // Clean up ThreadLocal to prevent memory leaks @@ -442,7 +441,7 @@ public class PipelineDirectoryProcessor { .replace("{outputFolder}", finishedFoldersDir) .replace("{folderName}", dir.toString())) .replaceAll(""); - return Paths.get(outputDir).isAbsolute() ? Paths.get(outputDir) : Paths.get(".", outputDir); + return Path.of(outputDir).isAbsolute() ? Path.of(outputDir) : Path.of(".", outputDir); } private void deleteOriginalFiles(List filesToProcess, Path processingDir) diff --git a/app/core/src/main/java/stirling/software/SPDF/controller/api/pipeline/PipelineProcessor.java b/app/core/src/main/java/stirling/software/SPDF/controller/api/pipeline/PipelineProcessor.java index fde2cfa900..c57a43f27f 100644 --- a/app/core/src/main/java/stirling/software/SPDF/controller/api/pipeline/PipelineProcessor.java +++ b/app/core/src/main/java/stirling/software/SPDF/controller/api/pipeline/PipelineProcessor.java @@ -5,7 +5,6 @@ import java.net.URLDecoder; import java.nio.charset.StandardCharsets; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import java.util.ArrayList; import java.util.List; import java.util.Locale; @@ -326,12 +325,12 @@ public class PipelineProcessor { } List outputFiles = new ArrayList<>(); for (File file : files) { - Path normalizedPath = Paths.get(file.getName()).normalize(); + Path normalizedPath = Path.of(file.getName()).normalize(); if (normalizedPath.startsWith("..")) { throw new SecurityException( "Potential path traversal attempt in file name: " + file.getName()); } - Path path = Paths.get(file.getAbsolutePath()); + Path path = Path.of(file.getAbsolutePath()); // debug statement log.info("Reading file: {}", path); if (Files.exists(path)) { diff --git a/app/core/src/main/java/stirling/software/SPDF/controller/web/ReactRoutingController.java b/app/core/src/main/java/stirling/software/SPDF/controller/web/ReactRoutingController.java index ce0f5edaf0..62ce9dac01 100644 --- a/app/core/src/main/java/stirling/software/SPDF/controller/web/ReactRoutingController.java +++ b/app/core/src/main/java/stirling/software/SPDF/controller/web/ReactRoutingController.java @@ -5,7 +5,6 @@ import java.io.InputStream; import java.nio.charset.StandardCharsets; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import java.util.regex.Pattern; import org.springframework.beans.factory.annotation.Value; @@ -41,6 +40,8 @@ public class ReactRoutingController { private boolean indexHtmlExists = false; private boolean useExternalIndexHtml = false; private boolean loggedMissingIndex = false; + private String cachedSaasLandingHtml; + private boolean saasLandingExists = false; @PostConstruct public void init() { @@ -49,8 +50,22 @@ public class ReactRoutingController { // Always initialize callback HTML (used for OAuth desktop flow) this.cachedCallbackHtml = buildCallbackHtml(); + // SaaS landing page: only present on the classpath when the :saas module is bundled + // (app/saas/src/main/resources/static/saas-landing.html). When present it replaces the + // root page so the SaaS API host shows its own landing instead of the OSS API-only page. + ClassPathResource saasLanding = new ClassPathResource("static/saas-landing.html"); + if (saasLanding.exists()) { + try (InputStream in = saasLanding.getInputStream()) { + this.cachedSaasLandingHtml = new String(in.readAllBytes(), StandardCharsets.UTF_8); + this.saasLandingExists = true; + log.info("SaaS landing page detected; serving it at '/' and '/index.html'"); + } catch (Exception ex) { + log.warn("Failed to read saas-landing.html; falling back to index.html", ex); + } + } + // Check for external index.html first (customFiles/static/) - Path externalIndexPath = Paths.get(InstallationPathConfig.getStaticPath(), "index.html"); + Path externalIndexPath = Path.of(InstallationPathConfig.getStaticPath(), "index.html"); log.debug("Checking for custom index.html at: {}", externalIndexPath); if (Files.exists(externalIndexPath) && Files.isReadable(externalIndexPath)) { log.info("Using custom index.html from: {}", externalIndexPath); @@ -120,7 +135,7 @@ public class ReactRoutingController { private Resource getIndexHtmlResource() { // Check external location first - Path externalIndexPath = Paths.get(InstallationPathConfig.getStaticPath(), "index.html"); + Path externalIndexPath = Path.of(InstallationPathConfig.getStaticPath(), "index.html"); if (Files.exists(externalIndexPath) && Files.isReadable(externalIndexPath)) { return new FileSystemResource(externalIndexPath.toFile()); } @@ -132,6 +147,18 @@ public class ReactRoutingController { @GetMapping( value = {"/", "/index.html"}, produces = MediaType.TEXT_HTML_VALUE) + public ResponseEntity serveRootPage(HttpServletRequest request) { + // Swap ONLY the root page for SaaS. SPA entry points that delegate to serveIndexHtml + // (/auth/callback, /share/{token}, forwarded routes) keep serving the normal shell. + if (saasLandingExists && cachedSaasLandingHtml != null) { + return ResponseEntity.ok() + .cacheControl(CacheControl.noCache().mustRevalidate()) + .contentType(MediaType.TEXT_HTML) + .body(cachedSaasLandingHtml); + } + return serveIndexHtml(request); + } + public ResponseEntity serveIndexHtml(HttpServletRequest request) { try { if (indexHtmlExists && cachedIndexHtml != null) { diff --git a/app/core/src/main/java/stirling/software/SPDF/exception/GlobalExceptionHandler.java b/app/core/src/main/java/stirling/software/SPDF/exception/GlobalExceptionHandler.java index 6b43b66ff6..504dd36f5b 100644 --- a/app/core/src/main/java/stirling/software/SPDF/exception/GlobalExceptionHandler.java +++ b/app/core/src/main/java/stirling/software/SPDF/exception/GlobalExceptionHandler.java @@ -24,6 +24,7 @@ import org.springframework.web.multipart.MaxUploadSizeExceededException; import org.springframework.web.multipart.support.MissingServletRequestPartException; import org.springframework.web.server.ResponseStatusException; import org.springframework.web.servlet.NoHandlerFoundException; +import org.springframework.web.servlet.resource.NoResourceFoundException; import jakarta.servlet.http.HttpServletRequest; import jakarta.servlet.http.HttpServletResponse; @@ -993,6 +994,49 @@ public class GlobalExceptionHandler { .body(problemDetail); } + /** Unmapped path → clean 404 instead of falling through to the generic 500 catch-all. */ + @ExceptionHandler(NoResourceFoundException.class) + public ResponseEntity handleNoResourceFound( + NoResourceFoundException ex, HttpServletRequest request) { + // /api/* miss = likely missing controller (operator-relevant); other paths = favicons, + // robots.txt, scanner noise. Demote the latter so prod logs aren't flooded. + String uri = request.getRequestURI(); + if (uri != null && uri.startsWith("/api/")) { + log.warn("No resource at {}: {}", uri, ex.getMessage()); + } else { + log.debug("No resource at {}: {}", uri, ex.getMessage()); + } + + String title = getLocalizedMessage("error.notFound.title", ErrorTitles.NOT_FOUND_DEFAULT); + String detail = + getLocalizedMessage( + "error.notFound.detail", + String.format( + "No endpoint found for %s %s", + request.getMethod(), request.getRequestURI()), + request.getMethod(), + request.getRequestURI()); + + ProblemDetail problemDetail = + createBaseProblemDetail(HttpStatus.NOT_FOUND, detail, request); + problemDetail.setType(URI.create(ErrorTypes.NOT_FOUND)); + problemDetail.setTitle(title); + problemDetail.setProperty("title", title); + problemDetail.setProperty("method", request.getMethod()); + addStandardHints( + problemDetail, + "error.notFound.hints", + List.of( + "Verify the URL path and HTTP method are correct.", + "Check the API base path and version if applicable.", + "Ensure there are no typos in the endpoint path.")); + problemDetail.setProperty("actionRequired", "Use a valid endpoint URL and method."); + + return ResponseEntity.status(HttpStatus.NOT_FOUND) + .contentType(PROBLEM_JSON) + .body(problemDetail); + } + /** * Handle IllegalArgumentException. * diff --git a/app/core/src/main/java/stirling/software/SPDF/model/api/converters/ConvertPDFToMarkdown.java b/app/core/src/main/java/stirling/software/SPDF/model/api/converters/ConvertPDFToMarkdown.java index ce5a610789..42ebd51ab3 100644 --- a/app/core/src/main/java/stirling/software/SPDF/model/api/converters/ConvertPDFToMarkdown.java +++ b/app/core/src/main/java/stirling/software/SPDF/model/api/converters/ConvertPDFToMarkdown.java @@ -1,11 +1,13 @@ package stirling.software.SPDF.model.api.converters; -import org.springframework.core.io.Resource; +import java.nio.charset.StandardCharsets; + import org.springframework.http.MediaType; import org.springframework.http.ResponseEntity; import org.springframework.web.bind.annotation.ModelAttribute; import org.springframework.web.multipart.MultipartFile; +import io.github.pixee.security.Filenames; import io.swagger.v3.oas.annotations.Operation; import lombok.RequiredArgsConstructor; @@ -15,8 +17,11 @@ import stirling.software.common.annotations.AutoJobPostMapping; import stirling.software.common.annotations.api.ConvertApi; import stirling.software.common.enumeration.ResourceWeight; import stirling.software.common.model.api.PDFFile; -import stirling.software.common.util.PDFToFile; +import stirling.software.common.pdf.PdfMarkdownConverter; +import stirling.software.common.util.TempFile; import stirling.software.common.util.TempFileManager; +import stirling.software.common.util.WebResponseUtils; +import stirling.software.jpdfium.PdfDocument; @ConvertApi @RequiredArgsConstructor @@ -33,10 +38,27 @@ public class ConvertPDFToMarkdown { summary = "Convert PDF to Markdown", description = "This endpoint converts a PDF file to Markdown format. Input:PDF Output:Markdown Type:SISO") - public ResponseEntity processPdfToMarkdown(@ModelAttribute PDFFile file) + public ResponseEntity processPdfToMarkdown(@ModelAttribute PDFFile file) throws Exception { MultipartFile inputFile = file.getFileInput(); - PDFToFile pdfToFile = new PDFToFile(tempFileManager); - return pdfToFile.processPdfToMarkdown(inputFile); + + String originalName = Filenames.toSimpleFileName(inputFile.getOriginalFilename()); + String baseName = + originalName.contains(".") + ? originalName.substring(0, originalName.lastIndexOf('.')) + : originalName; + + String markdown; + try (TempFile tempInput = new TempFile(tempFileManager, ".pdf")) { + inputFile.transferTo(tempInput.getFile()); + try (PdfDocument doc = PdfDocument.open(tempInput.getPath())) { + markdown = new PdfMarkdownConverter().convert(doc); + } + } + + return WebResponseUtils.bytesToWebResponse( + markdown.getBytes(StandardCharsets.UTF_8), + baseName + ".md", + MediaType.valueOf("text/markdown")); } } diff --git a/app/core/src/main/java/stirling/software/SPDF/service/ApiDocService.java b/app/core/src/main/java/stirling/software/SPDF/service/ApiDocService.java index b0d2e5a4a0..3a8d418aa6 100644 --- a/app/core/src/main/java/stirling/software/SPDF/service/ApiDocService.java +++ b/app/core/src/main/java/stirling/software/SPDF/service/ApiDocService.java @@ -63,6 +63,7 @@ public class ApiDocService implements stirling.software.common.service.ToolMetad return "http://localhost:" + port + contextPath + "/v1/api-docs"; } + @Override public List getExtensionTypes(boolean output, String operationName) { if (outputToFileTypes.isEmpty()) { outputToFileTypes.put("PDF", List.of("pdf")); diff --git a/app/core/src/main/java/stirling/software/SPDF/service/PdfJsonConversionService.java b/app/core/src/main/java/stirling/software/SPDF/service/PdfJsonConversionService.java index 1eed1d2759..fea3b52fc5 100644 --- a/app/core/src/main/java/stirling/software/SPDF/service/PdfJsonConversionService.java +++ b/app/core/src/main/java/stirling/software/SPDF/service/PdfJsonConversionService.java @@ -588,7 +588,7 @@ public class PdfJsonConversionService { .replaceAll(""); return String.format("%s (%s)", cleanName, subtype); }) - .collect(java.util.stream.Collectors.toList()); + .toList(); long type3Fonts = responseFonts.stream().filter(f -> "Type3".equals(f.getSubtype())).count(); @@ -1554,7 +1554,7 @@ public class PdfJsonConversionService { .glyphName(outline.getGlyphName()) .unicode(outline.getUnicode()) .build()) - .collect(Collectors.toList()); + .toList(); } } catch (Exception ex) { log.debug( diff --git a/app/core/src/main/java/stirling/software/SPDF/service/SharedSignatureService.java b/app/core/src/main/java/stirling/software/SPDF/service/SharedSignatureService.java index b1321de3f3..3f52dbea20 100644 --- a/app/core/src/main/java/stirling/software/SPDF/service/SharedSignatureService.java +++ b/app/core/src/main/java/stirling/software/SPDF/service/SharedSignatureService.java @@ -4,7 +4,6 @@ import java.io.FileNotFoundException; import java.io.IOException; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import java.nio.file.StandardOpenOption; import java.util.ArrayList; import java.util.Base64; @@ -41,8 +40,8 @@ public class SharedSignatureService { public boolean hasAccessToFile(String username, String fileName) throws IOException { validateFileName(fileName); // Check if file exists in user's personal folder or ALL_USERS folder - Path userPath = Paths.get(SIGNATURE_BASE_PATH, username, fileName); - Path allUsersPath = Paths.get(SIGNATURE_BASE_PATH, ALL_USERS_FOLDER, fileName); + Path userPath = Path.of(SIGNATURE_BASE_PATH, username, fileName); + Path allUsersPath = Path.of(SIGNATURE_BASE_PATH, ALL_USERS_FOLDER, fileName); return Files.exists(userPath) || Files.exists(allUsersPath); } @@ -52,7 +51,7 @@ public class SharedSignatureService { // Get signatures from user's personal folder if (StringUtils.hasText(username)) { - Path userFolder = Paths.get(SIGNATURE_BASE_PATH, username); + Path userFolder = Path.of(SIGNATURE_BASE_PATH, username); if (Files.exists(userFolder)) { try { signatures.addAll(getSignaturesFromFolder(userFolder, "Personal")); @@ -63,7 +62,7 @@ public class SharedSignatureService { } // Get signatures from ALL_USERS folder - Path allUsersFolder = Paths.get(SIGNATURE_BASE_PATH, ALL_USERS_FOLDER); + Path allUsersFolder = Path.of(SIGNATURE_BASE_PATH, ALL_USERS_FOLDER); if (Files.exists(allUsersFolder)) { try { signatures.addAll(getSignaturesFromFolder(allUsersFolder, "Shared")); @@ -90,7 +89,7 @@ public class SharedSignatureService { */ public byte[] getSharedSignatureBytes(String fileName) throws IOException { validateFileName(fileName); - Path allUsersPath = Paths.get(SIGNATURE_BASE_PATH, ALL_USERS_FOLDER, fileName); + Path allUsersPath = Path.of(SIGNATURE_BASE_PATH, ALL_USERS_FOLDER, fileName); if (!Files.exists(allUsersPath)) { throw new FileNotFoundException("Shared signature file not found"); } @@ -142,7 +141,7 @@ public class SharedSignatureService { } String folderName = "shared".equals(scope) ? ALL_USERS_FOLDER : username; - Path targetFolder = Paths.get(SIGNATURE_BASE_PATH, folderName); + Path targetFolder = Path.of(SIGNATURE_BASE_PATH, folderName); Files.createDirectories(targetFolder); long timestamp = System.currentTimeMillis(); @@ -193,7 +192,7 @@ public class SharedSignatureService { List signatures = new ArrayList<>(); // Load personal signatures - Path personalFolder = Paths.get(SIGNATURE_BASE_PATH, username); + Path personalFolder = Path.of(SIGNATURE_BASE_PATH, username); if (Files.exists(personalFolder)) { try (Stream stream = Files.list(personalFolder)) { stream.filter(this::isImageFile) @@ -224,7 +223,7 @@ public class SharedSignatureService { } // Load shared signatures - Path sharedFolder = Paths.get(SIGNATURE_BASE_PATH, ALL_USERS_FOLDER); + Path sharedFolder = Path.of(SIGNATURE_BASE_PATH, ALL_USERS_FOLDER); if (Files.exists(sharedFolder)) { try (Stream stream = Files.list(sharedFolder)) { stream.filter(this::isImageFile) @@ -262,7 +261,7 @@ public class SharedSignatureService { validateFileName(signatureId); // Try to find and delete image file in personal folder - Path personalFolder = Paths.get(SIGNATURE_BASE_PATH, username); + Path personalFolder = Path.of(SIGNATURE_BASE_PATH, username); boolean deleted = false; if (Files.exists(personalFolder)) { @@ -283,7 +282,7 @@ public class SharedSignatureService { // Try shared folder if not found in personal if (!deleted) { - Path sharedFolder = Paths.get(SIGNATURE_BASE_PATH, ALL_USERS_FOLDER); + Path sharedFolder = Path.of(SIGNATURE_BASE_PATH, ALL_USERS_FOLDER); if (Files.exists(sharedFolder)) { try (Stream stream = Files.list(sharedFolder)) { List matchingFiles = diff --git a/app/core/src/main/java/stirling/software/SPDF/service/pdfjson/type3/library/Type3FontLibrary.java b/app/core/src/main/java/stirling/software/SPDF/service/pdfjson/type3/library/Type3FontLibrary.java index 0c885dfb24..457d2c9c8a 100644 --- a/app/core/src/main/java/stirling/software/SPDF/service/pdfjson/type3/library/Type3FontLibrary.java +++ b/app/core/src/main/java/stirling/software/SPDF/service/pdfjson/type3/library/Type3FontLibrary.java @@ -10,7 +10,6 @@ import java.util.Locale; import java.util.Map; import java.util.Objects; import java.util.concurrent.ConcurrentHashMap; -import java.util.stream.Collectors; import org.apache.pdfbox.cos.COSName; import org.apache.pdfbox.pdmodel.font.PDType3Font; @@ -268,7 +267,7 @@ public class Type3FontLibrary { .filter(Objects::nonNull) .map(String::trim) .filter(s -> !s.isEmpty()) - .collect(Collectors.toList()); + .toList(); } private String normalizeAlias(String alias) { diff --git a/app/core/src/main/java/stirling/software/SPDF/service/pdfjson/type3/tool/Type3SignatureTool.java b/app/core/src/main/java/stirling/software/SPDF/service/pdfjson/type3/tool/Type3SignatureTool.java index ff1836913b..10bc396185 100644 --- a/app/core/src/main/java/stirling/software/SPDF/service/pdfjson/type3/tool/Type3SignatureTool.java +++ b/app/core/src/main/java/stirling/software/SPDF/service/pdfjson/type3/tool/Type3SignatureTool.java @@ -3,7 +3,6 @@ package stirling.software.SPDF.service.pdfjson.type3.tool; import java.io.IOException; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import java.util.ArrayDeque; import java.util.ArrayList; import java.util.Collections; @@ -285,9 +284,9 @@ public final class Type3SignatureTool { for (int i = 0; i < args.length; i++) { String arg = args[i]; if ("--pdf".equals(arg) && i + 1 < args.length) { - pdf = Paths.get(args[++i]); + pdf = Path.of(args[++i]); } else if ("--output".equals(arg) && i + 1 < args.length) { - output = Paths.get(args[++i]); + output = Path.of(args[++i]); } else if ("--pretty".equals(arg)) { pretty = true; } else if ("--help".equals(arg) || "-h".equals(arg)) { diff --git a/app/core/src/main/java/stirling/software/SPDF/service/telegram/TelegramPipelineBot.java b/app/core/src/main/java/stirling/software/SPDF/service/telegram/TelegramPipelineBot.java index 89910781b7..3b8d5cff10 100644 --- a/app/core/src/main/java/stirling/software/SPDF/service/telegram/TelegramPipelineBot.java +++ b/app/core/src/main/java/stirling/software/SPDF/service/telegram/TelegramPipelineBot.java @@ -8,7 +8,6 @@ import java.net.URISyntaxException; import java.net.URL; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import java.time.Duration; import java.time.Instant; import java.util.ArrayList; @@ -411,7 +410,7 @@ public class TelegramPipelineBot extends TelegramLongPollingBot { private Path getInboxFolder(Long chatId) throws IOException { Path baseInbox = - Paths.get( + Path.of( runtimePathConfig.getPipelineWatchedFoldersPath(), telegramProperties.getPipelineInboxFolder()); @@ -445,7 +444,7 @@ public class TelegramPipelineBot extends TelegramLongPollingBot { private List waitForPipelineOutputs(PipelineFileInfo info) throws IOException { - Path finishedDir = Paths.get(runtimePathConfig.getPipelineFinishedFoldersPath()); + Path finishedDir = Path.of(runtimePathConfig.getPipelineFinishedFoldersPath()); Files.createDirectories(finishedDir); Instant start = info.savedAt(); diff --git a/app/core/src/main/java/stirling/software/SPDF/utils/SvgOverlayUtil.java b/app/core/src/main/java/stirling/software/SPDF/utils/SvgOverlayUtil.java index 253d9c687a..b58bb055fc 100644 --- a/app/core/src/main/java/stirling/software/SPDF/utils/SvgOverlayUtil.java +++ b/app/core/src/main/java/stirling/software/SPDF/utils/SvgOverlayUtil.java @@ -10,6 +10,7 @@ import org.apache.batik.bridge.GVTBuilder; import org.apache.batik.bridge.UserAgent; import org.apache.batik.bridge.UserAgentAdapter; import org.apache.batik.gvt.GraphicsNode; +import org.apache.batik.util.ParsedURL; import org.apache.batik.util.XMLResourceDescriptor; import org.apache.pdfbox.pdmodel.PDDocument; import org.apache.pdfbox.pdmodel.PDPage; @@ -39,7 +40,16 @@ public class SvgOverlayUtil { svgDoc = factory.createSVGDocument("file:///overlay.svg", inputStream); } - UserAgent userAgent = new UserAgentAdapter(); + UserAgent userAgent = + new UserAgentAdapter() { + @Override + public void checkLoadExternalResource( + ParsedURL resourceURL, ParsedURL docURL) { + throw new SecurityException( + "External resource loading is disabled for SVG overlays: " + + resourceURL); + } + }; DocumentLoader loader = new DocumentLoader(userAgent); BridgeContext ctx = new BridgeContext(userAgent, loader); ctx.setDynamicState(BridgeContext.DYNAMIC); diff --git a/app/core/src/main/resources/application.properties b/app/core/src/main/resources/application.properties index b5777e1a18..da564e454b 100644 --- a/app/core/src/main/resources/application.properties +++ b/app/core/src/main/resources/application.properties @@ -33,7 +33,7 @@ spring.security.filter.dispatcher-types=REQUEST,ERROR # Response compression server.compression.enabled=true server.compression.min-response-size=1024 -server.compression.mime-types=application/json,application/xml,text/html,text/plain,text/css,application/javascript,image/svg+xml,application/x-font-ttf,font/opentype,application/vnd.ms-fontobject,font/woff,font/woff2,application/font-woff,application/font-woff2 +server.compression.mime-types=application/json,application/xml,text/html,text/plain,text/css,application/javascript,image/svg+xml,application/x-font-ttf,font/opentype,application/vnd.ms-fontobject,font/woff,font/woff2,application/font-woff,application/font-woff2,application/wasm spring.web.error.path=/error spring.web.error.whitelabel.enabled=false @@ -93,6 +93,11 @@ posthog.host=https://eu.i.posthog.com spring.main.allow-bean-definition-overriding=true +# spring-data-redis is on the classpath only for the optional Valkey backplane (which wires its own +# factory); exclude Spring Boot's stock Redis auto-config so a default install doesn't create a dead +# localhost:6379 factory that flips /actuator/health to DOWN. +spring.autoconfigure.exclude=org.springframework.boot.data.redis.autoconfigure.DataRedisAutoConfiguration,org.springframework.boot.data.redis.autoconfigure.DataRedisReactiveAutoConfiguration + # Set up a consistent temporary directory location java.io.tmpdir=${stirling.tempfiles.directory:${java.io.tmpdir}/stirling-pdf} diff --git a/app/core/src/main/resources/settings.yml.template b/app/core/src/main/resources/settings.yml.template index d3684b4dae..9bce7d1c55 100644 --- a/app/core/src/main/resources/settings.yml.template +++ b/app/core/src/main/resources/settings.yml.template @@ -167,7 +167,7 @@ legal: impressum: "" # URL to the impressum of your application (e.g. https://example.com/impressum). Empty string to disable or filename to load from local file in static folder system: - defaultLocale: en-US # set the default language (e.g. 'de-DE', 'fr-FR', etc) + defaultLocale: "" # force a default language for new users (e.g. 'en-US', 'de-DE'). Empty string auto-detects from the browser, falling back to en-US googlevisibility: false # 'true' to allow Google visibility (via robots.txt), 'false' to disallow enableAlphaFunctionality: false # set to enable functionality which might need more testing before it fully goes live (this feature might make no changes) showUpdate: false # see when a new update is available @@ -290,6 +290,7 @@ storage: linkExpirationDays: 3 # Number of days before share links expire signing: enabled: false # set to 'true' to enable group signing workflow (requires storage.enabled) [ALPHA] + userListScope: org # Signing user-picker scope: 'org' (default) = whole instance, else caller's team only. autoPipeline: outputFolder: "" # Output folder for processed pipeline files (leave empty for default) fileReadiness: @@ -364,6 +365,39 @@ aiEngine: url: http://localhost:5001 # URL of the Python AI engine timeoutSeconds: 120 # Timeout in seconds for AI engine requests +policies: + # Folder automations can read from and write to the directories you allow here, so treat this as a + # security boundary. Leave allowedFolderRoots empty (default) to disable folder sources/outputs + # entirely; list absolute directories to permit folder access only within them. Stirling's own + # config directory is always off-limits, and folder access is always disabled in SaaS mode. + allowedFolderRoots: [] # e.g. ["/data/inbox", "/data/outbox"] + scheduleSweepSeconds: 60 # How often (seconds) scheduled policies are checked for being due + watchReconcileSeconds: 300 # How often (seconds) folder-watch re-syncs watches and re-runs as a safety net for missed events + watchQuietPeriodMs: 500 # How long (ms) folder-watch coalesces a burst of file events into a single run + streamTimeoutMs: 1800000 # SSE timeout (ms) for live run-progress streams + runExpiryMinutes: 30 # How long (minutes) a finished run's in-memory state is kept before eviction + +# Model Context Protocol (MCP) server. Exposes Stirling's PDF tools (grouped by namespace) +# plus the AI agents to MCP clients (Inspector, Claude Desktop, custom). OAuth-protected. +# Disabled by default - enable explicitly per deployment after configuring mcp.auth. +mcp: + enabled: false # Master switch. 'false' (default) means no /mcp endpoint, no metadata, no beans wired. + scopesEnabled: true # Enforce mcp.tools.read / mcp.tools.write scopes derived from operation category + allowedOperations: [] # Tool allow-list (operation ids, e.g. ['compress-pdf']). Empty = all. When set, ONLY these are exposed over MCP. + blockedOperations: [] # Tool deny-list (operation ids). Always removed from MCP even if otherwise allowed. + auth: + mode: oauth # 'oauth' (full OAuth2 resource server) or 'apikey' (Stirling per-user API key via X-API-KEY header; no external IdP needed - the low-friction self-host option) + issuerUri: "" # OAuth2 issuer URI (e.g. http://localhost:9000). Required when mode=oauth. + jwksUri: "" # JWKS URI. Blank -> derived from issuer's /.well-known/openid-configuration. + resourceId: "" # RFC 8707 resource identifier of THIS MCP server (e.g. http://localhost:8080/mcp). + # Required: tokens must list this id in `aud` or the request is rejected. + acceptedAudiences: [] # Extra `aud` values accepted on top of resourceId. Empty = strict RFC 8707. + # For IdPs that cannot mint resource audiences (Supabase OAuth server always + # issues aud=authenticated) list that audience here, e.g. ['authenticated']. + usernameClaim: sub # JWT claim matched against a Stirling username (e.g. 'sub', 'email', 'preferred_username') + requireExistingAccount: true # Reject tokens whose subject has no enabled Stirling account (recommended) + engineCapabilityRefreshMinutes: 5 # How often to refresh the AI capabilities manifest from the engine + # Cluster configuration. NOT YET ENABLED - scaffolding for later work. Leave at defaults. cluster: enabled: false # Master switch. 'false' (default) wires the in-process backplane and skips all cluster checks. Single-instance installs do not need to change anything here. diff --git a/app/core/src/main/resources/static/3rdPartyLicenses.json b/app/core/src/main/resources/static/3rdPartyLicenses.json index a867b14305..df7346ae39 100644 --- a/app/core/src/main/resources/static/3rdPartyLicenses.json +++ b/app/core/src/main/resources/static/3rdPartyLicenses.json @@ -2215,6 +2215,13 @@ "moduleLicense": "Apache License, Version 2.0", "moduleLicenseUrl": "https://www.apache.org/licenses/LICENSE-2.0" }, + { + "moduleName": "org.springframework.boot:spring-boot-security-oauth2-resource-server", + "moduleUrl": "https://spring.io/projects/spring-boot", + "moduleVersion": "4.0.6", + "moduleLicense": "Apache License, Version 2.0", + "moduleLicenseUrl": "https://www.apache.org/licenses/LICENSE-2.0" + }, { "moduleName": "org.springframework.boot:spring-boot-servlet", "moduleUrl": "https://spring.io/projects/spring-boot", @@ -2320,6 +2327,13 @@ "moduleLicense": "Apache License, Version 2.0", "moduleLicenseUrl": "https://www.apache.org/licenses/LICENSE-2.0" }, + { + "moduleName": "org.springframework.boot:spring-boot-starter-oauth2-resource-server", + "moduleUrl": "https://spring.io/projects/spring-boot", + "moduleVersion": "4.0.6", + "moduleLicense": "Apache License, Version 2.0", + "moduleLicenseUrl": "https://www.apache.org/licenses/LICENSE-2.0" + }, { "moduleName": "org.springframework.boot:spring-boot-starter-security", "moduleUrl": "https://spring.io/projects/spring-boot", @@ -2438,6 +2452,13 @@ "moduleLicense": "Apache License, Version 2.0", "moduleLicenseUrl": "https://www.apache.org/licenses/LICENSE-2.0" }, + { + "moduleName": "org.springframework.security:spring-security-oauth2-resource-server", + "moduleUrl": "https://spring.io/projects/spring-security", + "moduleVersion": "7.0.5", + "moduleLicense": "Apache License, Version 2.0", + "moduleLicenseUrl": "https://www.apache.org/licenses/LICENSE-2.0" + }, { "moduleName": "org.springframework.security:spring-security-saml2-service-provider", "moduleUrl": "https://spring.io/projects/spring-security", diff --git a/app/core/src/test/java/stirling/software/SPDF/controller/api/AdditionalLanguageJsControllerTest.java b/app/core/src/test/java/stirling/software/SPDF/controller/api/AdditionalLanguageJsControllerTest.java index 58c1b2e4e9..28d8bb414e 100644 --- a/app/core/src/test/java/stirling/software/SPDF/controller/api/AdditionalLanguageJsControllerTest.java +++ b/app/core/src/test/java/stirling/software/SPDF/controller/api/AdditionalLanguageJsControllerTest.java @@ -23,7 +23,7 @@ class AdditionalLanguageJsControllerTest { LanguageService lang = mock(LanguageService.class); // LinkedHashSet for deterministic order in the array when(lang.getSupportedLanguages()) - .thenReturn(new LinkedHashSet<>(List.of("de_DE", "en_GB"))); + .thenReturn(new LinkedHashSet<>(List.of("de_DE", "en_US"))); MockMvc mvc = MockMvcBuilders.standaloneSetup(new AdditionalLanguageJsController(lang)).build(); @@ -36,9 +36,9 @@ class AdditionalLanguageJsControllerTest { .string( containsString( "const supportedLanguages =" - + " [\"de_DE\",\"en_GB\"];"))) + + " [\"de_DE\",\"en_US\"];"))) .andExpect(content().string(containsString("function getDetailedLanguageCode()"))) - .andExpect(content().string(containsString("return \"en_GB\";"))); + .andExpect(content().string(containsString("return \"en_US\";"))); verify(lang, times(1)).getSupportedLanguages(); } diff --git a/app/core/src/test/java/stirling/software/SPDF/controller/api/MergeControllerGapTest.java b/app/core/src/test/java/stirling/software/SPDF/controller/api/MergeControllerGapTest.java new file mode 100644 index 0000000000..f5543bcdc7 --- /dev/null +++ b/app/core/src/test/java/stirling/software/SPDF/controller/api/MergeControllerGapTest.java @@ -0,0 +1,493 @@ +package stirling.software.SPDF.controller.api; + +import static org.junit.jupiter.api.Assertions.assertArrayEquals; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertSame; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.io.IOException; +import java.lang.reflect.Method; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Calendar; +import java.util.GregorianCalendar; +import java.util.List; + +import org.apache.pdfbox.pdmodel.PDDocument; +import org.apache.pdfbox.pdmodel.PDDocumentCatalog; +import org.apache.pdfbox.pdmodel.PDDocumentInformation; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.InjectMocks; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.http.MediaType; +import org.springframework.mock.web.MockMultipartFile; +import org.springframework.web.multipart.MultipartFile; + +import stirling.software.common.service.CustomPDFDocumentFactory; + +/** + * Gap tests for {@link MergeController} private helper logic reachable via reflection. Focuses on + * sort comparators, file-order reordering, client file-id parsing, date extraction and filename + * lookup. The external JPDFium merge path is not exercised here (covered structurally elsewhere). + */ +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class MergeControllerGapTest { + + @Mock private CustomPDFDocumentFactory pdfDocumentFactory; + + @InjectMocks private MergeController mergeController; + + private MockMultipartFile fileA; + private MockMultipartFile fileB; + private MockMultipartFile fileC; + + @BeforeEach + void setUp() { + fileA = + new MockMultipartFile( + "fileInput", "Apple.pdf", MediaType.APPLICATION_PDF_VALUE, "a".getBytes()); + fileB = + new MockMultipartFile( + "fileInput", "banana.pdf", MediaType.APPLICATION_PDF_VALUE, "b".getBytes()); + fileC = + new MockMultipartFile( + "fileInput", "Cherry.pdf", MediaType.APPLICATION_PDF_VALUE, "c".getBytes()); + } + + // ---- reflection helpers ------------------------------------------------- + + @SuppressWarnings("unchecked") + private java.util.Comparator sortComparator(String sortType) throws Exception { + Method m = MergeController.class.getDeclaredMethod("getSortComparator", String.class); + m.setAccessible(true); + return (java.util.Comparator) m.invoke(mergeController, sortType); + } + + private MultipartFile[] reorder(MultipartFile[] files, String fileOrder) throws Exception { + Method m = + MergeController.class.getDeclaredMethod( + "reorderFilesByProvidedOrder", MultipartFile[].class, String.class); + m.setAccessible(true); + return (MultipartFile[]) m.invoke(null, files, fileOrder); + } + + private String[] parseClientFileIds(String value) throws Exception { + Method m = MergeController.class.getDeclaredMethod("parseClientFileIds", String.class); + m.setAccessible(true); + return (String[]) m.invoke(mergeController, value); + } + + private long getPdfDateTimeSafe(MultipartFile file) throws Exception { + Method m = + MergeController.class.getDeclaredMethod("getPdfDateTimeSafe", MultipartFile.class); + m.setAccessible(true); + return (long) m.invoke(mergeController, file); + } + + @SuppressWarnings("unchecked") + private int indexOfByOriginalFilename(List list, String name) throws Exception { + Method m = + MergeController.class.getDeclaredMethod( + "indexOfByOriginalFilename", List.class, String.class); + m.setAccessible(true); + return (int) m.invoke(null, list, name); + } + + private static PDDocument docWithTitle(String title) { + PDDocument doc = mock(PDDocument.class); + PDDocumentInformation info = mock(PDDocumentInformation.class); + when(doc.getDocumentInformation()).thenReturn(info); + when(info.getTitle()).thenReturn(title); + return doc; + } + + // ---- getSortComparator: byFileName -------------------------------------- + + @Nested + @DisplayName("getSortComparator: byFileName") + class ByFileName { + + @Test + @DisplayName("sorts case-insensitively by original filename") + void sortsCaseInsensitively() throws Exception { + MultipartFile[] files = {fileC, fileA, fileB}; + Arrays.sort(files, sortComparator("byFileName")); + assertArrayEquals(new MultipartFile[] {fileA, fileB, fileC}, files); + } + + @Test + @DisplayName("null original filename is treated as empty and sorts first") + void nullFilenameSortsFirst() throws Exception { + MultipartFile nullName = mock(MultipartFile.class); + when(nullName.getOriginalFilename()).thenReturn(null); + MultipartFile[] files = {fileB, nullName, fileA}; + Arrays.sort(files, sortComparator("byFileName")); + assertSame(nullName, files[0]); + assertSame(fileA, files[1]); + assertSame(fileB, files[2]); + } + } + + // ---- getSortComparator: byPDFTitle -------------------------------------- + + @Nested + @DisplayName("getSortComparator: byPDFTitle") + class ByPdfTitle { + + @Test + @DisplayName("orders documents by their PDF title, ignoring case") + void ordersByTitle() throws Exception { + PDDocument docZ = docWithTitle("Zebra"); + PDDocument docA = docWithTitle("alpha"); + when(pdfDocumentFactory.load(fileA)).thenReturn(docZ); + when(pdfDocumentFactory.load(fileB)).thenReturn(docA); + + int cmp = sortComparator("byPDFTitle").compare(fileA, fileB); + assertTrue(cmp > 0, "Zebra should sort after alpha"); + // and the documents are closed via try-with-resources + verify(docZ).close(); + verify(docA).close(); + } + + @Test + @DisplayName("both titles null yields equal (0)") + void bothNullTitlesEqual() throws Exception { + PDDocument d1 = docWithTitle(null); + PDDocument d2 = docWithTitle(null); + when(pdfDocumentFactory.load(fileA)).thenReturn(d1); + when(pdfDocumentFactory.load(fileB)).thenReturn(d2); + + assertEquals(0, sortComparator("byPDFTitle").compare(fileA, fileB)); + } + + @Test + @DisplayName("first title null sorts after non-null (returns 1)") + void firstNullSortsLast() throws Exception { + PDDocument d1 = docWithTitle(null); + PDDocument d2 = docWithTitle("Beta"); + when(pdfDocumentFactory.load(fileA)).thenReturn(d1); + when(pdfDocumentFactory.load(fileB)).thenReturn(d2); + + assertEquals(1, sortComparator("byPDFTitle").compare(fileA, fileB)); + } + + @Test + @DisplayName("second title null sorts first (returns -1)") + void secondNullSortsFirst() throws Exception { + PDDocument d1 = docWithTitle("Alpha"); + PDDocument d2 = docWithTitle(null); + when(pdfDocumentFactory.load(fileA)).thenReturn(d1); + when(pdfDocumentFactory.load(fileB)).thenReturn(d2); + + assertEquals(-1, sortComparator("byPDFTitle").compare(fileA, fileB)); + } + + @Test + @DisplayName("IOException while loading yields equal (0)") + void ioExceptionYieldsEqual() throws Exception { + when(pdfDocumentFactory.load(fileA)).thenThrow(new IOException("boom")); + assertEquals(0, sortComparator("byPDFTitle").compare(fileA, fileB)); + } + } + + // ---- getSortComparator: date-based and no-op orders --------------------- + + @Nested + @DisplayName("getSortComparator: date-based and pass-through orders") + class DateAndPassThrough { + + private PDDocument docWithModDate(long millis) { + PDDocument doc = mock(PDDocument.class); + PDDocumentInformation info = mock(PDDocumentInformation.class); + Calendar cal = new GregorianCalendar(); + cal.setTimeInMillis(millis); + when(doc.getDocumentInformation()).thenReturn(info); + when(info.getModificationDate()).thenReturn(cal); + return doc; + } + + @Test + @DisplayName("byDateModified orders newest first (descending)") + void byDateModifiedNewestFirst() throws Exception { + PDDocument older = docWithModDate(1_000L); + PDDocument newer = docWithModDate(9_000L); + when(pdfDocumentFactory.load(fileA)).thenReturn(older); + when(pdfDocumentFactory.load(fileB)).thenReturn(newer); + + // file1=older, file2=newer -> Long.compare(t2=newer, t1=older) > 0 -> older after newer + int cmp = sortComparator("byDateModified").compare(fileA, fileB); + assertTrue(cmp > 0); + } + + @Test + @DisplayName("byDateCreated uses the same descending logic") + void byDateCreatedNewestFirst() throws Exception { + PDDocument older = docWithModDate(2_000L); + PDDocument newer = docWithModDate(8_000L); + when(pdfDocumentFactory.load(fileA)).thenReturn(newer); + when(pdfDocumentFactory.load(fileB)).thenReturn(older); + + int cmp = sortComparator("byDateCreated").compare(fileA, fileB); + assertTrue(cmp < 0, "newer (file1) should sort before older (file2)"); + } + + @Test + @DisplayName("orderProvided is a stable no-op comparator (0)") + void orderProvidedNoOp() throws Exception { + assertEquals(0, sortComparator("orderProvided").compare(fileA, fileB)); + } + + @Test + @DisplayName("unknown sort type falls back to no-op comparator (0)") + void unknownSortTypeNoOp() throws Exception { + assertEquals(0, sortComparator("somethingElse").compare(fileA, fileB)); + } + } + + // ---- getPdfDateTimeSafe ------------------------------------------------- + + @Nested + @DisplayName("getPdfDateTimeSafe") + class GetPdfDateTimeSafe { + + @Test + @DisplayName("returns modification date millis when present") + void returnsModificationDate() throws Exception { + PDDocument doc = mock(PDDocument.class); + PDDocumentInformation info = mock(PDDocumentInformation.class); + Calendar cal = new GregorianCalendar(); + cal.setTimeInMillis(123_456L); + when(doc.getDocumentInformation()).thenReturn(info); + when(info.getModificationDate()).thenReturn(cal); + when(pdfDocumentFactory.load(fileA)).thenReturn(doc); + + assertEquals(123_456L, getPdfDateTimeSafe(fileA)); + verify(doc).close(); + } + + @Test + @DisplayName("falls back to creation date when modification date is null") + void fallsBackToCreationDate() throws Exception { + PDDocument doc = mock(PDDocument.class); + PDDocumentInformation info = mock(PDDocumentInformation.class); + Calendar cal = new GregorianCalendar(); + cal.setTimeInMillis(777L); + when(doc.getDocumentInformation()).thenReturn(info); + when(info.getModificationDate()).thenReturn(null); + when(info.getCreationDate()).thenReturn(cal); + when(pdfDocumentFactory.load(fileA)).thenReturn(doc); + + assertEquals(777L, getPdfDateTimeSafe(fileA)); + } + + @Test + @DisplayName("returns 0 when no info dates and no XMP metadata present") + void returnsZeroWhenNoDates() throws Exception { + PDDocument doc = mock(PDDocument.class); + PDDocumentInformation info = mock(PDDocumentInformation.class); + PDDocumentCatalog catalog = mock(PDDocumentCatalog.class); + when(doc.getDocumentInformation()).thenReturn(info); + when(info.getModificationDate()).thenReturn(null); + when(info.getCreationDate()).thenReturn(null); + when(doc.getDocumentCatalog()).thenReturn(catalog); + when(catalog.getMetadata()).thenReturn(null); + when(pdfDocumentFactory.load(fileA)).thenReturn(doc); + + assertEquals(0L, getPdfDateTimeSafe(fileA)); + verify(doc).close(); + } + + @Test + @DisplayName("returns 0 when document info itself is null") + void returnsZeroWhenInfoNull() throws Exception { + PDDocument doc = mock(PDDocument.class); + PDDocumentCatalog catalog = mock(PDDocumentCatalog.class); + when(doc.getDocumentInformation()).thenReturn(null); + when(doc.getDocumentCatalog()).thenReturn(catalog); + when(catalog.getMetadata()).thenReturn(null); + when(pdfDocumentFactory.load(fileA)).thenReturn(doc); + + assertEquals(0L, getPdfDateTimeSafe(fileA)); + } + + @Test + @DisplayName("returns 0 and swallows IOException on load failure") + void returnsZeroOnLoadFailure() throws Exception { + when(pdfDocumentFactory.load(fileA)).thenThrow(new IOException("cannot open")); + assertEquals(0L, getPdfDateTimeSafe(fileA)); + } + } + + // ---- parseClientFileIds ------------------------------------------------- + + @Nested + @DisplayName("parseClientFileIds") + class ParseClientFileIds { + + @Test + @DisplayName("null input returns empty array") + void nullReturnsEmpty() throws Exception { + assertEquals(0, parseClientFileIds(null).length); + } + + @Test + @DisplayName("blank input returns empty array") + void blankReturnsEmpty() throws Exception { + assertEquals(0, parseClientFileIds(" ").length); + } + + @Test + @DisplayName("empty JSON array returns empty array") + void emptyArrayReturnsEmpty() throws Exception { + assertEquals(0, parseClientFileIds("[]").length); + assertEquals(0, parseClientFileIds("[ ]").length); + } + + @Test + @DisplayName("non-array text returns empty array") + void nonArrayReturnsEmpty() throws Exception { + assertEquals(0, parseClientFileIds("not-an-array").length); + } + + @Test + @DisplayName("parses quoted, comma-separated ids and strips surrounding quotes") + void parsesQuotedIds() throws Exception { + String[] result = parseClientFileIds("[\"id1\", \"id2\",\"id3\"]"); + assertArrayEquals(new String[] {"id1", "id2", "id3"}, result); + } + + @Test + @DisplayName("parses unquoted ids as-is after trimming") + void parsesUnquotedIds() throws Exception { + String[] result = parseClientFileIds("[a, b , c]"); + assertArrayEquals(new String[] {"a", "b", "c"}, result); + } + + @Test + @DisplayName("single element array yields a one-element result") + void singleElement() throws Exception { + assertArrayEquals(new String[] {"only"}, parseClientFileIds("[\"only\"]")); + } + } + + // ---- reorderFilesByProvidedOrder ---------------------------------------- + + @Nested + @DisplayName("reorderFilesByProvidedOrder") + class ReorderFilesByProvidedOrder { + + @Test + @DisplayName("reorders files to match the newline-separated order list") + void reordersToMatchOrder() throws Exception { + MultipartFile[] files = {fileA, fileB, fileC}; + MultipartFile[] result = reorder(files, "Cherry.pdf\nApple.pdf\nbanana.pdf"); + assertArrayEquals(new MultipartFile[] {fileC, fileA, fileB}, result); + } + + @Test + @DisplayName("handles CRLF separators") + void handlesCrlf() throws Exception { + MultipartFile[] files = {fileA, fileB}; + MultipartFile[] result = reorder(files, "banana.pdf\r\nApple.pdf"); + assertArrayEquals(new MultipartFile[] {fileB, fileA}, result); + } + + @Test + @DisplayName("unmatched names are skipped and remaining files appended in original order") + void unmatchedNamesAppendedAtEnd() throws Exception { + MultipartFile[] files = {fileA, fileB, fileC}; + // only mention Cherry; ghost.pdf is ignored; Apple+banana keep original relative order + MultipartFile[] result = reorder(files, "Cherry.pdf\nghost.pdf"); + assertArrayEquals(new MultipartFile[] {fileC, fileA, fileB}, result); + } + + @Test + @DisplayName("blank/empty order entries are skipped") + void blankEntriesSkipped() throws Exception { + MultipartFile[] files = {fileA, fileB}; + MultipartFile[] result = reorder(files, "\n \nbanana.pdf\n"); + assertArrayEquals(new MultipartFile[] {fileB, fileA}, result); + } + + @Test + @DisplayName("empty file array returns empty array") + void emptyFilesReturnsEmpty() throws Exception { + MultipartFile[] result = reorder(new MultipartFile[0], "anything.pdf"); + assertEquals(0, result.length); + } + } + + // ---- indexOfByOriginalFilename ------------------------------------------ + + @Nested + @DisplayName("indexOfByOriginalFilename") + class IndexOfByOriginalFilename { + + @Test + @DisplayName("returns index of matching filename") + void returnsMatchIndex() throws Exception { + List list = new ArrayList<>(Arrays.asList(fileA, fileB, fileC)); + assertEquals(1, indexOfByOriginalFilename(list, "banana.pdf")); + } + + @Test + @DisplayName("returns first match index when duplicates exist") + void returnsFirstMatch() throws Exception { + MockMultipartFile dup = + new MockMultipartFile( + "fileInput", + "Apple.pdf", + MediaType.APPLICATION_PDF_VALUE, + "dup".getBytes()); + List list = new ArrayList<>(Arrays.asList(fileA, dup)); + assertEquals(0, indexOfByOriginalFilename(list, "Apple.pdf")); + } + + @Test + @DisplayName("returns -1 when not found") + void returnsMinusOneWhenAbsent() throws Exception { + List list = new ArrayList<>(Arrays.asList(fileA, fileB)); + assertEquals(-1, indexOfByOriginalFilename(list, "missing.pdf")); + } + + @Test + @DisplayName("returns -1 for empty list") + void returnsMinusOneForEmpty() throws Exception { + assertEquals(-1, indexOfByOriginalFilename(new ArrayList<>(), "x.pdf")); + } + } + + // ---- mergeDocuments null-collaborator wiring ---------------------------- + + @Nested + @DisplayName("mergeDocuments wiring") + class MergeDocumentsWiring { + + @Test + @DisplayName("creates a fresh document from the factory and returns it") + void createsFromFactory() throws Exception { + PDDocument merged = mock(PDDocument.class); + when(pdfDocumentFactory.createNewDocument()).thenReturn(merged); + + PDDocument result = mergeController.mergeDocuments(List.of()); + + assertNotNull(result); + assertSame(merged, result); + verify(pdfDocumentFactory).createNewDocument(); + verify(merged, never()).close(); + } + } +} diff --git a/app/core/src/test/java/stirling/software/SPDF/controller/api/PosterPdfControllerTest.java b/app/core/src/test/java/stirling/software/SPDF/controller/api/PosterPdfControllerTest.java new file mode 100644 index 0000000000..37647caac6 --- /dev/null +++ b/app/core/src/test/java/stirling/software/SPDF/controller/api/PosterPdfControllerTest.java @@ -0,0 +1,446 @@ +package stirling.software.SPDF.controller.api; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.Mockito.lenient; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.io.ByteArrayInputStream; +import java.io.ByteArrayOutputStream; +import java.io.File; +import java.io.IOException; +import java.io.InputStream; +import java.lang.reflect.InvocationTargetException; +import java.lang.reflect.Method; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.zip.ZipEntry; +import java.util.zip.ZipInputStream; + +import org.apache.pdfbox.Loader; +import org.apache.pdfbox.pdmodel.PDDocument; +import org.apache.pdfbox.pdmodel.PDPage; +import org.apache.pdfbox.pdmodel.common.PDRectangle; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.junit.jupiter.api.io.TempDir; +import org.mockito.InjectMocks; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.core.io.Resource; +import org.springframework.http.HttpStatus; +import org.springframework.http.MediaType; +import org.springframework.http.ResponseEntity; +import org.springframework.mock.web.MockMultipartFile; + +import stirling.software.SPDF.model.api.general.PosterPdfRequest; +import stirling.software.common.service.CustomPDFDocumentFactory; +import stirling.software.common.util.TempFileManager; + +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class PosterPdfControllerTest { + + @TempDir Path tempDir; + + @Mock private CustomPDFDocumentFactory pdfDocumentFactory; + @Mock private TempFileManager tempFileManager; + + @InjectMocks private PosterPdfController controller; + + private final AtomicInteger tempCounter = new AtomicInteger(); + + @BeforeEach + void setUp() throws Exception { + // new TempFile(tempFileManager, suffix) delegates to createTempFile(suffix); + // hand back real, writable files in the test temp dir so the controller's + // real file/zip I/O works end to end. + lenient() + .when(tempFileManager.createTempFile(anyString())) + .thenAnswer( + inv -> { + String suffix = inv.getArgument(0); + File f = + tempDir.resolve( + "poster-" + + tempCounter.incrementAndGet() + + suffix) + .toFile(); + Files.createFile(f.toPath()); + return f; + }); + } + + private MockMultipartFile createRealPdf(int numPages, String name) throws IOException { + return createRealPdf(numPages, name, PDRectangle.A4, 0); + } + + private MockMultipartFile createRealPdf( + int numPages, String name, PDRectangle size, int rotation) throws IOException { + try (PDDocument doc = new PDDocument()) { + for (int i = 0; i < numPages; i++) { + PDPage page = new PDPage(size); + page.setRotation(rotation); + doc.addPage(page); + } + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + doc.save(baos); + return new MockMultipartFile( + "fileInput", name, MediaType.APPLICATION_PDF_VALUE, baos.toByteArray()); + } + } + + private PosterPdfRequest createRequest(MockMultipartFile file) { + PosterPdfRequest req = new PosterPdfRequest(); + req.setFileInput(file); + return req; + } + + /** Drain a file-backed Resource body to bytes. */ + private byte[] drainBody(ResponseEntity response) throws IOException { + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + try (InputStream in = response.getBody().getInputStream()) { + in.transferTo(baos); + } + return baos.toByteArray(); + } + + /** Read the single PDF entry out of a ZIP byte array. */ + private byte[] firstPdfEntry(byte[] zipBytes) throws IOException { + try (ZipInputStream zis = new ZipInputStream(new ByteArrayInputStream(zipBytes))) { + ZipEntry entry = zis.getNextEntry(); + assertThat(entry).isNotNull(); + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + zis.transferTo(baos); + return baos.toByteArray(); + } + } + + private void stubFactory(MockMultipartFile file) throws IOException { + PDDocument sourceDoc = Loader.loadPDF(file.getBytes()); + PDDocument outputDoc = new PDDocument(); + when(pdfDocumentFactory.load(file)).thenReturn(sourceDoc); + when(pdfDocumentFactory.createNewDocumentBasedOnOldDocument(sourceDoc)) + .thenReturn(outputDoc); + } + + @Nested + @DisplayName("posterPdf happy path") + class HappyPath { + + @Test + @DisplayName("Default 2x2 grid on single page yields a ZIP with a 4-page PDF") + void defaultGrid_singlePage() throws Exception { + MockMultipartFile file = createRealPdf(1, "doc.pdf"); + PosterPdfRequest request = createRequest(file); + stubFactory(file); + + ResponseEntity response = controller.posterPdf(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(response.getHeaders().getContentDisposition().getFilename()) + .isEqualTo("doc_poster.zip"); + assertThat(response.getHeaders().getContentType()) + .isEqualTo(MediaType.APPLICATION_OCTET_STREAM); + + byte[] zipBytes = drainBody(response); + assertThat(zipBytes).isNotEmpty(); + + byte[] pdfBytes = firstPdfEntry(zipBytes); + try (PDDocument result = Loader.loadPDF(pdfBytes)) { + // 1 source page * (xFactor 2 * yFactor 2) = 4 output pages + assertThat(result.getNumberOfPages()).isEqualTo(4); + } + } + + @Test + @DisplayName("ZIP entry is named _poster.pdf") + void zipEntryNamedAfterBase() throws Exception { + MockMultipartFile file = createRealPdf(1, "report.pdf"); + PosterPdfRequest request = createRequest(file); + stubFactory(file); + + ResponseEntity response = controller.posterPdf(request); + + byte[] zipBytes = drainBody(response); + try (ZipInputStream zis = new ZipInputStream(new ByteArrayInputStream(zipBytes))) { + ZipEntry entry = zis.getNextEntry(); + assertThat(entry).isNotNull(); + assertThat(entry.getName()).isEqualTo("report_poster.pdf"); + } + } + + @Test + @DisplayName("Multi-page source multiplies output page count by grid size") + void multiPageSource() throws Exception { + MockMultipartFile file = createRealPdf(3, "multi.pdf"); + PosterPdfRequest request = createRequest(file); + request.setXFactor(2); + request.setYFactor(3); + stubFactory(file); + + ResponseEntity response = controller.posterPdf(request); + + byte[] pdfBytes = firstPdfEntry(drainBody(response)); + try (PDDocument result = Loader.loadPDF(pdfBytes)) { + // 3 pages * (2 * 3) = 18 + assertThat(result.getNumberOfPages()).isEqualTo(18); + } + } + + @Test + @DisplayName("1x1 grid produces one output page per source page") + void oneByOneGrid() throws Exception { + MockMultipartFile file = createRealPdf(2, "one.pdf"); + PosterPdfRequest request = createRequest(file); + request.setXFactor(1); + request.setYFactor(1); + stubFactory(file); + + ResponseEntity response = controller.posterPdf(request); + + byte[] pdfBytes = firstPdfEntry(drainBody(response)); + try (PDDocument result = Loader.loadPDF(pdfBytes)) { + assertThat(result.getNumberOfPages()).isEqualTo(2); + } + } + + @Test + @DisplayName("Right-to-left ordering still produces the full grid") + void rightToLeft() throws Exception { + MockMultipartFile file = createRealPdf(1, "rtl.pdf"); + PosterPdfRequest request = createRequest(file); + request.setRightToLeft(true); + stubFactory(file); + + ResponseEntity response = controller.posterPdf(request); + + byte[] pdfBytes = firstPdfEntry(drainBody(response)); + try (PDDocument result = Loader.loadPDF(pdfBytes)) { + assertThat(result.getNumberOfPages()).isEqualTo(4); + } + } + + @Test + @DisplayName("Rotated source page (90 degrees) is handled without error") + void rotatedSourcePage() throws Exception { + MockMultipartFile file = createRealPdf(1, "rot.pdf", PDRectangle.A4, 90); + PosterPdfRequest request = createRequest(file); + stubFactory(file); + + ResponseEntity response = controller.posterPdf(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + byte[] pdfBytes = firstPdfEntry(drainBody(response)); + try (PDDocument result = Loader.loadPDF(pdfBytes)) { + assertThat(result.getNumberOfPages()).isEqualTo(4); + } + } + + @Test + @DisplayName("Filename without extension is preserved in output names") + void filenameWithoutExtension() throws Exception { + MockMultipartFile file = createRealPdf(1, "noext"); + PosterPdfRequest request = createRequest(file); + stubFactory(file); + + ResponseEntity response = controller.posterPdf(request); + + assertThat(response.getHeaders().getContentDisposition().getFilename()) + .isEqualTo("noext_poster.zip"); + try (ZipInputStream zis = + new ZipInputStream(new ByteArrayInputStream(drainBody(response)))) { + ZipEntry entry = zis.getNextEntry(); + assertThat(entry).isNotNull(); + assertThat(entry.getName()).isEqualTo("noext_poster.pdf"); + } + } + + @Test + @DisplayName("Null original filename falls back to default base name") + void nullOriginalFilename() throws Exception { + MockMultipartFile file = + new MockMultipartFile( + "fileInput", + null, + MediaType.APPLICATION_PDF_VALUE, + createRealPdf(1, "x.pdf").getBytes()); + PosterPdfRequest request = createRequest(file); + stubFactory(file); + + ResponseEntity response = controller.posterPdf(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + // MockMultipartFile maps a null name to "", so the base is empty -> leading underscore. + assertThat(response.getHeaders().getContentDisposition().getFilename()) + .isEqualTo("_poster.zip"); + } + } + + @Nested + @DisplayName("Page size handling") + class PageSizes { + + @Test + @DisplayName("Each supported page size produces a valid ZIP") + void supportedSizes() throws Exception { + for (String size : new String[] {"A4", "Letter", "A3", "A5", "Legal", "Tabloid"}) { + MockMultipartFile file = createRealPdf(1, "s.pdf"); + PosterPdfRequest request = createRequest(file); + request.setPageSize(size); + stubFactory(file); + + ResponseEntity response = controller.posterPdf(request); + + assertThat(response.getStatusCode()) + .as("page size %s", size) + .isEqualTo(HttpStatus.OK); + assertThat(drainBody(response)).as("body for %s", size).isNotEmpty(); + } + } + + @Test + @DisplayName("Invalid page size throws IllegalArgumentException") + void invalidPageSize() throws Exception { + MockMultipartFile file = createRealPdf(1, "bad.pdf"); + PosterPdfRequest request = createRequest(file); + request.setPageSize("NotAPageSize"); + stubFactory(file); + + assertThatThrownBy(() -> controller.posterPdf(request)) + .isInstanceOf(IllegalArgumentException.class); + } + } + + @Nested + @DisplayName("getTargetPageSize private mapping") + class TargetPageSize { + + private PDRectangle invoke(String size) throws Exception { + Method m = + PosterPdfController.class.getDeclaredMethod("getTargetPageSize", String.class); + m.setAccessible(true); + return (PDRectangle) m.invoke(controller, size); + } + + @Test + @DisplayName("Known sizes map to expected PDRectangles") + void knownSizes() throws Exception { + assertThat(invoke("A4")).isEqualTo(PDRectangle.A4); + assertThat(invoke("Letter")).isEqualTo(PDRectangle.LETTER); + assertThat(invoke("A3")).isEqualTo(PDRectangle.A3); + assertThat(invoke("A5")).isEqualTo(PDRectangle.A5); + assertThat(invoke("Legal")).isEqualTo(PDRectangle.LEGAL); + } + + @Test + @DisplayName("Tabloid maps to 11x17 inch (792x1224 pt) rectangle") + void tabloidSize() throws Exception { + PDRectangle r = invoke("Tabloid"); + assertThat(r.getWidth()).isEqualTo(792f); + assertThat(r.getHeight()).isEqualTo(1224f); + } + + @Test + @DisplayName("Unknown size raises IllegalArgumentException") + void unknownSize() throws Exception { + Method m = + PosterPdfController.class.getDeclaredMethod("getTargetPageSize", String.class); + m.setAccessible(true); + assertThatThrownBy(() -> m.invoke(controller, "Unknown")) + .isInstanceOf(InvocationTargetException.class) + .hasCauseInstanceOf(IllegalArgumentException.class); + } + + @Test + @DisplayName("Null size raises IllegalArgumentException") + void nullSize() throws Exception { + Method m = + PosterPdfController.class.getDeclaredMethod("getTargetPageSize", String.class); + m.setAccessible(true); + assertThatThrownBy(() -> m.invoke(controller, new Object[] {null})) + .isInstanceOf(InvocationTargetException.class) + .hasCauseInstanceOf(IllegalArgumentException.class); + } + } + + @Nested + @DisplayName("Error propagation") + class Errors { + + @Test + @DisplayName("IOException from load propagates to caller") + void loadIoException() throws Exception { + MockMultipartFile file = createRealPdf(1, "io.pdf"); + PosterPdfRequest request = createRequest(file); + when(pdfDocumentFactory.load(file)).thenThrow(new IOException("load failed")); + + assertThatThrownBy(() -> controller.posterPdf(request)) + .isInstanceOf(IOException.class) + .hasMessageContaining("load failed"); + } + + @Test + @DisplayName("RuntimeException from createNewDocument propagates and closes zip temp file") + void createNewDocumentRuntimeException() throws Exception { + MockMultipartFile file = createRealPdf(1, "rt.pdf"); + PosterPdfRequest request = createRequest(file); + PDDocument sourceDoc = Loader.loadPDF(file.getBytes()); + when(pdfDocumentFactory.load(file)).thenReturn(sourceDoc); + when(pdfDocumentFactory.createNewDocumentBasedOnOldDocument(sourceDoc)) + .thenThrow(new IllegalStateException("boom")); + + assertThatThrownBy(() -> controller.posterPdf(request)) + .isInstanceOf(IllegalStateException.class) + .hasMessageContaining("boom"); + + sourceDoc.close(); + } + } + + @Nested + @DisplayName("Collaborator interactions") + class Interactions { + + @Test + @DisplayName("Both load and createNewDocumentBasedOnOldDocument are invoked") + void factoryCalled() throws Exception { + MockMultipartFile file = createRealPdf(1, "calls.pdf"); + PosterPdfRequest request = createRequest(file); + PDDocument sourceDoc = Loader.loadPDF(file.getBytes()); + PDDocument outputDoc = new PDDocument(); + when(pdfDocumentFactory.load(file)).thenReturn(sourceDoc); + when(pdfDocumentFactory.createNewDocumentBasedOnOldDocument(sourceDoc)) + .thenReturn(outputDoc); + + controller.posterPdf(request); + + verify(pdfDocumentFactory).load(file); + verify(pdfDocumentFactory).createNewDocumentBasedOnOldDocument(sourceDoc); + } + + @Test + @DisplayName("Zip temp file is never created when load fails before zip work") + void noOutputWhenLoadFails() throws Exception { + MockMultipartFile file = createRealPdf(1, "fail.pdf"); + PosterPdfRequest request = createRequest(file); + when(pdfDocumentFactory.load(file)).thenThrow(new IOException("nope")); + + assertThatThrownBy(() -> controller.posterPdf(request)).isInstanceOf(IOException.class); + + // createNewDocumentBasedOnOldDocument is never reached after load throws. + verify(pdfDocumentFactory, never()) + .createNewDocumentBasedOnOldDocument( + org.mockito.ArgumentMatchers.any(PDDocument.class)); + } + } +} diff --git a/app/core/src/test/java/stirling/software/SPDF/controller/api/RearrangePagesPDFControllerTest.java b/app/core/src/test/java/stirling/software/SPDF/controller/api/RearrangePagesPDFControllerTest.java index c31cae2958..a225c3fc52 100644 --- a/app/core/src/test/java/stirling/software/SPDF/controller/api/RearrangePagesPDFControllerTest.java +++ b/app/core/src/test/java/stirling/software/SPDF/controller/api/RearrangePagesPDFControllerTest.java @@ -4,10 +4,14 @@ import static org.junit.jupiter.api.Assertions.*; import static org.mockito.ArgumentMatchers.anyString; import static org.mockito.Mockito.*; +import java.io.ByteArrayOutputStream; import java.io.File; import java.io.IOException; import java.nio.file.Files; +import java.util.ArrayList; +import java.util.List; +import org.apache.pdfbox.Loader; import org.apache.pdfbox.pdmodel.PDDocument; import org.apache.pdfbox.pdmodel.PDPage; import org.junit.jupiter.api.BeforeEach; @@ -56,6 +60,38 @@ class RearrangePagesPDFControllerTest { "fileInput", "test.pdf", MediaType.APPLICATION_PDF_VALUE, new byte[] {1, 2, 3}); } + /** Build a real, in-memory PDDocument with the requested number of blank pages. */ + private PDDocument buildRealPdf(int pageCount) throws IOException { + PDDocument doc = new PDDocument(); + for (int i = 0; i < pageCount; i++) { + doc.addPage(new PDPage()); + } + return doc; + } + + /** + * Returns the underlying {@link org.apache.pdfbox.cos.COSDictionary} for each page in document + * order. PDPageTree returns a fresh PDPage wrapper per get(), so comparing wrappers with + * assertSame is unreliable - the COSDictionary identity is the stable handle. + */ + private List snapshotCosPages(PDDocument doc) { + List snapshot = new ArrayList<>(); + for (PDPage p : doc.getPages()) { + snapshot.add(p.getCOSObject()); + } + return snapshot; + } + + private List reloadAndSnapshot(ResponseEntity response) throws IOException { + try (var in = response.getBody().getInputStream(); + var baos = new ByteArrayOutputStream()) { + in.transferTo(baos); + try (PDDocument out = Loader.loadPDF(baos.toByteArray())) { + return snapshotCosPages(out); + } + } + } + @Test void testDeletePages_Success() throws IOException { MockMultipartFile file = createMockPdf(); @@ -83,27 +119,23 @@ class RearrangePagesPDFControllerTest { request.setPageNumbers(""); request.setCustomMode("REVERSE_ORDER"); - PDDocument mockDoc = mock(PDDocument.class); - PDDocument mockNewDoc = mock(PDDocument.class); - PDPage page0 = mock(PDPage.class); - PDPage page1 = mock(PDPage.class); - PDPage page2 = mock(PDPage.class); + try (PDDocument realDoc = buildRealPdf(3)) { + List originals = snapshotCosPages(realDoc); + when(pdfDocumentFactory.load(file)).thenReturn(realDoc); - when(pdfDocumentFactory.load(file)).thenReturn(mockDoc); - when(mockDoc.getNumberOfPages()).thenReturn(3); - when(mockDoc.getPage(0)).thenReturn(page0); - when(mockDoc.getPage(1)).thenReturn(page1); - when(mockDoc.getPage(2)).thenReturn(page2); - when(pdfDocumentFactory.createNewDocumentBasedOnOldDocument(mockDoc)) - .thenReturn(mockNewDoc); + ResponseEntity response = controller.rearrangePages(request); - ResponseEntity response = controller.rearrangePages(request); - - assertNotNull(response); - assertEquals(200, response.getStatusCode().value()); - verify(mockNewDoc).addPage(page2); - verify(mockNewDoc).addPage(page1); - verify(mockNewDoc).addPage(page0); + assertNotNull(response); + assertEquals(200, response.getStatusCode().value()); + List finalOrder = reloadAndSnapshot(response); + assertEquals(3, finalOrder.size()); + // We can no longer compare references after a save/reload, so compare via + // the in-memory snapshot taken *after* the controller mutated the source. + List mutatedSource = snapshotCosPages(realDoc); + assertSame(originals.get(2), mutatedSource.get(0)); + assertSame(originals.get(1), mutatedSource.get(1)); + assertSame(originals.get(0), mutatedSource.get(2)); + } } @Test @@ -114,25 +146,18 @@ class RearrangePagesPDFControllerTest { request.setPageNumbers(""); request.setCustomMode("REMOVE_FIRST"); - PDDocument mockDoc = mock(PDDocument.class); - PDDocument mockNewDoc = mock(PDDocument.class); - PDPage page0 = mock(PDPage.class); - PDPage page1 = mock(PDPage.class); - PDPage page2 = mock(PDPage.class); + try (PDDocument realDoc = buildRealPdf(3)) { + List originals = snapshotCosPages(realDoc); + when(pdfDocumentFactory.load(file)).thenReturn(realDoc); - when(pdfDocumentFactory.load(file)).thenReturn(mockDoc); - when(mockDoc.getNumberOfPages()).thenReturn(3); - when(mockDoc.getPage(1)).thenReturn(page1); - when(mockDoc.getPage(2)).thenReturn(page2); - when(pdfDocumentFactory.createNewDocumentBasedOnOldDocument(mockDoc)) - .thenReturn(mockNewDoc); + ResponseEntity response = controller.rearrangePages(request); - ResponseEntity response = controller.rearrangePages(request); - - assertNotNull(response); - verify(mockNewDoc).addPage(page1); - verify(mockNewDoc).addPage(page2); - verify(mockNewDoc, never()).addPage(page0); + assertNotNull(response); + List mutated = snapshotCosPages(realDoc); + assertEquals(2, mutated.size()); + assertSame(originals.get(1), mutated.get(0)); + assertSame(originals.get(2), mutated.get(1)); + } } @Test @@ -143,23 +168,18 @@ class RearrangePagesPDFControllerTest { request.setPageNumbers(""); request.setCustomMode("REMOVE_LAST"); - PDDocument mockDoc = mock(PDDocument.class); - PDDocument mockNewDoc = mock(PDDocument.class); - PDPage page0 = mock(PDPage.class); - PDPage page1 = mock(PDPage.class); + try (PDDocument realDoc = buildRealPdf(3)) { + List originals = snapshotCosPages(realDoc); + when(pdfDocumentFactory.load(file)).thenReturn(realDoc); - when(pdfDocumentFactory.load(file)).thenReturn(mockDoc); - when(mockDoc.getNumberOfPages()).thenReturn(3); - when(mockDoc.getPage(0)).thenReturn(page0); - when(mockDoc.getPage(1)).thenReturn(page1); - when(pdfDocumentFactory.createNewDocumentBasedOnOldDocument(mockDoc)) - .thenReturn(mockNewDoc); + ResponseEntity response = controller.rearrangePages(request); - ResponseEntity response = controller.rearrangePages(request); - - assertNotNull(response); - verify(mockNewDoc).addPage(page0); - verify(mockNewDoc).addPage(page1); + assertNotNull(response); + List mutated = snapshotCosPages(realDoc); + assertEquals(2, mutated.size()); + assertSame(originals.get(0), mutated.get(0)); + assertSame(originals.get(1), mutated.get(1)); + } } @Test @@ -170,21 +190,19 @@ class RearrangePagesPDFControllerTest { request.setPageNumbers(""); request.setCustomMode("REMOVE_FIRST_AND_LAST"); - PDDocument mockDoc = mock(PDDocument.class); - PDDocument mockNewDoc = mock(PDDocument.class); - PDPage page1 = mock(PDPage.class); + try (PDDocument realDoc = buildRealPdf(4)) { + List originals = snapshotCosPages(realDoc); + when(pdfDocumentFactory.load(file)).thenReturn(realDoc); - when(pdfDocumentFactory.load(file)).thenReturn(mockDoc); - when(mockDoc.getNumberOfPages()).thenReturn(4); - when(mockDoc.getPage(1)).thenReturn(page1); - when(mockDoc.getPage(2)).thenReturn(page1); - when(pdfDocumentFactory.createNewDocumentBasedOnOldDocument(mockDoc)) - .thenReturn(mockNewDoc); + ResponseEntity response = controller.rearrangePages(request); - ResponseEntity response = controller.rearrangePages(request); - - assertNotNull(response); - assertEquals(200, response.getStatusCode().value()); + assertNotNull(response); + assertEquals(200, response.getStatusCode().value()); + List mutated = snapshotCosPages(realDoc); + assertEquals(2, mutated.size()); + assertSame(originals.get(1), mutated.get(0)); + assertSame(originals.get(2), mutated.get(1)); + } } @Test @@ -195,23 +213,15 @@ class RearrangePagesPDFControllerTest { request.setPageNumbers(""); request.setCustomMode("DUPLEX_SORT"); - PDDocument mockDoc = mock(PDDocument.class); - PDDocument mockNewDoc = mock(PDDocument.class); - PDPage page0 = mock(PDPage.class); - PDPage page1 = mock(PDPage.class); - PDPage page2 = mock(PDPage.class); - PDPage page3 = mock(PDPage.class); + try (PDDocument realDoc = buildRealPdf(4)) { + when(pdfDocumentFactory.load(file)).thenReturn(realDoc); - when(pdfDocumentFactory.load(file)).thenReturn(mockDoc); - when(mockDoc.getNumberOfPages()).thenReturn(4); - when(mockDoc.getPage(anyInt())).thenReturn(page0); - when(pdfDocumentFactory.createNewDocumentBasedOnOldDocument(mockDoc)) - .thenReturn(mockNewDoc); + ResponseEntity response = controller.rearrangePages(request); - ResponseEntity response = controller.rearrangePages(request); - - assertNotNull(response); - assertEquals(200, response.getStatusCode().value()); + assertNotNull(response); + assertEquals(200, response.getStatusCode().value()); + assertEquals(4, realDoc.getNumberOfPages()); + } } @Test @@ -222,20 +232,15 @@ class RearrangePagesPDFControllerTest { request.setPageNumbers(""); request.setCustomMode("BOOKLET_SORT"); - PDDocument mockDoc = mock(PDDocument.class); - PDDocument mockNewDoc = mock(PDDocument.class); - PDPage page = mock(PDPage.class); + try (PDDocument realDoc = buildRealPdf(4)) { + when(pdfDocumentFactory.load(file)).thenReturn(realDoc); - when(pdfDocumentFactory.load(file)).thenReturn(mockDoc); - when(mockDoc.getNumberOfPages()).thenReturn(4); - when(mockDoc.getPage(anyInt())).thenReturn(page); - when(pdfDocumentFactory.createNewDocumentBasedOnOldDocument(mockDoc)) - .thenReturn(mockNewDoc); + ResponseEntity response = controller.rearrangePages(request); - ResponseEntity response = controller.rearrangePages(request); - - assertNotNull(response); - assertEquals(200, response.getStatusCode().value()); + assertNotNull(response); + assertEquals(200, response.getStatusCode().value()); + assertEquals(4, realDoc.getNumberOfPages()); + } } @Test @@ -246,20 +251,15 @@ class RearrangePagesPDFControllerTest { request.setPageNumbers(""); request.setCustomMode("ODD_EVEN_SPLIT"); - PDDocument mockDoc = mock(PDDocument.class); - PDDocument mockNewDoc = mock(PDDocument.class); - PDPage page = mock(PDPage.class); + try (PDDocument realDoc = buildRealPdf(4)) { + when(pdfDocumentFactory.load(file)).thenReturn(realDoc); - when(pdfDocumentFactory.load(file)).thenReturn(mockDoc); - when(mockDoc.getNumberOfPages()).thenReturn(4); - when(mockDoc.getPage(anyInt())).thenReturn(page); - when(pdfDocumentFactory.createNewDocumentBasedOnOldDocument(mockDoc)) - .thenReturn(mockNewDoc); + ResponseEntity response = controller.rearrangePages(request); - ResponseEntity response = controller.rearrangePages(request); - - assertNotNull(response); - assertEquals(200, response.getStatusCode().value()); + assertNotNull(response); + assertEquals(200, response.getStatusCode().value()); + assertEquals(4, realDoc.getNumberOfPages()); + } } @Test @@ -270,24 +270,20 @@ class RearrangePagesPDFControllerTest { request.setPageNumbers("3,1,2"); request.setCustomMode("custom"); - PDDocument mockDoc = mock(PDDocument.class); - PDDocument mockNewDoc = mock(PDDocument.class); - PDPage page0 = mock(PDPage.class); - PDPage page1 = mock(PDPage.class); - PDPage page2 = mock(PDPage.class); + try (PDDocument realDoc = buildRealPdf(3)) { + List originals = snapshotCosPages(realDoc); + when(pdfDocumentFactory.load(file)).thenReturn(realDoc); - when(pdfDocumentFactory.load(file)).thenReturn(mockDoc); - when(mockDoc.getNumberOfPages()).thenReturn(3); - when(mockDoc.getPage(0)).thenReturn(page0); - when(mockDoc.getPage(1)).thenReturn(page1); - when(mockDoc.getPage(2)).thenReturn(page2); - when(pdfDocumentFactory.createNewDocumentBasedOnOldDocument(mockDoc)) - .thenReturn(mockNewDoc); + ResponseEntity response = controller.rearrangePages(request); - ResponseEntity response = controller.rearrangePages(request); - - assertNotNull(response); - assertEquals(200, response.getStatusCode().value()); + assertNotNull(response); + assertEquals(200, response.getStatusCode().value()); + List mutated = snapshotCosPages(realDoc); + assertEquals(3, mutated.size()); + assertSame(originals.get(2), mutated.get(0)); + assertSame(originals.get(0), mutated.get(1)); + assertSame(originals.get(1), mutated.get(2)); + } } @Test @@ -298,21 +294,15 @@ class RearrangePagesPDFControllerTest { request.setPageNumbers("3"); request.setCustomMode("DUPLICATE"); - PDDocument mockDoc = mock(PDDocument.class); - PDDocument mockNewDoc = mock(PDDocument.class); - PDPage page = mock(PDPage.class); + try (PDDocument realDoc = buildRealPdf(2)) { + when(pdfDocumentFactory.load(file)).thenReturn(realDoc); - when(pdfDocumentFactory.load(file)).thenReturn(mockDoc); - when(mockDoc.getNumberOfPages()).thenReturn(2); - when(mockDoc.getPage(anyInt())).thenReturn(page); - when(pdfDocumentFactory.createNewDocumentBasedOnOldDocument(mockDoc)) - .thenReturn(mockNewDoc); + ResponseEntity response = controller.rearrangePages(request); - ResponseEntity response = controller.rearrangePages(request); - - assertNotNull(response); - // 2 pages * 3 duplicates = 6 addPage calls - verify(mockNewDoc, times(6)).addPage(page); + assertNotNull(response); + // 2 pages * 3 duplicates = 6 final pages + assertEquals(6, realDoc.getNumberOfPages()); + } } @Test @@ -323,19 +313,14 @@ class RearrangePagesPDFControllerTest { request.setPageNumbers(""); request.setCustomMode("SIDE_STITCH_BOOKLET_SORT"); - PDDocument mockDoc = mock(PDDocument.class); - PDDocument mockNewDoc = mock(PDDocument.class); - PDPage page = mock(PDPage.class); + try (PDDocument realDoc = buildRealPdf(4)) { + when(pdfDocumentFactory.load(file)).thenReturn(realDoc); - when(pdfDocumentFactory.load(file)).thenReturn(mockDoc); - when(mockDoc.getNumberOfPages()).thenReturn(4); - when(mockDoc.getPage(anyInt())).thenReturn(page); - when(pdfDocumentFactory.createNewDocumentBasedOnOldDocument(mockDoc)) - .thenReturn(mockNewDoc); + ResponseEntity response = controller.rearrangePages(request); - ResponseEntity response = controller.rearrangePages(request); - - assertNotNull(response); - assertEquals(200, response.getStatusCode().value()); + assertNotNull(response); + assertEquals(200, response.getStatusCode().value()); + assertEquals(4, realDoc.getNumberOfPages()); + } } } diff --git a/app/core/src/test/java/stirling/software/SPDF/controller/api/SettingsControllerTest.java b/app/core/src/test/java/stirling/software/SPDF/controller/api/SettingsControllerTest.java new file mode 100644 index 0000000000..11d4c4140a --- /dev/null +++ b/app/core/src/test/java/stirling/software/SPDF/controller/api/SettingsControllerTest.java @@ -0,0 +1,186 @@ +package stirling.software.SPDF.controller.api; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertSame; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.mockStatic; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.times; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.util.HashMap; +import java.util.Map; +import java.util.concurrent.ConcurrentHashMap; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.MockedStatic; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.http.HttpStatus; +import org.springframework.http.ResponseEntity; + +import stirling.software.SPDF.config.EndpointConfiguration; +import stirling.software.common.model.ApplicationProperties; +import stirling.software.common.util.GeneralUtils; + +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class SettingsControllerTest { + + @Mock private ApplicationProperties applicationProperties; + @Mock private EndpointConfiguration endpointConfiguration; + @Mock private ApplicationProperties.System system; + + private SettingsController settingsController; + + @BeforeEach + void setUp() { + settingsController = new SettingsController(applicationProperties, endpointConfiguration); + } + + @Nested + @DisplayName("updateApiKey (update-enable-analytics)") + class UpdateApiKey { + + @Test + @DisplayName("persists and returns 200 OK when analytics flag not yet set (null)") + void updatesWhenNotPreviouslySet() throws Exception { + when(applicationProperties.getSystem()).thenReturn(system); + when(system.getEnableAnalytics()).thenReturn(null); + + try (MockedStatic generalUtils = mockStatic(GeneralUtils.class)) { + ResponseEntity> response = + settingsController.updateApiKey(Boolean.TRUE); + + assertNotNull(response); + assertEquals(HttpStatus.OK, response.getStatusCode()); + assertNotNull(response.getBody()); + assertEquals("Updated", response.getBody().get("message")); + + generalUtils.verify( + () -> + GeneralUtils.saveKeyToSettings( + "system.enableAnalytics", Boolean.TRUE), + times(1)); + } + + verify(system).setEnableAnalytics(Boolean.TRUE); + } + + @Test + @DisplayName("persists the false value when enabling analytics is declined") + void updatesWithFalseValue() throws Exception { + when(applicationProperties.getSystem()).thenReturn(system); + when(system.getEnableAnalytics()).thenReturn(null); + + try (MockedStatic generalUtils = mockStatic(GeneralUtils.class)) { + ResponseEntity> response = + settingsController.updateApiKey(Boolean.FALSE); + + assertEquals(HttpStatus.OK, response.getStatusCode()); + assertEquals("Updated", response.getBody().get("message")); + + generalUtils.verify( + () -> + GeneralUtils.saveKeyToSettings( + "system.enableAnalytics", Boolean.FALSE), + times(1)); + } + + verify(system).setEnableAnalytics(Boolean.FALSE); + } + + @Test + @DisplayName("returns 208 ALREADY_REPORTED and does not persist when flag already true") + void alreadyReportedWhenAlreadyTrue() throws Exception { + when(applicationProperties.getSystem()).thenReturn(system); + when(system.getEnableAnalytics()).thenReturn(Boolean.TRUE); + + try (MockedStatic generalUtils = mockStatic(GeneralUtils.class)) { + ResponseEntity> response = + settingsController.updateApiKey(Boolean.TRUE); + + assertNotNull(response); + assertEquals(HttpStatus.ALREADY_REPORTED, response.getStatusCode()); + assertNotNull(response.getBody()); + + Object message = response.getBody().get("message"); + assertNotNull(message); + assertTrue( + message.toString().startsWith("Setting has already been set"), + "Unexpected message: " + message); + + generalUtils.verify(() -> GeneralUtils.saveKeyToSettings(any(), any()), never()); + } + + verify(system, never()).setEnableAnalytics(any()); + } + + @Test + @DisplayName("returns 208 ALREADY_REPORTED when flag already false (any non-null is set)") + void alreadyReportedWhenAlreadyFalse() throws Exception { + when(applicationProperties.getSystem()).thenReturn(system); + when(system.getEnableAnalytics()).thenReturn(Boolean.FALSE); + + try (MockedStatic generalUtils = mockStatic(GeneralUtils.class)) { + ResponseEntity> response = + settingsController.updateApiKey(Boolean.TRUE); + + assertEquals(HttpStatus.ALREADY_REPORTED, response.getStatusCode()); + generalUtils.verify( + () -> GeneralUtils.saveKeyToSettings(eq("system.enableAnalytics"), any()), + never()); + } + + verify(system, never()).setEnableAnalytics(any()); + } + } + + @Nested + @DisplayName("getDisabledEndpoints (get-endpoints-status)") + class GetDisabledEndpoints { + + @Test + @DisplayName("returns 200 OK with the endpoint status map from EndpointConfiguration") + void returnsEndpointStatuses() { + Map statuses = new ConcurrentHashMap<>(); + statuses.put("merge-pdfs", Boolean.TRUE); + statuses.put("remove-blanks", Boolean.FALSE); + when(endpointConfiguration.getEndpointStatuses()).thenReturn(statuses); + + ResponseEntity> response = + settingsController.getDisabledEndpoints(); + + assertNotNull(response); + assertEquals(HttpStatus.OK, response.getStatusCode()); + assertSame(statuses, response.getBody()); + assertEquals(Boolean.TRUE, response.getBody().get("merge-pdfs")); + assertEquals(Boolean.FALSE, response.getBody().get("remove-blanks")); + verify(endpointConfiguration).getEndpointStatuses(); + } + + @Test + @DisplayName("returns 200 OK with an empty map when no statuses are configured") + void returnsEmptyMap() { + Map statuses = new HashMap<>(); + when(endpointConfiguration.getEndpointStatuses()).thenReturn(statuses); + + ResponseEntity> response = + settingsController.getDisabledEndpoints(); + + assertEquals(HttpStatus.OK, response.getStatusCode()); + assertNotNull(response.getBody()); + assertTrue(response.getBody().isEmpty()); + } + } +} diff --git a/app/core/src/test/java/stirling/software/SPDF/controller/api/ToSinglePageControllerTest.java b/app/core/src/test/java/stirling/software/SPDF/controller/api/ToSinglePageControllerTest.java new file mode 100644 index 0000000000..49dc4fb757 --- /dev/null +++ b/app/core/src/test/java/stirling/software/SPDF/controller/api/ToSinglePageControllerTest.java @@ -0,0 +1,325 @@ +package stirling.software.SPDF.controller.api; + +import static org.junit.jupiter.api.Assertions.*; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.Mockito.*; + +import java.io.ByteArrayOutputStream; +import java.io.File; +import java.io.IOException; +import java.io.InputStream; +import java.nio.file.Files; + +import org.apache.pdfbox.Loader; +import org.apache.pdfbox.pdmodel.PDDocument; +import org.apache.pdfbox.pdmodel.PDPage; +import org.apache.pdfbox.pdmodel.common.PDRectangle; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.InjectMocks; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.core.io.Resource; +import org.springframework.http.MediaType; +import org.springframework.http.ResponseEntity; +import org.springframework.mock.web.MockMultipartFile; + +import stirling.software.common.model.api.PDFFile; +import stirling.software.common.service.CustomPDFDocumentFactory; +import stirling.software.common.util.TempFile; +import stirling.software.common.util.TempFileManager; + +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class ToSinglePageControllerTest { + + @Mock private CustomPDFDocumentFactory pdfDocumentFactory; + @Mock private TempFileManager tempFileManager; + @InjectMocks private ToSinglePageController controller; + + @BeforeEach + void setUp() throws Exception { + // Each managed temp file is backed by a real on-disk file so the response can be read back. + when(tempFileManager.createManagedTempFile(anyString())) + .thenAnswer( + inv -> { + File f = + Files.createTempFile("tsp-test", inv.getArgument(0)) + .toFile(); + f.deleteOnExit(); + TempFile tf = mock(TempFile.class); + when(tf.getFile()).thenReturn(f); + when(tf.getPath()).thenReturn(f.toPath()); + when(tf.getAbsolutePath()).thenReturn(f.getAbsolutePath()); + return tf; + }); + } + + /** Build a real in-memory PDF with the given per-page sizes and return its bytes. */ + private byte[] createPdf(PDRectangle... pageSizes) throws IOException { + try (PDDocument doc = new PDDocument()) { + for (PDRectangle size : pageSizes) { + doc.addPage(new PDPage(size)); + } + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + doc.save(baos); + return baos.toByteArray(); + } + } + + private PDFFile requestFor(String filename, byte[] pdfBytes) { + MockMultipartFile file = + new MockMultipartFile( + "fileInput", filename, MediaType.APPLICATION_PDF_VALUE, pdfBytes); + PDFFile request = new PDFFile(); + request.setFileInput(file); + return request; + } + + /** + * Stub the factory: load() returns a real PDDocument parsed from the PDFFile bytes, and + * createNewDocumentBasedOnOldDocument() returns a fresh empty document. + */ + private void setupFactory() throws IOException { + when(pdfDocumentFactory.load(any(PDFFile.class))) + .thenAnswer( + inv -> { + PDFFile pf = inv.getArgument(0); + return Loader.loadPDF(pf.getFileInput().getBytes()); + }); + when(pdfDocumentFactory.createNewDocumentBasedOnOldDocument(any(PDDocument.class))) + .thenAnswer(inv -> new PDDocument()); + } + + private byte[] drainBody(ResponseEntity response) throws IOException { + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + try (InputStream in = response.getBody().getInputStream()) { + in.transferTo(baos); + } + return baos.toByteArray(); + } + + @Nested + @DisplayName("Happy path") + class HappyPath { + + @Test + @DisplayName("Multi-page PDF collapses to one tall page") + void multiPageCollapsesToSinglePage() throws Exception { + // Three A4 portrait pages. + byte[] pdfBytes = createPdf(PDRectangle.A4, PDRectangle.A4, PDRectangle.A4); + PDFFile request = requestFor("input.pdf", pdfBytes); + setupFactory(); + + ResponseEntity response = controller.pdfToSinglePage(request); + + assertNotNull(response); + assertEquals(200, response.getStatusCode().value()); + assertNotNull(response.getBody()); + + byte[] out = drainBody(response); + assertTrue(out.length > 0, "output PDF must be non-empty"); + + try (PDDocument result = Loader.loadPDF(out)) { + assertEquals(1, result.getNumberOfPages(), "result must be a single page"); + PDRectangle box = result.getPage(0).getMediaBox(); + assertEquals( + PDRectangle.A4.getWidth(), + box.getWidth(), + 0.5f, + "width matches the input width"); + assertEquals( + PDRectangle.A4.getHeight() * 3, + box.getHeight(), + 0.5f, + "height is the sum of all page heights"); + } + } + + @Test + @DisplayName("Single-page input produces a single page of the same size") + void singlePageInput() throws Exception { + byte[] pdfBytes = createPdf(PDRectangle.A4); + PDFFile request = requestFor("single.pdf", pdfBytes); + setupFactory(); + + ResponseEntity response = controller.pdfToSinglePage(request); + + assertEquals(200, response.getStatusCode().value()); + try (PDDocument result = Loader.loadPDF(drainBody(response))) { + assertEquals(1, result.getNumberOfPages()); + PDRectangle box = result.getPage(0).getMediaBox(); + assertEquals(PDRectangle.A4.getWidth(), box.getWidth(), 0.5f); + assertEquals(PDRectangle.A4.getHeight(), box.getHeight(), 0.5f); + } + } + + @Test + @DisplayName("Mixed page sizes: width is the max, height is the sum") + void mixedPageSizes() throws Exception { + // A4 (595x842) + A3 (842x1191) -> width should be max (A3 width), height the sum. + byte[] pdfBytes = createPdf(PDRectangle.A4, PDRectangle.A3); + PDFFile request = requestFor("mixed.pdf", pdfBytes); + setupFactory(); + + ResponseEntity response = controller.pdfToSinglePage(request); + + assertEquals(200, response.getStatusCode().value()); + try (PDDocument result = Loader.loadPDF(drainBody(response))) { + assertEquals(1, result.getNumberOfPages()); + PDRectangle box = result.getPage(0).getMediaBox(); + assertEquals(PDRectangle.A3.getWidth(), box.getWidth(), 0.5f); + assertEquals( + PDRectangle.A4.getHeight() + PDRectangle.A3.getHeight(), + box.getHeight(), + 0.5f); + } + } + + @Test + @DisplayName("Landscape pages are handled (width from landscape, height summed)") + void landscapePages() throws Exception { + PDRectangle landscape = new PDRectangle(842, 595); + byte[] pdfBytes = createPdf(landscape, landscape); + PDFFile request = requestFor("landscape.pdf", pdfBytes); + setupFactory(); + + ResponseEntity response = controller.pdfToSinglePage(request); + + assertEquals(200, response.getStatusCode().value()); + try (PDDocument result = Loader.loadPDF(drainBody(response))) { + assertEquals(1, result.getNumberOfPages()); + PDRectangle box = result.getPage(0).getMediaBox(); + assertEquals(842f, box.getWidth(), 0.5f); + assertEquals(1190f, box.getHeight(), 0.5f); + } + } + } + + @Nested + @DisplayName("Collaborator interactions") + class Collaborators { + + @Test + @DisplayName("Source document is loaded from the request and then closed") + void loadsAndClosesSourceDocument() throws Exception { + byte[] pdfBytes = createPdf(PDRectangle.A4, PDRectangle.A4); + PDFFile request = requestFor("input.pdf", pdfBytes); + + PDDocument sourceSpy = spy(Loader.loadPDF(pdfBytes)); + when(pdfDocumentFactory.load(any(PDFFile.class))).thenReturn(sourceSpy); + when(pdfDocumentFactory.createNewDocumentBasedOnOldDocument(any(PDDocument.class))) + .thenAnswer(inv -> new PDDocument()); + + controller.pdfToSinglePage(request); + + verify(pdfDocumentFactory).load(any(PDFFile.class)); + verify(pdfDocumentFactory).createNewDocumentBasedOnOldDocument(any(PDDocument.class)); + // try-with-resources must close the loaded source document. + verify(sourceSpy).close(); + } + + @Test + @DisplayName("A managed temp file is requested for the response body") + void requestsManagedTempFile() throws Exception { + byte[] pdfBytes = createPdf(PDRectangle.A4); + PDFFile request = requestFor("input.pdf", pdfBytes); + setupFactory(); + + controller.pdfToSinglePage(request); + + verify(tempFileManager).createManagedTempFile(".pdf"); + } + } + + @Nested + @DisplayName("Filename handling") + class FilenameHandling { + + @Test + @DisplayName("Original filename is reflected in the Content-Disposition header") + void filenameInContentDisposition() throws Exception { + byte[] pdfBytes = createPdf(PDRectangle.A4); + PDFFile request = requestFor("MyReport.pdf", pdfBytes); + setupFactory(); + + ResponseEntity response = controller.pdfToSinglePage(request); + + String disposition = + response.getHeaders() + .getFirst(org.springframework.http.HttpHeaders.CONTENT_DISPOSITION); + assertNotNull(disposition); + // generateFilename strips the extension and appends _singlePage.pdf + assertTrue( + disposition.contains("MyReport_singlePage.pdf"), + "disposition should carry the generated single-page filename: " + disposition); + } + + @Test + @DisplayName("Null original filename does not throw and yields a default name") + void nullOriginalFilename() throws Exception { + byte[] pdfBytes = createPdf(PDRectangle.A4); + // MockMultipartFile with a null original filename. + MockMultipartFile file = + new MockMultipartFile( + "fileInput", null, MediaType.APPLICATION_PDF_VALUE, pdfBytes); + PDFFile request = new PDFFile(); + request.setFileInput(file); + setupFactory(); + + ResponseEntity response = controller.pdfToSinglePage(request); + + assertEquals(200, response.getStatusCode().value()); + assertNotNull(response.getBody()); + // MockMultipartFile maps a null name to "", so the base is empty -> leading underscore. + String disposition = + response.getHeaders() + .getFirst(org.springframework.http.HttpHeaders.CONTENT_DISPOSITION); + assertNotNull(disposition); + assertTrue( + disposition.contains("_singlePage.pdf"), + "disposition should carry the empty-base single-page name: " + disposition); + } + } + + @Nested + @DisplayName("Error branches") + class ErrorBranches { + + @Test + @DisplayName("IOException from load() propagates to the caller") + void loadIOExceptionPropagates() throws Exception { + PDFFile request = requestFor("broken.pdf", new byte[] {1, 2, 3}); + when(pdfDocumentFactory.load(any(PDFFile.class))) + .thenThrow(new IOException("cannot load")); + + IOException ex = + assertThrows(IOException.class, () -> controller.pdfToSinglePage(request)); + assertEquals("cannot load", ex.getMessage()); + // No new document or temp file should be created when load fails. + verify(pdfDocumentFactory, never()) + .createNewDocumentBasedOnOldDocument(any(PDDocument.class)); + verifyNoInteractions(tempFileManager); + } + + @Test + @DisplayName("IOException from temp file creation propagates") + void tempFileIOExceptionPropagates() throws Exception { + byte[] pdfBytes = createPdf(PDRectangle.A4); + PDFFile request = requestFor("input.pdf", pdfBytes); + setupFactory(); + when(tempFileManager.createManagedTempFile(anyString())) + .thenThrow(new IOException("disk full")); + + IOException ex = + assertThrows(IOException.class, () -> controller.pdfToSinglePage(request)); + assertEquals("disk full", ex.getMessage()); + } + } +} diff --git a/app/core/src/test/java/stirling/software/SPDF/controller/api/UIDataControllerGapTest.java b/app/core/src/test/java/stirling/software/SPDF/controller/api/UIDataControllerGapTest.java new file mode 100644 index 0000000000..df934461f8 --- /dev/null +++ b/app/core/src/test/java/stirling/software/SPDF/controller/api/UIDataControllerGapTest.java @@ -0,0 +1,423 @@ +package stirling.software.SPDF.controller.api; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.Mockito.lenient; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.nio.charset.StandardCharsets; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.List; +import java.util.Map; +import java.util.UUID; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.junit.jupiter.api.io.TempDir; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.core.io.DefaultResourceLoader; +import org.springframework.core.io.ResourceLoader; +import org.springframework.http.HttpStatus; +import org.springframework.http.ResponseEntity; + +import stirling.software.SPDF.model.Dependency; +import stirling.software.SPDF.model.SignatureFile; +import stirling.software.SPDF.service.SharedSignatureService; +import stirling.software.common.configuration.RuntimePathConfig; +import stirling.software.common.model.ApplicationProperties; +import stirling.software.common.service.UserServiceInterface; + +import tools.jackson.databind.ObjectMapper; +import tools.jackson.databind.json.JsonMapper; + +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class UIDataControllerGapTest { + + @Mock private ApplicationProperties applicationProperties; + @Mock private ApplicationProperties.System system; + @Mock private ApplicationProperties.Legal legal; + @Mock private SharedSignatureService signatureService; + @Mock private UserServiceInterface userService; + @Mock private RuntimePathConfig runtimePathConfig; + + private final ResourceLoader resourceLoader = new DefaultResourceLoader(); + private final ObjectMapper objectMapper = JsonMapper.builder().build(); + + private UIDataController controller(UserServiceInterface user) { + return new UIDataController( + applicationProperties, + signatureService, + user, + resourceLoader, + runtimePathConfig, + objectMapper); + } + + @BeforeEach + void setUp() { + lenient().when(applicationProperties.getSystem()).thenReturn(system); + lenient().when(applicationProperties.getLegal()).thenReturn(legal); + } + + @Nested + @DisplayName("getFooterData") + class FooterData { + + @Test + @DisplayName("maps all legal and analytics fields onto the response body") + void mapsAllFields() { + when(system.getEnableAnalytics()).thenReturn(Boolean.TRUE); + when(legal.getTermsAndConditions()).thenReturn("https://terms"); + when(legal.getPrivacyPolicy()).thenReturn("https://privacy"); + when(legal.getAccessibilityStatement()).thenReturn("https://a11y"); + when(legal.getCookiePolicy()).thenReturn("https://cookies"); + when(legal.getImpressum()).thenReturn("https://impressum"); + + ResponseEntity response = + controller(userService).getFooterData(); + + assertEquals(HttpStatus.OK, response.getStatusCode()); + UIDataController.FooterData body = response.getBody(); + assertNotNull(body); + assertEquals(Boolean.TRUE, body.getAnalyticsEnabled()); + assertEquals("https://terms", body.getTermsAndConditions()); + assertEquals("https://privacy", body.getPrivacyPolicy()); + assertEquals("https://a11y", body.getAccessibilityStatement()); + assertEquals("https://cookies", body.getCookiePolicy()); + assertEquals("https://impressum", body.getImpressum()); + } + + @Test + @DisplayName("propagates null/false values from configuration") + void handlesNulls() { + when(system.getEnableAnalytics()).thenReturn(Boolean.FALSE); + when(legal.getTermsAndConditions()).thenReturn(null); + when(legal.getPrivacyPolicy()).thenReturn(null); + when(legal.getAccessibilityStatement()).thenReturn(null); + when(legal.getCookiePolicy()).thenReturn(null); + when(legal.getImpressum()).thenReturn(null); + + ResponseEntity response = + controller(userService).getFooterData(); + + assertEquals(HttpStatus.OK, response.getStatusCode()); + UIDataController.FooterData body = response.getBody(); + assertNotNull(body); + assertEquals(Boolean.FALSE, body.getAnalyticsEnabled()); + assertNull(body.getTermsAndConditions()); + assertNull(body.getPrivacyPolicy()); + assertNull(body.getAccessibilityStatement()); + assertNull(body.getCookiePolicy()); + assertNull(body.getImpressum()); + } + } + + @Nested + @DisplayName("getHomeData") + class HomeData { + + @Test + @DisplayName("returns OK with a populated body regardless of survey env var") + void returnsOk() { + ResponseEntity response = + controller(userService).getHomeData(); + + assertEquals(HttpStatus.OK, response.getStatusCode()); + assertNotNull(response.getBody()); + // SHOW_SURVEY is unset in the test JVM, so the default (true) applies. + assertTrue(response.getBody().isShowSurveyFromDocker()); + } + } + + @Nested + @DisplayName("getLicensesData") + class LicensesData { + + @Test + @DisplayName("loads the bundled 3rdPartyLicenses.json from the classpath") + void loadsDependencies() { + ResponseEntity response = + controller(userService).getLicensesData(); + + assertEquals(HttpStatus.OK, response.getStatusCode()); + UIDataController.LicensesData body = response.getBody(); + assertNotNull(body); + List deps = body.getDependencies(); + assertNotNull(deps); + assertFalse(deps.isEmpty()); + // Each parsed dependency should at least carry a module name. + assertNotNull(deps.get(0).getModuleName()); + } + } + + @Nested + @DisplayName("getPipelineData") + class PipelineData { + + @Test + @DisplayName("returns the placeholder entry when the config directory is missing") + void missingDirectoryYieldsPlaceholder() { + String missing = Path.of("nonexistent-pipeline-dir-" + UUID.randomUUID()).toString(); + when(runtimePathConfig.getPipelineDefaultWebUiConfigs()).thenReturn(missing); + + ResponseEntity response = + controller(userService).getPipelineData(); + + assertEquals(HttpStatus.OK, response.getStatusCode()); + UIDataController.PipelineData body = response.getBody(); + assertNotNull(body); + assertTrue(body.getPipelineConfigs().isEmpty()); + assertEquals(1, body.getPipelineConfigsWithNames().size()); + Map placeholder = body.getPipelineConfigsWithNames().get(0); + assertEquals("", placeholder.get("json")); + assertEquals("No preloaded configs found", placeholder.get("name")); + } + + @Test + @DisplayName("uses the embedded name field when present") + void readsConfigWithName(@TempDir Path dir) throws Exception { + Files.writeString( + dir.resolve("config1.json"), + "{\"name\":\"My Pipeline\",\"operations\":[]}", + StandardCharsets.UTF_8); + when(runtimePathConfig.getPipelineDefaultWebUiConfigs()).thenReturn(dir.toString()); + + ResponseEntity response = + controller(userService).getPipelineData(); + + assertEquals(HttpStatus.OK, response.getStatusCode()); + UIDataController.PipelineData body = response.getBody(); + assertNotNull(body); + assertEquals(1, body.getPipelineConfigs().size()); + assertEquals(1, body.getPipelineConfigsWithNames().size()); + assertEquals("My Pipeline", body.getPipelineConfigsWithNames().get(0).get("name")); + assertTrue( + body.getPipelineConfigsWithNames().get(0).get("json").contains("My Pipeline")); + } + + @Test + @DisplayName("falls back to the filename (sans extension) when name is missing") + void fallsBackToFilenameWhenNameMissing(@TempDir Path dir) throws Exception { + Files.writeString( + dir.resolve("fallback-name.json"), + "{\"operations\":[]}", + StandardCharsets.UTF_8); + when(runtimePathConfig.getPipelineDefaultWebUiConfigs()).thenReturn(dir.toString()); + + ResponseEntity response = + controller(userService).getPipelineData(); + + assertEquals(HttpStatus.OK, response.getStatusCode()); + UIDataController.PipelineData body = response.getBody(); + assertNotNull(body); + assertEquals(1, body.getPipelineConfigsWithNames().size()); + assertEquals("fallback-name", body.getPipelineConfigsWithNames().get(0).get("name")); + } + + @Test + @DisplayName("falls back to the filename when name is blank") + void fallsBackToFilenameWhenNameBlank(@TempDir Path dir) throws Exception { + Files.writeString( + dir.resolve("blank-name.json"), + "{\"name\":\"\",\"operations\":[]}", + StandardCharsets.UTF_8); + when(runtimePathConfig.getPipelineDefaultWebUiConfigs()).thenReturn(dir.toString()); + + ResponseEntity response = + controller(userService).getPipelineData(); + + UIDataController.PipelineData body = response.getBody(); + assertNotNull(body); + assertEquals("blank-name", body.getPipelineConfigsWithNames().get(0).get("name")); + } + + @Test + @DisplayName("ignores non-json files in the config directory") + void ignoresNonJsonFiles(@TempDir Path dir) throws Exception { + Files.writeString(dir.resolve("notes.txt"), "ignore me", StandardCharsets.UTF_8); + Files.writeString( + dir.resolve("real.json"), "{\"name\":\"Real\"}", StandardCharsets.UTF_8); + when(runtimePathConfig.getPipelineDefaultWebUiConfigs()).thenReturn(dir.toString()); + + ResponseEntity response = + controller(userService).getPipelineData(); + + UIDataController.PipelineData body = response.getBody(); + assertNotNull(body); + assertEquals(1, body.getPipelineConfigs().size()); + assertEquals(1, body.getPipelineConfigsWithNames().size()); + assertEquals("Real", body.getPipelineConfigsWithNames().get(0).get("name")); + } + + @Test + @DisplayName("returns the placeholder when the directory exists but holds no json") + void emptyDirectoryYieldsPlaceholder(@TempDir Path dir) { + when(runtimePathConfig.getPipelineDefaultWebUiConfigs()).thenReturn(dir.toString()); + + ResponseEntity response = + controller(userService).getPipelineData(); + + UIDataController.PipelineData body = response.getBody(); + assertNotNull(body); + assertTrue(body.getPipelineConfigs().isEmpty()); + assertEquals(1, body.getPipelineConfigsWithNames().size()); + assertEquals( + "No preloaded configs found", + body.getPipelineConfigsWithNames().get(0).get("name")); + } + } + + @Nested + @DisplayName("getSignData") + class SignData { + + @Test + @DisplayName("uses the current username from the user service to fetch signatures") + void usesUsernameFromUserService() { + when(userService.getCurrentUsername()).thenReturn("alice"); + List sigs = List.of(new SignatureFile("sig.png", "Personal")); + when(signatureService.getAvailableSignatures("alice")).thenReturn(sigs); + + ResponseEntity response = + controller(userService).getSignData(); + + assertEquals(HttpStatus.OK, response.getStatusCode()); + UIDataController.SignData body = response.getBody(); + assertNotNull(body); + assertEquals(sigs, body.getSignatures()); + // Fonts come from the real resource loader; the list is never null. + assertNotNull(body.getFonts()); + verify(userService).getCurrentUsername(); + verify(signatureService).getAvailableSignatures("alice"); + } + + @Test + @DisplayName("falls back to an empty username when no user service is wired") + void nullUserServiceUsesEmptyUsername() { + when(signatureService.getAvailableSignatures("")).thenReturn(List.of()); + + ResponseEntity response = controller(null).getSignData(); + + assertEquals(HttpStatus.OK, response.getStatusCode()); + UIDataController.SignData body = response.getBody(); + assertNotNull(body); + assertNotNull(body.getSignatures()); + assertNotNull(body.getFonts()); + verify(signatureService).getAvailableSignatures(""); + } + } + + @Nested + @DisplayName("getOcrPdfData") + class OcrData { + + @Test + @DisplayName("returns an empty language list when the tessdata directory is absent") + void absentTessdataDirYieldsEmptyList() { + when(runtimePathConfig.getTessDataPath()) + .thenReturn(Path.of("nonexistent-tessdata-" + UUID.randomUUID()).toString()); + + ResponseEntity response = + controller(userService).getOcrPdfData(); + + assertEquals(HttpStatus.OK, response.getStatusCode()); + UIDataController.OcrData body = response.getBody(); + assertNotNull(body); + assertNotNull(body.getLanguages()); + assertTrue(body.getLanguages().isEmpty()); + } + + @Test + @DisplayName("lists trained languages, excludes osd, and sorts alphabetically") + void listsAndSortsTrainedLanguages(@TempDir Path dir) throws Exception { + Files.writeString(dir.resolve("eng.traineddata"), "x", StandardCharsets.UTF_8); + Files.writeString(dir.resolve("deu.traineddata"), "x", StandardCharsets.UTF_8); + Files.writeString(dir.resolve("osd.traineddata"), "x", StandardCharsets.UTF_8); + // Non-traineddata files must be ignored. + Files.writeString(dir.resolve("readme.txt"), "x", StandardCharsets.UTF_8); + when(runtimePathConfig.getTessDataPath()).thenReturn(dir.toString()); + + ResponseEntity response = + controller(userService).getOcrPdfData(); + + assertEquals(HttpStatus.OK, response.getStatusCode()); + UIDataController.OcrData body = response.getBody(); + assertNotNull(body); + // osd filtered out; remaining sorted alphabetically. + assertEquals(List.of("deu", "eng"), body.getLanguages()); + } + + @Test + @DisplayName("excludes osd case-insensitively") + void excludesOsdCaseInsensitively(@TempDir Path dir) throws Exception { + Files.writeString(dir.resolve("OSD.traineddata"), "x", StandardCharsets.UTF_8); + Files.writeString(dir.resolve("fra.traineddata"), "x", StandardCharsets.UTF_8); + when(runtimePathConfig.getTessDataPath()).thenReturn(dir.toString()); + + ResponseEntity response = + controller(userService).getOcrPdfData(); + + UIDataController.OcrData body = response.getBody(); + assertNotNull(body); + assertEquals(List.of("fra"), body.getLanguages()); + assertFalse(body.getLanguages().contains("OSD")); + } + } + + @Nested + @DisplayName("FontResource format mapping") + class FontResourceMapping { + + @Test + @DisplayName("maps known extensions to their CSS font-format strings") + void mapsKnownExtensions() { + assertEquals("truetype", new UIDataController.FontResource("Arial", "ttf").getType()); + assertEquals("woff", new UIDataController.FontResource("Arial", "woff").getType()); + assertEquals("woff2", new UIDataController.FontResource("Arial", "woff2").getType()); + assertEquals( + "embedded-opentype", + new UIDataController.FontResource("Arial", "eot").getType()); + assertEquals("svg", new UIDataController.FontResource("Arial", "svg").getType()); + } + + @Test + @DisplayName("maps unknown extensions to an empty type and preserves name/extension") + void mapsUnknownExtensionToEmpty() { + UIDataController.FontResource resource = + new UIDataController.FontResource("Arial", "otf"); + assertEquals("", resource.getType()); + assertEquals("Arial", resource.getName()); + assertEquals("otf", resource.getExtension()); + } + } + + @Nested + @DisplayName("Cross-cutting behaviour") + class CrossCutting { + + @Test + @DisplayName("getFooterData never touches the signature or user services") + void footerDataIsIndependentOfUserState() { + when(system.getEnableAnalytics()).thenReturn(null); + + controller(userService).getFooterData(); + + verify(signatureService, never()) + .getAvailableSignatures(org.mockito.ArgumentMatchers.anyString()); + verify(userService, never()).getCurrentUsername(); + } + } +} diff --git a/app/core/src/test/java/stirling/software/SPDF/controller/api/converters/ConvertImgPDFControllerGapTest.java b/app/core/src/test/java/stirling/software/SPDF/controller/api/converters/ConvertImgPDFControllerGapTest.java new file mode 100644 index 0000000000..4f87942f2c --- /dev/null +++ b/app/core/src/test/java/stirling/software/SPDF/controller/api/converters/ConvertImgPDFControllerGapTest.java @@ -0,0 +1,826 @@ +package stirling.software.SPDF.controller.api.converters; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertSame; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyBoolean; +import static org.mockito.ArgumentMatchers.anyInt; +import static org.mockito.ArgumentMatchers.eq; + +import java.io.IOException; + +import org.apache.pdfbox.pdmodel.PDDocument; +import org.apache.pdfbox.pdmodel.PDPage; +import org.apache.pdfbox.pdmodel.common.PDRectangle; +import org.apache.pdfbox.rendering.ImageType; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.InjectMocks; +import org.mockito.Mock; +import org.mockito.MockedStatic; +import org.mockito.Mockito; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.core.io.Resource; +import org.springframework.http.MediaType; +import org.springframework.http.ResponseEntity; +import org.springframework.mock.web.MockMultipartFile; + +import stirling.software.SPDF.config.EndpointConfiguration; +import stirling.software.SPDF.model.api.converters.ConvertCbrToPdfRequest; +import stirling.software.SPDF.model.api.converters.ConvertCbzToPdfRequest; +import stirling.software.SPDF.model.api.converters.ConvertPdfToCbrRequest; +import stirling.software.SPDF.model.api.converters.ConvertPdfToCbzRequest; +import stirling.software.SPDF.model.api.converters.ConvertToImageRequest; +import stirling.software.common.service.CustomPDFDocumentFactory; +import stirling.software.common.util.CbrUtils; +import stirling.software.common.util.CbzUtils; +import stirling.software.common.util.CheckProgramInstall; +import stirling.software.common.util.GeneralUtils; +import stirling.software.common.util.PdfToCbrUtils; +import stirling.software.common.util.PdfToCbzUtils; +import stirling.software.common.util.PdfUtils; +import stirling.software.common.util.TempFile; +import stirling.software.common.util.TempFileManager; +import stirling.software.common.util.WebResponseUtils; + +/** + * Gap coverage for {@link ConvertImgPDFController}, exercising the comic-book and image converters + * left untested by ConvertImgPDFControllerTest. External binaries (Ghostscript, Python, RAR) are + * never invoked: the utility boundaries are mocked statically. + */ +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class ConvertImgPDFControllerGapTest { + + @Mock private CustomPDFDocumentFactory pdfDocumentFactory; + @Mock private TempFileManager tempFileManager; + @Mock private EndpointConfiguration endpointConfiguration; + + @InjectMocks private ConvertImgPDFController controller; + + /** Builds a tiny, valid single-page A4 PDF as bytes. */ + private static byte[] tinyPdfBytes(int pages) throws IOException { + try (PDDocument doc = new PDDocument(); + java.io.ByteArrayOutputStream baos = new java.io.ByteArrayOutputStream()) { + for (int i = 0; i < pages; i++) { + doc.addPage(new PDPage(PDRectangle.A4)); + } + doc.save(baos); + return baos.toByteArray(); + } + } + + /** Builds a fresh in-memory PDDocument (the factory must hand back a real, open document). */ + private static PDDocument tinyDocument(int pages) { + PDDocument doc = new PDDocument(); + for (int i = 0; i < pages; i++) { + doc.addPage(new PDPage(PDRectangle.A4)); + } + return doc; + } + + private static MockMultipartFile pdfFile(String name, byte[] bytes) { + return new MockMultipartFile("fileInput", name, "application/pdf", bytes); + } + + @Nested + @DisplayName("convertCbzToPdf") + class ConvertCbzToPdf { + + @Test + @DisplayName("disables ebook optimization when Ghostscript is not enabled") + void disablesOptimizationWhenGhostscriptMissing() throws Exception { + MockMultipartFile file = + new MockMultipartFile( + "fileInput", "book.cbz", "application/zip", new byte[] {1}); + ConvertCbzToPdfRequest request = new ConvertCbzToPdfRequest(); + request.setFileInput(file); + request.setOptimizeForEbook(true); + + Mockito.when(endpointConfiguration.isGroupEnabled("Ghostscript")).thenReturn(false); + + TempFile tempFile = Mockito.mock(TempFile.class); + @SuppressWarnings("unchecked") + ResponseEntity expected = Mockito.mock(ResponseEntity.class); + + try (MockedStatic cbz = Mockito.mockStatic(CbzUtils.class); + MockedStatic gu = Mockito.mockStatic(GeneralUtils.class); + MockedStatic wr = + Mockito.mockStatic(WebResponseUtils.class)) { + + cbz.when( + () -> + CbzUtils.convertCbzToPdf( + eq(file), + eq(pdfDocumentFactory), + eq(tempFileManager), + eq(false))) + .thenReturn(tempFile); + gu.when(() -> GeneralUtils.generateFilename("book", "_converted.pdf")) + .thenReturn("book_converted.pdf"); + wr.when(() -> WebResponseUtils.pdfFileToWebResponse(tempFile, "book_converted.pdf")) + .thenReturn(expected); + + ResponseEntity response = controller.convertCbzToPdf(request); + + assertSame(expected, response); + cbz.verify( + () -> + CbzUtils.convertCbzToPdf( + eq(file), + eq(pdfDocumentFactory), + eq(tempFileManager), + eq(false))); + } + } + + @Test + @DisplayName("keeps ebook optimization when Ghostscript is enabled") + void keepsOptimizationWhenGhostscriptEnabled() throws Exception { + MockMultipartFile file = + new MockMultipartFile( + "fileInput", "book.cbz", "application/zip", new byte[] {1}); + ConvertCbzToPdfRequest request = new ConvertCbzToPdfRequest(); + request.setFileInput(file); + request.setOptimizeForEbook(true); + + Mockito.when(endpointConfiguration.isGroupEnabled("Ghostscript")).thenReturn(true); + + TempFile tempFile = Mockito.mock(TempFile.class); + @SuppressWarnings("unchecked") + ResponseEntity expected = Mockito.mock(ResponseEntity.class); + + try (MockedStatic cbz = Mockito.mockStatic(CbzUtils.class); + MockedStatic gu = Mockito.mockStatic(GeneralUtils.class); + MockedStatic wr = + Mockito.mockStatic(WebResponseUtils.class)) { + + cbz.when( + () -> + CbzUtils.convertCbzToPdf( + eq(file), + eq(pdfDocumentFactory), + eq(tempFileManager), + eq(true))) + .thenReturn(tempFile); + gu.when(() -> GeneralUtils.generateFilename("book", "_converted.pdf")) + .thenReturn("book_converted.pdf"); + wr.when(() -> WebResponseUtils.pdfFileToWebResponse(tempFile, "book_converted.pdf")) + .thenReturn(expected); + + ResponseEntity response = controller.convertCbzToPdf(request); + + assertSame(expected, response); + cbz.verify( + () -> + CbzUtils.convertCbzToPdf( + eq(file), + eq(pdfDocumentFactory), + eq(tempFileManager), + eq(true))); + } + } + + @Test + @DisplayName("falls back to the default comic name when the original filename is null") + void usesDefaultNameWhenFilenameNull() throws Exception { + MockMultipartFile file = + new MockMultipartFile("fileInput", null, "application/zip", new byte[] {1}); + ConvertCbzToPdfRequest request = new ConvertCbzToPdfRequest(); + request.setFileInput(file); + request.setOptimizeForEbook(false); + + TempFile tempFile = Mockito.mock(TempFile.class); + @SuppressWarnings("unchecked") + ResponseEntity expected = Mockito.mock(ResponseEntity.class); + + try (MockedStatic cbz = Mockito.mockStatic(CbzUtils.class); + MockedStatic gu = Mockito.mockStatic(GeneralUtils.class); + MockedStatic wr = + Mockito.mockStatic(WebResponseUtils.class)) { + + cbz.when(() -> CbzUtils.convertCbzToPdf(any(), any(), any(), anyBoolean())) + .thenReturn(tempFile); + gu.when(() -> GeneralUtils.generateFilename("comic", "_converted.pdf")) + .thenReturn("comic_converted.pdf"); + wr.when( + () -> + WebResponseUtils.pdfFileToWebResponse( + tempFile, "comic_converted.pdf")) + .thenReturn(expected); + + ResponseEntity response = controller.convertCbzToPdf(request); + + assertSame(expected, response); + // Default comic name is resolved when the upload carries no filename. + gu.verify(() -> GeneralUtils.generateFilename("comic", "_converted.pdf")); + } + } + } + + @Nested + @DisplayName("convertCbrToPdf") + class ConvertCbrToPdf { + + @Test + @DisplayName("disables ebook optimization when Ghostscript is not enabled") + void disablesOptimizationWhenGhostscriptMissing() throws Exception { + MockMultipartFile file = + new MockMultipartFile( + "fileInput", "book.cbr", "application/x-rar", new byte[] {1}); + ConvertCbrToPdfRequest request = new ConvertCbrToPdfRequest(); + request.setFileInput(file); + request.setOptimizeForEbook(true); + + Mockito.when(endpointConfiguration.isGroupEnabled("Ghostscript")).thenReturn(false); + + byte[] pdfBytes = "converted-pdf".getBytes(); + ResponseEntity expected = ResponseEntity.ok(pdfBytes); + + try (MockedStatic cbr = Mockito.mockStatic(CbrUtils.class); + MockedStatic gu = Mockito.mockStatic(GeneralUtils.class); + MockedStatic wr = + Mockito.mockStatic(WebResponseUtils.class)) { + + cbr.when( + () -> + CbrUtils.convertCbrToPdf( + eq(file), + eq(pdfDocumentFactory), + eq(tempFileManager), + eq(false))) + .thenReturn(pdfBytes); + gu.when(() -> GeneralUtils.generateFilename("book", "_converted.pdf")) + .thenReturn("book_converted.pdf"); + wr.when(() -> WebResponseUtils.bytesToWebResponse(pdfBytes, "book_converted.pdf")) + .thenReturn(expected); + + ResponseEntity response = controller.convertCbrToPdf(request); + + assertSame(expected, response); + cbr.verify( + () -> + CbrUtils.convertCbrToPdf( + eq(file), + eq(pdfDocumentFactory), + eq(tempFileManager), + eq(false))); + } + } + + @Test + @DisplayName("keeps ebook optimization when Ghostscript is enabled") + void keepsOptimizationWhenGhostscriptEnabled() throws Exception { + MockMultipartFile file = + new MockMultipartFile( + "fileInput", "story.cbr", "application/x-rar", new byte[] {1}); + ConvertCbrToPdfRequest request = new ConvertCbrToPdfRequest(); + request.setFileInput(file); + request.setOptimizeForEbook(true); + + Mockito.when(endpointConfiguration.isGroupEnabled("Ghostscript")).thenReturn(true); + + byte[] pdfBytes = "converted-pdf".getBytes(); + ResponseEntity expected = ResponseEntity.ok(pdfBytes); + + try (MockedStatic cbr = Mockito.mockStatic(CbrUtils.class); + MockedStatic gu = Mockito.mockStatic(GeneralUtils.class); + MockedStatic wr = + Mockito.mockStatic(WebResponseUtils.class)) { + + cbr.when( + () -> + CbrUtils.convertCbrToPdf( + eq(file), + eq(pdfDocumentFactory), + eq(tempFileManager), + eq(true))) + .thenReturn(pdfBytes); + gu.when(() -> GeneralUtils.generateFilename("story", "_converted.pdf")) + .thenReturn("story_converted.pdf"); + wr.when(() -> WebResponseUtils.bytesToWebResponse(pdfBytes, "story_converted.pdf")) + .thenReturn(expected); + + ResponseEntity response = controller.convertCbrToPdf(request); + + assertSame(expected, response); + } + } + } + + @Nested + @DisplayName("convertPdfToCbz") + class ConvertPdfToCbz { + + @Test + @DisplayName("passes the requested DPI through and returns a zip response") + void passesThroughRequestedDpi() throws Exception { + MockMultipartFile file = pdfFile("doc.pdf", tinyPdfBytes(1)); + ConvertPdfToCbzRequest request = new ConvertPdfToCbzRequest(); + request.setFileInput(file); + request.setDpi(200); + + TempFile cbzFile = Mockito.mock(TempFile.class); + @SuppressWarnings("unchecked") + ResponseEntity expected = Mockito.mock(ResponseEntity.class); + + try (MockedStatic p2c = Mockito.mockStatic(PdfToCbzUtils.class); + MockedStatic gu = Mockito.mockStatic(GeneralUtils.class); + MockedStatic wr = + Mockito.mockStatic(WebResponseUtils.class)) { + + p2c.when( + () -> + PdfToCbzUtils.convertPdfToCbz( + eq(file), + eq(200), + eq(pdfDocumentFactory), + eq(tempFileManager))) + .thenReturn(cbzFile); + gu.when(() -> GeneralUtils.generateFilename("doc", "_converted.cbz")) + .thenReturn("doc_converted.cbz"); + wr.when(() -> WebResponseUtils.zipFileToWebResponse(cbzFile, "doc_converted.cbz")) + .thenReturn(expected); + + ResponseEntity response = controller.convertPdfToCbz(request); + + assertSame(expected, response); + } + } + + @Test + @DisplayName("defaults DPI to 300 when a non-positive value is supplied") + void defaultsDpiWhenNonPositive() throws Exception { + MockMultipartFile file = pdfFile("doc.pdf", tinyPdfBytes(1)); + ConvertPdfToCbzRequest request = new ConvertPdfToCbzRequest(); + request.setFileInput(file); + request.setDpi(0); + + TempFile cbzFile = Mockito.mock(TempFile.class); + @SuppressWarnings("unchecked") + ResponseEntity expected = Mockito.mock(ResponseEntity.class); + + try (MockedStatic p2c = Mockito.mockStatic(PdfToCbzUtils.class); + MockedStatic gu = Mockito.mockStatic(GeneralUtils.class); + MockedStatic wr = + Mockito.mockStatic(WebResponseUtils.class)) { + + p2c.when( + () -> + PdfToCbzUtils.convertPdfToCbz( + eq(file), + eq(300), + eq(pdfDocumentFactory), + eq(tempFileManager))) + .thenReturn(cbzFile); + gu.when(() -> GeneralUtils.generateFilename("doc", "_converted.cbz")) + .thenReturn("doc_converted.cbz"); + wr.when(() -> WebResponseUtils.zipFileToWebResponse(cbzFile, "doc_converted.cbz")) + .thenReturn(expected); + + ResponseEntity response = controller.convertPdfToCbz(request); + + assertSame(expected, response); + // Negative/zero DPI is replaced by the 300 default before delegating. + p2c.verify( + () -> + PdfToCbzUtils.convertPdfToCbz( + eq(file), + eq(300), + eq(pdfDocumentFactory), + eq(tempFileManager))); + } + } + } + + @Nested + @DisplayName("convertPdfToCbr") + class ConvertPdfToCbr { + + @Test + @DisplayName("passes the requested DPI and wraps bytes as an octet-stream response") + void passesThroughRequestedDpi() throws Exception { + MockMultipartFile file = pdfFile("doc.pdf", tinyPdfBytes(1)); + ConvertPdfToCbrRequest request = new ConvertPdfToCbrRequest(); + request.setFileInput(file); + request.setDpi(150); + + byte[] cbrBytes = "cbr-archive".getBytes(); + ResponseEntity expected = ResponseEntity.ok(cbrBytes); + + try (MockedStatic p2c = Mockito.mockStatic(PdfToCbrUtils.class); + MockedStatic gu = Mockito.mockStatic(GeneralUtils.class); + MockedStatic wr = + Mockito.mockStatic(WebResponseUtils.class)) { + + p2c.when( + () -> + PdfToCbrUtils.convertPdfToCbr( + eq(file), eq(150), eq(pdfDocumentFactory))) + .thenReturn(cbrBytes); + gu.when(() -> GeneralUtils.generateFilename("doc", "_converted.cbr")) + .thenReturn("doc_converted.cbr"); + wr.when( + () -> + WebResponseUtils.bytesToWebResponse( + eq(cbrBytes), + eq("doc_converted.cbr"), + eq(MediaType.APPLICATION_OCTET_STREAM))) + .thenReturn(expected); + + ResponseEntity response = controller.convertPdfToCbr(request); + + assertSame(expected, response); + } + } + + @Test + @DisplayName("defaults DPI to 300 when a non-positive value is supplied") + void defaultsDpiWhenNonPositive() throws Exception { + MockMultipartFile file = pdfFile("doc.pdf", tinyPdfBytes(1)); + ConvertPdfToCbrRequest request = new ConvertPdfToCbrRequest(); + request.setFileInput(file); + request.setDpi(-5); + + byte[] cbrBytes = "cbr-archive".getBytes(); + ResponseEntity expected = ResponseEntity.ok(cbrBytes); + + try (MockedStatic p2c = Mockito.mockStatic(PdfToCbrUtils.class); + MockedStatic gu = Mockito.mockStatic(GeneralUtils.class); + MockedStatic wr = + Mockito.mockStatic(WebResponseUtils.class)) { + + p2c.when( + () -> + PdfToCbrUtils.convertPdfToCbr( + eq(file), eq(300), eq(pdfDocumentFactory))) + .thenReturn(cbrBytes); + gu.when(() -> GeneralUtils.generateFilename("doc", "_converted.cbr")) + .thenReturn("doc_converted.cbr"); + wr.when( + () -> + WebResponseUtils.bytesToWebResponse( + eq(cbrBytes), + eq("doc_converted.cbr"), + eq(MediaType.APPLICATION_OCTET_STREAM))) + .thenReturn(expected); + + ResponseEntity response = controller.convertPdfToCbr(request); + + assertSame(expected, response); + p2c.verify( + () -> + PdfToCbrUtils.convertPdfToCbr( + eq(file), eq(300), eq(pdfDocumentFactory))); + } + } + } + + @Nested + @DisplayName("convertToImage") + class ConvertToImage { + + private MockMultipartFile imagePdf(byte[] bytes) { + return pdfFile("source.pdf", bytes); + } + + @Test + @DisplayName("single-image PNG path returns the rendered bytes") + void singleImagePng() throws Exception { + byte[] pdfBytes = tinyPdfBytes(1); + MockMultipartFile file = imagePdf(pdfBytes); + + ConvertToImageRequest request = new ConvertToImageRequest(); + request.setFileInput(file); + request.setImageFormat("png"); + request.setSingleOrMultiple("single"); + request.setColorType("color"); + request.setDpi(72); + request.setPageNumbers("all"); + request.setIncludeAnnotations(false); + + // rearrangePdfPages loads a real document; convertFromPdf is the boundary we stub. + Mockito.when(pdfDocumentFactory.load(any(MockMultipartFile.class))) + .thenReturn(tinyDocument(1)); + + byte[] imageBytes = "png-image".getBytes(); + ResponseEntity expected = ResponseEntity.ok(imageBytes); + + try (MockedStatic pu = Mockito.mockStatic(PdfUtils.class); + MockedStatic wr = + Mockito.mockStatic(WebResponseUtils.class)) { + + pu.when( + () -> + PdfUtils.convertFromPdf( + eq(pdfDocumentFactory), + any(byte[].class), + eq("PNG"), + eq(ImageType.RGB), + eq(true), + eq(72), + any(String.class), + eq(false))) + .thenReturn(imageBytes); + wr.when( + () -> + WebResponseUtils.bytesToWebResponse( + eq(imageBytes), + any(String.class), + any(MediaType.class))) + .thenReturn(expected); + + ResponseEntity response = controller.convertToImage(request); + + assertSame(expected, response); + } + } + + @Test + @DisplayName("multiple-image path zips the rendered output with octet-stream type") + void multipleImagesZip() throws Exception { + byte[] pdfBytes = tinyPdfBytes(2); + MockMultipartFile file = imagePdf(pdfBytes); + + ConvertToImageRequest request = new ConvertToImageRequest(); + request.setFileInput(file); + request.setImageFormat("jpg"); + request.setSingleOrMultiple("multiple"); + request.setColorType("greyscale"); + request.setDpi(72); + request.setPageNumbers("all"); + request.setIncludeAnnotations(true); + + Mockito.when(pdfDocumentFactory.load(any(MockMultipartFile.class))) + .thenReturn(tinyDocument(2)); + + byte[] zipBytes = "zip-bytes".getBytes(); + ResponseEntity expected = ResponseEntity.ok(zipBytes); + + try (MockedStatic pu = Mockito.mockStatic(PdfUtils.class); + MockedStatic wr = + Mockito.mockStatic(WebResponseUtils.class)) { + + // greyscale -> ImageType.GRAY, multiple -> singleImage=false + pu.when( + () -> + PdfUtils.convertFromPdf( + eq(pdfDocumentFactory), + any(byte[].class), + eq("JPG"), + eq(ImageType.GRAY), + eq(false), + eq(72), + any(String.class), + eq(true))) + .thenReturn(zipBytes); + wr.when( + () -> + WebResponseUtils.bytesToWebResponse( + eq(zipBytes), + any(String.class), + eq(MediaType.APPLICATION_OCTET_STREAM))) + .thenReturn(expected); + + ResponseEntity response = controller.convertToImage(request); + + assertSame(expected, response); + } + } + + @Test + @DisplayName("blackwhite color type maps to BINARY image type") + void blackwhiteMapsToBinary() throws Exception { + byte[] pdfBytes = tinyPdfBytes(1); + MockMultipartFile file = imagePdf(pdfBytes); + + ConvertToImageRequest request = new ConvertToImageRequest(); + request.setFileInput(file); + request.setImageFormat("png"); + request.setSingleOrMultiple("single"); + request.setColorType("blackwhite"); + request.setDpi(72); + request.setPageNumbers("all"); + request.setIncludeAnnotations(false); + + Mockito.when(pdfDocumentFactory.load(any(MockMultipartFile.class))) + .thenReturn(tinyDocument(1)); + + byte[] imageBytes = "bw-image".getBytes(); + ResponseEntity expected = ResponseEntity.ok(imageBytes); + + try (MockedStatic pu = Mockito.mockStatic(PdfUtils.class); + MockedStatic wr = + Mockito.mockStatic(WebResponseUtils.class)) { + + pu.when( + () -> + PdfUtils.convertFromPdf( + eq(pdfDocumentFactory), + any(byte[].class), + eq("PNG"), + eq(ImageType.BINARY), + eq(true), + eq(72), + any(String.class), + eq(false))) + .thenReturn(imageBytes); + wr.when( + () -> + WebResponseUtils.bytesToWebResponse( + eq(imageBytes), + any(String.class), + any(MediaType.class))) + .thenReturn(expected); + + ResponseEntity response = controller.convertToImage(request); + + assertSame(expected, response); + } + } + + @Test + @DisplayName("null page numbers fall back to all pages") + void nullPageNumbersDefaultsToAll() throws Exception { + byte[] pdfBytes = tinyPdfBytes(1); + MockMultipartFile file = imagePdf(pdfBytes); + + ConvertToImageRequest request = new ConvertToImageRequest(); + request.setFileInput(file); + request.setImageFormat("png"); + request.setSingleOrMultiple("single"); + request.setColorType("color"); + request.setDpi(72); + request.setPageNumbers(null); + request.setIncludeAnnotations(null); + + Mockito.when(pdfDocumentFactory.load(any(MockMultipartFile.class))) + .thenReturn(tinyDocument(1)); + + byte[] imageBytes = "png-image".getBytes(); + ResponseEntity expected = ResponseEntity.ok(imageBytes); + + try (MockedStatic pu = Mockito.mockStatic(PdfUtils.class); + MockedStatic wr = + Mockito.mockStatic(WebResponseUtils.class)) { + + // includeAnnotations null -> false + pu.when( + () -> + PdfUtils.convertFromPdf( + eq(pdfDocumentFactory), + any(byte[].class), + eq("PNG"), + eq(ImageType.RGB), + eq(true), + eq(72), + any(String.class), + eq(false))) + .thenReturn(imageBytes); + wr.when( + () -> + WebResponseUtils.bytesToWebResponse( + eq(imageBytes), + any(String.class), + any(MediaType.class))) + .thenReturn(expected); + + ResponseEntity response = controller.convertToImage(request); + + assertSame(expected, response); + } + } + + @Test + @DisplayName("webp requested without Python throws the python-required IOException") + void webpWithoutPythonThrows() throws Exception { + byte[] pdfBytes = tinyPdfBytes(1); + MockMultipartFile file = imagePdf(pdfBytes); + + ConvertToImageRequest request = new ConvertToImageRequest(); + request.setFileInput(file); + request.setImageFormat("webp"); + request.setSingleOrMultiple("single"); + request.setColorType("color"); + request.setDpi(72); + request.setPageNumbers("all"); + request.setIncludeAnnotations(false); + + Mockito.when(pdfDocumentFactory.load(any(MockMultipartFile.class))) + .thenReturn(tinyDocument(1)); + + try (MockedStatic pu = Mockito.mockStatic(PdfUtils.class); + MockedStatic cpi = + Mockito.mockStatic(CheckProgramInstall.class)) { + + // webp renders to PNG first, then requires Python for the final conversion. + pu.when( + () -> + PdfUtils.convertFromPdf( + eq(pdfDocumentFactory), + any(byte[].class), + eq("png"), + any(ImageType.class), + anyBoolean(), + anyInt(), + any(String.class), + anyBoolean())) + .thenReturn("png-image".getBytes()); + cpi.when(CheckProgramInstall::isPythonAvailable).thenReturn(false); + + assertThrows(IOException.class, () -> controller.convertToImage(request)); + } + } + } + + @Nested + @DisplayName("createConvertedFilename (via the comic converters)") + class CreateConvertedFilename { + + @Test + @DisplayName("strips only the trailing extension from a multi-dot filename") + void stripsTrailingExtensionOnly() throws Exception { + MockMultipartFile file = + new MockMultipartFile( + "fileInput", "my.archive.cbz", "application/zip", new byte[] {1}); + ConvertCbzToPdfRequest request = new ConvertCbzToPdfRequest(); + request.setFileInput(file); + request.setOptimizeForEbook(false); + + TempFile tempFile = Mockito.mock(TempFile.class); + @SuppressWarnings("unchecked") + ResponseEntity expected = Mockito.mock(ResponseEntity.class); + + try (MockedStatic cbz = Mockito.mockStatic(CbzUtils.class); + MockedStatic gu = Mockito.mockStatic(GeneralUtils.class); + MockedStatic wr = + Mockito.mockStatic(WebResponseUtils.class)) { + + cbz.when(() -> CbzUtils.convertCbzToPdf(any(), any(), any(), anyBoolean())) + .thenReturn(tempFile); + gu.when(() -> GeneralUtils.generateFilename("my.archive", "_converted.pdf")) + .thenReturn("my.archive_converted.pdf"); + wr.when( + () -> + WebResponseUtils.pdfFileToWebResponse( + tempFile, "my.archive_converted.pdf")) + .thenReturn(expected); + + ResponseEntity response = controller.convertCbzToPdf(request); + + assertSame(expected, response); + // Only the final ".cbz" is removed, the inner dot is preserved. + gu.verify(() -> GeneralUtils.generateFilename("my.archive", "_converted.pdf")); + } + } + + @Test + @DisplayName("uses the default comic name when stripping leaves a blank base name") + void fallsBackToComicWhenBaseNameBlank() throws Exception { + // ".cbz" strips to an empty base, which the controller replaces with "comic". + MockMultipartFile file = + new MockMultipartFile("fileInput", ".cbz", "application/zip", new byte[] {1}); + ConvertCbzToPdfRequest request = new ConvertCbzToPdfRequest(); + request.setFileInput(file); + request.setOptimizeForEbook(false); + + TempFile tempFile = Mockito.mock(TempFile.class); + @SuppressWarnings("unchecked") + ResponseEntity expected = Mockito.mock(ResponseEntity.class); + + try (MockedStatic cbz = Mockito.mockStatic(CbzUtils.class); + MockedStatic gu = Mockito.mockStatic(GeneralUtils.class); + MockedStatic wr = + Mockito.mockStatic(WebResponseUtils.class)) { + + cbz.when(() -> CbzUtils.convertCbzToPdf(any(), any(), any(), anyBoolean())) + .thenReturn(tempFile); + gu.when(() -> GeneralUtils.generateFilename("comic", "_converted.pdf")) + .thenReturn("comic_converted.pdf"); + wr.when( + () -> + WebResponseUtils.pdfFileToWebResponse( + tempFile, "comic_converted.pdf")) + .thenReturn(expected); + + ResponseEntity response = controller.convertCbzToPdf(request); + + assertSame(expected, response); + gu.verify(() -> GeneralUtils.generateFilename("comic", "_converted.pdf")); + } + } + } + + @Test + @DisplayName("controller is constructed with its injected collaborators") + void controllerIsConstructed() { + assertNotNull(controller); + assertEquals(0, 0); + } +} diff --git a/app/core/src/test/java/stirling/software/SPDF/controller/api/converters/ConvertPDFToPDFAGapTest.java b/app/core/src/test/java/stirling/software/SPDF/controller/api/converters/ConvertPDFToPDFAGapTest.java new file mode 100644 index 0000000000..b64776665b --- /dev/null +++ b/app/core/src/test/java/stirling/software/SPDF/controller/api/converters/ConvertPDFToPDFAGapTest.java @@ -0,0 +1,869 @@ +package stirling.software.SPDF.controller.api.converters; + +import static org.assertj.core.api.Assertions.*; +import static org.junit.jupiter.api.Assertions.*; +import static org.mockito.Mockito.*; + +import java.io.IOException; +import java.lang.reflect.InvocationTargetException; +import java.lang.reflect.Method; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.ArrayList; +import java.util.Collections; +import java.util.List; +import java.util.Set; + +import org.apache.pdfbox.cos.COSArray; +import org.apache.pdfbox.cos.COSDictionary; +import org.apache.pdfbox.cos.COSName; +import org.apache.pdfbox.pdmodel.PDDocument; +import org.apache.pdfbox.pdmodel.PDDocumentCatalog; +import org.apache.pdfbox.pdmodel.PDDocumentInformation; +import org.apache.pdfbox.pdmodel.PDPage; +import org.apache.pdfbox.pdmodel.PDPageContentStream; +import org.apache.pdfbox.pdmodel.PDResources; +import org.apache.pdfbox.pdmodel.common.PDRectangle; +import org.apache.pdfbox.pdmodel.font.PDType1Font; +import org.apache.pdfbox.pdmodel.font.Standard14Fonts; +import org.apache.pdfbox.pdmodel.interactive.annotation.PDAnnotation; +import org.apache.pdfbox.pdmodel.interactive.annotation.PDAnnotationLink; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.junit.jupiter.api.io.TempDir; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.http.HttpStatus; +import org.springframework.http.MediaType; +import org.springframework.mock.web.MockMultipartFile; +import org.springframework.web.server.ResponseStatusException; + +import stirling.software.SPDF.model.api.converters.PdfToPdfARequest; +import stirling.software.SPDF.model.api.security.PDFVerificationResult; +import stirling.software.SPDF.service.VeraPDFService; +import stirling.software.common.configuration.RuntimePathConfig; +import stirling.software.common.util.TempFileManager; + +/** + * Gap-filling unit tests for {@link ConvertPDFToPDFA}. Focuses on validation/option/error branches + * and the small pure helpers, mocking the collaborators so that no external binary (ghostscript, + * qpdf, libreoffice) or network is invoked. + */ +@DisplayName("ConvertPDFToPDFA gap tests") +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class ConvertPDFToPDFAGapTest { + + @TempDir Path tempDir; + + @Mock private RuntimePathConfig runtimePathConfig; + @Mock private VeraPDFService veraPDFService; + @Mock private TempFileManager tempFileManager; + + private ConvertPDFToPDFA newController() { + return new ConvertPDFToPDFA(runtimePathConfig, veraPDFService, tempFileManager); + } + + // ---- reflection helpers ---------------------------------------------------------------- + + @SuppressWarnings("unchecked") + private static T invokeStatic(String methodName, Object... args) throws Exception { + Method method = findMethod(methodName, args.length); + method.setAccessible(true); + try { + return (T) method.invoke(null, args); + } catch (InvocationTargetException e) { + throw unwrap(e); + } + } + + @SuppressWarnings("unchecked") + private static T invokeInstance(Object target, String methodName, Object... args) + throws Exception { + Method method = findMethod(methodName, args.length); + method.setAccessible(true); + try { + return (T) method.invoke(target, args); + } catch (InvocationTargetException e) { + throw unwrap(e); + } + } + + private static Method findMethod(String methodName, int argCount) { + for (Method method : ConvertPDFToPDFA.class.getDeclaredMethods()) { + if (method.getName().equals(methodName) && method.getParameterCount() == argCount) { + return method; + } + } + throw new IllegalStateException( + "No method named " + methodName + " with " + argCount + " params"); + } + + private static Exception unwrap(InvocationTargetException e) { + Throwable cause = e.getCause(); + if (cause instanceof Exception ex) { + return ex; + } + return new RuntimeException(cause); + } + + // ---- pdf builders ---------------------------------------------------------------------- + + private PDDocument simplePdf() throws IOException { + PDDocument document = new PDDocument(); + PDPage page = new PDPage(PDRectangle.A4); + document.addPage(page); + try (PDPageContentStream cs = new PDPageContentStream(document, page)) { + cs.beginText(); + cs.setFont(new PDType1Font(Standard14Fonts.FontName.HELVETICA), 12); + cs.newLineAtOffset(100, 700); + cs.showText("hello world"); + cs.endText(); + } + return document; + } + + private byte[] simplePdfBytes() throws IOException { + try (PDDocument document = simplePdf()) { + java.io.ByteArrayOutputStream baos = new java.io.ByteArrayOutputStream(); + document.save(baos); + return baos.toByteArray(); + } + } + + // ======================================================================================= + @Nested + @DisplayName("PdfaProfile.fromRequest token resolution") + class PdfaProfileResolution { + + // invokes the enum's static fromRequest via reflection, then reads getPart()/displayName + private Object resolveProfile(String token) throws Exception { + Class enumClass = null; + for (Class inner : ConvertPDFToPDFA.class.getDeclaredClasses()) { + if (inner.getSimpleName().equals("PdfaProfile")) { + enumClass = inner; + } + } + assertThat(enumClass).isNotNull(); + Method m = enumClass.getDeclaredMethod("fromRequest", String.class); + m.setAccessible(true); + return m.invoke(null, token); + } + + private int partOf(Object profile) throws Exception { + Method getPart = profile.getClass().getDeclaredMethod("getPart"); + getPart.setAccessible(true); + return (int) getPart.invoke(profile); + } + + private String suffixOf(Object profile) throws Exception { + Method m = profile.getClass().getDeclaredMethod("outputSuffix"); + m.setAccessible(true); + return (String) m.invoke(profile); + } + + @Test + @DisplayName("null token defaults to PDF/A-2b") + void nullDefaultsToPdfA2() throws Exception { + Object profile = resolveProfile(null); + assertThat(partOf(profile)).isEqualTo(2); + assertThat(suffixOf(profile)).isEqualTo("_PDFA-2b.pdf"); + } + + @Test + @DisplayName("'pdfa-1' resolves to PDF/A-1b") + void pdfa1Resolves() throws Exception { + assertThat(partOf(resolveProfile("pdfa-1"))).isEqualTo(1); + assertThat(suffixOf(resolveProfile("pdfa-1"))).isEqualTo("_PDFA-1b.pdf"); + } + + @Test + @DisplayName("'pdfa' and 'pdfa-2b' resolve to PDF/A-2b") + void pdfa2Resolves() throws Exception { + assertThat(partOf(resolveProfile("pdfa"))).isEqualTo(2); + assertThat(partOf(resolveProfile("pdfa-2b"))).isEqualTo(2); + } + + @Test + @DisplayName("'pdfa-3' and 'pdfa-3b' resolve to PDF/A-3b") + void pdfa3Resolves() throws Exception { + assertThat(partOf(resolveProfile("pdfa-3"))).isEqualTo(3); + assertThat(partOf(resolveProfile("pdfa-3b"))).isEqualTo(3); + assertThat(suffixOf(resolveProfile("pdfa-3"))).isEqualTo("_PDFA-3b.pdf"); + } + + @Test + @DisplayName("token is trimmed and case-insensitive") + void caseInsensitiveAndTrimmed() throws Exception { + assertThat(partOf(resolveProfile(" PDFA-1 "))).isEqualTo(1); + assertThat(partOf(resolveProfile("PDFA-3B"))).isEqualTo(3); + } + + @Test + @DisplayName("unknown token falls back to PDF/A-2b") + void unknownFallsBack() throws Exception { + assertThat(partOf(resolveProfile("not-a-real-format"))).isEqualTo(2); + } + } + + // ======================================================================================= + @Nested + @DisplayName("PdfXProfile.fromRequest token resolution") + class PdfXProfileResolution { + + private Object resolveProfile(String token) throws Exception { + Class enumClass = null; + for (Class inner : ConvertPDFToPDFA.class.getDeclaredClasses()) { + if (inner.getSimpleName().equals("PdfXProfile")) { + enumClass = inner; + } + } + assertThat(enumClass).isNotNull(); + Method m = enumClass.getDeclaredMethod("fromRequest", String.class); + m.setAccessible(true); + return m.invoke(null, token); + } + + private String suffixOf(Object profile) throws Exception { + Method m = profile.getClass().getDeclaredMethod("outputSuffix"); + m.setAccessible(true); + return (String) m.invoke(profile); + } + + @Test + @DisplayName("null token defaults to PDF/X") + void nullDefaultsToPdfX() throws Exception { + assertThat(suffixOf(resolveProfile(null))).isEqualTo("_PDFX.pdf"); + } + + @Test + @DisplayName("'pdfx' resolves to the PDF/X profile") + void pdfxResolves() throws Exception { + assertThat(suffixOf(resolveProfile("pdfx"))).isEqualTo("_PDFX.pdf"); + } + + @Test + @DisplayName("unknown token falls back to PDF/X") + void unknownFallsBack() throws Exception { + assertThat(suffixOf(resolveProfile("garbage"))).isEqualTo("_PDFX.pdf"); + } + } + + // ======================================================================================= + @Nested + @DisplayName("detectMimeTypeFromFilename") + class MimeTypeDetection { + + private String detect(String fileName) throws Exception { + return invokeInstance(newController(), "detectMimeTypeFromFilename", fileName); + } + + @Test + @DisplayName("known extensions map to their MIME type") + void knownExtensions() throws Exception { + assertThat(detect("data.xml")).isEqualTo("application/xml"); + assertThat(detect("data.json")).isEqualTo("application/json"); + assertThat(detect("notes.txt")).isEqualTo("text/plain"); + assertThat(detect("image.png")).isEqualTo("image/png"); + assertThat(detect("photo.JPEG")).isEqualTo("image/jpeg"); + assertThat(detect("archive.zip")).isEqualTo("application/zip"); + } + + @Test + @DisplayName("unknown extension yields octet-stream default") + void unknownExtension() throws Exception { + assertThat(detect("file.unknownext")).isEqualTo("application/octet-stream"); + } + + @Test + @DisplayName("null and empty file names yield octet-stream default") + void nullAndEmpty() throws Exception { + assertThat(detect(null)).isEqualTo("application/octet-stream"); + assertThat(detect("")).isEqualTo("application/octet-stream"); + } + } + + // ======================================================================================= + @Nested + @DisplayName("countGlyphs / buildStandardType1GlyphSet") + class GlyphHelpers { + + @Test + @DisplayName("countGlyphs counts forward slashes") + void countsSlashes() throws Exception { + assertThat((int) invokeStatic("countGlyphs", "/a/b/c")).isEqualTo(3); + assertThat((int) invokeStatic("countGlyphs", "no-slashes")).isEqualTo(0); + } + + @Test + @DisplayName("countGlyphs handles null and empty") + void countsNullEmpty() throws Exception { + assertThat((int) invokeStatic("countGlyphs", (Object) null)).isEqualTo(0); + assertThat((int) invokeStatic("countGlyphs", "")).isEqualTo(0); + } + + @Test + @DisplayName("standard glyph set is space-separated and contains core glyphs") + void standardGlyphSet() throws Exception { + String glyphs = invokeStatic("buildStandardType1GlyphSet"); + assertThat(glyphs).isNotBlank(); + assertThat(glyphs).contains(".notdef", "space", "A", "z", "zero", "period"); + // space-separated; no leading slash format here + assertThat(glyphs.split(" ").length).isGreaterThan(100); + } + } + + // ======================================================================================= + @Nested + @DisplayName("isType1Font / quad and rect validation") + class TypeAndGeometryHelpers { + + @Test + @DisplayName("isType1Font true for Standard14 Type1 font") + void isType1True() throws Exception { + PDType1Font font = new PDType1Font(Standard14Fonts.FontName.HELVETICA); + assertThat((boolean) invokeStatic("isType1Font", font)).isTrue(); + } + + @Test + @DisplayName("isValidQuadPoints accepts multiples of 8, rejects otherwise") + void quadValidation() throws Exception { + ConvertPDFToPDFA controller = newController(); + assertThat( + (boolean) + invokeInstance( + controller, "isValidQuadPoints", (Object) new float[8])) + .isTrue(); + assertThat( + (boolean) + invokeInstance( + controller, + "isValidQuadPoints", + (Object) new float[16])) + .isTrue(); + assertThat( + (boolean) + invokeInstance( + controller, "isValidQuadPoints", (Object) new float[5])) + .isFalse(); + assertThat((boolean) invokeInstance(controller, "isValidQuadPoints", (Object) null)) + .isFalse(); + } + + @Test + @DisplayName("isZeroSizeRect distinguishes collapsed and real rectangles") + void zeroSizeRect() throws Exception { + ConvertPDFToPDFA controller = newController(); + PDRectangle zero = new PDRectangle(10f, 10f, 0f, 0f); + PDRectangle real = new PDRectangle(0f, 0f, 100f, 50f); + assertThat((boolean) invokeInstance(controller, "isZeroSizeRect", zero)).isTrue(); + assertThat((boolean) invokeInstance(controller, "isZeroSizeRect", real)).isFalse(); + } + } + + // ======================================================================================= + @Nested + @DisplayName("findUnembeddedFontNames") + class UnembeddedFontDetection { + + @Test + @DisplayName("standard 14 fonts (not embedded) are reported as missing") + void detectsStandardFontAsUnembedded() throws Exception { + try (PDDocument document = simplePdf()) { + Set missing = invokeStatic("findUnembeddedFontNames", document); + assertThat(missing).isNotNull(); + assertThat(missing).anyMatch(name -> name.contains("Helvetica")); + } + } + + @Test + @DisplayName("page with no resources reports no missing fonts") + void pageWithoutResources() throws Exception { + try (PDDocument document = new PDDocument()) { + PDPage page = new PDPage(PDRectangle.A4); + page.setResources(new PDResources()); + document.addPage(page); + Set missing = invokeStatic("findUnembeddedFontNames", document); + assertThat(missing).isEmpty(); + } + } + } + + // ======================================================================================= + @Nested + @DisplayName("detectTransparentXObjects") + class TransparentXObjectDetection { + + @Test + @DisplayName("page with no resources returns empty set") + void noResources() throws Exception { + PDPage page = new PDPage(PDRectangle.A4); + Set result = invokeStatic("detectTransparentXObjects", page); + assertThat(result).isEmpty(); + } + + @Test + @DisplayName("page with empty resources returns empty set") + void emptyResources() throws Exception { + PDPage page = new PDPage(PDRectangle.A4); + page.setResources(new PDResources()); + Set result = invokeStatic("detectTransparentXObjects", page); + assertThat(result).isEmpty(); + } + } + + // ======================================================================================= + @Nested + @DisplayName("sanitizePdfA part-specific behaviour") + class SanitizePdfAExtras { + + @Test + @DisplayName("PDF/A-3 preserves embedded-file structures (FILESPEC type)") + void pdfA3PreservesFilespec() throws Exception { + COSDictionary dict = new COSDictionary(); + dict.setItem(COSName.TYPE, COSName.FILESPEC); + dict.setItem(COSName.EF, new COSDictionary()); + dict.setItem(COSName.URI, COSName.A); + + invokeStatic("sanitizePdfA", dict, 3); + + // For part 3, filespec dictionaries are skipped entirely, so URI survives. + assertThat(dict.containsKey(COSName.EF)).isTrue(); + assertThat(dict.containsKey(COSName.URI)).isTrue(); + } + + @Test + @DisplayName("PDF/A-3 keeps URI on a normal dictionary but still strips JavaScript") + void pdfA3KeepsUriStripsJs() throws Exception { + COSDictionary dict = new COSDictionary(); + dict.setString(COSName.URI, "http://example.com"); + dict.setString(COSName.JAVA_SCRIPT, "app.alert('x');"); + + invokeStatic("sanitizePdfA", dict, 3); + + assertThat(dict.containsKey(COSName.URI)).isTrue(); + assertThat(dict.containsKey(COSName.JAVA_SCRIPT)).isFalse(); + } + + @Test + @DisplayName("recurses into nested arrays and dictionaries") + void recursesIntoNested() throws Exception { + COSDictionary child = new COSDictionary(); + child.setString(COSName.JAVA_SCRIPT, "code"); + COSArray array = new COSArray(); + array.add(child); + COSDictionary parent = new COSDictionary(); + parent.setItem(COSName.getPDFName("Kids"), array); + + invokeStatic("sanitizePdfA", parent, 2); + + assertThat(child.containsKey(COSName.JAVA_SCRIPT)).isFalse(); + } + + @Test + @DisplayName("PDF/A-2 does NOT remove SMask/CA (those are only stripped for part 1)") + void pdfA2KeepsTransparencyEntries() throws Exception { + COSDictionary dict = new COSDictionary(); + dict.setItem(COSName.SMASK, new COSArray()); + dict.setFloat(COSName.CA, 0.5f); + + invokeStatic("sanitizePdfA", dict, 2); + + assertThat(dict.containsKey(COSName.SMASK)).isTrue(); + assertThat(dict.containsKey(COSName.CA)).isTrue(); + } + } + + // ======================================================================================= + @Nested + @DisplayName("sanitizeMetadata / removeForbiddenActions") + class MetadataAndActions { + + @Test + @DisplayName("sanitizeMetadata strips non-printable chars and sets producer") + void sanitizeMetadataCleans() throws Exception { + try (PDDocument document = simplePdf()) { + PDDocumentInformation info = new PDDocumentInformation(); + info.setCustomMetadataValue("Custom", "cleanvalue"); + document.setDocumentInformation(info); + + invokeInstance(newController(), "sanitizeMetadata", document); + + PDDocumentInformation result = document.getDocumentInformation(); + assertThat(result.getProducer()).isEqualTo("Stirling-PDF Sanitizer"); + assertThat(result.getCustomMetadataValue("Custom")).isEqualTo("cleanvalue"); + } + } + + @Test + @DisplayName("sanitizeMetadata always overwrites producer to the sanitizer marker") + void sanitizeMetadataOverwritesProducer() throws Exception { + try (PDDocument document = simplePdf()) { + PDDocumentInformation info = new PDDocumentInformation(); + info.setProducer("Some Other Producer"); + document.setDocumentInformation(info); + + invokeInstance(newController(), "sanitizeMetadata", document); + + assertThat(document.getDocumentInformation().getProducer()) + .isEqualTo("Stirling-PDF Sanitizer"); + } + } + + @Test + @DisplayName("removeForbiddenActions clears open action and JavaScript") + void removeForbiddenActions() throws Exception { + try (PDDocument document = simplePdf()) { + PDDocumentCatalog catalog = document.getDocumentCatalog(); + catalog.getCOSObject().setItem(COSName.JAVA_SCRIPT, new COSDictionary()); + + invokeInstance(newController(), "removeForbiddenActions", document); + + assertThat(catalog.getCOSObject().containsKey(COSName.JAVA_SCRIPT)).isFalse(); + assertThat(catalog.getOpenAction()).isNull(); + } + } + } + + // ======================================================================================= + @Nested + @DisplayName("fixOptionalContentGroups") + class OptionalContentGroups { + + @Test + @DisplayName("no OCProperties is a no-op (does not throw)") + void noOcProperties() throws Exception { + try (PDDocument document = simplePdf()) { + assertThatCode(() -> invokeStatic("fixOptionalContentGroups", document)) + .doesNotThrowAnyException(); + } + } + } + + // ======================================================================================= + @Nested + @DisplayName("addWhiteBackground") + class WhiteBackground { + + @Test + @DisplayName("adds a prepended content stream without changing page count") + void addsBackground() throws Exception { + try (PDDocument document = simplePdf()) { + int pagesBefore = document.getNumberOfPages(); + assertThatCode( + () -> + invokeInstance( + newController(), "addWhiteBackground", document)) + .doesNotThrowAnyException(); + assertThat(document.getNumberOfPages()).isEqualTo(pagesBefore); + // page still has content streams after prepending background + assertThat(document.getPage(0).hasContents()).isTrue(); + } + } + } + + // ======================================================================================= + @Nested + @DisplayName("ensureAnnotationAppearances") + class AnnotationAppearances { + + @Test + @DisplayName("Link annotations are skipped (kept) by appearance enforcement") + void linkAnnotationsKept() throws Exception { + try (PDDocument document = simplePdf()) { + PDPage page = document.getPage(0); + PDAnnotationLink link = new PDAnnotationLink(); + link.setRectangle(new PDRectangle(0, 0, 50, 50)); + List annotations = new ArrayList<>(); + annotations.add(link); + page.setAnnotations(annotations); + + invokeInstance(newController(), "ensureAnnotationAppearances", document); + + assertThat(page.getAnnotations()).hasSize(1); + } + } + } + + // ======================================================================================= + @Nested + @DisplayName("ensureEmbeddedFileCompliance / addICCProfileIfNotPresent") + class EmbeddedAndIcc { + + @Test + @DisplayName("ensureEmbeddedFileCompliance returns quietly with no names dictionary") + void noNamesDictionary() throws Exception { + try (PDDocument document = simplePdf()) { + assertThatCode( + () -> + invokeInstance( + newController(), + "ensureEmbeddedFileCompliance", + document)) + .doesNotThrowAnyException(); + } + } + + @Test + @DisplayName("addICCProfileIfNotPresent adds an sRGB output intent") + void addsIccOutputIntent() throws Exception { + try (PDDocument document = simplePdf()) { + assertThat(document.getDocumentCatalog().getOutputIntents()).isEmpty(); + + invokeInstance(newController(), "addICCProfileIfNotPresent", document); + + assertThat(document.getDocumentCatalog().getOutputIntents()).hasSize(1); + assertThat(document.getDocumentCatalog().getOutputIntents().get(0).getInfo()) + .contains("sRGB"); + } + } + + @Test + @DisplayName("addICCProfileIfNotPresent does not add a second intent when one exists") + void doesNotDuplicateIntent() throws Exception { + try (PDDocument document = simplePdf()) { + invokeInstance(newController(), "addICCProfileIfNotPresent", document); + invokeInstance(newController(), "addICCProfileIfNotPresent", document); + assertThat(document.getDocumentCatalog().getOutputIntents()).hasSize(1); + } + } + } + + // ======================================================================================= + @Nested + @DisplayName("performBasicPdfAValidation / buildComprehensiveValidationMessage") + class ValidationHelpers { + + @Test + @DisplayName("basic validation flags missing XMP and output intent for a plain PDF") + void basicValidationFlagsMissing() throws Exception { + Path pdf = tempDir.resolve("plain.pdf"); + Files.write(pdf, simplePdfBytes()); + + // PdfaProfile.PDF_A_2B has no preflight format -> basic validation path. + Object profile = resolvePdfaProfile("pdfa-2b"); + org.apache.pdfbox.preflight.ValidationResult result = + invokeStaticWithTypes( + "performBasicPdfAValidation", + new Class[] {Path.class, profile.getClass()}, + pdf, + profile); + + assertThat(result).isNotNull(); + assertThat(result.isValid()).isFalse(); + assertThat(result.getErrorsList()).isNotEmpty(); + } + + @Test + @DisplayName("comprehensive message summarises error count and codes") + void comprehensiveMessage() throws Exception { + Object profile = resolvePdfaProfile("pdfa-1"); + org.apache.pdfbox.preflight.ValidationResult result = + new org.apache.pdfbox.preflight.ValidationResult(false); + result.addError( + new org.apache.pdfbox.preflight.ValidationResult.ValidationError( + "CODE_A", "first problem")); + result.addError( + new org.apache.pdfbox.preflight.ValidationResult.ValidationError( + "CODE_B", "second problem")); + + String message = + invokeStaticWithTypes( + "buildComprehensiveValidationMessage", + new Class[] { + org.apache.pdfbox.preflight.ValidationResult.class, + profile.getClass() + }, + result, + profile); + + assertThat(message).contains("PDF/A-1b"); + assertThat(message).contains("2 errors"); + assertThat(message).contains("CODE_A"); + } + + // helper: resolve enum constant + call typed static method + private Object resolvePdfaProfile(String token) throws Exception { + Class enumClass = null; + for (Class inner : ConvertPDFToPDFA.class.getDeclaredClasses()) { + if (inner.getSimpleName().equals("PdfaProfile")) { + enumClass = inner; + } + } + Method m = enumClass.getDeclaredMethod("fromRequest", String.class); + m.setAccessible(true); + return m.invoke(null, token); + } + + @SuppressWarnings("unchecked") + private T invokeStaticWithTypes(String name, Class[] types, Object... args) + throws Exception { + Method m = ConvertPDFToPDFA.class.getDeclaredMethod(name, types); + m.setAccessible(true); + try { + return (T) m.invoke(null, args); + } catch (InvocationTargetException e) { + throw unwrap(e); + } + } + } + + // ======================================================================================= + @Nested + @DisplayName("verifyStrictCompliance (VeraPDFService mocked)") + class StrictCompliance { + + @Test + @DisplayName("compliant result passes without throwing") + void compliantPasses() throws Exception { + PDFVerificationResult ok = new PDFVerificationResult(); + ok.setCompliant(true); + ok.setStandard("1b"); + ok.setComplianceSummary("PDF/A-1b compliant"); + when(veraPDFService.validatePDF(any())).thenReturn(List.of(ok)); + + ConvertPDFToPDFA controller = newController(); + assertThatCode( + () -> + invokeInstance( + controller, + "verifyStrictCompliance", + (Object) "dummy".getBytes())) + .doesNotThrowAnyException(); + } + + @Test + @DisplayName("non-compliant result throws 400 ResponseStatusException with details") + void nonCompliantThrowsBadRequest() throws Exception { + PDFVerificationResult bad = new PDFVerificationResult(); + bad.setCompliant(false); + bad.setStandard("1b"); + bad.setComplianceSummary("PDF/A-1b with errors"); + when(veraPDFService.validatePDF(any())).thenReturn(List.of(bad)); + + ConvertPDFToPDFA controller = newController(); + ResponseStatusException ex = + (ResponseStatusException) + catchThrowable( + () -> + invokeInstance( + controller, + "verifyStrictCompliance", + (Object) "dummy".getBytes())); + assertThat(ex).isNotNull(); + assertThat(ex.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + assertThat(ex.getReason()).contains("PDF/A-1b with errors"); + } + + @Test + @DisplayName("empty result list is treated as non-compliant -> 400") + void emptyResultsTreatedNonCompliant() throws Exception { + when(veraPDFService.validatePDF(any())).thenReturn(Collections.emptyList()); + + ConvertPDFToPDFA controller = newController(); + ResponseStatusException ex = + (ResponseStatusException) + catchThrowable( + () -> + invokeInstance( + controller, + "verifyStrictCompliance", + (Object) "dummy".getBytes())); + assertThat(ex).isNotNull(); + assertThat(ex.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + } + + @Test + @DisplayName("service error is wrapped as 500 ResponseStatusException") + void serviceErrorWrappedAs500() throws Exception { + when(veraPDFService.validatePDF(any())).thenThrow(new IOException("boom")); + + ConvertPDFToPDFA controller = newController(); + ResponseStatusException ex = + (ResponseStatusException) + catchThrowable( + () -> + invokeInstance( + controller, + "verifyStrictCompliance", + (Object) "dummy".getBytes())); + assertThat(ex).isNotNull(); + assertThat(ex.getStatusCode()).isEqualTo(HttpStatus.INTERNAL_SERVER_ERROR); + } + } + + // ======================================================================================= + @Nested + @DisplayName("pdfToPdfA controller validation branch") + class ControllerValidation { + + @Test + @DisplayName("non-PDF content type throws PDF-required exception before any conversion") + void nonPdfContentTypeRejected() { + MockMultipartFile file = + new MockMultipartFile( + "fileInput", + "input.txt", + MediaType.TEXT_PLAIN_VALUE, + "not a pdf".getBytes()); + PdfToPdfARequest request = new PdfToPdfARequest(); + request.setFileInput(file); + request.setOutputFormat("pdfa-2b"); + + ConvertPDFToPDFA controller = newController(); + + assertThatThrownBy(() -> controller.pdfToPdfA(request)) + .isInstanceOf(IllegalArgumentException.class); + + // collaborators must not be touched on the validation-failure path + verifyNoInteractions(tempFileManager, veraPDFService); + } + + @Test + @DisplayName("null content type is also rejected as not-a-PDF") + void nullContentTypeRejected() { + MockMultipartFile file = + new MockMultipartFile("fileInput", "input.bin", null, "data".getBytes()); + PdfToPdfARequest request = new PdfToPdfARequest(); + request.setFileInput(file); + request.setOutputFormat("pdfa"); + + ConvertPDFToPDFA controller = newController(); + + assertThatThrownBy(() -> controller.pdfToPdfA(request)) + .isInstanceOf(IllegalArgumentException.class); + } + } + + // ======================================================================================= + @Nested + @DisplayName("ensureEmbeddedFilesAFRelationship / isTransparencyGroup") + class StaticEdgeCases { + + @Test + @DisplayName("ensureEmbeddedFilesAFRelationship is a no-op when no names dictionary") + void afRelationshipNoNames() throws Exception { + try (PDDocument document = simplePdf()) { + assertThatCode(() -> invokeStatic("ensureEmbeddedFilesAFRelationship", document)) + .doesNotThrowAnyException(); + } + } + + @Test + @DisplayName("isTransparencyGroup true only for /S /Transparency group dictionaries") + void transparencyGroup() throws Exception { + COSDictionary withGroup = new COSDictionary(); + COSDictionary groupDict = new COSDictionary(); + groupDict.setItem(COSName.S, COSName.TRANSPARENCY); + withGroup.setItem(COSName.GROUP, groupDict); + assertThat((boolean) invokeStatic("isTransparencyGroup", withGroup)).isTrue(); + + COSDictionary withoutGroup = new COSDictionary(); + assertThat((boolean) invokeStatic("isTransparencyGroup", withoutGroup)).isFalse(); + } + } +} diff --git a/app/core/src/test/java/stirling/software/SPDF/controller/api/converters/ConvertPdfToVideoControllerTest.java b/app/core/src/test/java/stirling/software/SPDF/controller/api/converters/ConvertPdfToVideoControllerTest.java new file mode 100644 index 0000000000..7a46611443 --- /dev/null +++ b/app/core/src/test/java/stirling/software/SPDF/controller/api/converters/ConvertPdfToVideoControllerTest.java @@ -0,0 +1,450 @@ +package stirling.software.SPDF.controller.api.converters; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.when; + +import java.awt.Color; +import java.awt.Graphics2D; +import java.awt.image.BufferedImage; +import java.io.File; +import java.io.IOException; +import java.lang.reflect.InvocationTargetException; +import java.lang.reflect.Method; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.List; + +import org.apache.pdfbox.Loader; +import org.apache.pdfbox.pdmodel.PDDocument; +import org.apache.pdfbox.pdmodel.PDPage; +import org.apache.pdfbox.pdmodel.common.PDRectangle; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.junit.jupiter.api.io.TempDir; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.http.MediaType; + +import stirling.software.common.service.CustomPDFDocumentFactory; +import stirling.software.common.util.TempFile; +import stirling.software.common.util.TempFileManager; + +/** + * Unit tests for {@link ConvertPdfToVideoController}. + * + *

The public {@code convertPdfToVideo} endpoint is commented out in production (ffmpeg disabled + * due to CVEs), so these tests exercise the remaining private helper methods directly via + * reflection. No external processes (ffmpeg) are ever spawned: {@code buildFfmpegCommand} only + * constructs the argument list, and {@code generateFrames} renders PDF pages to PNG files on disk. + */ +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class ConvertPdfToVideoControllerTest { + + @Mock private CustomPDFDocumentFactory pdfDocumentFactory; + @Mock private TempFileManager tempFileManager; + + private ConvertPdfToVideoController controller; + + @BeforeEach + void setUp() { + controller = new ConvertPdfToVideoController(pdfDocumentFactory, tempFileManager); + } + + // ---- reflection helpers ------------------------------------------------- + + private Object invokePrivate(String name, Class[] types, Object... args) throws Exception { + Method method = ConvertPdfToVideoController.class.getDeclaredMethod(name, types); + method.setAccessible(true); + try { + return method.invoke(controller, args); + } catch (InvocationTargetException e) { + Throwable cause = e.getCause(); + if (cause instanceof Exception ex) { + throw ex; + } + throw e; + } + } + + private String normalizeFormat(String requested) throws Exception { + return (String) invokePrivate("normalizeFormat", new Class[] {String.class}, requested); + } + + private MediaType getMediaType(String format) throws Exception { + return (MediaType) invokePrivate("getMediaType", new Class[] {String.class}, format); + } + + private int getMaxDpi() throws Exception { + return (int) invokePrivate("getMaxDpi", new Class[] {}); + } + + @SuppressWarnings("unchecked") + private List buildFfmpegCommand( + String format, String resolution, String frameRate, TempFile outputVideo) + throws Exception { + return (List) + invokePrivate( + "buildFfmpegCommand", + new Class[] {String.class, String.class, String.class, TempFile.class}, + format, + resolution, + frameRate, + outputVideo); + } + + private void applyWatermark(BufferedImage image, float opacity, String text) throws Exception { + invokePrivate( + "applyWatermark", + new Class[] {BufferedImage.class, float.class, String.class}, + image, + opacity, + text); + } + + private void generateFrames( + Path inputPdf, + Path outputDir, + int dpi, + float opacity, + String watermarkText, + boolean watermarkEnabled) + throws Exception { + invokePrivate( + "generateFrames", + new Class[] { + Path.class, Path.class, int.class, float.class, String.class, boolean.class + }, + inputPdf, + outputDir, + dpi, + opacity, + watermarkText, + watermarkEnabled); + } + + // ---- PDF builder helper ------------------------------------------------- + + private byte[] buildPdf(int pageCount) throws IOException { + try (PDDocument document = new PDDocument()) { + for (int i = 0; i < pageCount; i++) { + document.addPage(new PDPage(PDRectangle.A4)); + } + java.io.ByteArrayOutputStream baos = new java.io.ByteArrayOutputStream(); + document.save(baos); + return baos.toByteArray(); + } + } + + // ---- normalizeFormat ---------------------------------------------------- + + @Nested + @DisplayName("normalizeFormat") + class NormalizeFormat { + + @Test + @DisplayName("null defaults to mp4") + void nullDefaultsToMp4() throws Exception { + assertEquals("mp4", normalizeFormat(null)); + } + + @Test + @DisplayName("mp4 passes through") + void mp4PassesThrough() throws Exception { + assertEquals("mp4", normalizeFormat("mp4")); + } + + @Test + @DisplayName("webm passes through") + void webmPassesThrough() throws Exception { + assertEquals("webm", normalizeFormat("webm")); + } + + @Test + @DisplayName("uppercase is lowercased") + void uppercaseIsLowercased() throws Exception { + assertEquals("mp4", normalizeFormat("MP4")); + assertEquals("webm", normalizeFormat("WEBM")); + } + + @Test + @DisplayName("mixed case is normalized") + void mixedCaseIsNormalized() throws Exception { + assertEquals("webm", normalizeFormat("WeBm")); + } + + @Test + @DisplayName("unsupported format falls back to mp4") + void unsupportedFallsBackToMp4() throws Exception { + assertEquals("mp4", normalizeFormat("avi")); + assertEquals("mp4", normalizeFormat("gif")); + assertEquals("mp4", normalizeFormat("")); + } + } + + // ---- getMediaType ------------------------------------------------------- + + @Nested + @DisplayName("getMediaType") + class GetMediaType { + + @Test + @DisplayName("webm maps to video/webm") + void webmMapsToVideoWebm() throws Exception { + assertEquals(MediaType.valueOf("video/webm"), getMediaType("webm")); + } + + @Test + @DisplayName("mp4 maps to video/mp4") + void mp4MapsToVideoMp4() throws Exception { + assertEquals(MediaType.valueOf("video/mp4"), getMediaType("mp4")); + } + + @Test + @DisplayName("unknown format defaults to video/mp4") + void unknownDefaultsToVideoMp4() throws Exception { + assertEquals(MediaType.valueOf("video/mp4"), getMediaType("avi")); + } + } + + // ---- getMaxDpi ---------------------------------------------------------- + + @Nested + @DisplayName("getMaxDpi") + class GetMaxDpi { + + @Test + @DisplayName("returns 500 fallback when no Spring context is available") + void returnsFallbackWithoutContext() throws Exception { + // No ApplicationContext is set in this unit test, so getBean returns null and the + // method falls back to the hardcoded default of 500. + assertEquals(500, getMaxDpi()); + } + } + + // ---- buildFfmpegCommand ------------------------------------------------- + + @Nested + @DisplayName("buildFfmpegCommand") + class BuildFfmpegCommand { + + private TempFile newTempFile(File backing) throws IOException { + when(tempFileManager.createTempFile(any())).thenReturn(backing); + return new TempFile(tempFileManager, ".mp4"); + } + + @Test + @DisplayName("mp4 command includes libx264 and faststart flags") + void mp4Command(@TempDir Path dir) throws Exception { + File backing = dir.resolve("out.mp4").toFile(); + TempFile outputVideo = newTempFile(backing); + + List command = buildFfmpegCommand("mp4", "ORIGINAL", "0.333333", outputVideo); + + assertEquals("ffmpeg", command.get(0)); + assertTrue(command.contains("-y")); + assertTrue(command.contains("-framerate")); + assertTrue(command.contains("0.333333")); + assertTrue(command.contains("frame_%05d.png")); + assertTrue(command.contains("-vf")); + assertTrue(command.contains("libx264")); + assertTrue(command.contains("yuv420p")); + assertTrue(command.contains("+faststart")); + assertFalse(command.contains("libvpx-vp9")); + // Output path is always the last argument. + assertEquals(backing.getAbsolutePath(), command.get(command.size() - 1)); + } + + @Test + @DisplayName("webm command includes libvpx-vp9 and crf flags") + void webmCommand(@TempDir Path dir) throws Exception { + File backing = dir.resolve("out.webm").toFile(); + TempFile outputVideo = newTempFile(backing); + + List command = buildFfmpegCommand("webm", "720P", "0.5", outputVideo); + + assertTrue(command.contains("libvpx-vp9")); + assertTrue(command.contains("-crf")); + assertTrue(command.contains("30")); + assertFalse(command.contains("libx264")); + assertFalse(command.contains("+faststart")); + assertEquals(backing.getAbsolutePath(), command.get(command.size() - 1)); + } + + @Test + @DisplayName("framerate value is placed right after -framerate") + void framerateOrdering(@TempDir Path dir) throws Exception { + TempFile outputVideo = newTempFile(dir.resolve("o.mp4").toFile()); + + List command = buildFfmpegCommand("mp4", "ORIGINAL", "0.25", outputVideo); + + int idx = command.indexOf("-framerate"); + assertTrue(idx >= 0); + assertEquals("0.25", command.get(idx + 1)); + } + + @Test + @DisplayName("known resolution applies the matching scale filter") + void knownResolutionFilter(@TempDir Path dir) throws Exception { + TempFile outputVideo = newTempFile(dir.resolve("o.mp4").toFile()); + + List command = buildFfmpegCommand("mp4", "1080P", "0.5", outputVideo); + + int idx = command.indexOf("-vf"); + assertEquals("scale=-2:1080,setsar=1", command.get(idx + 1)); + } + + @Test + @DisplayName("unknown resolution falls back to the ORIGINAL filter") + void unknownResolutionFallsBackToOriginal(@TempDir Path dir) throws Exception { + TempFile outputVideo = newTempFile(dir.resolve("o.mp4").toFile()); + + List command = buildFfmpegCommand("mp4", "NONSENSE", "0.5", outputVideo); + + int idx = command.indexOf("-vf"); + assertEquals("scale=trunc(iw/2)*2:trunc(ih/2)*2,setsar=1", command.get(idx + 1)); + } + } + + // ---- applyWatermark ----------------------------------------------------- + + @Nested + @DisplayName("applyWatermark") + class ApplyWatermark { + + private BufferedImage solidImage(int w, int h, Color color) { + BufferedImage image = new BufferedImage(w, h, BufferedImage.TYPE_INT_RGB); + Graphics2D g = image.createGraphics(); + g.setColor(color); + g.fillRect(0, 0, w, h); + g.dispose(); + return image; + } + + @Test + @DisplayName("modifies pixels at full opacity without throwing") + void modifiesPixels() throws Exception { + // The watermark is drawn through the image centre, so a small image still gets pixels + // painted; the smaller buffer keeps the two getRGB snapshots cheap. + BufferedImage image = solidImage(100, 80, Color.RED); + int[] before = + image.getRGB( + 0, 0, image.getWidth(), image.getHeight(), null, 0, image.getWidth()); + + applyWatermark(image, 1.0f, "CONFIDENTIAL"); + + int[] after = + image.getRGB( + 0, 0, image.getWidth(), image.getHeight(), null, 0, image.getWidth()); + boolean changed = false; + for (int i = 0; i < before.length; i++) { + if (before[i] != after[i]) { + changed = true; + break; + } + } + assertTrue(changed, "watermark should have altered at least one pixel"); + } + + @Test + @DisplayName("does not throw on a small square image") + void handlesSmallImage() throws Exception { + BufferedImage image = solidImage(50, 50, Color.BLUE); + applyWatermark(image, 0.5f, "X"); + assertNotNull(image); + } + + @Test + @DisplayName("does not throw at zero opacity") + void handlesZeroOpacity() throws Exception { + BufferedImage image = solidImage(120, 80, Color.GREEN); + applyWatermark(image, 0.0f, "WM"); + assertNotNull(image); + } + } + + // ---- generateFrames ----------------------------------------------------- + + @Nested + @DisplayName("generateFrames") + class GenerateFrames { + + @Test + @DisplayName("renders one PNG frame per page") + void rendersOneFramePerPage(@TempDir Path dir) throws Exception { + Path inputPdf = dir.resolve("input.pdf"); + Files.write(inputPdf, buildPdf(3)); + Path outputDir = Files.createDirectory(dir.resolve("frames")); + + // The factory must return a fresh, real document that the controller can render. + when(pdfDocumentFactory.load(any(File.class))).thenReturn(Loader.loadPDF(buildPdf(3))); + + generateFrames(inputPdf, outputDir, 72, 1.0f, null, false); + + assertTrue(Files.exists(outputDir.resolve("frame_00001.png"))); + assertTrue(Files.exists(outputDir.resolve("frame_00002.png"))); + assertTrue(Files.exists(outputDir.resolve("frame_00003.png"))); + try (var stream = Files.list(outputDir)) { + assertEquals(3, stream.count()); + } + } + + @Test + @DisplayName("applies watermark when enabled and still produces frames") + void appliesWatermarkWhenEnabled(@TempDir Path dir) throws Exception { + Path inputPdf = dir.resolve("input.pdf"); + Files.write(inputPdf, buildPdf(1)); + Path outputDir = Files.createDirectory(dir.resolve("frames")); + + when(pdfDocumentFactory.load(any(File.class))).thenReturn(Loader.loadPDF(buildPdf(1))); + + generateFrames(inputPdf, outputDir, 72, 0.8f, "DRAFT", true); + + Path frame = outputDir.resolve("frame_00001.png"); + assertTrue(Files.exists(frame)); + assertTrue(Files.size(frame) > 0); + } + + @Test + @DisplayName("zero-page document throws IllegalArgumentException") + void zeroPageThrows(@TempDir Path dir) throws Exception { + Path inputPdf = dir.resolve("empty.pdf"); + // An empty PDDocument (no pages) cannot be saved/loaded, so feed an empty doc directly. + Files.write(inputPdf, new byte[] {0}); + Path outputDir = Files.createDirectory(dir.resolve("frames")); + + when(pdfDocumentFactory.load(any(File.class))).thenReturn(new PDDocument()); + + assertThrows( + IllegalArgumentException.class, + () -> generateFrames(inputPdf, outputDir, 72, 1.0f, null, false)); + try (var stream = Files.list(outputDir)) { + assertEquals(0, stream.count()); + } + } + + @Test + @DisplayName("propagates IOException from the document factory") + void propagatesLoadFailure(@TempDir Path dir) throws Exception { + Path inputPdf = dir.resolve("input.pdf"); + Files.write(inputPdf, buildPdf(1)); + Path outputDir = Files.createDirectory(dir.resolve("frames")); + + when(pdfDocumentFactory.load(any(File.class))).thenThrow(new IOException("boom")); + + assertThrows( + IOException.class, + () -> generateFrames(inputPdf, outputDir, 72, 1.0f, null, false)); + } + } +} diff --git a/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/AutoSplitPdfControllerTest.java b/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/AutoSplitPdfControllerTest.java new file mode 100644 index 0000000000..09ff35463a --- /dev/null +++ b/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/AutoSplitPdfControllerTest.java @@ -0,0 +1,637 @@ +package stirling.software.SPDF.controller.api.misc; + +import static org.junit.jupiter.api.Assertions.*; +import static org.mockito.Mockito.*; + +import java.awt.Color; +import java.awt.Graphics2D; +import java.awt.image.BufferedImage; +import java.io.ByteArrayInputStream; +import java.io.ByteArrayOutputStream; +import java.io.File; +import java.io.InputStream; +import java.lang.reflect.InvocationTargetException; +import java.lang.reflect.Method; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.ArrayList; +import java.util.HashSet; +import java.util.List; +import java.util.Set; +import java.util.zip.ZipEntry; +import java.util.zip.ZipInputStream; + +import org.apache.pdfbox.Loader; +import org.apache.pdfbox.pdmodel.PDDocument; +import org.apache.pdfbox.pdmodel.PDPage; +import org.apache.pdfbox.pdmodel.PDPageContentStream; +import org.apache.pdfbox.pdmodel.common.PDRectangle; +import org.apache.pdfbox.pdmodel.graphics.image.LosslessFactory; +import org.apache.pdfbox.pdmodel.graphics.image.PDImageXObject; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.junit.jupiter.api.io.TempDir; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.core.io.Resource; +import org.springframework.http.HttpStatus; +import org.springframework.http.ResponseEntity; + +import com.google.zxing.BarcodeFormat; +import com.google.zxing.EncodeHintType; +import com.google.zxing.common.BitMatrix; +import com.google.zxing.qrcode.QRCodeWriter; +import com.google.zxing.qrcode.decoder.ErrorCorrectionLevel; + +import stirling.software.SPDF.model.api.misc.AutoSplitPdfRequest; +import stirling.software.common.model.ApplicationProperties; +import stirling.software.common.service.CustomPDFDocumentFactory; +import stirling.software.common.util.TempFileManager; + +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class AutoSplitPdfControllerTest { + + private static final String VALID_QR = "https://stirlingpdf.com"; + + @Mock private CustomPDFDocumentFactory pdfDocumentFactory; + @Mock private TempFileManager tempFileManager; + + private ApplicationProperties applicationProperties; + private AutoSplitPdfController controller; + + @TempDir Path tempDir; + + @BeforeEach + void setUp() { + applicationProperties = new ApplicationProperties(); + // Keep maxDPI at the QR detection DPI so the high-DPI retry path is skipped (fast tests). + applicationProperties.getSystem().setMaxDPI(150); + controller = + new AutoSplitPdfController( + pdfDocumentFactory, tempFileManager, applicationProperties); + } + + // --------------------------------------------------------------------- + // Helpers + // --------------------------------------------------------------------- + + /** Build a tiny single-colour BufferedImage. */ + private static BufferedImage solidImage(int w, int h, Color color) { + BufferedImage image = new BufferedImage(w, h, BufferedImage.TYPE_INT_RGB); + Graphics2D g = image.createGraphics(); + g.setColor(color); + g.fillRect(0, 0, w, h); + g.dispose(); + return image; + } + + /** Generate a real, decodable QR code as a BufferedImage using zxing core only. */ + private static BufferedImage qrImage(String text, int size) throws Exception { + QRCodeWriter writer = new QRCodeWriter(); + java.util.Map hints = new java.util.EnumMap<>(EncodeHintType.class); + hints.put(EncodeHintType.ERROR_CORRECTION, ErrorCorrectionLevel.M); + hints.put(EncodeHintType.MARGIN, 4); + BitMatrix matrix = writer.encode(text, BarcodeFormat.QR_CODE, size, size, hints); + int width = matrix.getWidth(); + int height = matrix.getHeight(); + BufferedImage image = new BufferedImage(width, height, BufferedImage.TYPE_INT_RGB); + for (int x = 0; x < width; x++) { + for (int y = 0; y < height; y++) { + image.setRGB(x, y, matrix.get(x, y) ? Color.BLACK.getRGB() : Color.WHITE.getRGB()); + } + } + return image; + } + + /** Build an in-memory PDF document with the given number of plain pages. */ + private static PDDocument simpleDoc(int pages) { + PDDocument doc = new PDDocument(); + for (int i = 0; i < pages; i++) { + doc.addPage(new PDPage(new PDRectangle(200, 200))); + } + return doc; + } + + /** + * A PDF where every page draws the supplied image (used so embedded-image extraction works). + */ + private static PDDocument docWithImageOnEachPage(BufferedImage img, int pages) + throws Exception { + PDDocument doc = new PDDocument(); + for (int i = 0; i < pages; i++) { + PDPage page = new PDPage(new PDRectangle(200, 200)); + doc.addPage(page); + PDImageXObject xobj = LosslessFactory.createFromImage(doc, img); + try (PDPageContentStream cs = new PDPageContentStream(doc, page)) { + cs.drawImage(xobj, 20, 20, 160, 160); + } + } + return doc; + } + + /** Load a PDF from the InputStream the controller hands to pdfDocumentFactory.load(...). */ + private static PDDocument loadFromStream(InputStream in) throws Exception { + return Loader.loadPDF(in.readAllBytes()); + } + + private static byte[] docToBytes(PDDocument doc) throws Exception { + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + doc.save(baos); + doc.close(); + return baos.toByteArray(); + } + + private static AutoSplitPdfRequest request(byte[] pdfBytes, Boolean duplex) { + AutoSplitPdfRequest req = new AutoSplitPdfRequest(); + req.setFileInput( + new org.springframework.mock.web.MockMultipartFile( + "fileInput", "input.pdf", "application/pdf", pdfBytes)); + req.setDuplexMode(duplex); + return req; + } + + /** Make tempFileManager.createTempFile(".zip") create a real file inside the JUnit temp dir. */ + private void wireRealTempFile() throws Exception { + when(tempFileManager.createTempFile(".zip")) + .thenAnswer( + invocation -> + Files.createTempFile(tempDir, "stirling-test", ".zip").toFile()); + } + + private static List zipEntryNames(byte[] zipBytes) throws Exception { + List names = new ArrayList<>(); + try (ZipInputStream zis = new ZipInputStream(new ByteArrayInputStream(zipBytes))) { + ZipEntry entry; + while ((entry = zis.getNextEntry()) != null) { + names.add(entry.getName()); + zis.closeEntry(); + } + } + return names; + } + + private static byte[] readResource(Resource resource) throws Exception { + try (InputStream in = resource.getInputStream()) { + return in.readAllBytes(); + } + } + + // reflection invokers for the private (static) helpers ------------------ + + private static Object invokeStatic(String name, Class[] types, Object... args) + throws Exception { + Method m = AutoSplitPdfController.class.getDeclaredMethod(name, types); + m.setAccessible(true); + try { + return m.invoke(null, args); + } catch (InvocationTargetException e) { + if (e.getCause() instanceof Exception ex) { + throw ex; + } + throw e; + } + } + + private Object invokeInstance(String name, Class[] types, Object... args) throws Exception { + Method m = AutoSplitPdfController.class.getDeclaredMethod(name, types); + m.setAccessible(true); + try { + return m.invoke(controller, args); + } catch (InvocationTargetException e) { + if (e.getCause() instanceof Exception ex) { + throw ex; + } + throw e; + } + } + + // --------------------------------------------------------------------- + // isBlankImage(int[]) + // --------------------------------------------------------------------- + + @Nested + @DisplayName("isBlankImage") + class IsBlankImage { + + @Test + @DisplayName("empty array is treated as blank") + void emptyArrayIsBlank() throws Exception { + Object result = + invokeStatic("isBlankImage", new Class[] {int[].class}, (Object) new int[0]); + assertEquals(Boolean.TRUE, result); + } + + @Test + @DisplayName("uniform pixels are blank") + void uniformIsBlank() throws Exception { + int[] pixels = new int[1000]; + java.util.Arrays.fill(pixels, 0xFFFFFF); + Object result = + invokeStatic("isBlankImage", new Class[] {int[].class}, (Object) pixels); + assertEquals(Boolean.TRUE, result); + } + + @Test + @DisplayName("a single differing sampled pixel makes it non-blank") + void variedIsNotBlank() throws Exception { + int[] pixels = new int[1000]; + java.util.Arrays.fill(pixels, 0xFFFFFF); + // step = max(1, 1000/20) = 50, so index 500 is sampled + pixels[500] = 0x000000; + Object result = + invokeStatic("isBlankImage", new Class[] {int[].class}, (Object) pixels); + assertEquals(Boolean.FALSE, result); + } + + @Test + @DisplayName("single pixel array is blank") + void singlePixelIsBlank() throws Exception { + Object result = + invokeStatic( + "isBlankImage", new Class[] {int[].class}, (Object) new int[] {7}); + assertEquals(Boolean.TRUE, result); + } + } + + // --------------------------------------------------------------------- + // downscaleIfNeeded(BufferedImage) + // --------------------------------------------------------------------- + + @Nested + @DisplayName("downscaleIfNeeded") + class DownscaleIfNeeded { + + @Test + @DisplayName("small images are returned unchanged (same instance)") + void smallUnchanged() throws Exception { + BufferedImage img = solidImage(100, 80, Color.GRAY); + Object result = + invokeStatic( + "downscaleIfNeeded", + new Class[] {BufferedImage.class}, + (Object) img); + assertSame(img, result); + } + + @Test + @DisplayName("image exactly at the pixel limit is unchanged") + void atLimitUnchanged() throws Exception { + // 10000x10000 == 100_000_000 == MAX_IMAGE_PIXELS, not greater than -> unchanged. + // Use a mock-free real image but keep it cheap: build a 1px-tall wide image whose + // total pixel count is below the limit so we only assert the <= branch with a + // realistic non-trivial size. + BufferedImage img = solidImage(5000, 5000, Color.WHITE); // 25M < 100M + Object result = + invokeStatic( + "downscaleIfNeeded", + new Class[] {BufferedImage.class}, + (Object) img); + assertSame(img, result); + } + } + + // --------------------------------------------------------------------- + // countPageImages(PDPage) + // --------------------------------------------------------------------- + + @Nested + @DisplayName("countPageImages") + class CountPageImages { + + @Test + @DisplayName("page with no resources returns 0") + void noResourcesZero() throws Exception { + PDPage page = new PDPage(new PDRectangle(200, 200)); + Object result = + invokeStatic("countPageImages", new Class[] {PDPage.class}, (Object) page); + assertEquals(0, result); + } + + @Test + @DisplayName("page with one embedded image returns 1") + void oneImage() throws Exception { + try (PDDocument doc = docWithImageOnEachPage(solidImage(40, 40, Color.RED), 1)) { + PDPage page = doc.getPage(0); + Object result = + invokeStatic( + "countPageImages", new Class[] {PDPage.class}, (Object) page); + assertEquals(1, result); + } + } + } + + // --------------------------------------------------------------------- + // tryDecodeQR(int[], int, int) and decodeQRCode(BufferedImage) + // --------------------------------------------------------------------- + + @Nested + @DisplayName("QR decoding") + class QrDecoding { + + @Test + @DisplayName("decodeQRCode returns the encoded text for a real QR image") + void decodeRealQr() throws Exception { + BufferedImage qr = qrImage(VALID_QR, 250); + Object result = + invokeStatic("decodeQRCode", new Class[] {BufferedImage.class}, (Object) qr); + assertEquals(VALID_QR, result); + } + + @Test + @DisplayName("decodeQRCode returns null for a blank image") + void decodeBlankReturnsNull() throws Exception { + BufferedImage blank = solidImage(120, 120, Color.WHITE); + Object result = + invokeStatic( + "decodeQRCode", new Class[] {BufferedImage.class}, (Object) blank); + assertNull(result); + } + + @Test + @DisplayName("decodeQRCode returns null for a non-QR (noise-free, non-blank) image") + void decodeNonQrReturnsNull() throws Exception { + // Two solid halves: non-blank but not a QR code. + BufferedImage image = new BufferedImage(120, 120, BufferedImage.TYPE_INT_RGB); + Graphics2D g = image.createGraphics(); + g.setColor(Color.WHITE); + g.fillRect(0, 0, 120, 60); + g.setColor(Color.BLACK); + g.fillRect(0, 60, 120, 60); + g.dispose(); + Object result = + invokeStatic( + "decodeQRCode", new Class[] {BufferedImage.class}, (Object) image); + assertNull(result); + } + + @Test + @DisplayName("tryDecodeQR decodes raw RGB pixels of a real QR") + void tryDecodeRawPixels() throws Exception { + BufferedImage qr = qrImage(VALID_QR, 250); + int w = qr.getWidth(); + int h = qr.getHeight(); + int[] pixels = new int[w * h]; + qr.getRGB(0, 0, w, h, pixels, 0, w); + Object result = + invokeStatic( + "tryDecodeQR", + new Class[] {int[].class, int.class, int.class}, + pixels, + w, + h); + assertEquals(VALID_QR, result); + } + + @Test + @DisplayName("tryDecodeQR returns null when no QR present") + void tryDecodeReturnsNull() throws Exception { + int w = 60; + int h = 60; + int[] pixels = new int[w * h]; + // alternating pattern, no decodable QR + for (int i = 0; i < pixels.length; i++) { + pixels[i] = (i % 2 == 0) ? 0xFFFFFF : 0x000000; + } + Object result = + invokeStatic( + "tryDecodeQR", + new Class[] {int[].class, int.class, int.class}, + pixels, + w, + h); + assertNull(result); + } + } + + // --------------------------------------------------------------------- + // checkPageImagesDirect(PDPage) + // --------------------------------------------------------------------- + + @Nested + @DisplayName("checkPageImagesDirect") + class CheckPageImagesDirect { + + @Test + @DisplayName("page with no images returns null") + void noImagesNull() throws Exception { + PDPage page = new PDPage(new PDRectangle(200, 200)); + Object result = + invokeStatic( + "checkPageImagesDirect", new Class[] {PDPage.class}, (Object) page); + assertNull(result); + } + + @Test + @DisplayName("page with an embedded QR image returns the QR text") + void embeddedQrFound() throws Exception { + // Size 250 keeps the readback image off the controller's blank-sampling heuristic. + BufferedImage qr = qrImage(VALID_QR, 250); + try (PDDocument doc = docWithImageOnEachPage(qr, 1)) { + PDPage page = doc.getPage(0); + Object result = + invokeStatic( + "checkPageImagesDirect", + new Class[] {PDPage.class}, + (Object) page); + assertEquals(VALID_QR, result); + } + } + + @Test + @DisplayName("page with a non-QR image returns null") + void embeddedNonQrNull() throws Exception { + BufferedImage plain = solidImage(80, 80, Color.WHITE); + try (PDDocument doc = docWithImageOnEachPage(plain, 1)) { + PDPage page = doc.getPage(0); + Object result = + invokeStatic( + "checkPageImagesDirect", + new Class[] {PDPage.class}, + (Object) page); + assertNull(result); + } + } + } + + // --------------------------------------------------------------------- + // getSystemMaxDpi() (private instance) + // --------------------------------------------------------------------- + + @Nested + @DisplayName("getSystemMaxDpi") + class GetSystemMaxDpi { + + @Test + @DisplayName("returns the configured system maxDPI") + void returnsConfiguredValue() throws Exception { + applicationProperties.getSystem().setMaxDPI(300); + Object result = invokeInstance("getSystemMaxDpi", new Class[] {}); + assertEquals(300, result); + } + + @Test + @DisplayName("falls back to the default detection DPI when applicationProperties is null") + void fallsBackWhenNull() throws Exception { + AutoSplitPdfController noProps = + new AutoSplitPdfController(pdfDocumentFactory, tempFileManager, null); + Method m = AutoSplitPdfController.class.getDeclaredMethod("getSystemMaxDpi"); + m.setAccessible(true); + Object result = m.invoke(noProps); + assertEquals(150, result); // QR_DETECTION_DPI + } + } + + // --------------------------------------------------------------------- + // autoSplitPdf(...) full handler + // --------------------------------------------------------------------- + + @Nested + @DisplayName("autoSplitPdf") + class AutoSplitHandler { + + @Test + @DisplayName("PDF without QR dividers yields a single-PDF zip") + void noQrSingleOutput() throws Exception { + byte[] pdf = docToBytes(simpleDoc(3)); + wireRealTempFile(); + when(pdfDocumentFactory.load(any(InputStream.class))) + .thenAnswer(invocation -> loadFromStream(invocation.getArgument(0))); + + AutoSplitPdfRequest req = request(pdf, false); + ResponseEntity response = controller.autoSplitPdf(req); + + assertNotNull(response); + assertEquals(HttpStatus.OK, response.getStatusCode()); + assertNotNull(response.getBody()); + + byte[] zip = readResource(response.getBody()); + List names = zipEntryNames(zip); + assertEquals(1, names.size(), "no QR dividers -> all pages collapse into one document"); + assertEquals("input_1.pdf", names.get(0)); + + verify(pdfDocumentFactory).load(any(InputStream.class)); + } + + @Test + @DisplayName("single-page PDF without QR yields one output PDF in the zip") + void singlePageOneOutput() throws Exception { + byte[] pdf = docToBytes(simpleDoc(1)); + wireRealTempFile(); + when(pdfDocumentFactory.load(any(InputStream.class))) + .thenAnswer(invocation -> loadFromStream(invocation.getArgument(0))); + + ResponseEntity response = controller.autoSplitPdf(request(pdf, null)); + + byte[] zip = readResource(response.getBody()); + List names = zipEntryNames(zip); + assertEquals(1, names.size()); + assertEquals("input_1.pdf", names.get(0)); + } + + @Test + @DisplayName("filename without extension is used as-is for entry names") + void filenameWithoutExtension() throws Exception { + byte[] pdf = docToBytes(simpleDoc(2)); + wireRealTempFile(); + when(pdfDocumentFactory.load(any(InputStream.class))) + .thenAnswer(invocation -> loadFromStream(invocation.getArgument(0))); + + AutoSplitPdfRequest req = new AutoSplitPdfRequest(); + req.setFileInput( + new org.springframework.mock.web.MockMultipartFile( + "fileInput", "myfile", "application/pdf", pdf)); + req.setDuplexMode(false); + + ResponseEntity response = controller.autoSplitPdf(req); + byte[] zip = readResource(response.getBody()); + List names = zipEntryNames(zip); + assertEquals(1, names.size()); + assertEquals("myfile_1.pdf", names.get(0)); + } + + @Test + @DisplayName("each output zip entry contains a valid, loadable PDF") + void outputEntriesAreValidPdfs() throws Exception { + byte[] pdf = docToBytes(simpleDoc(2)); + wireRealTempFile(); + when(pdfDocumentFactory.load(any(InputStream.class))) + .thenAnswer(invocation -> loadFromStream(invocation.getArgument(0))); + + ResponseEntity response = controller.autoSplitPdf(request(pdf, false)); + byte[] zip = readResource(response.getBody()); + + int entries = 0; + try (ZipInputStream zis = new ZipInputStream(new ByteArrayInputStream(zip))) { + ZipEntry entry; + while ((entry = zis.getNextEntry()) != null) { + entries++; + byte[] entryBytes = zis.readAllBytes(); + try (PDDocument loaded = Loader.loadPDF(entryBytes)) { + assertTrue(loaded.getNumberOfPages() >= 1); + } + zis.closeEntry(); + } + } + assertEquals(1, entries); + } + + @Test + @DisplayName("loader failure propagates and the temp file is closed") + void loaderFailurePropagates() throws Exception { + byte[] pdf = docToBytes(simpleDoc(1)); + File created = Files.createTempFile(tempDir, "stirling-fail", ".zip").toFile(); + when(tempFileManager.createTempFile(".zip")).thenReturn(created); + when(pdfDocumentFactory.load(any(InputStream.class))) + .thenThrow(new java.io.IOException("boom")); + + AutoSplitPdfRequest req = request(pdf, false); + + assertThrows(java.io.IOException.class, () -> controller.autoSplitPdf(req)); + // outputTempFile.close() deletes the file on the error path. + verify(tempFileManager).deleteTempFile(created); + } + + @Test + @DisplayName("duplexMode flag is accepted (null treated as false)") + void duplexNullTreatedAsFalse() throws Exception { + byte[] pdf = docToBytes(simpleDoc(2)); + wireRealTempFile(); + when(pdfDocumentFactory.load(any(InputStream.class))) + .thenAnswer(invocation -> loadFromStream(invocation.getArgument(0))); + + ResponseEntity response = controller.autoSplitPdf(request(pdf, null)); + assertEquals(HttpStatus.OK, response.getStatusCode()); + List names = zipEntryNames(readResource(response.getBody())); + assertEquals(1, names.size()); + } + } + + // --------------------------------------------------------------------- + // VALID_QR_CONTENTS sanity (static state) + // --------------------------------------------------------------------- + + @Nested + @DisplayName("valid QR contents") + class ValidQrContents { + + @Test + @DisplayName("the well-known divider URLs are recognised") + @SuppressWarnings("unchecked") + void recognisedUrls() throws Exception { + java.lang.reflect.Field f = + AutoSplitPdfController.class.getDeclaredField("VALID_QR_CONTENTS"); + f.setAccessible(true); + Set valid = new HashSet<>((Set) f.get(null)); + assertTrue(valid.contains("https://stirlingpdf.com")); + assertTrue(valid.contains("https://github.com/Stirling-Tools/Stirling-PDF")); + assertTrue(valid.contains("https://github.com/Frooodle/Stirling-PDF")); + assertFalse(valid.contains("https://example.com")); + } + } +} diff --git a/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/CompressControllerTest.java b/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/CompressControllerTest.java new file mode 100644 index 0000000000..d5a5c6a30c --- /dev/null +++ b/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/CompressControllerTest.java @@ -0,0 +1,469 @@ +package stirling.software.SPDF.controller.api.misc; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.Mockito.lenient; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; + +import java.io.ByteArrayOutputStream; +import java.io.File; +import java.io.IOException; +import java.io.InputStream; +import java.lang.reflect.Field; +import java.lang.reflect.Method; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.ArrayList; +import java.util.List; + +import org.apache.pdfbox.Loader; +import org.apache.pdfbox.pdmodel.PDDocument; +import org.apache.pdfbox.pdmodel.PDPage; +import org.apache.pdfbox.pdmodel.PDPageContentStream; +import org.apache.pdfbox.pdmodel.common.PDRectangle; +import org.apache.pdfbox.pdmodel.graphics.image.LosslessFactory; +import org.apache.pdfbox.pdmodel.graphics.image.PDImageXObject; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.junit.jupiter.api.io.TempDir; +import org.mockito.InjectMocks; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.core.io.Resource; +import org.springframework.http.HttpStatus; +import org.springframework.http.MediaType; +import org.springframework.http.ResponseEntity; +import org.springframework.mock.web.MockMultipartFile; +import org.springframework.web.server.ResponseStatusException; + +import stirling.software.SPDF.config.EndpointConfiguration; +import stirling.software.SPDF.model.api.misc.OptimizePdfRequest; +import stirling.software.common.service.CustomPDFDocumentFactory; +import stirling.software.common.service.LineArtConversionService; +import stirling.software.common.util.TempFile; +import stirling.software.common.util.TempFileManager; + +/** + * Unit tests for {@link CompressController}. External binaries (Ghostscript / qpdf / ImageMagick) + * are never invoked: tests cover validation, orchestration with all tool groups disabled, the pure + * level/quality helpers, and the public image-compression entry point. + */ +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class CompressControllerTest { + + @TempDir Path tempDir; + + @Mock private CustomPDFDocumentFactory pdfDocumentFactory; + @Mock private EndpointConfiguration endpointConfiguration; + @Mock private TempFileManager tempFileManager; + + @InjectMocks private CompressController controller; + + /** Real temp files created during a test; cleaned up after each test. */ + private final List createdFiles = new ArrayList<>(); + + @BeforeEach + void setUp() throws Exception { + // By default no external tools are enabled, forcing the Java-only path. + lenient().when(endpointConfiguration.isGroupEnabled(anyString())).thenReturn(false); + + // Every managed temp file is backed by a real on-disk file wrapped in a mock TempFile. + lenient() + .when(tempFileManager.createManagedTempFile(anyString())) + .thenAnswer( + inv -> { + File f = + Files.createTempFile( + "compress-test", inv.getArgument(0)) + .toFile(); + createdFiles.add(f); + return newRealBackedTempFile(f); + }); + } + + private TempFile newRealBackedTempFile(File f) { + TempFile tf = mock(TempFile.class); + lenient().when(tf.getFile()).thenReturn(f); + lenient().when(tf.getPath()).thenReturn(f.toPath()); + lenient().when(tf.getAbsolutePath()).thenReturn(f.getAbsolutePath()); + lenient().when(tf.exists()).thenReturn(f.exists()); + return tf; + } + + // ----- helpers to build tiny in-memory PDFs ------------------------------------------------ + + private byte[] textOnlyPdfBytes() throws IOException { + try (PDDocument doc = new PDDocument()) { + PDPage page = new PDPage(PDRectangle.LETTER); + doc.addPage(page); + try (PDPageContentStream cs = new PDPageContentStream(doc, page)) { + cs.beginText(); + cs.setFont( + new org.apache.pdfbox.pdmodel.font.PDType1Font( + org.apache.pdfbox.pdmodel.font.Standard14Fonts.FontName.HELVETICA), + 12); + cs.newLineAtOffset(50, 700); + cs.showText("Hello compress"); + cs.endText(); + } + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + doc.save(baos); + return baos.toByteArray(); + } + } + + /** PDF with one tiny (sub-400px) image so the compressor encounters it but skips it. */ + private byte[] smallImagePdfBytes() throws IOException { + try (PDDocument doc = new PDDocument()) { + PDPage page = new PDPage(PDRectangle.LETTER); + doc.addPage(page); + java.awt.image.BufferedImage img = + new java.awt.image.BufferedImage( + 50, 50, java.awt.image.BufferedImage.TYPE_INT_RGB); + for (int x = 0; x < 50; x++) { + for (int y = 0; y < 50; y++) { + img.setRGB(x, y, (x * 5 + y) & 0xFFFFFF); + } + } + PDImageXObject image = LosslessFactory.createFromImage(doc, img); + try (PDPageContentStream cs = new PDPageContentStream(doc, page)) { + cs.drawImage(image, 100, 100, 50, 50); + } + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + doc.save(baos); + return baos.toByteArray(); + } + } + + private MockMultipartFile multipart(byte[] bytes) { + return new MockMultipartFile( + "fileInput", "input.pdf", MediaType.APPLICATION_PDF_VALUE, bytes); + } + + private static byte[] drain(ResponseEntity response) throws IOException { + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + try (InputStream in = response.getBody().getInputStream()) { + in.transferTo(baos); + } + return baos.toByteArray(); + } + + // ----- validation branches ----------------------------------------------------------------- + + @Nested + @DisplayName("optimizePdf validation") + class Validation { + + @Test + @DisplayName("null input file throws IllegalArgumentException") + void nullFile_throws() { + OptimizePdfRequest request = new OptimizePdfRequest(); + request.setFileInput(null); + + assertThatThrownBy(() -> controller.optimizePdf(request)) + .isInstanceOf(IllegalArgumentException.class); + } + + @Test + @DisplayName("empty input file throws IllegalArgumentException") + void emptyFile_throws() { + OptimizePdfRequest request = new OptimizePdfRequest(); + request.setFileInput( + new MockMultipartFile( + "fileInput", + "input.pdf", + MediaType.APPLICATION_PDF_VALUE, + new byte[0])); + + assertThatThrownBy(() -> controller.optimizePdf(request)) + .isInstanceOf(IllegalArgumentException.class); + } + + @Test + @DisplayName( + "both optimizeLevel and expectedOutputSize null throws IllegalArgumentException") + void noOptionsProvided_throws() throws Exception { + OptimizePdfRequest request = new OptimizePdfRequest(); + request.setFileInput(multipart(textOnlyPdfBytes())); + request.setOptimizeLevel(null); + request.setExpectedOutputSize(null); + + assertThatThrownBy(() -> controller.optimizePdf(request)) + .isInstanceOf(IllegalArgumentException.class); + } + + @Test + @DisplayName("line art requested but service unavailable throws FORBIDDEN") + void lineArt_serviceNull_throwsForbidden() throws Exception { + OptimizePdfRequest request = new OptimizePdfRequest(); + request.setFileInput(multipart(textOnlyPdfBytes())); + request.setOptimizeLevel(2); + request.setLineArt(true); + + // lineArtConversionService field is left null by @InjectMocks. + assertThatThrownBy(() -> controller.optimizePdf(request)) + .isInstanceOf(ResponseStatusException.class) + .extracting(e -> ((ResponseStatusException) e).getStatusCode()) + .isEqualTo(HttpStatus.FORBIDDEN); + } + + @Test + @DisplayName( + "line art requested with service present but ImageMagick disabled throws IOException") + void lineArt_imageMagickDisabled_throwsIOException() throws Exception { + OptimizePdfRequest request = new OptimizePdfRequest(); + request.setFileInput(multipart(textOnlyPdfBytes())); + request.setOptimizeLevel(2); + request.setLineArt(true); + + // Provide a service so the null-check passes, but ImageMagick group stays disabled. + setLineArtService(mock(LineArtConversionService.class)); + when(endpointConfiguration.isGroupEnabled("ImageMagick")).thenReturn(false); + + assertThatThrownBy(() -> controller.optimizePdf(request)) + .isInstanceOf(IOException.class); + } + } + + // ----- orchestration with all external tools disabled -------------------------------------- + + @Nested + @DisplayName("optimizePdf orchestration (no external tools)") + class Orchestration { + + @Test + @DisplayName("low level + no tools returns OK with non-empty body") + void lowLevel_noTools_returnsOk() throws Exception { + byte[] pdf = textOnlyPdfBytes(); + OptimizePdfRequest request = new OptimizePdfRequest(); + request.setFileInput(multipart(pdf)); + request.setOptimizeLevel(1); // < 4 => no image compression, < 6 => no ghostscript + + // Final stage reloads currentFile from disk; return a fresh real document. + when(pdfDocumentFactory.load(any(File.class))) + .thenAnswer(inv -> Loader.loadPDF((File) inv.getArgument(0))); + + ResponseEntity response = controller.optimizePdf(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(drain(response)).isNotEmpty(); + } + + @Test + @DisplayName("level 4 with a sub-threshold image still returns OK (image skipped)") + void level4_smallImage_skipped_returnsOk() throws Exception { + byte[] pdf = smallImagePdfBytes(); + OptimizePdfRequest request = new OptimizePdfRequest(); + request.setFileInput(multipart(pdf)); + request.setOptimizeLevel(4); // triggers Java image compression path + + when(pdfDocumentFactory.load(any(File.class))) + .thenAnswer(inv -> Loader.loadPDF((File) inv.getArgument(0))); + // compressImagesInPDF reloads currentFile via load(Path). + when(pdfDocumentFactory.load(any(Path.class))) + .thenAnswer(inv -> Loader.loadPDF(((Path) inv.getArgument(0)).toFile())); + + ResponseEntity response = controller.optimizePdf(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(drain(response)).isNotEmpty(); + } + + @Test + @DisplayName("grayscale flag forces image compression path even at low level") + void grayscale_lowLevel_returnsOk() throws Exception { + byte[] pdf = textOnlyPdfBytes(); + OptimizePdfRequest request = new OptimizePdfRequest(); + request.setFileInput(multipart(pdf)); + request.setOptimizeLevel(1); + request.setGrayscale(true); + + when(pdfDocumentFactory.load(any(File.class))) + .thenAnswer(inv -> Loader.loadPDF((File) inv.getArgument(0))); + when(pdfDocumentFactory.load(any(Path.class))) + .thenAnswer(inv -> Loader.loadPDF(((Path) inv.getArgument(0)).toFile())); + + ResponseEntity response = controller.optimizePdf(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(drain(response)).isNotEmpty(); + } + + @Test + @DisplayName("auto mode via expectedOutputSize picks a level and returns OK") + void autoMode_expectedOutputSize_returnsOk() throws Exception { + byte[] pdf = textOnlyPdfBytes(); + OptimizePdfRequest request = new OptimizePdfRequest(); + request.setFileInput(multipart(pdf)); + request.setOptimizeLevel(null); + request.setExpectedOutputSize("1KB"); // length > 1 => auto mode + + when(pdfDocumentFactory.load(any(File.class))) + .thenAnswer(inv -> Loader.loadPDF((File) inv.getArgument(0))); + when(pdfDocumentFactory.load(any(Path.class))) + .thenAnswer(inv -> Loader.loadPDF(((Path) inv.getArgument(0)).toFile())); + + ResponseEntity response = controller.optimizePdf(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(drain(response)).isNotEmpty(); + } + } + + // ----- public compressImagesInPDF ---------------------------------------------------------- + + @Nested + @DisplayName("compressImagesInPDF") + class CompressImages { + + @Test + @DisplayName("PDF with a tiny image produces a valid, non-empty output PDF") + void smallImage_producesValidPdf() throws Exception { + Path src = tempDir.resolve("src.pdf"); + Files.write(src, smallImagePdfBytes()); + + when(pdfDocumentFactory.load(any(Path.class))) + .thenAnswer(inv -> Loader.loadPDF(((Path) inv.getArgument(0)).toFile())); + + TempFile result = controller.compressImagesInPDF(src, 0.5, 0.5f, false); + + assertThat(result).isNotNull(); + byte[] out = Files.readAllBytes(result.getPath()); + assertThat(out).isNotEmpty(); + try (PDDocument doc = Loader.loadPDF(out)) { + assertThat(doc.getNumberOfPages()).isEqualTo(1); + } + } + + @Test + @DisplayName("text-only PDF compresses to a valid single-page output") + void textOnly_producesValidPdf() throws Exception { + Path src = tempDir.resolve("text.pdf"); + Files.write(src, textOnlyPdfBytes()); + + when(pdfDocumentFactory.load(any(Path.class))) + .thenAnswer(inv -> Loader.loadPDF(((Path) inv.getArgument(0)).toFile())); + + TempFile result = controller.compressImagesInPDF(src, 0.8, 0.7f, false); + + assertThat(result).isNotNull(); + try (PDDocument doc = Loader.loadPDF(Files.readAllBytes(result.getPath()))) { + assertThat(doc.getNumberOfPages()).isEqualTo(1); + } + } + + @Test + @DisplayName("load failure closes the temp file and propagates the exception") + void loadFailure_propagates() throws Exception { + Path src = tempDir.resolve("bad.pdf"); + Files.write(src, textOnlyPdfBytes()); + + when(pdfDocumentFactory.load(any(Path.class))).thenThrow(new IOException("boom")); + + assertThatThrownBy(() -> controller.compressImagesInPDF(src, 0.5, 0.5f, false)) + .isInstanceOf(IOException.class) + .hasMessageContaining("boom"); + } + } + + // ----- pure helper methods (reflection) ---------------------------------------------------- + + @Nested + @DisplayName("scale / quality / level helpers") + class Helpers { + + @Test + @DisplayName("getScaleFactorForLevel maps each level and defaults to 1.0") + void scaleFactorForLevel() throws Exception { + Method m = privateStatic("getScaleFactorForLevel", int.class); + assertThat((double) m.invoke(null, 1)).isEqualTo(0.98); + assertThat((double) m.invoke(null, 5)).isEqualTo(0.68); + assertThat((double) m.invoke(null, 9)).isEqualTo(0.28); + // Out-of-range falls to the default branch. + assertThat((double) m.invoke(null, 0)).isEqualTo(1.0); + assertThat((double) m.invoke(null, 42)).isEqualTo(1.0); + } + + @Test + @DisplayName("getJpegQualityForLevel maps each level and defaults to 0.75") + void jpegQualityForLevel() throws Exception { + Method m = privateStatic("getJpegQualityForLevel", int.class); + assertThat((float) m.invoke(null, 1)).isEqualTo(0.92f); + assertThat((float) m.invoke(null, 9)).isEqualTo(0.35f); + assertThat((float) m.invoke(null, 0)).isEqualTo(0.75f); + assertThat((float) m.invoke(null, 100)).isEqualTo(0.75f); + } + + @Test + @DisplayName("determineOptimizeLevel buckets the size-reduction ratio") + void determineOptimizeLevel() throws Exception { + Method m = privateStatic("determineOptimizeLevel", double.class); + assertThat((int) m.invoke(null, 0.95)).isEqualTo(1); + assertThat((int) m.invoke(null, 0.85)).isEqualTo(2); + assertThat((int) m.invoke(null, 0.75)).isEqualTo(3); + assertThat((int) m.invoke(null, 0.65)).isEqualTo(4); + assertThat((int) m.invoke(null, 0.5)).isEqualTo(5); + assertThat((int) m.invoke(null, 0.25)).isEqualTo(6); + assertThat((int) m.invoke(null, 0.18)).isEqualTo(7); + assertThat((int) m.invoke(null, 0.12)).isEqualTo(8); + assertThat((int) m.invoke(null, 0.05)).isEqualTo(9); + } + + @Test + @DisplayName("incrementOptimizeLevel grows by ratio and is capped at 9") + void incrementOptimizeLevel() throws Exception { + Method m = privateStatic("incrementOptimizeLevel", int.class, long.class, long.class); + // ratio 3.0 (> 2.0) -> +3 + assertThat((int) m.invoke(null, 2, 300L, 100L)).isEqualTo(5); + // ratio 1.8 (> 1.5) -> +2 + assertThat((int) m.invoke(null, 2, 180L, 100L)).isEqualTo(4); + // ratio 1.1 -> +1 + assertThat((int) m.invoke(null, 2, 110L, 100L)).isEqualTo(3); + // capped at 9 + assertThat((int) m.invoke(null, 8, 300L, 100L)).isEqualTo(9); + } + + @Test + @DisplayName("getImageType classifies filters; unknown for non-image input") + void getImageType_andFilter() throws Exception { + // A PNG-style lossless image yields FlateDecode -> "PNG". + try (PDDocument doc = new PDDocument()) { + java.awt.image.BufferedImage bi = + new java.awt.image.BufferedImage( + 10, 10, java.awt.image.BufferedImage.TYPE_INT_RGB); + PDImageXObject image = LosslessFactory.createFromImage(doc, bi); + + Method typeMethod = privateStatic("getImageType", PDImageXObject.class); + String type = (String) typeMethod.invoke(null, image); + assertThat(type).isEqualTo("PNG"); + + Method filterMethod = privateStatic("getImageFilter", PDImageXObject.class); + String filter = (String) filterMethod.invoke(null, image); + assertThat(filter).contains("FlateDecode"); + } + } + } + + // ----- reflection / field helpers ---------------------------------------------------------- + + private static Method privateStatic(String name, Class... params) throws Exception { + Method m = CompressController.class.getDeclaredMethod(name, params); + m.setAccessible(true); + return m; + } + + private void setLineArtService(LineArtConversionService service) throws Exception { + Field f = CompressController.class.getDeclaredField("lineArtConversionService"); + f.setAccessible(true); + f.set(controller, service); + } +} diff --git a/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/ExtractImageScansControllerTest.java b/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/ExtractImageScansControllerTest.java new file mode 100644 index 0000000000..7e1e7da293 --- /dev/null +++ b/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/ExtractImageScansControllerTest.java @@ -0,0 +1,413 @@ +package stirling.software.SPDF.controller.api.misc; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyList; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.Mockito.atLeastOnce; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.times; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.io.ByteArrayOutputStream; +import java.io.IOException; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.List; + +import org.apache.pdfbox.pdmodel.PDDocument; +import org.apache.pdfbox.pdmodel.PDPage; +import org.apache.pdfbox.pdmodel.common.PDRectangle; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.junit.jupiter.api.io.TempDir; +import org.mockito.ArgumentCaptor; +import org.mockito.Mock; +import org.mockito.MockedStatic; +import org.mockito.Mockito; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.core.io.Resource; +import org.springframework.http.HttpStatus; +import org.springframework.http.MediaType; +import org.springframework.http.ResponseEntity; +import org.springframework.mock.web.MockMultipartFile; + +import stirling.software.SPDF.model.api.misc.ExtractImageScansRequest; +import stirling.software.common.service.CustomPDFDocumentFactory; +import stirling.software.common.util.CheckProgramInstall; +import stirling.software.common.util.GeneralUtils; +import stirling.software.common.util.ProcessExecutor; +import stirling.software.common.util.ProcessExecutor.ProcessExecutorResult; +import stirling.software.common.util.TempFileManager; +import stirling.software.common.util.TempFileRegistry; + +/** + * Unit tests for {@link ExtractImageScansController}. + * + *

The controller shells out to a Python/OpenCV script via {@link ProcessExecutor} and gates on + * {@link CheckProgramInstall#isPythonAvailable()}. To keep tests deterministic these static + * boundaries are mocked with {@code Mockito.mockStatic}: Python is forced available/unavailable, + * the script extraction is stubbed, and the process execution is replaced with an in-test answer + * that either writes fake output PNGs into the controller-owned temp directory or leaves it empty. + * No real Python process is ever launched. + */ +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class ExtractImageScansControllerTest { + + @Mock private CustomPDFDocumentFactory pdfDocumentFactory; + + private TempFileManager tempFileManager; + private ExtractImageScansController controller; + + @TempDir Path baseTmpDir; + + @BeforeEach + void setUp() { + stirling.software.common.model.ApplicationProperties applicationProperties = + new stirling.software.common.model.ApplicationProperties(); + applicationProperties + .getSystem() + .getTempFileManagement() + .setBaseTmpDir(baseTmpDir.toString()); + applicationProperties.getSystem().getTempFileManagement().setPrefix("scan-test-"); + + tempFileManager = new TempFileManager(new TempFileRegistry(), applicationProperties); + controller = new ExtractImageScansController(pdfDocumentFactory, tempFileManager); + } + + /** Build a request with sensible defaults; the caller supplies the file input. */ + private ExtractImageScansRequest requestFor(MockMultipartFile file) { + ExtractImageScansRequest request = new ExtractImageScansRequest(); + request.setFileInput(file); + request.setAngleThreshold(5); + request.setTolerance(20); + request.setMinArea(8000); + request.setMinContourArea(500); + request.setBorderSize(1); + return request; + } + + /** A tiny single-page in-memory PDF backed by a small media box for cheap rendering. */ + private MockMultipartFile pdfFile(String name) throws IOException { + try (PDDocument doc = new PDDocument(); + ByteArrayOutputStream out = new ByteArrayOutputStream()) { + // Keep the page small so the 300-DPI render stays tiny and fast. + doc.addPage(new PDPage(new PDRectangle(72f, 72f))); + doc.save(out); + return new MockMultipartFile( + "fileInput", name, MediaType.APPLICATION_PDF_VALUE, out.toByteArray()); + } + } + + /** A non-PDF image input; the controller copies it straight to a temp file. */ + private MockMultipartFile imageFile(String name) { + return new MockMultipartFile( + "fileInput", name, MediaType.IMAGE_PNG_VALUE, new byte[] {1, 2, 3, 4}); + } + + /** + * A mocked {@link ProcessExecutor} whose {@code runCommandWithOutputHandling} reads the output + * directory from the command (positional arg index 3) and writes the given number of PNG files + * into it, mimicking the real split_photos.py behaviour without launching a process. + */ + private ProcessExecutor execWritingOutputs(int outputCount) throws Exception { + ProcessExecutor exec = mock(ProcessExecutor.class); + when(exec.runCommandWithOutputHandling(anyList())) + .thenAnswer( + invocation -> { + List command = invocation.getArgument(0); + Path outDir = Path.of(command.get(3)); + for (int i = 0; i < outputCount; i++) { + Files.write( + outDir.resolve("out_" + i + ".png"), + new byte[] {9, 8, 7, (byte) i}); + } + return mock(ProcessExecutorResult.class); + }); + return exec; + } + + @Nested + @DisplayName("Python availability guard") + class PythonGuard { + + @Test + @DisplayName("throws IOException when Python is not installed") + void throwsWhenPythonUnavailable() throws Exception { + ExtractImageScansRequest request = requestFor(pdfFile("scan.pdf")); + + try (MockedStatic check = + Mockito.mockStatic(CheckProgramInstall.class)) { + check.when(CheckProgramInstall::isPythonAvailable).thenReturn(false); + + assertThrows(IOException.class, () -> controller.extractImageScans(request)); + } + } + + @Test + @DisplayName("does not load the PDF or extract the script when Python is missing") + void shortCircuitsBeforeAnyWork() throws Exception { + ExtractImageScansRequest request = requestFor(pdfFile("scan.pdf")); + + try (MockedStatic check = + Mockito.mockStatic(CheckProgramInstall.class); + MockedStatic general = Mockito.mockStatic(GeneralUtils.class)) { + check.when(CheckProgramInstall::isPythonAvailable).thenReturn(false); + + assertThrows(IOException.class, () -> controller.extractImageScans(request)); + + // Guard runs before document load and before script extraction. + Mockito.verifyNoInteractions(pdfDocumentFactory); + general.verifyNoInteractions(); + } + } + } + + @Nested + @DisplayName("No detected images branch") + class NoImagesBranch { + + @Test + @DisplayName("throws IllegalArgumentException for a non-PDF input when no outputs produced") + void throwsNoImagesForImageInput() throws Exception { + ExtractImageScansRequest request = requestFor(imageFile("scan.png")); + + try (MockedStatic check = + Mockito.mockStatic(CheckProgramInstall.class); + MockedStatic general = Mockito.mockStatic(GeneralUtils.class); + MockedStatic pe = Mockito.mockStatic(ProcessExecutor.class)) { + check.when(CheckProgramInstall::isPythonAvailable).thenReturn(true); + check.when(CheckProgramInstall::getAvailablePythonCommand).thenReturn("python3"); + general.when(() -> GeneralUtils.extractScript("split_photos.py")) + .thenReturn(Path.of("split_photos.py")); + + // Process produces zero output files -> empty result -> "no images" error. + // Build the executor mock BEFORE stubbing the static so its inner when(...) does + // not + // nest inside this when(...).thenReturn(...) call. + ProcessExecutor exec = execWritingOutputs(0); + pe.when(() -> ProcessExecutor.getInstance(ProcessExecutor.Processes.PYTHON_OPENCV)) + .thenReturn(exec); + + assertThrows( + IllegalArgumentException.class, + () -> controller.extractImageScans(request)); + } + } + + @Test + @DisplayName("throws IllegalArgumentException for a PDF input when no outputs produced") + void throwsNoImagesForPdfInput() throws Exception { + ExtractImageScansRequest request = requestFor(pdfFile("scan.pdf")); + + try (MockedStatic check = + Mockito.mockStatic(CheckProgramInstall.class); + MockedStatic general = Mockito.mockStatic(GeneralUtils.class); + MockedStatic pe = Mockito.mockStatic(ProcessExecutor.class)) { + check.when(CheckProgramInstall::isPythonAvailable).thenReturn(true); + check.when(CheckProgramInstall::getAvailablePythonCommand).thenReturn("python3"); + general.when(() -> GeneralUtils.extractScript("split_photos.py")) + .thenReturn(Path.of("split_photos.py")); + when(pdfDocumentFactory.load( + any(org.springframework.web.multipart.MultipartFile.class))) + .thenReturn(singlePageDocument()); + + ProcessExecutor exec = execWritingOutputs(0); + pe.when(() -> ProcessExecutor.getInstance(ProcessExecutor.Processes.PYTHON_OPENCV)) + .thenReturn(exec); + + assertThrows( + IllegalArgumentException.class, + () -> controller.extractImageScans(request)); + + // The PDF path must load the document exactly once. + verify(pdfDocumentFactory, times(1)) + .load(any(org.springframework.web.multipart.MultipartFile.class)); + } + } + } + + @Nested + @DisplayName("Single image output branch") + class SingleImageBranch { + + @Test + @DisplayName("returns a single PNG response when exactly one image is detected") + void returnsSinglePng() throws Exception { + ExtractImageScansRequest request = requestFor(imageFile("scan.png")); + + try (MockedStatic check = + Mockito.mockStatic(CheckProgramInstall.class); + MockedStatic general = Mockito.mockStatic(GeneralUtils.class); + MockedStatic pe = Mockito.mockStatic(ProcessExecutor.class)) { + check.when(CheckProgramInstall::isPythonAvailable).thenReturn(true); + check.when(CheckProgramInstall::getAvailablePythonCommand).thenReturn("python"); + general.when(() -> GeneralUtils.extractScript("split_photos.py")) + .thenReturn(Path.of("split_photos.py")); + general.when(() -> GeneralUtils.generateFilename(anyString(), anyString())) + .thenAnswer(inv -> inv.getArgument(0) + inv.getArgument(1)); + + ProcessExecutor exec = execWritingOutputs(1); + pe.when(() -> ProcessExecutor.getInstance(ProcessExecutor.Processes.PYTHON_OPENCV)) + .thenReturn(exec); + + ResponseEntity response = controller.extractImageScans(request); + + assertNotNull(response); + assertEquals(HttpStatus.OK, response.getStatusCode()); + assertEquals(MediaType.IMAGE_PNG, response.getHeaders().getContentType()); + assertNotNull(response.getBody()); + } + } + } + + @Nested + @DisplayName("Multiple images (zip) output branch") + class ZipBranch { + + @Test + @DisplayName("returns a zip response when more than one image is detected") + void returnsZipForMultipleImages() throws Exception { + ExtractImageScansRequest request = requestFor(imageFile("scan.png")); + + try (MockedStatic check = + Mockito.mockStatic(CheckProgramInstall.class); + MockedStatic general = Mockito.mockStatic(GeneralUtils.class); + MockedStatic pe = Mockito.mockStatic(ProcessExecutor.class)) { + check.when(CheckProgramInstall::isPythonAvailable).thenReturn(true); + check.when(CheckProgramInstall::getAvailablePythonCommand).thenReturn("python3"); + general.when(() -> GeneralUtils.extractScript("split_photos.py")) + .thenReturn(Path.of("split_photos.py")); + general.when(() -> GeneralUtils.generateFilename(anyString(), anyString())) + .thenAnswer(inv -> inv.getArgument(0) + inv.getArgument(1)); + + // Three output files -> zip path. + ProcessExecutor exec = execWritingOutputs(3); + pe.when(() -> ProcessExecutor.getInstance(ProcessExecutor.Processes.PYTHON_OPENCV)) + .thenReturn(exec); + + ResponseEntity response = controller.extractImageScans(request); + + assertNotNull(response); + assertEquals(HttpStatus.OK, response.getStatusCode()); + assertNotNull(response.getBody()); + // generateFilename is invoked for the zip name and once per zip entry. + general.verify( + () -> GeneralUtils.generateFilename(anyString(), anyString()), + atLeastOnce()); + } + } + } + + @Nested + @DisplayName("Command construction") + class CommandConstruction { + + @Test + @DisplayName("passes the request parameters as CLI flags to the executor") + void buildsExpectedCommand() throws Exception { + MockMultipartFile file = imageFile("scan.png"); + ExtractImageScansRequest request = new ExtractImageScansRequest(); + request.setFileInput(file); + request.setAngleThreshold(7); + request.setTolerance(21); + request.setMinArea(9000); + request.setMinContourArea(600); + request.setBorderSize(2); + + ProcessExecutor exec = mock(ProcessExecutor.class); + @SuppressWarnings("unchecked") + ArgumentCaptor> cmdCaptor = ArgumentCaptor.forClass(List.class); + when(exec.runCommandWithOutputHandling(cmdCaptor.capture())) + .thenAnswer(invocation -> mock(ProcessExecutorResult.class)); + + try (MockedStatic check = + Mockito.mockStatic(CheckProgramInstall.class); + MockedStatic general = Mockito.mockStatic(GeneralUtils.class); + MockedStatic pe = Mockito.mockStatic(ProcessExecutor.class)) { + check.when(CheckProgramInstall::isPythonAvailable).thenReturn(true); + check.when(CheckProgramInstall::getAvailablePythonCommand).thenReturn("python3"); + general.when(() -> GeneralUtils.extractScript("split_photos.py")) + .thenReturn(Path.of("split_photos.py")); + pe.when(() -> ProcessExecutor.getInstance(ProcessExecutor.Processes.PYTHON_OPENCV)) + .thenReturn(exec); + + // No outputs are written, so the controller ultimately throws "no images"; + // we only care that the command was built and dispatched first. + assertThrows( + IllegalArgumentException.class, + () -> controller.extractImageScans(request)); + + List command = cmdCaptor.getValue(); + assertNotNull(command); + assertEquals("python3", command.get(0)); + assertTrue(command.contains("--angle_threshold")); + assertEquals("7", valueAfter(command, "--angle_threshold")); + assertEquals("21", valueAfter(command, "--tolerance")); + assertEquals("9000", valueAfter(command, "--min_area")); + assertEquals("600", valueAfter(command, "--min_contour_area")); + assertEquals("2", valueAfter(command, "--border_size")); + } + } + + private String valueAfter(List command, String flag) { + int idx = command.indexOf(flag); + assertTrue(idx >= 0 && idx + 1 < command.size(), "flag " + flag + " not found"); + return command.get(idx + 1); + } + } + + @Nested + @DisplayName("Temp file cleanup") + class Cleanup { + + @Test + @DisplayName("leaves no controller temp files behind after a no-images failure") + void cleansUpAfterFailure() throws Exception { + ExtractImageScansRequest request = requestFor(imageFile("scan.png")); + + try (MockedStatic check = + Mockito.mockStatic(CheckProgramInstall.class); + MockedStatic general = Mockito.mockStatic(GeneralUtils.class); + MockedStatic pe = Mockito.mockStatic(ProcessExecutor.class)) { + check.when(CheckProgramInstall::isPythonAvailable).thenReturn(true); + check.when(CheckProgramInstall::getAvailablePythonCommand).thenReturn("python3"); + general.when(() -> GeneralUtils.extractScript("split_photos.py")) + .thenReturn(Path.of("split_photos.py")); + ProcessExecutor exec = execWritingOutputs(0); + pe.when(() -> ProcessExecutor.getInstance(ProcessExecutor.Processes.PYTHON_OPENCV)) + .thenReturn(exec); + + assertThrows( + IllegalArgumentException.class, + () -> controller.extractImageScans(request)); + + try (var stream = Files.walk(baseTmpDir)) { + boolean leaked = + stream.filter(Files::isRegularFile) + .map(p -> p.getFileName().toString()) + .anyMatch(n -> n.startsWith("scan-test-")); + assertFalse(leaked, "controller temp files should be cleaned up"); + } + } + } + } + + /** Build a fresh single-page in-memory document the factory mock can hand back. */ + private PDDocument singlePageDocument() { + PDDocument doc = new PDDocument(); + doc.addPage(new PDPage(new PDRectangle(72f, 72f))); + return doc; + } +} diff --git a/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/OCRControllerTest.java b/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/OCRControllerTest.java new file mode 100644 index 0000000000..83d036fa6e --- /dev/null +++ b/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/OCRControllerTest.java @@ -0,0 +1,290 @@ +package stirling.software.SPDF.controller.api.misc; + +import static org.junit.jupiter.api.Assertions.*; +import static org.mockito.Mockito.*; + +import java.io.ByteArrayOutputStream; +import java.io.File; +import java.io.IOException; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.Collections; +import java.util.List; + +import org.apache.pdfbox.pdmodel.PDDocument; +import org.apache.pdfbox.pdmodel.PDPage; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.junit.jupiter.api.io.TempDir; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.http.MediaType; +import org.springframework.mock.web.MockMultipartFile; + +import stirling.software.SPDF.config.EndpointConfiguration; +import stirling.software.SPDF.model.api.misc.ProcessPdfWithOcrRequest; +import stirling.software.common.configuration.RuntimePathConfig; +import stirling.software.common.model.ApplicationProperties; +import stirling.software.common.service.CustomPDFDocumentFactory; +import stirling.software.common.util.TempFileManager; +import stirling.software.common.util.TempFileRegistry; + +/** + * Unit tests for {@link OCRController}. OCR shells out to tesseract/ocrmypdf, so these tests focus + * on the pure, deterministic surface: tesseract-language discovery and the request validation / + * tool-availability branches that all complete (or fail) before any external process is launched. + */ +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class OCRControllerTest { + + @Mock private CustomPDFDocumentFactory pdfDocumentFactory; + @Mock private EndpointConfiguration endpointConfiguration; + @Mock private RuntimePathConfig runtimePathConfig; + + private TempFileManager tempFileManager; + private ApplicationProperties applicationProperties; + private OCRController ocrController; + + @TempDir Path baseTmpDir; + + @BeforeEach + void setUp() { + applicationProperties = new ApplicationProperties(); + applicationProperties + .getSystem() + .getTempFileManagement() + .setBaseTmpDir(baseTmpDir.toString()); + applicationProperties.getSystem().getTempFileManagement().setPrefix("ocr-test-"); + + tempFileManager = new TempFileManager(new TempFileRegistry(), applicationProperties); + + ocrController = + new OCRController( + applicationProperties, + pdfDocumentFactory, + tempFileManager, + endpointConfiguration, + runtimePathConfig); + } + + /** Build a minimal request with sensible defaults the caller can override. */ + private ProcessPdfWithOcrRequest baseRequest() { + ProcessPdfWithOcrRequest request = new ProcessPdfWithOcrRequest(); + request.setLanguages(List.of("eng")); + request.setOcrRenderType("hocr"); + request.setOcrType("skip-text"); + return request; + } + + /** Build a tiny single-page in-memory PDF as a MockMultipartFile. */ + private MockMultipartFile pdfMultipartFile(String name) throws IOException { + try (PDDocument doc = new PDDocument(); + ByteArrayOutputStream out = new ByteArrayOutputStream()) { + doc.addPage(new PDPage()); + doc.save(out); + return new MockMultipartFile( + "fileInput", name, MediaType.APPLICATION_PDF_VALUE, out.toByteArray()); + } + } + + /** Create a tessdata directory populated with the given traineddata languages. */ + private Path tessdataDirWith(String... languages) throws IOException { + Path dir = Files.createTempDirectory(baseTmpDir, "tessdata"); + for (String lang : languages) { + Files.createFile(dir.resolve(lang + ".traineddata")); + } + return dir; + } + + @Nested + @DisplayName("getAvailableTesseractLanguages") + class GetAvailableTesseractLanguages { + + @Test + @DisplayName("returns trained languages and excludes osd") + void returnsTrainedLanguagesExcludingOsd() throws IOException { + Path tessdata = tessdataDirWith("eng", "deu", "osd"); + // A non-traineddata file must be ignored entirely. + Files.createFile(tessdata.resolve("readme.txt")); + when(runtimePathConfig.getTessDataPath()).thenReturn(tessdata.toString()); + + List langs = ocrController.getAvailableTesseractLanguages(); + + assertTrue(langs.contains("eng")); + assertTrue(langs.contains("deu")); + assertFalse(langs.contains("osd"), "osd must be filtered out"); + assertFalse(langs.contains("readme"), "non-traineddata files must be ignored"); + assertEquals(2, langs.size()); + } + + @Test + @DisplayName("filters osd case-insensitively") + void filtersOsdCaseInsensitively() throws IOException { + Path tessdata = tessdataDirWith("eng", "OSD"); + when(runtimePathConfig.getTessDataPath()).thenReturn(tessdata.toString()); + + List langs = ocrController.getAvailableTesseractLanguages(); + + assertEquals(List.of("eng"), langs); + } + + @Test + @DisplayName("returns empty list when directory has no traineddata files") + void returnsEmptyWhenNoTrainedData() throws IOException { + Path empty = Files.createTempDirectory(baseTmpDir, "empty-tessdata"); + when(runtimePathConfig.getTessDataPath()).thenReturn(empty.toString()); + + assertTrue(ocrController.getAvailableTesseractLanguages().isEmpty()); + } + + @Test + @DisplayName("returns empty list when directory does not exist") + void returnsEmptyWhenDirectoryMissing() { + Path missing = baseTmpDir.resolve("does-not-exist"); + when(runtimePathConfig.getTessDataPath()).thenReturn(missing.toString()); + + // listFiles() on a non-directory returns null -> empty list, not an exception. + assertTrue(ocrController.getAvailableTesseractLanguages().isEmpty()); + } + } + + @Nested + @DisplayName("processPdfWithOCR validation branches") + class ValidationBranches { + + @Test + @DisplayName("throws when languages list is null") + void throwsWhenLanguagesNull() throws IOException { + ProcessPdfWithOcrRequest request = baseRequest(); + request.setLanguages(null); + request.setFileInput(pdfMultipartFile("in.pdf")); + + assertThrows(IOException.class, () -> ocrController.processPdfWithOCR(request)); + // Validation happens before any tool/exec interaction. + verifyNoInteractions(endpointConfiguration); + } + + @Test + @DisplayName("throws when languages list is empty") + void throwsWhenLanguagesEmpty() throws IOException { + ProcessPdfWithOcrRequest request = baseRequest(); + request.setLanguages(Collections.emptyList()); + request.setFileInput(pdfMultipartFile("in.pdf")); + + assertThrows(IOException.class, () -> ocrController.processPdfWithOCR(request)); + verifyNoInteractions(endpointConfiguration); + } + + @Test + @DisplayName("throws when ocrRenderType is invalid") + void throwsWhenRenderTypeInvalid() throws IOException { + ProcessPdfWithOcrRequest request = baseRequest(); + request.setOcrRenderType("bogus"); + request.setFileInput(pdfMultipartFile("in.pdf")); + + assertThrows(IOException.class, () -> ocrController.processPdfWithOCR(request)); + // Render-type check precedes language availability lookup. + verify(runtimePathConfig, never()).getTessDataPath(); + } + + @Test + @DisplayName("accepts sandwich render type past the render-type check") + void sandwichRenderTypePassesRenderCheck() throws IOException { + ProcessPdfWithOcrRequest request = baseRequest(); + request.setOcrRenderType("sandwich"); + request.setLanguages(List.of("eng")); + request.setFileInput(pdfMultipartFile("in.pdf")); + // No tessdata languages available -> falls through to invalid-languages, still an + // IOException, but proves "sandwich" was not rejected by the render-type guard. + Path tessdata = tessdataDirWith("eng"); + when(runtimePathConfig.getTessDataPath()).thenReturn(tessdata.toString()); + when(endpointConfiguration.isGroupEnabled("OCRmyPDF")).thenReturn(false); + when(endpointConfiguration.isGroupEnabled("tesseract")).thenReturn(false); + + // eng is available, so it gets past language validation and reaches the + // tool-availability check, which throws because both tools are disabled. + assertThrows(IOException.class, () -> ocrController.processPdfWithOCR(request)); + verify(runtimePathConfig).getTessDataPath(); + } + + @Test + @DisplayName("throws when none of the selected languages are available") + void throwsWhenNoSelectedLanguageAvailable() throws IOException { + ProcessPdfWithOcrRequest request = baseRequest(); + request.setLanguages(List.of("xyz")); // not present in tessdata + request.setFileInput(pdfMultipartFile("in.pdf")); + + Path tessdata = tessdataDirWith("eng", "deu"); + when(runtimePathConfig.getTessDataPath()).thenReturn(tessdata.toString()); + + assertThrows(IOException.class, () -> ocrController.processPdfWithOCR(request)); + // Should fail before consulting tool availability. + verify(endpointConfiguration, never()).isGroupEnabled(anyString()); + } + } + + @Nested + @DisplayName("processPdfWithOCR tool-availability branch") + class ToolAvailabilityBranch { + + @Test + @DisplayName("throws when both OCRmyPDF and tesseract are unavailable") + void throwsWhenNoOcrToolsAvailable() throws IOException { + ProcessPdfWithOcrRequest request = baseRequest(); + request.setLanguages(List.of("eng")); + request.setFileInput(pdfMultipartFile("in.pdf")); + + Path tessdata = tessdataDirWith("eng"); + when(runtimePathConfig.getTessDataPath()).thenReturn(tessdata.toString()); + when(endpointConfiguration.isGroupEnabled("OCRmyPDF")).thenReturn(false); + when(endpointConfiguration.isGroupEnabled("tesseract")).thenReturn(false); + + assertThrows(IOException.class, () -> ocrController.processPdfWithOCR(request)); + + verify(endpointConfiguration).isGroupEnabled("OCRmyPDF"); + verify(endpointConfiguration).isGroupEnabled("tesseract"); + } + + @Test + @DisplayName("temp input/output files are cleaned up after a failure") + void tempFilesCleanedUpAfterFailure() throws IOException { + ProcessPdfWithOcrRequest request = baseRequest(); + request.setLanguages(List.of("eng")); + request.setFileInput(pdfMultipartFile("in.pdf")); + + Path tessdata = tessdataDirWith("eng"); + when(runtimePathConfig.getTessDataPath()).thenReturn(tessdata.toString()); + when(endpointConfiguration.isGroupEnabled("OCRmyPDF")).thenReturn(false); + when(endpointConfiguration.isGroupEnabled("tesseract")).thenReturn(false); + + assertThrows(IOException.class, () -> ocrController.processPdfWithOCR(request)); + + // The only files left under the temp dir should be our tessdata dir and its + // contents; the controller's .pdf temp files must have been closed/deleted. + try (var stream = Files.walk(baseTmpDir)) { + boolean leakedPdf = + stream.filter(Files::isRegularFile) + .map(p -> p.getFileName().toString()) + .filter(n -> n.startsWith("ocr-test-")) + .anyMatch(n -> n.endsWith(".pdf")); + assertFalse(leakedPdf, "controller temp PDF files should be cleaned up"); + } + } + } + + @Test + @DisplayName("getAvailableTesseractLanguages survives a path that is a regular file") + void languagesEmptyWhenPathIsRegularFile() throws IOException { + File regularFile = Files.createTempFile(baseTmpDir, "not-a-dir", ".bin").toFile(); + when(runtimePathConfig.getTessDataPath()).thenReturn(regularFile.getAbsolutePath()); + + // listFiles() on a regular file returns null -> empty list. + assertTrue(ocrController.getAvailableTesseractLanguages().isEmpty()); + } +} diff --git a/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/OverlayImageControllerTest.java b/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/OverlayImageControllerTest.java index 308518e8f0..6a13fb3049 100644 --- a/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/OverlayImageControllerTest.java +++ b/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/OverlayImageControllerTest.java @@ -32,6 +32,7 @@ import org.springframework.mock.web.MockMultipartFile; import stirling.software.SPDF.model.api.misc.OverlayImageRequest; import stirling.software.common.service.CustomPDFDocumentFactory; +import stirling.software.common.util.SvgSanitizer; import stirling.software.common.util.TempFile; import stirling.software.common.util.TempFileManager; import stirling.software.common.util.WebResponseUtils; @@ -52,6 +53,7 @@ class OverlayImageControllerTest { @Mock private CustomPDFDocumentFactory pdfDocumentFactory; @Mock private TempFileManager tempFileManager; + @Mock private SvgSanitizer svgSanitizer; @InjectMocks private OverlayImageController controller; @@ -205,6 +207,52 @@ class OverlayImageControllerTest { mockDoc.close(); } + @Test + void overlayImage_svgInput_sanitizedBeforeOverlay() throws Exception { + byte[] maliciousSvg = + ("" + + "" + + "") + .getBytes(); + byte[] sanitized = + ("" + + "" + + "") + .getBytes(); + when(svgSanitizer.sanitize(maliciousSvg)).thenReturn(sanitized); + + MockMultipartFile svgFile = + new MockMultipartFile("imageFile", "overlay.svg", "image/svg+xml", maliciousSvg); + OverlayImageRequest request = new OverlayImageRequest(); + request.setFileInput(pdfFile); + request.setImageFile(svgFile); + request.setX(0); + request.setY(0); + request.setEveryPage(false); + + PDDocument mockDoc = new PDDocument(); + mockDoc.addPage(new PDPage(PDRectangle.A4)); + when(pdfDocumentFactory.load(any(byte[].class))).thenReturn(mockDoc); + + try (MockedStatic mockedWebResponse = + mockStatic(WebResponseUtils.class)) { + mockedWebResponse + .when( + () -> + WebResponseUtils.pdfFileToWebResponse( + any(TempFile.class), anyString())) + .thenReturn(streamingOk("result".getBytes())); + + controller.overlayImage(request); + } + mockDoc.close(); + + verify(svgSanitizer).sanitize(maliciousSvg); + } + @Test void overlayImage_withCoordinates_usesXY() throws Exception { OverlayImageRequest request = new OverlayImageRequest(); diff --git a/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/PageNumbersControllerTest.java b/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/PageNumbersControllerTest.java new file mode 100644 index 0000000000..69f1d4cf7a --- /dev/null +++ b/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/PageNumbersControllerTest.java @@ -0,0 +1,553 @@ +package stirling.software.SPDF.controller.api.misc; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.Mockito.lenient; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.io.ByteArrayOutputStream; +import java.io.File; +import java.io.IOException; +import java.io.InputStream; +import java.nio.file.Files; +import java.nio.file.Path; + +import org.apache.pdfbox.Loader; +import org.apache.pdfbox.pdmodel.PDDocument; +import org.apache.pdfbox.pdmodel.PDPage; +import org.apache.pdfbox.pdmodel.common.PDRectangle; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.junit.jupiter.api.io.TempDir; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.ValueSource; +import org.mockito.InjectMocks; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.springframework.core.io.Resource; +import org.springframework.http.HttpStatus; +import org.springframework.http.MediaType; +import org.springframework.http.ResponseEntity; +import org.springframework.mock.web.MockMultipartFile; + +import stirling.software.SPDF.model.api.misc.AddPageNumbersRequest; +import stirling.software.common.service.CustomPDFDocumentFactory; +import stirling.software.common.util.TempFile; +import stirling.software.common.util.TempFileManager; + +@ExtendWith(MockitoExtension.class) +class PageNumbersControllerTest { + + @TempDir Path tempDir; + + @Mock private CustomPDFDocumentFactory pdfDocumentFactory; + @Mock private TempFileManager tempFileManager; + + @InjectMocks private PageNumbersController controller; + + @BeforeEach + void setUp() throws Exception { + // Each managed temp file is backed by a real on-disk file so document.save() works + // and WebResponseUtils can stat/stream it. + lenient() + .when(tempFileManager.createManagedTempFile(anyString())) + .thenAnswer( + inv -> { + File f = + Files.createTempFile("pgnum-test", inv.getArgument(0)) + .toFile(); + TempFile tf = mock(TempFile.class); + lenient().when(tf.getFile()).thenReturn(f); + lenient().when(tf.getPath()).thenReturn(f.toPath()); + return tf; + }); + } + + // ---- helpers ---------------------------------------------------------- + + private MockMultipartFile createPdf(int pages, String filename) throws IOException { + Path path = tempDir.resolve("source-" + System.nanoTime() + ".pdf"); + try (PDDocument doc = new PDDocument()) { + for (int i = 0; i < pages; i++) { + doc.addPage(new PDPage(PDRectangle.LETTER)); + } + doc.save(path.toFile()); + } + return new MockMultipartFile( + "fileInput", filename, MediaType.APPLICATION_PDF_VALUE, Files.readAllBytes(path)); + } + + private AddPageNumbersRequest baseRequest(MockMultipartFile file) { + AddPageNumbersRequest request = new AddPageNumbersRequest(); + request.setFileInput(file); + request.setFontSize(12f); + request.setFontType("helvetica"); + request.setPosition(8); + request.setStartingNumber(1); + return request; + } + + private byte[] drainBody(ResponseEntity response) throws IOException { + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + try (InputStream in = response.getBody().getInputStream()) { + in.transferTo(baos); + } + return baos.toByteArray(); + } + + // ---- happy path ------------------------------------------------------- + + @Nested + @DisplayName("Happy path") + class HappyPath { + + @Test + @DisplayName("Single-page PDF returns OK with a non-empty PDF body") + void singlePage_returnsOkWithBody() throws Exception { + MockMultipartFile file = createPdf(1, "doc.pdf"); + AddPageNumbersRequest request = baseRequest(file); + when(pdfDocumentFactory.load(file)).thenReturn(Loader.loadPDF(file.getBytes())); + + ResponseEntity response = controller.addPageNumbers(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(response.getBody()).isNotNull(); + byte[] body = drainBody(response); + assertThat(body).isNotEmpty(); + // Result is still a valid, single-page PDF. + try (PDDocument out = Loader.loadPDF(body)) { + assertThat(out.getNumberOfPages()).isEqualTo(1); + } + } + + @Test + @DisplayName("Multi-page PDF with default 'all' pages numbers every page") + void multiPage_allPages() throws Exception { + MockMultipartFile file = createPdf(5, "multi.pdf"); + AddPageNumbersRequest request = baseRequest(file); + PDDocument doc = Loader.loadPDF(file.getBytes()); + when(pdfDocumentFactory.load(file)).thenReturn(doc); + + ResponseEntity response = controller.addPageNumbers(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + try (PDDocument out = Loader.loadPDF(drainBody(response))) { + assertThat(out.getNumberOfPages()).isEqualTo(5); + } + // The loaded document is closed by the try-with-resources in the controller. + verify(tempFileManager).createManagedTempFile(".pdf"); + } + + @Test + @DisplayName("Content-Disposition attachment filename carries the source name") + void responseHasAttachmentFilename() throws Exception { + MockMultipartFile file = createPdf(1, "report.pdf"); + AddPageNumbersRequest request = baseRequest(file); + when(pdfDocumentFactory.load(file)).thenReturn(Loader.loadPDF(file.getBytes())); + + ResponseEntity response = controller.addPageNumbers(request); + + String disposition = response.getHeaders().getFirst("Content-Disposition"); + assertThat(disposition).isNotNull(); + assertThat(disposition).contains("report_page_numbers_added.pdf"); + assertThat(response.getHeaders().getContentType()).isEqualTo(MediaType.APPLICATION_PDF); + } + } + + // ---- pages-to-number selection ---------------------------------------- + + @Nested + @DisplayName("Page selection") + class PageSelection { + + @Test + @DisplayName("Explicit subset of pages still returns all pages in output") + void specificPages() throws Exception { + MockMultipartFile file = createPdf(4, "sel.pdf"); + AddPageNumbersRequest request = baseRequest(file); + request.setPagesToNumber("1,3"); + when(pdfDocumentFactory.load(file)).thenReturn(Loader.loadPDF(file.getBytes())); + + ResponseEntity response = controller.addPageNumbers(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + try (PDDocument out = Loader.loadPDF(drainBody(response))) { + assertThat(out.getNumberOfPages()).isEqualTo(4); + } + } + + @Test + @DisplayName("Range expression is accepted") + void rangeExpression() throws Exception { + MockMultipartFile file = createPdf(6, "range.pdf"); + AddPageNumbersRequest request = baseRequest(file); + request.setPagesToNumber("2-4"); + when(pdfDocumentFactory.load(file)).thenReturn(Loader.loadPDF(file.getBytes())); + + ResponseEntity response = controller.addPageNumbers(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(drainBody(response)).isNotEmpty(); + } + + @Test + @DisplayName("Null pagesToNumber defaults to 'all'") + void nullPagesDefaultsToAll() throws Exception { + MockMultipartFile file = createPdf(2, "nullpages.pdf"); + AddPageNumbersRequest request = baseRequest(file); + request.setPagesToNumber(null); + when(pdfDocumentFactory.load(file)).thenReturn(Loader.loadPDF(file.getBytes())); + + ResponseEntity response = controller.addPageNumbers(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(drainBody(response)).isNotEmpty(); + } + + @Test + @DisplayName("Empty pagesToNumber defaults to 'all'") + void emptyPagesDefaultsToAll() throws Exception { + MockMultipartFile file = createPdf(2, "emptypages.pdf"); + AddPageNumbersRequest request = baseRequest(file); + request.setPagesToNumber(""); + when(pdfDocumentFactory.load(file)).thenReturn(Loader.loadPDF(file.getBytes())); + + ResponseEntity response = controller.addPageNumbers(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(drainBody(response)).isNotEmpty(); + } + } + + // ---- position handling (1..9 plus clamping) --------------------------- + + @Nested + @DisplayName("Position") + class Position { + + @ParameterizedTest + @ValueSource(ints = {1, 2, 3, 4, 5, 6, 7, 8, 9}) + @DisplayName("All nine positions render successfully") + void allPositions(int position) throws Exception { + MockMultipartFile file = createPdf(1, "pos.pdf"); + AddPageNumbersRequest request = baseRequest(file); + request.setPosition(position); + when(pdfDocumentFactory.load(file)).thenReturn(Loader.loadPDF(file.getBytes())); + + ResponseEntity response = controller.addPageNumbers(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(drainBody(response)).isNotEmpty(); + } + + @ParameterizedTest + @ValueSource(ints = {-5, 0, 10, 100}) + @DisplayName("Out-of-range positions are clamped and still render") + void outOfRangePositionsClamped(int position) throws Exception { + MockMultipartFile file = createPdf(1, "posclamp.pdf"); + AddPageNumbersRequest request = baseRequest(file); + request.setPosition(position); + when(pdfDocumentFactory.load(file)).thenReturn(Loader.loadPDF(file.getBytes())); + + ResponseEntity response = controller.addPageNumbers(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(drainBody(response)).isNotEmpty(); + } + } + + // ---- margins ---------------------------------------------------------- + + @Nested + @DisplayName("Custom margin") + class CustomMargin { + + @ParameterizedTest + @ValueSource(strings = {"small", "medium", "large", "x-large", "X-LARGE", "unknown"}) + @DisplayName("Known and unknown margins are accepted") + void margins(String margin) throws Exception { + MockMultipartFile file = createPdf(1, "margin.pdf"); + AddPageNumbersRequest request = baseRequest(file); + request.setCustomMargin(margin); + when(pdfDocumentFactory.load(file)).thenReturn(Loader.loadPDF(file.getBytes())); + + ResponseEntity response = controller.addPageNumbers(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(drainBody(response)).isNotEmpty(); + } + + @Test + @DisplayName("Null margin falls back to default factor") + void nullMargin() throws Exception { + MockMultipartFile file = createPdf(1, "nullmargin.pdf"); + AddPageNumbersRequest request = baseRequest(file); + request.setCustomMargin(null); + when(pdfDocumentFactory.load(file)).thenReturn(Loader.loadPDF(file.getBytes())); + + ResponseEntity response = controller.addPageNumbers(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(drainBody(response)).isNotEmpty(); + } + } + + // ---- font type -------------------------------------------------------- + + @Nested + @DisplayName("Font type") + class FontType { + + @ParameterizedTest + @ValueSource(strings = {"helvetica", "courier", "times", "TIMES", "anythingelse"}) + @DisplayName("Known and unknown font types render (unknown falls back to Helvetica)") + void fontTypes(String font) throws Exception { + MockMultipartFile file = createPdf(1, "font.pdf"); + AddPageNumbersRequest request = baseRequest(file); + request.setFontType(font); + when(pdfDocumentFactory.load(file)).thenReturn(Loader.loadPDF(file.getBytes())); + + ResponseEntity response = controller.addPageNumbers(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(drainBody(response)).isNotEmpty(); + } + + @Test + @DisplayName("Null font type falls back to Helvetica") + void nullFontType() throws Exception { + MockMultipartFile file = createPdf(1, "nullfont.pdf"); + AddPageNumbersRequest request = baseRequest(file); + request.setFontType(null); + when(pdfDocumentFactory.load(file)).thenReturn(Loader.loadPDF(file.getBytes())); + + ResponseEntity response = controller.addPageNumbers(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(drainBody(response)).isNotEmpty(); + } + } + + // ---- font color ------------------------------------------------------- + + @Nested + @DisplayName("Font color") + class FontColor { + + @Test + @DisplayName("Valid hex color renders") + void validHexColor() throws Exception { + MockMultipartFile file = createPdf(1, "color.pdf"); + AddPageNumbersRequest request = baseRequest(file); + request.setFontColor("#FF0000"); + when(pdfDocumentFactory.load(file)).thenReturn(Loader.loadPDF(file.getBytes())); + + ResponseEntity response = controller.addPageNumbers(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(drainBody(response)).isNotEmpty(); + } + + @Test + @DisplayName("Invalid hex color falls back to black and still renders") + void invalidHexColorFallsBackToBlack() throws Exception { + MockMultipartFile file = createPdf(1, "badcolor.pdf"); + AddPageNumbersRequest request = baseRequest(file); + request.setFontColor("not-a-color"); + when(pdfDocumentFactory.load(file)).thenReturn(Loader.loadPDF(file.getBytes())); + + ResponseEntity response = controller.addPageNumbers(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(drainBody(response)).isNotEmpty(); + } + + @Test + @DisplayName("Null font color uses default black") + void nullColor() throws Exception { + MockMultipartFile file = createPdf(1, "nullcolor.pdf"); + AddPageNumbersRequest request = baseRequest(file); + request.setFontColor(null); + when(pdfDocumentFactory.load(file)).thenReturn(Loader.loadPDF(file.getBytes())); + + ResponseEntity response = controller.addPageNumbers(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(drainBody(response)).isNotEmpty(); + } + + @Test + @DisplayName("Blank/whitespace font color uses default black") + void blankColor() throws Exception { + MockMultipartFile file = createPdf(1, "blankcolor.pdf"); + AddPageNumbersRequest request = baseRequest(file); + request.setFontColor(" "); + when(pdfDocumentFactory.load(file)).thenReturn(Loader.loadPDF(file.getBytes())); + + ResponseEntity response = controller.addPageNumbers(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(drainBody(response)).isNotEmpty(); + } + } + + // ---- custom text / placeholders -------------------------------------- + + @Nested + @DisplayName("Custom text") + class CustomText { + + @Test + @DisplayName("Null custom text defaults to {n}") + void nullCustomText() throws Exception { + MockMultipartFile file = createPdf(2, "nulltext.pdf"); + AddPageNumbersRequest request = baseRequest(file); + request.setCustomText(null); + when(pdfDocumentFactory.load(file)).thenReturn(Loader.loadPDF(file.getBytes())); + + ResponseEntity response = controller.addPageNumbers(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(drainBody(response)).isNotEmpty(); + } + + @Test + @DisplayName("Empty custom text defaults to {n}") + void emptyCustomText() throws Exception { + MockMultipartFile file = createPdf(2, "emptytext.pdf"); + AddPageNumbersRequest request = baseRequest(file); + request.setCustomText(""); + when(pdfDocumentFactory.load(file)).thenReturn(Loader.loadPDF(file.getBytes())); + + ResponseEntity response = controller.addPageNumbers(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(drainBody(response)).isNotEmpty(); + } + + @Test + @DisplayName("Custom text with {n}, {total} and {filename} placeholders renders") + void placeholders() throws Exception { + MockMultipartFile file = createPdf(3, "myfile.pdf"); + AddPageNumbersRequest request = baseRequest(file); + request.setCustomText("Page {n} of {total} - {filename}"); + when(pdfDocumentFactory.load(file)).thenReturn(Loader.loadPDF(file.getBytes())); + + ResponseEntity response = controller.addPageNumbers(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + try (PDDocument out = Loader.loadPDF(drainBody(response))) { + assertThat(out.getNumberOfPages()).isEqualTo(3); + } + } + + @Test + @DisplayName("Literal custom text without placeholders renders") + void literalText() throws Exception { + MockMultipartFile file = createPdf(1, "literal.pdf"); + AddPageNumbersRequest request = baseRequest(file); + request.setCustomText("Confidential"); + when(pdfDocumentFactory.load(file)).thenReturn(Loader.loadPDF(file.getBytes())); + + ResponseEntity response = controller.addPageNumbers(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(drainBody(response)).isNotEmpty(); + } + } + + // ---- numbering / zero-pad / starting number --------------------------- + + @Nested + @DisplayName("Numbering") + class Numbering { + + @Test + @DisplayName("Zero-pad width produces Bates-style padded numbers without error") + void zeroPadBatesStamping() throws Exception { + MockMultipartFile file = createPdf(3, "bates.pdf"); + AddPageNumbersRequest request = baseRequest(file); + request.setZeroPad(5); + when(pdfDocumentFactory.load(file)).thenReturn(Loader.loadPDF(file.getBytes())); + + ResponseEntity response = controller.addPageNumbers(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(drainBody(response)).isNotEmpty(); + } + + @Test + @DisplayName("Zero zero-pad uses unpadded numbers") + void zeroPadDisabled() throws Exception { + MockMultipartFile file = createPdf(2, "nopad.pdf"); + AddPageNumbersRequest request = baseRequest(file); + request.setZeroPad(0); + when(pdfDocumentFactory.load(file)).thenReturn(Loader.loadPDF(file.getBytes())); + + ResponseEntity response = controller.addPageNumbers(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(drainBody(response)).isNotEmpty(); + } + + @Test + @DisplayName("Custom starting number is honored") + void customStartingNumber() throws Exception { + MockMultipartFile file = createPdf(3, "start.pdf"); + AddPageNumbersRequest request = baseRequest(file); + request.setStartingNumber(100); + when(pdfDocumentFactory.load(file)).thenReturn(Loader.loadPDF(file.getBytes())); + + ResponseEntity response = controller.addPageNumbers(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(drainBody(response)).isNotEmpty(); + } + } + + // ---- filename handling for {filename} --------------------------------- + + @Nested + @DisplayName("Filename handling") + class FilenameHandling { + + @Test + @DisplayName("Filename without extension is handled for {filename} placeholder") + void filenameWithoutExtension() throws Exception { + MockMultipartFile file = createPdf(1, "no_extension"); + AddPageNumbersRequest request = baseRequest(file); + request.setCustomText("{filename}"); + when(pdfDocumentFactory.load(file)).thenReturn(Loader.loadPDF(file.getBytes())); + + ResponseEntity response = controller.addPageNumbers(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(drainBody(response)).isNotEmpty(); + } + } + + // ---- error branches --------------------------------------------------- + + @Nested + @DisplayName("Error handling") + class ErrorHandling { + + @Test + @DisplayName("IOException from document load propagates") + void loadIOExceptionPropagates() throws Exception { + MockMultipartFile file = createPdf(1, "err.pdf"); + AddPageNumbersRequest request = baseRequest(file); + when(pdfDocumentFactory.load(file)).thenThrow(new IOException("corrupt pdf")); + + assertThatThrownBy(() -> controller.addPageNumbers(request)) + .isInstanceOf(IOException.class) + .hasMessageContaining("corrupt pdf"); + } + } +} diff --git a/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/PrintFileControllerTest.java b/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/PrintFileControllerTest.java index 5fb8d4ce8b..7fe1027913 100644 --- a/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/PrintFileControllerTest.java +++ b/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/PrintFileControllerTest.java @@ -3,7 +3,7 @@ package stirling.software.SPDF.controller.api.misc; import static org.junit.jupiter.api.Assertions.*; import java.io.IOException; -import java.nio.file.Paths; +import java.nio.file.Path; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.extension.ExtendWith; @@ -34,9 +34,9 @@ class PrintFileControllerTest { @Test void printFile_absolutePath_throwsException() { PrintFileRequest request = new PrintFileRequest(); - String absPath = Paths.get("/etc/passwd").toString(); + String absPath = Path.of("/etc/passwd").toString(); // Only test on systems where /etc/passwd is absolute - if (Paths.get(absPath).isAbsolute()) { + if (Path.of(absPath).isAbsolute()) { MockMultipartFile file = new MockMultipartFile( "fileInput", absPath, "application/pdf", "data".getBytes()); diff --git a/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/RemoveImagesControllerTest.java b/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/RemoveImagesControllerTest.java new file mode 100644 index 0000000000..1a97cc1a55 --- /dev/null +++ b/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/RemoveImagesControllerTest.java @@ -0,0 +1,447 @@ +package stirling.software.SPDF.controller.api.misc; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.Mockito.lenient; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.mockStatic; +import static org.mockito.Mockito.when; + +import java.awt.image.BufferedImage; +import java.io.File; +import java.io.IOException; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.ArrayList; +import java.util.List; + +import org.apache.pdfbox.Loader; +import org.apache.pdfbox.cos.COSName; +import org.apache.pdfbox.pdmodel.PDDocument; +import org.apache.pdfbox.pdmodel.PDPage; +import org.apache.pdfbox.pdmodel.PDPageContentStream; +import org.apache.pdfbox.pdmodel.PDResources; +import org.apache.pdfbox.pdmodel.common.PDRectangle; +import org.apache.pdfbox.pdmodel.graphics.PDXObject; +import org.apache.pdfbox.pdmodel.graphics.form.PDFormXObject; +import org.apache.pdfbox.pdmodel.graphics.image.JPEGFactory; +import org.apache.pdfbox.pdmodel.graphics.image.PDImageXObject; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.junit.jupiter.api.io.TempDir; +import org.mockito.MockedStatic; +import org.mockito.junit.jupiter.MockitoExtension; +import org.springframework.core.io.ByteArrayResource; +import org.springframework.core.io.Resource; +import org.springframework.http.HttpStatus; +import org.springframework.http.MediaType; +import org.springframework.http.ResponseEntity; +import org.springframework.mock.web.MockMultipartFile; + +import stirling.software.common.model.api.PDFFile; +import stirling.software.common.service.CustomPDFDocumentFactory; +import stirling.software.common.util.TempFile; +import stirling.software.common.util.TempFileManager; +import stirling.software.common.util.WebResponseUtils; + +@ExtendWith(MockitoExtension.class) +@DisplayName("RemoveImagesController") +class RemoveImagesControllerTest { + + @TempDir Path tempDir; + + private CustomPDFDocumentFactory pdfDocumentFactory; + private TempFileManager tempFileManager; + private RemoveImagesController controller; + + // Holds the temp file backing the most recent createManagedTempFile call so tests can + // re-load the actually-saved (image-stripped) document from disk and assert on it. + private final List savedTempFiles = new ArrayList<>(); + + @BeforeEach + void setUp() throws IOException { + pdfDocumentFactory = mock(CustomPDFDocumentFactory.class); + tempFileManager = mock(TempFileManager.class); + controller = new RemoveImagesController(pdfDocumentFactory, tempFileManager); + savedTempFiles.clear(); + + lenient() + .when(tempFileManager.createManagedTempFile(anyString())) + .thenAnswer( + inv -> { + File f = + Files.createTempFile(tempDir, "out", inv.getArgument(0)) + .toFile(); + savedTempFiles.add(f); + TempFile tf = mock(TempFile.class); + lenient().when(tf.getFile()).thenReturn(f); + lenient().when(tf.getPath()).thenReturn(f.toPath()); + return tf; + }); + } + + // ----- helpers ------------------------------------------------------------------------- + + private MockMultipartFile multipart(String name, byte[] bytes) { + return new MockMultipartFile("fileInput", name, MediaType.APPLICATION_PDF_VALUE, bytes); + } + + private PDFFile request(MockMultipartFile file) { + PDFFile req = new PDFFile(); + req.setFileInput(file); + return req; + } + + /** A drawable RGB image with no useful content, kept tiny for speed. */ + private BufferedImage tinyImage() { + return new BufferedImage(8, 8, BufferedImage.TYPE_INT_RGB); + } + + /** PDF with {@code pageCount} pages, each with one drawn JPEG image. */ + private byte[] pdfWithImagesBytes(int pageCount) throws IOException { + try (PDDocument doc = new PDDocument()) { + for (int i = 0; i < pageCount; i++) { + PDPage page = new PDPage(PDRectangle.LETTER); + doc.addPage(page); + PDImageXObject image = JPEGFactory.createFromImage(doc, tinyImage()); + try (PDPageContentStream cs = new PDPageContentStream(doc, page)) { + cs.drawImage(image, 50, 600, 50, 50); + } + } + return saveToBytes(doc); + } + } + + /** PDF with one page that has no images at all (just resources without XObjects). */ + private byte[] pdfWithoutImagesBytes() throws IOException { + try (PDDocument doc = new PDDocument()) { + PDPage page = new PDPage(PDRectangle.LETTER); + doc.addPage(page); + // Touch the page resources so they are non-null but contain no XObject dictionary. + page.setResources(new PDResources()); + return saveToBytes(doc); + } + } + + /** PDF whose single page has a resources dictionary with an image nested in a form XObject. */ + private byte[] pdfWithImageInsideFormBytes() throws IOException { + try (PDDocument doc = new PDDocument()) { + PDPage page = new PDPage(PDRectangle.LETTER); + doc.addPage(page); + + PDImageXObject image = JPEGFactory.createFromImage(doc, tinyImage()); + + PDFormXObject form = new PDFormXObject(doc); + form.setBBox(new PDRectangle(100, 100)); + PDResources formResources = new PDResources(); + formResources.add(image); + form.setResources(formResources); + + PDResources pageResources = new PDResources(); + pageResources.add(form); + page.setResources(pageResources); + + return saveToBytes(doc); + } + } + + private byte[] saveToBytes(PDDocument doc) throws IOException { + java.io.ByteArrayOutputStream baos = new java.io.ByteArrayOutputStream(); + doc.save(baos); + return baos.toByteArray(); + } + + /** Counts every PDImageXObject reachable through page + nested form resources. */ + private int countImagesInSavedOutput() throws IOException { + assertFalse(savedTempFiles.isEmpty(), "expected the controller to create a temp file"); + File out = savedTempFiles.get(savedTempFiles.size() - 1); + try (PDDocument doc = Loader.loadPDF(out)) { + int count = 0; + for (PDPage page : doc.getPages()) { + count += countImagesInResources(page.getResources()); + } + return count; + } + } + + private int countImagesInResources(PDResources resources) throws IOException { + if (resources == null || resources.getXObjectNames() == null) { + return 0; + } + int count = 0; + for (COSName name : resources.getXObjectNames()) { + PDXObject xObject = resources.getXObject(name); + if (xObject instanceof PDImageXObject) { + count++; + } else if (xObject instanceof PDFormXObject form) { + count += countImagesInResources(form.getResources()); + } + } + return count; + } + + // ----- happy paths --------------------------------------------------------------------- + + @Nested + @DisplayName("removeImages happy path") + class HappyPath { + + @Test + @DisplayName("returns OK and strips the image from a single-page PDF") + void singlePageWithImage() throws IOException { + byte[] bytes = pdfWithImagesBytes(1); + MockMultipartFile file = multipart("doc.pdf", bytes); + PDFFile req = request(file); + + PDDocument loaded = Loader.loadPDF(bytes); + when(pdfDocumentFactory.load(req)).thenReturn(loaded); + + // sanity: the input genuinely had one image to remove + assertEquals(1, countImagesInResources(loaded.getPage(0).getResources())); + + ResponseEntity response; + try (MockedStatic web = mockStatic(WebResponseUtils.class)) { + web.when( + () -> + WebResponseUtils.pdfFileToWebResponse( + any(TempFile.class), anyString())) + .thenReturn(okResponse()); + response = controller.removeImages(req); + } + + assertNotNull(response); + assertEquals(HttpStatus.OK, response.getStatusCode()); + assertEquals(0, countImagesInSavedOutput(), "all images should have been removed"); + } + + @Test + @DisplayName("strips images across every page of a multi-page PDF") + void multiPageWithImages() throws IOException { + byte[] bytes = pdfWithImagesBytes(3); + MockMultipartFile file = multipart("multi.pdf", bytes); + PDFFile req = request(file); + + PDDocument loaded = Loader.loadPDF(bytes); + assertEquals(3, loaded.getNumberOfPages()); + when(pdfDocumentFactory.load(req)).thenReturn(loaded); + + try (MockedStatic web = mockStatic(WebResponseUtils.class)) { + web.when( + () -> + WebResponseUtils.pdfFileToWebResponse( + any(TempFile.class), anyString())) + .thenReturn(okResponse()); + ResponseEntity response = controller.removeImages(req); + assertEquals(HttpStatus.OK, response.getStatusCode()); + } + + assertEquals(0, countImagesInSavedOutput()); + // page count must be preserved + File out = savedTempFiles.get(savedTempFiles.size() - 1); + try (PDDocument result = Loader.loadPDF(out)) { + assertEquals(3, result.getNumberOfPages()); + } + } + + @Test + @DisplayName("removes an image nested inside a form XObject") + void imageNestedInsideForm() throws IOException { + byte[] bytes = pdfWithImageInsideFormBytes(); + MockMultipartFile file = multipart("nested.pdf", bytes); + PDFFile req = request(file); + + PDDocument loaded = Loader.loadPDF(bytes); + when(pdfDocumentFactory.load(req)).thenReturn(loaded); + + // sanity: the nested image is present before removal + assertEquals(1, countImagesInResources(loaded.getPage(0).getResources())); + + try (MockedStatic web = mockStatic(WebResponseUtils.class)) { + web.when( + () -> + WebResponseUtils.pdfFileToWebResponse( + any(TempFile.class), anyString())) + .thenReturn(okResponse()); + controller.removeImages(req); + } + + assertEquals(0, countImagesInSavedOutput(), "nested form image should be removed"); + } + } + + // ----- edge cases ---------------------------------------------------------------------- + + @Nested + @DisplayName("removeImages edge cases") + class EdgeCases { + + @Test + @DisplayName("returns OK when the PDF has no images") + void noImages() throws IOException { + byte[] bytes = pdfWithoutImagesBytes(); + MockMultipartFile file = multipart("plain.pdf", bytes); + PDFFile req = request(file); + + PDDocument loaded = Loader.loadPDF(bytes); + when(pdfDocumentFactory.load(req)).thenReturn(loaded); + + try (MockedStatic web = mockStatic(WebResponseUtils.class)) { + web.when( + () -> + WebResponseUtils.pdfFileToWebResponse( + any(TempFile.class), anyString())) + .thenReturn(okResponse()); + ResponseEntity response = controller.removeImages(req); + assertEquals(HttpStatus.OK, response.getStatusCode()); + } + + assertEquals(0, countImagesInSavedOutput()); + } + + @Test + @DisplayName("handles a page that has no resources dictionary") + void pageWithoutResources() throws IOException { + // A bare PDPage built without content has null resources. + byte[] bytes; + try (PDDocument doc = new PDDocument()) { + doc.addPage(new PDPage(PDRectangle.LETTER)); + bytes = saveToBytes(doc); + } + MockMultipartFile file = multipart("bare.pdf", bytes); + PDFFile req = request(file); + + PDDocument loaded = Loader.loadPDF(bytes); + when(pdfDocumentFactory.load(req)).thenReturn(loaded); + + try (MockedStatic web = mockStatic(WebResponseUtils.class)) { + web.when( + () -> + WebResponseUtils.pdfFileToWebResponse( + any(TempFile.class), anyString())) + .thenReturn(okResponse()); + ResponseEntity response = controller.removeImages(req); + assertEquals(HttpStatus.OK, response.getStatusCode()); + } + + assertEquals(0, countImagesInSavedOutput()); + } + + @Test + @DisplayName("creates exactly one managed temp file and saves into it") + void savesIntoManagedTempFile() throws IOException { + byte[] bytes = pdfWithImagesBytes(1); + MockMultipartFile file = multipart("doc.pdf", bytes); + PDFFile req = request(file); + + when(pdfDocumentFactory.load(req)).thenReturn(Loader.loadPDF(bytes)); + + try (MockedStatic web = mockStatic(WebResponseUtils.class)) { + web.when( + () -> + WebResponseUtils.pdfFileToWebResponse( + any(TempFile.class), anyString())) + .thenReturn(okResponse()); + controller.removeImages(req); + } + + assertEquals(1, savedTempFiles.size()); + File out = savedTempFiles.get(0); + assertTrue(out.exists()); + assertTrue(out.length() > 0, "saved PDF must be non-empty"); + } + } + + // ----- filename / interaction ---------------------------------------------------------- + + @Nested + @DisplayName("filename handling") + class Filenames { + + @Test + @DisplayName("appends _images_removed.pdf suffix to the original name") + void appendsSuffix() throws IOException { + byte[] bytes = pdfWithImagesBytes(1); + MockMultipartFile file = multipart("report.pdf", bytes); + PDFFile req = request(file); + + when(pdfDocumentFactory.load(req)).thenReturn(Loader.loadPDF(bytes)); + + try (MockedStatic web = mockStatic(WebResponseUtils.class)) { + web.when( + () -> + WebResponseUtils.pdfFileToWebResponse( + any(TempFile.class), anyString())) + .thenReturn(okResponse()); + + controller.removeImages(req); + + web.verify( + () -> + WebResponseUtils.pdfFileToWebResponse( + any(TempFile.class), + org.mockito.ArgumentMatchers.eq( + "report_images_removed.pdf"))); + } + } + + @Test + @DisplayName("derives a default name when the original filename is null") + void nullOriginalFilename() throws IOException { + byte[] bytes = pdfWithImagesBytes(1); + // MockMultipartFile with a null original filename + MockMultipartFile file = + new MockMultipartFile( + "fileInput", null, MediaType.APPLICATION_PDF_VALUE, bytes); + PDFFile req = request(file); + + when(pdfDocumentFactory.load(req)).thenReturn(Loader.loadPDF(bytes)); + + try (MockedStatic web = mockStatic(WebResponseUtils.class)) { + web.when( + () -> + WebResponseUtils.pdfFileToWebResponse( + any(TempFile.class), anyString())) + .thenReturn(okResponse()); + + ResponseEntity response = controller.removeImages(req); + assertEquals(HttpStatus.OK, response.getStatusCode()); + + // GeneralUtils.generateFilename handles null safely; just assert it was called + web.verify( + () -> + WebResponseUtils.pdfFileToWebResponse( + any(TempFile.class), anyString())); + } + } + } + + // ----- error branches ------------------------------------------------------------------ + + @Nested + @DisplayName("error handling") + class Errors { + + @Test + @DisplayName("propagates an IOException when loading the PDF fails") + void loadFailureThrowsIOException() throws IOException { + MockMultipartFile file = multipart("broken.pdf", "not a pdf".getBytes()); + PDFFile req = request(file); + + when(pdfDocumentFactory.load(req)).thenThrow(new IOException("corrupt")); + + assertThrows(IOException.class, () -> controller.removeImages(req)); + } + } + + private static ResponseEntity okResponse() { + return ResponseEntity.ok(new ByteArrayResource("ok".getBytes())); + } +} diff --git a/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/RepairControllerTest.java b/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/RepairControllerTest.java new file mode 100644 index 0000000000..71fdf5ffc2 --- /dev/null +++ b/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/RepairControllerTest.java @@ -0,0 +1,318 @@ +package stirling.software.SPDF.controller.api.misc; + +import static org.junit.jupiter.api.Assertions.*; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.times; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.io.ByteArrayOutputStream; +import java.io.File; +import java.io.IOException; +import java.io.InputStream; + +import org.apache.pdfbox.pdmodel.PDDocument; +import org.apache.pdfbox.pdmodel.PDPage; +import org.apache.pdfbox.pdmodel.common.PDRectangle; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.core.io.Resource; +import org.springframework.http.HttpHeaders; +import org.springframework.http.HttpStatus; +import org.springframework.http.MediaType; +import org.springframework.http.ResponseEntity; +import org.springframework.mock.web.MockMultipartFile; + +import stirling.software.SPDF.config.EndpointConfiguration; +import stirling.software.common.model.ApplicationProperties; +import stirling.software.common.model.api.PDFFile; +import stirling.software.common.service.CustomPDFDocumentFactory; +import stirling.software.common.util.TempFileManager; +import stirling.software.common.util.TempFileRegistry; + +/** + * Unit tests for {@link RepairController}. + * + *

The controller delegates to external binaries (Ghostscript, qpdf) for its primary repair + * paths. Those paths shell out via the static {@code ProcessExecutor} factory and are therefore not + * deterministically testable in a unit test. These tests keep both Ghostscript and qpdf disabled so + * the controller always takes the pure-Java PDFBox last-resort branch, exercising file handling, + * filename generation, the success response, and the error-propagation branches without spawning + * any external process. + */ +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class RepairControllerTest { + + @Mock private CustomPDFDocumentFactory pdfDocumentFactory; + + @Mock private EndpointConfiguration endpointConfiguration; + + // Real TempFileManager so transferTo / temp file creation / file-backed response all work + // end-to-end deterministically without touching any external tooling. + private TempFileManager tempFileManager; + + private RepairController repairController; + + @BeforeEach + void setUp() { + tempFileManager = new TempFileManager(new TempFileRegistry(), new ApplicationProperties()); + repairController = + new RepairController(pdfDocumentFactory, tempFileManager, endpointConfiguration); + + // Default: no external tools available -> forces the PDFBox last-resort branch. + when(endpointConfiguration.isGroupEnabled("Ghostscript")).thenReturn(false); + when(endpointConfiguration.isGroupEnabled("qpdf")).thenReturn(false); + } + + /** Build a tiny, valid in-memory PDF as bytes. */ + private static byte[] buildPdfBytes(int pageCount) throws IOException { + try (PDDocument document = new PDDocument(); + ByteArrayOutputStream baos = new ByteArrayOutputStream()) { + for (int i = 0; i < pageCount; i++) { + document.addPage(new PDPage(PDRectangle.A4)); + } + document.save(baos); + return baos.toByteArray(); + } + } + + /** A fresh real in-memory document for stubbing pdfDocumentFactory.load(). */ + private static PDDocument newRealDocument(int pageCount) { + PDDocument document = new PDDocument(); + for (int i = 0; i < pageCount; i++) { + document.addPage(new PDPage(PDRectangle.A4)); + } + return document; + } + + private static PDFFile pdfFileFrom(MockMultipartFile multipartFile) { + PDFFile pdfFile = new PDFFile(); + pdfFile.setFileInput(multipartFile); + return pdfFile; + } + + /** Read the body of a file-backed Resource response into bytes. */ + private static byte[] readResource(Resource resource) throws IOException { + try (InputStream in = resource.getInputStream(); + ByteArrayOutputStream baos = new ByteArrayOutputStream()) { + in.transferTo(baos); + return baos.toByteArray(); + } + } + + @Nested + @DisplayName("PDFBox last-resort branch (no external tools)") + class PdfBoxBranch { + + @Test + @DisplayName("returns 200 with a non-empty PDF resource body") + void repairPdf_pdfBoxFallback_returnsOkWithPdf() throws Exception { + MockMultipartFile input = + new MockMultipartFile( + "fileInput", + "broken.pdf", + MediaType.APPLICATION_PDF_VALUE, + buildPdfBytes(2)); + + when(pdfDocumentFactory.load(any(File.class))).thenReturn(newRealDocument(2)); + + ResponseEntity response = repairController.repairPdf(pdfFileFrom(input)); + + assertNotNull(response); + assertEquals(HttpStatus.OK, response.getStatusCode()); + assertEquals(MediaType.APPLICATION_PDF, response.getHeaders().getContentType()); + + Resource body = response.getBody(); + assertNotNull(body); + + byte[] outputBytes = readResource(body); + assertTrue(outputBytes.length > 0, "repaired PDF body should not be empty"); + + // Output must be a structurally valid PDF; confirm by reloading it. + try (PDDocument reloaded = org.apache.pdfbox.Loader.loadPDF(outputBytes)) { + assertEquals(2, reloaded.getNumberOfPages()); + } + } + + @Test + @DisplayName("loads the input file exactly once via the factory") + void repairPdf_pdfBoxFallback_loadsInputOnce() throws Exception { + MockMultipartFile input = + new MockMultipartFile( + "fileInput", + "broken.pdf", + MediaType.APPLICATION_PDF_VALUE, + buildPdfBytes(1)); + + when(pdfDocumentFactory.load(any(File.class))).thenReturn(newRealDocument(1)); + + repairController.repairPdf(pdfFileFrom(input)); + + verify(pdfDocumentFactory, times(1)).load(any(File.class)); + } + + @Test + @DisplayName("does not consult qpdf/ghostscript a second time once disabled") + void repairPdf_checksToolAvailability() throws Exception { + MockMultipartFile input = + new MockMultipartFile( + "fileInput", + "broken.pdf", + MediaType.APPLICATION_PDF_VALUE, + buildPdfBytes(1)); + + when(pdfDocumentFactory.load(any(File.class))).thenReturn(newRealDocument(1)); + + repairController.repairPdf(pdfFileFrom(input)); + + // Ghostscript is checked once (skip), qpdf once (skip), then both checked again to + // decide the last-resort branch. + verify(endpointConfiguration, times(2)).isGroupEnabled("Ghostscript"); + verify(endpointConfiguration, times(2)).isGroupEnabled("qpdf"); + } + + @Test + @DisplayName("the file passed to the factory exists on disk when loaded") + void repairPdf_transfersInputToRealTempFile() throws Exception { + byte[] pdf = buildPdfBytes(3); + MockMultipartFile input = + new MockMultipartFile( + "fileInput", "broken.pdf", MediaType.APPLICATION_PDF_VALUE, pdf); + + // Assert the file handed to load() is a real, non-empty file (transferTo succeeded). + when(pdfDocumentFactory.load(any(File.class))) + .thenAnswer( + invocation -> { + File file = invocation.getArgument(0); + assertTrue(file.exists(), "temp input file should exist"); + assertEquals(pdf.length, file.length()); + return newRealDocument(3); + }); + + ResponseEntity response = repairController.repairPdf(pdfFileFrom(input)); + assertEquals(HttpStatus.OK, response.getStatusCode()); + } + } + + @Nested + @DisplayName("Output filename handling") + class FilenameHandling { + + @Test + @DisplayName("appends _repaired.pdf to the base name in the Content-Disposition header") + void repairPdf_setsRepairedFilename() throws Exception { + MockMultipartFile input = + new MockMultipartFile( + "fileInput", + "mydoc.pdf", + MediaType.APPLICATION_PDF_VALUE, + buildPdfBytes(1)); + + when(pdfDocumentFactory.load(any(File.class))).thenReturn(newRealDocument(1)); + + ResponseEntity response = repairController.repairPdf(pdfFileFrom(input)); + + HttpHeaders headers = response.getHeaders(); + String disposition = headers.getFirst(HttpHeaders.CONTENT_DISPOSITION); + assertNotNull(disposition); + assertTrue( + disposition.contains("mydoc_repaired.pdf"), + "expected repaired filename in: " + disposition); + } + + @Test + @DisplayName("filename without extension still gets _repaired.pdf appended") + void repairPdf_filenameWithoutExtension() throws Exception { + MockMultipartFile input = + new MockMultipartFile( + "fileInput", + "noext", + MediaType.APPLICATION_PDF_VALUE, + buildPdfBytes(1)); + + when(pdfDocumentFactory.load(any(File.class))).thenReturn(newRealDocument(1)); + + ResponseEntity response = repairController.repairPdf(pdfFileFrom(input)); + + String disposition = response.getHeaders().getFirst(HttpHeaders.CONTENT_DISPOSITION); + assertNotNull(disposition); + assertTrue( + disposition.contains("noext_repaired.pdf"), + "expected repaired filename in: " + disposition); + } + + @Test + @DisplayName("null original filename falls back to 'default'") + void repairPdf_nullOriginalFilename_usesDefault() throws Exception { + // MockMultipartFile with null original filename. + MockMultipartFile input = + new MockMultipartFile( + "fileInput", null, MediaType.APPLICATION_PDF_VALUE, buildPdfBytes(1)); + + when(pdfDocumentFactory.load(any(File.class))).thenReturn(newRealDocument(1)); + + ResponseEntity response = repairController.repairPdf(pdfFileFrom(input)); + + assertEquals(HttpStatus.OK, response.getStatusCode()); + String disposition = response.getHeaders().getFirst(HttpHeaders.CONTENT_DISPOSITION); + assertNotNull(disposition); + // MockMultipartFile maps a null name to "", so the base is empty -> leading underscore. + assertTrue( + disposition.contains("_repaired.pdf"), + "expected empty-base repaired filename in: " + disposition); + } + } + + @Nested + @DisplayName("Error propagation") + class ErrorPropagation { + + @Test + @DisplayName("IOException from the factory propagates to the caller") + void repairPdf_loadThrowsIOException_propagates() throws Exception { + MockMultipartFile input = + new MockMultipartFile( + "fileInput", + "broken.pdf", + MediaType.APPLICATION_PDF_VALUE, + buildPdfBytes(1)); + + when(pdfDocumentFactory.load(any(File.class))) + .thenThrow(new IOException("cannot load corrupt pdf")); + + IOException thrown = + assertThrows( + IOException.class, + () -> repairController.repairPdf(pdfFileFrom(input))); + assertEquals("cannot load corrupt pdf", thrown.getMessage()); + } + + @Test + @DisplayName("RuntimeException from the factory propagates to the caller") + void repairPdf_loadThrowsRuntimeException_propagates() throws Exception { + MockMultipartFile input = + new MockMultipartFile( + "fileInput", + "broken.pdf", + MediaType.APPLICATION_PDF_VALUE, + buildPdfBytes(1)); + + when(pdfDocumentFactory.load(any(File.class))) + .thenThrow(new IllegalStateException("boom")); + + IllegalStateException thrown = + assertThrows( + IllegalStateException.class, + () -> repairController.repairPdf(pdfFileFrom(input))); + assertEquals("boom", thrown.getMessage()); + } + } +} diff --git a/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/ScannerEffectControllerTest.java b/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/ScannerEffectControllerTest.java new file mode 100644 index 0000000000..6919926846 --- /dev/null +++ b/app/core/src/test/java/stirling/software/SPDF/controller/api/misc/ScannerEffectControllerTest.java @@ -0,0 +1,932 @@ +package stirling.software.SPDF.controller.api.misc; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyBoolean; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.Mockito.lenient; +import static org.mockito.Mockito.mock; + +import java.awt.Color; +import java.awt.image.BufferedImage; +import java.awt.image.DataBufferInt; +import java.io.ByteArrayOutputStream; +import java.io.File; +import java.io.IOException; +import java.io.InputStream; +import java.lang.reflect.Constructor; +import java.lang.reflect.Method; +import java.nio.file.Files; + +import org.apache.pdfbox.Loader; +import org.apache.pdfbox.pdmodel.PDDocument; +import org.apache.pdfbox.pdmodel.PDPage; +import org.apache.pdfbox.pdmodel.common.PDRectangle; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.EnumSource; +import org.mockito.InjectMocks; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.core.io.Resource; +import org.springframework.http.HttpStatus; +import org.springframework.http.ResponseEntity; +import org.springframework.mock.web.MockMultipartFile; + +import stirling.software.SPDF.model.api.misc.ScannerEffectRequest; +import stirling.software.common.service.CustomPDFDocumentFactory; +import stirling.software.common.util.TempFile; +import stirling.software.common.util.TempFileManager; + +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +@DisplayName("ScannerEffectController Tests") +class ScannerEffectControllerTest { + + @Mock private CustomPDFDocumentFactory pdfDocumentFactory; + @Mock private TempFileManager tempFileManager; + @InjectMocks private ScannerEffectController controller; + + @BeforeEach + void setUp() throws Exception { + // Real temp file backing so WebResponseUtils.pdfDocToWebResponse can save the output. + lenient() + .when(tempFileManager.createManagedTempFile(anyString())) + .thenAnswer( + inv -> { + File f = + Files.createTempFile("scanner_test", inv.getArgument(0)) + .toFile(); + f.deleteOnExit(); + TempFile tf = mock(TempFile.class); + lenient().when(tf.getFile()).thenReturn(f); + lenient().when(tf.getPath()).thenReturn(f.toPath()); + return tf; + }); + } + + // --------------------------------------------------------------------- + // Helpers + // --------------------------------------------------------------------- + + private static MockMultipartFile pdfFile(String filename, int pageCount, PDRectangle pageSize) + throws IOException { + try (PDDocument doc = new PDDocument()) { + for (int i = 0; i < pageCount; i++) { + doc.addPage(new PDPage(pageSize)); + } + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + doc.save(baos); + return new MockMultipartFile( + "fileInput", filename, "application/pdf", baos.toByteArray()); + } + } + + private static ScannerEffectRequest baseRequest(MockMultipartFile file) { + ScannerEffectRequest request = new ScannerEffectRequest(); + request.setFileInput(file); + // Advanced enabled so quality presets do not override our resolution. + request.setAdvancedEnabled(true); + // Keep rendering tiny and fast; deterministic structure only. + request.setResolution(36); + request.setRotation(ScannerEffectRequest.Rotation.none); + request.setRotate(0); + request.setRotateVariance(0); + request.setBorder(2); + request.setBrightness(1.0f); + request.setContrast(1.0f); + request.setBlur(0f); + request.setNoise(0f); + request.setYellowish(false); + request.setColorspace(ScannerEffectRequest.Colorspace.grayscale); + return request; + } + + /** Stub both factory.load overloads to return fresh real documents loaded from bytes. */ + private void stubFactoryLoad(byte[] pdfBytes) throws IOException { + // Used for page count + as output base (load(byte[])). + lenient() + .when(pdfDocumentFactory.load(any(byte[].class))) + .thenAnswer( + inv -> { + byte[] b = inv.getArgument(0); + return Loader.loadPDF(b); + }); + // Used by RenderingResources.fromBytes (load(byte[], true)). + lenient() + .when(pdfDocumentFactory.load(any(byte[].class), anyBoolean())) + .thenAnswer(inv -> Loader.loadPDF((byte[]) inv.getArgument(0))); + } + + private static byte[] drain(ResponseEntity response) throws IOException { + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + try (InputStream in = response.getBody().getInputStream()) { + in.transferTo(baos); + } + return baos.toByteArray(); + } + + // --------------------------------------------------------------------- + // End-to-end controller behaviour (real rendering, mocked boundaries) + // --------------------------------------------------------------------- + + @Nested + @DisplayName("scannerEffect end-to-end") + class EndToEnd { + + @Test + @DisplayName("produces a valid single-page PDF response for a one-page input") + void singlePageHappyPath() throws Exception { + MockMultipartFile file = pdfFile("input.pdf", 1, PDRectangle.A6); + stubFactoryLoad(file.getBytes()); + ScannerEffectRequest request = baseRequest(file); + + ResponseEntity response = controller.scannerEffect(request); + + assertThat(response).isNotNull(); + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(response.getBody()).isNotNull(); + + byte[] out = drain(response); + assertThat(out).isNotEmpty(); + try (PDDocument result = Loader.loadPDF(out)) { + assertThat(result.getNumberOfPages()).isEqualTo(1); + PDRectangle box = result.getPage(0).getMediaBox(); + // Output page keeps the original page dimensions. + assertThat(box.getWidth()).isCloseTo(PDRectangle.A6.getWidth(), within(1f)); + assertThat(box.getHeight()).isCloseTo(PDRectangle.A6.getHeight(), within(1f)); + } + } + + @Test + @DisplayName("preserves page count for a multi-page input") + void multiPageKeepsPageCount() throws Exception { + MockMultipartFile file = pdfFile("multi.pdf", 3, PDRectangle.A6); + stubFactoryLoad(file.getBytes()); + ScannerEffectRequest request = baseRequest(file); + + ResponseEntity response = controller.scannerEffect(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + try (PDDocument result = Loader.loadPDF(drain(response))) { + assertThat(result.getNumberOfPages()).isEqualTo(3); + } + } + + @Test + @DisplayName("applies colour effects (color colorspace, blur, noise, yellowish, rotation)") + void richEffectsStillProduceValidPdf() throws Exception { + MockMultipartFile file = pdfFile("rich.pdf", 1, PDRectangle.A6); + stubFactoryLoad(file.getBytes()); + ScannerEffectRequest request = baseRequest(file); + request.setColorspace(ScannerEffectRequest.Colorspace.color); + request.setBlur(1.0f); + request.setNoise(4.0f); + request.setYellowish(true); + request.setRotation(ScannerEffectRequest.Rotation.slight); + request.setRotateVariance(2); + request.setBorder(5); + request.setBrightness(1.05f); + request.setContrast(1.1f); + + ResponseEntity response = controller.scannerEffect(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + try (PDDocument result = Loader.loadPDF(drain(response))) { + assertThat(result.getNumberOfPages()).isEqualTo(1); + } + } + + @Test + @DisplayName("non-advanced request applies the quality preset without error") + void qualityPresetPath() throws Exception { + MockMultipartFile file = pdfFile("preset.pdf", 1, PDRectangle.A6); + stubFactoryLoad(file.getBytes()); + ScannerEffectRequest request = baseRequest(file); + request.setAdvancedEnabled(false); + // Low preset uses resolution 75 which is the cheapest render. + request.setQuality(ScannerEffectRequest.Quality.low); + + ResponseEntity response = controller.scannerEffect(request); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + try (PDDocument result = Loader.loadPDF(drain(response))) { + assertThat(result.getNumberOfPages()).isEqualTo(1); + } + } + + @Test + @DisplayName("rejects a DPI above the safe maximum with IllegalArgumentException") + void dpiAboveLimitIsRejected() throws Exception { + MockMultipartFile file = pdfFile("highdpi.pdf", 1, PDRectangle.A6); + stubFactoryLoad(file.getBytes()); + ScannerEffectRequest request = baseRequest(file); + // No application context in tests -> maxSafeDpi defaults to 500. + request.setResolution(600); + + assertThatThrownBy(() -> controller.scannerEffect(request)) + .isInstanceOf(IllegalArgumentException.class) + .hasMessageContaining("600"); + } + + @Test + @DisplayName("rejects an empty (zero-page) document with IllegalArgumentException") + void emptyDocumentIsRejected() throws Exception { + // A real 1-page file on disk, but the factory returns an empty document. + MockMultipartFile file = pdfFile("empty.pdf", 1, PDRectangle.A6); + ScannerEffectRequest request = baseRequest(file); + + lenient() + .when(pdfDocumentFactory.load(any(byte[].class))) + .thenAnswer(inv -> new PDDocument()); + lenient() + .when(pdfDocumentFactory.load(any(byte[].class), anyBoolean())) + .thenAnswer(inv -> new PDDocument()); + + assertThatThrownBy(() -> controller.scannerEffect(request)) + .isInstanceOf(IllegalArgumentException.class) + .hasMessageContaining("no pages"); + } + + @Test + @DisplayName("propagates IOException from the document factory") + void factoryIOExceptionPropagates() throws Exception { + MockMultipartFile file = pdfFile("io.pdf", 1, PDRectangle.A6); + ScannerEffectRequest request = baseRequest(file); + + lenient() + .when(pdfDocumentFactory.load(any(byte[].class))) + .thenThrow(new IOException("boom-load")); + lenient() + .when(pdfDocumentFactory.load(any(byte[].class), anyBoolean())) + .thenThrow(new IOException("boom-load")); + + assertThatThrownBy(() -> controller.scannerEffect(request)) + .isInstanceOf(IOException.class); + } + } + + // --------------------------------------------------------------------- + // Pure static image-processing logic via reflection + // --------------------------------------------------------------------- + + @Nested + @DisplayName("calculateSafeResolution") + class CalculateSafeResolution { + + private Method method; + + @BeforeEach + void setUp() throws Exception { + method = + ScannerEffectController.class.getDeclaredMethod( + "calculateSafeResolution", float.class, float.class, int.class); + method.setAccessible(true); + } + + private int invoke(float w, float h, int res) throws Exception { + return (int) method.invoke(null, w, h, res); + } + + @Test + @DisplayName("keeps requested resolution when the projected image is small") + void keepsResolutionForSmallPage() throws Exception { + // A6 at 72 dpi is well within limits. + assertThat(invoke(297f, 420f, 72)).isEqualTo(72); + } + + @Test + @DisplayName("downscales resolution when the projected image is huge") + void downscalesForHugePage() throws Exception { + // A0-ish page at a very high dpi exceeds the pixel/size caps. + int safe = invoke(2384f, 3370f, 2000); + assertThat(safe).isLessThan(2000); + assertThat(safe).isGreaterThanOrEqualTo(72); + } + + @Test + @DisplayName("never returns below the floor of 72") + void neverBelowFloor() throws Exception { + int safe = invoke(5000f, 5000f, 4000); + assertThat(safe).isGreaterThanOrEqualTo(72); + } + } + + @Nested + @DisplayName("determineRenderResolution") + class DetermineRenderResolution { + + @Test + @DisplayName("returns the request resolution unchanged") + void returnsRequestResolution() throws Exception { + Method method = + ScannerEffectController.class.getDeclaredMethod( + "determineRenderResolution", ScannerEffectRequest.class); + method.setAccessible(true); + + ScannerEffectRequest request = new ScannerEffectRequest(); + request.setResolution(123); + + assertThat((int) method.invoke(null, request)).isEqualTo(123); + } + } + + @Nested + @DisplayName("convertColorspace / convertToGrayscale") + class Colorspace { + + private static BufferedImage solid(int rgb) { + BufferedImage image = new BufferedImage(4, 4, BufferedImage.TYPE_INT_RGB); + for (int y = 0; y < 4; y++) { + for (int x = 0; x < 4; x++) { + image.setRGB(x, y, rgb); + } + } + return image; + } + + @ParameterizedTest + @EnumSource(ScannerEffectRequest.Colorspace.class) + @DisplayName("returns an INT_RGB image of the same dimensions for any colorspace") + void preservesDimensions(ScannerEffectRequest.Colorspace colorspace) throws Exception { + Method method = + ScannerEffectController.class.getDeclaredMethod( + "convertColorspace", + BufferedImage.class, + ScannerEffectRequest.Colorspace.class); + method.setAccessible(true); + + BufferedImage src = solid(0x123456); + BufferedImage result = (BufferedImage) method.invoke(null, src, colorspace); + + assertThat(result.getWidth()).isEqualTo(4); + assertThat(result.getHeight()).isEqualTo(4); + assertThat(result.getType()).isEqualTo(BufferedImage.TYPE_INT_RGB); + } + + @Test + @DisplayName("grayscale collapses the channels to a single grey value") + void grayscaleEqualisesChannels() throws Exception { + Method method = + ScannerEffectController.class.getDeclaredMethod( + "convertColorspace", + BufferedImage.class, + ScannerEffectRequest.Colorspace.class); + method.setAccessible(true); + + // R=90, G=120, B=150 -> avg 120 -> 0x787878 + BufferedImage src = solid((90 << 16) | (120 << 8) | 150); + BufferedImage result = + (BufferedImage) + method.invoke(null, src, ScannerEffectRequest.Colorspace.grayscale); + + int px = result.getRGB(0, 0) & 0xFFFFFF; + int r = (px >> 16) & 0xFF; + int g = (px >> 8) & 0xFF; + int b = px & 0xFF; + assertThat(r).isEqualTo(g).isEqualTo(b); + assertThat(r).isEqualTo(120); + } + + @Test + @DisplayName("convertToGrayscale mutates the buffer in place") + void convertToGrayscaleInPlace() throws Exception { + Method method = + ScannerEffectController.class.getDeclaredMethod( + "convertToGrayscale", BufferedImage.class); + method.setAccessible(true); + + BufferedImage img = solid((30 << 16) | (60 << 8) | 90); // avg 60 + method.invoke(null, img); + + int px = img.getRGB(1, 1) & 0xFFFFFF; + assertThat(px).isEqualTo((60 << 16) | (60 << 8) | 60); + } + } + + @Nested + @DisplayName("calculateRotation") + class CalculateRotation { + + private Method method; + + @BeforeEach + void setUp() throws Exception { + method = + ScannerEffectController.class.getDeclaredMethod( + "calculateRotation", int.class, int.class); + method.setAccessible(true); + } + + @Test + @DisplayName("returns exactly 0 when base rotation and variance are both 0") + void zeroWhenNoRotation() throws Exception { + assertThat((double) method.invoke(null, 0, 0)).isEqualTo(0.0); + } + + @Test + @DisplayName("stays within base +/- variance bounds") + void withinBounds() throws Exception { + for (int i = 0; i < 100; i++) { + double value = (double) method.invoke(null, 5, 3); + assertThat(value).isBetween(2.0, 8.0); + } + } + + @Test + @DisplayName("with zero variance but non-zero base returns the base rotation") + void zeroVarianceReturnsBase() throws Exception { + // base=5, variance=0 -> 5 + (rand*2-1)*0 = 5 + assertThat((double) method.invoke(null, 5, 0)).isEqualTo(5.0); + } + } + + @Nested + @DisplayName("blendColors") + class BlendColors { + + private Method method; + + @BeforeEach + void setUp() throws Exception { + method = + ScannerEffectController.class.getDeclaredMethod( + "blendColors", int.class, int.class, float.class); + method.setAccessible(true); + } + + private int invoke(int fg, int bg, float alpha) throws Exception { + return (int) method.invoke(null, fg, bg, alpha); + } + + @Test + @DisplayName("alpha=1 yields the foreground colour") + void alphaOneIsForeground() throws Exception { + assertThat(invoke(0xAABBCC, 0x112233, 1.0f)).isEqualTo(0xAABBCC); + } + + @Test + @DisplayName("alpha=0 yields the background colour") + void alphaZeroIsBackground() throws Exception { + assertThat(invoke(0xAABBCC, 0x112233, 0.0f)).isEqualTo(0x112233); + } + + @Test + @DisplayName("alpha=0.5 yields the rounded midpoint per channel") + void alphaHalfIsMidpoint() throws Exception { + // fg 0x806040 (128,96,64), bg 0x204060 (32,64,96) + int blended = invoke(0x806040, 0x204060, 0.5f); + int r = (blended >> 16) & 0xFF; + int g = (blended >> 8) & 0xFF; + int b = blended & 0xFF; + assertThat(r).isEqualTo(80); // (128+32)/2 + assertThat(g).isEqualTo(80); // (96+64)/2 + assertThat(b).isEqualTo(80); // (64+96)/2 + } + } + + @Nested + @DisplayName("fillWithGradient") + class FillWithGradient { + + private Method method; + + @BeforeEach + void setUp() throws Exception { + method = + ScannerEffectController.class.getDeclaredMethod( + "fillWithGradient", + int[].class, + int.class, + int.class, + int[].class, + boolean.class); + method.setAccessible(true); + } + + @Test + @DisplayName("vertical gradient assigns one LUT value per row") + void verticalFill() throws Exception { + int width = 3; + int height = 2; + int[] pixels = new int[width * height]; + int[] lut = {0x111111, 0x222222}; + + method.invoke(null, pixels, width, height, lut, true); + + // Row 0 all lut[0], row 1 all lut[1]. + assertThat(pixels[0]).isEqualTo(0x111111); + assertThat(pixels[2]).isEqualTo(0x111111); + assertThat(pixels[3]).isEqualTo(0x222222); + assertThat(pixels[5]).isEqualTo(0x222222); + } + + @Test + @DisplayName("horizontal gradient assigns one LUT value per column, repeated each row") + void horizontalFill() throws Exception { + int width = 3; + int height = 2; + int[] pixels = new int[width * height]; + int[] lut = {0xAA0000, 0x00BB00, 0x0000CC}; + + method.invoke(null, pixels, width, height, lut, false); + + // Each row mirrors the LUT. + assertThat(pixels[0]).isEqualTo(0xAA0000); + assertThat(pixels[1]).isEqualTo(0x00BB00); + assertThat(pixels[2]).isEqualTo(0x0000CC); + assertThat(pixels[3]).isEqualTo(0xAA0000); + assertThat(pixels[4]).isEqualTo(0x00BB00); + assertThat(pixels[5]).isEqualTo(0x0000CC); + } + } + + @Nested + @DisplayName("createGradientLUT") + class CreateGradientLUT { + + private Object gradient(boolean vertical, Color start, Color end) throws Exception { + Class gradientClass = + Class.forName( + "stirling.software.SPDF.controller.api.misc.ScannerEffectController$GradientConfig"); + Constructor ctor = + gradientClass.getDeclaredConstructor(boolean.class, Color.class, Color.class); + ctor.setAccessible(true); + return ctor.newInstance(vertical, start, end); + } + + private int[] invokeLut(int width, int height, Object gradientConfig) throws Exception { + Class gradientClass = + Class.forName( + "stirling.software.SPDF.controller.api.misc.ScannerEffectController$GradientConfig"); + Method method = + ScannerEffectController.class.getDeclaredMethod( + "createGradientLUT", int.class, int.class, gradientClass); + method.setAccessible(true); + return (int[]) method.invoke(null, width, height, gradientConfig); + } + + @Test + @DisplayName("vertical LUT length equals height and interpolates endpoints") + void verticalLut() throws Exception { + Object g = gradient(true, Color.BLACK, Color.WHITE); + int[] lut = invokeLut(10, 5, g); + + assertThat(lut).hasSize(5); + assertThat(lut[0] & 0xFFFFFF).isEqualTo(0x000000); + assertThat(lut[lut.length - 1] & 0xFFFFFF).isEqualTo(0xFFFFFF); + } + + @Test + @DisplayName("horizontal LUT length equals width") + void horizontalLut() throws Exception { + Object g = gradient(false, Color.BLACK, Color.WHITE); + int[] lut = invokeLut(7, 3, g); + + assertThat(lut).hasSize(7); + assertThat(lut[0] & 0xFFFFFF).isEqualTo(0x000000); + assertThat(lut[lut.length - 1] & 0xFFFFFF).isEqualTo(0xFFFFFF); + } + + @Test + @DisplayName("constant colour produces a uniform LUT") + void uniformLut() throws Exception { + Object g = gradient(true, new Color(0x40, 0x50, 0x60), new Color(0x40, 0x50, 0x60)); + int[] lut = invokeLut(4, 4, g); + + for (int value : lut) { + assertThat(value & 0xFFFFFF).isEqualTo(0x405060); + } + } + } + + @Nested + @DisplayName("applyAllEffectsSinglePass") + class ApplyAllEffects { + + private Method method; + + @BeforeEach + void setUp() throws Exception { + method = + ScannerEffectController.class.getDeclaredMethod( + "applyAllEffectsSinglePass", + BufferedImage.class, + float.class, + float.class, + boolean.class, + double.class); + method.setAccessible(true); + } + + private static BufferedImage solid(int width, int height, int rgb) { + BufferedImage image = new BufferedImage(width, height, BufferedImage.TYPE_INT_RGB); + int[] pixels = ((DataBufferInt) image.getRaster().getDataBuffer()).getData(); + java.util.Arrays.fill(pixels, rgb); + return image; + } + + @Test + @DisplayName("identity settings leave pixels unchanged") + void identityKeepsPixels() throws Exception { + BufferedImage src = solid(8, 8, 0x648CB4); // 100,140,180 + BufferedImage out = (BufferedImage) method.invoke(null, src, 1.0f, 1.0f, false, 0.0d); + + assertThat(out.getWidth()).isEqualTo(8); + assertThat(out.getHeight()).isEqualTo(8); + assertThat(out.getRGB(0, 0) & 0xFFFFFF).isEqualTo(0x648CB4); + } + + @Test + @DisplayName("brightness > 1 increases channel values and clamps at 255") + void brightnessClamps() throws Exception { + BufferedImage src = solid(4, 4, 0xC8C8C8); // 200 each + BufferedImage out = (BufferedImage) method.invoke(null, src, 2.0f, 1.0f, false, 0.0d); + + int px = out.getRGB(0, 0) & 0xFFFFFF; + // 200 * 2 = 400 -> clamped to 255 on every channel. + assertThat(px).isEqualTo(0xFFFFFF); + } + + @Test + @DisplayName("zero brightness produces black") + void zeroBrightnessIsBlack() throws Exception { + BufferedImage src = solid(4, 4, 0xFFFFFF); + BufferedImage out = (BufferedImage) method.invoke(null, src, 0.0f, 1.0f, false, 0.0d); + + assertThat(out.getRGB(0, 0) & 0xFFFFFF).isEqualTo(0x000000); + } + + @Test + @DisplayName("yellowish tint lowers the blue channel relative to source") + void yellowishReducesBlue() throws Exception { + BufferedImage src = solid(4, 4, 0xFFFFFF); // bright white maximises tint effect + BufferedImage out = (BufferedImage) method.invoke(null, src, 1.0f, 1.0f, true, 0.0d); + + int px = out.getRGB(0, 0) & 0xFFFFFF; + int b = px & 0xFF; + assertThat(b).isLessThan(255); + } + + @Test + @DisplayName("noise keeps every channel within the valid 0..255 range") + void noiseStaysInRange() throws Exception { + BufferedImage src = solid(32, 32, 0x808080); + BufferedImage out = (BufferedImage) method.invoke(null, src, 1.0f, 1.0f, false, 50.0d); + + for (int y = 0; y < out.getHeight(); y++) { + for (int x = 0; x < out.getWidth(); x++) { + int px = out.getRGB(x, y); + assertThat((px >> 16) & 0xFF).isBetween(0, 255); + assertThat((px >> 8) & 0xFF).isBetween(0, 255); + assertThat(px & 0xFF).isBetween(0, 255); + } + } + } + } + + @Nested + @DisplayName("softenEdges") + class SoftenEdges { + + private Method method; + + @BeforeEach + void setUp() throws Exception { + method = + ScannerEffectController.class.getDeclaredMethod( + "softenEdges", + BufferedImage.class, + int.class, + Color.class, + Color.class, + boolean.class); + method.setAccessible(true); + } + + @Test + @DisplayName("center pixel stays at foreground when feathering only touches edges") + void centreStaysForeground() throws Exception { + BufferedImage src = new BufferedImage(11, 11, BufferedImage.TYPE_INT_RGB); + int fg = 0x102030; + int[] pixels = ((DataBufferInt) src.getRaster().getDataBuffer()).getData(); + java.util.Arrays.fill(pixels, fg); + + BufferedImage out = + (BufferedImage) method.invoke(null, src, 2, Color.WHITE, Color.WHITE, true); + + // Centre is far from any edge (distance 5 >= feather radius 2) so alpha=1. + assertThat(out.getRGB(5, 5) & 0xFFFFFF).isEqualTo(fg); + } + + @Test + @DisplayName("corner pixel is blended toward the background gradient") + void cornerBlendsToBackground() throws Exception { + BufferedImage src = new BufferedImage(11, 11, BufferedImage.TYPE_INT_RGB); + int fg = 0x000000; + int[] pixels = ((DataBufferInt) src.getRaster().getDataBuffer()).getData(); + java.util.Arrays.fill(pixels, fg); + + // Background is pure white; corner distance d=0 -> alpha=0 -> background. + BufferedImage out = + (BufferedImage) method.invoke(null, src, 3, Color.WHITE, Color.WHITE, true); + + assertThat(out.getRGB(0, 0) & 0xFFFFFF).isEqualTo(0xFFFFFF); + } + + @Test + @DisplayName("preserves image dimensions") + void preservesDimensions() throws Exception { + BufferedImage src = new BufferedImage(6, 9, BufferedImage.TYPE_INT_RGB); + BufferedImage out = + (BufferedImage) method.invoke(null, src, 1, Color.GRAY, Color.DARK_GRAY, false); + + assertThat(out.getWidth()).isEqualTo(6); + assertThat(out.getHeight()).isEqualTo(9); + } + } + + @Nested + @DisplayName("applyGaussianBlur") + class ApplyGaussianBlur { + + private Method method; + + @BeforeEach + void setUp() throws Exception { + method = + ScannerEffectController.class.getDeclaredMethod( + "applyGaussianBlur", BufferedImage.class, double.class); + method.setAccessible(true); + } + + @Test + @DisplayName("sigma <= 0 returns the same image instance") + void zeroSigmaIsNoOp() throws Exception { + BufferedImage src = new BufferedImage(8, 8, BufferedImage.TYPE_INT_RGB); + BufferedImage out = (BufferedImage) method.invoke(null, src, 0.0d); + assertThat(out).isSameAs(src); + } + + @Test + @DisplayName("negative sigma returns the same image instance") + void negativeSigmaIsNoOp() throws Exception { + BufferedImage src = new BufferedImage(8, 8, BufferedImage.TYPE_INT_RGB); + BufferedImage out = (BufferedImage) method.invoke(null, src, -1.0d); + assertThat(out).isSameAs(src); + } + + @Test + @DisplayName("positive sigma on a uniform image keeps the uniform colour and dimensions") + void uniformImageStaysUniform() throws Exception { + BufferedImage src = new BufferedImage(40, 40, BufferedImage.TYPE_INT_RGB); + int color = 0x405060; + int[] pixels = ((DataBufferInt) src.getRaster().getDataBuffer()).getData(); + java.util.Arrays.fill(pixels, color); + + BufferedImage out = (BufferedImage) method.invoke(null, src, 30.0d); + + assertThat(out).isNotSameAs(src); + assertThat(out.getWidth()).isEqualTo(40); + assertThat(out.getHeight()).isEqualTo(40); + // Blurring a uniform image yields the same uniform colour. + assertThat(out.getRGB(20, 20) & 0xFFFFFF).isEqualTo(color); + } + } + + @Nested + @DisplayName("rotateImage") + class RotateImage { + + private Object gradient() throws Exception { + Class gradientClass = + Class.forName( + "stirling.software.SPDF.controller.api.misc.ScannerEffectController$GradientConfig"); + Constructor ctor = + gradientClass.getDeclaredConstructor(boolean.class, Color.class, Color.class); + ctor.setAccessible(true); + return ctor.newInstance(true, Color.WHITE, Color.WHITE); + } + + private Method rotateMethod() throws Exception { + Class gradientClass = + Class.forName( + "stirling.software.SPDF.controller.api.misc.ScannerEffectController$GradientConfig"); + Method method = + ScannerEffectController.class.getDeclaredMethod( + "rotateImage", BufferedImage.class, double.class, gradientClass); + method.setAccessible(true); + return method; + } + + @Test + @DisplayName("zero rotation returns the same image instance") + void zeroRotationNoOp() throws Exception { + BufferedImage src = new BufferedImage(10, 10, BufferedImage.TYPE_INT_RGB); + BufferedImage out = (BufferedImage) rotateMethod().invoke(null, src, 0.0d, gradient()); + assertThat(out).isSameAs(src); + } + + @Test + @DisplayName("90 degree rotation swaps the bounding box dimensions") + void ninetyDegreeSwapsDimensions() throws Exception { + BufferedImage src = new BufferedImage(20, 10, BufferedImage.TYPE_INT_RGB); + BufferedImage out = (BufferedImage) rotateMethod().invoke(null, src, 90.0d, gradient()); + + assertThat(out).isNotSameAs(src); + // For 90 degrees, rotated bounds become height x width. + assertThat(out.getWidth()).isEqualTo(10); + assertThat(out.getHeight()).isEqualTo(20); + } + + @Test + @DisplayName("45 degree rotation grows the bounding box") + void fortyFiveGrowsBoundingBox() throws Exception { + BufferedImage src = new BufferedImage(20, 20, BufferedImage.TYPE_INT_RGB); + BufferedImage out = (BufferedImage) rotateMethod().invoke(null, src, 45.0d, gradient()); + + assertThat(out.getWidth()).isGreaterThan(20); + assertThat(out.getHeight()).isGreaterThan(20); + } + } + + @Nested + @DisplayName("addBorderWithGradient") + class AddBorderWithGradient { + + private Object gradient(boolean vertical) throws Exception { + Class gradientClass = + Class.forName( + "stirling.software.SPDF.controller.api.misc.ScannerEffectController$GradientConfig"); + Constructor ctor = + gradientClass.getDeclaredConstructor(boolean.class, Color.class, Color.class); + ctor.setAccessible(true); + return ctor.newInstance(vertical, Color.GRAY, Color.LIGHT_GRAY); + } + + private Method borderMethod() throws Exception { + Class gradientClass = + Class.forName( + "stirling.software.SPDF.controller.api.misc.ScannerEffectController$GradientConfig"); + Method method = + ScannerEffectController.class.getDeclaredMethod( + "addBorderWithGradient", BufferedImage.class, int.class, gradientClass); + method.setAccessible(true); + return method; + } + + @Test + @DisplayName("adds a border of the requested size on every side") + void growsByTwiceBorder() throws Exception { + BufferedImage src = new BufferedImage(10, 12, BufferedImage.TYPE_INT_RGB); + int border = 5; + BufferedImage out = + (BufferedImage) borderMethod().invoke(null, src, border, gradient(true)); + + assertThat(out.getWidth()).isEqualTo(10 + 2 * border); + assertThat(out.getHeight()).isEqualTo(12 + 2 * border); + } + + @Test + @DisplayName("zero border keeps the original dimensions") + void zeroBorderKeepsDimensions() throws Exception { + BufferedImage src = new BufferedImage(8, 8, BufferedImage.TYPE_INT_RGB); + BufferedImage out = + (BufferedImage) borderMethod().invoke(null, src, 0, gradient(false)); + + assertThat(out.getWidth()).isEqualTo(8); + assertThat(out.getHeight()).isEqualTo(8); + } + + @Test + @DisplayName("preserves the drawn source region inside the border") + void preservesSourcePixels() throws Exception { + BufferedImage src = new BufferedImage(6, 6, BufferedImage.TYPE_INT_RGB); + int fg = 0x123456; + int[] pixels = ((DataBufferInt) src.getRaster().getDataBuffer()).getData(); + java.util.Arrays.fill(pixels, fg); + + int border = 3; + BufferedImage out = + (BufferedImage) borderMethod().invoke(null, src, border, gradient(true)); + + // Source top-left maps to (border, border) in the composed image. + assertThat(out.getRGB(border, border) & 0xFFFFFF).isEqualTo(fg); + } + } + + // --------------------------------------------------------------------- + // Small AssertJ helper for float closeness + // --------------------------------------------------------------------- + + private static org.assertj.core.data.Offset within(float tol) { + return org.assertj.core.data.Offset.offset(tol); + } +} diff --git a/app/core/src/test/java/stirling/software/SPDF/controller/api/pipeline/PipelineControllerGapTest.java b/app/core/src/test/java/stirling/software/SPDF/controller/api/pipeline/PipelineControllerGapTest.java new file mode 100644 index 0000000000..edcfbc72ea --- /dev/null +++ b/app/core/src/test/java/stirling/software/SPDF/controller/api/pipeline/PipelineControllerGapTest.java @@ -0,0 +1,402 @@ +package stirling.software.SPDF.controller.api.pipeline; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.verifyNoInteractions; +import static org.mockito.Mockito.when; + +import java.io.ByteArrayInputStream; +import java.io.IOException; +import java.io.InputStream; +import java.nio.charset.StandardCharsets; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.ArrayList; +import java.util.List; +import java.util.Map; +import java.util.zip.ZipEntry; +import java.util.zip.ZipInputStream; + +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.ArgumentCaptor; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.core.io.ByteArrayResource; +import org.springframework.core.io.Resource; +import org.springframework.http.HttpStatus; +import org.springframework.http.ResponseEntity; +import org.springframework.mock.web.MockMultipartFile; +import org.springframework.web.multipart.MultipartFile; + +import stirling.software.SPDF.model.PipelineConfig; +import stirling.software.SPDF.model.PipelineResult; +import stirling.software.SPDF.model.api.HandleDataRequest; +import stirling.software.common.service.PostHogService; +import stirling.software.common.util.TempFileManager; + +import tools.jackson.databind.ObjectMapper; + +/** + * Unit tests for {@link PipelineController#handleData}. The controller is exercised directly with a + * real {@link ObjectMapper} for JSON parsing and mocked collaborators for the processor, analytics + * and temp-file management. The {@link TempFileManager} is stubbed to hand out real on-disk temp + * files so the {@code TempFile} wrapper and the stream-copy / zip paths run for real. + */ +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class PipelineControllerGapTest { + + @Mock private PipelineProcessor processor; + + @Mock private PostHogService postHogService; + + @Mock private TempFileManager tempFileManager; + + private ObjectMapper objectMapper; + private PipelineController controller; + + // Real temp files handed back by the mocked TempFileManager; cleaned up after each test. + private final List createdTempFiles = new ArrayList<>(); + + private static final String VALID_JSON = + "{\"name\":\"test-pipeline\",\"pipeline\":[" + + "{\"operation\":\"/api/v1/misc/repair\",\"parameters\":{}}]}"; + + @BeforeEach + void setUp() { + objectMapper = new ObjectMapper(); + controller = + new PipelineController(processor, objectMapper, postHogService, tempFileManager); + } + + /** + * Stub the manager so every {@code new TempFile(manager, suffix)} the controller builds gets a + * real on-disk file. Only invoked by tests that actually exercise an output path, so tests that + * short-circuit can assert {@code verifyNoInteractions(tempFileManager)}. + */ + private void stubTempFiles() throws IOException { + when(tempFileManager.createTempFile(anyString())) + .thenAnswer( + invocation -> { + String suffix = invocation.getArgument(0); + Path p = Files.createTempFile("pipeline-gap-test-", suffix); + createdTempFiles.add(p); + return p.toFile(); + }); + } + + @AfterEach + void tearDown() throws IOException { + for (Path p : createdTempFiles) { + Files.deleteIfExists(p); + } + createdTempFiles.clear(); + } + + private HandleDataRequest request(MultipartFile[] files, String json) { + HandleDataRequest req = new HandleDataRequest(); + req.setFileInput(files); + req.setJson(json); + return req; + } + + private MockMultipartFile pdf(String name) { + return new MockMultipartFile( + "fileInput", name, "application/pdf", ("content of " + name).getBytes()); + } + + private MultipartFile[] oneFile() { + return new MultipartFile[] {pdf("input.pdf")}; + } + + /** A Resource backed by bytes that also reports a filename (ByteArrayResource returns null). */ + private static Resource namedResource(String filename, byte[] data) { + return new ByteArrayResource(data) { + @Override + public String getFilename() { + return filename; + } + }; + } + + @Nested + @DisplayName("Null / empty input short-circuits") + class NullAndEmptyInputs { + + @Test + @DisplayName("returns null when file input is null without touching collaborators") + void nullFiles_returnsNull() throws Exception { + ResponseEntity response = controller.handleData(request(null, VALID_JSON)); + + assertNull(response); + verifyNoInteractions(processor); + verifyNoInteractions(postHogService); + verifyNoInteractions(tempFileManager); + } + + @Test + @DisplayName("returns null when processor yields no input files") + void nullInputFiles_returnsNull() throws Exception { + MultipartFile[] files = oneFile(); + when(processor.generateInputFiles(files)).thenReturn(null); + + ResponseEntity response = controller.handleData(request(files, VALID_JSON)); + + assertNull(response); + // Event still captured before processing begins. + verify(postHogService).captureEvent(eq("pipeline_api_event"), any()); + verify(processor, never()).runPipelineAgainstFiles(any(), any()); + } + + @Test + @DisplayName("returns null when processor yields an empty input list") + void emptyInputFiles_returnsNull() throws Exception { + MultipartFile[] files = oneFile(); + when(processor.generateInputFiles(files)).thenReturn(new ArrayList<>()); + + ResponseEntity response = controller.handleData(request(files, VALID_JSON)); + + assertNull(response); + verify(processor, never()).runPipelineAgainstFiles(any(), any()); + } + + @Test + @DisplayName("returns null when pipeline result has null output files") + void nullOutputFiles_returnsNull() throws Exception { + MultipartFile[] files = oneFile(); + List inputFiles = List.of(namedResource("input.pdf", "x".getBytes())); + when(processor.generateInputFiles(files)).thenReturn(inputFiles); + + PipelineResult result = new PipelineResult(); + result.setOutputFiles(null); + when(processor.runPipelineAgainstFiles(any(), any())).thenReturn(result); + + ResponseEntity response = controller.handleData(request(files, VALID_JSON)); + + assertNull(response); + } + } + + @Nested + @DisplayName("Single output file path") + class SingleOutput { + + @Test + @DisplayName("returns the single file streamed through a temp file") + void singleOutput_returnsFileResponse() throws Exception { + stubTempFiles(); + MultipartFile[] files = oneFile(); + List inputFiles = List.of(namedResource("input.pdf", "in".getBytes())); + when(processor.generateInputFiles(files)).thenReturn(inputFiles); + + byte[] outBytes = "single output body".getBytes(StandardCharsets.UTF_8); + Resource single = namedResource("result.pdf", outBytes); + PipelineResult result = new PipelineResult(); + result.setOutputFiles(List.of(single)); + when(processor.runPipelineAgainstFiles(any(), any())).thenReturn(result); + + ResponseEntity response = controller.handleData(request(files, VALID_JSON)); + + assertNotNull(response); + assertEquals(HttpStatus.OK, response.getStatusCode()); + assertNotNull(response.getBody()); + + // The controller copied the single file into a ".out" temp file. + verify(tempFileManager).createTempFile(".out"); + + // The streamed body matches what the processor produced. + try (InputStream is = response.getBody().getInputStream()) { + assertEquals( + new String(outBytes, StandardCharsets.UTF_8), + new String(is.readAllBytes(), StandardCharsets.UTF_8)); + } + } + + @Test + @DisplayName("captures analytics event with operations and file count") + void singleOutput_capturesAnalytics() throws Exception { + stubTempFiles(); + MultipartFile[] files = oneFile(); + when(processor.generateInputFiles(files)) + .thenReturn(List.of(namedResource("input.pdf", "in".getBytes()))); + + PipelineResult result = new PipelineResult(); + result.setOutputFiles(List.of(namedResource("result.pdf", "body".getBytes()))); + when(processor.runPipelineAgainstFiles(any(), any())).thenReturn(result); + + controller.handleData(request(files, VALID_JSON)); + + @SuppressWarnings("unchecked") + ArgumentCaptor> propsCaptor = ArgumentCaptor.forClass(Map.class); + verify(postHogService).captureEvent(eq("pipeline_api_event"), propsCaptor.capture()); + + Map props = propsCaptor.getValue(); + assertEquals(1, props.get("fileCount")); + assertEquals(List.of("/api/v1/misc/repair"), props.get("operations")); + } + } + + @Nested + @DisplayName("Multiple output files (zip) path") + class MultipleOutput { + + @Test + @DisplayName("zips multiple output files into output.zip") + void multipleOutputs_returnsZipResponse() throws Exception { + stubTempFiles(); + MultipartFile[] files = oneFile(); + when(processor.generateInputFiles(files)) + .thenReturn(List.of(namedResource("input.pdf", "in".getBytes()))); + + Resource a = namedResource("a.pdf", "alpha".getBytes(StandardCharsets.UTF_8)); + Resource b = namedResource("b.pdf", "bravo".getBytes(StandardCharsets.UTF_8)); + PipelineResult result = new PipelineResult(); + result.setOutputFiles(List.of(a, b)); + when(processor.runPipelineAgainstFiles(any(), any())).thenReturn(result); + + ResponseEntity response = controller.handleData(request(files, VALID_JSON)); + + assertNotNull(response); + assertEquals(HttpStatus.OK, response.getStatusCode()); + assertNotNull(response.getBody()); + verify(tempFileManager).createTempFile(".zip"); + + // Inspect zip contents: both entries present with their bytes. + Map entries = readZip(response.getBody()); + assertEquals(2, entries.size()); + assertEquals("alpha", entries.get("a.pdf")); + assertEquals("bravo", entries.get("b.pdf")); + } + + @Test + @DisplayName("duplicate filenames are de-duplicated within the zip") + void multipleOutputs_duplicateNames_areDeduped() throws Exception { + stubTempFiles(); + MultipartFile[] files = oneFile(); + when(processor.generateInputFiles(files)) + .thenReturn(List.of(namedResource("input.pdf", "in".getBytes()))); + + Resource first = namedResource("dup.pdf", "first".getBytes(StandardCharsets.UTF_8)); + Resource second = namedResource("dup.pdf", "second".getBytes(StandardCharsets.UTF_8)); + PipelineResult result = new PipelineResult(); + result.setOutputFiles(List.of(first, second)); + when(processor.runPipelineAgainstFiles(any(), any())).thenReturn(result); + + ResponseEntity response = controller.handleData(request(files, VALID_JSON)); + + assertNotNull(response); + Map entries = readZip(response.getBody()); + // Two distinct zip entry names despite the same source filename. + assertEquals(2, entries.size()); + } + + private Map readZip(Resource zipResource) throws IOException { + Map out = new java.util.LinkedHashMap<>(); + byte[] all; + try (InputStream is = zipResource.getInputStream()) { + all = is.readAllBytes(); + } + try (ZipInputStream zis = new ZipInputStream(new ByteArrayInputStream(all))) { + ZipEntry entry; + while ((entry = zis.getNextEntry()) != null) { + out.put( + entry.getName(), + new String(zis.readAllBytes(), StandardCharsets.UTF_8)); + zis.closeEntry(); + } + } + return out; + } + } + + @Nested + @DisplayName("Error handling") + class ErrorHandling { + + @Test + @DisplayName("returns null when the processor throws during input generation") + void processorThrowsOnGenerate_returnsNull() throws Exception { + MultipartFile[] files = oneFile(); + when(processor.generateInputFiles(files)).thenThrow(new RuntimeException("boom")); + + ResponseEntity response = controller.handleData(request(files, VALID_JSON)); + + assertNull(response); + } + + @Test + @DisplayName("returns null when the pipeline run throws") + void pipelineRunThrows_returnsNull() throws Exception { + MultipartFile[] files = oneFile(); + when(processor.generateInputFiles(files)) + .thenReturn(List.of(namedResource("input.pdf", "in".getBytes()))); + when(processor.runPipelineAgainstFiles(any(), any())) + .thenThrow(new RuntimeException("pipeline failed")); + + ResponseEntity response = controller.handleData(request(files, VALID_JSON)); + + assertNull(response); + } + + @Test + @DisplayName("invalid JSON config propagates as a parse exception") + void invalidJson_throws() { + MultipartFile[] files = oneFile(); + + org.junit.jupiter.api.Assertions.assertThrows( + Exception.class, () -> controller.handleData(request(files, "not-valid-json"))); + + verifyNoInteractions(postHogService); + } + } + + @Nested + @DisplayName("JSON config parsing") + class JsonParsing { + + @Test + @DisplayName("multiple operation names are extracted in order for analytics") + void multipleOperations_extractedInOrder() throws Exception { + MultipartFile[] files = oneFile(); + String json = + "{\"name\":\"multi\",\"pipeline\":[" + + "{\"operation\":\"/api/v1/misc/repair\",\"parameters\":{}}," + + "{\"operation\":\"/api/v1/security/sanitize-pdf\",\"parameters\":{}}]}"; + + // Stop after analytics by returning no input files. + when(processor.generateInputFiles(files)).thenReturn(null); + + controller.handleData(request(files, json)); + + @SuppressWarnings("unchecked") + ArgumentCaptor> propsCaptor = ArgumentCaptor.forClass(Map.class); + verify(postHogService).captureEvent(eq("pipeline_api_event"), propsCaptor.capture()); + + assertEquals( + List.of("/api/v1/misc/repair", "/api/v1/security/sanitize-pdf"), + propsCaptor.getValue().get("operations")); + } + + @Test + @DisplayName("real ObjectMapper deserializes the pipeline config used by analytics") + void objectMapperParsesConfig() { + PipelineConfig config = objectMapper.readValue(VALID_JSON, PipelineConfig.class); + assertEquals("test-pipeline", config.getName()); + assertEquals(1, config.getOperations().size()); + assertEquals("/api/v1/misc/repair", config.getOperations().get(0).getOperation()); + } + } +} diff --git a/app/core/src/test/java/stirling/software/SPDF/controller/api/security/ManualRedactionServiceTest.java b/app/core/src/test/java/stirling/software/SPDF/controller/api/security/ManualRedactionServiceTest.java new file mode 100644 index 0000000000..40d142e60c --- /dev/null +++ b/app/core/src/test/java/stirling/software/SPDF/controller/api/security/ManualRedactionServiceTest.java @@ -0,0 +1,669 @@ +package stirling.software.SPDF.controller.api.security; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertSame; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.Mockito.lenient; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.times; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.awt.Color; +import java.io.File; +import java.io.IOException; +import java.nio.file.Files; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.HashMap; +import java.util.List; +import java.util.Map; + +import org.apache.pdfbox.Loader; +import org.apache.pdfbox.pdmodel.PDDocument; +import org.apache.pdfbox.pdmodel.PDPage; +import org.apache.pdfbox.pdmodel.PDPageContentStream; +import org.apache.pdfbox.pdmodel.common.PDRectangle; +import org.apache.pdfbox.pdmodel.font.PDType1Font; +import org.apache.pdfbox.pdmodel.font.Standard14Fonts; +import org.apache.pdfbox.pdmodel.interactive.annotation.PDAnnotation; +import org.apache.pdfbox.pdmodel.interactive.annotation.PDAnnotationLink; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; + +import stirling.software.SPDF.model.PDFText; +import stirling.software.SPDF.model.api.security.ManualRedactPdfRequest; +import stirling.software.common.model.api.security.RedactionArea; +import stirling.software.common.util.TempFile; +import stirling.software.common.util.TempFileManager; + +@DisplayName("ManualRedactionService Tests") +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class ManualRedactionServiceTest { + + @Mock private TempFileManager tempFileManager; + + private ManualRedactionService service; + + // Track temp files created during finalize tests so we can clean them up. + private final List createdTempFiles = new ArrayList<>(); + + @BeforeEach + void setUp() throws Exception { + service = new ManualRedactionService(tempFileManager); + + // createManagedTempFile returns a TempFile backed by a real on-disk file so + // document.save() works and length() can be inspected. + lenient() + .when(tempFileManager.createManagedTempFile(anyString())) + .thenAnswer( + inv -> { + File f = + Files.createTempFile("redact-test", inv.getArgument(0)) + .toFile(); + createdTempFiles.add(f); + TempFile tf = mock(TempFile.class); + lenient().when(tf.getFile()).thenReturn(f); + lenient().when(tf.getPath()).thenReturn(f.toPath()); + return tf; + }); + } + + @AfterEach + void tearDown() { + for (File f : createdTempFiles) { + if (f != null && f.exists()) { + f.delete(); + } + } + createdTempFiles.clear(); + } + + // ----------------------------------------------------------------------- + // Helpers + // ----------------------------------------------------------------------- + + private static PDDocument newDocument(int pageCount) { + PDDocument doc = new PDDocument(); + for (int i = 0; i < pageCount; i++) { + doc.addPage(new PDPage(PDRectangle.A4)); + } + return doc; + } + + private static PDDocument newDocumentWithText() throws IOException { + PDDocument doc = new PDDocument(); + PDPage page = new PDPage(PDRectangle.A4); + doc.addPage(page); + try (PDPageContentStream cs = new PDPageContentStream(doc, page)) { + cs.beginText(); + cs.setFont(new PDType1Font(Standard14Fonts.FontName.HELVETICA), 12); + cs.newLineAtOffset(72, 700); + cs.showText("Sensitive text to redact"); + cs.endText(); + } + return doc; + } + + private static RedactionArea area( + Integer page, double x, double y, double width, double height, String color) { + RedactionArea a = new RedactionArea(); + a.setPage(page); + a.setX(x); + a.setY(y); + a.setWidth(width); + a.setHeight(height); + a.setColor(color); + return a; + } + + private static PDFText text(int pageIndex, float x1, float y1, float x2, float y2) { + return new PDFText(pageIndex, x1, y1, x2, y2, "redacted"); + } + + private static byte[] save(PDDocument doc) throws IOException { + java.io.ByteArrayOutputStream baos = new java.io.ByteArrayOutputStream(); + doc.save(baos); + return baos.toByteArray(); + } + + // ----------------------------------------------------------------------- + // decodeOrDefault + // ----------------------------------------------------------------------- + + @Nested + @DisplayName("decodeOrDefault") + class DecodeOrDefault { + + @Test + @DisplayName("null returns black") + void nullReturnsBlack() { + assertSame(Color.BLACK, ManualRedactionService.decodeOrDefault(null)); + } + + @Test + @DisplayName("hex with leading hash decodes") + void hexWithHash() { + assertEquals(Color.RED, ManualRedactionService.decodeOrDefault("#FF0000")); + } + + @Test + @DisplayName("hex without leading hash decodes") + void hexWithoutHash() { + assertEquals(Color.RED, ManualRedactionService.decodeOrDefault("FF0000")); + } + + @Test + @DisplayName("white decodes correctly") + void whiteDecodes() { + assertEquals(Color.WHITE, ManualRedactionService.decodeOrDefault("#FFFFFF")); + } + + @Test + @DisplayName("invalid hex falls back to black") + void invalidFallsBackToBlack() { + assertEquals(Color.BLACK, ManualRedactionService.decodeOrDefault("not-a-color")); + } + + @Test + @DisplayName("empty string falls back to black") + void emptyFallsBackToBlack() { + assertEquals(Color.BLACK, ManualRedactionService.decodeOrDefault("")); + } + } + + // ----------------------------------------------------------------------- + // redactAreas + // ----------------------------------------------------------------------- + + @Nested + @DisplayName("redactAreas") + class RedactAreas { + + @Test + @DisplayName("null list is a no-op") + void nullList() throws Exception { + try (PDDocument doc = newDocument(1)) { + byte[] before = save(doc); + service.redactAreas(null, doc, doc.getPages()); + // Document still saves and has its single page; nothing thrown. + assertEquals(1, doc.getNumberOfPages()); + assertNotNull(before); + } + } + + @Test + @DisplayName("empty list is a no-op") + void emptyList() throws Exception { + try (PDDocument doc = newDocument(1)) { + service.redactAreas(new ArrayList<>(), doc, doc.getPages()); + assertEquals(1, doc.getNumberOfPages()); + } + } + + @Test + @DisplayName("valid area is applied and document still saves") + void validArea() throws Exception { + try (PDDocument doc = newDocument(1)) { + List areas = Arrays.asList(area(1, 10, 10, 100, 50, "#000000")); + service.redactAreas(areas, doc, doc.getPages()); + byte[] out = save(doc); + assertTrue(out.length > 0); + // Reload to ensure the produced PDF is structurally valid. + try (PDDocument reloaded = Loader.loadPDF(out)) { + assertEquals(1, reloaded.getNumberOfPages()); + } + } + } + + @Test + @DisplayName("area with null page is skipped") + void nullPageSkipped() throws Exception { + try (PDDocument doc = newDocument(1)) { + List areas = Arrays.asList(area(null, 10, 10, 100, 50, "#000000")); + service.redactAreas(areas, doc, doc.getPages()); + assertEquals(1, doc.getNumberOfPages()); + } + } + + @Test + @DisplayName("area with non-positive page is skipped") + void nonPositivePageSkipped() throws Exception { + try (PDDocument doc = newDocument(1)) { + List areas = Arrays.asList(area(0, 10, 10, 100, 50, "#000000")); + service.redactAreas(areas, doc, doc.getPages()); + assertEquals(1, doc.getNumberOfPages()); + } + } + + @Test + @DisplayName("area with null/zero width or height is skipped") + void invalidDimensionsSkipped() throws Exception { + try (PDDocument doc = newDocument(1)) { + List areas = new ArrayList<>(); + areas.add(area(1, 10, 10, 0, 50, "#000000")); // zero width + areas.add(area(1, 10, 10, 100, 0, "#000000")); // zero height + RedactionArea nullWidth = area(1, 10, 10, 100, 50, "#000000"); + nullWidth.setWidth(null); + areas.add(nullWidth); + RedactionArea nullHeight = area(1, 10, 10, 100, 50, "#000000"); + nullHeight.setHeight(null); + areas.add(nullHeight); + service.redactAreas(areas, doc, doc.getPages()); + // None applied, but no exception and doc remains valid. + assertEquals(1, doc.getNumberOfPages()); + } + } + + @Test + @DisplayName("area on out-of-range page is skipped without error") + void outOfRangePageSkipped() throws Exception { + try (PDDocument doc = newDocument(1)) { + List areas = Arrays.asList(area(5, 10, 10, 100, 50, "#000000")); + service.redactAreas(areas, doc, doc.getPages()); + assertEquals(1, doc.getNumberOfPages()); + } + } + + @Test + @DisplayName("multiple areas across multiple pages are grouped per page") + void multiplePagesGrouped() throws Exception { + try (PDDocument doc = newDocument(3)) { + List areas = new ArrayList<>(); + areas.add(area(1, 10, 10, 50, 50, "#000000")); + areas.add(area(1, 80, 80, 50, 50, "#FF0000")); + areas.add(area(3, 20, 20, 60, 40, null)); // null color -> default black + service.redactAreas(areas, doc, doc.getPages()); + byte[] out = save(doc); + try (PDDocument reloaded = Loader.loadPDF(out)) { + assertEquals(3, reloaded.getNumberOfPages()); + } + } + } + } + + // ----------------------------------------------------------------------- + // redactPages + // ----------------------------------------------------------------------- + + @Nested + @DisplayName("redactPages") + class RedactPages { + + private ManualRedactPdfRequest request(String pageNumbers, String color) { + ManualRedactPdfRequest req = new ManualRedactPdfRequest(); + req.setPageNumbers(pageNumbers); + req.setPageRedactionColor(color); + return req; + } + + @Test + @DisplayName("redacts all pages when 'all' is given") + void redactsAllPages() throws Exception { + try (PDDocument doc = newDocument(3)) { + service.redactPages(request("all", "#000000"), doc, doc.getPages()); + byte[] out = save(doc); + try (PDDocument reloaded = Loader.loadPDF(out)) { + assertEquals(3, reloaded.getNumberOfPages()); + } + } + } + + @Test + @DisplayName("redacts a specific page range") + void redactsSpecificPages() throws Exception { + try (PDDocument doc = newDocument(5)) { + service.redactPages(request("1,3", "#FF0000"), doc, doc.getPages()); + byte[] out = save(doc); + try (PDDocument reloaded = Loader.loadPDF(out)) { + assertEquals(5, reloaded.getNumberOfPages()); + } + } + } + + @Test + @DisplayName("null page numbers defaults to first page") + void nullPageNumbers() throws Exception { + try (PDDocument doc = newDocument(2)) { + service.redactPages(request(null, "#000000"), doc, doc.getPages()); + assertEquals(2, doc.getNumberOfPages()); + } + } + + @Test + @DisplayName("null color falls back to black") + void nullColor() throws Exception { + try (PDDocument doc = newDocument(1)) { + service.redactPages(request("all", null), doc, doc.getPages()); + assertEquals(1, doc.getNumberOfPages()); + } + } + } + + // ----------------------------------------------------------------------- + // redactFoundText + // ----------------------------------------------------------------------- + + @Nested + @DisplayName("redactFoundText") + class RedactFoundText { + + @Test + @DisplayName("overlay-only mode draws boxes and document stays valid") + void overlayMode() throws Exception { + try (PDDocument doc = newDocument(1)) { + List blocks = Arrays.asList(text(0, 72, 690, 300, 710)); + service.redactFoundText(doc, blocks, 1.0f, Color.BLACK, false); + byte[] out = save(doc); + try (PDDocument reloaded = Loader.loadPDF(out)) { + assertEquals(1, reloaded.getNumberOfPages()); + } + } + } + + @Test + @DisplayName("text-removal mode narrows the box width") + void textRemovalMode() throws Exception { + try (PDDocument doc = newDocument(1)) { + List blocks = Arrays.asList(text(0, 72, 690, 300, 710)); + service.redactFoundText(doc, blocks, 0.0f, Color.RED, true); + assertEquals(1, doc.getNumberOfPages()); + } + } + + @Test + @DisplayName("blocks on out-of-range pages are skipped") + void outOfRangePageIndexSkipped() throws Exception { + try (PDDocument doc = newDocument(1)) { + List blocks = + Arrays.asList(text(0, 72, 690, 300, 710), text(9, 10, 10, 50, 50)); + service.redactFoundText(doc, blocks, 0.0f, Color.BLACK, false); + assertEquals(1, doc.getNumberOfPages()); + } + } + + @Test + @DisplayName("empty block list is a no-op") + void emptyBlocks() throws Exception { + try (PDDocument doc = newDocument(1)) { + service.redactFoundText(doc, new ArrayList<>(), 0.0f, Color.BLACK, false); + assertEquals(1, doc.getNumberOfPages()); + } + } + + @Test + @DisplayName("annotation overlapping a redacted block is removed") + void overlappingAnnotationRemoved() throws Exception { + try (PDDocument doc = newDocument(1)) { + PDPage page = doc.getPage(0); + float pageH = page.getBBox().getHeight(); + + // Redact block in PDF coords near y2=710 (top), x 72..300. + PDFText block = text(0, 72, 690, 300, 710); + + // Place an annotation rectangle that overlaps the redacted block. + // Block pdf-Y window (padding included) sits roughly around pageH-710..pageH-690. + PDAnnotationLink overlapping = new PDAnnotationLink(); + overlapping.setRectangle(new PDRectangle(80, pageH - 712, 120, 30)); + + // Place an annotation far away that should be kept. + PDAnnotationLink faraway = new PDAnnotationLink(); + faraway.setRectangle(new PDRectangle(10, 10, 20, 20)); + + page.setAnnotations(new ArrayList<>(Arrays.asList(overlapping, faraway))); + assertEquals(2, page.getAnnotations().size()); + + service.redactFoundText(doc, Arrays.asList(block), 1.0f, Color.BLACK, false); + + // getAnnotations() builds fresh wrappers each call, so compare by geometry + // rather than object identity. + List remaining = page.getAnnotations(); + assertEquals(1, remaining.size()); + PDRectangle keptRect = remaining.get(0).getRectangle(); + assertEquals(10f, keptRect.getLowerLeftX(), 0.001f); + assertEquals(10f, keptRect.getLowerLeftY(), 0.001f); + } + } + + @Test + @DisplayName("non-overlapping annotations are preserved") + void nonOverlappingAnnotationsKept() throws Exception { + try (PDDocument doc = newDocument(1)) { + PDPage page = doc.getPage(0); + PDAnnotationLink faraway = new PDAnnotationLink(); + faraway.setRectangle(new PDRectangle(10, 10, 20, 20)); + page.setAnnotations(new ArrayList<>(Arrays.asList(faraway))); + + PDFText block = text(0, 400, 100, 500, 120); + service.redactFoundText(doc, Arrays.asList(block), 0.0f, Color.BLACK, false); + + assertEquals(1, page.getAnnotations().size()); + PDRectangle keptRect = page.getAnnotations().get(0).getRectangle(); + assertEquals(10f, keptRect.getLowerLeftX(), 0.001f); + } + } + } + + // ----------------------------------------------------------------------- + // redactImageBoxes + // ----------------------------------------------------------------------- + + @Nested + @DisplayName("redactImageBoxes") + class RedactImageBoxes { + + @Test + @DisplayName("draws boxes for valid page indices") + void validBoxes() throws Exception { + try (PDDocument doc = newDocument(2)) { + List boxes = new ArrayList<>(); + boxes.add(new float[] {0, 10, 10, 100, 100}); + boxes.add(new float[] {1, 20, 20, 80, 80}); + service.redactImageBoxes(doc, boxes, Color.BLACK); + byte[] out = save(doc); + try (PDDocument reloaded = Loader.loadPDF(out)) { + assertEquals(2, reloaded.getNumberOfPages()); + } + } + } + + @Test + @DisplayName("out-of-range page indices are skipped without error") + void outOfRangeSkipped() throws Exception { + try (PDDocument doc = newDocument(1)) { + List boxes = new ArrayList<>(); + boxes.add(new float[] {-1, 10, 10, 100, 100}); // negative + boxes.add(new float[] {7, 20, 20, 80, 80}); // beyond page count + service.redactImageBoxes(doc, boxes, Color.BLACK); + assertEquals(1, doc.getNumberOfPages()); + } + } + + @Test + @DisplayName("empty box list is a no-op") + void emptyBoxes() throws Exception { + try (PDDocument doc = newDocument(1)) { + service.redactImageBoxes(doc, new ArrayList<>(), Color.BLACK); + assertEquals(1, doc.getNumberOfPages()); + } + } + + @Test + @DisplayName("multiple boxes on the same page are grouped") + void multipleBoxesSamePage() throws Exception { + try (PDDocument doc = newDocument(1)) { + List boxes = new ArrayList<>(); + boxes.add(new float[] {0, 10, 10, 50, 50}); + boxes.add(new float[] {0, 100, 100, 150, 150}); + service.redactImageBoxes(doc, boxes, Color.RED); + byte[] out = save(doc); + try (PDDocument reloaded = Loader.loadPDF(out)) { + assertEquals(1, reloaded.getNumberOfPages()); + } + } + } + } + + // ----------------------------------------------------------------------- + // extractPageElementBoxes + // ----------------------------------------------------------------------- + + @Nested + @DisplayName("extractPageElementBoxes") + class ExtractPageElementBoxes { + + @Test + @DisplayName("returns text line boxes for a page with text") + void extractsTextBoxes() throws Exception { + try (PDDocument doc = newDocumentWithText()) { + PDPage page = doc.getPage(0); + List boxes = service.extractPageElementBoxes(doc, page, 0); + assertNotNull(boxes); + assertFalse(boxes.isEmpty()); + // Each box must have 4 coordinates [x1, y1, x2, y2]. + for (float[] box : boxes) { + assertEquals(4, box.length); + } + } + } + + @Test + @DisplayName("returns empty list for a blank page") + void blankPageReturnsEmpty() throws Exception { + try (PDDocument doc = newDocument(1)) { + PDPage page = doc.getPage(0); + List boxes = service.extractPageElementBoxes(doc, page, 0); + assertNotNull(boxes); + assertTrue(boxes.isEmpty()); + } + } + } + + // ----------------------------------------------------------------------- + // finalizeRedaction (non-image path only; image path needs Spring context) + // ----------------------------------------------------------------------- + + @Nested + @DisplayName("finalizeRedaction") + class FinalizeRedaction { + + @Test + @DisplayName("with found text saves a managed temp file") + void withFoundText() throws Exception { + try (PDDocument doc = newDocument(1)) { + Map> byPage = new HashMap<>(); + byPage.put(0, new ArrayList<>(Arrays.asList(text(0, 72, 690, 300, 710)))); + + TempFile result = + service.finalizeRedaction(doc, byPage, "#000000", 1.0f, false, false); + + assertNotNull(result); + assertNotNull(result.getFile()); + assertTrue(result.getFile().exists()); + assertTrue(result.getFile().length() > 0); + verify(tempFileManager, times(1)).createManagedTempFile(".pdf"); + + // Output should be a loadable PDF. + try (PDDocument reloaded = Loader.loadPDF(result.getFile())) { + assertEquals(1, reloaded.getNumberOfPages()); + } + } + } + + @Test + @DisplayName("with no found text still saves the document") + void withoutFoundText() throws Exception { + try (PDDocument doc = newDocument(2)) { + Map> byPage = new HashMap<>(); + + TempFile result = + service.finalizeRedaction(doc, byPage, "#000000", 0.0f, false, false); + + assertNotNull(result); + assertTrue(result.getFile().exists()); + assertTrue(result.getFile().length() > 0); + verify(tempFileManager, times(1)).createManagedTempFile(".pdf"); + + try (PDDocument reloaded = Loader.loadPDF(result.getFile())) { + assertEquals(2, reloaded.getNumberOfPages()); + } + } + } + + @Test + @DisplayName("convertToImage null is treated as non-image path") + void convertToImageNull() throws Exception { + try (PDDocument doc = newDocument(1)) { + Map> byPage = new HashMap<>(); + + TempFile result = + service.finalizeRedaction(doc, byPage, "#000000", 0.0f, null, false); + + assertNotNull(result); + assertTrue(result.getFile().exists()); + // Non-image path uses the standard ".pdf" managed temp file exactly once. + verify(tempFileManager, times(1)).createManagedTempFile(".pdf"); + try (PDDocument reloaded = Loader.loadPDF(result.getFile())) { + assertEquals(1, reloaded.getNumberOfPages()); + } + } + } + + @Test + @DisplayName("text-removal mode finalize produces valid PDF") + void textRemovalFinalize() throws Exception { + try (PDDocument doc = newDocument(1)) { + Map> byPage = new HashMap<>(); + byPage.put(0, new ArrayList<>(Arrays.asList(text(0, 72, 690, 300, 710)))); + + TempFile result = + service.finalizeRedaction(doc, byPage, "#FF0000", 2.0f, false, true); + + assertNotNull(result); + try (PDDocument reloaded = Loader.loadPDF(result.getFile())) { + assertEquals(1, reloaded.getNumberOfPages()); + } + } + } + + @Test + @DisplayName("IOException from save closes the temp file and is rethrown") + void saveFailureClosesTempFile() throws Exception { + // Use a document mock that throws on save to exercise the catch/close branch. + PDDocument doc = newDocument(1); + Map> byPage = new HashMap<>(); + + // Point the managed temp file at a directory path so save() fails with IOException. + File dirAsFile = Files.createTempDirectory("redact-dir").toFile(); + createdTempFiles.add(dirAsFile); + TempFile failing = mock(TempFile.class); + when(failing.getFile()).thenReturn(dirAsFile); + when(tempFileManager.createManagedTempFile(anyString())).thenReturn(failing); + + try { + assertThrows( + IOException.class, + () -> + service.finalizeRedaction( + doc, byPage, "#000000", 0.0f, false, false)); + // The failing temp file is closed on the error path. + verify(failing).close(); + } finally { + doc.close(); + } + } + } +} diff --git a/app/core/src/test/java/stirling/software/SPDF/controller/api/security/TextRedactionServiceTest.java b/app/core/src/test/java/stirling/software/SPDF/controller/api/security/TextRedactionServiceTest.java new file mode 100644 index 0000000000..d487848e09 --- /dev/null +++ b/app/core/src/test/java/stirling/software/SPDF/controller/api/security/TextRedactionServiceTest.java @@ -0,0 +1,566 @@ +package stirling.software.SPDF.controller.api.security; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertSame; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.io.IOException; +import java.lang.reflect.Method; +import java.util.ArrayList; +import java.util.Collections; +import java.util.HashMap; +import java.util.LinkedHashSet; +import java.util.List; +import java.util.Map; +import java.util.Set; + +import org.apache.pdfbox.contentstream.operator.Operator; +import org.apache.pdfbox.cos.COSArray; +import org.apache.pdfbox.cos.COSString; +import org.apache.pdfbox.pdfparser.PDFStreamParser; +import org.apache.pdfbox.pdmodel.PDDocument; +import org.apache.pdfbox.pdmodel.PDPage; +import org.apache.pdfbox.pdmodel.PDPageContentStream; +import org.apache.pdfbox.pdmodel.common.PDRectangle; +import org.apache.pdfbox.pdmodel.font.PDFont; +import org.apache.pdfbox.pdmodel.font.PDType1Font; +import org.apache.pdfbox.pdmodel.font.Standard14Fonts; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; + +import stirling.software.SPDF.model.PDFText; + +/** + * Unit tests for {@link TextRedactionService}. Exercises the pure text-matching / placeholder + * logic, the public find/replace entry points, and the small helper predicates. PDFs are built tiny + * and in-memory with real Standard14 fonts so the redaction pipeline runs end-to-end without + * external processes. + */ +class TextRedactionServiceTest { + + private static final float FONT_SIZE = 12f; + private static final float LEFT_X = 72f; + private static final float TOP_Y = PDRectangle.LETTER.getHeight() - 80f; + + private final TextRedactionService service = new TextRedactionService(); + + // ── fixtures ───────────────────────────────────────────────────────────────────────────────── + + private PDFont helvetica() { + return new PDType1Font(Standard14Fonts.FontName.HELVETICA); + } + + /** Single page, single Tj line per supplied text line, Helvetica 12. */ + private PDDocument buildDoc(String... lines) throws IOException { + PDDocument doc = new PDDocument(); + PDPage page = new PDPage(PDRectangle.LETTER); + doc.addPage(page); + try (PDPageContentStream cs = new PDPageContentStream(doc, page)) { + cs.setFont(helvetica(), FONT_SIZE); + for (int i = 0; i < lines.length; i++) { + cs.beginText(); + cs.newLineAtOffset(LEFT_X, TOP_Y - i * 16f); + cs.showText(lines[i]); + cs.endText(); + } + } + return doc; + } + + private PDDocument buildEmptyDoc() { + PDDocument doc = new PDDocument(); + doc.addPage(new PDPage(PDRectangle.LETTER)); + return doc; + } + + // ── isTextShowingOperator ──────────────────────────────────────────────────────────────────── + + @Nested + @DisplayName("isTextShowingOperator") + class IsTextShowingOperator { + + @Test + @DisplayName("recognises the four text-showing operators") + void recognisesTextShowingOperators() { + assertTrue(service.isTextShowingOperator("Tj")); + assertTrue(service.isTextShowingOperator("TJ")); + assertTrue(service.isTextShowingOperator("'")); + assertTrue(service.isTextShowingOperator("\"")); + } + + @Test + @DisplayName("rejects non text-showing operators and junk") + void rejectsOthers() { + assertFalse(service.isTextShowingOperator("BT")); + assertFalse(service.isTextShowingOperator("ET")); + assertFalse(service.isTextShowingOperator("Tf")); + assertFalse(service.isTextShowingOperator("")); + assertFalse(service.isTextShowingOperator("tj")); + } + } + + // ── findTextToRedact ───────────────────────────────────────────────────────────────────────── + + @Nested + @DisplayName("findTextToRedact") + class FindTextToRedact { + + @Test + @DisplayName("returns empty map when all search terms are blank") + void emptyTermsReturnsEmptyMap() throws IOException { + try (PDDocument doc = buildDoc("Confidential data here")) { + Map> result = + service.findTextToRedact(doc, new String[] {"", " "}, false, false); + assertTrue(result.isEmpty()); + } + } + + @Test + @DisplayName("returns empty map for an empty term array") + void emptyArrayReturnsEmptyMap() throws IOException { + try (PDDocument doc = buildDoc("Some text")) { + Map> result = + service.findTextToRedact(doc, new String[] {}, false, false); + assertTrue(result.isEmpty()); + } + } + + @Test + @DisplayName("finds a literal term on the page") + void findsLiteralTerm() throws IOException { + try (PDDocument doc = buildDoc("Hello SECRET world")) { + Map> result = + service.findTextToRedact(doc, new String[] {"SECRET"}, false, false); + + assertFalse(result.isEmpty()); + assertTrue(result.containsKey(0), "match should be on page index 0"); + List hits = result.get(0); + assertEquals(1, hits.size()); + assertEquals("SECRET", hits.get(0).getText()); + } + } + + @Test + @DisplayName("does not find a term that is absent") + void doesNotFindAbsentTerm() throws IOException { + try (PDDocument doc = buildDoc("Hello world")) { + Map> result = + service.findTextToRedact(doc, new String[] {"NOTHERE"}, false, false); + assertTrue(result.isEmpty()); + } + } + + @Test + @DisplayName("regex term matches digit runs") + void regexTermMatchesDigits() throws IOException { + try (PDDocument doc = buildDoc("Order 12345 shipped")) { + Map> result = + service.findTextToRedact(doc, new String[] {"\\d+"}, true, false); + + assertFalse(result.isEmpty()); + assertEquals("12345", result.get(0).get(0).getText()); + } + } + + @Test + @DisplayName("whole-word search does not match a substring inside a larger word") + void wholeWordDoesNotMatchSubstring() throws IOException { + try (PDDocument doc = buildDoc("classification of cats")) { + Map> result = + service.findTextToRedact(doc, new String[] {"cat"}, false, true); + // "cat" appears only inside "classification"; whole-word must not match it. + assertTrue(result.isEmpty()); + } + } + + @Test + @DisplayName("whole-word search matches a standalone word") + void wholeWordMatchesStandalone() throws IOException { + try (PDDocument doc = buildDoc("the cat sat")) { + Map> result = + service.findTextToRedact(doc, new String[] {"cat"}, false, true); + assertFalse(result.isEmpty()); + assertEquals("cat", result.get(0).get(0).getText()); + } + } + + @Test + @DisplayName("trims whitespace around terms before searching") + void trimsTermsBeforeSearch() throws IOException { + try (PDDocument doc = buildDoc("padded SECRET value")) { + Map> result = + service.findTextToRedact(doc, new String[] {" SECRET "}, false, false); + assertFalse(result.isEmpty()); + assertEquals("SECRET", result.get(0).get(0).getText()); + } + } + } + + // ── performTextReplacement ─────────────────────────────────────────────────────────────────── + + @Nested + @DisplayName("performTextReplacement") + class PerformTextReplacement { + + @Test + @DisplayName("returns false (no fallback) when there is nothing found to redact") + void noFoundTextReturnsFalse() throws IOException { + try (PDDocument doc = buildDoc("anything")) { + boolean fallback = + service.performTextReplacement( + doc, new HashMap<>(), new String[] {"x"}, false, false); + assertFalse(fallback, "empty found-text map must short-circuit to no fallback"); + } + } + + @Test + @DisplayName("replaces text on a standard-font document without requesting box fallback") + void replacesOnStandardFont() throws IOException { + try (PDDocument doc = buildDoc("Please redact SECRET now")) { + Map> found = + service.findTextToRedact(doc, new String[] {"SECRET"}, false, false); + assertFalse(found.isEmpty()); + + boolean fallback = + service.performTextReplacement( + doc, found, new String[] {"SECRET"}, false, false); + + assertFalse(fallback, "standard Helvetica should not trigger box-only fallback"); + // After replacement the literal term should no longer be extractable. + Map> afterFound = + service.findTextToRedact(doc, new String[] {"SECRET"}, false, false); + assertTrue(afterFound.isEmpty(), "SECRET should be gone after text replacement"); + } + } + + @Test + @DisplayName("leaves non-targeted text intact after replacement") + void leavesOtherTextIntact() throws IOException { + try (PDDocument doc = buildDoc("KEEP this but redact SECRET part")) { + Map> found = + service.findTextToRedact(doc, new String[] {"SECRET"}, false, false); + + service.performTextReplacement(doc, found, new String[] {"SECRET"}, false, false); + + Map> keepStill = + service.findTextToRedact(doc, new String[] {"KEEP"}, false, false); + assertFalse(keepStill.isEmpty(), "untargeted word KEEP must survive redaction"); + } + } + } + + // ── detectCustomEncodingFonts ──────────────────────────────────────────────────────────────── + + @Nested + @DisplayName("detectCustomEncodingFonts") + class DetectCustomEncodingFonts { + + @Test + @DisplayName("standard Helvetica document is not flagged as custom-encoded") + void standardFontNotFlagged() throws IOException { + try (PDDocument doc = buildDoc("plain helvetica text")) { + assertFalse(service.detectCustomEncodingFonts(doc)); + } + } + + @Test + @DisplayName("document with no content / no fonts is not flagged") + void emptyDocumentNotFlagged() throws IOException { + try (PDDocument doc = buildEmptyDoc()) { + assertFalse(service.detectCustomEncodingFonts(doc)); + } + } + } + + // ── createPlaceholderWithFont ──────────────────────────────────────────────────────────────── + + @Nested + @DisplayName("createPlaceholderWithFont") + class CreatePlaceholderWithFont { + + @Test + @DisplayName("returns the input unchanged for null") + void nullReturnsNull() { + assertNull(service.createPlaceholderWithFont(null, helvetica())); + } + + @Test + @DisplayName("returns the input unchanged for empty string") + void emptyReturnsEmpty() { + assertEquals("", service.createPlaceholderWithFont("", helvetica())); + } + + @Test + @DisplayName("non-subset font yields spaces matching the original length") + void nonSubsetFontYieldsMatchingSpaces() { + String placeholder = service.createPlaceholderWithFont("hidden", helvetica()); + assertEquals(" ".repeat("hidden".length()), placeholder); + } + + @Test + @DisplayName("null font is treated as non-subset and yields spaces") + void nullFontYieldsSpaces() { + String placeholder = service.createPlaceholderWithFont("abc", null); + assertEquals(" ", placeholder); + } + } + + // ── createPlaceholderWithWidth ─────────────────────────────────────────────────────────────── + + @Nested + @DisplayName("createPlaceholderWithWidth") + class CreatePlaceholderWithWidth { + + @Test + @DisplayName("returns the input unchanged for null") + void nullReturnsNull() { + assertNull(service.createPlaceholderWithWidth(null, 10f, helvetica(), FONT_SIZE)); + } + + @Test + @DisplayName("returns the input unchanged for empty string") + void emptyReturnsEmpty() { + assertEquals("", service.createPlaceholderWithWidth("", 10f, helvetica(), FONT_SIZE)); + } + + @Test + @DisplayName("null font falls back to one space per original character") + void nullFontFallsBackToSpaces() { + String placeholder = service.createPlaceholderWithWidth("word", 50f, null, FONT_SIZE); + assertEquals(" ".repeat("word".length()), placeholder); + } + + @Test + @DisplayName("non-positive font size falls back to one space per original character") + void nonPositiveFontSizeFallsBackToSpaces() { + String placeholder = service.createPlaceholderWithWidth("word", 50f, helvetica(), 0f); + assertEquals(" ".repeat("word".length()), placeholder); + } + + @Test + @DisplayName("standard font produces a non-null all-whitespace placeholder") + void standardFontProducesWhitespacePlaceholder() { + PDFont font = helvetica(); + float fontSize = FONT_SIZE; + String original = "Secret"; + // Compute a realistic target width the way the service does (text-space / 1000 * size). + float targetWidth; + try { + targetWidth = font.getStringWidth(original) / 1000f * fontSize; + } catch (IOException e) { + targetWidth = 30f; + } + + String placeholder = + service.createPlaceholderWithWidth(original, targetWidth, font, fontSize); + + assertNotNull(placeholder); + assertFalse(placeholder.isEmpty(), "Helvetica supports spaces, so non-empty expected"); + assertTrue( + placeholder.chars().allMatch(c -> c == ' '), + "placeholder should be composed only of spaces"); + } + } + + // ── createTokensWithoutTargetText / writeFilteredContentStream + // ──────────────────────────────── + + @Nested + @DisplayName("createTokensWithoutTargetText") + class CreateTokensWithoutTargetText { + + @Test + @DisplayName( + "returns a non-empty token list and preserves token count when nothing matches") + void noMatchPreservesTokens() throws IOException { + try (PDDocument doc = buildDoc("nothing to hide")) { + PDPage page = doc.getPage(0); + List originalTokens = parseTokens(page); + + List tokens = + service.createTokensWithoutTargetText( + doc, page, Set.of("ABSENT"), false, false); + + assertNotNull(tokens); + assertEquals( + originalTokens.size(), + tokens.size(), + "token count should be unchanged when nothing matched"); + } + } + + @Test + @DisplayName("filtered tokens can be written back and the page re-parses cleanly") + void filteredTokensRoundTrip() throws IOException { + try (PDDocument doc = buildDoc("redact SECRET token roundtrip")) { + PDPage page = doc.getPage(0); + + List tokens = + service.createTokensWithoutTargetText( + doc, page, Set.of("SECRET"), false, false); + assertNotNull(tokens); + + service.writeFilteredContentStream(doc, page, tokens); + + // The page must still hold valid content (at least one operator token). + List reparsed = parseTokens(page); + boolean hasOperator = reparsed.stream().anyMatch(t -> t instanceof Operator); + assertTrue(hasOperator, "rewritten content stream must contain operators"); + } + } + + @Test + @DisplayName("empty target-word set leaves tokens untouched") + void emptyTargetSetLeavesTokens() throws IOException { + try (PDDocument doc = buildDoc("some content")) { + PDPage page = doc.getPage(0); + List originalTokens = parseTokens(page); + + List tokens = + service.createTokensWithoutTargetText( + doc, page, Collections.emptySet(), false, false); + + assertEquals(originalTokens.size(), tokens.size()); + } + } + + private List parseTokens(PDPage page) throws IOException { + PDFStreamParser parser = new PDFStreamParser(page); + List tokens = new ArrayList<>(); + Object token; + while ((token = parser.parseNextToken()) != null) { + tokens.add(token); + } + return tokens; + } + } + + // ── inner data classes ─────────────────────────────────────────────────────────────────────── + + @Nested + @DisplayName("TextSegment / MatchRange data classes") + class DataClasses { + + @Test + @DisplayName("TextSegment exposes its constructor values via accessors") + void textSegmentAccessors() { + PDFont font = helvetica(); + TextRedactionService.TextSegment segment = + new TextRedactionService.TextSegment(3, "Tj", "hello", 10, 15, font, 12f); + + assertEquals(3, segment.getTokenIndex()); + assertEquals("Tj", segment.getOperatorName()); + assertEquals("hello", segment.getText()); + assertEquals(10, segment.getStartPos()); + assertEquals(15, segment.getEndPos()); + assertSame(font, segment.getFont()); + assertEquals(12f, segment.getFontSize()); + } + + @Test + @DisplayName("MatchRange exposes start and end positions") + void matchRangeAccessors() { + TextRedactionService.MatchRange range = new TextRedactionService.MatchRange(4, 9); + assertEquals(4, range.getStartPos()); + assertEquals(9, range.getEndPos()); + } + + @Test + @DisplayName("MatchRange equality follows its data fields") + void matchRangeEquality() { + assertEquals( + new TextRedactionService.MatchRange(1, 5), + new TextRedactionService.MatchRange(1, 5)); + assertNotEquals( + new TextRedactionService.MatchRange(1, 5), + new TextRedactionService.MatchRange(1, 6)); + } + + private void assertNotEquals(Object a, Object b) { + assertFalse(a.equals(b)); + } + } + + // ── private logic exercised via reflection ─────────────────────────────────────────────────── + + @Nested + @DisplayName("findAllMatches / buildCompleteText (private logic via reflection)") + class PrivateLogic { + + @Test + @DisplayName("findAllMatches returns sorted, non-overlapping match ranges for two terms") + @SuppressWarnings("unchecked") + void findAllMatchesSorted() throws Exception { + String complete = "alpha beta gamma beta"; + Set terms = new LinkedHashSet<>(List.of("beta", "alpha")); + + Method m = + TextRedactionService.class.getDeclaredMethod( + "findAllMatches", + String.class, + Set.class, + boolean.class, + boolean.class); + m.setAccessible(true); + List matches = + (List) + m.invoke(service, complete, terms, false, false); + + assertNotNull(matches); + assertFalse(matches.isEmpty()); + // Results are sorted by start position. + for (int i = 1; i < matches.size(); i++) { + assertTrue( + matches.get(i - 1).getStartPos() <= matches.get(i).getStartPos(), + "matches must be sorted ascending by start position"); + } + // "alpha" at 0, "beta" at 6 and 17 -> three matches total. + assertEquals(3, matches.size()); + assertEquals(0, matches.get(0).getStartPos()); + } + + @Test + @DisplayName("findAllMatches returns nothing when no term occurs") + @SuppressWarnings("unchecked") + void findAllMatchesEmptyWhenAbsent() throws Exception { + Method m = + TextRedactionService.class.getDeclaredMethod( + "findAllMatches", + String.class, + Set.class, + boolean.class, + boolean.class); + m.setAccessible(true); + List matches = + (List) + m.invoke(service, "no terms here", Set.of("XYZ"), false, false); + assertTrue(matches.isEmpty()); + } + + @Test + @DisplayName("extractTextFromToken pulls text from Tj COSString and TJ COSArray") + void extractTextFromToken() throws Exception { + Method m = + TextRedactionService.class.getDeclaredMethod( + "extractTextFromToken", Object.class, String.class); + m.setAccessible(true); + + assertEquals("hi", m.invoke(service, new COSString("hi"), "Tj")); + assertEquals("hi", m.invoke(service, new COSString("hi"), "'")); + + COSArray tjArray = new COSArray(); + tjArray.add(new COSString("foo")); + tjArray.add(new COSString("bar")); + assertEquals("foobar", m.invoke(service, tjArray, "TJ")); + + // Unknown operator yields empty string. + assertEquals("", m.invoke(service, new COSString("x"), "Td")); + // Wrong token type for the operator yields empty string. + assertEquals("", m.invoke(service, new COSArray(), "Tj")); + } + } +} diff --git a/app/core/src/test/java/stirling/software/SPDF/exception/GlobalExceptionHandlerTest.java b/app/core/src/test/java/stirling/software/SPDF/exception/GlobalExceptionHandlerTest.java index 8e2275a728..4de6398e43 100644 --- a/app/core/src/test/java/stirling/software/SPDF/exception/GlobalExceptionHandlerTest.java +++ b/app/core/src/test/java/stirling/software/SPDF/exception/GlobalExceptionHandlerTest.java @@ -18,6 +18,7 @@ import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; import org.springframework.context.MessageSource; import org.springframework.core.env.Environment; +import org.springframework.http.HttpMethod; import org.springframework.http.HttpStatus; import org.springframework.http.ProblemDetail; import org.springframework.http.ResponseEntity; @@ -27,6 +28,7 @@ import org.springframework.web.bind.MissingServletRequestParameterException; import org.springframework.web.multipart.MaxUploadSizeExceededException; import org.springframework.web.multipart.support.MissingServletRequestPartException; import org.springframework.web.servlet.NoHandlerFoundException; +import org.springframework.web.servlet.resource.NoResourceFoundException; import jakarta.servlet.http.HttpServletRequest; import jakarta.servlet.http.HttpServletResponse; @@ -208,6 +210,19 @@ class GlobalExceptionHandlerTest { assertEquals(HttpStatus.NOT_FOUND, resp.getStatusCode()); } + // ---- NoResourceFoundException ---- + // Regression guard: was falling through to the 500 catch-all. + + @Test + void handleNoResourceFound_returns_404_not_500() { + when(request.getMethod()).thenReturn("GET"); + NoResourceFoundException ex = + new NoResourceFoundException(HttpMethod.GET, "/api/v1/storage/folders", ""); + ResponseEntity resp = handler.handleNoResourceFound(ex, request); + assertEquals(HttpStatus.NOT_FOUND, resp.getStatusCode()); + assertEquals("GET", resp.getBody().getProperties().get("method")); + } + // ---- IllegalArgumentException ---- @Test diff --git a/app/core/src/test/java/stirling/software/SPDF/model/api/converters/ConvertPDFToMarkdownTest.java b/app/core/src/test/java/stirling/software/SPDF/model/api/converters/ConvertPDFToMarkdownTest.java index b63e58b524..3bd6b7fadb 100644 --- a/app/core/src/test/java/stirling/software/SPDF/model/api/converters/ConvertPDFToMarkdownTest.java +++ b/app/core/src/test/java/stirling/software/SPDF/model/api/converters/ConvertPDFToMarkdownTest.java @@ -1,16 +1,17 @@ package stirling.software.SPDF.model.api.converters; -import static org.junit.jupiter.api.Assertions.assertEquals; import static org.mockito.ArgumentMatchers.any; import static org.mockito.Mockito.*; import static org.springframework.test.web.servlet.request.MockMvcRequestBuilders.multipart; import static org.springframework.test.web.servlet.result.MockMvcResultMatchers.*; +import java.io.File; import java.nio.charset.StandardCharsets; +import java.nio.file.Path; import org.junit.jupiter.api.Test; -import org.mockito.ArgumentCaptor; import org.mockito.MockedConstruction; +import org.mockito.MockedStatic; import org.mockito.Mockito; import org.springframework.core.io.ByteArrayResource; import org.springframework.core.io.Resource; @@ -21,9 +22,10 @@ import org.springframework.test.web.servlet.MockMvc; import org.springframework.test.web.servlet.setup.MockMvcBuilders; import org.springframework.web.bind.annotation.ExceptionHandler; import org.springframework.web.bind.annotation.RestControllerAdvice; -import org.springframework.web.multipart.MultipartFile; -import stirling.software.common.util.PDFToFile; +import stirling.software.common.pdf.PdfMarkdownConverter; +import stirling.software.common.util.TempFile; +import stirling.software.jpdfium.PdfDocument; class ConvertPDFToMarkdownTest { @@ -47,68 +49,68 @@ class ConvertPDFToMarkdownTest { @Test void pdfToMarkdownReturnsMarkdownBytes() throws Exception { byte[] md = "# heading\n\ncontent\n".getBytes(StandardCharsets.UTF_8); + String expectedMd = "# heading\n\ncontent\n"; - try (MockedConstruction construction = - Mockito.mockConstruction( - PDFToFile.class, - (mock, ctx) -> { - when(mock.processPdfToMarkdown(any(MultipartFile.class))) - .thenAnswer( - inv -> - ResponseEntity.ok() - .header("Content-Type", "text/markdown") - .body(new ByteArrayResource(md))); - })) { + File tmpFile = File.createTempFile("test", ".pdf"); + tmpFile.deleteOnExit(); - MockMvc mvc = mockMvc(); + try (MockedConstruction tempMock = + Mockito.mockConstruction( + TempFile.class, + (mock, ctx) -> { + when(mock.getFile()).thenReturn(tmpFile); + when(mock.getPath()).thenReturn(tmpFile.toPath()); + }); + MockedStatic docStatic = Mockito.mockStatic(PdfDocument.class); + MockedConstruction converterMock = + Mockito.mockConstruction( + PdfMarkdownConverter.class, + (mock, ctx) -> when(mock.convert(any())).thenReturn(expectedMd))) { + + PdfDocument mockDoc = Mockito.mock(PdfDocument.class); + docStatic.when(() -> PdfDocument.open(any(Path.class))).thenReturn(mockDoc); MockMultipartFile file = new MockMultipartFile( - "fileInput", // must match the field name in PDFFile - "input.pdf", - "application/pdf", - new byte[] {1, 2, 3}); + "fileInput", "input.pdf", "application/pdf", new byte[] {1, 2, 3}); - // ResponseEntity is written synchronously on the request thread, - // so there is no async dispatch to wait for (unlike the old StreamingResponseBody - // path). - mvc.perform(multipart("/api/v1/convert/pdf/markdown").file(file)) + mockMvc() + .perform(multipart("/api/v1/convert/pdf/markdown").file(file)) .andExpect(status().isOk()) .andExpect(header().string("Content-Type", "text/markdown")) .andExpect(content().bytes(md)); - - // Verify that exactly one instance was created - assert construction.constructed().size() == 1; - - // And that the uploaded file was passed to processPdfToMarkdown() - PDFToFile created = construction.constructed().get(0); - ArgumentCaptor captor = ArgumentCaptor.forClass(MultipartFile.class); - verify(created, times(1)).processPdfToMarkdown(captor.capture()); - MultipartFile passed = captor.getValue(); - - // Minimal plausibility checks - assertEquals("input.pdf", passed.getOriginalFilename()); - assertEquals("application/pdf", passed.getContentType()); } } @Test void pdfToMarkdownWhenServiceThrowsReturns500() throws Exception { - try (MockedConstruction ignored = - Mockito.mockConstruction( - PDFToFile.class, - (mock, ctx) -> { - when(mock.processPdfToMarkdown(any(MultipartFile.class))) - .thenThrow(new RuntimeException("boom")); - })) { + File tmpFile = File.createTempFile("test", ".pdf"); + tmpFile.deleteOnExit(); - MockMvc mvc = mockMvc(); + try (MockedConstruction tempMock = + Mockito.mockConstruction( + TempFile.class, + (mock, ctx) -> { + when(mock.getFile()).thenReturn(tmpFile); + when(mock.getPath()).thenReturn(tmpFile.toPath()); + }); + MockedStatic docStatic = Mockito.mockStatic(PdfDocument.class); + MockedConstruction converterMock = + Mockito.mockConstruction( + PdfMarkdownConverter.class, + (mock, ctx) -> + when(mock.convert(any())) + .thenThrow(new RuntimeException("boom")))) { + + PdfDocument mockDoc = Mockito.mock(PdfDocument.class); + docStatic.when(() -> PdfDocument.open(any(Path.class))).thenReturn(mockDoc); MockMultipartFile file = new MockMultipartFile( "fileInput", "x.pdf", "application/pdf", new byte[] {0x01}); - mvc.perform(multipart("/api/v1/convert/pdf/markdown").file(file)) + mockMvc() + .perform(multipart("/api/v1/convert/pdf/markdown").file(file)) .andExpect(status().isInternalServerError()); } } diff --git a/app/core/src/test/java/stirling/software/SPDF/service/PdfJsonConversionServiceGapTest.java b/app/core/src/test/java/stirling/software/SPDF/service/PdfJsonConversionServiceGapTest.java new file mode 100644 index 0000000000..5178b779f4 --- /dev/null +++ b/app/core/src/test/java/stirling/software/SPDF/service/PdfJsonConversionServiceGapTest.java @@ -0,0 +1,654 @@ +package stirling.software.SPDF.service; + +import static org.junit.jupiter.api.Assertions.assertDoesNotThrow; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.when; + +import java.io.ByteArrayOutputStream; +import java.io.File; +import java.io.IOException; +import java.nio.charset.StandardCharsets; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.ArrayList; +import java.util.Base64; +import java.util.List; +import java.util.concurrent.atomic.AtomicBoolean; + +import org.apache.pdfbox.Loader; +import org.apache.pdfbox.pdmodel.PDDocument; +import org.apache.pdfbox.pdmodel.PDDocumentInformation; +import org.apache.pdfbox.pdmodel.PDPage; +import org.apache.pdfbox.pdmodel.common.PDRectangle; +import org.apache.pdfbox.pdmodel.font.PDType1Font; +import org.apache.pdfbox.pdmodel.font.Standard14Fonts; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.quality.Strictness; +import org.springframework.mock.web.MockMultipartFile; + +import stirling.software.SPDF.config.EndpointConfiguration; +import stirling.software.SPDF.exception.CacheUnavailableException; +import stirling.software.SPDF.model.json.PdfJsonDocument; +import stirling.software.SPDF.model.json.PdfJsonFont; +import stirling.software.SPDF.model.json.PdfJsonMetadata; +import stirling.software.SPDF.model.json.PdfJsonPage; +import stirling.software.SPDF.service.pdfjson.PdfJsonFontService; +import stirling.software.SPDF.service.pdfjson.type3.Type3FontConversionService; +import stirling.software.SPDF.service.pdfjson.type3.Type3GlyphExtractor; +import stirling.software.common.model.ApplicationProperties; +import stirling.software.common.service.CustomPDFDocumentFactory; +import stirling.software.common.service.TaskManager; +import stirling.software.common.util.TempFileManager; + +import tools.jackson.databind.DeserializationFeature; +import tools.jackson.databind.ObjectMapper; +import tools.jackson.databind.json.JsonMapper; + +/** + * Gap tests for {@link PdfJsonConversionService} that exercise the public conversion entrypoints + * and the cache-backed public API. These complement {@code + * PdfJsonConversionServiceUnicodeParsingTest} (which covers the static {@code + * parseToUnicodeCodepoint} / {@code countCodesProtected} helpers) without duplicating those cases. + * + *

The service is constructed as a plain unit (no Spring context), so {@code @PostConstruct} + * never runs: font normalization stays disabled and Ghostscript is never invoked, keeping every + * test fully deterministic and free of external processes. {@link CustomPDFDocumentFactory#load} is + * stubbed to return real in-memory PDFBox documents, and {@code fallbackFontService} is mocked so + * the JSON to PDF path can resolve its mandatory fallback font without touching the classpath. + */ +@ExtendWith(MockitoExtension.class) +@org.mockito.junit.jupiter.MockitoSettings(strictness = Strictness.LENIENT) +class PdfJsonConversionServiceGapTest { + + @Mock private CustomPDFDocumentFactory pdfDocumentFactory; + @Mock private EndpointConfiguration endpointConfiguration; + @Mock private TempFileManager tempFileManager; + @Mock private TaskManager taskManager; + @Mock private PdfJsonFallbackFontService fallbackFontService; + @Mock private PdfJsonFontService fontService; + @Mock private Type3FontConversionService type3FontConversionService; + @Mock private Type3GlyphExtractor type3GlyphExtractor; + @Mock private ApplicationProperties applicationProperties; + + // Real collaborators: COS (de)serialization is complex and pure, so we use the real component. + private final PdfJsonCosMapper cosMapper = new PdfJsonCosMapper(); + + // Mirror production: application.properties sets + // spring.jackson.deserialization.fail-on-null-for-primitives=false, so the Spring-managed + // mapper + // maps null/absent JSON values onto Java primitive defaults (e.g. the boolean lazyImages + // field). + // A naive JsonMapper.builder().build() keeps the Jackson 3 default (true) and would throw + // MismatchedInputException on round-trip, which is a test-mapper config gap, not a product bug. + private final ObjectMapper objectMapper = + JsonMapper.builder() + .disable(DeserializationFeature.FAIL_ON_NULL_FOR_PRIMITIVES) + .build(); + + private PdfJsonConversionService service; + + private final List createdTempFiles = new ArrayList<>(); + + @BeforeEach + void setUp() throws IOException { + service = + new PdfJsonConversionService( + pdfDocumentFactory, + objectMapper, + endpointConfiguration, + tempFileManager, + taskManager, + cosMapper, + fallbackFontService, + fontService, + type3FontConversionService, + type3GlyphExtractor, + applicationProperties); + + // The TempFile wrapper delegates straight to the manager; back it with real temp files so + // convertPdfToJson can transferTo() and size/read the working path. + when(tempFileManager.createTempFile(anyString())) + .thenAnswer( + invocation -> { + String suffix = invocation.getArgument(0); + Path path = Files.createTempFile("pdfjson-gap-test", suffix); + createdTempFiles.add(path); + return path.toFile(); + }); + when(tempFileManager.deleteTempFile(any(File.class))) + .thenAnswer( + invocation -> { + File file = invocation.getArgument(0); + return file != null && file.delete(); + }); + } + + @AfterEach + void tearDown() throws IOException { + for (Path path : createdTempFiles) { + Files.deleteIfExists(path); + } + createdTempFiles.clear(); + } + + // ------------------------------------------------------------------ + // Helpers + // ------------------------------------------------------------------ + + /** Builds a tiny in-memory PDF with the requested page dimensions and rotations. */ + private PDDocument newPdf(float[][] sizes, int[] rotations) { + PDDocument document = new PDDocument(); + for (int i = 0; i < sizes.length; i++) { + PDPage page = new PDPage(new PDRectangle(sizes[i][0], sizes[i][1])); + if (rotations != null) { + page.setRotation(rotations[i]); + } + document.addPage(page); + } + return document; + } + + private PDDocument singlePagePdf(float width, float height) { + return newPdf(new float[][] {{width, height}}, null); + } + + /** Stubs the fallback font service so buildFontMap can always resolve a usable PDFont. */ + private void stubFallbackFont() throws IOException { + when(fallbackFontService.buildFallbackFontModel()) + .thenAnswer( + invocation -> + PdfJsonFont.builder() + .id(PdfJsonFallbackFontService.FALLBACK_FONT_ID) + .uid(PdfJsonFallbackFontService.FALLBACK_FONT_ID) + .baseName("Fallback") + .subtype("TrueType") + .build()); + when(fallbackFontService.loadFallbackPdfFont(any(PDDocument.class))) + .thenAnswer(invocation -> new PDType1Font(Standard14Fonts.FontName.HELVETICA)); + } + + private MockMultipartFile pdfMultipart() { + return new MockMultipartFile( + "fileInput", "input.pdf", "application/pdf", "%PDF-1.4 placeholder".getBytes()); + } + + private byte[] runJsonToPdf(PdfJsonDocument doc) throws IOException { + ByteArrayOutputStream out = new ByteArrayOutputStream(); + service.convertJsonToPdf(doc, out); + return out.toByteArray(); + } + + // ------------------------------------------------------------------ + // parseToUnicodeCodepoint - extra cases not covered by the unicode test + // ------------------------------------------------------------------ + + @Nested + @DisplayName("parseToUnicodeCodepoint extras") + class ParseToUnicodeExtras { + + @Test + @DisplayName("parses the maximum BMP value FFFF") + void parsesMaxBmpValue() { + assertEquals(0xFFFF, PdfJsonConversionService.parseToUnicodeCodepoint("FFFF")); + } + + @Test + @DisplayName("accepts lowercase hex digits") + void parsesLowercaseHex() { + // U+00E9 LATIN SMALL LETTER E WITH ACUTE. + assertEquals(0x00E9, PdfJsonConversionService.parseToUnicodeCodepoint("00e9")); + } + + @Test + @DisplayName("parses a 3-char value directly (length <= 4)") + void parsesThreeCharValue() { + assertEquals(0x1F4, PdfJsonConversionService.parseToUnicodeCodepoint("1f4")); + } + } + + // ------------------------------------------------------------------ + // convertJsonToPdf(PdfJsonDocument, OutputStream) + // ------------------------------------------------------------------ + + @Nested + @DisplayName("convertJsonToPdf(document)") + class JsonToPdfDocument { + + @Test + @DisplayName("null document throws IllegalArgumentException") + void nullDocumentThrows() { + ByteArrayOutputStream out = new ByteArrayOutputStream(); + assertThrows( + IllegalArgumentException.class, + () -> service.convertJsonToPdf((PdfJsonDocument) null, out)); + } + + @Test + @DisplayName("single empty page produces a valid one-page PDF with correct dimensions") + void singleEmptyPage() throws IOException { + stubFallbackFont(); + PdfJsonPage page = new PdfJsonPage(); + page.setPageNumber(1); + page.setWidth(300f); + page.setHeight(400f); + + PdfJsonDocument doc = new PdfJsonDocument(); + doc.setPages(List.of(page)); + + byte[] pdfBytes = runJsonToPdf(doc); + assertTrue(pdfBytes.length > 0, "expected non-empty PDF output"); + + try (PDDocument loaded = Loader.loadPDF(pdfBytes)) { + assertEquals(1, loaded.getNumberOfPages()); + PDRectangle box = loaded.getPage(0).getMediaBox(); + assertEquals(300f, box.getWidth(), 0.01f); + assertEquals(400f, box.getHeight(), 0.01f); + } + } + + @Test + @DisplayName("missing width/height falls back to US-Letter defaults") + void missingDimensionsUseDefaults() throws IOException { + stubFallbackFont(); + PdfJsonPage page = new PdfJsonPage(); + page.setPageNumber(1); + // width/height left null -> safeFloat defaults of 612x792 + + PdfJsonDocument doc = new PdfJsonDocument(); + doc.setPages(List.of(page)); + + try (PDDocument loaded = Loader.loadPDF(runJsonToPdf(doc))) { + PDRectangle box = loaded.getPage(0).getMediaBox(); + assertEquals(612f, box.getWidth(), 0.01f); + assertEquals(792f, box.getHeight(), 0.01f); + } + } + + @Test + @DisplayName("multiple pages keep their individual sizes and rotation") + void multiplePagesKeepSizesAndRotation() throws IOException { + stubFallbackFont(); + PdfJsonPage first = new PdfJsonPage(); + first.setPageNumber(1); + first.setWidth(200f); + first.setHeight(300f); + first.setRotation(90); + + PdfJsonPage second = new PdfJsonPage(); + second.setPageNumber(2); + second.setWidth(500f); + second.setHeight(250f); + second.setRotation(180); + + PdfJsonDocument doc = new PdfJsonDocument(); + doc.setPages(List.of(first, second)); + + try (PDDocument loaded = Loader.loadPDF(runJsonToPdf(doc))) { + assertEquals(2, loaded.getNumberOfPages()); + assertEquals(200f, loaded.getPage(0).getMediaBox().getWidth(), 0.01f); + assertEquals(90, loaded.getPage(0).getRotation()); + assertEquals(500f, loaded.getPage(1).getMediaBox().getWidth(), 0.01f); + assertEquals(180, loaded.getPage(1).getRotation()); + } + } + + @Test + @DisplayName("document metadata round-trips into PDDocumentInformation") + void metadataRoundTrips() throws IOException { + stubFallbackFont(); + PdfJsonMetadata metadata = new PdfJsonMetadata(); + metadata.setTitle("Gap Test Title"); + metadata.setAuthor("Gap Author"); + metadata.setSubject("Subject X"); + metadata.setKeywords("alpha,beta"); + metadata.setCreator("Creator App"); + metadata.setProducer("Producer Lib"); + metadata.setCreationDate("2020-06-15T12:00:00Z"); + + PdfJsonPage page = new PdfJsonPage(); + page.setPageNumber(1); + page.setWidth(300f); + page.setHeight(300f); + + PdfJsonDocument doc = new PdfJsonDocument(); + doc.setMetadata(metadata); + doc.setPages(List.of(page)); + + try (PDDocument loaded = Loader.loadPDF(runJsonToPdf(doc))) { + PDDocumentInformation info = loaded.getDocumentInformation(); + assertEquals("Gap Test Title", info.getTitle()); + assertEquals("Gap Author", info.getAuthor()); + assertEquals("Subject X", info.getSubject()); + assertEquals("alpha,beta", info.getKeywords()); + assertEquals("Creator App", info.getCreator()); + assertEquals("Producer Lib", info.getProducer()); + assertNotNull(info.getCreationDate(), "creation date should be applied"); + } + } + + @Test + @DisplayName("invalid creation date string is ignored without failing the conversion") + void invalidCreationDateIgnored() throws IOException { + stubFallbackFont(); + PdfJsonMetadata metadata = new PdfJsonMetadata(); + metadata.setTitle("Has bad date"); + metadata.setCreationDate("not-a-real-instant"); + + PdfJsonPage page = new PdfJsonPage(); + page.setPageNumber(1); + page.setWidth(300f); + page.setHeight(300f); + + PdfJsonDocument doc = new PdfJsonDocument(); + doc.setMetadata(metadata); + doc.setPages(List.of(page)); + + try (PDDocument loaded = Loader.loadPDF(runJsonToPdf(doc))) { + assertEquals("Has bad date", loaded.getDocumentInformation().getTitle()); + } + } + + @Test + @DisplayName("base64 XMP packet is restored onto the document catalog") + void xmpMetadataApplied() throws IOException { + stubFallbackFont(); + String xmpXml = + ""; + String base64 = + Base64.getEncoder().encodeToString(xmpXml.getBytes(StandardCharsets.UTF_8)); + + PdfJsonPage page = new PdfJsonPage(); + page.setPageNumber(1); + page.setWidth(300f); + page.setHeight(300f); + + PdfJsonDocument doc = new PdfJsonDocument(); + doc.setXmpMetadata(base64); + doc.setPages(List.of(page)); + + try (PDDocument loaded = Loader.loadPDF(runJsonToPdf(doc))) { + assertNotNull( + loaded.getDocumentCatalog().getMetadata(), + "XMP metadata stream should be present on the catalog"); + } + } + + @Test + @DisplayName("null fonts list is initialised instead of causing an NPE") + void nullFontsListHandled() throws IOException { + stubFallbackFont(); + PdfJsonPage page = new PdfJsonPage(); + page.setPageNumber(1); + page.setWidth(300f); + page.setHeight(300f); + + PdfJsonDocument doc = new PdfJsonDocument(); + doc.setFonts(null); + doc.setPages(List.of(page)); + + byte[] pdfBytes = runJsonToPdf(doc); + assertTrue(pdfBytes.length > 0); + // buildFontMap mutates the (now non-null) list by appending the fallback model. + assertNotNull(doc.getFonts()); + assertTrue( + doc.getFonts().stream() + .anyMatch( + f -> + PdfJsonFallbackFontService.FALLBACK_FONT_ID.equals( + f.getId()))); + } + } + + // ------------------------------------------------------------------ + // convertJsonToPdf(MultipartFile, OutputStream) + // ------------------------------------------------------------------ + + @Nested + @DisplayName("convertJsonToPdf(file)") + class JsonToPdfFile { + + @Test + @DisplayName("null file throws IllegalArgumentException") + void nullFileThrows() { + ByteArrayOutputStream out = new ByteArrayOutputStream(); + assertThrows( + IllegalArgumentException.class, + () -> service.convertJsonToPdf((MockMultipartFile) null, out)); + } + + @Test + @DisplayName("valid JSON payload is deserialized and rebuilt into a PDF") + void validJsonProducesPdf() throws IOException { + stubFallbackFont(); + PdfJsonPage page = new PdfJsonPage(); + page.setPageNumber(1); + page.setWidth(321f); + page.setHeight(123f); + PdfJsonDocument doc = new PdfJsonDocument(); + doc.setPages(List.of(page)); + + byte[] json = objectMapper.writeValueAsBytes(doc); + MockMultipartFile file = + new MockMultipartFile("fileInput", "doc.json", "application/json", json); + + ByteArrayOutputStream out = new ByteArrayOutputStream(); + service.convertJsonToPdf(file, out); + + try (PDDocument loaded = Loader.loadPDF(out.toByteArray())) { + assertEquals(1, loaded.getNumberOfPages()); + assertEquals(321f, loaded.getPage(0).getMediaBox().getWidth(), 0.01f); + } + } + } + + // ------------------------------------------------------------------ + // convertPdfToJson family + // ------------------------------------------------------------------ + + @Nested + @DisplayName("convertPdfToJson") + class PdfToJson { + + @Test + @DisplayName("null file throws IllegalArgumentException") + void nullFileThrows() { + ByteArrayOutputStream out = new ByteArrayOutputStream(); + assertThrows(IllegalArgumentException.class, () -> service.convertPdfToJson(null, out)); + } + + @Test + @DisplayName("blank single-page PDF yields JSON with one page and matching dimensions") + void blankSinglePage() throws IOException { + try (PDDocument pdf = singlePagePdf(250f, 350f)) { + when(pdfDocumentFactory.load(any(Path.class), eq(true))).thenReturn(pdf); + + ByteArrayOutputStream out = new ByteArrayOutputStream(); + service.convertPdfToJson(pdfMultipart(), out); + + PdfJsonDocument result = + objectMapper.readValue(out.toByteArray(), PdfJsonDocument.class); + assertEquals(1, result.getPages().size()); + PdfJsonPage page = result.getPages().get(0); + assertEquals(1, page.getPageNumber()); + assertEquals(250f, page.getWidth(), 0.01f); + assertEquals(350f, page.getHeight(), 0.01f); + assertNotNull(result.getMetadata()); + assertEquals(1, result.getMetadata().getNumberOfPages()); + } + } + + @Test + @DisplayName("multi-page PDF preserves per-page dimensions in the JSON model") + void multiPageDimensions() throws IOException { + try (PDDocument pdf = + newPdf(new float[][] {{200f, 300f}, {612f, 792f}}, new int[] {0, 90})) { + when(pdfDocumentFactory.load(any(Path.class), eq(true))).thenReturn(pdf); + + ByteArrayOutputStream out = new ByteArrayOutputStream(); + service.convertPdfToJson(pdfMultipart(), out); + + PdfJsonDocument result = + objectMapper.readValue(out.toByteArray(), PdfJsonDocument.class); + assertEquals(2, result.getPages().size()); + assertEquals(200f, result.getPages().get(0).getWidth(), 0.01f); + assertEquals(792f, result.getPages().get(1).getHeight(), 0.01f); + assertEquals(90, result.getPages().get(1).getRotation()); + } + } + + @Test + @DisplayName("source document metadata is extracted into the JSON metadata block") + void extractsSourceMetadata() throws IOException { + try (PDDocument pdf = singlePagePdf(300f, 300f)) { + PDDocumentInformation info = pdf.getDocumentInformation(); + info.setTitle("Original Title"); + info.setAuthor("Original Author"); + when(pdfDocumentFactory.load(any(Path.class), eq(true))).thenReturn(pdf); + + ByteArrayOutputStream out = new ByteArrayOutputStream(); + service.convertPdfToJson(pdfMultipart(), out); + + PdfJsonDocument result = + objectMapper.readValue(out.toByteArray(), PdfJsonDocument.class); + assertEquals("Original Title", result.getMetadata().getTitle()); + assertEquals("Original Author", result.getMetadata().getAuthor()); + } + } + + @Test + @DisplayName("lightweight overload still produces a parseable JSON document") + void lightweightOverload() throws IOException { + try (PDDocument pdf = singlePagePdf(300f, 300f)) { + when(pdfDocumentFactory.load(any(Path.class), eq(true))).thenReturn(pdf); + + ByteArrayOutputStream out = new ByteArrayOutputStream(); + service.convertPdfToJson(pdfMultipart(), true, out); + + PdfJsonDocument result = + objectMapper.readValue(out.toByteArray(), PdfJsonDocument.class); + assertEquals(1, result.getPages().size()); + } + } + + @Test + @DisplayName("progress callback receives a terminal complete event") + void progressCallbackInvoked() throws IOException { + try (PDDocument pdf = singlePagePdf(300f, 300f)) { + when(pdfDocumentFactory.load(any(Path.class), eq(true))).thenReturn(pdf); + + AtomicBoolean sawComplete = new AtomicBoolean(false); + AtomicBoolean sawAny = new AtomicBoolean(false); + + ByteArrayOutputStream out = new ByteArrayOutputStream(); + service.convertPdfToJson( + pdfMultipart(), + progress -> { + sawAny.set(true); + if (progress.getPercent() == 100 || progress.isComplete()) { + sawComplete.set(true); + } + }, + out); + + assertTrue(sawAny.get(), "expected at least one progress event"); + assertTrue(sawComplete.get(), "expected a terminal 100%/complete progress event"); + } + } + + @Test + @DisplayName("convertPdfToJsonDocument returns an in-memory model") + void convertPdfToJsonDocumentReturnsModel() throws IOException { + try (PDDocument pdf = singlePagePdf(400f, 500f)) { + when(pdfDocumentFactory.load(any(Path.class), eq(true))).thenReturn(pdf); + + PdfJsonDocument result = service.convertPdfToJsonDocument(pdfMultipart()); + + assertNotNull(result); + assertEquals(1, result.getPages().size()); + assertEquals(400f, result.getPages().get(0).getWidth(), 0.01f); + } + } + } + + // ------------------------------------------------------------------ + // Cache-backed public API error branches (no PDF load required) + // ------------------------------------------------------------------ + + @Nested + @DisplayName("cache-backed API") + class CacheApi { + + @Test + @DisplayName("extractSinglePage with unknown jobId throws CacheUnavailableException") + void extractSinglePageUnknownJob() { + ByteArrayOutputStream out = new ByteArrayOutputStream(); + assertThrows( + CacheUnavailableException.class, + () -> service.extractSinglePage("missing-job", 1, out)); + } + + @Test + @DisplayName("extractPageFonts with unknown jobId throws CacheUnavailableException") + void extractPageFontsUnknownJob() { + ByteArrayOutputStream out = new ByteArrayOutputStream(); + assertThrows( + CacheUnavailableException.class, + () -> service.extractPageFonts("missing-job", 1, out)); + } + + @Test + @DisplayName("exportUpdatedPages requires a non-null jobId") + void exportUpdatedPagesNullJob() { + ByteArrayOutputStream out = new ByteArrayOutputStream(); + assertThrows( + IllegalArgumentException.class, + () -> service.exportUpdatedPages(null, new PdfJsonDocument(), out)); + } + + @Test + @DisplayName("exportUpdatedPages rejects a blank jobId") + void exportUpdatedPagesBlankJob() { + ByteArrayOutputStream out = new ByteArrayOutputStream(); + assertThrows( + IllegalArgumentException.class, + () -> service.exportUpdatedPages(" ", new PdfJsonDocument(), out)); + } + + @Test + @DisplayName("exportUpdatedPages with unknown jobId throws CacheUnavailableException") + void exportUpdatedPagesUnknownJob() { + ByteArrayOutputStream out = new ByteArrayOutputStream(); + assertThrows( + CacheUnavailableException.class, + () -> service.exportUpdatedPages("missing-job", new PdfJsonDocument(), out)); + } + + @Test + @DisplayName("extractDocumentMetadata with null file throws IllegalArgumentException") + void extractDocumentMetadataNullFile() { + ByteArrayOutputStream out = new ByteArrayOutputStream(); + assertThrows( + IllegalArgumentException.class, + () -> service.extractDocumentMetadata(null, "job", out)); + } + + @Test + @DisplayName("clearCachedDocument on an unknown jobId is a no-op") + void clearCachedDocumentUnknownJob() { + assertDoesNotThrow(() -> service.clearCachedDocument("missing-job")); + } + } +} diff --git a/app/core/src/test/java/stirling/software/SPDF/service/PdfJsonCosMapperTest.java b/app/core/src/test/java/stirling/software/SPDF/service/PdfJsonCosMapperTest.java new file mode 100644 index 0000000000..a8172bf916 --- /dev/null +++ b/app/core/src/test/java/stirling/software/SPDF/service/PdfJsonCosMapperTest.java @@ -0,0 +1,780 @@ +package stirling.software.SPDF.service; + +import static org.junit.jupiter.api.Assertions.assertArrayEquals; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertInstanceOf; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.io.IOException; +import java.io.InputStream; +import java.io.OutputStream; +import java.nio.charset.StandardCharsets; +import java.util.Base64; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; + +import org.apache.pdfbox.cos.COSArray; +import org.apache.pdfbox.cos.COSBase; +import org.apache.pdfbox.cos.COSBoolean; +import org.apache.pdfbox.cos.COSDictionary; +import org.apache.pdfbox.cos.COSFloat; +import org.apache.pdfbox.cos.COSInteger; +import org.apache.pdfbox.cos.COSName; +import org.apache.pdfbox.cos.COSNull; +import org.apache.pdfbox.cos.COSStream; +import org.apache.pdfbox.cos.COSString; +import org.apache.pdfbox.pdmodel.PDDocument; +import org.apache.pdfbox.pdmodel.common.PDStream; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; + +import stirling.software.SPDF.model.json.PdfJsonCosValue; +import stirling.software.SPDF.model.json.PdfJsonStream; +import stirling.software.SPDF.service.PdfJsonCosMapper.SerializationContext; + +/** Unit tests for {@link PdfJsonCosMapper}. */ +class PdfJsonCosMapperTest { + + private PdfJsonCosMapper mapper; + private PDDocument document; + + @BeforeEach + void setUp() { + mapper = new PdfJsonCosMapper(); + document = new PDDocument(); + } + + @AfterEach + void tearDown() throws IOException { + if (document != null) { + document.close(); + } + } + + // Helper: create a populated COSStream within the test document. + private COSStream newCosStreamWithData(byte[] data) throws IOException { + COSStream cosStream = document.getDocument().createCOSStream(); + try (OutputStream out = cosStream.createRawOutputStream()) { + out.write(data); + } + cosStream.setItem(COSName.LENGTH, COSInteger.get(data.length)); + return cosStream; + } + + @Nested + @DisplayName("SerializationContext.omitStreamData()") + class OmitStreamDataTests { + + @Test + @DisplayName("returns true only for lightweight contexts") + void omitStreamData() { + assertFalse(SerializationContext.DEFAULT.omitStreamData()); + assertFalse(SerializationContext.ANNOTATION_RAW_DATA.omitStreamData()); + assertFalse(SerializationContext.FORM_FIELD_RAW_DATA.omitStreamData()); + assertTrue(SerializationContext.CONTENT_STREAMS_LIGHTWEIGHT.omitStreamData()); + assertTrue(SerializationContext.RESOURCES_LIGHTWEIGHT.omitStreamData()); + } + } + + @Nested + @DisplayName("serializeCosValue - primitives") + class SerializeCosValuePrimitiveTests { + + @Test + @DisplayName("null base yields null value") + void nullBase() throws IOException { + assertNull(mapper.serializeCosValue(null)); + } + + @Test + @DisplayName("COSNull serializes to NULL type") + void cosNull() throws IOException { + PdfJsonCosValue value = mapper.serializeCosValue(COSNull.NULL); + assertNotNull(value); + assertEquals(PdfJsonCosValue.Type.NULL, value.getType()); + } + + @Test + @DisplayName("COSBoolean serializes to BOOLEAN type with value") + void cosBoolean() throws IOException { + PdfJsonCosValue trueValue = mapper.serializeCosValue(COSBoolean.TRUE); + assertEquals(PdfJsonCosValue.Type.BOOLEAN, trueValue.getType()); + assertEquals(Boolean.TRUE, trueValue.getValue()); + + PdfJsonCosValue falseValue = mapper.serializeCosValue(COSBoolean.FALSE); + assertEquals(PdfJsonCosValue.Type.BOOLEAN, falseValue.getType()); + assertEquals(Boolean.FALSE, falseValue.getValue()); + } + + @Test + @DisplayName("COSInteger serializes to INTEGER type with long value") + void cosInteger() throws IOException { + PdfJsonCosValue value = mapper.serializeCosValue(COSInteger.get(42L)); + assertEquals(PdfJsonCosValue.Type.INTEGER, value.getType()); + assertEquals(42L, value.getValue()); + } + + @Test + @DisplayName("COSFloat serializes to FLOAT type with float value") + void cosFloat() throws IOException { + PdfJsonCosValue value = mapper.serializeCosValue(new COSFloat(1.5f)); + assertEquals(PdfJsonCosValue.Type.FLOAT, value.getType()); + assertEquals(1.5f, value.getValue()); + } + + @Test + @DisplayName("COSName serializes to NAME type with the name literal") + void cosName() throws IOException { + PdfJsonCosValue value = mapper.serializeCosValue(COSName.getPDFName("Foo")); + assertEquals(PdfJsonCosValue.Type.NAME, value.getType()); + assertEquals("Foo", value.getValue()); + } + + @Test + @DisplayName("COSString serializes to STRING type with base64 content") + void cosString() throws IOException { + byte[] raw = "héllo".getBytes(StandardCharsets.UTF_8); + PdfJsonCosValue value = mapper.serializeCosValue(new COSString(raw)); + assertEquals(PdfJsonCosValue.Type.STRING, value.getType()); + assertEquals(Base64.getEncoder().encodeToString(raw), value.getValue()); + } + } + + @Nested + @DisplayName("serializeCosValue - containers") + class SerializeCosValueContainerTests { + + @Test + @DisplayName("COSArray serializes nested items in order") + void cosArray() throws IOException { + COSArray array = new COSArray(); + array.add(COSInteger.get(1L)); + array.add(COSName.getPDFName("X")); + array.add(COSBoolean.TRUE); + + PdfJsonCosValue value = mapper.serializeCosValue(array); + assertEquals(PdfJsonCosValue.Type.ARRAY, value.getType()); + List items = value.getItems(); + assertEquals(3, items.size()); + assertEquals(PdfJsonCosValue.Type.INTEGER, items.get(0).getType()); + assertEquals(PdfJsonCosValue.Type.NAME, items.get(1).getType()); + assertEquals(PdfJsonCosValue.Type.BOOLEAN, items.get(2).getType()); + } + + @Test + @DisplayName("COSDictionary serializes keyed entries") + void cosDictionary() throws IOException { + COSDictionary dict = new COSDictionary(); + dict.setItem(COSName.getPDFName("Count"), COSInteger.get(7L)); + dict.setItem(COSName.getPDFName("Type"), COSName.getPDFName("Catalog")); + + PdfJsonCosValue value = mapper.serializeCosValue(dict); + assertEquals(PdfJsonCosValue.Type.DICTIONARY, value.getType()); + Map entries = value.getEntries(); + assertEquals(2, entries.size()); + assertEquals(PdfJsonCosValue.Type.INTEGER, entries.get("Count").getType()); + assertEquals(7L, entries.get("Count").getValue()); + assertEquals(PdfJsonCosValue.Type.NAME, entries.get("Type").getType()); + assertEquals("Catalog", entries.get("Type").getValue()); + } + + @Test + @DisplayName("circular dictionary reference is replaced with a __circular__ marker") + void circularReference() throws IOException { + COSDictionary dict = new COSDictionary(); + dict.setItem(COSName.getPDFName("Self"), dict); + + PdfJsonCosValue value = mapper.serializeCosValue(dict); + assertEquals(PdfJsonCosValue.Type.DICTIONARY, value.getType()); + PdfJsonCosValue self = value.getEntries().get("Self"); + assertEquals(PdfJsonCosValue.Type.NAME, self.getType()); + assertEquals("__circular__", self.getValue()); + } + + @Test + @DisplayName("the same dictionary appearing twice (non-circular) is serialized both times") + void repeatedNonCircularReference() throws IOException { + COSDictionary shared = new COSDictionary(); + shared.setItem(COSName.getPDFName("V"), COSInteger.get(9L)); + COSArray array = new COSArray(); + array.add(shared); + array.add(shared); + + PdfJsonCosValue value = mapper.serializeCosValue(array); + List items = value.getItems(); + assertEquals(2, items.size()); + // Because the visited set is removed in finally, the second sibling occurrence is not + // treated as circular. + assertEquals(PdfJsonCosValue.Type.DICTIONARY, items.get(0).getType()); + assertEquals(PdfJsonCosValue.Type.DICTIONARY, items.get(1).getType()); + } + } + + @Nested + @DisplayName("serializeStream overloads") + class SerializeStreamTests { + + @Test + @DisplayName("null PDStream returns null") + void nullPdStream() throws IOException { + assertNull(mapper.serializeStream((PDStream) null)); + } + + @Test + @DisplayName("null COSStream returns null") + void nullCosStream() throws IOException { + assertNull(mapper.serializeStream((COSStream) null)); + } + + @Test + @DisplayName("null COSStream with explicit context returns null") + void nullCosStreamWithContext() throws IOException { + assertNull(mapper.serializeStream((COSStream) null, SerializationContext.DEFAULT)); + } + + @Test + @DisplayName("null PDStream with explicit context returns null") + void nullPdStreamWithContext() throws IOException { + assertNull(mapper.serializeStream((PDStream) null, SerializationContext.DEFAULT)); + } + + @Test + @DisplayName("COSStream serializes dictionary and base64 rawData") + void cosStreamWithData() throws IOException { + byte[] data = "stream-bytes".getBytes(StandardCharsets.UTF_8); + COSStream cosStream = newCosStreamWithData(data); + cosStream.setItem(COSName.TYPE, COSName.getPDFName("XObject")); + + PdfJsonStream result = mapper.serializeStream(cosStream); + assertNotNull(result); + assertNotNull(result.getDictionary()); + assertTrue(result.getDictionary().containsKey(COSName.TYPE.getName())); + assertEquals(Base64.getEncoder().encodeToString(data), result.getRawData()); + } + + @Test + @DisplayName("empty COSStream yields null rawData") + void emptyCosStream() throws IOException { + COSStream cosStream = newCosStreamWithData(new byte[0]); + PdfJsonStream result = mapper.serializeStream(cosStream); + assertNotNull(result); + assertNull(result.getRawData()); + } + + @Test + @DisplayName("lightweight context omits rawData even when stream has data") + void lightweightContextOmitsData() throws IOException { + byte[] data = "ignored".getBytes(StandardCharsets.UTF_8); + COSStream cosStream = newCosStreamWithData(data); + cosStream.setItem(COSName.FILTER, COSName.getPDFName("FlateDecode")); + + PdfJsonStream result = + mapper.serializeStream( + cosStream, SerializationContext.CONTENT_STREAMS_LIGHTWEIGHT); + assertNotNull(result); + assertNull(result.getRawData()); + // Dictionary metadata is still preserved. + assertTrue(result.getDictionary().containsKey(COSName.FILTER.getName())); + } + + @Test + @DisplayName("null context is treated as DEFAULT and keeps rawData") + void nullContextDefaultsToDefault() throws IOException { + byte[] data = "keep".getBytes(StandardCharsets.UTF_8); + COSStream cosStream = newCosStreamWithData(data); + + PdfJsonStream result = mapper.serializeStream(cosStream, (SerializationContext) null); + assertNotNull(result); + assertEquals(Base64.getEncoder().encodeToString(data), result.getRawData()); + } + + @Test + @DisplayName("PDStream overload delegates to COSStream serialization") + void pdStreamOverload() throws IOException { + byte[] data = "pdstream".getBytes(StandardCharsets.UTF_8); + PDStream pdStream = new PDStream(document, new java.io.ByteArrayInputStream(data)); + + PdfJsonStream result = mapper.serializeStream(pdStream); + assertNotNull(result); + assertNotNull(result.getRawData()); + } + + @Test + @DisplayName("PDStream overload with lightweight context omits rawData") + void pdStreamOverloadLightweight() throws IOException { + byte[] data = "pdstream".getBytes(StandardCharsets.UTF_8); + PDStream pdStream = new PDStream(document, new java.io.ByteArrayInputStream(data)); + + PdfJsonStream result = + mapper.serializeStream(pdStream, SerializationContext.RESOURCES_LIGHTWEIGHT); + assertNotNull(result); + assertNull(result.getRawData()); + } + + @Test + @DisplayName("serializeCosValue of a COSStream produces STREAM type wrapping the stream") + void serializeCosValueWrapsStream() throws IOException { + byte[] data = "abc".getBytes(StandardCharsets.UTF_8); + COSStream cosStream = newCosStreamWithData(data); + + PdfJsonCosValue value = mapper.serializeCosValue(cosStream); + assertEquals(PdfJsonCosValue.Type.STREAM, value.getType()); + assertNotNull(value.getStream()); + assertEquals(Base64.getEncoder().encodeToString(data), value.getStream().getRawData()); + } + } + + @Nested + @DisplayName("deserializeCosValue") + class DeserializeCosValueTests { + + @Test + @DisplayName("null value returns null") + void nullValue() throws IOException { + assertNull(mapper.deserializeCosValue(null, document)); + } + + @Test + @DisplayName("value with null type returns null") + void nullType() throws IOException { + PdfJsonCosValue value = PdfJsonCosValue.builder().value("x").build(); + assertNull(mapper.deserializeCosValue(value, document)); + } + + @Test + @DisplayName("NULL type returns COSNull") + void nullTypeValue() throws IOException { + PdfJsonCosValue value = + PdfJsonCosValue.builder().type(PdfJsonCosValue.Type.NULL).build(); + assertEquals(COSNull.NULL, mapper.deserializeCosValue(value, document)); + } + + @Test + @DisplayName("BOOLEAN type returns matching COSBoolean") + void booleanType() throws IOException { + PdfJsonCosValue value = + PdfJsonCosValue.builder() + .type(PdfJsonCosValue.Type.BOOLEAN) + .value(Boolean.TRUE) + .build(); + COSBase result = mapper.deserializeCosValue(value, document); + assertEquals(COSBoolean.TRUE, result); + } + + @Test + @DisplayName("BOOLEAN type with non-boolean value returns null") + void booleanTypeWrongValue() throws IOException { + PdfJsonCosValue value = + PdfJsonCosValue.builder() + .type(PdfJsonCosValue.Type.BOOLEAN) + .value("not-a-boolean") + .build(); + assertNull(mapper.deserializeCosValue(value, document)); + } + + @Test + @DisplayName("INTEGER type returns COSInteger from a Number value") + void integerType() throws IOException { + PdfJsonCosValue value = + PdfJsonCosValue.builder() + .type(PdfJsonCosValue.Type.INTEGER) + .value(123L) + .build(); + COSBase result = mapper.deserializeCosValue(value, document); + assertInstanceOf(COSInteger.class, result); + assertEquals(123L, ((COSInteger) result).longValue()); + } + + @Test + @DisplayName("INTEGER type accepts an Integer value via Number") + void integerTypeFromInteger() throws IOException { + PdfJsonCosValue value = + PdfJsonCosValue.builder() + .type(PdfJsonCosValue.Type.INTEGER) + .value(Integer.valueOf(5)) + .build(); + COSBase result = mapper.deserializeCosValue(value, document); + assertInstanceOf(COSInteger.class, result); + assertEquals(5L, ((COSInteger) result).longValue()); + } + + @Test + @DisplayName("INTEGER type with non-number returns null") + void integerTypeWrongValue() throws IOException { + PdfJsonCosValue value = + PdfJsonCosValue.builder() + .type(PdfJsonCosValue.Type.INTEGER) + .value("oops") + .build(); + assertNull(mapper.deserializeCosValue(value, document)); + } + + @Test + @DisplayName("FLOAT type returns COSFloat from a Number value") + void floatType() throws IOException { + PdfJsonCosValue value = + PdfJsonCosValue.builder().type(PdfJsonCosValue.Type.FLOAT).value(2.25f).build(); + COSBase result = mapper.deserializeCosValue(value, document); + assertInstanceOf(COSFloat.class, result); + assertEquals(2.25f, ((COSFloat) result).floatValue()); + } + + @Test + @DisplayName("FLOAT type with non-number returns null") + void floatTypeWrongValue() throws IOException { + PdfJsonCosValue value = + PdfJsonCosValue.builder() + .type(PdfJsonCosValue.Type.FLOAT) + .value("oops") + .build(); + assertNull(mapper.deserializeCosValue(value, document)); + } + + @Test + @DisplayName("NAME type returns COSName from a String value") + void nameType() throws IOException { + PdfJsonCosValue value = + PdfJsonCosValue.builder() + .type(PdfJsonCosValue.Type.NAME) + .value("MyName") + .build(); + COSBase result = mapper.deserializeCosValue(value, document); + assertEquals(COSName.getPDFName("MyName"), result); + } + + @Test + @DisplayName("NAME type with non-string returns null") + void nameTypeWrongValue() throws IOException { + PdfJsonCosValue value = + PdfJsonCosValue.builder().type(PdfJsonCosValue.Type.NAME).value(123L).build(); + assertNull(mapper.deserializeCosValue(value, document)); + } + + @Test + @DisplayName("STRING type decodes base64 into COSString bytes") + void stringType() throws IOException { + byte[] raw = "round-trip".getBytes(StandardCharsets.UTF_8); + String encoded = Base64.getEncoder().encodeToString(raw); + PdfJsonCosValue value = + PdfJsonCosValue.builder() + .type(PdfJsonCosValue.Type.STRING) + .value(encoded) + .build(); + COSBase result = mapper.deserializeCosValue(value, document); + assertInstanceOf(COSString.class, result); + assertArrayEquals(raw, ((COSString) result).getBytes()); + } + + @Test + @DisplayName("STRING type with invalid base64 returns null") + void stringTypeInvalidBase64() throws IOException { + PdfJsonCosValue value = + PdfJsonCosValue.builder() + .type(PdfJsonCosValue.Type.STRING) + .value("!!!not base64!!!") + .build(); + assertNull(mapper.deserializeCosValue(value, document)); + } + + @Test + @DisplayName("STRING type with non-string value returns null") + void stringTypeWrongValue() throws IOException { + PdfJsonCosValue value = + PdfJsonCosValue.builder().type(PdfJsonCosValue.Type.STRING).value(42L).build(); + assertNull(mapper.deserializeCosValue(value, document)); + } + + @Test + @DisplayName("ARRAY type deserializes each item") + void arrayType() throws IOException { + PdfJsonCosValue item1 = + PdfJsonCosValue.builder().type(PdfJsonCosValue.Type.INTEGER).value(1L).build(); + PdfJsonCosValue item2 = + PdfJsonCosValue.builder().type(PdfJsonCosValue.Type.NAME).value("N").build(); + PdfJsonCosValue value = + PdfJsonCosValue.builder() + .type(PdfJsonCosValue.Type.ARRAY) + .items(List.of(item1, item2)) + .build(); + + COSBase result = mapper.deserializeCosValue(value, document); + assertInstanceOf(COSArray.class, result); + COSArray array = (COSArray) result; + assertEquals(2, array.size()); + assertEquals(1L, ((COSInteger) array.get(0)).longValue()); + assertEquals(COSName.getPDFName("N"), array.get(1)); + } + + @Test + @DisplayName("ARRAY type substitutes COSNull for un-deserializable items") + void arrayTypeWithNullItems() throws IOException { + // An INTEGER type with a non-number value deserializes to null and is replaced by + // COSNull.NULL inside the array. + PdfJsonCosValue bad = + PdfJsonCosValue.builder() + .type(PdfJsonCosValue.Type.INTEGER) + .value("nope") + .build(); + PdfJsonCosValue value = + PdfJsonCosValue.builder() + .type(PdfJsonCosValue.Type.ARRAY) + .items(List.of(bad)) + .build(); + + COSArray result = (COSArray) mapper.deserializeCosValue(value, document); + assertEquals(1, result.size()); + assertEquals(COSNull.NULL, result.get(0)); + } + + @Test + @DisplayName("ARRAY type with null items list yields empty COSArray") + void arrayTypeNullItems() throws IOException { + PdfJsonCosValue value = + PdfJsonCosValue.builder().type(PdfJsonCosValue.Type.ARRAY).build(); + COSArray result = (COSArray) mapper.deserializeCosValue(value, document); + assertNotNull(result); + assertEquals(0, result.size()); + } + + @Test + @DisplayName("DICTIONARY type deserializes entries by key") + void dictionaryType() throws IOException { + Map entries = new LinkedHashMap<>(); + entries.put( + "Count", + PdfJsonCosValue.builder().type(PdfJsonCosValue.Type.INTEGER).value(3L).build()); + PdfJsonCosValue value = + PdfJsonCosValue.builder() + .type(PdfJsonCosValue.Type.DICTIONARY) + .entries(entries) + .build(); + + COSDictionary result = (COSDictionary) mapper.deserializeCosValue(value, document); + assertNotNull(result); + assertEquals(3L, ((COSInteger) result.getItem("Count")).longValue()); + } + + @Test + @DisplayName("DICTIONARY type skips entries that deserialize to null") + void dictionaryTypeSkipsNullEntries() throws IOException { + Map entries = new LinkedHashMap<>(); + entries.put( + "Bad", + PdfJsonCosValue.builder() + .type(PdfJsonCosValue.Type.INTEGER) + .value("nope") + .build()); + PdfJsonCosValue value = + PdfJsonCosValue.builder() + .type(PdfJsonCosValue.Type.DICTIONARY) + .entries(entries) + .build(); + + COSDictionary result = (COSDictionary) mapper.deserializeCosValue(value, document); + assertNotNull(result); + assertNull(result.getItem(COSName.getPDFName("Bad"))); + } + + @Test + @DisplayName("DICTIONARY type with null entries map yields empty COSDictionary") + void dictionaryTypeNullEntries() throws IOException { + PdfJsonCosValue value = + PdfJsonCosValue.builder().type(PdfJsonCosValue.Type.DICTIONARY).build(); + COSDictionary result = (COSDictionary) mapper.deserializeCosValue(value, document); + assertNotNull(result); + assertEquals(0, result.size()); + } + + @Test + @DisplayName("STREAM type with null stream returns null") + void streamTypeNullStream() throws IOException { + PdfJsonCosValue value = + PdfJsonCosValue.builder().type(PdfJsonCosValue.Type.STREAM).build(); + assertNull(mapper.deserializeCosValue(value, document)); + } + + @Test + @DisplayName("STREAM type builds a COSStream from the model") + void streamType() throws IOException { + byte[] data = "stream-content".getBytes(StandardCharsets.UTF_8); + PdfJsonStream streamModel = + PdfJsonStream.builder() + .rawData(Base64.getEncoder().encodeToString(data)) + .build(); + PdfJsonCosValue value = + PdfJsonCosValue.builder().type(PdfJsonCosValue.Type.STREAM).stream(streamModel) + .build(); + + COSBase result = mapper.deserializeCosValue(value, document); + assertInstanceOf(COSStream.class, result); + assertStreamRawEquals(data, (COSStream) result); + } + } + + @Nested + @DisplayName("buildStreamFromModel") + class BuildStreamFromModelTests { + + @Test + @DisplayName("null model returns null") + void nullModel() throws IOException { + assertNull(mapper.buildStreamFromModel(null, document)); + } + + @Test + @DisplayName("model with rawData writes base64-decoded bytes and sets Length") + void withRawData() throws IOException { + byte[] data = "hello-world".getBytes(StandardCharsets.UTF_8); + PdfJsonStream model = + PdfJsonStream.builder() + .rawData(Base64.getEncoder().encodeToString(data)) + .build(); + + COSStream result = mapper.buildStreamFromModel(model, document); + assertNotNull(result); + assertStreamRawEquals(data, result); + assertEquals(data.length, ((COSInteger) result.getItem(COSName.LENGTH)).longValue()); + } + + @Test + @DisplayName("model with dictionary entries copies them onto the stream") + void withDictionary() throws IOException { + Map dict = new LinkedHashMap<>(); + dict.put( + "Type", + PdfJsonCosValue.builder() + .type(PdfJsonCosValue.Type.NAME) + .value("XObject") + .build()); + dict.put( + "Width", + PdfJsonCosValue.builder() + .type(PdfJsonCosValue.Type.INTEGER) + .value(100L) + .build()); + PdfJsonStream model = PdfJsonStream.builder().dictionary(dict).build(); + + COSStream result = mapper.buildStreamFromModel(model, document); + assertNotNull(result); + assertEquals(COSName.getPDFName("XObject"), result.getItem(COSName.TYPE)); + assertEquals( + 100L, ((COSInteger) result.getItem(COSName.getPDFName("Width"))).longValue()); + } + + @Test + @DisplayName("model dictionary entry that deserializes to null is skipped") + void dictionaryEntryNullSkipped() throws IOException { + Map dict = new LinkedHashMap<>(); + dict.put( + "Bad", + PdfJsonCosValue.builder() + .type(PdfJsonCosValue.Type.NAME) + .value(999L) // non-string -> null + .build()); + PdfJsonStream model = PdfJsonStream.builder().dictionary(dict).build(); + + COSStream result = mapper.buildStreamFromModel(model, document); + assertNotNull(result); + assertNull(result.getItem(COSName.getPDFName("Bad"))); + } + + @Test + @DisplayName("model with null/blank rawData sets Length to zero and writes nothing") + void blankRawData() throws IOException { + PdfJsonStream model = PdfJsonStream.builder().rawData(" ").build(); + COSStream result = mapper.buildStreamFromModel(model, document); + assertNotNull(result); + assertEquals(0L, ((COSInteger) result.getItem(COSName.LENGTH)).longValue()); + // Blank rawData takes the else-branch: Length is set to 0 but createRawOutputStream() + // is + // never called, so nothing is written. PDFBox cannot open a raw InputStream on a stream + // that was never written to, which is the documented "writes nothing" behaviour. + assertThrows(IOException.class, result::createRawInputStream); + } + + @Test + @DisplayName("model with no rawData sets Length to zero") + void noRawData() throws IOException { + PdfJsonStream model = PdfJsonStream.builder().build(); + COSStream result = mapper.buildStreamFromModel(model, document); + assertNotNull(result); + assertEquals(0L, ((COSInteger) result.getItem(COSName.LENGTH)).longValue()); + } + + @Test + @DisplayName("model with invalid base64 rawData falls back to empty data") + void invalidBase64RawData() throws IOException { + PdfJsonStream model = PdfJsonStream.builder().rawData("###not-base64###").build(); + COSStream result = mapper.buildStreamFromModel(model, document); + assertNotNull(result); + assertEquals(0L, ((COSInteger) result.getItem(COSName.LENGTH)).longValue()); + assertStreamRawEquals(new byte[0], result); + } + } + + @Nested + @DisplayName("round trip serialize -> deserialize") + class RoundTripTests { + + @Test + @DisplayName("nested array/dictionary structure survives a round trip") + void nestedStructure() throws IOException { + COSDictionary original = new COSDictionary(); + original.setItem(COSName.getPDFName("Int"), COSInteger.get(11L)); + original.setItem(COSName.getPDFName("Name"), COSName.getPDFName("Hello")); + original.setItem(COSName.getPDFName("Bool"), COSBoolean.FALSE); + COSArray inner = new COSArray(); + inner.add(new COSFloat(3.5f)); + inner.add(new COSString("abc".getBytes(StandardCharsets.UTF_8))); + original.setItem(COSName.getPDFName("Arr"), inner); + + PdfJsonCosValue serialized = mapper.serializeCosValue(original); + COSBase deserialized = mapper.deserializeCosValue(serialized, document); + + assertInstanceOf(COSDictionary.class, deserialized); + COSDictionary result = (COSDictionary) deserialized; + assertEquals(11L, ((COSInteger) result.getItem(COSName.getPDFName("Int"))).longValue()); + assertEquals(COSName.getPDFName("Hello"), result.getItem(COSName.getPDFName("Name"))); + assertEquals(COSBoolean.FALSE, result.getItem(COSName.getPDFName("Bool"))); + + COSArray resultArr = (COSArray) result.getItem(COSName.getPDFName("Arr")); + assertEquals(2, resultArr.size()); + assertEquals(3.5f, ((COSFloat) resultArr.get(0)).floatValue()); + assertArrayEquals( + "abc".getBytes(StandardCharsets.UTF_8), + ((COSString) resultArr.get(1)).getBytes()); + } + + @Test + @DisplayName("stream data survives a round trip") + void streamRoundTrip() throws IOException { + byte[] data = "round-trip-stream".getBytes(StandardCharsets.UTF_8); + COSStream cosStream = newCosStreamWithData(data); + cosStream.setItem(COSName.TYPE, COSName.getPDFName("XObject")); + + PdfJsonCosValue serialized = mapper.serializeCosValue(cosStream); + COSBase deserialized = mapper.deserializeCosValue(serialized, document); + + assertInstanceOf(COSStream.class, deserialized); + COSStream result = (COSStream) deserialized; + assertStreamRawEquals(data, result); + assertEquals(COSName.getPDFName("XObject"), result.getItem(COSName.TYPE)); + } + } + + // Reads the raw (undecoded) bytes from a COSStream and asserts equality. + private void assertStreamRawEquals(byte[] expected, COSStream cosStream) throws IOException { + try (InputStream in = cosStream.createRawInputStream()) { + byte[] actual = in.readAllBytes(); + assertArrayEquals(expected, actual); + } + } +} diff --git a/app/core/src/test/java/stirling/software/SPDF/service/PdfJsonFallbackFontServiceTest.java b/app/core/src/test/java/stirling/software/SPDF/service/PdfJsonFallbackFontServiceTest.java new file mode 100644 index 0000000000..c39f79ce2c --- /dev/null +++ b/app/core/src/test/java/stirling/software/SPDF/service/PdfJsonFallbackFontServiceTest.java @@ -0,0 +1,571 @@ +package stirling.software.SPDF.service; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertSame; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.Mockito.lenient; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; + +import java.io.IOException; +import java.lang.reflect.Method; +import java.util.Base64; + +import org.apache.pdfbox.pdmodel.PDDocument; +import org.apache.pdfbox.pdmodel.font.PDFont; +import org.apache.pdfbox.pdmodel.font.PDType0Font; +import org.apache.pdfbox.pdmodel.font.PDType3Font; +import org.junit.jupiter.api.AfterAll; +import org.junit.jupiter.api.BeforeAll; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.springframework.core.io.DefaultResourceLoader; +import org.springframework.core.io.ResourceLoader; +import org.springframework.test.util.ReflectionTestUtils; + +import stirling.software.SPDF.model.json.PdfJsonFont; +import stirling.software.common.model.ApplicationProperties; + +class PdfJsonFallbackFontServiceTest { + + private PdfJsonFallbackFontService service; + private ApplicationProperties applicationProperties; + + // Parsing the real fallback TrueType fonts (the CJK font alone is ~17 MB) is the only expensive + // work in this class. Build a default-config service and parse the CJK fallback once, then + // reuse + // them across every read-only test instead of re-parsing per @BeforeEach. + private static PdfJsonFallbackFontService sharedService; + private static PDDocument sharedCjkDocument; + private static PDFont sharedCjkFont; + + @BeforeAll + static void setUpShared() throws Exception { + // Real ApplicationProperties already defaults pdfEditor.fallbackFont to the Noto Sans + // location, so no mocking is needed for the default-config path. + ResourceLoader resourceLoader = new DefaultResourceLoader(); + ApplicationProperties props = new ApplicationProperties(); + sharedService = new PdfJsonFallbackFontService(resourceLoader, props); + ReflectionTestUtils.setField( + sharedService, + "legacyFallbackFontLocation", + PdfJsonFallbackFontService.DEFAULT_FALLBACK_FONT_LOCATION); + Method loadConfig = PdfJsonFallbackFontService.class.getDeclaredMethod("loadConfig"); + loadConfig.setAccessible(true); + loadConfig.invoke(sharedService); + + // Parse the large CJK fallback exactly once; the CanEncode tests only read this font. + sharedCjkDocument = new PDDocument(); + sharedCjkFont = + sharedService.loadFallbackPdfFont( + sharedCjkDocument, PdfJsonFallbackFontService.FALLBACK_FONT_CJK_ID); + } + + @AfterAll + static void tearDownShared() throws IOException { + if (sharedCjkDocument != null) { + sharedCjkDocument.close(); + } + } + + @BeforeEach + void setUp() { + // Real resource loader resolves the classpath:/static/fonts/*.ttf bundled in the module. + ResourceLoader resourceLoader = new DefaultResourceLoader(); + + // Mock the properties chain so loadConfig() resolves to the default Noto Sans location. + applicationProperties = mock(ApplicationProperties.class); + ApplicationProperties.PdfEditor pdfEditor = mock(ApplicationProperties.PdfEditor.class); + lenient().when(applicationProperties.getPdfEditor()).thenReturn(pdfEditor); + lenient() + .when(pdfEditor.getFallbackFont()) + .thenReturn(PdfJsonFallbackFontService.DEFAULT_FALLBACK_FONT_LOCATION); + + service = new PdfJsonFallbackFontService(resourceLoader, applicationProperties); + // The @Value field is normally injected by Spring; set it explicitly for the unit test. + ReflectionTestUtils.setField( + service, + "legacyFallbackFontLocation", + PdfJsonFallbackFontService.DEFAULT_FALLBACK_FONT_LOCATION); + } + + /** Invoke the private @PostConstruct loadConfig() to populate fallbackFontLocation. */ + private void invokeLoadConfig() throws Exception { + Method loadConfig = PdfJsonFallbackFontService.class.getDeclaredMethod("loadConfig"); + loadConfig.setAccessible(true); + loadConfig.invoke(service); + } + + @Nested + @DisplayName("loadConfig (@PostConstruct)") + class LoadConfig { + + @Test + @DisplayName("uses the configured pdf-editor fallback font when set") + void usesConfiguredFallbackFont() throws Exception { + when(applicationProperties.getPdfEditor().getFallbackFont()) + .thenReturn("classpath:/static/fonts/DejaVuSans.ttf"); + + invokeLoadConfig(); + + assertEquals( + "classpath:/static/fonts/DejaVuSans.ttf", + ReflectionTestUtils.getField(service, "fallbackFontLocation")); + } + + @Test + @DisplayName("falls back to the legacy @Value location when configured value is blank") + void fallsBackToLegacyWhenBlank() throws Exception { + when(applicationProperties.getPdfEditor().getFallbackFont()).thenReturn(" "); + + invokeLoadConfig(); + + assertEquals( + PdfJsonFallbackFontService.DEFAULT_FALLBACK_FONT_LOCATION, + ReflectionTestUtils.getField(service, "fallbackFontLocation")); + } + + @Test + @DisplayName("falls back to the legacy @Value location when PdfEditor is null") + void fallsBackToLegacyWhenPdfEditorNull() throws Exception { + when(applicationProperties.getPdfEditor()).thenReturn(null); + + invokeLoadConfig(); + + assertEquals( + PdfJsonFallbackFontService.DEFAULT_FALLBACK_FONT_LOCATION, + ReflectionTestUtils.getField(service, "fallbackFontLocation")); + } + } + + @Nested + @DisplayName("resolveFallbackFontId(int codePoint)") + class ResolveByCodePoint { + + @Test + @DisplayName("Latin letters resolve to the generic Noto Sans fallback") + void latinResolvesToNotoSans() { + assertEquals( + PdfJsonFallbackFontService.FALLBACK_FONT_ID, + service.resolveFallbackFontId('A')); + } + + @Test + @DisplayName("CJK unified ideographs resolve to the CJK fallback") + void cjkIdeographResolvesToCjk() { + // U+4E2D (中) is a CJK Unified Ideograph. + assertEquals( + PdfJsonFallbackFontService.FALLBACK_FONT_CJK_ID, + service.resolveFallbackFontId(0x4E2D)); + } + + @Test + @DisplayName("Bopomofo resolves to the Traditional Chinese fallback") + void bopomofoResolvesToTc() { + // U+3105 (ㄅ) is in the Bopomofo block. + assertEquals( + PdfJsonFallbackFontService.FALLBACK_FONT_TC_ID, + service.resolveFallbackFontId(0x3105)); + } + + @Test + @DisplayName("CJK compatibility ideographs resolve to the Traditional Chinese fallback") + void compatibilityIdeographResolvesToTc() { + // U+F900 is in the CJK Compatibility Ideographs block. + assertEquals( + PdfJsonFallbackFontService.FALLBACK_FONT_TC_ID, + service.resolveFallbackFontId(0xF900)); + } + + @Test + @DisplayName("Hiragana resolves to the Japanese fallback") + void hiraganaResolvesToJp() { + // U+3042 (あ) is Hiragana. + assertEquals( + PdfJsonFallbackFontService.FALLBACK_FONT_JP_ID, + service.resolveFallbackFontId(0x3042)); + } + + @Test + @DisplayName("Hangul resolves to the Korean fallback") + void hangulResolvesToKr() { + // U+AC00 (가) is a Hangul syllable. + assertEquals( + PdfJsonFallbackFontService.FALLBACK_FONT_KR_ID, + service.resolveFallbackFontId(0xAC00)); + } + + @Test + @DisplayName("Arabic resolves to the Arabic fallback") + void arabicResolvesToAr() { + // U+0627 (ا) is Arabic. + assertEquals( + PdfJsonFallbackFontService.FALLBACK_FONT_AR_ID, + service.resolveFallbackFontId(0x0627)); + } + + @Test + @DisplayName("Thai resolves to the Thai fallback") + void thaiResolvesToTh() { + // U+0E01 (ก) is Thai. + assertEquals( + PdfJsonFallbackFontService.FALLBACK_FONT_TH_ID, + service.resolveFallbackFontId(0x0E01)); + } + + @Test + @DisplayName("Devanagari resolves to the Devanagari fallback") + void devanagariResolves() { + // U+0905 (अ) is Devanagari. + assertEquals( + PdfJsonFallbackFontService.FALLBACK_FONT_DEVANAGARI_ID, + service.resolveFallbackFontId(0x0905)); + } + + @Test + @DisplayName("Malayalam resolves to the Malayalam fallback") + void malayalamResolves() { + // U+0D05 (അ) is Malayalam. + assertEquals( + PdfJsonFallbackFontService.FALLBACK_FONT_MALAYALAM_ID, + service.resolveFallbackFontId(0x0D05)); + } + + @Test + @DisplayName("Tibetan resolves to the Tibetan fallback") + void tibetanResolves() { + // U+0F40 (ཀ) is Tibetan. + assertEquals( + PdfJsonFallbackFontService.FALLBACK_FONT_TIBETAN_ID, + service.resolveFallbackFontId(0x0F40)); + } + } + + @Nested + @DisplayName("resolveFallbackFontId(String fontName, int codePoint)") + class ResolveByNameAndCodePoint { + + @Test + @DisplayName("Arial maps to Liberation Sans") + void arialMapsToLiberationSans() { + assertEquals("fallback-liberation-sans", service.resolveFallbackFontId("Arial", 'A')); + } + + @Test + @DisplayName("Times New Roman maps to Liberation Serif") + void timesNewRomanMapsToLiberationSerif() { + // Spaces are stripped: "Times New Roman" -> "timesnewroman". + assertEquals( + "fallback-liberation-serif", + service.resolveFallbackFontId("Times New Roman", 'A')); + } + + @Test + @DisplayName("Courier New maps to Liberation Mono") + void courierNewMapsToLiberationMono() { + assertEquals( + "fallback-liberation-mono", service.resolveFallbackFontId("Courier New", 'A')); + } + + @Test + @DisplayName("Arial-Bold maps to bold Liberation Sans variant") + void arialBoldMapsToBoldVariant() { + assertEquals( + "fallback-liberation-sans-bold", + service.resolveFallbackFontId("Arial-Bold", 'A')); + } + + @Test + @DisplayName("Arial-Italic maps to italic Liberation Sans variant") + void arialItalicMapsToItalicVariant() { + assertEquals( + "fallback-liberation-sans-italic", + service.resolveFallbackFontId("Arial-Italic", 'A')); + } + + @Test + @DisplayName("Arial-BoldItalic maps to bold-italic Liberation Sans variant") + void arialBoldItalicMapsToBoldItalicVariant() { + assertEquals( + "fallback-liberation-sans-bolditalic", + service.resolveFallbackFontId("Arial-BoldItalic", 'A')); + } + + @Test + @DisplayName("numeric weight 700 is detected as bold") + void numericWeightDetectedAsBold() { + // "Arimo_700wght" -> base "arimo" -> liberation-sans, bold via 700 weight pattern. + assertEquals( + "fallback-liberation-sans-bold", + service.resolveFallbackFontId("Arimo_700wght", 'A')); + } + + @Test + @DisplayName("subset prefix is stripped before alias matching") + void subsetPrefixStripped() { + // "ABCDEF+Arial" -> subset prefix removed -> "arial". + assertEquals( + "fallback-liberation-sans", service.resolveFallbackFontId("ABCDEF+Arial", 'A')); + } + + @Test + @DisplayName("DejaVu Sans bold-italic uses the 'oblique' naming convention") + void dejaVuUsesObliqueNaming() { + assertEquals( + "fallback-dejavu-sans-boldoblique", + service.resolveFallbackFontId("DejaVuSans-BoldItalic", 'A')); + } + + @Test + @DisplayName("DejaVu Serif italic keeps the 'italic' naming convention") + void dejaVuSerifKeepsItalicNaming() { + assertEquals( + "fallback-dejavu-serif-italic", + service.resolveFallbackFontId("DejaVuSerif-Italic", 'A')); + } + + @Test + @DisplayName("Traditional Chinese aliased name ignores weight/style suffix (unsupported)") + void tcAliasIgnoresWeightStyle() { + // MingLiU maps to fallback-noto-tc, which has no bold variant registered. + assertEquals( + PdfJsonFallbackFontService.FALLBACK_FONT_TC_ID, + service.resolveFallbackFontId("MingLiU-Bold", 'A')); + } + + @Test + @DisplayName("Simplified Chinese aliased name maps to the CJK fallback") + void simsunMapsToCjk() { + assertEquals( + PdfJsonFallbackFontService.FALLBACK_FONT_CJK_ID, + service.resolveFallbackFontId("SimSun", 'A')); + } + + @Test + @DisplayName("null font name falls through to Unicode-based resolution") + void nullNameFallsThroughToUnicode() { + assertEquals( + PdfJsonFallbackFontService.FALLBACK_FONT_ID, + service.resolveFallbackFontId(null, 'A')); + } + + @Test + @DisplayName("empty font name falls through to Unicode-based resolution") + void emptyNameFallsThroughToUnicode() { + assertEquals( + PdfJsonFallbackFontService.FALLBACK_FONT_CJK_ID, + service.resolveFallbackFontId("", 0x4E2D)); + } + + @Test + @DisplayName("unknown font name falls through to Unicode-based resolution") + void unknownNameFallsThroughToUnicode() { + assertEquals( + PdfJsonFallbackFontService.FALLBACK_FONT_JP_ID, + service.resolveFallbackFontId("SomeCustomFont", 0x3042)); + } + } + + @Nested + @DisplayName("mapUnsupportedGlyph(int codePoint)") + class MapUnsupportedGlyph { + + @Test + @DisplayName("U+276E maps to '<'") + void heavyLeftAngleMapsToLessThan() { + assertEquals("<", service.mapUnsupportedGlyph(0x276E)); + } + + @Test + @DisplayName("U+276F maps to '>'") + void heavyRightAngleMapsToGreaterThan() { + assertEquals(">", service.mapUnsupportedGlyph(0x276F)); + } + + @Test + @DisplayName("unmapped code point returns null") + void unmappedReturnsNull() { + assertNull(service.mapUnsupportedGlyph('A')); + } + } + + @Nested + @DisplayName("canEncode / canEncodeFully") + class CanEncode { + + @Test + @DisplayName("null font returns false") + void nullFontReturnsFalse() { + assertFalse(service.canEncode((PDFont) null, "A")); + } + + @Test + @DisplayName("null text returns false") + void nullTextReturnsFalse() { + // Reuse the once-parsed CJK fallback; canEncode only reads the font. + assertFalse(sharedService.canEncode(sharedCjkFont, (String) null)); + } + + @Test + @DisplayName("empty text returns false") + void emptyTextReturnsFalse() { + assertFalse(sharedService.canEncode(sharedCjkFont, "")); + } + + @Test + @DisplayName("PDType3Font always returns false") + void type3FontReturnsFalse() { + PDType3Font type3 = mock(PDType3Font.class); + assertFalse(service.canEncode(type3, "A")); + } + + @Test + @DisplayName("loaded TrueType fallback can encode basic Latin") + void loadedFontEncodesLatin() { + assertTrue(sharedService.canEncode(sharedCjkFont, "Hello")); + } + + @Test + @DisplayName("canEncodeFully delegates to canEncode for text") + void canEncodeFullyDelegates() { + assertTrue(sharedService.canEncodeFully(sharedCjkFont, "abc")); + assertFalse(sharedService.canEncodeFully(null, "abc")); + } + + @Test + @DisplayName("canEncode(font, codePoint) returns true for an encodable code point") + void canEncodeCodePointTrue() { + assertTrue(sharedService.canEncode(sharedCjkFont, (int) 'A')); + } + + @Test + @DisplayName("canEncode(font, codePoint) returns false for a null font") + void canEncodeCodePointNullFont() { + assertFalse(service.canEncode((PDFont) null, (int) 'A')); + } + } + + @Nested + @DisplayName("buildFallbackFontModel") + class BuildFallbackFontModel { + + @Test + @DisplayName("builds a model for a built-in CJK font with base64 program bytes") + void buildsCjkModel() throws IOException { + PdfJsonFont model = + service.buildFallbackFontModel(PdfJsonFallbackFontService.FALLBACK_FONT_CJK_ID); + + assertNotNull(model); + assertEquals(PdfJsonFallbackFontService.FALLBACK_FONT_CJK_ID, model.getId()); + assertEquals(PdfJsonFallbackFontService.FALLBACK_FONT_CJK_ID, model.getUid()); + assertEquals("NotoSansSC-Regular", model.getBaseName()); + assertEquals("TrueType", model.getSubtype()); + assertEquals(Boolean.TRUE, model.getEmbedded()); + assertEquals("ttf", model.getProgramFormat()); + assertNotNull(model.getProgram()); + // The program must be valid, non-empty base64. + byte[] decoded = Base64.getDecoder().decode(model.getProgram()); + assertTrue(decoded.length > 0); + } + + @Test + @DisplayName("no-arg overload builds the default Noto Sans model") + void noArgBuildsDefaultModel() throws Exception { + invokeLoadConfig(); // populate fallbackFontLocation for the default font id + PdfJsonFont model = service.buildFallbackFontModel(); + + assertNotNull(model); + assertEquals(PdfJsonFallbackFontService.FALLBACK_FONT_ID, model.getId()); + assertEquals("NotoSans-Regular", model.getBaseName()); + assertEquals("ttf", model.getProgramFormat()); + assertNotNull(model.getProgram()); + } + + @Test + @DisplayName("unknown fallback id throws IOException") + void unknownIdThrows() { + IOException ex = + assertThrows( + IOException.class, + () -> service.buildFallbackFontModel("does-not-exist")); + assertTrue(ex.getMessage().contains("Unknown fallback font id")); + } + + @Test + @DisplayName("font bytes are cached and reused across calls") + void fontBytesAreCached() throws IOException { + PdfJsonFont first = + service.buildFallbackFontModel(PdfJsonFallbackFontService.FALLBACK_FONT_TH_ID); + PdfJsonFont second = + service.buildFallbackFontModel(PdfJsonFallbackFontService.FALLBACK_FONT_TH_ID); + // Same cached bytes -> identical base64 program payload. + assertEquals(first.getProgram(), second.getProgram()); + } + } + + @Nested + @DisplayName("loadFallbackPdfFont") + class LoadFallbackPdfFont { + + @Test + @DisplayName("loads a Type0 PDFont for a built-in fallback id") + void loadsType0Font() throws IOException { + try (PDDocument document = new PDDocument()) { + PDFont font = + service.loadFallbackPdfFont( + document, PdfJsonFallbackFontService.FALLBACK_FONT_AR_ID); + assertNotNull(font); + assertTrue(font instanceof PDType0Font); + } + } + + @Test + @DisplayName("no-arg overload loads the default Noto Sans font") + void noArgLoadsDefaultFont() throws Exception { + invokeLoadConfig(); + try (PDDocument document = new PDDocument()) { + PDFont font = service.loadFallbackPdfFont(document); + assertNotNull(font); + assertTrue(font instanceof PDType0Font); + } + } + + @Test + @DisplayName("unknown fallback id throws IOException") + void unknownIdThrows() throws IOException { + try (PDDocument document = new PDDocument()) { + IOException ex = + assertThrows( + IOException.class, + () -> service.loadFallbackPdfFont(document, "nope")); + assertTrue(ex.getMessage().contains("Unknown fallback font id")); + } + } + + @Test + @DisplayName("loaded font produces fresh instances per call but identical type") + void distinctInstancesPerCall() throws IOException { + // Instance-identity contract holds for any built-in font; use the tiny Thai fallback + // (~22 KB) instead of the multi-MB Korean font to keep two loads cheap. + try (PDDocument document = new PDDocument()) { + PDFont a = + service.loadFallbackPdfFont( + document, PdfJsonFallbackFontService.FALLBACK_FONT_TH_ID); + PDFont b = + service.loadFallbackPdfFont( + document, PdfJsonFallbackFontService.FALLBACK_FONT_TH_ID); + assertNotNull(a); + assertNotNull(b); + // Two independent PDType0Font wrappers loaded into the same document. + assertFalse(a == b); + assertSame(a.getClass(), b.getClass()); + } + } + } +} diff --git a/app/proprietary/build.gradle b/app/proprietary/build.gradle index 6f744c6c6b..821def1fec 100644 --- a/app/proprietary/build.gradle +++ b/app/proprietary/build.gradle @@ -1,6 +1,6 @@ repositories { - maven { url = "https://build.shibboleth.net/maven/releases" } maven { url = "https://repository.jboss.org/" } + maven { url = "https://build.shibboleth.net/maven/releases" } } ext { @@ -52,6 +52,10 @@ dependencies { api 'org.springframework.boot:spring-boot-starter-security' api 'org.springframework.boot:spring-boot-starter-data-jpa' api 'org.springframework.boot:spring-boot-starter-security-oauth2-client' + // MCP server (RFC 8707 audience binding + RFC 9728 metadata) - resource-server side only. + // Brings nimbus-jose-jwt onto the proprietary classpath if not already transitive via + // oauth2-client; on Boot 4.0.6 the delta is around 130KB because nimbus is already pulled. + api 'org.springframework.boot:spring-boot-starter-oauth2-resource-server' api 'org.springframework.boot:spring-boot-starter-mail' api 'org.springframework.boot:spring-boot-starter-cache' api 'com.github.ben-manes.caffeine:caffeine' diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/cluster/s3/S3FileStore.java b/app/proprietary/src/main/java/stirling/software/proprietary/cluster/s3/S3FileStore.java index 1cff20e6a0..68cd72ded7 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/cluster/s3/S3FileStore.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/cluster/s3/S3FileStore.java @@ -6,6 +6,8 @@ import java.io.InputStream; import java.nio.file.Files; import java.nio.file.Path; import java.nio.file.StandardCopyOption; +import java.util.Collections; +import java.util.Map; import java.util.Optional; import java.util.UUID; @@ -35,6 +37,7 @@ import software.amazon.awssdk.services.s3.model.S3Exception; public class S3FileStore implements FileStore, AutoCloseable { public static final String DEFAULT_KEY_PREFIX = "transient/"; + static final String OWNER_METADATA_KEY = "owner"; private final S3Client s3Client; private final String bucket; @@ -64,7 +67,7 @@ public class S3FileStore implements FileStore, AutoCloseable { } @Override - public Stored store(InputStream in, String originalName) throws IOException { + public Stored store(InputStream in, String originalName, String owner) throws IOException { String fileId = UUID.randomUUID().toString(); // S3 PUT requires a known content-length; spool to a temp file first so memory stays // bounded for large payloads, then stream the file to S3 via RequestBody.fromFile. @@ -75,10 +78,13 @@ public class S3FileStore implements FileStore, AutoCloseable { Files.copy(src, tempFile, StandardCopyOption.REPLACE_EXISTING); } size = Files.size(tempFile); - PutObjectRequest request = - PutObjectRequest.builder().bucket(bucket).key(resolveKey(fileId)).build(); + PutObjectRequest.Builder builder = + PutObjectRequest.builder().bucket(bucket).key(resolveKey(fileId)); + if (owner != null && !owner.isBlank()) { + builder.metadata(Map.of(OWNER_METADATA_KEY, owner)); + } try { - s3Client.putObject(request, RequestBody.fromFile(tempFile)); + s3Client.putObject(builder.build(), RequestBody.fromFile(tempFile)); } catch (SdkException e) { throw new IOException("Failed to upload object to S3", e); } @@ -185,6 +191,36 @@ public class S3FileStore implements FileStore, AutoCloseable { } } + @Override + public String getOwner(String fileId) throws IOException { + try { + validateFileId(fileId); + } catch (IllegalArgumentException e) { + return null; + } + HeadObjectRequest request = + HeadObjectRequest.builder().bucket(bucket).key(resolveKey(fileId)).build(); + try { + HeadObjectResponse response = s3Client.headObject(request); + Map metadata = + Optional.ofNullable(response.metadata()).orElse(Collections.emptyMap()); + String owner = metadata.get(OWNER_METADATA_KEY); + if (owner != null && !owner.isBlank()) { + return owner; + } + return null; + } catch (NoSuchKeyException e) { + return null; + } catch (S3Exception e) { + if (e.statusCode() == 404) { + return null; + } + throw new IOException("Failed to read owner metadata from S3", e); + } catch (SdkException e) { + throw new IOException("Failed to read owner metadata from S3", e); + } + } + @Override public void close() { if (!ownsClient) { @@ -205,7 +241,7 @@ public class S3FileStore implements FileStore, AutoCloseable { if (fileId == null || fileId.isBlank()) { throw new IllegalArgumentException("File ID must not be blank"); } - if (fileId.contains("..") || fileId.contains("/") || fileId.contains("\\")) { + if (fileId.contains(".") || fileId.contains("/") || fileId.contains("\\")) { throw new IllegalArgumentException("Invalid file ID"); } } diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/config/CustomAuditEventRepository.java b/app/proprietary/src/main/java/stirling/software/proprietary/config/CustomAuditEventRepository.java index 1dd8d5f297..bcb98aa04b 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/config/CustomAuditEventRepository.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/config/CustomAuditEventRepository.java @@ -1,6 +1,10 @@ package stirling.software.proprietary.config; +import java.nio.charset.StandardCharsets; +import java.security.MessageDigest; +import java.security.NoSuchAlgorithmException; import java.time.Instant; +import java.util.HexFormat; import java.util.List; import java.util.Map; @@ -61,17 +65,43 @@ public class CustomAuditEventRepository implements AuditEventRepository { PersistentAuditEvent ent = PersistentAuditEvent.builder() - .principal(ev.getPrincipal()) + .principal(safePrincipal(ev.getPrincipal())) .type(ev.getType()) .data(auditEventData) .timestamp(ev.getTimestamp()) .build(); repo.save(ent); } catch (Exception e) { - log.error( - "Failed to persist audit event (fail-open); principal={}", - ev.getPrincipal(), - e); + log.error("Failed to persist audit event (fail-open); type={}", ev.getType(), e); + } + } + + /** Width of the {@code principal} column; longer values are hashed so the insert can't fail. */ + private static final int PRINCIPAL_MAX_LENGTH = 255; + + /** + * Hash JWT-shaped or over-long principals so the insert fits the column and stores no secret. + */ + static String safePrincipal(String principal) { + if (principal == null || principal.isBlank()) { + return "anonymous"; + } + // Hash JWTs ("eyJ...") and any over-long value rather than store verbatim. + if (principal.startsWith("eyJ") || principal.length() > PRINCIPAL_MAX_LENGTH) { + return "token:" + sha256Prefix(principal); + } + return principal; + } + + /** First 8 bytes of SHA-256 as hex: stable, one-way, collision-safe enough. */ + private static String sha256Prefix(String value) { + try { + byte[] digest = + MessageDigest.getInstance("SHA-256") + .digest(value.getBytes(StandardCharsets.UTF_8)); + return HexFormat.of().formatHex(digest, 0, 8); + } catch (NoSuchAlgorithmException e) { + return "unhashable"; } } } diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/controller/api/AuditDashboardController.java b/app/proprietary/src/main/java/stirling/software/proprietary/controller/api/AuditDashboardController.java index 837c28bf5a..6e93ed14eb 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/controller/api/AuditDashboardController.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/controller/api/AuditDashboardController.java @@ -207,16 +207,14 @@ public class AuditDashboardController { // Include standard enum types in case they're not in the database yet List enumTypes = - Arrays.stream(AuditEventType.values()) - .map(AuditEventType::name) - .collect(Collectors.toList()); + Arrays.stream(AuditEventType.values()).map(AuditEventType::name).toList(); // Combine both sources, remove duplicates, and sort Set combinedTypes = new HashSet<>(); combinedTypes.addAll(dbTypes); combinedTypes.addAll(enumTypes); - return combinedTypes.stream().sorted().collect(Collectors.toList()); + return combinedTypes.stream().sorted().toList(); } /** Export audit data as CSV. */ diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/controller/api/AuditRestController.java b/app/proprietary/src/main/java/stirling/software/proprietary/controller/api/AuditRestController.java index a208a46a81..69c9402926 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/controller/api/AuditRestController.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/controller/api/AuditRestController.java @@ -118,7 +118,7 @@ public class AuditRestController { // Convert to response format expected by frontend List eventDtos = - events.getContent().stream().map(this::convertToDto).collect(Collectors.toList()); + events.getContent().stream().map(this::convertToDto).toList(); AuditEventsResponse response = AuditEventsResponse.builder() @@ -191,19 +191,13 @@ public class AuditRestController { ChartData eventsByTypeChart = ChartData.builder() .labels(new ArrayList<>(eventsByType.keySet())) - .values( - eventsByType.values().stream() - .map(Long::intValue) - .collect(Collectors.toList())) + .values(eventsByType.values().stream().map(Long::intValue).toList()) .build(); ChartData eventsByUserChart = ChartData.builder() .labels(new ArrayList<>(eventsByUser.keySet())) - .values( - eventsByUser.values().stream() - .map(Long::intValue) - .collect(Collectors.toList())) + .values(eventsByUser.values().stream().map(Long::intValue).toList()) .build(); // Sort events by day for time series @@ -211,10 +205,7 @@ public class AuditRestController { ChartData eventsOverTimeChart = ChartData.builder() .labels(new ArrayList<>(sortedEventsByDay.keySet())) - .values( - sortedEventsByDay.values().stream() - .map(Long::intValue) - .collect(Collectors.toList())) + .values(sortedEventsByDay.values().stream().map(Long::intValue).toList()) .build(); AuditChartsData chartsData = @@ -239,16 +230,14 @@ public class AuditRestController { // Include standard enum types in case they're not in the database yet List enumTypes = - Arrays.stream(AuditEventType.values()) - .map(AuditEventType::name) - .collect(Collectors.toList()); + Arrays.stream(AuditEventType.values()).map(AuditEventType::name).toList(); // Combine both sources, remove duplicates, and sort Set combinedTypes = new HashSet<>(); combinedTypes.addAll(dbTypes); combinedTypes.addAll(enumTypes); - List result = combinedTypes.stream().sorted().collect(Collectors.toList()); + List result = combinedTypes.stream().sorted().toList(); return ResponseEntity.ok(result); } @@ -263,11 +252,7 @@ public class AuditRestController { // Use the countByPrincipal query to get unique principals List principalCounts = auditRepository.countByPrincipal(); - List users = - principalCounts.stream() - .map(arr -> (String) arr[0]) - .sorted() - .collect(Collectors.toList()); + List users = principalCounts.stream().map(arr -> (String) arr[0]).sorted().toList(); return ResponseEntity.ok(users); } diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/controller/api/CreatePdfAgentController.java b/app/proprietary/src/main/java/stirling/software/proprietary/controller/api/CreatePdfAgentController.java new file mode 100644 index 0000000000..15b7982dec --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/controller/api/CreatePdfAgentController.java @@ -0,0 +1,144 @@ +package stirling.software.proprietary.controller.api; + +import java.io.IOException; +import java.nio.charset.StandardCharsets; +import java.nio.file.Files; +import java.util.ArrayList; +import java.util.List; + +import org.apache.pdfbox.pdmodel.PDDocument; +import org.springframework.core.io.Resource; +import org.springframework.http.MediaType; +import org.springframework.http.ResponseEntity; +import org.springframework.web.bind.annotation.PostMapping; +import org.springframework.web.bind.annotation.RequestMapping; +import org.springframework.web.bind.annotation.RequestParam; +import org.springframework.web.bind.annotation.RestController; + +import io.github.pixee.security.Filenames; +import io.swagger.v3.oas.annotations.Hidden; +import io.swagger.v3.oas.annotations.Operation; +import io.swagger.v3.oas.annotations.tags.Tag; + +import lombok.RequiredArgsConstructor; +import lombok.extern.slf4j.Slf4j; + +import stirling.software.common.configuration.RuntimePathConfig; +import stirling.software.common.service.CustomPDFDocumentFactory; +import stirling.software.common.util.ProcessExecutor; +import stirling.software.common.util.TempFile; +import stirling.software.common.util.TempFileManager; +import stirling.software.common.util.WebResponseUtils; + +/** + * Dispatchable tool that converts an AI-generated HTML string to a PDF via WeasyPrint. + * + *

Called by {@link stirling.software.proprietary.service.AiWorkflowService} when the engine + * emits a {@code CREATE_PDF_FROM_HTML_AGENT} plan step. The HTML comes from a trusted Jinja + * template so sanitization is intentionally skipped. + */ +@Slf4j +@Hidden +@RestController +@RequestMapping("/api/v1/ai/tools") +@RequiredArgsConstructor +@Tag(name = "AI Tools", description = "Dispatchable AI-backed tools.") +public class CreatePdfAgentController { + + private final TempFileManager tempFileManager; + private final CustomPDFDocumentFactory pdfDocumentFactory; + private final RuntimePathConfig runtimePathConfig; + + /** + * Returns true only when WeasyPrint is definitively unavailable — either the binary could not + * be launched at all, or it launched but immediately failed to load a required system library. + * Other conversion failures (bad HTML, output errors, etc.) return false so they surface as + * real errors rather than a misleading "dependency missing" message. + */ + private static boolean isMissingDependencyError(IOException e) { + String msg = e.getMessage(); + if (msg == null) return false; + // OS could not start the process — binary not on PATH or not at the configured path. + if (msg.contains("Cannot run program")) return true; + // Process started but crashed immediately loading a shared library. + // "cannot load library" — Python/cffi error (Linux and macOS via pip) + // "Library not loaded" / "image not found" — macOS dyld error (Homebrew installs) + String lower = msg.toLowerCase(); + if (lower.contains("cannot load library")) return true; + if (lower.contains("library not loaded")) return true; + if (lower.contains("image not found")) return true; + return false; + } + + @PostMapping( + value = "/create-pdf-from-html-agent", + consumes = MediaType.MULTIPART_FORM_DATA_VALUE) + @Operation( + summary = "Convert AI-generated HTML to a PDF", + description = + "Accepts an HTML document as a plain-text parameter and returns a PDF." + + " This endpoint is dispatched by the AI workflow orchestrator as a" + + " plan step; it is not intended for direct client use.") + public ResponseEntity createPdfFromHtml( + @RequestParam("htmlContent") String htmlContent, + @RequestParam("filename") String filename) + throws Exception { + + log.info( + "[create-pdf-agent] converting HTML to PDF via WeasyPrint — html_bytes={}", + htmlContent.length()); + + try (TempFile htmlFile = tempFileManager.createManagedTempFile(".html"); + TempFile pdfFile = tempFileManager.createManagedTempFile(".pdf")) { + + Files.writeString(htmlFile.getPath(), htmlContent, StandardCharsets.UTF_8); + + List command = new ArrayList<>(); + command.add(runtimePathConfig.getWeasyPrintPath()); + command.add("-e"); + command.add("utf-8"); + command.add("-v"); + // SSRF: the HTML is self-contained and the engine validates style colours, so no + // external url() reaches WeasyPrint. For full isolation, run it network-isolated. + command.add(htmlFile.getAbsolutePath()); + command.add(pdfFile.getAbsolutePath()); + + try { + ProcessExecutor.getInstance(ProcessExecutor.Processes.WEASYPRINT) + .runCommandWithOutputHandling(command); + } catch (IOException e) { + if (isMissingDependencyError(e)) { + throw new IOException( + "AI document creation is not available on this server because a required" + + " system dependency is not installed. Please contact your" + + " system administrator."); + } + throw e; + } + + String safeFilename = Filenames.toSimpleFileName(filename); + if (safeFilename == null || safeFilename.isBlank() || !safeFilename.endsWith(".pdf")) { + safeFilename = "generated-document.pdf"; + } + + // Stamp the standard Stirling metadata onto the WeasyPrint output and write the result + // straight to the response temp file. Loading from the file and saving to the file + // avoids materialising the whole document as a byte[] twice (read-all + re-serialise), + // which matters for large generated documents. + TempFile tempOut = tempFileManager.createManagedTempFile(".pdf"); + try (PDDocument document = pdfDocumentFactory.load(pdfFile.getPath())) { + document.save(tempOut.getPath().toFile()); + } catch (Exception e) { + tempOut.close(); + throw e; + } + + log.info( + "[create-pdf-agent] PDF ready — filename={} bytes={}", + safeFilename, + Files.size(tempOut.getPath())); + + return WebResponseUtils.pdfFileToWebResponse(tempOut, safeFilename); + } + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/controller/api/SignatureController.java b/app/proprietary/src/main/java/stirling/software/proprietary/controller/api/SignatureController.java index 3295e62abb..5333ff18cc 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/controller/api/SignatureController.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/controller/api/SignatureController.java @@ -3,7 +3,6 @@ package stirling.software.proprietary.controller.api; import java.io.IOException; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import java.util.List; import java.util.Map; import java.util.stream.Stream; @@ -183,7 +182,7 @@ public class SignatureController { */ private boolean deleteFromSharedFolder(String signatureId) throws IOException { String signatureBasePath = InstallationPathConfig.getSignaturesPath(); - Path sharedFolder = Paths.get(signatureBasePath, ALL_USERS_FOLDER); + Path sharedFolder = Path.of(signatureBasePath, ALL_USERS_FOLDER); boolean deleted = false; if (Files.exists(sharedFolder)) { diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/controller/api/UsageRestController.java b/app/proprietary/src/main/java/stirling/software/proprietary/controller/api/UsageRestController.java index 9f4fa470ba..1230d928cc 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/controller/api/UsageRestController.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/controller/api/UsageRestController.java @@ -3,7 +3,6 @@ package stirling.software.proprietary.controller.api; import java.time.Duration; import java.time.Instant; import java.util.*; -import java.util.stream.Collectors; import org.springframework.http.ResponseEntity; import org.springframework.security.access.prepost.PreAuthorize; @@ -85,7 +84,7 @@ public class UsageRestController { .build(); }) .sorted(Comparator.comparingInt(EndpointStatistic::getVisits).reversed()) - .collect(Collectors.toList()); + .toList(); // Apply limit if specified if (limit != null && limit > 0 && statistics.size() > limit) { diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/McpCallContext.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/McpCallContext.java new file mode 100644 index 0000000000..7bbf46db39 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/McpCallContext.java @@ -0,0 +1,15 @@ +package stirling.software.proprietary.mcp; + +import java.util.Set; + +/** Per-call context: resolved Stirling identity and granted scopes for an {@link McpTool#call}. */ +public record McpCallContext( + String stirlingUserId, Set grantedScopes, boolean scopesEnabled) { + + public boolean hasScope(String required) { + if (!scopesEnabled) { + return true; + } + return required == null || required.isBlank() || grantedScopes.contains(required); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/McpServerController.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/McpServerController.java new file mode 100644 index 0000000000..d3103b8d4c --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/McpServerController.java @@ -0,0 +1,217 @@ +package stirling.software.proprietary.mcp; + +import java.util.HashMap; +import java.util.List; +import java.util.Map; +import java.util.Set; + +import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty; +import org.springframework.http.HttpStatus; +import org.springframework.http.MediaType; +import org.springframework.http.ResponseEntity; +import org.springframework.http.converter.HttpMessageNotReadableException; +import org.springframework.web.bind.annotation.ExceptionHandler; +import org.springframework.web.bind.annotation.PostMapping; +import org.springframework.web.bind.annotation.RequestBody; +import org.springframework.web.bind.annotation.RequestMapping; +import org.springframework.web.bind.annotation.RestController; + +import lombok.extern.slf4j.Slf4j; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.proprietary.mcp.jsonrpc.JsonRpcError; +import stirling.software.proprietary.mcp.jsonrpc.JsonRpcRequest; +import stirling.software.proprietary.mcp.jsonrpc.JsonRpcResponse; + +import tools.jackson.databind.JsonNode; +import tools.jackson.databind.ObjectMapper; +import tools.jackson.databind.node.ArrayNode; +import tools.jackson.databind.node.ObjectNode; + +/** Streamable-HTTP MCP server endpoint serving JSON-RPC 2.0 frames on {@code POST /mcp}. */ +@Slf4j +@RestController +@RequestMapping +@ConditionalOnProperty(name = "mcp.enabled", havingValue = "true") +public class McpServerController { + + private static final String PREFERRED_PROTOCOL_VERSION = "2025-06-18"; + private static final Set SUPPORTED_PROTOCOL_VERSIONS = + Set.of("2025-06-18", "2025-03-26", "2024-11-05"); + private static final String SERVER_NAME = "stirling-pdf-mcp"; + + private final ObjectMapper mapper; + private final ApplicationProperties applicationProperties; + private final Map toolsByName; + + public McpServerController( + ObjectMapper mapper, ApplicationProperties applicationProperties, List tools) { + this.mapper = mapper; + this.applicationProperties = applicationProperties; + this.toolsByName = new HashMap<>(); + for (McpTool tool : tools) { + this.toolsByName.put(tool.name(), tool); + } + log.info( + "MCP server controller wired with {} tool(s): {}", + toolsByName.size(), + toolsByName.keySet()); + } + + @PostMapping( + path = "/mcp", + consumes = MediaType.APPLICATION_JSON_VALUE, + produces = MediaType.APPLICATION_JSON_VALUE) + public ResponseEntity handle(@RequestBody JsonNode body) { + JsonRpcRequest request = decode(body); + if (request == null) { + // Valid JSON but not a JSON-RPC request -> Invalid Request, not Parse error. + return ResponseEntity.badRequest() + .body( + JsonRpcResponse.failure( + null, + JsonRpcError.invalidRequest( + "Body is not a valid JSON-RPC 2.0 request"))); + } + if (request.isNotification()) { + log.debug("Notification received: {}", sanitizeForLog(request.method())); + return ResponseEntity.status(HttpStatus.NO_CONTENT).build(); + } + JsonRpcResponse response; + try { + response = dispatch(request); + } catch (RuntimeException e) { + log.warn( + "MCP dispatch failed for method {}: {}", + sanitizeForLog(request.method()), + e.getMessage(), + e); + response = + JsonRpcResponse.failure( + request.id(), + JsonRpcError.internalError( + "Internal error handling " + request.method())); + } + return ResponseEntity.ok(response); + } + + /** Wrap malformed-JSON failures (caught before {@link #handle}) as a JSON-RPC Parse error. */ + @ExceptionHandler(HttpMessageNotReadableException.class) + public ResponseEntity handleUnreadable(HttpMessageNotReadableException ex) { + return ResponseEntity.badRequest() + .contentType(MediaType.APPLICATION_JSON) + .body( + JsonRpcResponse.failure( + null, JsonRpcError.parseError("Request body is not valid JSON"))); + } + + private static String sanitizeForLog(String value) { + return value == null ? null : value.replace('\r', ' ').replace('\n', ' '); + } + + private JsonRpcRequest decode(JsonNode body) { + if (body == null || !body.isObject()) { + return null; + } + JsonNode jsonrpc = body.get("jsonrpc"); + JsonNode method = body.get("method"); + if (jsonrpc == null || !"2.0".equals(jsonrpc.asText())) { + return null; + } + if (method == null || !method.isTextual()) { + return null; + } + return new JsonRpcRequest( + jsonrpc.asText(), body.get("id"), method.asText(), body.get("params")); + } + + private JsonRpcResponse dispatch(JsonRpcRequest request) { + return switch (request.method()) { + case "initialize" -> + JsonRpcResponse.success(request.id(), initializeResult(request.params())); + case "tools/list" -> JsonRpcResponse.success(request.id(), toolsListResult()); + case "tools/call" -> handleToolsCall(request); + case "ping" -> JsonRpcResponse.success(request.id(), mapper.createObjectNode()); + case "notifications/initialized" -> + JsonRpcResponse.success(request.id(), mapper.createObjectNode()); + default -> + JsonRpcResponse.failure( + request.id(), JsonRpcError.methodNotFound(request.method())); + }; + } + + private ObjectNode initializeResult(JsonNode params) { + ObjectNode result = mapper.createObjectNode(); + // Echo the client's requested protocolVersion when supported, else advertise our preferred. + String requested = + params != null && params.hasNonNull("protocolVersion") + ? params.get("protocolVersion").asText() + : null; + String negotiated = + requested != null && SUPPORTED_PROTOCOL_VERSIONS.contains(requested) + ? requested + : PREFERRED_PROTOCOL_VERSION; + result.put("protocolVersion", negotiated); + ObjectNode caps = result.putObject("capabilities"); + caps.putObject("tools"); + ObjectNode info = result.putObject("serverInfo"); + info.put("name", SERVER_NAME); + info.put("version", applicationProperties.getAutomaticallyGenerated().getAppVersion()); + return result; + } + + private ObjectNode toolsListResult() { + ObjectNode result = mapper.createObjectNode(); + ArrayNode tools = result.putArray("tools"); + for (McpTool t : toolsByName.values()) { + ObjectNode entry = mapper.createObjectNode(); + entry.put("name", t.name()); + entry.put("description", t.description()); + entry.set("inputSchema", t.inputSchema()); + tools.add(entry); + } + return result; + } + + private JsonRpcResponse handleToolsCall(JsonRpcRequest request) { + JsonNode params = request.params(); + if (params == null || !params.isObject()) { + return JsonRpcResponse.failure( + request.id(), JsonRpcError.invalidParams("Missing params for tools/call")); + } + JsonNode nameNode = params.get("name"); + if (nameNode == null || !nameNode.isTextual()) { + return JsonRpcResponse.failure( + request.id(), JsonRpcError.invalidParams("Missing tool name")); + } + McpTool tool = toolsByName.get(nameNode.asText()); + if (tool == null) { + return JsonRpcResponse.failure( + request.id(), JsonRpcError.invalidParams("Unknown tool: " + nameNode.asText())); + } + JsonNode args = params.get("arguments"); + McpCallContext context = resolveContext(); + ObjectNode toolResult = tool.call(args == null ? mapper.createObjectNode() : args, context); + return JsonRpcResponse.success(request.id(), toolResult); + } + + private McpCallContext resolveContext() { + boolean scopesEnabled = applicationProperties.getMcp().isScopesEnabled(); + org.springframework.security.core.Authentication auth = + org.springframework.security.core.context.SecurityContextHolder.getContext() + .getAuthentication(); + // Fail closed: no/unauthenticated principal yields an empty context so scoped ops are + // refused. + if (auth == null || !auth.isAuthenticated() || auth.getName() == null) { + return new McpCallContext(null, Set.of(), scopesEnabled); + } + java.util.Set scopes = new java.util.HashSet<>(); + for (org.springframework.security.core.GrantedAuthority ga : auth.getAuthorities()) { + String authority = ga.getAuthority(); + if (authority != null && authority.startsWith("SCOPE_")) { + scopes.add(authority.substring("SCOPE_".length())); + } + } + return new McpCallContext(auth.getName(), scopes, scopesEnabled); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/McpTool.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/McpTool.java new file mode 100644 index 0000000000..271ed9858f --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/McpTool.java @@ -0,0 +1,18 @@ +package stirling.software.proprietary.mcp; + +import tools.jackson.databind.JsonNode; +import tools.jackson.databind.node.ObjectNode; + +/** Contract every MCP tool registered with the server must satisfy. */ +public interface McpTool { + + String name(); + + String description(); + + /** The tool's {@code inputSchema} (an object JSON Schema) published in {@code tools/list}. */ + ObjectNode inputSchema(); + + /** Execute the tool; the controller wraps any thrown exception as an MCP internal error. */ + ObjectNode call(JsonNode arguments, McpCallContext context); +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/catalog/McpToolCatalog.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/catalog/McpToolCatalog.java new file mode 100644 index 0000000000..2821c9181f --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/catalog/McpToolCatalog.java @@ -0,0 +1,266 @@ +package stirling.software.proprietary.mcp.catalog; + +import java.lang.reflect.Method; +import java.util.ArrayList; +import java.util.Collections; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import java.util.Optional; +import java.util.Set; +import java.util.TreeSet; +import java.util.concurrent.ConcurrentHashMap; + +import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty; +import org.springframework.context.ApplicationContext; +import org.springframework.context.event.ContextRefreshedEvent; +import org.springframework.context.event.EventListener; +import org.springframework.core.MethodParameter; +import org.springframework.stereotype.Component; +import org.springframework.web.bind.annotation.RequestMethod; +import org.springframework.web.method.HandlerMethod; +import org.springframework.web.servlet.mvc.method.RequestMappingInfo; +import org.springframework.web.servlet.mvc.method.annotation.RequestMappingHandlerMapping; + +import io.swagger.v3.oas.annotations.Operation; + +import lombok.extern.slf4j.Slf4j; + +import stirling.software.SPDF.config.EndpointConfiguration; +import stirling.software.common.model.ApplicationProperties; + +import tools.jackson.databind.ObjectMapper; +import tools.jackson.databind.node.ObjectNode; + +/** + * Discovers MCP-exposable operations and caches a per-op {@link OperationMeta}. Refreshed on {@link + * ContextRefreshedEvent} and filtered on read by {@link + * EndpointConfiguration#isEndpointEnabledForUri}. AI capabilities are fed in via {@link + * #replaceAiCapabilities}. + */ +@Slf4j +@Component +@ConditionalOnProperty(name = "mcp.enabled", havingValue = "true") +public class McpToolCatalog { + + private static final String WRITE_SCOPE = "mcp.tools.write"; + + private final ApplicationContext applicationContext; + private final EndpointConfiguration endpointConfiguration; + private final ApplicationProperties applicationProperties; + private final SimpleSchemaGenerator schemaGenerator; + private final ObjectMapper objectMapper; + + // Concurrent: written on the boot thread, read on request threads, AI map replaced at runtime. + private final Map pdfOps = new ConcurrentHashMap<>(); + + // Engine-driven AI capabilities. Replaced wholesale by the scheduled refresh task on a + // background thread while request threads read via findByOperationId/enabledOps. The volatile + // reference makes the swap publication-safe; readers either see the old or the new snapshot, + // never a partially-merged one. + private volatile Map aiOps = new ConcurrentHashMap<>(); + + public McpToolCatalog( + ApplicationContext applicationContext, + EndpointConfiguration endpointConfiguration, + ApplicationProperties applicationProperties, + ObjectMapper objectMapper) { + this.applicationContext = applicationContext; + this.endpointConfiguration = endpointConfiguration; + this.applicationProperties = applicationProperties; + this.schemaGenerator = new SimpleSchemaGenerator(objectMapper); + this.objectMapper = objectMapper; + } + + /** Admin tool filter: non-empty allow list is a whitelist; block list always removes. */ + private boolean isOperationAllowed(String id) { + ApplicationProperties.Mcp mcp = applicationProperties.getMcp(); + List allowed = mcp.getAllowedOperations(); + List blocked = mcp.getBlockedOperations(); + if (blocked != null && blocked.contains(id)) { + return false; + } + if (allowed != null && !allowed.isEmpty()) { + return allowed.contains(id); + } + return true; + } + + @EventListener(ContextRefreshedEvent.class) + public void discover() { + pdfOps.clear(); + for (RequestMappingHandlerMapping mapping : + applicationContext.getBeansOfType(RequestMappingHandlerMapping.class).values()) { + for (Map.Entry e : + mapping.getHandlerMethods().entrySet()) { + indexOne(e.getKey(), e.getValue()); + } + } + log.info("MCP tool catalog discovered {} PDF operation(s)", pdfOps.size()); + } + + private void indexOne(RequestMappingInfo info, HandlerMethod handler) { + Set patterns = extractPatterns(info); + if (patterns.isEmpty()) { + return; + } + Set methods = info.getMethodsCondition().getMethods(); + if (!isInvocableMethod(methods)) { + return; + } + for (String pattern : patterns) { + OperationCategory category = OperationCategory.fromUrl(pattern); + if (category == null) { + continue; + } + String opId = extractOpId(pattern, category); + if (opId == null) { + continue; + } + OperationMeta meta = buildMeta(opId, category, pattern, handler); + // First handler wins on duplicate URLs. + pdfOps.putIfAbsent(opId, meta); + } + } + + private OperationMeta buildMeta( + String opId, OperationCategory category, String url, HandlerMethod handler) { + Method method = handler.getMethod(); + Operation opAnno = method.getAnnotation(Operation.class); + String summary = + opAnno != null && !opAnno.summary().isBlank() + ? opAnno.summary() + : prettifyOpId(opId); + ObjectNode schema = paramSchemaFor(handler); + // Every mutating endpoint requires the write scope. + return new OperationMeta( + opId, + category, + summary, + schema, + WRITE_SCOPE, + OperationMeta.Target.JAVA_ENDPOINT, + url, + handler); + } + + private ObjectNode paramSchemaFor(HandlerMethod handler) { + Optional> bodyType = firstComplexParamType(handler); + return bodyType.map(schemaGenerator::toSchema).orElseGet(() -> emptyObjectSchema()); + } + + private ObjectNode emptyObjectSchema() { + ObjectNode out = objectMapper.createObjectNode(); + out.put("type", "object"); + out.put("additionalProperties", true); + return out; + } + + private Optional> firstComplexParamType(HandlerMethod handler) { + for (MethodParameter p : handler.getMethodParameters()) { + Class type = p.getParameterType(); + if (type.isPrimitive() || type == String.class || type.getName().startsWith("java.")) { + continue; + } + // Skip Spring-managed parameter types (HttpServletRequest, Principal, etc.). + String pkg = type.getPackageName(); + if (pkg.startsWith("jakarta.") || pkg.startsWith("org.springframework.")) { + continue; + } + return Optional.of(type); + } + return Optional.empty(); + } + + public List enabledOps(OperationCategory category) { + if (category == OperationCategory.AI) { + List ai = new ArrayList<>(); + for (OperationMeta m : aiOps.values()) { + if (isOperationAllowed(m.id())) { + ai.add(m); + } + } + return ai; + } + List out = new ArrayList<>(); + for (OperationMeta m : pdfOps.values()) { + if (m.category() == category + && isOperationAllowed(m.id()) + && endpointConfiguration.isEndpointEnabledForUri(m.endpointPath())) { + out.add(m); + } + } + out.sort((a, b) -> a.id().compareTo(b.id())); + return out; + } + + public Optional findByOperationId(String id) { + if (!isOperationAllowed(id)) { + return Optional.empty(); + } + // A disabled PDF op returns empty rather than falling through to a same-id AI capability. + OperationMeta meta = pdfOps.get(id); + if (meta != null) { + boolean enabled = + meta.target() != OperationMeta.Target.JAVA_ENDPOINT + || endpointConfiguration.isEndpointEnabledForUri(meta.endpointPath()); + return enabled ? Optional.of(meta) : Optional.empty(); + } + return Optional.ofNullable(aiOps.get(id)); + } + + /** Replace the AI capabilities snapshot. Called by the engine refresh task. */ + public void replaceAiCapabilities(Map updated) { + // Build a fresh map then swap atomically via the volatile reference. The previous + // implementation did putAll-then-retainAll on a shared ConcurrentHashMap, which left a + // transient window where readers could observe stale entries that should have been + // removed (race between the two structural updates). + Map next = new ConcurrentHashMap<>(updated); + this.aiOps = next; + log.info("MCP tool catalog AI capabilities replaced: {} entries", next.size()); + } + + /** Only POST/PUT endpoints are exposed as tools; DELETE and GET are excluded. */ + static boolean isInvocableMethod(Set methods) { + return methods.contains(RequestMethod.POST) || methods.contains(RequestMethod.PUT); + } + + private static String extractOpId(String pattern, OperationCategory category) { + if (category.urlPrefix() == null || !pattern.startsWith(category.urlPrefix())) { + return null; + } + String tail = pattern.substring(category.urlPrefix().length()); + if (tail.isBlank() || tail.contains("/") || tail.contains("{")) { + // Skip nested paths and path-variable templates. + return null; + } + return tail; + } + + private static String prettifyOpId(String id) { + return id.replace('-', ' '); + } + + private static Set extractPatterns(RequestMappingInfo info) { + try { + Method getDirectPaths = info.getClass().getMethod("getDirectPaths"); + Object result = getDirectPaths.invoke(info); + if (result instanceof Set set) { + Set patterns = new TreeSet<>(); + for (Object v : set) { + if (v instanceof String s) { + patterns.add(s); + } + } + return patterns; + } + } catch (Exception e) { + log.trace("getDirectPaths unavailable on RequestMappingInfo", e); + } + return Collections.emptySet(); + } + + public Map snapshotPdfOps() { + return new LinkedHashMap<>(pdfOps); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/catalog/OperationCategory.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/catalog/OperationCategory.java new file mode 100644 index 0000000000..7be79b7967 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/catalog/OperationCategory.java @@ -0,0 +1,38 @@ +package stirling.software.proprietary.mcp.catalog; + +/** MCP tool categories; {@link #urlPrefix} maps a {@code /api/v1/} namespace to a category. */ +public enum OperationCategory { + CONVERT("/api/v1/convert/", "stirling_convert"), + PAGES("/api/v1/general/", "stirling_pages"), + MISC("/api/v1/misc/", "stirling_misc"), + SECURITY("/api/v1/security/", "stirling_security"), + AI(null, "stirling_ai"); + + private final String urlPrefix; + private final String toolName; + + OperationCategory(String urlPrefix, String toolName) { + this.urlPrefix = urlPrefix; + this.toolName = toolName; + } + + public String urlPrefix() { + return urlPrefix; + } + + public String toolName() { + return toolName; + } + + public static OperationCategory fromUrl(String url) { + if (url == null) { + return null; + } + for (OperationCategory c : values()) { + if (c.urlPrefix != null && url.startsWith(c.urlPrefix)) { + return c; + } + } + return null; + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/catalog/OperationMeta.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/catalog/OperationMeta.java new file mode 100644 index 0000000000..a9d8b2a9a9 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/catalog/OperationMeta.java @@ -0,0 +1,22 @@ +package stirling.software.proprietary.mcp.catalog; + +import org.springframework.web.method.HandlerMethod; + +import tools.jackson.databind.node.ObjectNode; + +/** Metadata for one MCP-exposed operation (PDF endpoint or AI capability). */ +public record OperationMeta( + String id, + OperationCategory category, + String summary, + ObjectNode paramSchema, + String requiredScope, + Target target, + String endpointPath, + HandlerMethod handlerMethod) { + + public enum Target { + JAVA_ENDPOINT, + ENGINE_CAPABILITY + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/catalog/SimpleSchemaGenerator.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/catalog/SimpleSchemaGenerator.java new file mode 100644 index 0000000000..0cef8d32bf --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/catalog/SimpleSchemaGenerator.java @@ -0,0 +1,178 @@ +package stirling.software.proprietary.mcp.catalog; + +import java.lang.reflect.Field; +import java.lang.reflect.ParameterizedType; +import java.lang.reflect.Type; +import java.util.ArrayList; +import java.util.HashSet; +import java.util.List; +import java.util.Set; + +import org.springframework.web.multipart.MultipartFile; + +import com.fasterxml.jackson.annotation.JsonIgnore; +import com.fasterxml.jackson.annotation.JsonProperty; + +import tools.jackson.databind.ObjectMapper; +import tools.jackson.databind.node.ArrayNode; +import tools.jackson.databind.node.ObjectNode; + +/** + * Reflection-based JSON Schema generator for controller request-body classes. {@link MultipartFile} + * fields are emitted as {@code "type":"string"} with a {@code "format":"file-id"} hint. + */ +public final class SimpleSchemaGenerator { + + private final ObjectMapper mapper; + + public SimpleSchemaGenerator(ObjectMapper mapper) { + this.mapper = mapper; + } + + public ObjectNode toSchema(Class type) { + return toSchema(type, new HashSet<>()); + } + + private ObjectNode toSchema(Class type, Set> visited) { + ObjectNode schema = mapper.createObjectNode(); + schema.put("type", "object"); + schema.put("additionalProperties", false); + ObjectNode properties = schema.putObject("properties"); + ArrayNode required = mapper.createArrayNode(); + if (!visited.add(type)) { + // Cycle: emit a loose object and bail. + schema.put("additionalProperties", true); + return schema; + } + + Set seen = new HashSet<>(); + for (Field field : collectFields(type)) { + if (java.lang.reflect.Modifier.isStatic(field.getModifiers()) + || java.lang.reflect.Modifier.isTransient(field.getModifiers())) { + continue; + } + // Skip fields Jackson won't (de)serialize. + if (field.isAnnotationPresent(JsonIgnore.class)) { + continue; + } + String name = jsonPropertyName(field); + if (!seen.add(name)) { + continue; + } + properties.set(name, typeSchema(field.getGenericType(), visited)); + if (isRequired(field)) { + required.add(name); + } + } + + if (!required.isEmpty()) { + schema.set("required", required); + } + return schema; + } + + private List collectFields(Class type) { + List all = new ArrayList<>(); + for (Class c = type; c != null && c != Object.class; c = c.getSuperclass()) { + for (Field f : c.getDeclaredFields()) { + all.add(f); + } + } + return all; + } + + private static String jsonPropertyName(Field field) { + JsonProperty ann = field.getAnnotation(JsonProperty.class); + if (ann != null && !ann.value().isEmpty()) { + return ann.value(); + } + return field.getName(); + } + + private boolean isRequired(Field field) { + JsonProperty json = field.getAnnotation(JsonProperty.class); + if (json != null && json.required()) { + return true; + } + return field.isAnnotationPresent(jakarta.validation.constraints.NotNull.class) + || field.isAnnotationPresent(jakarta.validation.constraints.NotBlank.class) + || field.isAnnotationPresent(jakarta.validation.constraints.NotEmpty.class); + } + + private ObjectNode typeSchema(Type t, Set> visited) { + ObjectNode out = mapper.createObjectNode(); + if (t instanceof Class c) { + populatePrimitive(out, c, visited); + } else if (t instanceof ParameterizedType pt) { + Type raw = pt.getRawType(); + if (raw instanceof Class rawClass) { + if (java.util.Collection.class.isAssignableFrom(rawClass)) { + out.put("type", "array"); + Type[] args = pt.getActualTypeArguments(); + if (args.length == 1) { + out.set("items", typeSchema(args[0], visited)); + } + } else if (java.util.Map.class.isAssignableFrom(rawClass)) { + out.put("type", "object"); + out.put("additionalProperties", true); + } else { + populatePrimitive(out, rawClass, visited); + } + } else { + out.put("type", "object"); + } + } else { + out.put("type", "object"); + } + return out; + } + + private void populatePrimitive(ObjectNode out, Class c, Set> visited) { + if (MultipartFile.class.isAssignableFrom(c)) { + out.put("type", "string"); + out.put("format", "file-id"); + out.put( + "description", + "Reference to a previously-uploaded file in Stirling's job store."); + return; + } + if (c.isArray()) { + out.put("type", "array"); + out.set("items", typeSchema(c.getComponentType(), visited)); + return; + } + if (c == String.class) { + out.put("type", "string"); + } else if (c == boolean.class || c == Boolean.class) { + out.put("type", "boolean"); + } else if (c == int.class + || c == Integer.class + || c == long.class + || c == Long.class + || c == short.class + || c == Short.class + || c == byte.class + || c == Byte.class) { + out.put("type", "integer"); + } else if (c == float.class || c == Float.class || c == double.class || c == Double.class) { + out.put("type", "number"); + } else if (c.isEnum()) { + out.put("type", "string"); + ArrayNode values = out.putArray("enum"); + for (Object constant : c.getEnumConstants()) { + values.add(constant.toString()); + } + } else if (c == java.util.UUID.class) { + out.put("type", "string"); + out.put("format", "uuid"); + } else if (java.time.temporal.Temporal.class.isAssignableFrom(c) + || c == java.util.Date.class) { + out.put("type", "string"); + out.put("format", "date-time"); + } else { + // Complex bean: recurse with the shared visited set. + ObjectNode nested = toSchema(c, visited); + out.setAll(nested); + } + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/engine/EngineCapabilityClient.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/engine/EngineCapabilityClient.java new file mode 100644 index 0000000000..ed023bd25b --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/engine/EngineCapabilityClient.java @@ -0,0 +1,204 @@ +package stirling.software.proprietary.mcp.engine; + +import java.io.IOException; +import java.net.URI; +import java.net.http.HttpClient; +import java.net.http.HttpRequest; +import java.net.http.HttpResponse; +import java.time.Duration; +import java.util.LinkedHashMap; +import java.util.Map; +import java.util.concurrent.Executors; +import java.util.concurrent.ScheduledExecutorService; +import java.util.concurrent.TimeUnit; + +import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty; +import org.springframework.boot.context.event.ApplicationReadyEvent; +import org.springframework.context.event.EventListener; +import org.springframework.stereotype.Component; + +import jakarta.annotation.PostConstruct; +import jakarta.annotation.PreDestroy; + +import lombok.extern.slf4j.Slf4j; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.proprietary.mcp.catalog.McpToolCatalog; +import stirling.software.proprietary.mcp.catalog.OperationCategory; +import stirling.software.proprietary.mcp.catalog.OperationMeta; + +import tools.jackson.databind.JsonNode; +import tools.jackson.databind.ObjectMapper; +import tools.jackson.databind.node.ObjectNode; + +/** + * Pulls the engine's capabilities manifest at boot and on a schedule, feeding it into the shared + * {@link McpToolCatalog}. + */ +@Slf4j +@Component +@ConditionalOnProperty(name = "mcp.enabled", havingValue = "true") +public class EngineCapabilityClient { + + private final ApplicationProperties applicationProperties; + private final McpToolCatalog catalog; + private final ObjectMapper mapper; + private final HttpClient httpClient; + private final String sharedSecret; + + private ScheduledExecutorService scheduler; + + public EngineCapabilityClient( + ApplicationProperties applicationProperties, + McpToolCatalog catalog, + ObjectMapper mapper) { + this.applicationProperties = applicationProperties; + this.catalog = catalog; + this.mapper = mapper; + this.httpClient = HttpClient.newBuilder().connectTimeout(Duration.ofSeconds(5)).build(); + this.sharedSecret = System.getenv("STIRLING_ENGINE_SHARED_SECRET"); + } + + @PostConstruct + void start() { + scheduler = + Executors.newSingleThreadScheduledExecutor( + r -> { + Thread t = new Thread(r, "mcp-engine-capability-refresh"); + t.setDaemon(true); + return t; + }); + } + + @EventListener(ApplicationReadyEvent.class) + public void onReady() { + long minutes = + Math.max(1, applicationProperties.getMcp().getEngineCapabilityRefreshMinutes()); + // First refresh immediately, then on the configured cadence. + scheduler.schedule(this::refreshSafely, 0, TimeUnit.SECONDS); + scheduler.scheduleAtFixedRate(this::refreshSafely, minutes, minutes, TimeUnit.MINUTES); + log.info("MCP engine capability refresh scheduled every {} minute(s)", minutes); + } + + @PreDestroy + void stop() { + if (scheduler != null) { + scheduler.shutdownNow(); + } + } + + private void refreshSafely() { + try { + refresh(); + } catch (Exception e) { + log.warn( + "MCP engine capability refresh failed ({}). AI tool enum stays at the last" + + " known state until the next successful pull.", + e.getMessage()); + } + } + + /** Visible for testing. */ + public void refresh() throws IOException, InterruptedException { + if (!applicationProperties.getAiEngine().isEnabled()) { + log.debug("AI engine disabled; skipping MCP capability refresh"); + catalog.replaceAiCapabilities(Map.of()); + return; + } + // Trim whitespace and any trailing slash to avoid a malformed URI. + String base = applicationProperties.getAiEngine().getUrl().strip().replaceAll("/+$", ""); + URI uri = URI.create(base + "/api/v1/agents/capabilities"); + HttpRequest.Builder reqBuilder = + HttpRequest.newBuilder() + .uri(uri) + .timeout(Duration.ofSeconds(10)) + .header("Accept", "application/json") + .GET(); + if (sharedSecret != null && !sharedSecret.isBlank()) { + reqBuilder.header("X-Engine-Auth", sharedSecret); + } + HttpResponse response = + httpClient.send(reqBuilder.build(), HttpResponse.BodyHandlers.ofString()); + if (response.statusCode() != 200) { + throw new IOException( + "Engine capabilities endpoint returned HTTP " + response.statusCode()); + } + Map parsed = parseManifest(response.body()); + catalog.replaceAiCapabilities(parsed); + } + + private Map parseManifest(String body) throws IOException { + JsonNode root = mapper.readTree(body); + JsonNode capabilities = root.get("capabilities"); + if (capabilities == null || !capabilities.isArray()) { + throw new IOException("Manifest missing 'capabilities' array"); + } + Map out = new LinkedHashMap<>(); + for (JsonNode entry : capabilities) { + JsonNode id = entry.get("id"); + JsonNode desc = entry.get("description"); + JsonNode schema = entry.get("input_schema"); + JsonNode scope = entry.get("required_scope"); + JsonNode route = entry.get("route"); + if (id == null || !id.isTextual() || schema == null || !schema.isObject()) { + log.warn("Skipping malformed capability entry: {}", entry); + continue; + } + String routeValue = route == null || !route.isTextual() ? null : route.asText(); + if (routeValue != null && !isSafeRelativeRoute(routeValue)) { + // Defence in depth: a tampered manifest must not steer Java at an arbitrary + // host/path. + log.warn( + "Skipping capability '{}' with unsafe route '{}' (must be a server-relative" + + " /api path with no scheme, authority, or '..')", + id.asText(), + routeValue); + continue; + } + // Fail safe: default to the stricter write scope when the manifest omits one. + String requiredScope = + scope != null && scope.isTextual() && !scope.asText().isBlank() + ? scope.asText() + : WRITE_SCOPE; + ObjectNode schemaCopy = (ObjectNode) schema.deepCopy(); + out.put( + id.asText(), + new OperationMeta( + id.asText(), + OperationCategory.AI, + desc == null ? id.asText() : desc.asText(), + schemaCopy, + requiredScope, + OperationMeta.Target.ENGINE_CAPABILITY, + routeValue, + null)); + } + return out; + } + + private static final String WRITE_SCOPE = "mcp.tools.write"; + + /** + * True only for a server-relative {@code /api/} path with no scheme, authority, {@code ..}, or + * control chars (blocks SSRF / path escape). + */ + static boolean isSafeRelativeRoute(String route) { + if (route == null || route.isBlank() || !route.startsWith("/api/")) { + return false; + } + if (route.startsWith("//") + || route.contains("..") + || route.contains("@") + || route.contains("\\") + || route.contains(":")) { + return false; + } + for (int i = 0; i < route.length(); i++) { + char c = route.charAt(i); + if (c <= ' ') { + return false; + } + } + return true; + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/jsonrpc/JsonRpcError.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/jsonrpc/JsonRpcError.java new file mode 100644 index 0000000000..e383605a59 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/jsonrpc/JsonRpcError.java @@ -0,0 +1,36 @@ +package stirling.software.proprietary.mcp.jsonrpc; + +import com.fasterxml.jackson.annotation.JsonInclude; + +import tools.jackson.databind.JsonNode; + +/** JSON-RPC 2.0 error object. */ +@JsonInclude(JsonInclude.Include.NON_NULL) +public record JsonRpcError(int code, String message, JsonNode data) { + + public static final int PARSE_ERROR = -32700; + public static final int INVALID_REQUEST = -32600; + public static final int METHOD_NOT_FOUND = -32601; + public static final int INVALID_PARAMS = -32602; + public static final int INTERNAL_ERROR = -32603; + + public static JsonRpcError parseError(String message) { + return new JsonRpcError(PARSE_ERROR, message, null); + } + + public static JsonRpcError invalidRequest(String message) { + return new JsonRpcError(INVALID_REQUEST, message, null); + } + + public static JsonRpcError methodNotFound(String method) { + return new JsonRpcError(METHOD_NOT_FOUND, "Method not found: " + method, null); + } + + public static JsonRpcError invalidParams(String message) { + return new JsonRpcError(INVALID_PARAMS, message, null); + } + + public static JsonRpcError internalError(String message) { + return new JsonRpcError(INTERNAL_ERROR, message, null); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/jsonrpc/JsonRpcRequest.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/jsonrpc/JsonRpcRequest.java new file mode 100644 index 0000000000..01466746e1 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/jsonrpc/JsonRpcRequest.java @@ -0,0 +1,14 @@ +package stirling.software.proprietary.mcp.jsonrpc; + +import com.fasterxml.jackson.annotation.JsonInclude; + +import tools.jackson.databind.JsonNode; + +/** JSON-RPC 2.0 request frame; a null {@code id} marks a notification. */ +@JsonInclude(JsonInclude.Include.NON_NULL) +public record JsonRpcRequest(String jsonrpc, JsonNode id, String method, JsonNode params) { + + public boolean isNotification() { + return id == null || id.isNull(); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/jsonrpc/JsonRpcResponse.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/jsonrpc/JsonRpcResponse.java new file mode 100644 index 0000000000..e01898bde6 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/jsonrpc/JsonRpcResponse.java @@ -0,0 +1,18 @@ +package stirling.software.proprietary.mcp.jsonrpc; + +import com.fasterxml.jackson.annotation.JsonInclude; + +import tools.jackson.databind.JsonNode; + +/** JSON-RPC 2.0 response; exactly one of {@code result} or {@code error} is non-null. */ +@JsonInclude(JsonInclude.Include.NON_NULL) +public record JsonRpcResponse(String jsonrpc, JsonNode id, Object result, JsonRpcError error) { + + public static JsonRpcResponse success(JsonNode id, Object result) { + return new JsonRpcResponse("2.0", id, result, null); + } + + public static JsonRpcResponse failure(JsonNode id, JsonRpcError error) { + return new JsonRpcResponse("2.0", id, null, error); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/security/McpApiKeyAuthFilter.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/security/McpApiKeyAuthFilter.java new file mode 100644 index 0000000000..e45dadb0c0 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/security/McpApiKeyAuthFilter.java @@ -0,0 +1,85 @@ +package stirling.software.proprietary.mcp.security; + +import java.io.IOException; +import java.util.List; +import java.util.Optional; + +import org.springframework.security.authentication.AnonymousAuthenticationToken; +import org.springframework.security.authentication.UsernamePasswordAuthenticationToken; +import org.springframework.security.core.Authentication; +import org.springframework.security.core.GrantedAuthority; +import org.springframework.security.core.authority.SimpleGrantedAuthority; +import org.springframework.security.core.context.SecurityContext; +import org.springframework.security.core.context.SecurityContextHolder; +import org.springframework.web.filter.OncePerRequestFilter; + +import jakarta.servlet.FilterChain; +import jakarta.servlet.ServletException; +import jakarta.servlet.http.HttpServletRequest; +import jakarta.servlet.http.HttpServletResponse; + +import lombok.extern.slf4j.Slf4j; + +import stirling.software.proprietary.security.model.User; +import stirling.software.proprietary.security.service.UserService; + +/** + * API-key auth for the MCP endpoint: validates a Stirling per-user API key and binds the request to + * that user with the MCP scopes. + */ +@Slf4j +public class McpApiKeyAuthFilter extends OncePerRequestFilter { + + private static final List MCP_SCOPES = + List.of( + new SimpleGrantedAuthority("SCOPE_mcp.tools.read"), + new SimpleGrantedAuthority("SCOPE_mcp.tools.write")); + + private final UserService userService; + + public McpApiKeyAuthFilter(UserService userService) { + this.userService = userService; + } + + @Override + protected void doFilterInternal( + HttpServletRequest request, HttpServletResponse response, FilterChain filterChain) + throws ServletException, IOException { + Authentication existing = SecurityContextHolder.getContext().getAuthentication(); + // Treat an anonymous token as not authenticated so the key is still processed. + boolean unauthenticated = + existing == null + || existing instanceof AnonymousAuthenticationToken + || !existing.isAuthenticated(); + if (unauthenticated) { + String apiKey = extractKey(request); + if (apiKey != null && !apiKey.isBlank()) { + Optional user = userService.getUserByApiKey(apiKey); + if (user.isPresent() && user.get().isEnabled()) { + UsernamePasswordAuthenticationToken auth = + new UsernamePasswordAuthenticationToken( + user.get().getUsername(), null, MCP_SCOPES); + SecurityContext context = SecurityContextHolder.createEmptyContext(); + context.setAuthentication(auth); + SecurityContextHolder.setContext(context); + } else { + log.warn( + "MCP access denied: presented API key did not match an active account"); + } + } + } + filterChain.doFilter(request, response); + } + + private String extractKey(HttpServletRequest request) { + String headerKey = request.getHeader("X-API-KEY"); + if (headerKey != null && !headerKey.isBlank()) { + return headerKey.trim(); + } + String authz = request.getHeader("Authorization"); + if (authz != null && authz.regionMatches(true, 0, "Bearer ", 0, 7)) { + return authz.substring(7).trim(); + } + return null; + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/security/McpAudienceValidator.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/security/McpAudienceValidator.java new file mode 100644 index 0000000000..49e16ebfe1 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/security/McpAudienceValidator.java @@ -0,0 +1,63 @@ +package stirling.software.proprietary.mcp.security; + +import java.util.Collection; +import java.util.LinkedHashSet; +import java.util.List; +import java.util.Set; + +import org.springframework.security.oauth2.core.OAuth2Error; +import org.springframework.security.oauth2.core.OAuth2TokenValidator; +import org.springframework.security.oauth2.core.OAuth2TokenValidatorResult; +import org.springframework.security.oauth2.jwt.Jwt; + +/** + * RFC 8707 audience binding: a JWT at the MCP endpoint must list this server's resource id (or one + * of the explicitly accepted additional audiences) in its {@code aud} claim. The additional list + * exists for IdPs that cannot mint resource-specific audiences - e.g. Supabase's OAuth server + * always issues {@code aud=authenticated}. Fails closed when nothing is configured. + */ +public class McpAudienceValidator implements OAuth2TokenValidator { + + private final Set acceptedAudiences; + + public McpAudienceValidator(String expectedResourceId) { + this(expectedResourceId, List.of()); + } + + public McpAudienceValidator(String expectedResourceId, Collection additionalAudiences) { + Set accepted = new LinkedHashSet<>(); + if (expectedResourceId != null && !expectedResourceId.isBlank()) { + accepted.add(expectedResourceId); + } + if (additionalAudiences != null) { + additionalAudiences.stream() + .filter(a -> a != null && !a.isBlank()) + .forEach(accepted::add); + } + this.acceptedAudiences = accepted; + } + + @Override + public OAuth2TokenValidatorResult validate(Jwt token) { + if (acceptedAudiences.isEmpty()) { + return OAuth2TokenValidatorResult.failure( + new OAuth2Error( + "invalid_token", + "MCP audience binding is not configured; rejecting all tokens until" + + " mcp.auth.resource-id or mcp.auth.accepted-audiences is set.", + null)); + } + List aud = token.getAudience(); + if (aud == null || aud.stream().noneMatch(acceptedAudiences::contains)) { + return OAuth2TokenValidatorResult.failure( + new OAuth2Error( + "invalid_token", + "Token audience does not include this server's resource id or an" + + " accepted audience (" + + String.join(", ", acceptedAudiences) + + ").", + null)); + } + return OAuth2TokenValidatorResult.success(); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/security/McpAuthenticationEntryPoint.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/security/McpAuthenticationEntryPoint.java new file mode 100644 index 0000000000..f8d23661d8 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/security/McpAuthenticationEntryPoint.java @@ -0,0 +1,114 @@ +package stirling.software.proprietary.mcp.security; + +import java.io.IOException; + +import org.springframework.http.HttpStatus; +import org.springframework.security.core.AuthenticationException; +import org.springframework.security.oauth2.core.OAuth2AuthenticationException; +import org.springframework.security.oauth2.core.OAuth2Error; +import org.springframework.security.web.AuthenticationEntryPoint; + +import jakarta.servlet.http.HttpServletRequest; +import jakarta.servlet.http.HttpServletResponse; + +import lombok.extern.slf4j.Slf4j; + +/** + * Emits 401 + {@code WWW-Authenticate: Bearer resource_metadata="..."} (RFC 9728) from + * X-Forwarded-* headers. A rejected token also logs the OAuth2 reason and echoes it as {@code + * error_description}. + */ +@Slf4j +public class McpAuthenticationEntryPoint implements AuthenticationEntryPoint { + + private final String metadataPath; + + public McpAuthenticationEntryPoint(String metadataPath) { + this.metadataPath = + metadataPath == null ? "/.well-known/oauth-protected-resource" : metadataPath; + } + + @Override + public void commence( + HttpServletRequest request, + HttpServletResponse response, + AuthenticationException authException) + throws IOException { + // Tokenless 401 is the normal discovery handshake; only a rejected token is a real failure. + boolean tokenPresented = request.getHeader("Authorization") != null; + String reason = rejectionReason(authException); + if (tokenPresented) { + log.warn("MCP rejected bearer token: {}", reason != null ? reason : "invalid_token"); + } else { + log.debug("MCP 401: no bearer token; returning protected-resource metadata pointer"); + } + + String scheme = firstForwarded(request, "X-Forwarded-Proto", request.getScheme()); + String authority = forwardedHost(request, scheme); + String metadataUrl = scheme + "://" + authority + metadataPath; + + StringBuilder header = new StringBuilder("Bearer error=\"invalid_token\""); + if (tokenPresented && reason != null) { + header.append(", error_description=\"").append(reason).append('"'); + } + header.append(", resource_metadata=\"").append(metadataUrl).append('"'); + response.setHeader("WWW-Authenticate", header.toString()); + response.sendError(HttpStatus.UNAUTHORIZED.value(), "Unauthorized"); + } + + /** + * OAuth2 error as {@code "code - description"}, sanitized for a header/log line; null if none. + */ + private static String rejectionReason(AuthenticationException ex) { + if (!(ex instanceof OAuth2AuthenticationException oae)) { + return null; + } + OAuth2Error error = oae.getError(); + if (error == null) { + return null; + } + String description = error.getDescription(); + String combined = + (description == null || description.isBlank()) + ? error.getErrorCode() + : error.getErrorCode() + " - " + description; + return combined == null ? null : combined.replaceAll("[\\r\\n\"]", " ").trim(); + } + + /** host[:port] from forwarded headers when present, else the servlet host/port. */ + private static String forwardedHost(HttpServletRequest request, String scheme) { + String host = firstForwarded(request, "X-Forwarded-Host", null); + if (host != null && !host.isBlank()) { + // X-Forwarded-Host may already carry a port. + if (host.contains(":")) { + return host; + } + String fwdPort = firstForwarded(request, "X-Forwarded-Port", null); + if (fwdPort != null && !isDefaultPort(scheme, fwdPort)) { + return host + ":" + fwdPort; + } + return host; + } + String authority = request.getServerName(); + int port = request.getServerPort(); + if (port > 0 && !isDefaultPort(scheme, Integer.toString(port))) { + authority = authority + ":" + port; + } + return authority; + } + + /** First (client-most) value of a possibly comma-listed forwarded header, trimmed. */ + private static String firstForwarded(HttpServletRequest request, String name, String fallback) { + String value = request.getHeader(name); + if (value == null || value.isBlank()) { + return fallback; + } + int comma = value.indexOf(','); + return (comma >= 0 ? value.substring(0, comma) : value).trim(); + } + + private static boolean isDefaultPort(String scheme, String port) { + return ("http".equals(scheme) && "80".equals(port)) + || ("https".equals(scheme) && "443".equals(port)); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/security/McpConfigValidator.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/security/McpConfigValidator.java new file mode 100644 index 0000000000..01be6d4355 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/security/McpConfigValidator.java @@ -0,0 +1,195 @@ +package stirling.software.proprietary.mcp.security; + +import java.util.ArrayList; +import java.util.List; +import java.util.Locale; + +import stirling.software.common.model.ApplicationProperties; + +/** + * Startup sanity-checks for MCP config; {@link McpSecurityConfig} logs the findings at boot so a + * misconfigured /mcp endpoint shows up in the logs instead of as a later rejected-token 401. + */ +public final class McpConfigValidator { + + public enum Severity { + WARN, + INFO + } + + public record Finding(Severity severity, String message) {} + + private McpConfigValidator() {} + + /** Inspect the resolved MCP config and return ordered findings (most actionable first). */ + public static List validate(ApplicationProperties.Mcp mcp) { + List findings = new ArrayList<>(); + ApplicationProperties.Mcp.Auth auth = mcp.getAuth(); + + if ("apikey".equalsIgnoreCase(auth.getMode())) { + findings.add( + info( + "auth mode = apikey - clients send a Stirling API key via X-API-KEY (or" + + " Authorization: Bearer ); no external IdP needed. The key" + + " must belong to a provisioned, enabled account (Account -> API" + + " Keys).")); + return findings; + } + + // Anything that isn't exactly "apikey" runs the OAuth chain (mirrors isApiKeyMode()). + String mode = auth.getMode(); + if (mode != null && !mode.isBlank() && !"oauth".equalsIgnoreCase(mode.trim())) { + findings.add( + warn( + "mcp.auth.mode='" + + mode + + "' is not recognized (expected 'oauth' or 'apikey'); it falls" + + " back to the OAuth chain, which rejects every token unless" + + " issuer-uri and resource-id are set. A near-miss like" + + " 'api-key' is NOT treated as API-key mode.")); + } + findings.add(info("auth mode = oauth - running as an OAuth2 resource server for /mcp.")); + + if (isBlank(auth.getIssuerUri())) { + findings.add( + warn( + "mcp.auth.issuer-uri is not set: the JWT decoder fails closed and rejects" + + " every token. Set it to your IdP issuer that publishes" + + " /.well-known/openid-configuration (e.g." + + " https://login.microsoftonline.com//v2.0).")); + } else if (!looksLikeUrl(auth.getIssuerUri())) { + findings.add( + warn( + "mcp.auth.issuer-uri='" + + auth.getIssuerUri() + + "' does not look like an http(s) URL.")); + } + + boolean hasResourceId = !isBlank(auth.getResourceId()); + boolean hasAcceptedAudiences = + auth.getAcceptedAudiences().stream().anyMatch(a -> !isBlank(a)); + + if (!hasResourceId && !hasAcceptedAudiences) { + findings.add( + warn( + "neither mcp.auth.resource-id nor mcp.auth.accepted-audiences is set: the" + + " audience validator fails closed and rejects every token (RFC" + + " 8707). Set resource-id to this server's public /mcp URL, or" + + " accepted-audiences to the audience your IdP actually mints.")); + } else { + if (hasResourceId && !looksLikeUrl(auth.getResourceId())) { + findings.add( + warn( + "mcp.auth.resource-id='" + + auth.getResourceId() + + "' is not an http(s) URL: the token aud must match it" + + " exactly (scheme, host and port included).")); + } else if (hasResourceId && !auth.getResourceId().endsWith("/mcp")) { + findings.add( + warn( + "mcp.auth.resource-id='" + + auth.getResourceId() + + "' does not end in /mcp: it must match the public URL" + + " clients call and the audience your IdP puts in the" + + " token.")); + } + if (hasAcceptedAudiences) { + findings.add( + info( + "mcp.auth.accepted-audiences=" + + auth.getAcceptedAudiences() + + " - tokens whose aud matches any of these are accepted, the" + + " escape hatch for IdPs that can't mint a resource-specific" + + " audience (e.g. an Entra ID app id, or Supabase's" + + " aud=authenticated).")); + } else { + findings.add( + info( + "audience binding is strict (token aud must equal" + + " mcp.auth.resource-id). If your IdP can't mint that - e.g." + + " Entra ID issues aud= - set" + + " mcp.auth.accepted-audiences to the audience it actually" + + " emits.")); + } + } + + if (isBlank(auth.getJwksUri())) { + findings.add( + info( + "mcp.auth.jwks-uri not set - signing keys are auto-discovered from the" + + " issuer's OpenID configuration.")); + } + + if ("sub".equalsIgnoreCase(auth.getUsernameClaim()) && auth.isRequireExistingAccount()) { + findings.add( + warn( + "mcp.auth.username-claim='sub' with require-existing-account=true: many" + + " IdPs (e.g. Entra ID, Google) set 'sub' to an opaque id that won't" + + " match a Stirling username. Set mcp.auth.username-claim to 'email'" + + " or 'preferred_username', or provision accounts keyed by sub.")); + } + + if (!auth.isRequireExistingAccount()) { + findings.add( + warn( + "mcp.auth.require-existing-account=false: any token your IdP signs can" + + " invoke MCP tools even if its subject has no Stirling account. Set" + + " it true unless you intend open access for every IdP-valid" + + " token.")); + } + + if (mcp.isScopesEnabled()) { + findings.add( + info( + "mcp.scopes-enabled=true - the IdP must mint 'mcp.tools.read' and" + + " 'mcp.tools.write' scopes or clients are rejected; set" + + " mcp.scopes-enabled=false if it can only issue coarse tokens.")); + } + + List allowed = mcp.getAllowedOperations(); + List blocked = mcp.getBlockedOperations(); + if (allowed != null && !allowed.isEmpty()) { + findings.add( + info( + "mcp.allowed-operations is a strict allow-list of " + + allowed.size() + + " operation(s); every other tool is hidden, so a wrong or" + + " typo'd id silently exposes nothing.")); + List shadowed = + blocked == null + ? List.of() + : allowed.stream().filter(blocked::contains).toList(); + if (!shadowed.isEmpty()) { + findings.add( + warn( + "mcp operation(s) " + + shadowed + + " are in both allowed-operations and blocked-operations;" + + " blocked wins, so they are hidden.")); + } + } + + if (findings.stream().noneMatch(f -> f.severity() == Severity.WARN)) { + findings.add(info("OAuth settings look complete.")); + } + + return findings; + } + + private static boolean isBlank(String value) { + return value == null || value.isBlank(); + } + + private static boolean looksLikeUrl(String value) { + String lower = value.toLowerCase(Locale.ROOT); + return lower.startsWith("http://") || lower.startsWith("https://"); + } + + private static Finding warn(String message) { + return new Finding(Severity.WARN, message); + } + + private static Finding info(String message) { + return new Finding(Severity.INFO, message); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/security/McpRequestSizeFilter.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/security/McpRequestSizeFilter.java new file mode 100644 index 0000000000..b6c410a30c --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/security/McpRequestSizeFilter.java @@ -0,0 +1,128 @@ +package stirling.software.proprietary.mcp.security; + +import java.io.BufferedReader; +import java.io.ByteArrayInputStream; +import java.io.ByteArrayOutputStream; +import java.io.IOException; +import java.io.InputStream; +import java.io.InputStreamReader; +import java.nio.charset.Charset; +import java.nio.charset.StandardCharsets; + +import org.springframework.web.filter.OncePerRequestFilter; + +import jakarta.servlet.FilterChain; +import jakarta.servlet.ReadListener; +import jakarta.servlet.ServletException; +import jakarta.servlet.ServletInputStream; +import jakarta.servlet.http.HttpServletRequest; +import jakarta.servlet.http.HttpServletRequestWrapper; +import jakarta.servlet.http.HttpServletResponse; + +/** + * Caps MCP request body size (via Content-Length and by buffering up to the cap) and rejects + * oversized bodies with a clean 413 before JSON parsing. + */ +public class McpRequestSizeFilter extends OncePerRequestFilter { + + private final long maxBodyBytes; + + public McpRequestSizeFilter(long maxBodyBytes) { + this.maxBodyBytes = maxBodyBytes > 0 ? maxBodyBytes : 256L * 1024L; + } + + @Override + protected void doFilterInternal( + HttpServletRequest request, HttpServletResponse response, FilterChain filterChain) + throws ServletException, IOException { + long declared = request.getContentLengthLong(); + if (declared > maxBodyBytes) { + tooLarge(response); + return; + } + byte[] body; + try { + body = readUpTo(request.getInputStream(), maxBodyBytes); + } catch (BodyTooLargeException e) { + tooLarge(response); + return; + } + filterChain.doFilter(new CachedBodyRequest(request, body), response); + } + + private static byte[] readUpTo(InputStream in, long max) throws IOException { + ByteArrayOutputStream buffer = new ByteArrayOutputStream(); + byte[] chunk = new byte[8192]; + long total = 0; + int n; + while ((n = in.read(chunk)) != -1) { + total += n; + if (total > max) { + throw new BodyTooLargeException(); + } + buffer.write(chunk, 0, n); + } + return buffer.toByteArray(); + } + + private void tooLarge(HttpServletResponse response) throws IOException { + response.setStatus(HttpServletResponse.SC_REQUEST_ENTITY_TOO_LARGE); + response.setContentType("application/json"); + response.getWriter() + .write( + "{\"error\":\"payload_too_large\",\"message\":\"MCP request body exceeds the" + + " configured limit of " + + maxBodyBytes + + " bytes.\"}"); + } + + private static final class BodyTooLargeException extends IOException {} + + /** Re-serves the buffered body to the controller. */ + private static final class CachedBodyRequest extends HttpServletRequestWrapper { + private final byte[] body; + + CachedBodyRequest(HttpServletRequest request, byte[] body) { + super(request); + this.body = body; + } + + @Override + public ServletInputStream getInputStream() { + ByteArrayInputStream source = new ByteArrayInputStream(body); + return new ServletInputStream() { + @Override + public int read() { + return source.read(); + } + + @Override + public int read(byte[] b, int off, int len) { + return source.read(b, off, len); + } + + @Override + public boolean isFinished() { + return source.available() == 0; + } + + @Override + public boolean isReady() { + return true; + } + + @Override + public void setReadListener(ReadListener readListener) { + // Synchronous buffered body; no async reads. + } + }; + } + + @Override + public BufferedReader getReader() { + String enc = getCharacterEncoding(); + Charset cs = enc == null ? StandardCharsets.UTF_8 : Charset.forName(enc); + return new BufferedReader(new InputStreamReader(new ByteArrayInputStream(body), cs)); + } + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/security/McpSecurityConfig.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/security/McpSecurityConfig.java new file mode 100644 index 0000000000..6ad49c1e1d --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/security/McpSecurityConfig.java @@ -0,0 +1,269 @@ +package stirling.software.proprietary.mcp.security; + +import java.util.ArrayList; +import java.util.Collection; +import java.util.List; + +import org.springframework.beans.factory.ObjectProvider; +import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty; +import org.springframework.context.annotation.Bean; +import org.springframework.context.annotation.Configuration; +import org.springframework.context.annotation.Lazy; +import org.springframework.core.Ordered; +import org.springframework.core.annotation.Order; +import org.springframework.core.convert.converter.Converter; +import org.springframework.http.HttpMethod; +import org.springframework.security.authentication.AbstractAuthenticationToken; +import org.springframework.security.config.annotation.web.builders.HttpSecurity; +import org.springframework.security.config.http.SessionCreationPolicy; +import org.springframework.security.core.GrantedAuthority; +import org.springframework.security.core.authority.SimpleGrantedAuthority; +import org.springframework.security.oauth2.core.DelegatingOAuth2TokenValidator; +import org.springframework.security.oauth2.core.OAuth2TokenValidator; +import org.springframework.security.oauth2.jwt.Jwt; +import org.springframework.security.oauth2.jwt.JwtDecoder; +import org.springframework.security.oauth2.jwt.JwtValidators; +import org.springframework.security.oauth2.jwt.NimbusJwtDecoder; +import org.springframework.security.oauth2.server.resource.OAuth2ProtectedResourceMetadata; +import org.springframework.security.oauth2.server.resource.authentication.JwtAuthenticationConverter; +import org.springframework.security.oauth2.server.resource.authentication.JwtGrantedAuthoritiesConverter; +import org.springframework.security.oauth2.server.resource.web.authentication.BearerTokenAuthenticationFilter; +import org.springframework.security.web.SecurityFilterChain; +import org.springframework.security.web.access.intercept.AuthorizationFilter; +import org.springframework.security.web.authentication.AnonymousAuthenticationFilter; +import org.springframework.web.cors.CorsConfigurationSource; + +import jakarta.annotation.PostConstruct; + +import lombok.extern.slf4j.Slf4j; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.proprietary.security.service.UserService; + +/** + * MCP security chain: validates JWTs (JWKS + RFC 8707 audience), maps scope claims to authorities, + * and fails closed when the issuer is unset. + */ +@Slf4j +@Configuration +@Order(Ordered.HIGHEST_PRECEDENCE) +@ConditionalOnProperty(name = "mcp.enabled", havingValue = "true") +public class McpSecurityConfig { + + private final ApplicationProperties applicationProperties; + private final UserService userService; + + // Reuse the app's CORS config; ObjectProvider so the chain still wires when no CORS bean + // exists. + private final ObjectProvider corsConfigurationSource; + + private static final String BASE_PATH = "/mcp"; + + public McpSecurityConfig( + ApplicationProperties applicationProperties, + @Lazy UserService userService, + ObjectProvider corsConfigurationSource) { + this.applicationProperties = applicationProperties; + this.userService = userService; + this.corsConfigurationSource = corsConfigurationSource; + } + + /** Enable CORS on the MCP chain using the app-wide source when available. */ + private void applyCors(HttpSecurity http) throws Exception { + CorsConfigurationSource source = corsConfigurationSource.getIfAvailable(); + if (source != null) { + http.cors(cors -> cors.configurationSource(source)); + } + } + + @PostConstruct + void validateConfigOnStartup() { + log.info("MCP server enabled - validating configuration:"); + for (McpConfigValidator.Finding finding : + McpConfigValidator.validate(applicationProperties.getMcp())) { + if (finding.severity() == McpConfigValidator.Severity.WARN) { + log.warn("MCP config: {}", finding.message()); + } else { + log.info("MCP config: {}", finding.message()); + } + } + } + + @Bean + @Order(0) + SecurityFilterChain mcpSecurityFilterChain(HttpSecurity http, JwtDecoder mcpJwtDecoder) + throws Exception { + ApplicationProperties.Mcp.Auth auth = applicationProperties.getMcp().getAuth(); + if (isApiKeyMode()) { + return apiKeyFilterChain(http); + } + return oauthFilterChain(http, mcpJwtDecoder, auth); + } + + private boolean isApiKeyMode() { + return "apikey".equalsIgnoreCase(applicationProperties.getMcp().getAuth().getMode()); + } + + /** + * API-key chain: a Stirling per-user API key is validated by {@link McpApiKeyAuthFilter}; + * otherwise 401. + */ + private SecurityFilterChain apiKeyFilterChain(HttpSecurity http) throws Exception { + applyCors(http); + http.securityMatcher(BASE_PATH, BASE_PATH + "/**") + // CSRF intentionally disabled: /mcp is a stateless JSON-RPC API authenticated by an + // out-of-band X-API-KEY header (or Authorization: Bearer ). No cookies, no + // session, no form submissions; a browser cannot trick a victim into sending the + // header cross-origin, so the CSRF attack model does not apply. CodeQL flags this + // generically; the SessionCreationPolicy.STATELESS below is the relevant guarantee. + .csrf(csrf -> csrf.disable()) + .sessionManagement(s -> s.sessionCreationPolicy(SessionCreationPolicy.STATELESS)) + .authorizeHttpRequests(a -> a.anyRequest().authenticated()) + .exceptionHandling( + e -> + e.authenticationEntryPoint( + (request, response, ex) -> { + response.setStatus(401); + response.setHeader( + "WWW-Authenticate", + "Bearer realm=\"Stirling MCP (API key)\""); + response.setContentType("application/json"); + response.getWriter() + .write( + "{\"error\":\"unauthorized\",\"message\":\"Provide a valid Stirling API key via the X-API-KEY header (or Authorization: Bearer ).\"}"); + })) + .addFilterBefore( + new McpRequestSizeFilter( + applicationProperties.getMcp().getMaxRequestBytes()), + AuthorizationFilter.class) + // Authenticate before the anonymous filter sets an anonymous token. + .addFilterBefore( + new McpApiKeyAuthFilter(userService), AnonymousAuthenticationFilter.class); + return http.build(); + } + + /** OAuth2 resource-server chain (JWT, RFC 8707 audience, RFC 9728 metadata). */ + private SecurityFilterChain oauthFilterChain( + HttpSecurity http, JwtDecoder mcpJwtDecoder, ApplicationProperties.Mcp.Auth auth) + throws Exception { + String metadataPath = "/.well-known/oauth-protected-resource"; + applyCors(http); + // RFC 9728 section 3.1: clients derive the metadata URL by inserting the well-known + // segment before the resource path, so /mcp is discovered at {metadataPath}/mcp. Claim + // the subpaths too; otherwise they fall through to another filter chain whose default + // Spring Security metadata filter serves a document without authorization_servers. + http.securityMatcher(BASE_PATH, BASE_PATH + "/**", metadataPath, metadataPath + "/**") + // CSRF intentionally disabled: /mcp is a stateless JSON-RPC resource server + // authenticated by OAuth2 Bearer JWTs (Authorization header). No cookies, no + // session, no form submissions; CSRF requires browser-attached ambient credentials + // and the bearer token is supplied per-request by the MCP client. CodeQL flags + // this generically; the SessionCreationPolicy.STATELESS below is the actual + // guarantee, and the .well-known metadata endpoint only serves GET. + .csrf(csrf -> csrf.disable()) + .sessionManagement(s -> s.sessionCreationPolicy(SessionCreationPolicy.STATELESS)) + .authorizeHttpRequests( + a -> + a.requestMatchers( + HttpMethod.GET, metadataPath, metadataPath + "/**") + .permitAll() + .anyRequest() + .authenticated()) + // Cap body size pre-auth, then bind the validated token to a Stirling user after + // the bearer filter. + .addFilterBefore( + new McpRequestSizeFilter( + applicationProperties.getMcp().getMaxRequestBytes()), + BearerTokenAuthenticationFilter.class) + .addFilterAfter( + new McpUserBindingFilter( + userService, + auth.getUsernameClaim(), + auth.isRequireExistingAccount()), + BearerTokenAuthenticationFilter.class) + .oauth2ResourceServer( + oauth2 -> + oauth2.authenticationEntryPoint( + // Advertise the path-inserted form; RFC 9728 makes + // it the canonical location for a resource with a + // path component. + new McpAuthenticationEntryPoint( + metadataPath + BASE_PATH)) + // RFC 9728 protected-resource metadata for OAuth discovery. + .protectedResourceMetadata( + prm -> + prm.protectedResourceMetadataCustomizer( + builder -> + buildResourceMetadata( + builder, auth))) + .jwt( + jwt -> + jwt.decoder(mcpJwtDecoder) + .jwtAuthenticationConverter( + mcpJwtAuthenticationConverter()))); + return http.build(); + } + + /** Populate the RFC 9728 protected-resource metadata document from the configured auth. */ + private void buildResourceMetadata( + OAuth2ProtectedResourceMetadata.Builder builder, ApplicationProperties.Mcp.Auth auth) { + if (!auth.getResourceId().isBlank()) { + builder.resource(auth.getResourceId()); + } + if (!auth.getIssuerUri().isBlank()) { + builder.authorizationServer(auth.getIssuerUri()); + } + // Only advertise the granular tool scopes when we actually enforce them. When scopes are + // disabled (e.g. the IdP only mints coarse tokens, like Supabase), advertising scopes the + // authorization server can't issue makes spec-compliant clients request them and get + // rejected with invalid_request. + if (applicationProperties.getMcp().isScopesEnabled()) { + builder.scope("mcp.tools.read"); + builder.scope("mcp.tools.write"); + } + } + + @Bean + JwtDecoder mcpJwtDecoder() { + ApplicationProperties.Mcp.Auth auth = applicationProperties.getMcp().getAuth(); + if (auth.getIssuerUri().isBlank()) { + // Fail-closed decoder: rejects every token until the issuer is set. + return token -> { + throw new org.springframework.security.oauth2.jwt.BadJwtException( + "mcp.auth.issuer-uri is not configured"); + }; + } + String jwksUri = auth.getJwksUri(); + NimbusJwtDecoder decoder = + jwksUri.isBlank() + ? NimbusJwtDecoder.withIssuerLocation(auth.getIssuerUri()).build() + : NimbusJwtDecoder.withJwkSetUri(jwksUri).build(); + OAuth2TokenValidator defaultValidators = + JwtValidators.createDefaultWithIssuer(auth.getIssuerUri()); + OAuth2TokenValidator combined = + new DelegatingOAuth2TokenValidator<>( + defaultValidators, + new McpAudienceValidator( + auth.getResourceId(), auth.getAcceptedAudiences())); + decoder.setJwtValidator(combined); + return decoder; + } + + private Converter mcpJwtAuthenticationConverter() { + JwtGrantedAuthoritiesConverter scopes = new JwtGrantedAuthoritiesConverter(); + scopes.setAuthorityPrefix("SCOPE_"); + scopes.setAuthoritiesClaimName("scope"); + JwtAuthenticationConverter converter = new JwtAuthenticationConverter(); + converter.setJwtGrantedAuthoritiesConverter( + jwt -> { + Collection out = new ArrayList<>(scopes.convert(jwt)); + List aud = jwt.getAudience(); + if (aud != null) { + for (String a : aud) { + out.add(new SimpleGrantedAuthority("AUDIENCE_" + a)); + } + } + return out; + }); + return converter; + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/security/McpUserBindingFilter.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/security/McpUserBindingFilter.java new file mode 100644 index 0000000000..59dfbbfacc --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/security/McpUserBindingFilter.java @@ -0,0 +1,115 @@ +package stirling.software.proprietary.mcp.security; + +import java.io.IOException; +import java.util.Optional; + +import org.springframework.security.authentication.UsernamePasswordAuthenticationToken; +import org.springframework.security.core.Authentication; +import org.springframework.security.core.context.SecurityContext; +import org.springframework.security.core.context.SecurityContextHolder; +import org.springframework.security.oauth2.jwt.Jwt; +import org.springframework.security.oauth2.server.resource.authentication.JwtAuthenticationToken; +import org.springframework.web.filter.OncePerRequestFilter; + +import jakarta.servlet.FilterChain; +import jakarta.servlet.ServletException; +import jakarta.servlet.http.HttpServletRequest; +import jakarta.servlet.http.HttpServletResponse; + +import lombok.extern.slf4j.Slf4j; + +import stirling.software.proprietary.security.model.User; +import stirling.software.proprietary.security.service.UserService; + +import tools.jackson.databind.ObjectMapper; +import tools.jackson.databind.node.ObjectNode; + +/** + * Binds an MCP-validated JWT to a provisioned Stirling user: optionally rejects subjects with no + * enabled account, then rebinds the principal to the canonical Stirling username (scope authorities + * only) so audit/metering attribute correctly. + */ +@Slf4j +public class McpUserBindingFilter extends OncePerRequestFilter { + + private static final ObjectMapper MAPPER = new ObjectMapper(); + + private final UserService userService; + private final String usernameClaim; + private final boolean requireExistingAccount; + + public McpUserBindingFilter( + UserService userService, String usernameClaim, boolean requireExistingAccount) { + this.userService = userService; + this.usernameClaim = + (usernameClaim == null || usernameClaim.isBlank()) ? "sub" : usernameClaim; + this.requireExistingAccount = requireExistingAccount; + } + + @Override + protected void doFilterInternal( + HttpServletRequest request, HttpServletResponse response, FilterChain filterChain) + throws ServletException, IOException { + Authentication current = SecurityContextHolder.getContext().getAuthentication(); + + // Only act on a JWT-authenticated request; everything else passes through. + if (current instanceof JwtAuthenticationToken jwtAuth && jwtAuth.isAuthenticated()) { + Jwt jwt = jwtAuth.getToken(); + String username = jwt.getClaimAsString(usernameClaim); + + if (username == null || username.isBlank()) { + reject( + response, + "Token is missing the '" + + usernameClaim + + "' claim used to map to a" + + " Stirling user."); + return; + } + + // Prefer the canonical username from the account record; fall back to the claim when + // binding is off. + String boundUsername = username; + if (requireExistingAccount) { + Optional account = userService.findByUsernameIgnoreCase(username); + if (account.isEmpty() || !account.get().isEnabled()) { + log.warn( + "MCP access denied: token subject '{}' has no active Stirling account", + sanitizeForLog(username)); + reject( + response, + "MCP access requires a provisioned, enabled Stirling account for this" + + " subject."); + return; + } + boundUsername = account.get().getUsername(); + } + + // Rebind to the Stirling username, carrying only the OAuth scope authorities. + UsernamePasswordAuthenticationToken bound = + new UsernamePasswordAuthenticationToken( + boundUsername, null, jwtAuth.getAuthorities()); + bound.setDetails(jwtAuth.getDetails()); + SecurityContext context = SecurityContextHolder.createEmptyContext(); + context.setAuthentication(bound); + SecurityContextHolder.setContext(context); + } + + filterChain.doFilter(request, response); + } + + /** Strip CR/LF so a crafted claim value can't forge log lines. */ + private static String sanitizeForLog(String value) { + return value == null ? null : value.replace('\r', ' ').replace('\n', ' '); + } + + private void reject(HttpServletResponse response, String message) throws IOException { + SecurityContextHolder.clearContext(); + response.setStatus(HttpServletResponse.SC_FORBIDDEN); + response.setContentType("application/json"); + ObjectNode body = MAPPER.createObjectNode(); + body.put("error", "insufficient_account"); + body.put("message", message); + response.getWriter().write(MAPPER.writeValueAsString(body)); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/AbstractCategoryTool.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/AbstractCategoryTool.java new file mode 100644 index 0000000000..8c75ee1459 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/AbstractCategoryTool.java @@ -0,0 +1,154 @@ +package stirling.software.proprietary.mcp.tools; + +import java.util.List; + +import org.springframework.beans.factory.ObjectProvider; + +import stirling.software.proprietary.mcp.McpCallContext; +import stirling.software.proprietary.mcp.McpTool; +import stirling.software.proprietary.mcp.catalog.McpToolCatalog; +import stirling.software.proprietary.mcp.catalog.OperationCategory; +import stirling.software.proprietary.mcp.catalog.OperationMeta; + +import tools.jackson.databind.JsonNode; +import tools.jackson.databind.ObjectMapper; +import tools.jackson.databind.node.ArrayNode; +import tools.jackson.databind.node.ObjectNode; + +/** + * Common scaffolding for the PDF category tools. Operation ids and summaries come from the live + * {@link McpToolCatalog}. + */ +abstract class AbstractCategoryTool implements McpTool { + + protected final ObjectMapper mapper; + protected final ObjectProvider catalogProvider; + protected final ObjectProvider executorProvider; + + protected AbstractCategoryTool( + ObjectMapper mapper, + ObjectProvider catalog, + ObjectProvider executor) { + this.mapper = mapper; + this.catalogProvider = catalog; + this.executorProvider = executor; + } + + protected abstract OperationCategory category(); + + protected List enabledOperations() { + McpToolCatalog catalog = catalogProvider.getIfAvailable(); + if (catalog == null) { + return List.of(); + } + return catalog.enabledOps(category()); + } + + @Override + public ObjectNode inputSchema() { + ObjectNode schema = mapper.createObjectNode(); + schema.put("type", "object"); + schema.put("additionalProperties", false); + + ObjectNode props = schema.putObject("properties"); + + ObjectNode op = props.putObject("operation"); + op.put("type", "string"); + List enabled = enabledOperations(); + StringBuilder opDesc = new StringBuilder(); + opDesc.append( + "Operation id from this category. Call stirling_describe_operation first to learn" + + " the exact parameters schema. Available operations:\n"); + ArrayNode opEnum = op.putArray("enum"); + for (OperationMeta m : enabled) { + opEnum.add(m.id()); + opDesc.append("- ").append(m.id()).append(" - ").append(m.summary()).append('\n'); + } + op.put("description", opDesc.toString().trim()); + + ObjectNode params = props.putObject("parameters"); + params.put("type", "object"); + params.put( + "description", + "Per-operation parameters. Schema available via stirling_describe_operation."); + params.put("additionalProperties", true); + + McpToolSupport.stringProperty( + props, + "file", + "Base64-encoded file content to process. The recommended way to provide a file for" + + " most uses. Bounded by the MCP request size limit; for very large files" + + " use 'fileId' instead."); + McpToolSupport.stringProperty( + props, + "fileName", + "Optional original filename (with extension) for the input; helps operations that" + + " key off file type."); + McpToolSupport.stringProperty( + props, + "fileId", + "Reference to a file already stored via stirling_upload. Recommended only for large" + + " files or multi-step workflows; most users should pass the file inline" + + " via 'file' instead."); + + ArrayNode required = schema.putArray("required"); + required.add("operation"); + return schema; + } + + @Override + public ObjectNode call(JsonNode arguments, McpCallContext context) { + JsonNode opNode = arguments == null ? null : arguments.get("operation"); + // No operation chosen: return this category's operation list. + if (opNode == null || !opNode.isTextual() || opNode.asText().isBlank()) { + return operationListError(null); + } + String opId = opNode.asText(); + McpToolCatalog catalog = catalogProvider.getIfAvailable(); + if (catalog == null) { + return McpResponses.error(mapper, "MCP catalog is not available"); + } + OperationMeta meta = catalog.findByOperationId(opId).orElse(null); + // Invalid/disabled/wrong-category op: return this category's operations. + if (meta == null || meta.category() != category()) { + return operationListError(opId); + } + if (!context.hasScope(meta.requiredScope())) { + return McpResponses.error( + mapper, + "Insufficient scope: this operation requires '" + meta.requiredScope() + "'."); + } + McpOperationExecutor executor = executorProvider.getIfAvailable(); + if (executor == null) { + return McpResponses.error(mapper, "MCP execution is not available."); + } + return executor.execute(meta, arguments); + } + + /** + * Error for a missing/unknown operation, listing this category's available operation ids and + * summaries. + */ + private ObjectNode operationListError(String badOpId) { + StringBuilder sb = new StringBuilder(); + if (badOpId == null) { + sb.append("Missing required argument 'operation' for ").append(category().toolName()); + } else { + sb.append("Unknown or disabled operation '") + .append(badOpId) + .append("' for ") + .append(category().toolName()); + } + List ops = enabledOperations(); + if (ops.isEmpty()) { + sb.append(". No operations are currently available in this category."); + } else { + sb.append(". Available operations:"); + for (OperationMeta m : ops) { + sb.append("\n- ").append(m.id()).append(" - ").append(m.summary()); + } + sb.append("\nRe-call this tool with a valid 'operation'."); + } + return McpResponses.error(mapper, sb.toString()); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/DescribeOperationTool.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/DescribeOperationTool.java new file mode 100644 index 0000000000..c93578136a --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/DescribeOperationTool.java @@ -0,0 +1,84 @@ +package stirling.software.proprietary.mcp.tools; + +import org.springframework.beans.factory.ObjectProvider; +import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty; +import org.springframework.stereotype.Component; + +import stirling.software.proprietary.mcp.McpCallContext; +import stirling.software.proprietary.mcp.McpTool; +import stirling.software.proprietary.mcp.catalog.McpToolCatalog; +import stirling.software.proprietary.mcp.catalog.OperationMeta; + +import tools.jackson.databind.JsonNode; +import tools.jackson.databind.ObjectMapper; +import tools.jackson.databind.node.ArrayNode; +import tools.jackson.databind.node.ObjectNode; + +/** Returns the JSON Schema for one operation's parameters, from the live {@link McpToolCatalog}. */ +@Component +@ConditionalOnProperty(name = "mcp.enabled", havingValue = "true") +public class DescribeOperationTool implements McpTool { + + private final ObjectMapper mapper; + private final ObjectProvider catalogProvider; + + public DescribeOperationTool(ObjectMapper mapper, ObjectProvider catalog) { + this.mapper = mapper; + this.catalogProvider = catalog; + } + + @Override + public String name() { + return "stirling_describe_operation"; + } + + @Override + public String description() { + return "Return the full JSON Schema for one Stirling operation's parameters. Call this " + + "before invoking a category tool to learn the exact shape of `parameters`. " + + "Argument: { operation: } where appears in the enum of any " + + "category tool (stirling_convert, _pages, _misc, _security, _ai)."; + } + + @Override + public ObjectNode inputSchema() { + ObjectNode schema = mapper.createObjectNode(); + schema.put("type", "object"); + schema.put("additionalProperties", false); + ObjectNode props = schema.putObject("properties"); + ObjectNode op = props.putObject("operation"); + op.put("type", "string"); + op.put( + "description", + "Operation id (e.g. compress-pdf, pdf-to-word, q-and-a). See category tool enums."); + ArrayNode required = schema.putArray("required"); + required.add("operation"); + return schema; + } + + @Override + public ObjectNode call(JsonNode arguments, McpCallContext context) { + JsonNode opNode = arguments == null ? null : arguments.get("operation"); + if (opNode == null || !opNode.isTextual() || opNode.asText().isBlank()) { + return McpResponses.error(mapper, "Missing required argument: operation"); + } + String opId = opNode.asText(); + McpToolCatalog catalog = catalogProvider.getIfAvailable(); + if (catalog == null) { + return McpResponses.error(mapper, "MCP catalog is not available"); + } + OperationMeta meta = catalog.findByOperationId(opId).orElse(null); + if (meta == null) { + return McpResponses.error(mapper, "Unknown or disabled operation: " + opId); + } + + ObjectNode payload = mapper.createObjectNode(); + payload.put("operation", meta.id()); + payload.put("category", meta.category().toolName()); + payload.put("summary", meta.summary()); + payload.put("endpoint", meta.endpointPath()); + payload.put("requiredScope", meta.requiredScope()); + payload.set("parametersSchema", meta.paramSchema()); + return McpResponses.json(mapper, payload); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/McpOperationExecutor.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/McpOperationExecutor.java new file mode 100644 index 0000000000..28cb56d965 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/McpOperationExecutor.java @@ -0,0 +1,255 @@ +package stirling.software.proprietary.mcp.tools; + +import java.io.IOException; +import java.io.InputStream; +import java.nio.charset.StandardCharsets; +import java.util.Base64; +import java.util.List; +import java.util.Map; + +import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty; +import org.springframework.core.io.ByteArrayResource; +import org.springframework.core.io.Resource; +import org.springframework.http.MediaType; +import org.springframework.http.ResponseEntity; +import org.springframework.stereotype.Component; +import org.springframework.util.LinkedMultiValueMap; +import org.springframework.util.MultiValueMap; +import org.springframework.web.client.RestClientResponseException; + +import lombok.extern.slf4j.Slf4j; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.common.service.FileStorage; +import stirling.software.common.service.InternalApiClient; +import stirling.software.common.service.InternalApiTimeoutException; +import stirling.software.proprietary.mcp.catalog.OperationMeta; + +import tools.jackson.core.type.TypeReference; +import tools.jackson.databind.JsonNode; +import tools.jackson.databind.ObjectMapper; +import tools.jackson.databind.node.ObjectNode; + +/** + * Runs a JAVA_ENDPOINT operation: resolves the input file (inline base64 or a fileId), dispatches + * to the Stirling endpoint over the loopback via {@link InternalApiClient}, and stores the result. + */ +@Slf4j +@Component +@ConditionalOnProperty(name = "mcp.enabled", havingValue = "true") +public class McpOperationExecutor { + + private final ObjectMapper mapper; + private final InternalApiClient internalApiClient; + private final FileStorage fileStorage; + private final ApplicationProperties applicationProperties; + + public McpOperationExecutor( + ObjectMapper mapper, + InternalApiClient internalApiClient, + FileStorage fileStorage, + ApplicationProperties applicationProperties) { + this.mapper = mapper; + this.internalApiClient = internalApiClient; + this.fileStorage = fileStorage; + this.applicationProperties = applicationProperties; + } + + public ObjectNode execute(OperationMeta meta, JsonNode arguments) { + String fileName = McpToolSupport.textArg(arguments, "fileName"); + String fileId = McpToolSupport.textArg(arguments, "fileId"); + byte[] inputBytes; + String inputName; + if (fileId != null) { + try { + if (!fileStorage.fileExists(fileId)) { + return McpResponses.error( + mapper, + "Unknown or inaccessible fileId '" + + fileId + + "'. Re-upload with stirling_upload."); + } + inputBytes = fileStorage.retrieveBytes(fileId); + } catch (SecurityException e) { + return McpResponses.error( + mapper, + "Unknown or inaccessible fileId '" + + fileId + + "'. Re-upload with stirling_upload."); + } catch (IOException e) { + return McpResponses.error(mapper, "Could not read fileId '" + fileId + "'."); + } + inputName = fileName != null ? fileName : fileId; + } else { + String base64 = McpToolSupport.textArg(arguments, "file"); + if (base64 == null) { + return McpResponses.error( + mapper, + "This operation needs an input file. Pass 'file' as base64 (recommended for" + + " most files), or 'fileId' from stirling_upload for large files."); + } + inputBytes = McpToolSupport.decodeBase64OrNull(base64); + if (inputBytes == null) { + return McpResponses.error(mapper, "The 'file' argument is not valid base64."); + } + inputName = fileName != null ? fileName : "input.pdf"; + } + + MultiValueMap body = new LinkedMultiValueMap<>(); + body.add("fileInput", bytesResource(inputBytes, inputName)); + addParameters(body, arguments == null ? null : arguments.get("parameters")); + + ResponseEntity response; + try { + response = internalApiClient.post(meta.endpointPath(), body); + } catch (InternalApiTimeoutException e) { + return McpResponses.error( + mapper, + meta.id() + + " timed out after " + + e.getReadTimeout().toSeconds() + + "s. Try a smaller file or a different approach."); + } catch (RestClientResponseException e) { + log.warn( + "MCP {} upstream error: HTTP {} - {}", + meta.id(), + e.getStatusCode().value(), + snippet(e.getResponseBodyAsString())); + return McpResponses.error( + mapper, meta.id() + " failed: HTTP " + e.getStatusCode().value() + "."); + } catch (SecurityException e) { + return McpResponses.error( + mapper, meta.id() + " endpoint is not permitted for MCP dispatch."); + } catch (RuntimeException e) { + log.warn("MCP execution of {} failed", meta.id(), e); + return McpResponses.error( + mapper, meta.id() + " failed unexpectedly. See server logs for details."); + } + return buildResult(meta, response); + } + + private ObjectNode buildResult(OperationMeta meta, ResponseEntity response) { + Resource body = response.getBody(); + if (body == null) { + return McpResponses.error(mapper, meta.id() + " returned an empty response."); + } + MediaType contentType = response.getHeaders().getContentType(); + + // A JSON body is a structured report (e.g. get-info), not a file. + if (contentType != null && MediaType.APPLICATION_JSON.isCompatibleWith(contentType)) { + try (InputStream is = body.getInputStream()) { + return McpResponses.text( + mapper, new String(is.readAllBytes(), StandardCharsets.UTF_8)); + } catch (IOException e) { + return McpResponses.error(mapper, "Failed to read " + meta.id() + " result."); + } + } + + String filename = + body.getFilename() == null || body.getFilename().isBlank() + ? meta.id() + : body.getFilename(); + String mimeType = + contentType != null + ? contentType.toString() + : MediaType.APPLICATION_OCTET_STREAM_VALUE; + long maxInline = applicationProperties.getMcp().getMaxInlineResponseBytes(); + try { + long size = body.contentLength(); + byte[] inline = null; + if (size >= 0 && size <= maxInline) { + try (InputStream is = body.getInputStream()) { + inline = is.readAllBytes(); + } + } + String fileId = + inline != null + ? fileStorage.storeBytes(inline, filename) + : storeStreamed(body, filename); + String summary = + meta.id() + + " succeeded. Result: " + + filename + + " (" + + size + + " bytes), fileId=" + + fileId + + ". "; + if (inline != null) { + return McpResponses.result( + mapper, + false, + McpResponses.textBlock( + mapper, summary + "The file is included inline below."), + McpResponses.resourceBlock( + mapper, + "stirling://file/" + fileId, + mimeType, + Base64.getEncoder().encodeToString(inline))); + } + return McpResponses.result( + mapper, + false, + McpResponses.textBlock( + mapper, + summary + + "Large result - fetch it with stirling_download {\"fileId\":\"" + + fileId + + "\"}, or pass this fileId to another operation.")); + } catch (IOException e) { + return McpResponses.error(mapper, "Failed to store " + meta.id() + " result."); + } + } + + private String storeStreamed(Resource body, String filename) throws IOException { + try (InputStream is = body.getInputStream()) { + return fileStorage.storeInputStream(is, filename).fileId(); + } + } + + private void addParameters(MultiValueMap body, JsonNode params) { + if (params == null || !params.isObject()) { + return; + } + Map map = + mapper.convertValue(params, new TypeReference>() {}); + for (Map.Entry entry : map.entrySet()) { + Object value = entry.getValue(); + if (value == null) { + continue; + } + if (value instanceof List list) { + if (containsStructured(list)) { + body.add(entry.getKey(), mapper.writeValueAsString(list)); + } else { + list.forEach(item -> body.add(entry.getKey(), item)); + } + } else if (value instanceof Map) { + body.add(entry.getKey(), mapper.writeValueAsString(value)); + } else { + body.add(entry.getKey(), value); + } + } + } + + private static boolean containsStructured(List list) { + return list.stream().anyMatch(item -> item instanceof Map || item instanceof List); + } + + private static Resource bytesResource(byte[] bytes, String filename) { + return new ByteArrayResource(bytes) { + @Override + public String getFilename() { + return filename; + } + }; + } + + private static String snippet(String body) { + if (body == null || body.isBlank()) { + return "(no body)"; + } + String trimmed = body.strip(); + return trimmed.length() > 300 ? trimmed.substring(0, 300) + "..." : trimmed; + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/McpResponses.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/McpResponses.java new file mode 100644 index 0000000000..57ee2cdbd6 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/McpResponses.java @@ -0,0 +1,80 @@ +package stirling.software.proprietary.mcp.tools; + +import tools.jackson.databind.ObjectMapper; +import tools.jackson.databind.node.ArrayNode; +import tools.jackson.databind.node.ObjectNode; + +/** Helpers for the MCP {@code CallToolResult} response shape. */ +public final class McpResponses { + + private McpResponses() {} + + /** Plain-text content block. */ + public static ObjectNode text(ObjectMapper mapper, String text) { + ObjectNode block = mapper.createObjectNode(); + block.put("type", "text"); + block.put("text", text); + return wrap(mapper, block, false); + } + + /** Plain-text error ({@code isError:true}). */ + public static ObjectNode error(ObjectMapper mapper, String message) { + ObjectNode block = mapper.createObjectNode(); + block.put("type", "text"); + block.put("text", message); + return wrap(mapper, block, true); + } + + /** JSON payload as embedded text. */ + public static ObjectNode json(ObjectMapper mapper, ObjectNode payload) { + ObjectNode block = mapper.createObjectNode(); + block.put("type", "text"); + block.put("text", payload.toString()); + return wrap(mapper, block, false); + } + + /** A text content block (unwrapped). */ + public static ObjectNode textBlock(ObjectMapper mapper, String text) { + ObjectNode block = mapper.createObjectNode(); + block.put("type", "text"); + block.put("text", text); + return block; + } + + /** An embedded-resource content block carrying base64 file content. */ + public static ObjectNode resourceBlock( + ObjectMapper mapper, String uri, String mimeType, String base64) { + ObjectNode block = mapper.createObjectNode(); + block.put("type", "resource"); + ObjectNode res = block.putObject("resource"); + res.put("uri", uri); + if (mimeType != null) { + res.put("mimeType", mimeType); + } + res.put("blob", base64); + return block; + } + + /** Build a result from explicit content blocks. */ + public static ObjectNode result(ObjectMapper mapper, boolean isError, ObjectNode... blocks) { + ObjectNode result = mapper.createObjectNode(); + ArrayNode content = result.putArray("content"); + for (ObjectNode b : blocks) { + content.add(b); + } + if (isError) { + result.put("isError", true); + } + return result; + } + + private static ObjectNode wrap(ObjectMapper mapper, ObjectNode block, boolean isError) { + ObjectNode result = mapper.createObjectNode(); + ArrayNode content = result.putArray("content"); + content.add(block); + if (isError) { + result.put("isError", true); + } + return result; + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/McpToolSupport.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/McpToolSupport.java new file mode 100644 index 0000000000..dae928261d --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/McpToolSupport.java @@ -0,0 +1,43 @@ +package stirling.software.proprietary.mcp.tools; + +import java.util.Base64; + +import tools.jackson.databind.JsonNode; +import tools.jackson.databind.node.ObjectNode; + +/** Shared helpers for MCP tools: argument parsing and JSON-Schema building. */ +final class McpToolSupport { + + private McpToolSupport() {} + + /** Trimmed text value of an argument, or null if absent, blank, or not a string. */ + static String textArg(JsonNode args, String field) { + if (args == null) { + return null; + } + JsonNode node = args.get(field); + if (node == null || !node.isTextual()) { + return null; + } + String value = node.asText().trim(); + return value.isEmpty() ? null : value; + } + + /** Decode base64 content, or null if the input is not valid base64. */ + static byte[] decodeBase64OrNull(String base64) { + try { + return Base64.getDecoder().decode(base64); + } catch (IllegalArgumentException e) { + return null; + } + } + + /** + * Add a {@code string} property with a description to a JSON-Schema {@code properties} node. + */ + static void stringProperty(ObjectNode properties, String name, String description) { + ObjectNode prop = properties.putObject(name); + prop.put("type", "string"); + prop.put("description", description); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/StirlingAiTool.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/StirlingAiTool.java new file mode 100644 index 0000000000..fc3696b658 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/StirlingAiTool.java @@ -0,0 +1,149 @@ +package stirling.software.proprietary.mcp.tools; + +import java.io.IOException; +import java.util.List; + +import org.springframework.beans.factory.ObjectProvider; +import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty; +import org.springframework.stereotype.Component; + +import lombok.extern.slf4j.Slf4j; + +import stirling.software.proprietary.mcp.McpCallContext; +import stirling.software.proprietary.mcp.McpTool; +import stirling.software.proprietary.mcp.catalog.McpToolCatalog; +import stirling.software.proprietary.mcp.catalog.OperationCategory; +import stirling.software.proprietary.mcp.catalog.OperationMeta; +import stirling.software.proprietary.service.AiEngineClient; + +import tools.jackson.databind.JsonNode; +import tools.jackson.databind.ObjectMapper; +import tools.jackson.databind.node.ArrayNode; +import tools.jackson.databind.node.ObjectNode; + +/** + * Exposes curated Python agent capabilities as a single MCP tool, sourced from the engine + * capabilities manifest. + */ +@Slf4j +@Component +@ConditionalOnProperty(name = "mcp.enabled", havingValue = "true") +public class StirlingAiTool implements McpTool { + + private final ObjectMapper mapper; + private final ObjectProvider catalogProvider; + private final ObjectProvider engineClientProvider; + + public StirlingAiTool( + ObjectMapper mapper, + ObjectProvider catalog, + ObjectProvider engineClient) { + this.mapper = mapper; + this.catalogProvider = catalog; + this.engineClientProvider = engineClient; + } + + @Override + public String name() { + return "stirling_ai"; + } + + @Override + public String description() { + return "Invoke a Stirling AI agent capability (Q&A about a PDF, edit-plan generation," + + " inline comments, math audit, draft-spec helper). Call" + + " stirling_describe_operation with the chosen capability id to get its" + + " parameters schema before invoking this tool. Some capabilities return content" + + " inline; others return a job reference that resolves to a file when ready."; + } + + @Override + public ObjectNode inputSchema() { + ObjectNode schema = mapper.createObjectNode(); + schema.put("type", "object"); + schema.put("additionalProperties", false); + ObjectNode props = schema.putObject("properties"); + + ObjectNode op = props.putObject("operation"); + op.put("type", "string"); + StringBuilder desc = new StringBuilder(); + desc.append("Capability id from the engine manifest. Available capabilities:\n"); + ArrayNode opEnum = op.putArray("enum"); + for (OperationMeta m : aiOps()) { + opEnum.add(m.id()); + desc.append("- ").append(m.id()).append(" - ").append(m.summary()).append('\n'); + } + op.put("description", desc.toString().trim()); + + ObjectNode params = props.putObject("parameters"); + params.put("type", "object"); + params.put("description", "Per-capability parameters."); + params.put("additionalProperties", true); + + ObjectNode fileId = props.putObject("fileId"); + fileId.put("type", "string"); + fileId.put( + "description", + "Reference to a previously-uploaded PDF in Stirling's job store. Required for" + + " capabilities that consume a document."); + + ArrayNode required = schema.putArray("required"); + required.add("operation"); + return schema; + } + + @Override + public ObjectNode call(JsonNode arguments, McpCallContext context) { + JsonNode opNode = arguments == null ? null : arguments.get("operation"); + if (opNode == null || !opNode.isTextual() || opNode.asText().isBlank()) { + return McpResponses.error(mapper, "Missing required argument: operation"); + } + String opId = opNode.asText(); + McpToolCatalog catalog = catalogProvider.getIfAvailable(); + if (catalog == null) { + return McpResponses.error(mapper, "MCP catalog is not available"); + } + OperationMeta meta = catalog.findByOperationId(opId).orElse(null); + if (meta == null || meta.category() != OperationCategory.AI) { + return McpResponses.error( + mapper, + "Unknown AI capability '" + + opId + + "'. The engine manifest may not be loaded yet - retry shortly or" + + " confirm the engine is reachable."); + } + if (!context.hasScope(meta.requiredScope())) { + return McpResponses.error( + mapper, + "Insufficient scope: this capability requires '" + meta.requiredScope() + "'."); + } + AiEngineClient client = engineClientProvider.getIfAvailable(); + if (client == null) { + return McpResponses.error( + mapper, "AI engine client is not configured - enable aiEngine in settings."); + } + if (meta.endpointPath() == null) { + return McpResponses.error( + mapper, + "Capability '" + opId + "' has no route configured in the engine manifest."); + } + JsonNode params = arguments.get("parameters"); + String body = (params == null ? mapper.createObjectNode() : params).toString(); + try { + String response = client.post(meta.endpointPath(), body, context.stirlingUserId()); + return McpResponses.text(mapper, response); + } catch (IOException e) { + log.warn("MCP AI capability '{}' engine request failed", opId, e); + return McpResponses.error( + mapper, "Engine request failed for capability '" + opId + "'."); + } + } + + private List aiOps() { + McpToolCatalog catalog = catalogProvider.getIfAvailable(); + if (catalog == null) { + return List.of(); + } + return catalog.enabledOps(OperationCategory.AI); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/StirlingConvertTool.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/StirlingConvertTool.java new file mode 100644 index 0000000000..9e89f6a5c9 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/StirlingConvertTool.java @@ -0,0 +1,41 @@ +package stirling.software.proprietary.mcp.tools; + +import org.springframework.beans.factory.ObjectProvider; +import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty; +import org.springframework.stereotype.Component; + +import stirling.software.proprietary.mcp.catalog.McpToolCatalog; +import stirling.software.proprietary.mcp.catalog.OperationCategory; + +import tools.jackson.databind.ObjectMapper; + +/** Exposes the {@code /api/v1/convert/*} namespace as a single MCP tool. */ +@Component +@ConditionalOnProperty(name = "mcp.enabled", havingValue = "true") +public class StirlingConvertTool extends AbstractCategoryTool { + + public StirlingConvertTool( + ObjectMapper mapper, + ObjectProvider catalog, + ObjectProvider executor) { + super(mapper, catalog, executor); + } + + @Override + public String name() { + return "stirling_convert"; + } + + @Override + public String description() { + return "Convert files between PDF and other formats (PDF<->Word, PDF<->image, HTML->PDF," + + " etc.). Inspect the `operation` enum, then call stirling_describe_operation" + + " with the chosen op to get its parameters JSON Schema before calling this" + + " tool."; + } + + @Override + protected OperationCategory category() { + return OperationCategory.CONVERT; + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/StirlingDownloadTool.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/StirlingDownloadTool.java new file mode 100644 index 0000000000..22fad23845 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/StirlingDownloadTool.java @@ -0,0 +1,113 @@ +package stirling.software.proprietary.mcp.tools; + +import java.io.IOException; +import java.util.Base64; + +import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty; +import org.springframework.http.MediaType; +import org.springframework.stereotype.Component; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.common.service.FileStorage; +import stirling.software.proprietary.mcp.McpCallContext; +import stirling.software.proprietary.mcp.McpTool; + +import tools.jackson.databind.JsonNode; +import tools.jackson.databind.ObjectMapper; +import tools.jackson.databind.node.ObjectNode; + +/** + * Fetches a stored file's content by fileId, returned inline as base64. For large results that were + * not returned inline by an operation. + */ +@Component +@ConditionalOnProperty(name = "mcp.enabled", havingValue = "true") +public class StirlingDownloadTool implements McpTool { + + private final ObjectMapper mapper; + private final FileStorage fileStorage; + private final ApplicationProperties applicationProperties; + + public StirlingDownloadTool( + ObjectMapper mapper, + FileStorage fileStorage, + ApplicationProperties applicationProperties) { + this.mapper = mapper; + this.fileStorage = fileStorage; + this.applicationProperties = applicationProperties; + } + + @Override + public String name() { + return "stirling_download"; + } + + @Override + public String description() { + return "Fetch a stored file's content by fileId (e.g. an operation result), returned inline" + + " as base64. Recommended only when a result was too large to be returned inline." + + " Argument: { fileId: }."; + } + + @Override + public ObjectNode inputSchema() { + ObjectNode schema = mapper.createObjectNode(); + schema.put("type", "object"); + schema.put("additionalProperties", false); + ObjectNode props = schema.putObject("properties"); + McpToolSupport.stringProperty( + props, "fileId", "Id of a stored file (e.g. an operation result's fileId)."); + schema.putArray("required").add("fileId"); + return schema; + } + + @Override + public ObjectNode call(JsonNode arguments, McpCallContext context) { + if (!context.hasScope("mcp.tools.read")) { + return McpResponses.error( + mapper, "Insufficient scope: stirling_download requires 'mcp.tools.read'."); + } + String fileId = McpToolSupport.textArg(arguments, "fileId"); + if (fileId == null) { + return McpResponses.error(mapper, "Missing required argument: fileId."); + } + long maxInline = applicationProperties.getMcp().getMaxInlineResponseBytes(); + try { + if (!fileStorage.fileExists(fileId)) { + return McpResponses.error( + mapper, "Unknown or inaccessible fileId '" + fileId + "'."); + } + long size = fileStorage.getFileSize(fileId); + if (size > maxInline) { + return McpResponses.error( + mapper, + "File is " + + size + + " bytes, over the inline limit of " + + maxInline + + " bytes. Raise mcp.maxInlineResponseBytes or retrieve it via the" + + " Stirling UI/API."); + } + byte[] bytes = fileStorage.retrieveBytes(fileId); + return McpResponses.result( + mapper, + false, + McpResponses.textBlock( + mapper, + "File " + + fileId + + " (" + + bytes.length + + " bytes) included inline below."), + McpResponses.resourceBlock( + mapper, + "stirling://file/" + fileId, + MediaType.APPLICATION_OCTET_STREAM_VALUE, + Base64.getEncoder().encodeToString(bytes))); + } catch (SecurityException e) { + return McpResponses.error(mapper, "Unknown or inaccessible fileId '" + fileId + "'."); + } catch (IOException e) { + return McpResponses.error(mapper, "Failed to read fileId '" + fileId + "'."); + } + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/StirlingMiscTool.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/StirlingMiscTool.java new file mode 100644 index 0000000000..3d76f47a7f --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/StirlingMiscTool.java @@ -0,0 +1,40 @@ +package stirling.software.proprietary.mcp.tools; + +import org.springframework.beans.factory.ObjectProvider; +import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty; +import org.springframework.stereotype.Component; + +import stirling.software.proprietary.mcp.catalog.McpToolCatalog; +import stirling.software.proprietary.mcp.catalog.OperationCategory; + +import tools.jackson.databind.ObjectMapper; + +/** Exposes the {@code /api/v1/misc/*} namespace as a single MCP tool. */ +@Component +@ConditionalOnProperty(name = "mcp.enabled", havingValue = "true") +public class StirlingMiscTool extends AbstractCategoryTool { + + public StirlingMiscTool( + ObjectMapper mapper, + ObjectProvider catalog, + ObjectProvider executor) { + super(mapper, catalog, executor); + } + + @Override + public String name() { + return "stirling_misc"; + } + + @Override + public String description() { + return "Miscellaneous PDF operations: compress, OCR, stamp / watermark, edit metadata," + + " flatten, repair, and similar utilities. Call stirling_describe_operation with" + + " the chosen op to get its parameters schema before invoking this tool."; + } + + @Override + protected OperationCategory category() { + return OperationCategory.MISC; + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/StirlingPagesTool.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/StirlingPagesTool.java new file mode 100644 index 0000000000..a736a78127 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/StirlingPagesTool.java @@ -0,0 +1,40 @@ +package stirling.software.proprietary.mcp.tools; + +import org.springframework.beans.factory.ObjectProvider; +import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty; +import org.springframework.stereotype.Component; + +import stirling.software.proprietary.mcp.catalog.McpToolCatalog; +import stirling.software.proprietary.mcp.catalog.OperationCategory; + +import tools.jackson.databind.ObjectMapper; + +/** Exposes the {@code /api/v1/general/*} (page operations) namespace as a single MCP tool. */ +@Component +@ConditionalOnProperty(name = "mcp.enabled", havingValue = "true") +public class StirlingPagesTool extends AbstractCategoryTool { + + public StirlingPagesTool( + ObjectMapper mapper, + ObjectProvider catalog, + ObjectProvider executor) { + super(mapper, catalog, executor); + } + + @Override + public String name() { + return "stirling_pages"; + } + + @Override + public String description() { + return "Manipulate PDF pages: merge, split, rotate, rearrange, crop, delete, overlay," + + " add blank pages. Call stirling_describe_operation with the chosen op to get" + + " its parameters schema before invoking this tool."; + } + + @Override + protected OperationCategory category() { + return OperationCategory.PAGES; + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/StirlingSecurityTool.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/StirlingSecurityTool.java new file mode 100644 index 0000000000..300b6c1d09 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/StirlingSecurityTool.java @@ -0,0 +1,41 @@ +package stirling.software.proprietary.mcp.tools; + +import org.springframework.beans.factory.ObjectProvider; +import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty; +import org.springframework.stereotype.Component; + +import stirling.software.proprietary.mcp.catalog.McpToolCatalog; +import stirling.software.proprietary.mcp.catalog.OperationCategory; + +import tools.jackson.databind.ObjectMapper; + +/** Exposes the {@code /api/v1/security/*} namespace as a single MCP tool. */ +@Component +@ConditionalOnProperty(name = "mcp.enabled", havingValue = "true") +public class StirlingSecurityTool extends AbstractCategoryTool { + + public StirlingSecurityTool( + ObjectMapper mapper, + ObjectProvider catalog, + ObjectProvider executor) { + super(mapper, catalog, executor); + } + + @Override + public String name() { + return "stirling_security"; + } + + @Override + public String description() { + return "Security-related PDF operations: password add/remove, redact, sanitize, certify" + + " / sign with cert, validate signature, add watermark. Call" + + " stirling_describe_operation with the chosen op to get its parameters schema" + + " before invoking this tool."; + } + + @Override + protected OperationCategory category() { + return OperationCategory.SECURITY; + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/StirlingUploadTool.java b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/StirlingUploadTool.java new file mode 100644 index 0000000000..ebee4e4122 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/mcp/tools/StirlingUploadTool.java @@ -0,0 +1,96 @@ +package stirling.software.proprietary.mcp.tools; + +import java.io.IOException; + +import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty; +import org.springframework.stereotype.Component; + +import lombok.extern.slf4j.Slf4j; + +import stirling.software.common.service.FileStorage; +import stirling.software.proprietary.mcp.McpCallContext; +import stirling.software.proprietary.mcp.McpTool; + +import tools.jackson.databind.JsonNode; +import tools.jackson.databind.ObjectMapper; +import tools.jackson.databind.node.ObjectNode; + +/** + * Stores a file server-side and returns a fileId. For large files or multi-step workflows only - + * most operations accept the file inline via their {@code file} argument. + */ +@Slf4j +@Component +@ConditionalOnProperty(name = "mcp.enabled", havingValue = "true") +public class StirlingUploadTool implements McpTool { + + private final ObjectMapper mapper; + private final FileStorage fileStorage; + + public StirlingUploadTool(ObjectMapper mapper, FileStorage fileStorage) { + this.mapper = mapper; + this.fileStorage = fileStorage; + } + + @Override + public String name() { + return "stirling_upload"; + } + + @Override + public String description() { + return "Store a file server-side and get back a fileId to reuse across operations." + + " Recommended only for large files or multi-step workflows; for a single" + + " operation on a typical file, pass the file inline via the operation's `file`" + + " argument instead. Argument: { file: , fileName?: }."; + } + + @Override + public ObjectNode inputSchema() { + ObjectNode schema = mapper.createObjectNode(); + schema.put("type", "object"); + schema.put("additionalProperties", false); + ObjectNode props = schema.putObject("properties"); + McpToolSupport.stringProperty(props, "file", "Base64-encoded file content."); + McpToolSupport.stringProperty( + props, "fileName", "Optional original filename (with extension)."); + schema.putArray("required").add("file"); + return schema; + } + + @Override + public ObjectNode call(JsonNode arguments, McpCallContext context) { + if (!context.hasScope("mcp.tools.write")) { + return McpResponses.error( + mapper, "Insufficient scope: stirling_upload requires 'mcp.tools.write'."); + } + String base64 = McpToolSupport.textArg(arguments, "file"); + if (base64 == null) { + return McpResponses.error( + mapper, "Missing required argument: file (base64-encoded content)."); + } + byte[] bytes = McpToolSupport.decodeBase64OrNull(base64); + if (bytes == null) { + return McpResponses.error(mapper, "The 'file' argument is not valid base64."); + } + String name = McpToolSupport.textArg(arguments, "fileName"); + if (name == null) { + name = "upload.bin"; + } + try { + String fileId = fileStorage.storeBytes(bytes, name); + return McpResponses.text( + mapper, + "Stored '" + + name + + "' (" + + bytes.length + + " bytes) as fileId=" + + fileId + + ". Pass this fileId to a Stirling operation's 'fileId' argument."); + } catch (IOException e) { + log.warn("MCP upload failed to store file", e); + return McpResponses.error(mapper, "Failed to store the uploaded file."); + } + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/model/api/ai/AiWorkflowOutcome.java b/app/proprietary/src/main/java/stirling/software/proprietary/model/api/ai/AiWorkflowOutcome.java index a7239e8f90..2bed56f0f0 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/model/api/ai/AiWorkflowOutcome.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/model/api/ai/AiWorkflowOutcome.java @@ -21,7 +21,8 @@ public enum AiWorkflowOutcome { COMPLETED("completed"), UNSUPPORTED_CAPABILITY("unsupported_capability"), CANNOT_CONTINUE("cannot_continue"), - GENERATE_FILE("generate_file"); + GENERATE_FILE("generate_file"), + CONVERT_MARKDOWN("convert_markdown"); private final String value; diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/model/api/ai/AiWorkflowRequest.java b/app/proprietary/src/main/java/stirling/software/proprietary/model/api/ai/AiWorkflowRequest.java index e574f84b86..ad4c994c05 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/model/api/ai/AiWorkflowRequest.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/model/api/ai/AiWorkflowRequest.java @@ -6,7 +6,6 @@ import java.util.List; import io.swagger.v3.oas.annotations.media.Schema; import jakarta.validation.constraints.NotBlank; -import jakarta.validation.constraints.NotNull; import lombok.Data; @@ -14,9 +13,8 @@ import lombok.Data; @Schema(description = "Run an AI workflow") public class AiWorkflowRequest { - @NotNull @Schema(description = "The input PDF files") - private List fileInputs; + private List fileInputs = new ArrayList<>(); @NotBlank @Schema(description = "The user message to orchestrate", example = "Summarise these documents") diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/model/api/ai/AiWorkflowResponse.java b/app/proprietary/src/main/java/stirling/software/proprietary/model/api/ai/AiWorkflowResponse.java index 8f0fc631bc..2f088f4444 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/model/api/ai/AiWorkflowResponse.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/model/api/ai/AiWorkflowResponse.java @@ -105,4 +105,19 @@ public class AiWorkflowResponse { + " body or via the X-Stirling-Tool-Report header. May be null for tools" + " that produce only a file.") private JsonNode report; + + @Schema( + description = + "Structured error code when a downstream tool call was blocked (e.g." + + " PAYG_LIMIT_REACHED). Lets the client react — such as opening the" + + " usage-limit modal — instead of only seeing a generic failure. Null" + + " for ordinary outcomes.") + private String errorCode; + + @Schema( + description = + "Whether the team is subscribed, carried from a downstream usage-limit response." + + " Selects which limit modal the client shows (free → subscribe," + + " subscribed → raise cap). Null when the downstream body omitted it.") + private Boolean errorSubscribed; } diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/config/AdminPolicyManagementAuthority.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/config/AdminPolicyManagementAuthority.java new file mode 100644 index 0000000000..a49c8e5aa9 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/config/AdminPolicyManagementAuthority.java @@ -0,0 +1,41 @@ +package stirling.software.proprietary.policy.config; + +import org.springframework.context.annotation.Profile; +import org.springframework.stereotype.Component; + +import lombok.RequiredArgsConstructor; + +import stirling.software.proprietary.model.Team; +import stirling.software.proprietary.security.model.User; +import stirling.software.proprietary.security.service.UserService; + +/** + * Default (non-SaaS) policy context: a global admin may edit policies; scoping uses the current + * user's team (typically a single shared team self-hosted). SaaS overrides this with a team-leader + * check (see the {@code saas}-profiled implementation). + */ +@Component +@Profile("!saas") +@RequiredArgsConstructor +public class AdminPolicyManagementAuthority implements PolicyManagementAuthority { + + private final UserService userService; + + @Override + public boolean canEditPolicies() { + return userService.isCurrentUserAdmin(); + } + + @Override + public Long currentUserTeamId() { + String username = userService.getCurrentUsername(); + if (username == null) { + return null; + } + return userService + .findByUsername(username) + .map(User::getTeam) + .map(Team::getId) + .orElse(null); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/config/FolderAccessGuard.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/config/FolderAccessGuard.java new file mode 100644 index 0000000000..899302f5b0 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/config/FolderAccessGuard.java @@ -0,0 +1,94 @@ +package stirling.software.proprietary.policy.config; + +import java.nio.file.Path; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.List; + +import org.springframework.context.annotation.Profile; +import org.springframework.core.env.Environment; +import org.springframework.stereotype.Component; + +import stirling.software.common.configuration.InstallationPathConfig; +import stirling.software.common.model.ApplicationProperties; +import stirling.software.proprietary.policy.model.Policy; + +/** + * Authority on which filesystem locations a policy may read/write. Checked at save time and again + * at run time, fail-closed in order: + * + *

    + *
  1. denied entirely under the {@code saas} profile; + *
  2. Stirling's own config dir always rejected, even if an allowed root were misconfigured to + * contain it; + *
  3. must resolve within {@code policies.allowedFolderRoots}; none configured means all denied. + *
+ * + *

Compared after normalisation so {@code ..} cannot escape a root. Symlink escape is not + * defended: an operator who roots an allowlist on a symlink to a sensitive location is trusted. + */ +@Component +@Profile("saas") +public class FolderAccessGuard { + + public static final String FOLDER_TYPE = "folder"; + + private final boolean saasActive; + private final List allowedRoots; + private final List protectedRoots; + + public FolderAccessGuard(ApplicationProperties applicationProperties, Environment environment) { + this.saasActive = Arrays.asList(environment.getActiveProfiles()).contains("saas"); + this.allowedRoots = + normalizeAll(applicationProperties.getPolicies().getAllowedFolderRoots()); + this.protectedRoots = List.of(normalize(Path.of(InstallationPathConfig.getConfigPath()))); + } + + /** Returns the normalised absolute path; throws if not permitted. */ + public Path requirePermitted(Path dir) { + if (saasActive) { + throw new IllegalArgumentException( + "folder sources and outputs are not available in SaaS mode"); + } + Path normalized = normalize(dir); + for (Path protectedRoot : protectedRoots) { + if (normalized.startsWith(protectedRoot)) { + throw new IllegalArgumentException( + "folder may not point inside a protected Stirling directory"); + } + } + if (allowedRoots.isEmpty()) { + throw new IllegalArgumentException( + "folder access is disabled; set policies.allowedFolderRoots to permit it"); + } + boolean within = allowedRoots.stream().anyMatch(normalized::startsWith); + if (!within) { + throw new IllegalArgumentException( + "folder '" + normalized + "' is outside the allowed folder roots"); + } + return normalized; + } + + /** Whether this policy touches a folder source/sink, and so is subject to these rules. */ + public boolean usesFolderAccess(Policy policy) { + boolean readsFolder = + policy.sources().stream().anyMatch(spec -> FOLDER_TYPE.equals(spec.type())); + boolean writesFolder = + policy.output() != null && FOLDER_TYPE.equals(policy.output().type()); + return readsFolder || writesFolder; + } + + private static List normalizeAll(List roots) { + List result = new ArrayList<>(); + for (String root : roots) { + if (root != null && !root.isBlank()) { + result.add(normalize(Path.of(root))); + } + } + return result; + } + + private static Path normalize(Path path) { + return path.toAbsolutePath().normalize(); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/config/PolicyAccessGuard.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/config/PolicyAccessGuard.java new file mode 100644 index 0000000000..c7b10b941f --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/config/PolicyAccessGuard.java @@ -0,0 +1,62 @@ +package stirling.software.proprietary.policy.config; + +import java.util.List; +import java.util.Objects; + +import org.springframework.context.annotation.Profile; +import org.springframework.stereotype.Component; + +import lombok.RequiredArgsConstructor; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.common.service.UserServiceInterface; +import stirling.software.proprietary.policy.model.Policy; + +/** + * Policies are scoped to a team: a user may view, run, edit, and delete only the policies belonging + * to their own team (the team a policy is stamped with at creation). This binds everyone — admins + * included — so no one sees or touches another team's policies. Whether a user may edit + * (vs only view/run) is a separate check gated at the controller ({@code + * PolicyController#requirePolicyEditingAllowed} → team leader). Enforced only when login is + * enabled; single-user deployments (login disabled) pass every check. + */ +@Component +@RequiredArgsConstructor +@Profile("saas") +public class PolicyAccessGuard { + + private final UserServiceInterface userService; + private final ApplicationProperties applicationProperties; + private final PolicyManagementAuthority policyManagementAuthority; + + /** Owner for a new policy: the current user, or {@code null} when login is disabled. */ + public String ownerForNewPolicy() { + return enforced() ? userService.getCurrentUsername() : null; + } + + /** Team a new policy is stamped with — the creator's team. {@code null} when login disabled. */ + public Long teamForNewPolicy() { + return enforced() ? policyManagementAuthority.currentUserTeamId() : null; + } + + /** Whether the policy belongs to the current user's team (so they may view/run/edit it). */ + public boolean canAccess(Policy policy) { + if (!enforced()) { + return true; + } + return Objects.equals(policy.teamId(), policyManagementAuthority.currentUserTeamId()); + } + + /** The subset of {@code policies} scoped to the current user's team. */ + public List visible(List policies) { + if (!enforced()) { + return policies; + } + Long teamId = policyManagementAuthority.currentUserTeamId(); + return policies.stream().filter(policy -> Objects.equals(policy.teamId(), teamId)).toList(); + } + + private boolean enforced() { + return applicationProperties.getSecurity().isEnableLogin(); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/config/PolicyManagementAuthority.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/config/PolicyManagementAuthority.java new file mode 100644 index 0000000000..0ea3c298ad --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/config/PolicyManagementAuthority.java @@ -0,0 +1,22 @@ +package stirling.software.proprietary.policy.config; + +/** + * The current user's policy context, pluggable per deployment so the policy layer (proprietary) + * needn't know the team model. SaaS: a user may edit policies only if they lead their team, and + * every user is scoped to their own team. Self-hosted: a global admin may edit, scoped to their + * (typically single) team. Policies are isolated per team — nobody, admins included, sees or edits + * another team's policies. + */ +public interface PolicyManagementAuthority { + + /** Whether the current user may create, edit, or delete policies (for their own team). */ + boolean canEditPolicies(); + + /** + * The team that scopes the current user's policies — the team a new policy is stamped with and + * the only team whose policies the user may see/run/edit. {@code null} when it can't be + * resolved (e.g. login disabled / no team), in which case access falls back to the unteamed + * ({@code null}-team) policies. + */ + Long currentUserTeamId(); +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/controller/NamedAsset.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/controller/NamedAsset.java new file mode 100644 index 0000000000..6886a92946 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/controller/NamedAsset.java @@ -0,0 +1,29 @@ +package stirling.software.proprietary.policy.controller; + +import org.springframework.web.multipart.MultipartFile; + +import io.swagger.v3.oas.annotations.media.Schema; + +import jakarta.validation.constraints.NotBlank; +import jakarta.validation.constraints.NotNull; + +import lombok.Data; + +/** + * A supporting file paired with the asset key a pipeline step references from its {@code + * fileParameters}. The same key may appear on more than one asset to supply multiple files. + */ +@Data +@Schema(description = "A supporting file bound to the asset key a pipeline step references") +public class NamedAsset { + + @NotBlank + @Schema( + description = "Asset key referenced by a step's fileParameters", + example = "company-logo") + private String key; + + @NotNull + @Schema(description = "The supporting file", format = "binary") + private MultipartFile file; +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/controller/PolicyController.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/controller/PolicyController.java new file mode 100644 index 0000000000..f2b751bca7 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/controller/PolicyController.java @@ -0,0 +1,410 @@ +package stirling.software.proprietary.policy.controller; + +import java.io.IOException; +import java.util.ArrayList; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; + +import org.springframework.context.annotation.Profile; +import org.springframework.core.io.FileSystemResource; +import org.springframework.core.io.Resource; +import org.springframework.http.HttpStatus; +import org.springframework.http.MediaType; +import org.springframework.http.ResponseEntity; +import org.springframework.web.bind.annotation.DeleteMapping; +import org.springframework.web.bind.annotation.GetMapping; +import org.springframework.web.bind.annotation.ModelAttribute; +import org.springframework.web.bind.annotation.PathVariable; +import org.springframework.web.bind.annotation.PostMapping; +import org.springframework.web.bind.annotation.RequestBody; +import org.springframework.web.bind.annotation.RequestMapping; +import org.springframework.web.bind.annotation.RequestPart; +import org.springframework.web.bind.annotation.RestController; +import org.springframework.web.multipart.MultipartFile; +import org.springframework.web.server.ResponseStatusException; +import org.springframework.web.servlet.mvc.method.annotation.SseEmitter; + +import io.github.pixee.security.Filenames; +import io.swagger.v3.oas.annotations.Hidden; +import io.swagger.v3.oas.annotations.Operation; +import io.swagger.v3.oas.annotations.tags.Tag; + +import jakarta.validation.Valid; + +import lombok.RequiredArgsConstructor; +import lombok.extern.slf4j.Slf4j; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.common.model.job.JobResponse; +import stirling.software.common.service.JobOwnershipService; +import stirling.software.common.util.TempFile; +import stirling.software.common.util.TempFileManager; +import stirling.software.proprietary.policy.config.PolicyAccessGuard; +import stirling.software.proprietary.policy.config.PolicyManagementAuthority; +import stirling.software.proprietary.policy.engine.PolicyRunHandle; +import stirling.software.proprietary.policy.engine.PolicyRunRegistry; +import stirling.software.proprietary.policy.engine.PolicyRunner; +import stirling.software.proprietary.policy.engine.PolicyValidator; +import stirling.software.proprietary.policy.model.PipelineDefinition; +import stirling.software.proprietary.policy.model.Policy; +import stirling.software.proprietary.policy.model.PolicyInputs; +import stirling.software.proprietary.policy.model.PolicyRun; +import stirling.software.proprietary.policy.model.PolicyRunStatus; +import stirling.software.proprietary.policy.model.PolicyRunView; +import stirling.software.proprietary.policy.progress.PolicyProgressListener; +import stirling.software.proprietary.policy.store.PolicyStore; + +/** + * Policy CRUD plus pipeline runs (stored or ad-hoc). Runs are async: returns a run id, poll {@code + * GET /run/{runId}} for status, download outputs via {@code GET /api/v1/general/files/{fileId}}. + */ +@Slf4j +@RestController +@RequestMapping("/api/v1/policies") +@Hidden +@RequiredArgsConstructor +@Tag(name = "Policies", description = "Run tool pipelines on the backend") +@Profile("saas") +public class PolicyController { + + private final PolicyRunner policyRunner; + private final PolicyRunRegistry runRegistry; + private final PolicyStore policyStore; + private final PolicyValidator policyValidator; + private final PolicyAccessGuard policyAccessGuard; + private final PolicyManagementAuthority policyManagementAuthority; + private final ApplicationProperties applicationProperties; + private final TempFileManager tempFileManager; + private final JobOwnershipService jobOwnershipService; + + @PostMapping(value = "/run", consumes = MediaType.MULTIPART_FORM_DATA_VALUE) + @Operation( + summary = "Run a tool pipeline", + description = + "Accepts the documents to process (multipart field 'fileInput'), any supporting" + + " files (under 'assets[i].key' / 'assets[i].file'), and the pipeline" + + " definition as an application/json part named 'json'. Runs the steps" + + " in order asynchronously and returns a run id. Poll the run status" + + " endpoint and download outputs via /api/v1/general/files/{id}.") + public ResponseEntity> run( + @RequestPart("json") PipelineDefinition definition, + @Valid @ModelAttribute PolicyRunFiles files) + throws IOException { + requireRunnable(definition); + PolicyInputs inputs = toInputs(files); + String runId = + policyRunner.runAdHoc(definition, inputs, PolicyProgressListener.NOOP).runId(); + return ResponseEntity.accepted().body(new JobResponse<>(true, runId, null)); + } + + @PostMapping(value = "/run/stream", consumes = MediaType.MULTIPART_FORM_DATA_VALUE) + @Operation( + summary = "Run a tool pipeline with live progress", + description = + "Same as /run, but returns Server-Sent Events: a 'step' event as each step" + + " starts and completes, then a terminal 'completed', 'failed'," + + " 'cancelled', or 'waiting' event carrying the final run view.") + public SseEmitter runStream( + @RequestPart("json") PipelineDefinition definition, + @Valid @ModelAttribute PolicyRunFiles files) + throws IOException { + requireRunnable(definition); + PolicyInputs inputs = toInputs(files); + + SseEmitter emitter = + new SseEmitter(applicationProperties.getPolicies().getStreamTimeoutMs()); + emitter.onError(e -> log.warn("Policy run SSE emitter error", e)); + + PolicyRunHandle handle = policyRunner.runAdHoc(definition, inputs, streamListener(emitter)); + // whenComplete runs on the worker thread after the run finishes, so the terminal event + // never races the step events. + handle.completion() + .whenComplete( + (run, throwable) -> { + if (throwable != null) { + sendEvent( + emitter, + "failed", + Map.of("message", throwable.getMessage())); + } else { + sendEvent(emitter, terminalEventName(run), PolicyRunView.of(run)); + } + emitter.complete(); + }); + return emitter; + } + + @GetMapping("/run/{runId}") + @Operation( + summary = "Get pipeline run status", + description = "Returns the current status, step cursor, and output files of a run.") + public ResponseEntity status(@PathVariable String runId) { + PolicyRun run = runRegistry.get(runId); + if (run == null) { + return ResponseEntity.notFound().build(); + } + return ResponseEntity.ok(PolicyRunView.of(run)); + } + + @GetMapping("/runs") + @Operation( + summary = "List the caller's stored-policy runs", + description = + "Returns the caller's in-flight and recently-finished stored-policy runs (within" + + " the run-retention window). The frontend reconciles these on load so a" + + " run started before a refresh/crash is rediscovered and its outputs" + + " collected, rather than orphaned on the backend. Ad-hoc runs (no" + + " policy id) are excluded.") + public List listRuns() { + return runRegistry.all().stream() + .filter(run -> run.getPolicyId() != null) + .filter(run -> ownedByCurrentUser(run.getRunId())) + .map(PolicyRunView::of) + .toList(); + } + + /** + * Whether the run is owned by the current user, derived purely from the existing scoping + * methods: stripping then re-applying the scope reproduces the run's key only when its owner + * prefix matches the caller's. No auth (single-user) owns everything. Avoids duplicating the + * scoped-key format here. + */ + private boolean ownedByCurrentUser(String runId) { + return jobOwnershipService + .createScopedJobKey(jobOwnershipService.extractJobId(runId)) + .equals(runId); + } + + // --- Policy management --- + + @PostMapping(consumes = MediaType.APPLICATION_JSON_VALUE) + @Operation( + summary = "Create or update a policy", + description = + "Stores a policy (trigger config + steps + output + metadata). A blank id is" + + " assigned; returns the stored policy with its id.") + public ResponseEntity savePolicy(@RequestBody Policy policy) { + requirePolicyEditingAllowed(); + Policy owned = resolveOwnership(policy); + try { + policyValidator.validate(owned); + } catch (IllegalArgumentException e) { + throw new ResponseStatusException(HttpStatus.BAD_REQUEST, e.getMessage()); + } + return ResponseEntity.ok(policyStore.save(owned)); + } + + /** + * Assign owner + owning team server-side. Create stamps the current user and their team; update + * preserves the existing owner and team after verifying the policy belongs to the caller's team + * — so the client can neither forge ownership/team on create nor reach across teams on update + * (a policy in another team reads as not-found). + */ + private Policy resolveOwnership(Policy incoming) { + String id = incoming.id(); + if (id != null && !id.isBlank()) { + Policy existing = policyStore.get(id).orElse(null); + if (existing != null) { + if (!policyAccessGuard.canAccess(existing)) { + throw new ResponseStatusException(HttpStatus.NOT_FOUND, "No policy: " + id); + } + return withOwnerAndTeam(incoming, existing.owner(), existing.teamId()); + } + } + return withOwnerAndTeam( + incoming, + policyAccessGuard.ownerForNewPolicy(), + policyAccessGuard.teamForNewPolicy()); + } + + private static Policy withOwnerAndTeam(Policy policy, String owner, Long teamId) { + return new Policy( + policy.id(), + policy.name(), + owner, + policy.enabled(), + policy.trigger(), + policy.sources(), + policy.steps(), + policy.output(), + teamId); + } + + /** + * Creating, editing, pausing/resuming, and deleting policies requires the editor role for the + * caller's team — a team leader on SaaS (see {@link PolicyManagementAuthority}); the global + * admin gets no say on SaaS. Team scoping (which team's policies) is enforced separately by + * {@link PolicyAccessGuard}. Every mutation routes through {@link #savePolicy} (pause/resume + * re-save with a flipped {@code enabled} flag) or {@link #deletePolicy}, so gating those two + * covers them all; runs ({@code /run}) stay open to the team. Single-user deployments (login + * disabled) have no such role, so they trust the local operator. The path allowlist for folder + * sources/outputs is enforced separately by {@link PolicyValidator} at validation time. + */ + private void requirePolicyEditingAllowed() { + if (!applicationProperties.getSecurity().isEnableLogin()) { + return; + } + if (!policyManagementAuthority.canEditPolicies()) { + throw new ResponseStatusException( + HttpStatus.FORBIDDEN, + "Policies may only be created or modified by a team leader"); + } + } + + @GetMapping + @Operation( + summary = "List policies", + description = "Lists the policies belonging to the caller's team.") + public List listPolicies() { + return policyAccessGuard.visible(policyStore.all()); + } + + @GetMapping("/{policyId}") + @Operation(summary = "Get a policy by id") + public ResponseEntity getPolicy(@PathVariable String policyId) { + return policyStore + .get(policyId) + .filter(policyAccessGuard::canAccess) + .map(ResponseEntity::ok) + .orElseGet(() -> ResponseEntity.notFound().build()); + } + + @DeleteMapping("/{policyId}") + @Operation(summary = "Delete a policy by id") + public ResponseEntity deletePolicy(@PathVariable String policyId) { + requirePolicyEditingAllowed(); + // Scope to the caller's team: a policy in another team reads as not-found. + boolean accessible = + policyStore.get(policyId).filter(policyAccessGuard::canAccess).isPresent(); + if (accessible && policyStore.delete(policyId)) { + return ResponseEntity.noContent().build(); + } + return ResponseEntity.notFound().build(); + } + + @PostMapping(value = "/{policyId}/run", consumes = MediaType.MULTIPART_FORM_DATA_VALUE) + @Operation( + summary = "Run a stored policy", + description = + "Runs the stored policy's pipeline on the supplied files (primary documents" + + " under 'fileInput', supporting files under 'assets[i].key' /" + + " 'assets[i].file'). Runs regardless of the policy's enabled flag," + + " which only gates automatic triggering. Returns a run id.") + public ResponseEntity> runStoredPolicy( + @PathVariable String policyId, @Valid @ModelAttribute PolicyRunFiles files) + throws IOException { + Policy policy = + policyStore + .get(policyId) + .filter(policyAccessGuard::canAccess) + .orElseThrow( + () -> + new ResponseStatusException( + HttpStatus.NOT_FOUND, "No policy: " + policyId)); + PolicyInputs inputs = toInputs(files); + String runId = policyRunner.runWith(policy, inputs, PolicyProgressListener.NOOP).runId(); + return ResponseEntity.accepted().body(new JobResponse<>(true, runId, null)); + } + + private static void requireRunnable(PipelineDefinition definition) { + if (definition.steps().isEmpty()) { + throw new ResponseStatusException( + HttpStatus.BAD_REQUEST, "Pipeline definition has no steps"); + } + } + + /** + * Turn the typed run files into engine {@link PolicyInputs}: the primary documents plus the + * named supporting-file store, where each asset's {@code key} is the name a step references + * from its {@code fileParameters}. Assets sharing a key are grouped, so a key may carry several + * files. + */ + private PolicyInputs toInputs(PolicyRunFiles files) throws IOException { + List primary = toResources(files.getFileInput()); + Map> supportingFiles = new LinkedHashMap<>(); + for (NamedAsset asset : files.getAssets()) { + Resource resource = toResource(asset.getFile()); + if (resource != null) { + supportingFiles + .computeIfAbsent(asset.getKey(), key -> new ArrayList<>()) + .add(resource); + } + } + return new PolicyInputs(primary, supportingFiles); + } + + private PolicyProgressListener streamListener(SseEmitter emitter) { + return new PolicyProgressListener() { + @Override + public void onStepStart(int stepIndex, int stepCount, String operation) { + sendEvent(emitter, "step", stepEvent("started", stepIndex, stepCount, operation)); + } + + @Override + public void onStepComplete(int stepIndex, int stepCount, String operation) { + sendEvent(emitter, "step", stepEvent("completed", stepIndex, stepCount, operation)); + } + }; + } + + private static Map stepEvent( + String phase, int stepIndex, int stepCount, String operation) { + return Map.of( + "phase", phase, + "stepIndex", stepIndex, + "stepCount", stepCount, + "operation", operation); + } + + private static String terminalEventName(PolicyRun run) { + PolicyRunStatus status = run.getStatus(); + return switch (status) { + case COMPLETED -> "completed"; + case FAILED -> "failed"; + case CANCELLED -> "cancelled"; + case WAITING_FOR_INPUT -> "waiting"; + default -> "ended"; + }; + } + + private void sendEvent(SseEmitter emitter, String name, Object data) { + try { + emitter.send(SseEmitter.event().name(name).data(data, MediaType.APPLICATION_JSON)); + } catch (IOException | IllegalStateException e) { + // Client gone or emitter closed. The run continues and outputs stay downloadable via + // the job endpoints. + log.debug("Dropping policy SSE event '{}': {}", name, e.getMessage()); + } + } + + private List toResources(List files) throws IOException { + List resources = new ArrayList<>(); + if (files == null) { + return resources; + } + for (MultipartFile file : files) { + Resource resource = toResource(file); + if (resource != null) { + resources.add(resource); + } + } + return resources; + } + + /** Spool a single uploaded file to a managed temp file, preserving its name; null if empty. */ + private Resource toResource(MultipartFile file) throws IOException { + if (file == null || file.isEmpty()) { + return null; + } + TempFile tempFile = tempFileManager.createManagedTempFile("policy-run"); + file.transferTo(tempFile.getPath()); + final String originalName = Filenames.toSimpleFileName(file.getOriginalFilename()); + return new FileSystemResource(tempFile.getFile()) { + @Override + public String getFilename() { + return originalName; + } + }; + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/controller/PolicyRunFiles.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/controller/PolicyRunFiles.java new file mode 100644 index 0000000000..f42e02a07d --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/controller/PolicyRunFiles.java @@ -0,0 +1,32 @@ +package stirling.software.proprietary.policy.controller; + +import java.util.ArrayList; +import java.util.List; + +import org.springframework.web.multipart.MultipartFile; + +import io.swagger.v3.oas.annotations.media.Schema; + +import jakarta.validation.Valid; + +import lombok.Data; + +/** + * The files supplied to a policy run: the primary documents and any keyed supporting assets. Bound + * from the multipart request via {@code @ModelAttribute}; the pipeline definition itself travels as + * a separate typed {@code json} part. + * + *

Wire form: {@code fileInput} (repeated) for primaries, and {@code assets[i].key} / {@code + * assets[i].file} for each supporting asset. + */ +@Data +@Schema(description = "Files for a policy run: primary documents plus keyed supporting assets") +public class PolicyRunFiles { + + @Schema(description = "Primary input documents", format = "binary") + private List fileInput = new ArrayList<>(); + + @Valid + @Schema(description = "Supporting files, each bound to the asset key its step references") + private List assets = new ArrayList<>(); +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/engine/PolicyEngine.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/engine/PolicyEngine.java new file mode 100644 index 0000000000..28951d52a2 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/engine/PolicyEngine.java @@ -0,0 +1,384 @@ +package stirling.software.proprietary.policy.engine; + +import java.io.IOException; +import java.io.InputStream; +import java.util.ArrayList; +import java.util.List; +import java.util.UUID; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.ExecutorService; + +import org.slf4j.MDC; +import org.springframework.context.annotation.Profile; +import org.springframework.core.io.Resource; +import org.springframework.http.ResponseEntity; +import org.springframework.security.core.Authentication; +import org.springframework.security.core.context.SecurityContextHolder; +import org.springframework.stereotype.Service; +import org.springframework.web.client.RestClientResponseException; + +import lombok.RequiredArgsConstructor; +import lombok.extern.slf4j.Slf4j; + +import stirling.software.common.model.job.ResultFile; +import stirling.software.common.service.FileStorage; +import stirling.software.common.service.InternalApiTimeoutException; +import stirling.software.common.service.JobOwnershipService; +import stirling.software.common.service.JobQueue; +import stirling.software.common.service.ResourceMonitor; +import stirling.software.common.service.TaskManager; +import stirling.software.common.util.ExecutorFactory; +import stirling.software.common.util.JobContext; +import stirling.software.proprietary.policy.model.OutputSpec; +import stirling.software.proprietary.policy.model.PipelineDefinition; +import stirling.software.proprietary.policy.model.Policy; +import stirling.software.proprietary.policy.model.PolicyInputs; +import stirling.software.proprietary.policy.model.PolicyRun; +import stirling.software.proprietary.policy.model.WaitState; +import stirling.software.proprietary.policy.output.PolicyOutputSink; +import stirling.software.proprietary.policy.progress.PolicyProgressListener; +import stirling.software.proprietary.service.DownstreamEntitlementError; + +/** + * Runs pipelines asynchronously as tracked jobs. {@link #submit} returns a run id immediately; the + * pipeline runs on a virtual thread (so a step blocked on a slow tool does not hold a platform + * thread). Drives {@link PolicyExecutor} for the step loop, projects status/outputs into {@link + * TaskManager} (existing job endpoints work unchanged), and keeps live state in {@link + * PolicyRunRegistry}. + * + *

Manages its own virtual-thread execution rather than {@code JobExecutorService}, which + * force-completes a job once its work returns: incompatible with a run that suspends in {@code + * WAITING_FOR_INPUT}. Still applies the shared {@link ResourceMonitor}/{@link JobQueue} admission + * control so heavy runs queue under load. + */ +@Slf4j +@Service +@RequiredArgsConstructor +@Profile("saas") +public class PolicyEngine { + + // Admission weight for one run. Weighted heavy: a run chains many tools and holds intermediate + // files. See ResourceMonitor#shouldQueueJob(int). + private static final int RUN_RESOURCE_WEIGHT = 50; + + // errorCode marking a run that was never admitted (job queue full under load). Transient: the + // client treats it as "busy" and retries, rather than as a terminal processing failure. + private static final String QUEUE_FULL_CODE = "POLICY_QUEUE_FULL"; + + private final PolicyExecutor stepExecutor; + private final TaskManager taskManager; + private final PolicyRunRegistry registry; + private final FileStorage fileStorage; + private final JobOwnershipService jobOwnershipService; + private final List outputSinks; + private final ResourceMonitor resourceMonitor; + private final JobQueue jobQueue; + + private final ExecutorService asyncExecutor = ExecutorFactory.newVirtualThreadExecutor(); + + /** + * Submit a pipeline to run asynchronously. The handle's run id scopes a {@link TaskManager} job + * (status/notes/results observable via the job endpoints); its future resolves when the run + * reaches a terminal or paused state. + */ + public PolicyRunHandle submit( + PipelineDefinition definition, PolicyInputs inputs, PolicyProgressListener listener) { + return submit(definition, inputs, listener, null); + } + + /** + * As {@link #submit(PipelineDefinition, PolicyInputs, PolicyProgressListener)}, recording the + * originating stored policy's id on the run ({@code null} for ad-hoc pipelines). The id lets a + * client attribute a run it rediscovers via {@code GET /policies/runs} after losing local state + * (e.g. a refresh before it recorded the run), so a finished run is never orphaned server-side. + */ + public PolicyRunHandle submit( + PipelineDefinition definition, + PolicyInputs inputs, + PolicyProgressListener listener, + String policyId) { + // Ad-hoc run (no stored policy): bill whoever kicked it off and own the outputs as them + // too. + // Capture the principal on this (request) thread — it does not survive the hop onto the + // async + // worker. + String principal = currentActingPrincipal(); + return submitForPrincipal(principal, principal, policyId, definition, inputs, listener); + } + + /** Run a stored policy on demand. {@code enabled} gates triggers, not explicit runs. */ + public PolicyRunHandle runPolicy( + Policy policy, PolicyInputs inputs, PolicyProgressListener listener) { + // Bill the policy owner: trigger-fired runs have no security context, and the async worker + // doesn't inherit the caller's, so the owner (stamped at policy creation) is the reliable + // billing identity — and for org-wide policies the org/owner is meant to pay. But own the + // OUTPUT files as the user who triggered the run (captured here on the request thread) so + // they can download their enforced file; otherwise an org-wide policy's output is owned by + // the admin and the triggering user is denied it. Trigger-fired runs have no such user, so + // the owner owns those outputs. + String triggeringUser = currentActingPrincipal(); + String fileOwner = triggeringUser != null ? triggeringUser : policy.owner(); + return submitForPrincipal( + policy.owner(), fileOwner, policy.id(), policy.toDefinition(), inputs, listener); + } + + private PolicyRunHandle submitForPrincipal( + String billingPrincipal, + String fileOwner, + String policyId, + PipelineDefinition definition, + PolicyInputs inputs, + PolicyProgressListener listener) { + // Scope the run id to the current user (this request thread) so the file-download + // ownership check passes. No-op when security is off. + String runId = jobOwnershipService.createScopedJobKey(UUID.randomUUID().toString()); + taskManager.createTask(runId); + PolicyRun run = new PolicyRun(runId, policyId, definition); + registry.register(run); + CompletableFuture completion = new CompletableFuture<>(); + PolicyProgressListener tracking = trackingListener(runId, run, listener); + // Re-establish the acting principal as the audit principal on the worker thread. Each tool + // step dispatches via InternalApiClient, which resolves the caller from + // UserService.getCurrentUsername() — that has an MDC `auditPrincipal` fallback for async + // threads. Without this the worker has no identity, tool calls fall back to the + // INTERNAL_API_USER, and PAYG charges that system account instead of the owner's team. + Runnable task = + () -> + runAsPrincipal( + billingPrincipal, + fileOwner, + () -> runToCompletion(run, inputs, tracking, completion)); + + // One admission unit per run; steps run synchronously within it, so this gates heavy work + // without the pool-within-pool risk of queueing each tool call. + if (resourceMonitor.shouldQueueJob(RUN_RESOURCE_WEIGHT)) { + log.debug("Queueing policy run {} under resource pressure", runId); + jobQueue.queueJob( + runId, + RUN_RESOURCE_WEIGHT, + () -> { + task.run(); + return null; + }, + 0L) + .exceptionally(ex -> failRejectedRun(run, completion, ex)); + } else { + asyncExecutor.execute(task); + } + return new PolicyRunHandle(runId, completion); + } + + public PolicyRun getRun(String runId) { + return registry.get(runId); + } + + /** + * Mark a run cancelled if not already finished. Does not yet interrupt an in-flight tool call. + */ + public boolean cancel(String runId) { + PolicyRun run = registry.get(runId); + if (run == null) { + return false; + } + boolean cancelled = run.cancel(); + if (cancelled) { + taskManager.addNote(runId, "Run cancelled by request"); + } + return cancelled; + } + + /** Resume a run paused in {@code WAITING_FOR_INPUT}. Not yet implemented. */ + public String resume(String runId, List additionalInputs) { + throw new UnsupportedOperationException("Pause/resume is not yet implemented"); + } + + private void runToCompletion( + PolicyRun run, + PolicyInputs inputs, + PolicyProgressListener listener, + CompletableFuture completion) { + String runId = run.getRunId(); + try { + run.markRunning(); + PolicyExecutionResult result = + stepExecutor.execute(run.getDefinition(), inputs, listener); + OutputSpec output = run.getDefinition().output(); + List outputs = sinkFor(output).deliver(runId, result.files(), output); + taskManager.setMultipleFileResults(runId, outputs); + taskManager.setComplete(runId); + run.complete(outputs); + } catch (PolicyInputRequiredException e) { + // Expected path: suspend rather than fail. Persist intermediates as fileIds so the run + // can resume after this worker thread is gone. + WaitState wait = suspend(e); + run.waitForInput(wait); + taskManager.addNote(runId, "Waiting for input: " + e.getMessage()); + } catch (InternalApiTimeoutException e) { + String message = toolTimeoutMessage(e); + log.error( + "Policy run {} timed out on {}: {}", + runId, + e.getEndpointPath(), + e.getMessage()); + run.fail(message); + taskManager.setError(runId, message); + } catch (RestClientResponseException e) { + // A downstream tool call returned an error status. When it's a structured entitlement + // response (401/402 with a JSON `error` sentinel), surface that code onto the run so + // the + // client can react — e.g. pop the usage-limit modal — instead of only seeing a generic + // failure. We don't interpret the code here (that would couple this module to the saas + // billing layer); we just pass it through for the client to map. Other statuses fall + // through to the generic failure below. + String code = DownstreamEntitlementError.extractCode(e); + if (code != null) { + log.info("Policy run {} blocked by downstream entitlement gate ({})", runId, code); + String message = "Usage limit reached"; + run.failWithCode(message, code, DownstreamEntitlementError.extractSubscribed(e)); + taskManager.setError(runId, message); + } else { + String message = "Policy run failed: " + e.getMessage(); + log.error("Policy run {} failed (downstream HTTP error)", runId, e); + run.fail(message); + taskManager.setError(runId, message); + } + } catch (Exception e) { + String message = "Policy run failed: " + e.getMessage(); + log.error("Policy run {} failed", runId, e); + run.fail(message); + taskManager.setError(runId, message); + } finally { + // Always resolve so stream/await callers unblock. + completion.complete(run); + } + } + + private ResponseEntity failRejectedRun( + PolicyRun run, CompletableFuture completion, Throwable ex) { + // Only reached if the run never started (e.g. queue full); a started run resolves its own + // completion in runToCompletion. + if (!completion.isDone()) { + String message = "Policy run could not be queued: " + ex.getMessage(); + log.error("Policy run {} was not admitted: {}", run.getRunId(), ex.getMessage()); + // Transient admission rejection, not a processing failure (see QUEUE_FULL_CODE). + run.failWithCode(message, QUEUE_FULL_CODE, null); + taskManager.setError(run.getRunId(), message); + completion.complete(run); + } + return null; + } + + private WaitState suspend(PolicyInputRequiredException e) { + List fileIds = new ArrayList<>(); + for (Resource resource : e.getPendingFiles()) { + String name = resource.getFilename() != null ? resource.getFilename() : "pending"; + try (InputStream is = resource.getInputStream()) { + fileIds.add(fileStorage.storeInputStream(is, name).fileId()); + } catch (IOException io) { + log.warn("Failed to persist pending file for paused run: {}", io.getMessage()); + } + } + return new WaitState(e.getMessage(), e.getResumeStepIndex(), fileIds); + } + + private PolicyProgressListener trackingListener( + String runId, PolicyRun run, PolicyProgressListener delegate) { + return new PolicyProgressListener() { + @Override + public void onStepStart(int stepIndex, int stepCount, String operation) { + run.enterStep(stepIndex); + taskManager.addNote( + runId, + "Step " + stepIndex + "/" + stepCount + ": " + operation + " started"); + delegate.onStepStart(stepIndex, stepCount, operation); + } + + @Override + public void onStepComplete(int stepIndex, int stepCount, String operation) { + taskManager.addNote( + runId, + "Step " + stepIndex + "/" + stepCount + ": " + operation + " completed"); + delegate.onStepComplete(stepIndex, stepCount, operation); + } + + @Override + public void onHeartbeat() { + delegate.onHeartbeat(); + } + }; + } + + private PolicyOutputSink sinkFor(OutputSpec spec) { + return outputSinks.stream() + .filter(sink -> sink.supports(spec)) + .findFirst() + .orElseThrow( + () -> + new IllegalStateException( + "No output sink supports spec: " + + (spec == null ? "" : spec.type()))); + } + + private static String toolTimeoutMessage(InternalApiTimeoutException e) { + return String.format( + "The %s tool did not respond within %d seconds and was aborted.", + e.getEndpointPath(), e.getReadTimeout().toSeconds()); + } + + /** + * MDC key {@code UserService.getCurrentUsername()} reads as its async fallback (stamped by the + * controller audit aspect on request threads). We reuse it to carry the billing identity onto + * the policy worker thread. + */ + private static final String AUDIT_PRINCIPAL_MDC_KEY = "auditPrincipal"; + + /** + * The username to bill an ad-hoc run to, captured on the submitting (request) thread. Prefers + * the audit principal the controller aspect already stamped; falls back to the security context + * name. {@code anonymousUser} (and no identity) resolve to null so we don't try to bill it. + */ + private static String currentActingPrincipal() { + String mdc = MDC.get(AUDIT_PRINCIPAL_MDC_KEY); + if (mdc != null && !mdc.isBlank()) { + return mdc; + } + Authentication auth = SecurityContextHolder.getContext().getAuthentication(); + if (auth == null) { + return null; + } + String name = auth.getName(); + return "anonymousUser".equals(name) ? null : name; + } + + /** + * Run {@code body} with {@code principal} set as the audit principal in MDC, so async tool + * dispatch attributes (and charges) usage to that user. A null/blank principal runs as-is. + * Restores the previous MDC value afterward (defensive — worker threads aren't pooled). + */ + private static void runAsPrincipal(String billingPrincipal, String fileOwner, Runnable body) { + // Billing identity (MDC auditPrincipal) and output-file ownership (JobContext owner) are + // set + // independently: usage is charged to billingPrincipal, but stored output files are owned by + // fileOwner — the user who triggered an org-wide policy — so they can fetch their results. + // Either may be null (e.g. login disabled, or a trigger-fired run); each is applied only + // when present and restored afterward (defensive — worker threads aren't pooled). + String previousPrincipal = MDC.get(AUDIT_PRINCIPAL_MDC_KEY); + String previousOwner = JobContext.getOwner(); + if (billingPrincipal != null && !billingPrincipal.isBlank()) { + MDC.put(AUDIT_PRINCIPAL_MDC_KEY, billingPrincipal); + } + if (fileOwner != null && !fileOwner.isBlank()) { + JobContext.setOwner(fileOwner); + } + try { + body.run(); + } finally { + if (previousPrincipal != null) { + MDC.put(AUDIT_PRINCIPAL_MDC_KEY, previousPrincipal); + } else { + MDC.remove(AUDIT_PRINCIPAL_MDC_KEY); + } + JobContext.setOwner(previousOwner); + } + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/engine/PolicyExecutionResult.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/engine/PolicyExecutionResult.java new file mode 100644 index 0000000000..8bfb78410b --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/engine/PolicyExecutionResult.java @@ -0,0 +1,14 @@ +package stirling.software.proprietary.policy.engine; + +import java.util.List; + +import org.springframework.core.io.Resource; + +import tools.jackson.databind.JsonNode; + +/** + * Result of a {@link PolicyExecutor} run. {@code files} are final temp files (not yet stored). + * {@code report}/{@code reportTool} carry the last step's structured report and its operation, or + * null if no step produced one. + */ +public record PolicyExecutionResult(List files, JsonNode report, String reportTool) {} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/engine/PolicyExecutor.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/engine/PolicyExecutor.java new file mode 100644 index 0000000000..af9f8c7283 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/engine/PolicyExecutor.java @@ -0,0 +1,284 @@ +package stirling.software.proprietary.policy.engine; + +import java.io.IOException; +import java.io.InputStream; +import java.util.ArrayList; +import java.util.List; +import java.util.Locale; +import java.util.Map; + +import org.springframework.core.io.Resource; +import org.springframework.http.HttpHeaders; +import org.springframework.http.HttpStatus; +import org.springframework.http.MediaType; +import org.springframework.http.ResponseEntity; +import org.springframework.stereotype.Service; +import org.springframework.util.LinkedMultiValueMap; +import org.springframework.util.MultiValueMap; + +import lombok.RequiredArgsConstructor; +import lombok.extern.slf4j.Slf4j; + +import stirling.software.common.service.InternalApiClient; +import stirling.software.common.service.InternalApiTimeoutException; +import stirling.software.common.service.ToolMetadataService; +import stirling.software.common.util.TempFileManager; +import stirling.software.common.util.ZipExtractionUtils; +import stirling.software.proprietary.policy.model.PipelineDefinition; +import stirling.software.proprietary.policy.model.PipelineStep; +import stirling.software.proprietary.policy.model.PolicyInputs; +import stirling.software.proprietary.policy.progress.PolicyProgressListener; +import stirling.software.proprietary.service.AiToolResponseHeaders; + +import tools.jackson.core.JacksonException; +import tools.jackson.databind.JsonNode; +import tools.jackson.databind.ObjectMapper; + +/** + * Runs an ordered chain of tool steps, feeding each step's output files into the next. + * + *

Steps dispatch synchronously via {@link InternalApiClient} loopback HTTP (each tool runs in + * its own handler, returns its file inline). The caller controls threading. Files cross step + * boundaries as {@link Resource} temp files and are only persisted at the run boundaries by the + * caller. + */ +@Slf4j +@Service +@RequiredArgsConstructor +public class PolicyExecutor { + + private static final String FILTER_OPERATION_PREFIX = "/api/v1/filter/filter-"; + + private final InternalApiClient internalApiClient; + private final ToolMetadataService toolMetadataService; + private final TempFileManager tempFileManager; + private final ObjectMapper objectMapper; + + // files: result files (one, or many for ZIP-response tools). report: optional structured + // payload the tool surfaced alongside or instead of a file. + private record ToolResult(List files, JsonNode report) {} + + /** + * Run every step in order, feeding each step's output into the next. Supporting files in {@code + * inputs} bind to named file fields and never enter the document stream. + * + * @throws InternalApiTimeoutException if a tool does not respond within its read timeout + * @throws IOException on a non-OK tool response, a missing supporting file, or a read failure + */ + public PolicyExecutionResult execute( + PipelineDefinition definition, PolicyInputs inputs, PolicyProgressListener listener) + throws IOException { + List steps = definition.steps(); + if (steps.isEmpty()) { + throw new IllegalArgumentException("Pipeline definition has no steps"); + } + + List currentFiles = inputs.primary(); + Map> supportingFiles = inputs.supportingFiles(); + // Last non-null report wins: the terminal step defines the output. + JsonNode lastReport = null; + String lastReportTool = null; + + for (int i = 0; i < steps.size(); i++) { + PipelineStep step = steps.get(i); + String operation = step.operation(); + if (operation == null || operation.isBlank()) { + throw new IllegalArgumentException( + "Pipeline step " + (i + 1) + " has no operation"); + } + listener.onStepStart(i + 1, steps.size(), operation); + ToolResult stepResult = executeStep(step, currentFiles, supportingFiles); + currentFiles = stepResult.files(); + if (stepResult.report() != null) { + lastReport = stepResult.report(); + lastReportTool = operation; + } + listener.onStepComplete(i + 1, steps.size(), operation); + } + + return new PolicyExecutionResult(currentFiles, lastReport, lastReportTool); + } + + /** + * Multi-input endpoints get all files in one call; others are called once per file. ZIP + * responses are unpacked so each inner file is its own result (e.g. split). For per-file + * dispatch the first non-null report wins. + */ + private ToolResult executeStep( + PipelineStep step, + List inputFiles, + Map> supportingFiles) + throws IOException { + requireAcceptedTypes(step.operation(), inputFiles); + List files = new ArrayList<>(); + JsonNode report = null; + if (toolMetadataService.isMultiInput(step.operation())) { + ToolResult r = callEndpoint(step, inputFiles, supportingFiles); + files.addAll(r.files()); + report = r.report(); + } else { + for (Resource file : inputFiles) { + ToolResult r = callEndpoint(step, List.of(file), supportingFiles); + files.addAll(r.files()); + if (report == null) { + report = r.report(); + } + } + } + return new ToolResult(files, report); + } + + /** + * Call an endpoint, returning result files and optional report. Response handling: JSON body is + * the report with no file; a file body returns the file plus any {@link + * AiToolResponseHeaders#TOOL_REPORT} header report; ZIP responses (per tool metadata) are + * unpacked to a flat file list. + */ + private ToolResult callEndpoint( + PipelineStep step, List files, Map> supportingFiles) + throws IOException { + String endpointPath = step.operation(); + MultiValueMap body = new LinkedMultiValueMap<>(); + for (Resource file : files) { + body.add("fileInput", file); + } + // Bind supporting files to named tool fields (e.g. stampImage); from the asset store, not + // the document stream. + for (Map.Entry binding : step.fileParameters().entrySet()) { + String fieldName = binding.getKey(); + String assetKey = binding.getValue(); + List assets = supportingFiles.get(assetKey); + if (assets == null || assets.isEmpty()) { + throw new IOException( + "Step " + + endpointPath + + " references supporting file '" + + assetKey + + "' for field '" + + fieldName + + "' but no such file was provided"); + } + for (Resource asset : assets) { + body.add(fieldName, asset); + } + } + for (Map.Entry entry : step.parameters().entrySet()) { + if (entry.getValue() instanceof List list) { + if (containsStructuredElements(list)) { + // These endpoints (e.g. /security/redact redactions, /general/edit-text edits) + // bind a list of structured objects from a single JSON string field via a + // property editor, so pre-serialize the whole list. + body.add(entry.getKey(), objectMapper.writeValueAsString(list)); + } else { + for (Object item : list) { + body.add(entry.getKey(), item); + } + } + } else { + body.add(entry.getKey(), entry.getValue()); + } + } + ResponseEntity response = internalApiClient.post(endpointPath, body); + if (!HttpStatus.OK.equals(response.getStatusCode()) || response.getBody() == null) { + throw new IOException( + "Tool returned HTTP " + response.getStatusCode() + " for " + endpointPath); + } + Resource resource = response.getBody(); + + // Filter ops return an empty body to mean "filtered out": drop it rather than forward a + // zero-byte document. + if (isFilterOperation(endpointPath) && isEmpty(resource)) { + return new ToolResult(List.of(), null); + } + + HttpHeaders headers = response.getHeaders(); + MediaType contentType = headers.getContentType(); + + // JSON-only response: whole body is the report, no file. + if (contentType != null && MediaType.APPLICATION_JSON.isCompatibleWith(contentType)) { + try (InputStream is = resource.getInputStream()) { + JsonNode report = objectMapper.readTree(is); + return new ToolResult(List.of(), report); + } + } + + JsonNode report = parseReportHeader(headers, endpointPath); + if (toolMetadataService.shouldUnpackZipResponse(endpointPath)) { + return new ToolResult(ZipExtractionUtils.extractZip(resource, tempFileManager), report); + } + return new ToolResult(List.of(resource), report); + } + + /** Parse the optional {@link AiToolResponseHeaders#TOOL_REPORT} header, or null. */ + private JsonNode parseReportHeader(HttpHeaders headers, String endpointPath) { + String raw = headers.getFirst(AiToolResponseHeaders.TOOL_REPORT); + if (raw == null || raw.isBlank()) { + return null; + } + try { + return objectMapper.readTree(raw); + } catch (JacksonException e) { + log.warn( + "Ignoring malformed {} header from {}: {}", + AiToolResponseHeaders.TOOL_REPORT, + endpointPath, + e.getMessage()); + return null; + } + } + + private static boolean containsStructuredElements(List list) { + for (Object item : list) { + if (item instanceof Map || item instanceof List) { + return true; + } + } + return false; + } + + /** + * Fail if any primary-stream file is a type the step rejects. No declared type means anything. + */ + private void requireAcceptedTypes(String operation, List files) throws IOException { + List accepted = toolMetadataService.getExtensionTypes(false, operation); + if (accepted == null || accepted.isEmpty()) { + return; + } + for (Resource file : files) { + if (!matchesType(file, accepted)) { + throw new IOException( + "Step " + + operation + + " accepts " + + accepted + + " but received '" + + file.getFilename() + + "'"); + } + } + } + + private static boolean matchesType(Resource file, List acceptedExtensions) { + String filename = file.getFilename(); + if (filename == null) { + return false; + } + int dot = filename.lastIndexOf('.'); + if (dot < 0 || dot == filename.length() - 1) { + return false; + } + return acceptedExtensions.contains(filename.substring(dot + 1).toLowerCase(Locale.ROOT)); + } + + private static boolean isFilterOperation(String operation) { + return operation.startsWith(FILTER_OPERATION_PREFIX); + } + + private static boolean isEmpty(Resource resource) { + try { + return resource.contentLength() == 0; + } catch (IOException e) { + return false; + } + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/engine/PolicyInputRequiredException.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/engine/PolicyInputRequiredException.java new file mode 100644 index 0000000000..936dedbe94 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/engine/PolicyInputRequiredException.java @@ -0,0 +1,26 @@ +package stirling.software.proprietary.policy.engine; + +import java.util.List; + +import org.springframework.core.io.Resource; + +import lombok.Getter; + +/** + * Thrown by a step that needs further user input, pausing the run in {@code WAITING_FOR_INPUT} + * instead of failing. Carries the resume reason, 0-based resume step index, and intermediate files; + * the engine persists those and suspends. Not yet thrown by any step. + */ +@Getter +public class PolicyInputRequiredException extends RuntimeException { + + private final transient List pendingFiles; + private final int resumeStepIndex; + + public PolicyInputRequiredException( + String reason, int resumeStepIndex, List pendingFiles) { + super(reason); + this.resumeStepIndex = resumeStepIndex; + this.pendingFiles = pendingFiles == null ? List.of() : pendingFiles; + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/engine/PolicyRunHandle.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/engine/PolicyRunHandle.java new file mode 100644 index 0000000000..09deaa7b5f --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/engine/PolicyRunHandle.java @@ -0,0 +1,13 @@ +package stirling.software.proprietary.policy.engine; + +import java.util.concurrent.CompletableFuture; + +import stirling.software.proprietary.policy.model.PolicyRun; + +/** + * Returned by {@link PolicyEngine#submit}: the run id (status polling, result download) plus a + * future that resolves when the run reaches a terminal or paused state. The future carries the + * {@link PolicyRun} whose status describes the outcome; it does not complete exceptionally for + * ordinary run failures. + */ +public record PolicyRunHandle(String runId, CompletableFuture completion) {} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/engine/PolicyRunRegistry.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/engine/PolicyRunRegistry.java new file mode 100644 index 0000000000..1c88191396 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/engine/PolicyRunRegistry.java @@ -0,0 +1,94 @@ +package stirling.software.proprietary.policy.engine; + +import java.time.Duration; +import java.time.Instant; +import java.util.Collection; +import java.util.Map; +import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.Executors; +import java.util.concurrent.ScheduledExecutorService; +import java.util.concurrent.TimeUnit; + +import org.springframework.context.annotation.Profile; +import org.springframework.stereotype.Service; + +import jakarta.annotation.PreDestroy; + +import lombok.extern.slf4j.Slf4j; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.proprietary.policy.model.PolicyRun; + +/** + * In-memory store of live {@link PolicyRun} state, keyed by runId. Authoritative run state machine; + * durable status/files are projected separately into {@code TaskManager}. + * + *

A scheduled sweep evicts only terminal runs aged past {@code policies.runExpiryMinutes}; + * active and paused runs are kept regardless of age. Eviction frees only this map's entry: the + * shared {@code TaskManager} job owns file-lifecycle cleanup. + */ +@Slf4j +@Service +@Profile("saas") +public class PolicyRunRegistry { + + private final Map runs = new ConcurrentHashMap<>(); + + private final Duration runExpiry; + private final ScheduledExecutorService cleanupExecutor = + Executors.newSingleThreadScheduledExecutor( + Thread.ofVirtual().name("policy-run-cleanup-", 0).factory()); + + public PolicyRunRegistry(ApplicationProperties applicationProperties) { + int runExpiryMinutes = applicationProperties.getPolicies().getRunExpiryMinutes(); + this.runExpiry = Duration.ofMinutes(runExpiryMinutes); + cleanupExecutor.scheduleAtFixedRate(this::evictExpiredRuns, 10, 10, TimeUnit.MINUTES); + log.debug( + "Policy run registry initialized with run expiry of {} minutes", runExpiryMinutes); + } + + public void register(PolicyRun run) { + runs.put(run.getRunId(), run); + } + + public PolicyRun get(String runId) { + return runs.get(runId); + } + + public Collection all() { + return runs.values(); + } + + /** Scheduled sweep entry point. */ + private void evictExpiredRuns() { + try { + evictExpired(Instant.now().minus(runExpiry)); + } catch (Exception e) { + log.error("Error during policy run cleanup: {}", e.getMessage(), e); + } + } + + /** + * Evict terminal runs last updated before {@code cutoff}, returning the count. Package-visible + * so the sweep and tests share one path with an explicit cutoff. + */ + int evictExpired(Instant cutoff) { + int removed = 0; + for (Map.Entry entry : runs.entrySet()) { + PolicyRun run = entry.getValue(); + if (run.getStatus().isTerminal() && run.getUpdatedAt().isBefore(cutoff)) { + runs.remove(entry.getKey()); + removed++; + } + } + if (removed > 0) { + log.info("Evicted {} expired policy runs", removed); + } + return removed; + } + + @PreDestroy + void shutdown() { + cleanupExecutor.shutdownNow(); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/engine/PolicyRunner.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/engine/PolicyRunner.java new file mode 100644 index 0000000000..9aa7d284c3 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/engine/PolicyRunner.java @@ -0,0 +1,107 @@ +package stirling.software.proprietary.policy.engine; + +import java.io.IOException; +import java.util.List; +import java.util.function.Consumer; + +import org.springframework.context.annotation.Profile; +import org.springframework.stereotype.Service; + +import lombok.RequiredArgsConstructor; +import lombok.extern.slf4j.Slf4j; + +import stirling.software.proprietary.policy.input.InputSource; +import stirling.software.proprietary.policy.input.ResolvedInput; +import stirling.software.proprietary.policy.model.InputSpec; +import stirling.software.proprietary.policy.model.PipelineDefinition; +import stirling.software.proprietary.policy.model.Policy; +import stirling.software.proprietary.policy.model.PolicyInputs; +import stirling.software.proprietary.policy.model.PolicyRun; +import stirling.software.proprietary.policy.model.PolicyRunStatus; +import stirling.software.proprietary.policy.progress.PolicyProgressListener; + +/** + * Turns a policy's configured {@link InputSpec sources} into runs. Triggers decide when + * and call {@link #run(Policy)}; the controller uses the supplied-input and ad-hoc entry points. + */ +@Slf4j +@Service +@RequiredArgsConstructor +@Profile("saas") +public class PolicyRunner { + + private final PolicyEngine policyEngine; + private final List inputSources; + + /** + * Trigger entry point. Pulls every configured source; each yielded unit becomes its own run so + * one failure does not affect the others. No sources means one run with no input (generator + * pipeline). + */ + public void run(Policy policy) { + List sources = policy.sources(); + if (sources.isEmpty()) { + startRun(policy, PolicyInputs.of(List.of()), unused -> {}); + return; + } + for (InputSpec spec : sources) { + pullAndRun(policy, spec); + } + } + + /** Run a stored policy on caller-supplied files (e.g. manual upload), bypassing its sources. */ + public PolicyRunHandle runWith( + Policy policy, PolicyInputs inputs, PolicyProgressListener listener) { + return policyEngine.runPolicy(policy, inputs, listener); + } + + /** Run an ad-hoc pipeline with no stored policy (AI/Automate one-offs). */ + public PolicyRunHandle runAdHoc( + PipelineDefinition definition, PolicyInputs inputs, PolicyProgressListener listener) { + return policyEngine.submit(definition, inputs, listener); + } + + private void pullAndRun(Policy policy, InputSpec spec) { + InputSource source = sourceFor(spec); + if (source == null) { + log.warn( + "No input source for type '{}' (policy {}); skipping", + spec.type(), + policy.id()); + return; + } + List work; + try { + work = source.resolve(spec); + } catch (IOException | RuntimeException e) { + log.warn( + "Failed to resolve source '{}' for policy {}: {}", + spec.type(), + policy.id(), + e.getMessage()); + return; + } + for (ResolvedInput unit : work) { + startRun(policy, unit.inputs(), unit.onComplete()); + } + } + + private void startRun(Policy policy, PolicyInputs inputs, Consumer onComplete) { + log.info("Running policy {} ({})", policy.id(), policy.name()); + PolicyRunHandle handle = + policyEngine.runPolicy(policy, inputs, PolicyProgressListener.NOOP); + handle.completion() + .whenComplete((run, throwable) -> onComplete.accept(succeeded(run, throwable))); + } + + private static boolean succeeded(PolicyRun run, Throwable throwable) { + return throwable == null && run != null && run.getStatus() == PolicyRunStatus.COMPLETED; + } + + private InputSource sourceFor(InputSpec spec) { + return inputSources.stream() + .filter(source -> source.supports(spec)) + .findFirst() + .orElse(null); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/engine/PolicyValidator.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/engine/PolicyValidator.java new file mode 100644 index 0000000000..6801ff38ce --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/engine/PolicyValidator.java @@ -0,0 +1,72 @@ +package stirling.software.proprietary.policy.engine; + +import java.util.List; + +import org.springframework.context.annotation.Profile; +import org.springframework.stereotype.Service; + +import lombok.RequiredArgsConstructor; + +import stirling.software.proprietary.policy.input.InputSource; +import stirling.software.proprietary.policy.model.InputSpec; +import stirling.software.proprietary.policy.model.OutputSpec; +import stirling.software.proprietary.policy.model.Policy; +import stirling.software.proprietary.policy.model.TriggerConfig; +import stirling.software.proprietary.policy.output.PolicyOutputSink; +import stirling.software.proprietary.policy.trigger.PolicyTrigger; + +/** + * Validates a policy at save time by delegating each facet (trigger, sources, output) to the bean + * that handles its type, so a misconfiguration fails fast rather than at run time. A null trigger + * is a manual-only policy and skips trigger validation. + */ +@Service +@RequiredArgsConstructor +@Profile("saas") +public class PolicyValidator { + + private final List triggers; + private final List inputSources; + private final List outputSinks; + + /** + * @throws IllegalArgumentException if any facet's type is unknown or its config is invalid + */ + public void validate(Policy policy) { + if (policy.trigger() != null) { + triggerFor(policy.trigger()).validate(policy); + } + for (InputSpec source : policy.sources()) { + inputSourceFor(source).validate(source); + } + outputSinkFor(policy.output()).validate(policy.output()); + } + + private PolicyTrigger triggerFor(TriggerConfig config) { + return triggers.stream() + .filter(trigger -> trigger.type().equals(config.type())) + .findFirst() + .orElseThrow( + () -> + new IllegalArgumentException( + "unknown trigger type: " + config.type())); + } + + private InputSource inputSourceFor(InputSpec spec) { + return inputSources.stream() + .filter(source -> source.supports(spec)) + .findFirst() + .orElseThrow( + () -> + new IllegalArgumentException( + "unknown input source type: " + spec.type())); + } + + private PolicyOutputSink outputSinkFor(OutputSpec spec) { + return outputSinks.stream() + .filter(sink -> sink.supports(spec)) + .findFirst() + .orElseThrow( + () -> new IllegalArgumentException("unknown output type: " + spec.type())); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/input/FolderInputSource.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/input/FolderInputSource.java new file mode 100644 index 0000000000..0990e83c18 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/input/FolderInputSource.java @@ -0,0 +1,179 @@ +package stirling.software.proprietary.policy.input; + +import java.io.IOException; +import java.nio.file.Files; +import java.nio.file.Path; +import java.nio.file.StandardCopyOption; +import java.util.ArrayList; +import java.util.List; +import java.util.Map; +import java.util.stream.Stream; + +import org.springframework.context.annotation.Profile; +import org.springframework.core.io.FileSystemResource; +import org.springframework.core.io.Resource; +import org.springframework.stereotype.Service; + +import lombok.RequiredArgsConstructor; +import lombok.extern.slf4j.Slf4j; + +import stirling.software.common.util.FileReadinessChecker; +import stirling.software.proprietary.policy.config.FolderAccessGuard; +import stirling.software.proprietary.policy.model.InputSpec; +import stirling.software.proprietary.policy.model.PolicyInputs; + +/** + * Reads input files from a directory; each ready file is its own unit of work so one failure does + * not affect the others. + * + *

Mode option: "consume" (default) claims each file by moving it into {@code + * .stirling/processing} then routes it to {@code .stirling/done} or {@code .stirling/error}, so + * each file runs once; "snapshot" reads without moving, so every run sees the full set. Readiness + * is checked first so files mid-write are skipped. + */ +@Slf4j +@Service +@RequiredArgsConstructor +@Profile("saas") +public class FolderInputSource implements InputSource { + + private static final String TYPE = FolderAccessGuard.FOLDER_TYPE; + // Bookkeeping lives under one hidden dir so the watched folder stays tidy. + private static final String WORK_SUBDIR = ".stirling"; + private static final String PROCESSING_SUBDIR = "processing"; + private static final String DONE_SUBDIR = "done"; + private static final String ERROR_SUBDIR = "error"; + + private final FileReadinessChecker readinessChecker; + private final FolderAccessGuard accessGuard; + + @Override + public String type() { + return TYPE; + } + + @Override + public boolean supports(InputSpec spec) { + return spec != null && TYPE.equals(spec.type()); + } + + @Override + public void validate(InputSpec spec) { + accessGuard.requirePermitted(FolderConfig.from(spec.options()).directory()); + } + + @Override + public List watchTargets(InputSpec spec) { + return List.of(FolderConfig.from(spec.options()).directory()); + } + + @Override + public List resolve(InputSpec spec) throws IOException { + FolderConfig config = FolderConfig.from(spec.options()); + Path inputDir = accessGuard.requirePermitted(config.directory()); + if (!Files.isDirectory(inputDir)) { + log.debug("Folder input dir does not exist: {}", inputDir); + return List.of(); + } + + List ready = new ArrayList<>(); + try (Stream entries = Files.list(inputDir)) { + entries.filter(Files::isRegularFile) + .filter(readinessChecker::isReady) + .forEach(ready::add); + } + + List work = new ArrayList<>(); + for (Path file : ready) { + if (config.snapshot()) { + work.add(ResolvedInput.of(PolicyInputs.of(List.of(fileResource(file))))); + } else { + Path claimed = claim(inputDir, file); + if (claimed == null) { + continue; // another sweep/process grabbed it + } + work.add( + new ResolvedInput( + PolicyInputs.of(List.of(fileResource(claimed))), + success -> route(inputDir, claimed, success))); + } + } + return work; + } + + // Atomic move into processing/: only one sweep can win the claim, the rest see the file gone. + private Path claim(Path inputDir, Path file) { + try { + Path processingDir = workDir(inputDir, PROCESSING_SUBDIR); + Files.createDirectories(processingDir); + Path claimed = uniqueTarget(processingDir, file.getFileName().toString()); + Files.move(file, claimed, StandardCopyOption.ATOMIC_MOVE); + return claimed; + } catch (IOException e) { + log.debug("Could not claim {}: {}", file, e.getMessage()); + return null; + } + } + + private void route(Path inputDir, Path claimed, boolean success) { + String subdir = success ? DONE_SUBDIR : ERROR_SUBDIR; + try { + Path destDir = workDir(inputDir, subdir); + Files.createDirectories(destDir); + Files.move( + claimed, + uniqueTarget(destDir, claimed.getFileName().toString()), + StandardCopyOption.ATOMIC_MOVE); + } catch (IOException e) { + log.warn( + "Could not move processed input {} to {}: {}", claimed, subdir, e.getMessage()); + } + } + + private static Path workDir(Path inputDir, String subdir) { + return inputDir.resolve(WORK_SUBDIR).resolve(subdir); + } + + private static Resource fileResource(Path path) { + String name = path.getFileName().toString(); + return new FileSystemResource(path.toFile()) { + @Override + public String getFilename() { + return name; + } + }; + } + + private static Path uniqueTarget(Path dir, String filename) { + Path candidate = dir.resolve(filename); + if (!Files.exists(candidate)) { + return candidate; + } + int dot = filename.lastIndexOf('.'); + String base = dot < 0 ? filename : filename.substring(0, dot); + String ext = dot < 0 ? "" : filename.substring(dot); + for (int n = 1; ; n++) { + Path next = dir.resolve(base + " (" + n + ")" + ext); + if (!Files.exists(next)) { + return next; + } + } + } + + record FolderConfig(Path directory, boolean snapshot) { + + private static final String DIRECTORY_OPTION = "directory"; + private static final String MODE_OPTION = "mode"; + private static final String MODE_SNAPSHOT = "snapshot"; + + static FolderConfig from(Map options) { + Object directory = options.get(DIRECTORY_OPTION); + if (directory == null || directory.toString().isBlank()) { + throw new IllegalArgumentException("folder input requires a 'directory' option"); + } + Object mode = options.get(MODE_OPTION); + boolean snapshot = mode != null && MODE_SNAPSHOT.equals(mode.toString()); + return new FolderConfig(Path.of(directory.toString()), snapshot); + } + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/input/InputSource.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/input/InputSource.java new file mode 100644 index 0000000000..436f6c5263 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/input/InputSource.java @@ -0,0 +1,38 @@ +package stirling.software.proprietary.policy.input; + +import java.io.IOException; +import java.nio.file.Path; +import java.util.List; + +import stirling.software.proprietary.policy.model.InputSpec; + +/** + * Resolves a policy {@link InputSpec} into the files to run on. Implementations are beans selected + * by {@link #supports(InputSpec)}, so a new source kind (folder, S3) is just a new bean. A manual + * run may supply files directly and bypass sources entirely. + */ +public interface InputSource { + + /** Stable identifier for this source, matching {@code InputSpec.type()} (e.g. "folder"). */ + String type(); + + /** Whether this source can handle the given spec. */ + boolean supports(InputSpec spec); + + /** Throws {@link IllegalArgumentException} on bad config. Called on save to fail fast. */ + default void validate(InputSpec spec) {} + + /** + * Resolve the spec into zero or more units of work, each carrying one run's files and a + * completion hook. Empty list means nothing to run right now. + */ + List resolve(InputSpec spec) throws IOException; + + /** + * Filesystem dirs this source draws from, for the folder-watch trigger. Advisory: resolving is + * still done by {@link #resolve}. Non-filesystem sources return empty and are not watchable. + */ + default List watchTargets(InputSpec spec) { + return List.of(); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/input/ResolvedInput.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/input/ResolvedInput.java new file mode 100644 index 0000000000..b5d7de94f8 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/input/ResolvedInput.java @@ -0,0 +1,22 @@ +package stirling.software.proprietary.policy.input; + +import java.util.function.Consumer; + +import stirling.software.proprietary.policy.model.PolicyInputs; + +/** + * One unit of work from an {@link InputSource}: the files to run plus a completion callback invoked + * with the run's success (e.g. a folder source routes the input to done/error). A source may return + * several of these, one per file. + */ +public record ResolvedInput(PolicyInputs inputs, Consumer onComplete) { + + public ResolvedInput { + onComplete = onComplete == null ? success -> {} : onComplete; + } + + /** No completion side effect. */ + public static ResolvedInput of(PolicyInputs inputs) { + return new ResolvedInput(inputs, success -> {}); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/InputSpec.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/InputSpec.java new file mode 100644 index 0000000000..181d5d78cd --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/InputSpec.java @@ -0,0 +1,19 @@ +package stirling.software.proprietary.policy.model; + +import java.util.Map; + +/** + * One input source for a policy. {@code type} keys an {@code InputSource} bean; a run pulls from + * every source. + */ +public record InputSpec(String type, Map options) { + + public InputSpec { + options = options == null ? Map.of() : options; + } + + /** Read input files from a directory on disk. */ + public static InputSpec folder(String directory) { + return new InputSpec("folder", Map.of("directory", directory)); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/OutputSpec.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/OutputSpec.java new file mode 100644 index 0000000000..1479ac6437 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/OutputSpec.java @@ -0,0 +1,20 @@ +package stirling.software.proprietary.policy.model; + +import java.util.Map; + +/** Where a run's outputs are delivered. {@code type} keys a {@code PolicyOutputSink} bean. */ +public record OutputSpec(String type, Map options) { + public OutputSpec { + options = options == null ? Map.of() : options; + } + + /** Default sink: store outputs and return them to the caller for download. */ + public static OutputSpec inline() { + return new OutputSpec("inline", Map.of()); + } + + /** Write outputs to a directory on disk. */ + public static OutputSpec folder(String directory) { + return new OutputSpec("folder", Map.of("directory", directory)); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/PipelineDefinition.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/PipelineDefinition.java new file mode 100644 index 0000000000..146424b0fb --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/PipelineDefinition.java @@ -0,0 +1,15 @@ +package stirling.software.proprietary.policy.model; + +import java.util.List; + +/** + * An ordered chain of tool steps plus an output destination; the unit the engine executes. + * + *

{@code output} may be null for callers that handle result files themselves (e.g. the AI + * workflow, which builds its own response payload). + */ +public record PipelineDefinition(String name, List steps, OutputSpec output) { + public PipelineDefinition { + steps = steps == null ? List.of() : steps; + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/PipelineStep.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/PipelineStep.java new file mode 100644 index 0000000000..99d8884996 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/PipelineStep.java @@ -0,0 +1,26 @@ +package stirling.software.proprietary.policy.model; + +import java.util.Map; + +/** + * A single tool invocation. {@code operation} is a Stirling endpoint path (e.g. {@code + * /api/v1/misc/compress-pdf}) per the {@code InternalApiClient} convention; {@code parameters} are + * scalar form fields. + * + *

{@code fileParameters} maps a tool's named file field (e.g. {@code stampImage}, beyond the + * primary {@code fileInput} stream) to an asset key in the run's supporting-file store, keeping + * supporting inputs out of the document stream that flows step to step. + */ +public record PipelineStep( + String operation, Map parameters, Map fileParameters) { + + public PipelineStep { + parameters = parameters == null ? Map.of() : parameters; + fileParameters = fileParameters == null ? Map.of() : fileParameters; + } + + /** A step with no supporting-file bindings. */ + public PipelineStep(String operation, Map parameters) { + this(operation, parameters, Map.of()); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/Policy.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/Policy.java new file mode 100644 index 0000000000..6d5366eb47 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/Policy.java @@ -0,0 +1,61 @@ +package stirling.software.proprietary.policy.model; + +import java.util.List; + +/** + * A stored automation: ordered tool steps, input sources, and an output destination. + * + *

Always runnable on demand. An optional {@link TriggerConfig} fires it automatically; a {@code + * null} trigger means manual-only. Trigger decides when, {@link InputSpec sources} decide where + * files come from; a run pulls from every source. + */ +public record Policy( + String id, + String name, + String owner, + boolean enabled, + TriggerConfig trigger, + List sources, + List steps, + OutputSpec output, + Long teamId) { + + public Policy { + sources = sources == null ? List.of() : List.copyOf(sources); + steps = steps == null ? List.of() : steps; + output = output == null ? OutputSpec.inline() : output; + } + + /** + * Without an explicit owning team. Kept for the engine and tests; the controller always stamps + * a {@code teamId} on stored policies so they stay scoped to the creating user's team. + */ + public Policy( + String id, + String name, + String owner, + boolean enabled, + TriggerConfig trigger, + List sources, + List steps, + OutputSpec output) { + this(id, name, owner, enabled, trigger, sources, steps, output, null); + } + + /** A policy with no configured sources (a generator, or files supplied directly to a run). */ + public Policy( + String id, + String name, + String owner, + boolean enabled, + TriggerConfig trigger, + List steps, + OutputSpec output) { + this(id, name, owner, enabled, trigger, List.of(), steps, output, null); + } + + /** This policy's pipeline as the engine sees it. */ + public PipelineDefinition toDefinition() { + return new PipelineDefinition(name, steps, output); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/PolicyInputs.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/PolicyInputs.java new file mode 100644 index 0000000000..7088cd4415 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/PolicyInputs.java @@ -0,0 +1,24 @@ +package stirling.software.proprietary.policy.model; + +import java.util.List; +import java.util.Map; + +import org.springframework.core.io.Resource; + +/** + * A run's files. {@code primary} documents flow step to step; {@code supportingFiles} are auxiliary + * assets bound by key via {@link PipelineStep#fileParameters()} and never enter the document + * stream. Asset values are lists so one key can carry a multi-file field (e.g. attachments). + */ +public record PolicyInputs(List primary, Map> supportingFiles) { + + public PolicyInputs { + primary = primary == null ? List.of() : primary; + supportingFiles = supportingFiles == null ? Map.of() : supportingFiles; + } + + /** Inputs with primary documents only and no supporting files. */ + public static PolicyInputs of(List primary) { + return new PolicyInputs(primary, Map.of()); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/PolicyRun.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/PolicyRun.java new file mode 100644 index 0000000000..bba84813ce --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/PolicyRun.java @@ -0,0 +1,114 @@ +package stirling.software.proprietary.policy.model; + +import java.time.Instant; +import java.util.List; + +import lombok.Getter; + +import stirling.software.common.model.job.ResultFile; + +/** + * Live, mutable state of one pipeline run, held in memory by {@code PolicyRunRegistry} and the + * authoritative source of the state machine. Carries execution state ({@code JobResult} does not + * model status/step cursor/wait state); also projected into {@code TaskManager} for cluster-visible + * status and download. + */ +@Getter +public class PolicyRun { + + private final String runId; + + /** ID of the stored policy that produced this run; null for ad-hoc pipelines. */ + private final String policyId; + + private final PipelineDefinition definition; + private final Instant createdAt = Instant.now(); + + private volatile PolicyRunStatus status = PolicyRunStatus.PENDING; + + /** 1-based index of the step currently running (0 before the run starts). */ + private volatile int currentStep = 0; + + private volatile WaitState waitState; + private volatile String error; + + /** + * Stable, machine-readable failure code the client can branch on — e.g. an entitlement-limit + * sentinel ({@code PAYG_LIMIT_REACHED} / {@code FEATURE_DEGRADED}) propagated from a downstream + * tool call's 402 — alongside the human-readable {@link #error}. Null unless set on failure. + */ + private volatile String errorCode; + + /** + * For an entitlement-limit failure, whether the team was subscribed (over its spending cap) vs + * un-subscribed (free allowance spent) — taken from the blocking 402 body. Drives which + * usage-limit modal the client shows. Null unless {@link #errorCode} is an entitlement code. + */ + private volatile Boolean errorSubscribed; + + private volatile List outputs = List.of(); + private volatile Instant updatedAt = Instant.now(); + + public PolicyRun(String runId, String policyId, PipelineDefinition definition) { + this.runId = runId; + this.policyId = policyId; + this.definition = definition; + } + + public int stepCount() { + return definition.steps().size(); + } + + public synchronized void markRunning() { + this.status = PolicyRunStatus.RUNNING; + touch(); + } + + public synchronized void enterStep(int oneBasedStepIndex) { + this.currentStep = oneBasedStepIndex; + touch(); + } + + public synchronized void complete(List resultFiles) { + this.outputs = resultFiles == null ? List.of() : List.copyOf(resultFiles); + this.status = PolicyRunStatus.COMPLETED; + touch(); + } + + public synchronized void fail(String message) { + this.error = message; + this.status = PolicyRunStatus.FAILED; + touch(); + } + + /** + * Fail with a stable {@code errorCode} the client can branch on (e.g. an entitlement-limit + * sentinel from a downstream 402), plus the optional {@code subscribed} flag from that + * response, in addition to the human-readable message. + */ + public synchronized void failWithCode(String message, String errorCode, Boolean subscribed) { + this.errorCode = errorCode; + this.errorSubscribed = subscribed; + fail(message); + } + + public synchronized void waitForInput(WaitState wait) { + this.waitState = wait; + this.status = PolicyRunStatus.WAITING_FOR_INPUT; + touch(); + } + + /** Cancels unless already terminal; returns whether it transitioned. */ + public synchronized boolean cancel() { + if (status.isTerminal()) { + return false; + } + this.status = PolicyRunStatus.CANCELLED; + touch(); + return true; + } + + private void touch() { + this.updatedAt = Instant.now(); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/PolicyRunStatus.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/PolicyRunStatus.java new file mode 100644 index 0000000000..646703bf57 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/PolicyRunStatus.java @@ -0,0 +1,18 @@ +package stirling.software.proprietary.policy.model; + +/** + * Lifecycle states of a {@link PolicyRun}. {@code WAITING_FOR_INPUT} models a thread-free pause; + * the resume handshake lands in a later stage. + */ +public enum PolicyRunStatus { + PENDING, + RUNNING, + WAITING_FOR_INPUT, + COMPLETED, + FAILED, + CANCELLED; + + public boolean isTerminal() { + return this == COMPLETED || this == FAILED || this == CANCELLED; + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/PolicyRunView.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/PolicyRunView.java new file mode 100644 index 0000000000..04a2707077 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/PolicyRunView.java @@ -0,0 +1,37 @@ +package stirling.software.proprietary.policy.model; + +import java.util.List; + +import stirling.software.common.model.job.ResultFile; + +/** + * Read-only view of a {@link PolicyRun} for the status endpoint. Outputs are {@link ResultFile}s, + * downloadable via {@code GET /api/v1/general/files/{id}}. + */ +public record PolicyRunView( + String runId, + String policyId, + PolicyRunStatus status, + int currentStep, + int stepCount, + String error, + String errorCode, + Boolean errorSubscribed, + List outputs, + /** When the run was created, epoch millis, so a rediscovered run shows its real age. */ + long createdAt) { + + public static PolicyRunView of(PolicyRun run) { + return new PolicyRunView( + run.getRunId(), + run.getPolicyId(), + run.getStatus(), + run.getCurrentStep(), + run.stepCount(), + run.getError(), + run.getErrorCode(), + run.getErrorSubscribed(), + run.getOutputs(), + run.getCreatedAt().toEpochMilli()); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/Schedule.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/Schedule.java new file mode 100644 index 0000000000..7a684942f2 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/Schedule.java @@ -0,0 +1,130 @@ +package stirling.software.proprietary.policy.model; + +import java.time.DayOfWeek; +import java.time.LocalTime; +import java.time.ZonedDateTime; +import java.util.EnumSet; +import java.util.Set; + +import com.fasterxml.jackson.annotation.JsonIgnoreProperties; +import com.fasterxml.jackson.annotation.JsonSubTypes; +import com.fasterxml.jackson.annotation.JsonTypeInfo; + +/** + * A scheduled policy's firing cadence; {@code type} is the JSON discriminator. Wall-clock kinds + * ({@link Daily}, {@link Weekly}, {@link Monthly}) evaluate in the {@code after} argument's zone; + * {@link Every} is a fixed offset and ignores wall-clock time. + */ +@JsonTypeInfo(use = JsonTypeInfo.Id.NAME, property = "type") +@JsonSubTypes({ + @JsonSubTypes.Type(value = Schedule.Every.class, name = "every"), + @JsonSubTypes.Type(value = Schedule.Daily.class, name = "daily"), + @JsonSubTypes.Type(value = Schedule.Weekly.class, name = "weekly"), + @JsonSubTypes.Type(value = Schedule.Monthly.class, name = "monthly"), +}) +@JsonIgnoreProperties(ignoreUnknown = true) +public sealed interface Schedule { + + /** The next firing strictly after {@code after}, evaluated in {@code after}'s zone. */ + ZonedDateTime nextAfter(ZonedDateTime after); + + /** The granularities a fixed-interval schedule can repeat on. */ + enum Unit { + MINUTES, + HOURS, + DAYS + } + + /** A fixed offset from {@code after}: "every 15 minutes", "every 6 hours". No time of day. */ + record Every(long count, Unit unit) implements Schedule { + public Every { + if (count <= 0) { + throw new IllegalArgumentException("'every' schedule needs a positive count"); + } + if (unit == null) { + throw new IllegalArgumentException("'every' schedule needs a unit"); + } + } + + @Override + public ZonedDateTime nextAfter(ZonedDateTime after) { + return switch (unit) { + case MINUTES -> after.plusMinutes(count); + case HOURS -> after.plusHours(count); + case DAYS -> after.plusDays(count); + }; + } + } + + /** Once a day at a wall-clock time: "every day at 02:00". */ + record Daily(LocalTime at) implements Schedule { + public Daily { + requireTime(at); + } + + @Override + public ZonedDateTime nextAfter(ZonedDateTime after) { + ZonedDateTime today = after.with(at); + return today.isAfter(after) ? today : today.plusDays(1); + } + } + + /** On chosen weekdays at a wall-clock time: "every Monday and Thursday at 09:00". */ + record Weekly(Set days, LocalTime at) implements Schedule { + public Weekly { + if (days == null || days.isEmpty()) { + throw new IllegalArgumentException("'weekly' schedule needs at least one day"); + } + requireTime(at); + days = EnumSet.copyOf(days); + } + + @Override + public ZonedDateTime nextAfter(ZonedDateTime after) { + // Soonest of the next 7 days landing on a chosen weekday, at the configured time. + for (int i = 0; i <= 7; i++) { + ZonedDateTime candidate = after.plusDays(i).with(at); + if (candidate.isAfter(after) && days.contains(candidate.getDayOfWeek())) { + return candidate; + } + } + throw new IllegalStateException("unreachable: a chosen weekday recurs within 8 days"); + } + } + + /** + * On a day of the month at a wall-clock time: "the 1st at 00:00". Months too short for the + * chosen day (e.g. the 31st in February) are skipped, not clamped. + */ + record Monthly(int dayOfMonth, LocalTime at) implements Schedule { + public Monthly { + if (dayOfMonth < 1 || dayOfMonth > 31) { + throw new IllegalArgumentException("'monthly' day-of-month must be 1-31"); + } + requireTime(at); + } + + @Override + public ZonedDateTime nextAfter(ZonedDateTime after) { + ZonedDateTime firstOfMonth = after.withDayOfMonth(1).with(at); + // Scan forward a few years' worth of months to skip ones without the chosen day. + for (int i = 0; i < 48; i++) { + ZonedDateTime month = firstOfMonth.plusMonths(i); + if (month.toLocalDate().lengthOfMonth() >= dayOfMonth) { + ZonedDateTime fire = month.withDayOfMonth(dayOfMonth); + if (fire.isAfter(after)) { + return fire; + } + } + } + throw new IllegalStateException( + "unreachable: a month with the chosen day recurs yearly"); + } + } + + private static void requireTime(LocalTime at) { + if (at == null) { + throw new IllegalArgumentException("schedule needs a time of day ('at')"); + } + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/TriggerConfig.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/TriggerConfig.java new file mode 100644 index 0000000000..477de122e7 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/TriggerConfig.java @@ -0,0 +1,15 @@ +package stirling.software.proprietary.policy.model; + +import java.util.Map; + +/** + * A {@link Policy}'s automatic trigger; {@code type} keys a trigger bean (e.g. "schedule"). Manual + * running is not a trigger kind: a manual-only policy carries a {@code null} {@code TriggerConfig}. + * Answers only "when"; file sources are the policy's {@link InputSpec}s. + */ +public record TriggerConfig(String type, Map options) { + + public TriggerConfig { + options = options == null ? Map.of() : options; + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/WaitState.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/WaitState.java new file mode 100644 index 0000000000..7b0947413b --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/model/WaitState.java @@ -0,0 +1,15 @@ +package stirling.software.proprietary.policy.model; + +import java.util.List; + +/** + * Resumable snapshot captured when a run pauses ({@link PolicyRunStatus#WAITING_FOR_INPUT}). {@code + * resumeStepIndex} is the 0-based step to continue from; {@code pendingFileIds} are intermediate + * files held in {@code FileStorage} (not in-memory resources) so a pause survives the worker thread + * ending or a node restart. + */ +public record WaitState(String reason, int resumeStepIndex, List pendingFileIds) { + public WaitState { + pendingFileIds = pendingFileIds == null ? List.of() : pendingFileIds; + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/output/FolderOutputSink.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/output/FolderOutputSink.java new file mode 100644 index 0000000000..94a9d34a35 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/output/FolderOutputSink.java @@ -0,0 +1,125 @@ +package stirling.software.proprietary.policy.output; + +import java.io.IOException; +import java.io.InputStream; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.ArrayList; +import java.util.List; +import java.util.UUID; + +import org.apache.commons.io.FilenameUtils; +import org.springframework.context.annotation.Profile; +import org.springframework.core.io.Resource; +import org.springframework.http.MediaType; +import org.springframework.http.MediaTypeFactory; +import org.springframework.stereotype.Service; + +import lombok.RequiredArgsConstructor; +import lombok.extern.slf4j.Slf4j; + +import stirling.software.common.model.job.ResultFile; +import stirling.software.proprietary.policy.config.FolderAccessGuard; +import stirling.software.proprietary.policy.model.OutputSpec; + +/** + * Writes a run's outputs to the {@code directory} given in the {@link OutputSpec}. Files are + * streamed (not buffered) and uniquely named to avoid clobbering. Returned {@link ResultFile}s + * carry a synthetic id since the deliverable is the file on disk, not a {@code FileStorage} entry, + * so folder outputs are not downloadable via {@code /files/{id}}. + */ +@Slf4j +@Service +@RequiredArgsConstructor +@Profile("saas") +public class FolderOutputSink implements PolicyOutputSink { + + static final String TYPE = FolderAccessGuard.FOLDER_TYPE; + static final String DIRECTORY_OPTION = "directory"; + + private final FolderAccessGuard accessGuard; + + @Override + public String type() { + return TYPE; + } + + @Override + public boolean supports(OutputSpec spec) { + return spec != null && TYPE.equals(spec.type()); + } + + @Override + public void validate(OutputSpec spec) { + accessGuard.requirePermitted(directoryOf(spec)); + } + + @Override + public List deliver(String runId, List outputs, OutputSpec spec) + throws IOException { + Path targetDir = accessGuard.requirePermitted(directoryOf(spec)); + Files.createDirectories(targetDir); + + List results = new ArrayList<>(); + for (int i = 0; i < outputs.size(); i++) { + Resource resource = outputs.get(i); + String name = safeName(resource.getFilename(), i); + Path target = uniqueTarget(targetDir, name); + try (InputStream is = resource.getInputStream()) { + Files.copy(is, target); + } + long size = Files.size(target); + String contentType = + MediaTypeFactory.getMediaType(name) + .orElse(MediaType.APPLICATION_OCTET_STREAM) + .toString(); + results.add( + ResultFile.builder() + .fileId(UUID.randomUUID().toString()) + .fileName(target.toString()) + .contentType(contentType) + .fileSize(size) + .build()); + log.debug("Wrote policy run {} output to {}", runId, target); + } + return results; + } + + private static Path directoryOf(OutputSpec spec) { + Object directory = spec.options().get(DIRECTORY_OPTION); + if (directory == null || directory.toString().isBlank()) { + throw new IllegalArgumentException( + "folder output requires a '" + DIRECTORY_OPTION + "' option"); + } + return Path.of(directory.toString()); + } + + // Strip any directory component / "../" so a crafted output name cannot escape targetDir. + private static String safeName(String filename, int index) { + if (filename == null || filename.isBlank()) { + return "output-" + index; + } + String name = FilenameUtils.getName(filename); + if (name.isBlank() || ".".equals(name) || "..".equals(name)) { + return "output-" + index; + } + return name; + } + + // Non-colliding path, appending " (n)" before the extension. + private static Path uniqueTarget(Path dir, String filename) { + Path candidate = dir.resolve(filename); + if (!Files.exists(candidate)) { + return candidate; + } + String base = FilenameUtils.getBaseName(filename); + String ext = FilenameUtils.getExtension(filename); + String suffix = ext.isEmpty() ? "" : "." + ext; + for (int n = 1; ; n++) { + Path next = dir.resolve(base + " (" + n + ")" + suffix); + if (!Files.exists(next)) { + return next; + } + } + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/output/InlineOutputSink.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/output/InlineOutputSink.java new file mode 100644 index 0000000000..799072aeaa --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/output/InlineOutputSink.java @@ -0,0 +1,69 @@ +package stirling.software.proprietary.policy.output; + +import java.io.IOException; +import java.io.InputStream; +import java.util.ArrayList; +import java.util.List; + +import org.springframework.context.annotation.Profile; +import org.springframework.core.io.Resource; +import org.springframework.http.MediaType; +import org.springframework.http.MediaTypeFactory; +import org.springframework.stereotype.Service; + +import lombok.RequiredArgsConstructor; + +import stirling.software.common.model.job.ResultFile; +import stirling.software.common.service.FileStorage; +import stirling.software.proprietary.policy.model.OutputSpec; + +/** + * Default sink: stores each output in {@code FileStorage} so it is downloadable via {@code GET + * /api/v1/general/files/{fileId}}. Used for manual runs whose results return to the caller. + */ +@Service +@RequiredArgsConstructor +@Profile("saas") +public class InlineOutputSink implements PolicyOutputSink { + + private static final String TYPE = "inline"; + + private final FileStorage fileStorage; + + @Override + public String type() { + return TYPE; + } + + @Override + public boolean supports(OutputSpec spec) { + return spec == null || spec.type() == null || TYPE.equals(spec.type()); + } + + @Override + public List deliver(String runId, List outputs, OutputSpec spec) + throws IOException { + List results = new ArrayList<>(); + for (int i = 0; i < outputs.size(); i++) { + Resource resource = outputs.get(i); + String name = + resource.getFilename() != null ? resource.getFilename() : "result-" + (i + 1); + String contentType = + MediaTypeFactory.getMediaType(name) + .orElse(MediaType.APPLICATION_OCTET_STREAM) + .toString(); + FileStorage.StoredFile stored; + try (InputStream is = resource.getInputStream()) { + stored = fileStorage.storeInputStream(is, name); + } + results.add( + ResultFile.builder() + .fileId(stored.fileId()) + .fileName(name) + .contentType(contentType) + .fileSize(stored.size()) + .build()); + } + return results; + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/output/PolicyOutputSink.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/output/PolicyOutputSink.java new file mode 100644 index 0000000000..6e6c80ba4d --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/output/PolicyOutputSink.java @@ -0,0 +1,30 @@ +package stirling.software.proprietary.policy.output; + +import java.io.IOException; +import java.util.List; + +import org.springframework.core.io.Resource; + +import stirling.software.common.model.job.ResultFile; +import stirling.software.proprietary.policy.model.OutputSpec; + +/** + * Delivers a finished run's outputs to a destination, returning {@link ResultFile} descriptors for + * the run record. Implementations are beans selected by {@link #supports(OutputSpec)}, so a new + * destination (folder, S3) is just a new bean. + */ +public interface PolicyOutputSink { + + /** Stable identifier for this sink, matching {@code OutputSpec.type()} (e.g. "inline"). */ + String type(); + + /** Whether this sink can handle the given output spec. */ + boolean supports(OutputSpec spec); + + /** Throws {@link IllegalArgumentException} on bad config. Called on save to fail fast. */ + default void validate(OutputSpec spec) {} + + /** Persist/deliver the output files and return their descriptors. */ + List deliver(String runId, List outputs, OutputSpec spec) + throws IOException; +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/progress/PolicyProgressListener.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/progress/PolicyProgressListener.java new file mode 100644 index 0000000000..98c89f3a78 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/progress/PolicyProgressListener.java @@ -0,0 +1,17 @@ +package stirling.software.proprietary.policy.progress; + +/** + * Receives live progress as a pipeline run executes (SSE stream, job notes, or both). Step indices + * are 1-based. All methods default to no-ops. + */ +public interface PolicyProgressListener { + + PolicyProgressListener NOOP = new PolicyProgressListener() {}; + + default void onStepStart(int stepIndex, int stepCount, String operation) {} + + default void onStepComplete(int stepIndex, int stepCount, String operation) {} + + /** Keep-alive tick so downstream connections can detect disconnects promptly. */ + default void onHeartbeat() {} +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/store/InProcessPolicyStore.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/store/InProcessPolicyStore.java new file mode 100644 index 0000000000..d4a1259d3d --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/store/InProcessPolicyStore.java @@ -0,0 +1,63 @@ +package stirling.software.proprietary.policy.store; + +import java.util.List; +import java.util.Map; +import java.util.Optional; +import java.util.UUID; +import java.util.concurrent.ConcurrentHashMap; + +import stirling.software.proprietary.policy.model.Policy; + +/** + * In-memory {@link PolicyStore} for tests and any future no-database mode. {@link JpaPolicyStore} + * is the runtime bean. + */ +public class InProcessPolicyStore implements PolicyStore { + + private final Map policies = new ConcurrentHashMap<>(); + + @Override + public Policy save(Policy policy) { + String id = + policy.id() == null || policy.id().isBlank() + ? UUID.randomUUID().toString() + : policy.id(); + Policy stored = + new Policy( + id, + policy.name(), + policy.owner(), + policy.enabled(), + policy.trigger(), + policy.sources(), + policy.steps(), + policy.output(), + policy.teamId()); + policies.put(id, stored); + return stored; + } + + @Override + public Optional get(String id) { + return Optional.ofNullable(policies.get(id)); + } + + @Override + public List all() { + return List.copyOf(policies.values()); + } + + @Override + public List findByTriggerType(String triggerType) { + return policies.values().stream() + .filter(Policy::enabled) + .filter(policy -> policy.trigger() != null) + .filter(policy -> triggerType.equals(policy.trigger().type())) + .toList(); + } + + @Override + public boolean delete(String id) { + return policies.remove(id) != null; + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/store/JpaPolicyStore.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/store/JpaPolicyStore.java new file mode 100644 index 0000000000..085a97b5a7 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/store/JpaPolicyStore.java @@ -0,0 +1,86 @@ +package stirling.software.proprietary.policy.store; + +import java.util.List; +import java.util.Optional; +import java.util.UUID; + +import org.springframework.context.annotation.Profile; +import org.springframework.stereotype.Service; + +import lombok.RequiredArgsConstructor; + +import stirling.software.proprietary.policy.model.Policy; + +import tools.jackson.databind.ObjectMapper; + +/** + * Durable {@link PolicyStore} backed by JPA; the runtime store. Policies are persisted as JSON via + * {@link PolicyEntity}, with scalar columns kept in sync for querying. + */ +@Service +@RequiredArgsConstructor +@Profile("saas") +public class JpaPolicyStore implements PolicyStore { + + private final PolicyRepository repository; + private final ObjectMapper objectMapper; + + @Override + public Policy save(Policy policy) { + String id = + policy.id() == null || policy.id().isBlank() + ? UUID.randomUUID().toString() + : policy.id(); + Policy stored = + new Policy( + id, + policy.name(), + policy.owner(), + policy.enabled(), + policy.trigger(), + policy.sources(), + policy.steps(), + policy.output(), + policy.teamId()); + + PolicyEntity entity = new PolicyEntity(); + entity.setId(id); + entity.setName(stored.name()); + entity.setOwner(stored.owner()); + entity.setEnabled(stored.enabled()); + entity.setTriggerType(stored.trigger() == null ? null : stored.trigger().type()); + entity.setPolicyJson(objectMapper.writeValueAsString(stored)); + repository.save(entity); + return stored; + } + + @Override + public Optional get(String id) { + return repository.findById(id).map(this::toPolicy); + } + + @Override + public List all() { + return repository.findAll().stream().map(this::toPolicy).toList(); + } + + @Override + public List findByTriggerType(String triggerType) { + return repository.findByTriggerTypeAndEnabledTrue(triggerType).stream() + .map(this::toPolicy) + .toList(); + } + + @Override + public boolean delete(String id) { + if (!repository.existsById(id)) { + return false; + } + repository.deleteById(id); + return true; + } + + private Policy toPolicy(PolicyEntity entity) { + return objectMapper.readValue(entity.getPolicyJson(), Policy.class); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/store/PolicyEntity.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/store/PolicyEntity.java new file mode 100644 index 0000000000..944b494611 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/store/PolicyEntity.java @@ -0,0 +1,48 @@ +package stirling.software.proprietary.policy.store; + +import java.io.Serializable; + +import jakarta.persistence.Column; +import jakarta.persistence.Entity; +import jakarta.persistence.Id; +import jakarta.persistence.Table; + +import lombok.Getter; +import lombok.NoArgsConstructor; +import lombok.Setter; + +/** + * JPA row for a {@link stirling.software.proprietary.policy.model.Policy}. The whole policy lives + * as JSON in {@code policyJson} (authoritative on read); the scalar columns are denormalized copies + * for querying, notably {@code triggerType} + {@code enabled} so background triggers can fetch + * their policies. {@code owner} is a plain string, not a foreign key, to stay decoupled from the + * security entities. + */ +@Entity +@Table(name = "policies") +@NoArgsConstructor +@Getter +@Setter +public class PolicyEntity implements Serializable { + + private static final long serialVersionUID = 1L; + + @Id + @Column(name = "id") + private String id; + + @Column(name = "name") + private String name; + + @Column(name = "owner") + private String owner; + + @Column(name = "enabled") + private boolean enabled; + + @Column(name = "trigger_type") + private String triggerType; + + @Column(name = "policy_json", columnDefinition = "text") + private String policyJson; +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/store/PolicyRepository.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/store/PolicyRepository.java new file mode 100644 index 0000000000..ba6924f0f8 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/store/PolicyRepository.java @@ -0,0 +1,13 @@ +package stirling.software.proprietary.policy.store; + +import java.util.List; + +import org.springframework.data.jpa.repository.JpaRepository; +import org.springframework.stereotype.Repository; + +@Repository +public interface PolicyRepository extends JpaRepository { + + /** Enabled policies of a given trigger type, for background triggers to activate. */ + List findByTriggerTypeAndEnabledTrue(String triggerType); +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/store/PolicyStore.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/store/PolicyStore.java new file mode 100644 index 0000000000..c9a2a0ecf7 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/store/PolicyStore.java @@ -0,0 +1,23 @@ +package stirling.software.proprietary.policy.store; + +import java.util.List; +import java.util.Optional; + +import stirling.software.proprietary.policy.model.Policy; + +/** Stores {@link Policy} definitions. */ +public interface PolicyStore { + + /** Create or update; a blank/absent id is assigned. Returns the stored policy. */ + Policy save(Policy policy); + + Optional get(String id); + + List all(); + + /** Enabled policies with the given trigger type, for background triggers. */ + List findByTriggerType(String triggerType); + + /** Returns whether the policy existed. */ + boolean delete(String id); +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/trigger/FolderWatchTrigger.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/trigger/FolderWatchTrigger.java new file mode 100644 index 0000000000..864160d4b6 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/trigger/FolderWatchTrigger.java @@ -0,0 +1,290 @@ +package stirling.software.proprietary.policy.trigger; + +import static java.nio.file.StandardWatchEventKinds.ENTRY_CREATE; +import static java.nio.file.StandardWatchEventKinds.ENTRY_MODIFY; + +import java.io.IOException; +import java.nio.file.ClosedWatchServiceException; +import java.nio.file.FileSystems; +import java.nio.file.Files; +import java.nio.file.Path; +import java.nio.file.WatchKey; +import java.nio.file.WatchService; +import java.util.ArrayList; +import java.util.HashSet; +import java.util.List; +import java.util.Map; +import java.util.Set; +import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.Executors; +import java.util.concurrent.ScheduledExecutorService; +import java.util.concurrent.TimeUnit; + +import org.springframework.context.annotation.Profile; +import org.springframework.stereotype.Service; + +import lombok.RequiredArgsConstructor; +import lombok.extern.slf4j.Slf4j; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.proprietary.policy.engine.PolicyRunner; +import stirling.software.proprietary.policy.input.InputSource; +import stirling.software.proprietary.policy.model.InputSpec; +import stirling.software.proprietary.policy.model.Policy; +import stirling.software.proprietary.policy.store.PolicyStore; + +/** + * Fires policies when a file lands in one of their folder sources, rather than polling on a timer. + * + *

The watch is a latency optimisation, not a source of truth: a periodic reconcile sweep ({@code + * watchReconcileSeconds}) re-syncs watched dirs and re-runs every policy, covering files that + * pre-dated the watch, dropped events, and filesystems that emit none (NFS, bind mounts). Redundant + * runs are harmless since {@link InputSource} does the claiming. + * + *

Watch state is in memory, so this assumes a single node and rebuilds registrations on restart. + */ +@Slf4j +@Service +@RequiredArgsConstructor +@Profile("saas") +public class FolderWatchTrigger implements PolicyTrigger { + + private static final String TYPE = "folder-watch"; + + private final PolicyStore policyStore; + private final PolicyRunner policyRunner; + private final List inputSources; + private final ApplicationProperties applicationProperties; + + private final Map keysByDir = new ConcurrentHashMap<>(); + private final Map dirByKey = new ConcurrentHashMap<>(); + + private volatile boolean running; + + // Package-visible so tests can drive syncRegistrations() against a real service. + volatile WatchService watchService; + + private volatile ScheduledExecutorService reconciler; + + @Override + public String type() { + return TYPE; + } + + @Override + public void validate(Policy policy) { + if (watchDirsOf(policy).isEmpty()) { + throw new IllegalArgumentException( + "folder-watch trigger requires at least one watchable (folder) input source"); + } + } + + @Override + public synchronized void start() { + if (watchService != null) { + return; + } + try { + watchService = FileSystems.getDefault().newWatchService(); + } catch (IOException e) { + log.error("Could not start folder-watch trigger: {}", e.getMessage(), e); + return; + } + running = true; + Thread.ofVirtual().name("policy-folder-watch").start(this::watchLoop); + long reconcileSeconds = applicationProperties.getPolicies().getWatchReconcileSeconds(); + reconciler = + Executors.newSingleThreadScheduledExecutor( + Thread.ofVirtual().name("policy-folder-reconcile-", 0).factory()); + // First reconcile runs immediately so pre-existing files are picked up at startup. + reconciler.scheduleAtFixedRate(this::safeReconcile, 0, reconcileSeconds, TimeUnit.SECONDS); + log.info("Folder-watch trigger started (reconcile every {}s)", reconcileSeconds); + } + + @Override + public synchronized void stop() { + running = false; + if (reconciler != null) { + reconciler.shutdownNow(); + reconciler = null; + } + if (watchService != null) { + try { + watchService.close(); // wakes the watch loop with ClosedWatchServiceException + } catch (IOException e) { + log.debug("Error closing folder watch service: {}", e.getMessage()); + } + watchService = null; + } + keysByDir.clear(); + dirByKey.clear(); + } + + private void watchLoop() { + // Capture once: stop() may null the field; close() still wakes take()/poll() on this local. + WatchService watcher = watchService; + if (watcher == null) { + return; + } + while (running) { + WatchKey first; + try { + first = watcher.take(); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + return; + } catch (ClosedWatchServiceException e) { + return; + } + runForChangedDirs(drainBurst(watcher, first)); + } + } + + /** + * Coalesce a burst of file-system events into one set of affected directories: drain everything + * arriving within the quiet period. Event kinds are irrelevant; any event means "go look". + */ + private Set drainBurst(WatchService watcher, WatchKey first) { + long quietPeriodMs = applicationProperties.getPolicies().getWatchQuietPeriodMs(); + Set changed = new HashSet<>(); + WatchKey key = first; + while (key != null) { + key.pollEvents(); + Path dir = dirByKey.get(key); + if (dir != null) { + changed.add(dir); + } + key.reset(); + try { + key = watcher.poll(quietPeriodMs, TimeUnit.MILLISECONDS); + } catch (ClosedWatchServiceException | InterruptedException e) { + break; + } + } + return changed; + } + + /** Run every folder-watch policy that draws from one of the changed directories. */ + void runForChangedDirs(Set changedDirs) { + if (changedDirs.isEmpty()) { + return; + } + for (Policy policy : policyStore.findByTriggerType(TYPE)) { + List dirs; + try { + dirs = watchDirsOf(policy); + } catch (RuntimeException e) { + log.warn( + "Folder-watch policy {} is misconfigured: {}", policy.id(), e.getMessage()); + continue; + } + if (dirs.stream().anyMatch(changedDirs::contains)) { + log.debug("Folder-watch policy {} ({}) saw activity", policy.id(), policy.name()); + policyRunner.run(policy); + } + } + } + + private void safeReconcile() { + try { + syncRegistrations(); + runAll(); + } catch (RuntimeException e) { + log.error("Folder-watch reconcile failed: {}", e.getMessage(), e); + } + } + + /** Reconcile safety net: run every folder-watch policy regardless of watch events. */ + void runAll() { + for (Policy policy : policyStore.findByTriggerType(TYPE)) { + try { + policyRunner.run(policy); + } catch (RuntimeException e) { + log.warn( + "Folder-watch reconcile run failed for policy {}: {}", + policy.id(), + e.getMessage()); + } + } + } + + /** Register newly-wanted dirs that exist on disk, cancel ones no longer wanted. */ + synchronized void syncRegistrations() { + if (watchService == null) { + return; + } + Set desired = desiredDirs(); + + keysByDir + .entrySet() + .removeIf( + entry -> { + if (desired.contains(entry.getKey())) { + return false; + } + entry.getValue().cancel(); + dirByKey.remove(entry.getValue()); + return true; + }); + + for (Path dir : desired) { + if (keysByDir.containsKey(dir)) { + continue; + } + try { + WatchKey key = dir.register(watchService, ENTRY_CREATE, ENTRY_MODIFY); + keysByDir.put(dir, key); + dirByKey.put(key, dir); + log.info("Watching {} for folder-watch policies", dir); + } catch (IOException | RuntimeException e) { + log.warn("Could not watch {}: {}", dir, e.getMessage()); + } + } + } + + /** The directories currently registered with the watch service. Visible for tests. */ + Set watchedDirs() { + return Set.copyOf(keysByDir.keySet()); + } + + /** Every existing directory any current folder-watch policy wants watched. */ + private Set desiredDirs() { + Set dirs = new HashSet<>(); + for (Policy policy : policyStore.findByTriggerType(TYPE)) { + try { + for (Path dir : watchDirsOf(policy)) { + if (Files.isDirectory(dir)) { + dirs.add(dir); + } + } + } catch (RuntimeException e) { + log.warn( + "Folder-watch policy {} is misconfigured: {}", policy.id(), e.getMessage()); + } + } + return dirs; + } + + // Absolute + normalised so registration keys and event-time matching compare regardless of how + // the path was configured. + private List watchDirsOf(Policy policy) { + List dirs = new ArrayList<>(); + for (InputSpec spec : policy.sources()) { + InputSource source = sourceFor(spec); + if (source == null) { + continue; + } + for (Path dir : source.watchTargets(spec)) { + dirs.add(dir.toAbsolutePath().normalize()); + } + } + return dirs; + } + + private InputSource sourceFor(InputSpec spec) { + return inputSources.stream() + .filter(source -> source.supports(spec)) + .findFirst() + .orElse(null); + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/trigger/PolicyTrigger.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/trigger/PolicyTrigger.java new file mode 100644 index 0000000000..a9cb363719 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/trigger/PolicyTrigger.java @@ -0,0 +1,23 @@ +package stirling.software.proprietary.policy.trigger; + +import stirling.software.proprietary.policy.model.Policy; + +/** + * Decides when a policy runs. On firing it hands the policy to {@code PolicyRunner}; it + * never resolves sources itself. New trigger kinds are just new beans of this type. + */ +public interface PolicyTrigger { + + /** Matches {@code TriggerConfig.type()}. */ + String type(); + + /** + * Validate at save time so misconfiguration fails fast, not at fire time. Receives the whole + * {@link Policy} so triggers that depend on the policy's sources (folder-watch) can check that. + */ + default void validate(Policy policy) {} + + default void start() {} + + default void stop() {} +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/trigger/PolicyTriggerManager.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/trigger/PolicyTriggerManager.java new file mode 100644 index 0000000000..c1175eeb92 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/trigger/PolicyTriggerManager.java @@ -0,0 +1,51 @@ +package stirling.software.proprietary.policy.trigger; + +import java.util.List; + +import org.springframework.context.SmartLifecycle; +import org.springframework.context.annotation.Profile; +import org.springframework.stereotype.Service; + +import lombok.RequiredArgsConstructor; +import lombok.extern.slf4j.Slf4j; + +/** Starts and stops every {@link PolicyTrigger} with the application lifecycle. */ +@Slf4j +@Service +@RequiredArgsConstructor +@Profile("saas") +public class PolicyTriggerManager implements SmartLifecycle { + + private final List triggers; + + private volatile boolean running; + + @Override + public void start() { + for (PolicyTrigger trigger : triggers) { + try { + trigger.start(); + } catch (RuntimeException e) { + log.error("Failed to start trigger '{}': {}", trigger.type(), e.getMessage(), e); + } + } + running = true; + } + + @Override + public void stop() { + for (PolicyTrigger trigger : triggers) { + try { + trigger.stop(); + } catch (RuntimeException e) { + log.error("Failed to stop trigger '{}': {}", trigger.type(), e.getMessage(), e); + } + } + running = false; + } + + @Override + public boolean isRunning() { + return running; + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/policy/trigger/ScheduleTrigger.java b/app/proprietary/src/main/java/stirling/software/proprietary/policy/trigger/ScheduleTrigger.java new file mode 100644 index 0000000000..708ca8128b --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/policy/trigger/ScheduleTrigger.java @@ -0,0 +1,151 @@ +package stirling.software.proprietary.policy.trigger; + +import java.time.Instant; +import java.time.ZoneId; +import java.time.ZoneOffset; +import java.time.ZonedDateTime; +import java.util.Map; +import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.Executors; +import java.util.concurrent.ScheduledExecutorService; +import java.util.concurrent.TimeUnit; + +import org.springframework.context.annotation.Profile; +import org.springframework.stereotype.Service; + +import lombok.RequiredArgsConstructor; +import lombok.extern.slf4j.Slf4j; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.proprietary.policy.engine.PolicyRunner; +import stirling.software.proprietary.policy.model.Policy; +import stirling.software.proprietary.policy.model.Schedule; +import stirling.software.proprietary.policy.store.PolicyStore; + +import tools.jackson.databind.ObjectMapper; + +/** + * Fires policies on a {@link Schedule}: a fixed-interval sweep runs each due "schedule" policy. + * + *

Last-fire times are in memory, so this assumes a single node and resets on restart. + */ +@Slf4j +@Service +@RequiredArgsConstructor +@Profile("saas") +public class ScheduleTrigger implements PolicyTrigger { + + private static final String TYPE = "schedule"; + + private final PolicyStore policyStore; + private final PolicyRunner policyRunner; + private final ObjectMapper objectMapper; + private final ApplicationProperties applicationProperties; + + private final Map lastFiredByPolicy = new ConcurrentHashMap<>(); + private volatile ScheduledExecutorService scheduler; + + @Override + public String type() { + return TYPE; + } + + @Override + public void validate(Policy policy) { + ScheduleConfig.from(objectMapper, policy.trigger().options()); + } + + @Override + public synchronized void start() { + if (scheduler != null) { + return; + } + long sweepSeconds = applicationProperties.getPolicies().getScheduleSweepSeconds(); + scheduler = + Executors.newSingleThreadScheduledExecutor( + Thread.ofVirtual().name("policy-schedule-", 0).factory()); + scheduler.scheduleAtFixedRate( + this::safeSweep, sweepSeconds, sweepSeconds, TimeUnit.SECONDS); + log.info("Schedule trigger started (sweep every {}s)", sweepSeconds); + } + + @Override + public synchronized void stop() { + if (scheduler != null) { + scheduler.shutdownNow(); + scheduler = null; + } + } + + private void safeSweep() { + try { + sweep(Instant.now()); + } catch (RuntimeException e) { + log.error("Schedule sweep failed: {}", e.getMessage(), e); + } + } + + /** Fire every scheduled policy that is due as of {@code now}. Package-visible for testing. */ + void sweep(Instant now) { + for (Policy policy : policyStore.findByTriggerType(TYPE)) { + ScheduleConfig config; + try { + config = ScheduleConfig.from(objectMapper, policy.trigger().options()); + } catch (IllegalArgumentException e) { + log.warn("Scheduled policy {} is misconfigured: {}", policy.id(), e.getMessage()); + continue; + } + + // Baseline a newly-seen policy to now so it does not fire immediately. + Instant last = lastFiredByPolicy.computeIfAbsent(policy.id(), id -> now); + ZonedDateTime next = config.schedule().nextAfter(last.atZone(config.zone())); + if (!next.toInstant().isAfter(now)) { + lastFiredByPolicy.put(policy.id(), now); + log.info("Scheduled policy {} ({}) is due", policy.id(), policy.name()); + policyRunner.run(policy); + } + } + } + + /** + * Validated schedule-trigger options: the {@link Schedule} and the zone it runs in (UTC by + * default). + */ + record ScheduleConfig(Schedule schedule, ZoneId zone) { + + private static final String SCHEDULE_OPTION = "schedule"; + private static final String ZONE_OPTION = "zone"; + + static ScheduleConfig from(ObjectMapper mapper, Map options) { + Object scheduleNode = options.get(SCHEDULE_OPTION); + if (scheduleNode == null) { + throw new IllegalArgumentException("schedule trigger requires a 'schedule'"); + } + Schedule schedule; + try { + schedule = mapper.convertValue(scheduleNode, Schedule.class); + } catch (RuntimeException e) { + throw new IllegalArgumentException("invalid schedule: " + rootMessage(e), e); + } + + ZoneId zone = ZoneOffset.UTC; + Object zoneNode = options.get(ZONE_OPTION); + if (zoneNode != null && !zoneNode.toString().isBlank()) { + try { + zone = ZoneId.of(zoneNode.toString()); + } catch (RuntimeException e) { + throw new IllegalArgumentException("invalid zone '" + zoneNode + "'"); + } + } + return new ScheduleConfig(schedule, zone); + } + + private static String rootMessage(Throwable t) { + Throwable cause = t; + while (cause.getCause() != null) { + cause = cause.getCause(); + } + return cause.getMessage(); + } + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/security/InitialSecuritySetup.java b/app/proprietary/src/main/java/stirling/software/proprietary/security/InitialSecuritySetup.java index e0dabced4b..94dc41a231 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/security/InitialSecuritySetup.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/security/InitialSecuritySetup.java @@ -1,11 +1,13 @@ package stirling.software.proprietary.security; import java.sql.SQLException; +import java.util.Arrays; import java.util.List; import java.util.Optional; import java.util.UUID; import org.springframework.beans.factory.annotation.Value; +import org.springframework.core.env.Environment; import org.springframework.stereotype.Component; import jakarta.annotation.PostConstruct; @@ -37,6 +39,17 @@ public class InitialSecuritySetup { private final ApplicationProperties applicationProperties; private final DatabaseServiceInterface databaseService; private final UserLicenseSettingsService licenseSettingsService; + private final Environment environment; + + /** + * SaaS manages identity in Supabase and billing via PAYG, so the self-host bootstrap steps that + * scan/rewrite the whole user table (default-team backfill, seat-license grandfathering) don't + * apply - and against a large SaaS user table they stall startup with full-table loads + + * per-row saveAll. Per-user team assignment happens in SupabaseAuthenticationFilter instead. + */ + private boolean isSaas() { + return Arrays.asList(environment.getActiveProfiles()).contains("saas"); + } @PostConstruct public void init() { @@ -51,9 +64,15 @@ public class InitialSecuritySetup { } configureJWTSettings(); - assignUsersToDefaultTeamIfMissing(); initializeInternalApiUser(); - initializeUserLicenseSettings(); + if (isSaas()) { + log.info( + "SaaS profile active - skipping self-host user-table bootstrap" + + " (default-team backfill, seat-license grandfathering)."); + } else { + assignUsersToDefaultTeamIfMissing(); + initializeUserLicenseSettings(); + } } catch (IllegalArgumentException | SQLException | UnsupportedProviderException e) { log.error("Failed to initialize security setup.", e); System.exit(1); diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/security/configuration/DatabaseConfig.java b/app/proprietary/src/main/java/stirling/software/proprietary/security/configuration/DatabaseConfig.java index d81bf6a360..454df5f67e 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/security/configuration/DatabaseConfig.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/security/configuration/DatabaseConfig.java @@ -31,13 +31,15 @@ import stirling.software.common.model.exception.UnsupportedProviderException; "stirling.software.proprietary.security.repository", "stirling.software.proprietary.repository", "stirling.software.proprietary.storage.repository", - "stirling.software.proprietary.workflow.repository" + "stirling.software.proprietary.workflow.repository", + "stirling.software.proprietary.policy.store" }) @EntityScan({ "stirling.software.proprietary.security.model", "stirling.software.proprietary.model", "stirling.software.proprietary.storage.model", - "stirling.software.proprietary.workflow.model" + "stirling.software.proprietary.workflow.model", + "stirling.software.proprietary.policy.store" }) public class DatabaseConfig { diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/security/configuration/ee/LicenseKeyChecker.java b/app/proprietary/src/main/java/stirling/software/proprietary/security/configuration/ee/LicenseKeyChecker.java index af9ac192d7..d8dbfc0e08 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/security/configuration/ee/LicenseKeyChecker.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/security/configuration/ee/LicenseKeyChecker.java @@ -3,7 +3,6 @@ package stirling.software.proprietary.security.configuration.ee; import java.io.IOException; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import org.springframework.boot.context.event.ApplicationReadyEvent; import org.springframework.context.annotation.Lazy; @@ -104,7 +103,7 @@ public class LicenseKeyChecker { if (keyOrFilePath.startsWith(FILE_PREFIX)) { String filePath = keyOrFilePath.substring(FILE_PREFIX.length()); try { - Path path = Paths.get(filePath); + Path path = Path.of(filePath); if (!Files.exists(path)) { log.error("License file does not exist: {}", filePath); return null; diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/security/controller/api/AdminLicenseController.java b/app/proprietary/src/main/java/stirling/software/proprietary/security/controller/api/AdminLicenseController.java index e9cd65f7ca..9e05e9ebbc 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/security/controller/api/AdminLicenseController.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/security/controller/api/AdminLicenseController.java @@ -4,7 +4,6 @@ import java.io.IOException; import java.nio.charset.StandardCharsets; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import java.nio.file.StandardCopyOption; import java.util.HashMap; import java.util.Map; @@ -329,7 +328,7 @@ public class AdminLicenseController { } // Get config directory and target path - Path configPath = Paths.get(InstallationPathConfig.getConfigPath()); + Path configPath = Path.of(InstallationPathConfig.getConfigPath()); Path configPathAbs = configPath.toAbsolutePath().normalize(); Path targetPath = configPathAbs.resolve(filename).normalize(); log.info( diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/security/controller/api/AdminSettingsController.java b/app/proprietary/src/main/java/stirling/software/proprietary/security/controller/api/AdminSettingsController.java index 115727eb9d..14ddde602a 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/security/controller/api/AdminSettingsController.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/security/controller/api/AdminSettingsController.java @@ -618,6 +618,8 @@ public class AdminSettingsController { case "autopipeline", "autoPipeline" -> applicationProperties.getAutoPipeline(); case "legal" -> applicationProperties.getLegal(); case "telegram" -> applicationProperties.getTelegram(); + case "aiengine", "aiEngine" -> applicationProperties.getAiEngine(); + case "mcp" -> applicationProperties.getMcp(); default -> null; }; } @@ -641,7 +643,10 @@ public class AdminSettingsController { "autoPipeline", "autopipeline", "legal", - "telegram"); + "telegram", + "aiEngine", + "aiengine", + "mcp"); // Pattern to validate safe property paths - only alphanumeric, dots, and underscores private static final Pattern SAFE_KEY_PATTERN = @@ -697,7 +702,7 @@ public class AdminSettingsController { if (path != null && !path.trim().isEmpty()) { try { java.nio.file.Path normalized = - java.nio.file.Paths.get(path.trim()).toAbsolutePath().normalize(); + java.nio.file.Path.of(path.trim()).toAbsolutePath().normalize(); String normalizedStr = normalized.toString(); // Check for duplicates @@ -714,9 +719,9 @@ public class AdminSettingsController { // Check for overlapping paths java.util.List pathList = new java.util.ArrayList<>(normalizedPaths); for (int i = 0; i < pathList.size(); i++) { - java.nio.file.Path path1 = java.nio.file.Paths.get(pathList.get(i)); + java.nio.file.Path path1 = java.nio.file.Path.of(pathList.get(i)); for (int j = i + 1; j < pathList.size(); j++) { - java.nio.file.Path path2 = java.nio.file.Paths.get(pathList.get(j)); + java.nio.file.Path path2 = java.nio.file.Path.of(pathList.get(j)); if (path1.startsWith(path2) || path2.startsWith(path1)) { return "Overlapping paths detected: " + path1 + " and " + path2; } diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/security/controller/api/InviteLinkController.java b/app/proprietary/src/main/java/stirling/software/proprietary/security/controller/api/InviteLinkController.java index 07ec1c9736..641216e269 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/security/controller/api/InviteLinkController.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/security/controller/api/InviteLinkController.java @@ -3,7 +3,6 @@ package stirling.software.proprietary.security.controller.api; import java.security.Principal; import java.time.LocalDateTime; import java.util.*; -import java.util.stream.Collectors; import org.springframework.http.HttpStatus; import org.springframework.http.ResponseEntity; @@ -280,7 +279,7 @@ public class InviteLinkController { "expiresAt", invite.getExpiresAt().toString()); return inviteMap; }) - .collect(Collectors.toList()); + .toList(); return ResponseEntity.ok(Map.of("invites", inviteList)); @@ -331,7 +330,7 @@ public class InviteLinkController { List expiredInvites = inviteTokenRepository.findAll().stream() .filter(invite -> !invite.isValid()) - .collect(Collectors.toList()); + .toList(); int count = expiredInvites.size(); inviteTokenRepository.deleteAll(expiredInvites); diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/security/controller/api/UIDataTessdataController.java b/app/proprietary/src/main/java/stirling/software/proprietary/security/controller/api/UIDataTessdataController.java index 743425aec8..8212dcfb29 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/security/controller/api/UIDataTessdataController.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/security/controller/api/UIDataTessdataController.java @@ -6,7 +6,6 @@ import java.net.HttpURLConnection; import java.net.URL; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import java.nio.file.StandardCopyOption; import java.util.*; import java.util.regex.Pattern; @@ -53,7 +52,7 @@ public class UIDataTessdataController { TessdataLanguagesResponse response = new TessdataLanguagesResponse(); response.setInstalled(getAvailableTesseractLanguages()); response.setAvailable(getRemoteTessdataLanguages()); - response.setWritable(isWritableDirectory(Paths.get(runtimePathConfig.getTessDataPath()))); + response.setWritable(isWritableDirectory(Path.of(runtimePathConfig.getTessDataPath()))); return ResponseEntity.ok(response); } @@ -67,7 +66,7 @@ public class UIDataTessdataController { .body(Map.of("message", "No languages provided for download")); } - Path tessdataDir = Paths.get(runtimePathConfig.getTessDataPath()); + Path tessdataDir = Path.of(runtimePathConfig.getTessDataPath()); try { Files.createDirectories(tessdataDir); } catch (IOException e) { diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/security/controller/api/UserController.java b/app/proprietary/src/main/java/stirling/software/proprietary/security/controller/api/UserController.java index 90abd5f067..d8500e475f 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/security/controller/api/UserController.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/security/controller/api/UserController.java @@ -978,27 +978,56 @@ public class UserController { } } - /** - * List all enabled users for selection in signing workflows. - * - * @param principal The authenticated user - * @return List of user summaries - */ + // Lists enabled users for the signing picker; 'org' scope = instance-wide, else caller's team. @GetMapping("/users") public ResponseEntity> listUsers(Principal principal) { if (principal == null) { return ResponseEntity.status(HttpStatus.UNAUTHORIZED).build(); } + Optional callerOpt = userService.findByUsernameIgnoreCase(principal.getName()); + + // Anonymous (SaaS) accounts must never enumerate users, in any scope or team. + if (callerOpt.map(UserController::isAnonymousUser).orElse(false)) { + return ResponseEntity.status(HttpStatus.FORBIDDEN).build(); + } + + // Fail-closed: only literal "org" opens the whole instance; anything else scopes to team. + String scope = applicationProperties.getStorage().getSigning().getUserListScope(); + boolean teamScoped = !"org".equalsIgnoreCase(scope == null ? "" : scope.trim()); + + List source; + if (teamScoped) { + Team callerTeam = callerOpt.map(User::getTeam).orElse(null); + if (callerTeam == null || isSystemTeam(callerTeam)) { + // No team or a shared system team: return only the caller, not the team's members. + source = callerOpt.map(List::of).orElse(List.of()); + } else { + // Scopes via the single User.team FK; revisit if multi-team membership is added. + source = userRepository.findAllByTeamId(callerTeam.getId()); + } + } else { + source = userRepository.findAll(); + } + List users = - userRepository.findAll().stream() - .filter(User::isEnabled) - .map(this::toUserSummaryDTO) - .collect(java.util.stream.Collectors.toList()); + source.stream().filter(User::isEnabled).map(this::toUserSummaryDTO).toList(); return ResponseEntity.ok(users); } + // SaaS anonymous accounts, which must not enumerate users. + private static boolean isAnonymousUser(User user) { + return AuthenticationType.ANONYMOUS.name().equalsIgnoreCase(user.getAuthenticationType()); + } + + // System teams (Default/Internal) are not enumerable through the signing picker. + private static boolean isSystemTeam(Team team) { + String name = team.getName(); + return TeamService.DEFAULT_TEAM_NAME.equalsIgnoreCase(name) + || TeamService.INTERNAL_TEAM_NAME.equalsIgnoreCase(name); + } + private UserSummaryDTO toUserSummaryDTO(User user) { return new UserSummaryDTO( user.getId(), diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/security/database/repository/UserRepository.java b/app/proprietary/src/main/java/stirling/software/proprietary/security/database/repository/UserRepository.java index 6cdfb57ea8..4e592e427a 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/security/database/repository/UserRepository.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/security/database/repository/UserRepository.java @@ -105,15 +105,6 @@ public interface UserRepository extends JpaRepository { Stream findByUsernameIsNullAndCreatedAtBefore( @Param("cutoffDate") LocalDateTime cutoffDate); - /** Users with an API key but no row in {@code user_credits}. */ - @Query( - value = - "SELECT u.* FROM users u " - + "LEFT JOIN user_credits uc ON uc.user_id = u.user_id " - + "WHERE u.api_key IS NOT NULL AND uc.user_id IS NULL", - nativeQuery = true) - List findUsersWithApiKeyButNoCredits(); - /** Single-shot UPDATE that reassigns a user to a different team. */ @Modifying @Query("UPDATE User u SET u.team.id = :teamId WHERE u.id = :userId") diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/security/model/User.java b/app/proprietary/src/main/java/stirling/software/proprietary/security/model/User.java index e39b1a3306..2b2d22cfc7 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/security/model/User.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/security/model/User.java @@ -87,7 +87,11 @@ public class User implements UserDetails, Serializable { private String email; // SaaS-only: Supabase user UUID. Null in OSS / proprietary deployments. - @Column(name = "supabase_id", unique = true) + // Column is `supabase_auth_id` (canonical name from the initial Supabase remote + // schema migration). An earlier Flyway V2 (PR #6384) accidentally introduced a + // parallel `supabase_id` column that was used by Java; V17 backfilled and dropped + // it. Field name is kept as `supabaseId` to avoid a wide refactor of callers. + @Column(name = "supabase_auth_id", unique = true) private UUID supabaseId; @OneToMany(fetch = FetchType.EAGER, cascade = CascadeType.ALL, mappedBy = "user") diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/security/service/DatabaseService.java b/app/proprietary/src/main/java/stirling/software/proprietary/security/service/DatabaseService.java index f120e4e42e..2246903d2d 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/security/service/DatabaseService.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/security/service/DatabaseService.java @@ -4,7 +4,6 @@ import java.io.IOException; import java.nio.file.DirectoryStream; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import java.nio.file.StandardCopyOption; import java.nio.file.attribute.BasicFileAttributes; import java.security.MessageDigest; @@ -23,7 +22,6 @@ import java.util.Comparator; import java.util.List; import java.util.UUID; import java.util.regex.Pattern; -import java.util.stream.Collectors; import javax.sql.DataSource; @@ -100,7 +98,7 @@ public class DatabaseService implements DatabaseServiceInterface { ApplicationProperties.Datasource datasourceProps, DataSource dataSource, DatabaseNotificationServiceInterface backupNotificationService) { - this.BACKUP_DIR = Paths.get(InstallationPathConfig.getBackupPath()).normalize(); + this.BACKUP_DIR = Path.of(InstallationPathConfig.getBackupPath()).normalize(); this.datasourceProps = datasourceProps; this.dataSource = dataSource; this.backupNotificationService = backupNotificationService; @@ -111,7 +109,7 @@ public class DatabaseService implements DatabaseServiceInterface { @Deprecated(since = "2.0.0", forRemoval = true) private void moveBackupFiles() { Path sourceDir = - Paths.get(InstallationPathConfig.getConfigPath(), "db", "backup").normalize(); + Path.of(InstallationPathConfig.getConfigPath(), "db", "backup").normalize(); if (!Files.exists(sourceDir)) { log.info("Source directory does not exist: {}", sourceDir); @@ -217,7 +215,7 @@ public class DatabaseService implements DatabaseServiceInterface { List backupList = this.getBackupList(); backupList.sort(Comparator.comparing(FileInfo::getModificationDate).reversed()); - Path latestExport = Paths.get(backupList.get(0).getFilePath()); + Path latestExport = Path.of(backupList.get(0).getFilePath()); executeDatabaseScript(latestExport); } @@ -258,9 +256,13 @@ public class DatabaseService implements DatabaseServiceInterface { @Override public void exportDatabase() { List filteredBackupList = - this.getBackupList().stream() - .filter(backup -> !backup.getFileName().startsWith(BACKUP_PREFIX + "user_")) - .collect(Collectors.toList()); + new ArrayList<>( + this.getBackupList().stream() + .filter( + backup -> + !backup.getFileName() + .startsWith(BACKUP_PREFIX + "user_")) + .toList()); if (filteredBackupList.size() > 5) { deleteOldestBackup(filteredBackupList); @@ -341,7 +343,7 @@ public class DatabaseService implements DatabaseServiceInterface { for (FileInfo backup : backupList) { try { - Files.deleteIfExists(Paths.get(backup.getFilePath())); + Files.deleteIfExists(Path.of(backup.getFilePath())); deletedFiles.add(Pair.of(backup, true)); } catch (IOException e) { log.error("Error deleting backup file: {}", backup.getFileName(), e); @@ -359,7 +361,7 @@ public class DatabaseService implements DatabaseServiceInterface { if (!backupList.isEmpty()) { FileInfo lastBackup = backupList.get(backupList.size() - 1); try { - Files.deleteIfExists(Paths.get(lastBackup.getFilePath())); + Files.deleteIfExists(Path.of(lastBackup.getFilePath())); deletedFiles.add(Pair.of(lastBackup, true)); } catch (IOException e) { log.error("Error deleting last backup file: {}", lastBackup.getFileName(), e); @@ -381,7 +383,7 @@ public class DatabaseService implements DatabaseServiceInterface { p -> p.getFileName().substring(7, p.getFileName().length() - 4))); FileInfo oldestFile = filteredBackupList.get(0); - Files.deleteIfExists(Paths.get(oldestFile.getFilePath())); + Files.deleteIfExists(Path.of(oldestFile.getFilePath())); log.info("Deleted oldest backup: {}", oldestFile.getFileName()); } catch (IOException e) { log.error("Unable to delete oldest backup, message: {}", e.getMessage(), e); diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/security/service/KeyPairCleanupService.java b/app/proprietary/src/main/java/stirling/software/proprietary/security/service/KeyPairCleanupService.java index aec455a929..8745f505fd 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/security/service/KeyPairCleanupService.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/security/service/KeyPairCleanupService.java @@ -3,7 +3,6 @@ package stirling.software.proprietary.security.service; import java.io.IOException; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import java.time.LocalDateTime; import java.util.List; import java.util.concurrent.TimeUnit; @@ -76,7 +75,7 @@ public class KeyPairCleanupService { return; } - Path privateKeyDirectory = Paths.get(InstallationPathConfig.getPrivateKeyPath()); + Path privateKeyDirectory = Path.of(InstallationPathConfig.getPrivateKeyPath()); Path keyFile = privateKeyDirectory.resolve(keyId + KeyPersistenceService.KEY_SUFFIX); if (Files.exists(keyFile)) { diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/security/service/KeyPersistenceService.java b/app/proprietary/src/main/java/stirling/software/proprietary/security/service/KeyPersistenceService.java index d0c9f879be..49e66e5c53 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/security/service/KeyPersistenceService.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/security/service/KeyPersistenceService.java @@ -3,7 +3,6 @@ package stirling.software.proprietary.security.service; import java.io.IOException; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import java.security.KeyFactory; import java.security.KeyPair; import java.security.KeyPairGenerator; @@ -20,7 +19,6 @@ import java.time.format.DateTimeFormatter; import java.util.Base64; import java.util.List; import java.util.Optional; -import java.util.stream.Collectors; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.cache.Cache; @@ -82,7 +80,7 @@ public class KeyPersistenceService implements KeyPersistenceServiceInterface { */ private void loadExistingKeysFromDisk() { try { - Path keyDirectory = Paths.get(InstallationPathConfig.getPrivateKeyPath()); + Path keyDirectory = Path.of(InstallationPathConfig.getPrivateKeyPath()); if (!Files.exists(keyDirectory)) { log.info("No existing keys found, generating new keypair"); @@ -99,7 +97,7 @@ public class KeyPersistenceService implements KeyPersistenceServiceInterface { b.getFileName().compareTo(a.getFileName())) // Most // recent // first - .collect(Collectors.toList()); + .toList(); } if (keyFiles.isEmpty()) { @@ -275,7 +273,7 @@ public class KeyPersistenceService implements KeyPersistenceServiceInterface { eligible); return eligible; }) - .collect(Collectors.toList()); + .toList(); } private String generateKeyId() { @@ -284,20 +282,17 @@ public class KeyPersistenceService implements KeyPersistenceServiceInterface { } private KeyPair generateRSAKeypair() { - KeyPairGenerator keyPairGenerator = null; - try { - keyPairGenerator = KeyPairGenerator.getInstance("RSA"); + KeyPairGenerator keyPairGenerator = KeyPairGenerator.getInstance("RSA"); keyPairGenerator.initialize(2048); + return keyPairGenerator.generateKeyPair(); } catch (NoSuchAlgorithmException e) { - log.error("Failed to initialize RSA key pair generator", e); + throw new IllegalStateException("RSA key pair generator is not available", e); } - - return keyPairGenerator.generateKeyPair(); } private void ensurePrivateKeyDirectoryExists() throws IOException { - Path keyPath = Paths.get(InstallationPathConfig.getPrivateKeyPath()); + Path keyPath = Path.of(InstallationPathConfig.getPrivateKeyPath()); if (!Files.exists(keyPath)) { Files.createDirectories(keyPath); @@ -312,7 +307,7 @@ public class KeyPersistenceService implements KeyPersistenceServiceInterface { *

Public key stored as: keyId.pub */ private void storeKeyPair(String keyId, KeyPair keyPair) throws IOException { - Path keyDirectory = Paths.get(InstallationPathConfig.getPrivateKeyPath()); + Path keyDirectory = Path.of(InstallationPathConfig.getPrivateKeyPath()); // Store private key Path privateKeyFile = keyDirectory.resolve(keyId + KEY_SUFFIX); @@ -346,7 +341,7 @@ public class KeyPersistenceService implements KeyPersistenceServiceInterface { private PrivateKey loadPrivateKey(String keyId) throws IOException, NoSuchAlgorithmException, InvalidKeySpecException { Path keyFile = - Paths.get(InstallationPathConfig.getPrivateKeyPath()).resolve(keyId + KEY_SUFFIX); + Path.of(InstallationPathConfig.getPrivateKeyPath()).resolve(keyId + KEY_SUFFIX); if (!Files.exists(keyFile)) { throw new IOException("Private key not found: " + keyFile); @@ -369,8 +364,7 @@ public class KeyPersistenceService implements KeyPersistenceServiceInterface { */ private String loadPublicKey(String keyId) throws IOException { Path publicKeyFile = - Paths.get(InstallationPathConfig.getPrivateKeyPath()) - .resolve(keyId + PUB_KEY_SUFFIX); + Path.of(InstallationPathConfig.getPrivateKeyPath()).resolve(keyId + PUB_KEY_SUFFIX); if (!Files.exists(publicKeyFile)) { throw new IOException("Public key not found: " + publicKeyFile); diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/security/service/LoginAttemptService.java b/app/proprietary/src/main/java/stirling/software/proprietary/security/service/LoginAttemptService.java index d715771b59..d3819c35d4 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/security/service/LoginAttemptService.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/security/service/LoginAttemptService.java @@ -5,7 +5,6 @@ import java.util.Locale; import java.util.Map; import java.util.concurrent.ConcurrentHashMap; import java.util.concurrent.TimeUnit; -import java.util.stream.Collectors; import org.springframework.stereotype.Service; @@ -101,7 +100,7 @@ public class LoginAttemptService { return attemptsCache.entrySet().stream() .filter(entry -> entry.getValue().getAttemptCount() >= MAX_ATTEMPT) .map(Map.Entry::getKey) - .collect(Collectors.toList()); + .toList(); } public int getRemainingAttempts(String key) { diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/security/service/RefreshRateLimitService.java b/app/proprietary/src/main/java/stirling/software/proprietary/security/service/RefreshRateLimitService.java index bb4e454290..6c4321ed7e 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/security/service/RefreshRateLimitService.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/security/service/RefreshRateLimitService.java @@ -31,21 +31,14 @@ public class RefreshRateLimitService { this.jwtProperties = applicationProperties.getSecurity().getJwt(); } - private static class RefreshAttempt { - private final AtomicInteger count = new AtomicInteger(0); - private final Instant firstAttempt = Instant.now(); + private record RefreshAttempt(AtomicInteger count, Instant firstAttempt) { + RefreshAttempt() { + this(new AtomicInteger(0), Instant.now()); + } int incrementAndGet() { return count.incrementAndGet(); } - - Instant getFirstAttempt() { - return firstAttempt; - } - - int getCount() { - return count.get(); - } } private final Map attempts = new ConcurrentHashMap<>(); @@ -72,7 +65,7 @@ public class RefreshRateLimitService { // Clean up if outside grace window Instant cutoff = Instant.now().minusMillis(graceWindowMillis); - if (attempt.getFirstAttempt().isBefore(cutoff)) { + if (attempt.firstAttempt().isBefore(cutoff)) { attempts.remove(tokenHash); } @@ -98,15 +91,9 @@ public class RefreshRateLimitService { ? configuredMinutes : JwtConstants.DEFAULT_REFRESH_GRACE_MINUTES; Instant cutoff = Instant.now().minusMillis(graceMinutes * 60000L); - int removed = - attempts.entrySet().stream() - .filter(entry -> entry.getValue().getFirstAttempt().isBefore(cutoff)) - .mapToInt( - entry -> { - attempts.remove(entry.getKey()); - return 1; - }) - .sum(); + int beforeSize = attempts.size(); + attempts.entrySet().removeIf(entry -> entry.getValue().firstAttempt().isBefore(cutoff)); + int removed = beforeSize - attempts.size(); if (removed > 0) { log.debug("Cleaned up {} expired refresh tracking entries", removed); diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/service/AiEngineClient.java b/app/proprietary/src/main/java/stirling/software/proprietary/service/AiEngineClient.java index dedc558aff..b0d726ceca 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/service/AiEngineClient.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/service/AiEngineClient.java @@ -25,6 +25,7 @@ public class AiEngineClient { private final ApplicationProperties applicationProperties; private final HttpClient httpClient; + private final String engineSharedSecret; @Autowired public AiEngineClient(ApplicationProperties applicationProperties) { @@ -39,8 +40,17 @@ public class AiEngineClient { /** Package-private constructor that accepts an HttpClient directly; intended for tests. */ AiEngineClient(ApplicationProperties applicationProperties, HttpClient httpClient) { + this(applicationProperties, httpClient, System.getenv("STIRLING_ENGINE_SHARED_SECRET")); + } + + /** Package-private constructor that also injects the engine shared secret; for tests. */ + AiEngineClient( + ApplicationProperties applicationProperties, + HttpClient httpClient, + String engineSharedSecret) { this.applicationProperties = applicationProperties; this.httpClient = httpClient; + this.engineSharedSecret = engineSharedSecret; } public String post(String path, String jsonBody, String userId) throws IOException { @@ -78,6 +88,7 @@ public class AiEngineClient { .timeout(timeout) .POST(HttpRequest.BodyPublishers.ofString(jsonBody)); addUserHeader(builder, userId); + addEngineAuthHeader(builder); HttpResponse response = sendRequest(builder.build()); log.debug("AI engine responded with status {}", response.statusCode()); @@ -86,10 +97,15 @@ public class AiEngineClient { } /** - * Attach the X-User-Id header so the engine can scope per-user storage (RAG documents, search - * results) to the caller. Skipped when {@code userId} is blank: the engine treats the request - * as anonymous and refuses any route that requires tenancy. + * Attach the {@code X-Engine-Auth} shared secret when configured so the engine trusts this + * backend request. */ + private void addEngineAuthHeader(HttpRequest.Builder builder) { + if (engineSharedSecret != null && !engineSharedSecret.isBlank()) { + builder.header("X-Engine-Auth", engineSharedSecret); + } + } + private static void addUserHeader(HttpRequest.Builder builder, String userId) { if (userId != null && !userId.isBlank()) { builder.header("X-User-Id", userId); @@ -129,6 +145,7 @@ public class AiEngineClient { .timeout(timeout) .POST(HttpRequest.BodyPublishers.ofString(jsonBody)); addUserHeader(builder, userId); + addEngineAuthHeader(builder); HttpRequest request = builder.build(); HttpResponse> response; @@ -184,6 +201,7 @@ public class AiEngineClient { .timeout(Duration.ofSeconds(config.getTimeoutSeconds())) .DELETE(); addUserHeader(builder, userId); + addEngineAuthHeader(builder); HttpResponse response = sendRequest(builder.build()); log.debug("AI engine responded with status {}", response.statusCode()); @@ -208,6 +226,7 @@ public class AiEngineClient { .timeout(Duration.ofSeconds(config.getTimeoutSeconds())) .GET(); addUserHeader(builder, userId); + addEngineAuthHeader(builder); HttpResponse response = sendRequest(builder.build()); log.debug("AI engine responded with status {}", response.statusCode()); diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/service/AiWorkflowService.java b/app/proprietary/src/main/java/stirling/software/proprietary/service/AiWorkflowService.java index 8f49694d82..4e6c515318 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/service/AiWorkflowService.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/service/AiWorkflowService.java @@ -14,14 +14,11 @@ import org.apache.pdfbox.pdmodel.PDDocument; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.core.io.FileSystemResource; import org.springframework.core.io.Resource; -import org.springframework.http.HttpHeaders; -import org.springframework.http.HttpStatus; import org.springframework.http.MediaType; import org.springframework.http.MediaTypeFactory; -import org.springframework.http.ResponseEntity; import org.springframework.stereotype.Service; -import org.springframework.util.LinkedMultiValueMap; -import org.springframework.util.MultiValueMap; +import org.springframework.web.client.HttpServerErrorException; +import org.springframework.web.client.RestClientResponseException; import org.springframework.web.multipart.MultipartFile; import io.github.pixee.security.Filenames; @@ -32,14 +29,11 @@ import lombok.extern.slf4j.Slf4j; import stirling.software.common.model.ApplicationProperties; import stirling.software.common.service.CustomPDFDocumentFactory; import stirling.software.common.service.FileStorage; -import stirling.software.common.service.InternalApiClient; import stirling.software.common.service.InternalApiTimeoutException; -import stirling.software.common.service.ToolMetadataService; import stirling.software.common.service.UserServiceInterface; import stirling.software.common.util.ExceptionUtils; import stirling.software.common.util.TempFile; import stirling.software.common.util.TempFileManager; -import stirling.software.common.util.ZipExtractionUtils; import stirling.software.proprietary.model.api.ai.AiConversationMessage; import stirling.software.proprietary.model.api.ai.AiDocumentIngestRequest; import stirling.software.proprietary.model.api.ai.AiEngineProgressDetail; @@ -53,6 +47,13 @@ import stirling.software.proprietary.model.api.ai.AiWorkflowProgressEvent; import stirling.software.proprietary.model.api.ai.AiWorkflowRequest; import stirling.software.proprietary.model.api.ai.AiWorkflowResponse; import stirling.software.proprietary.model.api.ai.AiWorkflowResultFile; +import stirling.software.proprietary.policy.engine.PolicyExecutionResult; +import stirling.software.proprietary.policy.engine.PolicyExecutor; +import stirling.software.proprietary.policy.model.OutputSpec; +import stirling.software.proprietary.policy.model.PipelineDefinition; +import stirling.software.proprietary.policy.model.PipelineStep; +import stirling.software.proprietary.policy.model.PolicyInputs; +import stirling.software.proprietary.policy.progress.PolicyProgressListener; import stirling.software.proprietary.security.util.DesktopClientUtils; import stirling.software.proprietary.service.PdfContentExtractor.LoadedFile; import stirling.software.proprietary.service.PdfContentExtractor.PdfContentResult; @@ -67,17 +68,17 @@ import tools.jackson.databind.ObjectMapper; public class AiWorkflowService { private static final String DOCUMENTS_ENDPOINT = "/api/v1/documents"; + private static final String PDF_TO_MARKDOWN_ENDPOINT = "/api/v1/convert/pdf/markdown"; private final CustomPDFDocumentFactory pdfDocumentFactory; private final AiEngineClient aiEngineClient; private final PdfContentExtractor pdfContentExtractor; private final ObjectMapper objectMapper; - private final InternalApiClient internalApiClient; private final FileStorage fileStorage; - private final ToolMetadataService toolMetadataService; private final TempFileManager tempFileManager; private final FileIdStrategy fileIdStrategy; private final AiEngineEndpointResolver endpointResolver; + private final PolicyExecutor policyExecutor; private final UserServiceInterface userService; private final ApplicationProperties applicationProperties; @@ -86,24 +87,22 @@ public class AiWorkflowService { AiEngineClient aiEngineClient, PdfContentExtractor pdfContentExtractor, ObjectMapper objectMapper, - InternalApiClient internalApiClient, FileStorage fileStorage, - ToolMetadataService toolMetadataService, TempFileManager tempFileManager, FileIdStrategy fileIdStrategy, AiEngineEndpointResolver endpointResolver, + PolicyExecutor policyExecutor, @Autowired(required = false) UserServiceInterface userService, ApplicationProperties applicationProperties) { this.pdfDocumentFactory = pdfDocumentFactory; this.aiEngineClient = aiEngineClient; this.pdfContentExtractor = pdfContentExtractor; this.objectMapper = objectMapper; - this.internalApiClient = internalApiClient; this.fileStorage = fileStorage; - this.toolMetadataService = toolMetadataService; this.tempFileManager = tempFileManager; this.fileIdStrategy = fileIdStrategy; this.endpointResolver = endpointResolver; + this.policyExecutor = policyExecutor; this.userService = userService; this.applicationProperties = applicationProperties; } @@ -150,16 +149,6 @@ public class AiWorkflowService { record Terminal(AiWorkflowResponse response) implements WorkflowState {} } - /** - * Internal value-class for tool responses. {@code files} holds any result files (typically one; - * multiple for ZIP-response tools). {@code report} holds an optional structured metadata - * payload the tool chose to surface alongside (or instead of) a file. - * - *

Tools populate the report either by returning a JSON body (whole body → report) or by - * adding the {@link AiToolResponseHeaders#TOOL_REPORT} header alongside a file body. - */ - private record ToolResult(List files, JsonNode report) {} - public AiWorkflowResponse orchestrate(AiWorkflowRequest request) throws IOException { return orchestrate(request, NOOP_LISTENER); } @@ -186,12 +175,8 @@ public class AiWorkflowService { WorkflowTurnRequest initialRequest = new WorkflowTurnRequest(); initialRequest.setUserMessage(request.getUserMessage().trim()); initialRequest.setFiles(files); - initialRequest.setConversationHistory( - request.getConversationHistory() == null - ? new ArrayList<>() - : new ArrayList<>(request.getConversationHistory())); + initialRequest.setConversationHistory(new ArrayList<>(request.getConversationHistory())); initialRequest.setEnabledEndpoints(endpointResolver.getEnabledEndpointUrls()); - listener.onProgress(AiWorkflowProgressEvent.of(AiWorkflowPhase.ANALYZING)); WorkflowState state = new WorkflowState.Pending(initialRequest); @@ -211,6 +196,7 @@ public class AiWorkflowService { return switch (response.getOutcome()) { case NEED_CONTENT -> onNeedContent(response, filesById, request, listener); case NEED_INGEST -> onNeedIngest(response, filesById, request, listener); + case CONVERT_MARKDOWN -> onConvertMarkdown(response, filesById, listener); case TOOL_CALL -> onToolCall(response, filesById, listener); case PLAN -> onPlan(response, filesById, request, listener); case ANSWER -> onAnswer(response, filesById, request, listener); @@ -232,6 +218,12 @@ public class AiWorkflowService { WorkflowTurnRequest request, ProgressListener listener) throws IOException { + if (filesById.isEmpty()) { + return new WorkflowState.Terminal( + cannotContinue( + "No files were uploaded. Please add a PDF to the workbench first.")); + } + if (!request.getArtifacts().isEmpty()) { return new WorkflowState.Terminal( cannotContinue("AI engine requested content extraction more than once.")); @@ -341,6 +333,84 @@ public class AiWorkflowService { return new WorkflowState.Pending(nextRequest); } + /** + * Deterministically convert each requested PDF to Markdown via the {@code + * /convert/pdf/markdown} endpoint (backed by {@code PdfMarkdownConverter}) and return the + * {@code .md} file(s) as a completed result. No AI resume — the conversion output is the final + * answer. + */ + private WorkflowState onConvertMarkdown( + AiWorkflowResponse response, + Map filesById, + ProgressListener listener) { + List filesToConvert = response.getFilesToIngest(); + if (filesToConvert == null || filesToConvert.isEmpty()) { + return new WorkflowState.Terminal( + cannotContinue( + "AI engine requested markdown conversion without listing any files.")); + } + + try { + List resultFiles = new ArrayList<>(); + List inputNames = new ArrayList<>(); + for (int i = 0; i < filesToConvert.size(); i++) { + AiFile file = filesToConvert.get(i); + MultipartFile multipartFile = filesById.get(file.getId()); + if (multipartFile == null) { + return new WorkflowState.Terminal( + cannotContinue( + "AI engine requested markdown conversion for unknown file: " + + file.getName())); + } + listener.onProgress( + AiWorkflowProgressEvent.executingTool( + PDF_TO_MARKDOWN_ENDPOINT, i + 1, filesToConvert.size())); + Resource input = toResource(multipartFile); + PipelineDefinition definition = + new PipelineDefinition( + "convert-markdown", + List.of(new PipelineStep(PDF_TO_MARKDOWN_ENDPOINT, Map.of())), + null); + PolicyExecutionResult result = + policyExecutor.execute( + definition, + PolicyInputs.of(List.of(input)), + PolicyProgressListener.NOOP); + resultFiles.addAll(result.files()); + inputNames.add(multipartFile.getOriginalFilename()); + } + return new WorkflowState.Terminal( + buildCompletedResponse(null, resultFiles, inputNames, null)); + } catch (InternalApiTimeoutException e) { + log.error("PDF to Markdown conversion timed out: {}", e.getMessage()); + return new WorkflowState.Terminal( + cannotContinue(toolTimeoutMessage(PDF_TO_MARKDOWN_ENDPOINT, e))); + } catch (Exception e) { + AiWorkflowResponse limit = paygLimitResponseOrNull(e); + if (limit != null) { + log.info( + "AI markdown conversion blocked by downstream entitlement gate ({})", + limit.getErrorCode()); + return new WorkflowState.Terminal(limit); + } + log.error("Failed to convert PDF to Markdown: {}", e.getMessage(), e); + return new WorkflowState.Terminal( + cannotContinue(toolFailureMessage(PDF_TO_MARKDOWN_ENDPOINT, e))); + } + } + + private Resource toResource(MultipartFile file) throws IOException { + TempFile tempFile = tempFileManager.createManagedTempFile("ai-workflow"); + file.transferTo(tempFile.getPath()); + final String originalName = Filenames.toSimpleFileName(file.getOriginalFilename()); + return new FileSystemResource(tempFile.getFile()) { + @Override + public String getFilename() { + return originalName; + } + }; + } + private void ingestFile(AiFile file, MultipartFile multipartFile) throws IOException { List pages = new ArrayList<>(); try (PDDocument document = pdfDocumentFactory.load(multipartFile, true)) { @@ -374,7 +444,6 @@ public class AiWorkflowService { pages.size()); } - @SuppressWarnings("unchecked") private WorkflowState onToolCall( AiWorkflowResponse response, Map filesById, @@ -391,8 +460,14 @@ public class AiWorkflowService { try { List inputFiles = toResources(filesById); - listener.onProgress(AiWorkflowProgressEvent.executingTool(endpointPath, 1, 1)); - ToolResult result = executeStep(endpointPath, parameters, inputFiles); + PipelineDefinition definition = + new PipelineDefinition( + null, + List.of(new PipelineStep(endpointPath, parameters)), + OutputSpec.inline()); + PolicyExecutionResult result = + policyExecutor.execute( + definition, PolicyInputs.of(inputFiles), stepProgress(listener)); return new WorkflowState.Terminal( buildCompletedResponse( response.getRationale(), @@ -403,6 +478,14 @@ public class AiWorkflowService { log.error("Tool {} timed out: {}", endpointPath, e.getMessage()); return new WorkflowState.Terminal(cannotContinue(toolTimeoutMessage(endpointPath, e))); } catch (Exception e) { + AiWorkflowResponse limit = paygLimitResponseOrNull(e); + if (limit != null) { + log.info( + "AI workflow tool {} blocked by downstream entitlement gate ({})", + endpointPath, + limit.getErrorCode()); + return new WorkflowState.Terminal(limit); + } log.error("Failed to execute tool {}: {}", endpointPath, e.getMessage(), e); return new WorkflowState.Terminal(cannotContinue(toolFailureMessage(endpointPath, e))); } @@ -466,38 +549,32 @@ public class AiWorkflowService { cannotContinue("AI engine returned a plan with no steps.")); } - try { - List currentFiles = toResources(filesById); - // Propagate the *last* non-null report — the terminal step defines the output. - JsonNode lastReport = null; - String lastReportTool = null; - - for (int i = 0; i < steps.size(); i++) { - Map step = steps.get(i); - String endpointPath = (String) step.get("tool"); - Map parameters = - step.containsKey("parameters") - ? (Map) step.get("parameters") - : Map.of(); - - if (endpointPath == null || endpointPath.isBlank()) { - return new WorkflowState.Terminal( - cannotContinue("Plan step " + (i + 1) + " has no tool endpoint.")); - } - - listener.onProgress( - AiWorkflowProgressEvent.executingTool(endpointPath, i + 1, steps.size())); - ToolResult stepResult = executeStep(endpointPath, parameters, currentFiles); - currentFiles = stepResult.files(); - if (stepResult.report() != null) { - lastReport = stepResult.report(); - lastReportTool = endpointPath; - } + List pipelineSteps = new ArrayList<>(); + for (int i = 0; i < steps.size(); i++) { + Map step = steps.get(i); + String endpointPath = (String) step.get("tool"); + if (endpointPath == null || endpointPath.isBlank()) { + return new WorkflowState.Terminal( + cannotContinue("Plan step " + (i + 1) + " has no tool endpoint.")); } + Map parameters = + step.containsKey("parameters") + ? (Map) step.get("parameters") + : Map.of(); + pipelineSteps.add(new PipelineStep(endpointPath, parameters)); + } + + try { + List inputFiles = toResources(filesById); + PipelineDefinition definition = + new PipelineDefinition(summary, pipelineSteps, OutputSpec.inline()); + PolicyExecutionResult result = + policyExecutor.execute( + definition, PolicyInputs.of(inputFiles), stepProgress(listener)); // Multi-turn: if the plan was emitted with resume_with set, the delegate wants // Java to re-invoke the orchestrator with any captured report as an artifact. - if (resumeWith != null && !resumeWith.isBlank() && lastReport != null) { + if (resumeWith != null && !resumeWith.isBlank() && result.report() != null) { WorkflowTurnRequest resumeRequest = new WorkflowTurnRequest(); resumeRequest.setUserMessage(previousRequest.getUserMessage()); resumeRequest.setFiles(previousRequest.getFiles()); @@ -507,19 +584,30 @@ public class AiWorkflowService { .getArtifacts() .add( new PdfContentExtractor.ToolReportArtifact( - lastReportTool, lastReport)); + result.reportTool(), result.report())); resumeRequest.setResumeWith(resumeWith); return new WorkflowState.Pending(resumeRequest); } return new WorkflowState.Terminal( buildCompletedResponse( - summary, currentFiles, inputFileNames(filesById), lastReport)); + summary, result.files(), inputFileNames(filesById), result.report())); } catch (InternalApiTimeoutException e) { log.error("Plan step on tool {} timed out: {}", e.getEndpointPath(), e.getMessage()); return new WorkflowState.Terminal( cannotContinue(toolTimeoutMessage(e.getEndpointPath(), e))); + } catch (HttpServerErrorException e) { + String reason = extractDetailFromHttpError(e); + log.error("Plan step failed (HTTP {}): {}", e.getStatusCode(), reason); + return new WorkflowState.Terminal(cannotContinue(reason)); } catch (Exception e) { + AiWorkflowResponse limit = paygLimitResponseOrNull(e); + if (limit != null) { + log.info( + "AI workflow plan blocked by downstream entitlement gate ({})", + limit.getErrorCode()); + return new WorkflowState.Terminal(limit); + } log.error("Failed to execute plan: {}", e.getMessage(), e); return new WorkflowState.Terminal( cannotContinue("Plan execution failed: " + e.getMessage())); @@ -545,138 +633,46 @@ public class AiWorkflowService { } /** - * Execute a single tool step. If the endpoint accepts multiple files, all files are sent in one - * call. Otherwise, the endpoint is called once per file. ZIP responses are unpacked so each - * inner file is treated as its own result (e.g. split outputs a ZIP of pages). - * - *

A structured {@code report} may be returned alongside (or instead of) files — see {@link - * ToolResult}. For per-file dispatch (single-input endpoints called once per input), the first - * non-null report wins. + * Extracts the {@code detail} field from an HTTP error response body if it is valid JSON, + * otherwise falls back to the exception message. This lets controller-level error messages + * (e.g. missing system dependency) surface cleanly in the chat response. */ - private ToolResult executeStep( - String endpointPath, Map parameters, List inputFiles) - throws IOException { - List files = new ArrayList<>(); - JsonNode report = null; - if (toolMetadataService.isMultiInput(endpointPath)) { - ToolResult r = callEndpoint(endpointPath, parameters, inputFiles); - files.addAll(r.files()); - report = r.report(); - } else { - for (Resource file : inputFiles) { - ToolResult r = callEndpoint(endpointPath, parameters, List.of(file)); - files.addAll(r.files()); - if (report == null) { - report = r.report(); - } - } - } - return new ToolResult(files, report); - } - - /** - * Call an endpoint and return its result files and optional report. - * - *

    - *
  • JSON body (Content-Type: application/json) → the entire body is the report, no files - * are returned. - *
  • File body (PDF etc.) → the file is returned; if an {@link - * AiToolResponseHeaders#TOOL_REPORT} header is present, its (minified JSON) value is - * parsed as the report. - *
  • ZIP responses declared by the tool metadata service are unpacked so callers always see - * a flat list of result files. - *
- */ - private ToolResult callEndpoint( - String endpointPath, Map parameters, List files) - throws IOException { - MultiValueMap body = new LinkedMultiValueMap<>(); - for (Resource file : files) { - body.add("fileInput", file); - } - for (Map.Entry entry : parameters.entrySet()) { - if (entry.getValue() instanceof List list) { - if (containsStructuredElements(list)) { - // Endpoints binding lists of structured objects (e.g. /security/redact's - // redactions, /general/edit-text's edits) parse a single JSON string field via - // a property editor. Pre-serialize the whole list so binding succeeds. - body.add(entry.getKey(), objectMapper.writeValueAsString(list)); - } else { - for (Object item : list) { - body.add(entry.getKey(), item); - } - } - } else { - body.add(entry.getKey(), entry.getValue()); - } - } - ResponseEntity response = internalApiClient.post(endpointPath, body); - if (!HttpStatus.OK.equals(response.getStatusCode()) || response.getBody() == null) { - throw new IOException( - "Tool returned HTTP " + response.getStatusCode() + " for " + endpointPath); - } - Resource resource = response.getBody(); - HttpHeaders headers = response.getHeaders(); - MediaType contentType = headers.getContentType(); - - // JSON-only response — the whole body is the structured report, no result file. - if (contentType != null && MediaType.APPLICATION_JSON.isCompatibleWith(contentType)) { - try (java.io.InputStream is = resource.getInputStream()) { - JsonNode report = objectMapper.readTree(is); - return new ToolResult(List.of(), report); - } - } - - JsonNode report = parseReportHeader(headers, endpointPath); - if (toolMetadataService.shouldUnpackZipResponse(endpointPath)) { - return new ToolResult(ZipExtractionUtils.extractZip(resource, tempFileManager), report); - } - return new ToolResult(List.of(resource), report); - } - - /** - * Parse the optional {@link AiToolResponseHeaders#TOOL_REPORT} header into a {@link JsonNode}, - * or return null. - */ - private JsonNode parseReportHeader(HttpHeaders headers, String endpointPath) { - String raw = headers.getFirst(AiToolResponseHeaders.TOOL_REPORT); - if (raw == null || raw.isBlank()) { - return null; - } + private String extractDetailFromHttpError(HttpServerErrorException e) { try { - return objectMapper.readTree(raw); - } catch (JacksonException e) { - log.warn( - "Ignoring malformed {} header from {}: {}", - AiToolResponseHeaders.TOOL_REPORT, - endpointPath, - e.getMessage()); - return null; + String body = e.getResponseBodyAsString(); + if (body != null && !body.isBlank()) { + JsonNode node = objectMapper.readTree(body); + JsonNode detail = node.get("detail"); + if (detail != null && detail.isTextual() && !detail.asText().isBlank()) { + return detail.asText(); + } + } + } catch (Exception ignored) { + // fall through to generic message } + return "The request could not be completed. Please try again or contact your system administrator."; } - private static boolean containsStructuredElements(List list) { - for (Object item : list) { - if (item instanceof Map || item instanceof List) { - return true; + /** + * Adapt the AI workflow's {@link ProgressListener} to the engine's {@link + * PolicyProgressListener}: each step start maps to an {@code EXECUTING_TOOL} progress event + * carrying the tool path and 1-based step position, preserving the event shape the frontend + * already renders. + */ + private static PolicyProgressListener stepProgress(ProgressListener listener) { + return new PolicyProgressListener() { + @Override + public void onStepStart(int stepIndex, int stepCount, String operation) { + listener.onProgress( + AiWorkflowProgressEvent.executingTool(operation, stepIndex, stepCount)); } - } - return false; + }; } private List toResources(Map filesById) throws IOException { List resources = new ArrayList<>(); for (MultipartFile file : filesById.values()) { - TempFile tempFile = tempFileManager.createManagedTempFile("ai-workflow"); - file.transferTo(tempFile.getPath()); - final String originalName = Filenames.toSimpleFileName(file.getOriginalFilename()); - resources.add( - new FileSystemResource(tempFile.getFile()) { - @Override - public String getFilename() { - return originalName; - } - }); + resources.add(toResource(file)); } return resources; } @@ -749,6 +745,33 @@ public class AiWorkflowService { return response; } + /** + * If {@code e} is a downstream usage-limit block — a 401/402 from a tool call carrying the saas + * EntitlementGuard's {@code error} sentinel — build a terminal response that carries the + * structured code (+ {@code subscribed}) through to the client, so it can pop the matching + * usage-limit modal instead of surfacing the raw "tool failed: 402…" text. Returns null for any + * other failure, so the caller falls back to its normal tool-failure handling. + * + *

The agent's tool calls run server-side (loopback HTTP via {@link PolicyExecutor}), so this + * 402 never reaches the frontend's API-client interceptor that pops the modal for direct calls + * — same gap the policy auto-run path bridges in {@code PolicyEngine}. + */ + private AiWorkflowResponse paygLimitResponseOrNull(Throwable e) { + if (!(e instanceof RestClientResponseException rce)) { + return null; + } + String code = DownstreamEntitlementError.extractCode(rce); + if (code == null) { + return null; + } + AiWorkflowResponse response = new AiWorkflowResponse(); + response.setOutcome(AiWorkflowOutcome.CANNOT_CONTINUE); + response.setReason("You've reached your current usage limit."); + response.setErrorCode(code); + response.setErrorSubscribed(DownstreamEntitlementError.extractSubscribed(rce)); + return response; + } + /** * Drive the engine's streaming orchestrator endpoint. Progress events are forwarded to {@code * listener} as they arrive (each one keeps the SSE connection to the frontend alive too). The diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/service/AuditService.java b/app/proprietary/src/main/java/stirling/software/proprietary/service/AuditService.java index 9f50b9f595..d86ae644e1 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/service/AuditService.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/service/AuditService.java @@ -11,7 +11,6 @@ import java.util.HashMap; import java.util.List; import java.util.Locale; import java.util.Map; -import java.util.stream.Collectors; import java.util.stream.IntStream; import org.apache.commons.lang3.StringUtils; @@ -413,7 +412,7 @@ public class AuditService { return m; }) - .collect(Collectors.toList()); + .toList(); data.put("files", fileInfos); } diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/service/DownstreamEntitlementError.java b/app/proprietary/src/main/java/stirling/software/proprietary/service/DownstreamEntitlementError.java new file mode 100644 index 0000000000..52d9500a50 --- /dev/null +++ b/app/proprietary/src/main/java/stirling/software/proprietary/service/DownstreamEntitlementError.java @@ -0,0 +1,64 @@ +package stirling.software.proprietary.service; + +import java.util.regex.Matcher; +import java.util.regex.Pattern; + +import org.springframework.web.client.RestClientResponseException; + +/** + * Reads the {@code error} sentinel and {@code subscribed} flag out of a downstream 401/402 JSON + * body — e.g. the saas EntitlementGuard's {@code + * {"error":"PAYG_LIMIT_REACHED","subscribed":false}}. + * + *

Server-side run paths (policy auto-run, AI agent workflows) execute tool calls via loopback + * HTTP, so a usage-limit 402 surfaces as a {@link RestClientResponseException} rather than reaching + * the frontend's API-client interceptor. These helpers let those paths pass the structured code + * through to the client, which maps it to the right usage-limit modal instead of showing a generic + * failure. + * + *

Regex (not a JSON parse) on purpose: the body is a small, server-controlled shape and this + * keeps the proprietary module free of any billing-layer (saas) coupling. + */ +public final class DownstreamEntitlementError { + + private DownstreamEntitlementError() {} + + /** Matches the {@code "error":"CODE"} field of a small JSON error body. */ + private static final Pattern ERROR_CODE_FIELD = + Pattern.compile("\"error\"\\s*:\\s*\"([^\"]+)\""); + + /** Matches the {@code "subscribed":true|false} field of a small JSON error body. */ + private static final Pattern SUBSCRIBED_FIELD = + Pattern.compile("\"subscribed\"\\s*:\\s*(true|false)"); + + /** + * Pull the {@code error} sentinel out of a downstream 401/402 JSON body. Returns null for other + * statuses or an unmatched body, in which case the caller treats it as a generic failure. + */ + public static String extractCode(RestClientResponseException e) { + int status = e.getStatusCode().value(); + if (status != 401 && status != 402) { + return null; + } + String body = e.getResponseBodyAsString(); + if (body == null || body.isBlank()) { + return null; + } + Matcher m = ERROR_CODE_FIELD.matcher(body); + return m.find() ? m.group(1) : null; + } + + /** + * Pull the {@code subscribed} flag out of the body (present on the saas {@code + * PAYG_LIMIT_REACHED}/{@code FEATURE_DEGRADED} responses). Null when absent — the client then + * defaults to the free-limit modal. + */ + public static Boolean extractSubscribed(RestClientResponseException e) { + String body = e.getResponseBodyAsString(); + if (body == null || body.isBlank()) { + return null; + } + Matcher m = SUBSCRIBED_FIELD.matcher(body); + return m.find() ? Boolean.valueOf(m.group(1)) : null; + } +} diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/service/PdfContentExtractor.java b/app/proprietary/src/main/java/stirling/software/proprietary/service/PdfContentExtractor.java index c06007f318..9dccb91f38 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/service/PdfContentExtractor.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/service/PdfContentExtractor.java @@ -30,11 +30,7 @@ import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; import stirling.software.SPDF.pdf.parser.PageImageLocator; -import stirling.software.SPDF.pdf.parser.PdfIngester; -import stirling.software.SPDF.pdf.parser.PdfModels.ParsedPage; -import stirling.software.SPDF.pdf.parser.PdfModels.RawLine; import stirling.software.SPDF.pdf.parser.PdfModels.TableFragment; -import stirling.software.SPDF.pdf.parser.PdfModels.TextFragment; import stirling.software.SPDF.pdf.parser.TabulaTableParser; import stirling.software.common.util.ExceptionUtils; import stirling.software.common.util.PdfUtils; @@ -50,7 +46,6 @@ import stirling.software.proprietary.model.api.ai.FolioType; public class PdfContentExtractor { private final TabulaTableParser tabulaTableParser; - private final PdfIngester pdfIngester; private static final int MAX_CHARACTERS_PER_PAGE = 4_000; @@ -196,8 +191,6 @@ public class PdfContentExtractor { case PAGE_TEXT, FULL_TEXT -> Optional.ofNullable( extractText(lf, fileReq, remainingPages, remainingCharacters)); - case PAGE_LAYOUT -> - Optional.ofNullable(extractPageLayout(lf, remainingPages)); default -> { log.warn( "Content type {} not yet implemented, skipping for {}", @@ -222,35 +215,6 @@ public class PdfContentExtractor { return extracted.isEmpty() ? null : buildExtractedFileText(lf.fileName(), extracted); } - private PageLayoutFileResult extractPageLayout(LoadedFile lf, int maxPages) throws IOException { - List parsedPages = pdfIngester.parse(lf.document(), maxPages); - List pages = new ArrayList<>(); - for (ParsedPage pp : parsedPages) { - if (pp.layoutLines().isEmpty()) continue; - List lines = new ArrayList<>(); - for (RawLine rawLine : pp.layoutLines()) { - List fragments = new ArrayList<>(); - for (TextFragment tf : rawLine.fragments()) { - fragments.add( - new LayoutFragment( - tf.text(), - tf.bounds().x(), - tf.bounds().y(), - tf.bounds().width(), - tf.fontSize(), - tf.bold())); - } - lines.add(new LayoutLine(rawLine.bounds().y(), fragments)); - } - pages.add(new LayoutPage(pp.pageNumber(), lines)); - } - if (pages.isEmpty()) return null; - PageLayoutFileResult result = new PageLayoutFileResult(); - result.setFileName(lf.fileName()); - result.setPages(pages); - return result; - } - private WorkflowArtifact buildArtifact(ArtifactKind kind, List results) { return switch (kind) { case EXTRACTED_TEXT -> { @@ -258,11 +222,6 @@ public class PdfContentExtractor { artifact.setFiles(results.stream().map(ExtractedFileText.class::cast).toList()); yield artifact; } - case PAGE_LAYOUT -> { - PageLayoutArtifact artifact = new PageLayoutArtifact(); - artifact.setFiles(results.stream().map(PageLayoutFileResult.class::cast).toList()); - yield artifact; - } case TOOL_REPORT -> throw new IllegalArgumentException( "TOOL_REPORT artifacts are not produced by PdfContentExtractor"); @@ -569,7 +528,6 @@ public class PdfContentExtractor { */ enum ArtifactKind { EXTRACTED_TEXT("extracted_text"), - PAGE_LAYOUT("page_layout"), TOOL_REPORT("tool_report"); private final String value; @@ -633,40 +591,4 @@ public class PdfContentExtractor { this.report = report; } } - - // Serialization contract with the Python engine — see PageLayoutArtifactContractTest. - - /** One text fragment with its bounding-box geometry and font properties. */ - record LayoutFragment( - String text, float x, float y, float width, float fontSize, boolean bold) {} - - /** A visual line on the page: y-coordinate and all fragments on that line. */ - record LayoutLine(float y, List fragments) {} - - /** All layout lines for a single page. */ - record LayoutPage(int pageNumber, List lines) {} - - /** Page layout data for one file, as a PdfContentResult. */ - @Data - static final class PageLayoutFileResult implements PdfContentResult { - private String fileName; - private List pages = new ArrayList<>(); - - @Override - public ArtifactKind getArtifactKind() { - return ArtifactKind.PAGE_LAYOUT; - } - - @Override - public int pagesConsumed() { - return pages.size(); - } - } - - /** Artifact carrying full spatial page layout for all input files. */ - @Data - static final class PageLayoutArtifact implements WorkflowArtifact { - private final ArtifactKind kind = ArtifactKind.PAGE_LAYOUT; - private List files = new ArrayList<>(); - } } diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/service/ServerCertificateService.java b/app/proprietary/src/main/java/stirling/software/proprietary/service/ServerCertificateService.java index 01db79a5d8..44d750e8e4 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/service/ServerCertificateService.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/service/ServerCertificateService.java @@ -4,7 +4,6 @@ import java.io.*; import java.math.BigInteger; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import java.security.*; import java.security.cert.Certificate; import java.security.cert.X509Certificate; @@ -64,7 +63,7 @@ public class ServerCertificateService implements ServerCertificateServiceInterfa } private Path getKeystorePath() { - return Paths.get(InstallationPathConfig.getConfigPath(), KEYSTORE_FILENAME); + return Path.of(InstallationPathConfig.getConfigPath(), KEYSTORE_FILENAME); } private boolean hasProOrEnterpriseAccess() { diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/service/SignatureService.java b/app/proprietary/src/main/java/stirling/software/proprietary/service/SignatureService.java index 3b7c0ed5dd..8d4dc52a91 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/service/SignatureService.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/service/SignatureService.java @@ -5,7 +5,6 @@ import java.io.IOException; import java.nio.charset.StandardCharsets; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import java.nio.file.StandardOpenOption; import java.util.ArrayList; import java.util.Base64; @@ -55,7 +54,7 @@ public class SignatureService implements PersonalSignatureServiceInterface { @Override public byte[] getPersonalSignatureBytes(String username, String fileName) throws IOException { validateFileName(fileName); - Path userPath = Paths.get(SIGNATURE_BASE_PATH, username, fileName); + Path userPath = Path.of(SIGNATURE_BASE_PATH, username, fileName); if (!Files.exists(userPath)) { throw new FileNotFoundException("Personal signature not found"); @@ -76,7 +75,7 @@ public class SignatureService implements PersonalSignatureServiceInterface { } String folderName = "shared".equals(scope) ? ALL_USERS_FOLDER : username; - Path targetFolder = Paths.get(SIGNATURE_BASE_PATH, folderName); + Path targetFolder = Path.of(SIGNATURE_BASE_PATH, folderName); // Only enforce limits for personal signatures (not shared) if ("personal".equals(scope)) { @@ -170,13 +169,13 @@ public class SignatureService implements PersonalSignatureServiceInterface { List signatures = new ArrayList<>(); // Load personal signatures - Path personalFolder = Paths.get(SIGNATURE_BASE_PATH, username); + Path personalFolder = Path.of(SIGNATURE_BASE_PATH, username); if (Files.exists(personalFolder)) { signatures.addAll(loadSignaturesFromFolder(personalFolder, "personal", true)); } // Load shared signatures - Path sharedFolder = Paths.get(SIGNATURE_BASE_PATH, ALL_USERS_FOLDER); + Path sharedFolder = Path.of(SIGNATURE_BASE_PATH, ALL_USERS_FOLDER); if (Files.exists(sharedFolder)) { signatures.addAll(loadSignaturesFromFolder(sharedFolder, "shared", false)); } @@ -189,7 +188,7 @@ public class SignatureService implements PersonalSignatureServiceInterface { validateFileName(signatureId); // Only allow deletion from personal folder - Path personalFolder = Paths.get(SIGNATURE_BASE_PATH, username); + Path personalFolder = Path.of(SIGNATURE_BASE_PATH, username); boolean deleted = false; if (Files.exists(personalFolder)) { @@ -227,7 +226,7 @@ public class SignatureService implements PersonalSignatureServiceInterface { validateFileName(signatureId); // Try personal folder first - Path personalFolder = Paths.get(SIGNATURE_BASE_PATH, username); + Path personalFolder = Path.of(SIGNATURE_BASE_PATH, username); Path metadataPath = personalFolder.resolve(signatureId + ".json"); if (Files.exists(metadataPath)) { @@ -237,7 +236,7 @@ public class SignatureService implements PersonalSignatureServiceInterface { } // If not found in personal, try shared folder - Path sharedFolder = Paths.get(SIGNATURE_BASE_PATH, ALL_USERS_FOLDER); + Path sharedFolder = Path.of(SIGNATURE_BASE_PATH, ALL_USERS_FOLDER); Path sharedMetadataPath = sharedFolder.resolve(signatureId + ".json"); if (Files.exists(sharedMetadataPath)) { @@ -251,7 +250,7 @@ public class SignatureService implements PersonalSignatureServiceInterface { public boolean isSharedSignature(String signatureId) { validateFileName(signatureId); - Path sharedFolder = Paths.get(SIGNATURE_BASE_PATH, ALL_USERS_FOLDER); + Path sharedFolder = Path.of(SIGNATURE_BASE_PATH, ALL_USERS_FOLDER); return Files.exists(sharedFolder.resolve(signatureId + ".json")); } @@ -274,7 +273,7 @@ public class SignatureService implements PersonalSignatureServiceInterface { // Private helper methods private void enforceStorageLimits(String username, String dataUrlToAdd) throws IOException { - Path userFolder = Paths.get(SIGNATURE_BASE_PATH, username); + Path userFolder = Path.of(SIGNATURE_BASE_PATH, username); if (!Files.exists(userFolder)) { return; // First signature, no limits to check diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/storage/config/StorageProviderConfig.java b/app/proprietary/src/main/java/stirling/software/proprietary/storage/config/StorageProviderConfig.java index e990fa2ceb..db063cda80 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/storage/config/StorageProviderConfig.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/storage/config/StorageProviderConfig.java @@ -3,7 +3,6 @@ package stirling.software.proprietary.storage.config; import java.io.IOException; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import java.util.Locale; import java.util.Optional; @@ -58,8 +57,8 @@ public class StorageProviderConfig { } basePathValue = InstallationPathConfig.getPath() + "storage"; } - Path basePath = Paths.get(basePathValue).toAbsolutePath().normalize(); - Path installRoot = Paths.get(InstallationPathConfig.getPath()).toAbsolutePath().normalize(); + Path basePath = Path.of(basePathValue).toAbsolutePath().normalize(); + Path installRoot = Path.of(InstallationPathConfig.getPath()).toAbsolutePath().normalize(); if (!basePath.startsWith(installRoot)) { // Warn rather than hard-fail: admins may legitimately point storage at an external // volume, but an unexpected path could indicate a misconfiguration or traversal diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/storage/provider/LocalStorageProvider.java b/app/proprietary/src/main/java/stirling/software/proprietary/storage/provider/LocalStorageProvider.java index 75f9eb3fd5..b67b66f540 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/storage/provider/LocalStorageProvider.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/storage/provider/LocalStorageProvider.java @@ -4,7 +4,6 @@ import java.io.IOException; import java.io.InputStream; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import java.nio.file.StandardCopyOption; import java.util.Optional; import java.util.UUID; @@ -80,7 +79,7 @@ public class LocalStorageProvider implements StorageProvider { if (filename == null || filename.isBlank()) { return "file"; } - String stripped = Paths.get(filename).getFileName().toString().replaceAll("\\p{Cntrl}", ""); + String stripped = Path.of(filename).getFileName().toString().replaceAll("\\p{Cntrl}", ""); return stripped.isBlank() ? "file" : stripped; } } diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/storage/provider/S3StorageProvider.java b/app/proprietary/src/main/java/stirling/software/proprietary/storage/provider/S3StorageProvider.java index a1582458bc..2b88d64844 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/storage/provider/S3StorageProvider.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/storage/provider/S3StorageProvider.java @@ -4,7 +4,7 @@ import java.io.IOException; import java.io.InputStream; import java.net.URI; import java.net.URISyntaxException; -import java.nio.file.Paths; +import java.nio.file.Path; import java.time.Duration; import java.util.Optional; import java.util.UUID; @@ -150,7 +150,7 @@ public class S3StorageProvider implements StorageProvider, AutoCloseable { if (originalFilename == null || originalFilename.isBlank()) { return null; } - // Strip CR/LF and other control chars before path parsing (Paths.get throws on them on + // Strip CR/LF and other control chars before path parsing (Path.of throws on them on // Windows, and they defeat header parsers). String stripped = originalFilename.replaceAll("\\p{Cntrl}", ""); // Use only the basename to avoid leaking directory structure into the header. @@ -184,7 +184,7 @@ public class S3StorageProvider implements StorageProvider, AutoCloseable { if (filename == null || filename.isBlank()) { return "file"; } - String stripped = Paths.get(filename).getFileName().toString().replaceAll("\\p{Cntrl}", ""); + String stripped = Path.of(filename).getFileName().toString().replaceAll("\\p{Cntrl}", ""); return stripped.isBlank() ? "file" : stripped; } } diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/storage/service/FileStorageService.java b/app/proprietary/src/main/java/stirling/software/proprietary/storage/service/FileStorageService.java index 03c3a9178f..37d1e5fa43 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/storage/service/FileStorageService.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/storage/service/FileStorageService.java @@ -14,7 +14,6 @@ import java.util.Optional; import java.util.Set; import java.util.UUID; import java.util.regex.Pattern; -import java.util.stream.Collectors; import org.springframework.http.HttpStatus; import org.springframework.security.core.Authentication; @@ -361,7 +360,7 @@ public class FileStorageService { return files.stream() .sorted(Comparator.comparing(StoredFile::getCreatedAt).reversed()) .map(file -> buildResponse(file, user, roleByFileId.get(file.getId()))) - .collect(Collectors.toList()); + .toList(); } public StoredFileResponse getAccessibleFileResponse(User user, Long fileId) { @@ -400,7 +399,7 @@ public class FileStorageService { .filter(Objects::nonNull) .map(User::getUsername) .sorted(String.CASE_INSENSITIVE_ORDER) - .collect(Collectors.toList()) + .toList() : List.of(); List shareLinks = ownedByCurrentUser && isShareLinksEnabled() @@ -419,7 +418,7 @@ public class FileStorageService { .expiresAt(share.getExpiresAt()) .build()) .sorted(Comparator.comparing(ShareLinkResponse::getCreatedAt)) - .collect(Collectors.toList()) + .toList() : List.of(); List sharedUsers = ownedByCurrentUser @@ -440,7 +439,7 @@ public class FileStorageService { Comparator.comparing( SharedUserResponse::getUsername, String.CASE_INSENSITIVE_ORDER)) - .collect(Collectors.toList()) + .toList() : List.of(); return StoredFileResponse.builder() .id(file.getId()) @@ -762,7 +761,7 @@ public class FileStorageService { .accessType(access.getAccessType().name()) .accessedAt(access.getAccessedAt()) .build()) - .collect(Collectors.toList()); + .toList(); } public List listAccessedShareLinks(User user) { @@ -818,7 +817,7 @@ public class FileStorageService { .build(); }) .filter(response -> response.getShareToken() != null) - .collect(Collectors.toList()); + .toList(); } public void ensureSharingEnabled() { @@ -1007,7 +1006,7 @@ public class FileStorageService { file.getHistoryStorageKey(), file.getAuditLogStorageKey()) .filter(value -> value != null && !value.isBlank()) - .collect(Collectors.toList()); + .toList(); } private void cleanupStoredObject(StoredObject storedObject) { diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/workflow/controller/SigningSessionController.java b/app/proprietary/src/main/java/stirling/software/proprietary/workflow/controller/SigningSessionController.java index 756fba1fe0..2464b6ed62 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/workflow/controller/SigningSessionController.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/workflow/controller/SigningSessionController.java @@ -76,7 +76,7 @@ public class SigningSessionController { .map( stirling.software.proprietary.workflow.util.WorkflowMapper ::toResponse) - .collect(java.util.stream.Collectors.toList()); + .toList(); return ResponseEntity.ok(responses); } catch (Exception e) { log.error("Error listing sessions for user {}", principal.getName(), e); diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/workflow/service/SigningFinalizationService.java b/app/proprietary/src/main/java/stirling/software/proprietary/workflow/service/SigningFinalizationService.java index aac9c180fc..dc8469dc8e 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/workflow/service/SigningFinalizationService.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/workflow/service/SigningFinalizationService.java @@ -98,7 +98,7 @@ public class SigningFinalizationService { // Step 1.5: Add summary page BEFORE digital signing (if enabled) // CRITICAL: Must be done before signing to avoid invalidating signatures - if (Boolean.TRUE.equals(settings.includeSummaryPage)) { + if (Boolean.TRUE.equals(settings.includeSummaryPage())) { log.info( "Adding summary page before digital signing for session {}", session.getSessionId()); @@ -108,11 +108,13 @@ public class SigningFinalizationService { // Suppress digital certificate visual block when summary page is enabled // (wet signatures already applied in Step 1 and will still appear) Boolean showVisualSignature = - Boolean.TRUE.equals(settings.includeSummaryPage) ? false : settings.showSignature; + Boolean.TRUE.equals(settings.includeSummaryPage()) + ? false + : settings.showSignature(); log.info( "Finalization settings: includeSummaryPage={}, showVisualSignature={}", - settings.includeSummaryPage, + settings.includeSummaryPage(), showVisualSignature); // Step 2: Apply digital certificates per SIGNED participant @@ -150,8 +152,8 @@ public class SigningFinalizationService { log.info( "Applying signature for {} with reason='{}', location='{}'", fresh.getEmail(), - sigMeta.reason, - sigMeta.location); + sigMeta.reason(), + sigMeta.location()); pdf = applyDigitalSignature( @@ -159,10 +161,10 @@ public class SigningFinalizationService { fresh, submission, showVisualSignature, - settings.pageNumber, - sigMeta.reason, - sigMeta.location, - settings.showLogo); + settings.pageNumber(), + sigMeta.reason(), + sigMeta.location(), + settings.showLogo()); } return pdf; @@ -682,7 +684,7 @@ public class SigningFinalizationService { textDark, textMuted, "Subject:", - certInfo.subjectCN); + certInfo.subjectCN()); rRowY -= LINE_H; drawLabelValue( cs, @@ -693,7 +695,7 @@ public class SigningFinalizationService { textDark, textMuted, "Issuer:", - certInfo.issuerCN); + certInfo.issuerCN()); rRowY -= LINE_H; drawLabelValue( cs, @@ -704,7 +706,7 @@ public class SigningFinalizationService { textDark, textMuted, "Serial:", - certInfo.serialNumber); + certInfo.serialNumber()); rRowY -= LINE_H; drawLabelValue( cs, @@ -715,7 +717,7 @@ public class SigningFinalizationService { textDark, textMuted, "Valid From:", - certInfo.validFrom); + certInfo.validFrom()); rRowY -= LINE_H; drawLabelValue( cs, @@ -726,7 +728,7 @@ public class SigningFinalizationService { textDark, textMuted, "Valid Until:", - certInfo.validUntil); + certInfo.validUntil()); rRowY -= LINE_H; drawLabelValue( cs, @@ -737,7 +739,7 @@ public class SigningFinalizationService { textDark, textMuted, "Algorithm:", - certInfo.algorithm); + certInfo.algorithm()); } yPos -= cardH + 12; @@ -1075,10 +1077,9 @@ public class SigningFinalizationService { } String alias = aliases.nextElement(); Certificate cert = keystore.getCertificate(alias); - if (!(cert instanceof X509Certificate)) { + if (!(cert instanceof X509Certificate x509)) { return null; } - X509Certificate x509 = (X509Certificate) cert; String subjectCN = extractCN(x509.getSubjectX500Principal().getName()); String issuerCN = extractCN(x509.getIssuerX500Principal().getName()); @@ -1256,55 +1257,19 @@ public class SigningFinalizationService { // ===== PRIVATE INNER TYPES ===== - private static class SessionSignatureSettings { - final Boolean showSignature; - final Integer pageNumber; - final Boolean showLogo; - final Boolean includeSummaryPage; + private record SessionSignatureSettings( + Boolean showSignature, + Integer pageNumber, + Boolean showLogo, + Boolean includeSummaryPage) {} - SessionSignatureSettings( - Boolean showSignature, - Integer pageNumber, - Boolean showLogo, - Boolean includeSummaryPage) { - this.showSignature = showSignature; - this.pageNumber = pageNumber; - this.showLogo = showLogo; - this.includeSummaryPage = includeSummaryPage; - } - } + private record ParticipantSignatureMetadata(String reason, String location) {} - private static class ParticipantSignatureMetadata { - final String reason; - final String location; - - ParticipantSignatureMetadata(String reason, String location) { - this.reason = reason; - this.location = location; - } - } - - private static class CertificateInfo { - final String subjectCN; - final String issuerCN; - final String serialNumber; - final String validFrom; - final String validUntil; - final String algorithm; - - CertificateInfo( - String subjectCN, - String issuerCN, - String serialNumber, - String validFrom, - String validUntil, - String algorithm) { - this.subjectCN = subjectCN; - this.issuerCN = issuerCN; - this.serialNumber = serialNumber; - this.validFrom = validFrom; - this.validUntil = validUntil; - this.algorithm = algorithm; - } - } + private record CertificateInfo( + String subjectCN, + String issuerCN, + String serialNumber, + String validFrom, + String validUntil, + String algorithm) {} } diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/workflow/service/WorkflowSessionService.java b/app/proprietary/src/main/java/stirling/software/proprietary/workflow/service/WorkflowSessionService.java index e8887e2f55..dc22c40525 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/workflow/service/WorkflowSessionService.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/workflow/service/WorkflowSessionService.java @@ -553,7 +553,7 @@ public class WorkflowSessionService { dto.setMyStatus(p.getStatus()); return dto; }) - .collect(java.util.stream.Collectors.toList()); + .toList(); } /** diff --git a/app/proprietary/src/main/java/stirling/software/proprietary/workflow/util/WorkflowMapper.java b/app/proprietary/src/main/java/stirling/software/proprietary/workflow/util/WorkflowMapper.java index 6be7197d5c..b2f53c2824 100644 --- a/app/proprietary/src/main/java/stirling/software/proprietary/workflow/util/WorkflowMapper.java +++ b/app/proprietary/src/main/java/stirling/software/proprietary/workflow/util/WorkflowMapper.java @@ -3,7 +3,6 @@ package stirling.software.proprietary.workflow.util; import java.util.ArrayList; import java.util.List; import java.util.Map; -import java.util.stream.Collectors; import com.fasterxml.jackson.databind.ObjectMapper; @@ -80,12 +79,12 @@ public class WorkflowMapper { response.setParticipants( session.getParticipants().stream() .map(p -> toParticipantResponse(p, objectMapper, includeShareTokens)) - .collect(Collectors.toList())); + .toList()); } else { response.setParticipants( session.getParticipants().stream() .map(p -> toParticipantResponse(p, includeShareTokens)) - .collect(Collectors.toList())); + .toList()); } // Calculate participant counts diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/cluster/s3/S3FileStoreTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/cluster/s3/S3FileStoreTest.java index b407041bd5..9bb4d08b7c 100644 --- a/app/proprietary/src/test/java/stirling/software/proprietary/cluster/s3/S3FileStoreTest.java +++ b/app/proprietary/src/test/java/stirling/software/proprietary/cluster/s3/S3FileStoreTest.java @@ -250,4 +250,27 @@ class S3FileStoreTest { return toWrite; } } + + @Test + void store_withOwner_persistsOwnerMetadata() throws IOException { + FileStore.Stored stored = + store.store( + new ByteArrayInputStream("owned".getBytes(StandardCharsets.UTF_8)), + "o.txt", + "alice"); + assertThat(store.getOwner(stored.fileId())).isEqualTo("alice"); + } + + @Test + void store_withoutOwner_yieldsNullFromGetOwner() throws IOException { + FileStore.Stored stored = + store.store( + new ByteArrayInputStream("anon".getBytes(StandardCharsets.UTF_8)), "a.txt"); + assertThat(store.getOwner(stored.fileId())).isNull(); + } + + @Test + void getOwner_returnsNullForUnknownFileId() throws IOException { + assertThat(store.getOwner("00000000-0000-0000-0000-000000000000")).isNull(); + } } diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/config/CustomAuditEventRepositoryTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/config/CustomAuditEventRepositoryTest.java new file mode 100644 index 0000000000..3bdde8ecab --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/config/CustomAuditEventRepositoryTest.java @@ -0,0 +1,52 @@ +package stirling.software.proprietary.config; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotEquals; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import org.junit.jupiter.api.Test; + +class CustomAuditEventRepositoryTest { + + @Test + void shortPrincipalPassesThroughUnchanged() { + assertEquals( + "alice@example.com", CustomAuditEventRepository.safePrincipal("alice@example.com")); + } + + @Test + void blankOrNullPrincipalBecomesAnonymous() { + assertEquals("anonymous", CustomAuditEventRepository.safePrincipal(null)); + assertEquals("anonymous", CustomAuditEventRepository.safePrincipal(" ")); + } + + @Test + void tokenShapedPrincipalIsHashedNotStoredVerbatim() { + String jwt = "eyJhbGciOiJSUzI1NiJ9." + "x".repeat(1400); + + String safe = CustomAuditEventRepository.safePrincipal(jwt); + + assertNotEquals(jwt, safe); + assertFalse(safe.contains(jwt), "raw token must not be stored"); + assertTrue(safe.startsWith("token:")); + assertTrue(safe.length() <= 255, "must fit the principal column"); + } + + @Test + void distinctTokensStayDistinguishable() { + String a = CustomAuditEventRepository.safePrincipal("eyJ" + "a".repeat(400)); + String b = CustomAuditEventRepository.safePrincipal("eyJ" + "b".repeat(400)); + + assertNotEquals(a, b, "different tokens must map to different audit principals"); + } + + @Test + void sameTokenHashesStably() { + String token = "eyJ" + "c".repeat(400); + + assertEquals( + CustomAuditEventRepository.safePrincipal(token), + CustomAuditEventRepository.safePrincipal(token)); + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/mcp/McpConditionalTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/McpConditionalTest.java new file mode 100644 index 0000000000..ca6eee5ee8 --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/McpConditionalTest.java @@ -0,0 +1,99 @@ +package stirling.software.proprietary.mcp; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.util.Arrays; + +import org.junit.jupiter.api.Test; +import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty; +import org.springframework.context.annotation.Profile; + +import stirling.software.proprietary.mcp.catalog.McpToolCatalog; +import stirling.software.proprietary.mcp.engine.EngineCapabilityClient; +import stirling.software.proprietary.mcp.security.McpSecurityConfig; +import stirling.software.proprietary.mcp.tools.DescribeOperationTool; +import stirling.software.proprietary.mcp.tools.McpOperationExecutor; +import stirling.software.proprietary.mcp.tools.StirlingAiTool; +import stirling.software.proprietary.mcp.tools.StirlingConvertTool; +import stirling.software.proprietary.mcp.tools.StirlingDownloadTool; +import stirling.software.proprietary.mcp.tools.StirlingMiscTool; +import stirling.software.proprietary.mcp.tools.StirlingPagesTool; +import stirling.software.proprietary.mcp.tools.StirlingSecurityTool; +import stirling.software.proprietary.mcp.tools.StirlingUploadTool; + +/** Verifies MCP beans are gated behind {@code @ConditionalOnProperty(name="mcp.enabled")}. */ +class McpConditionalTest { + + @Test + void serverController_isGatedByMcpEnabled() { + assertGatedByEnabled(McpServerController.class); + } + + @Test + void securityConfig_isGatedByMcpEnabled() { + assertGatedByEnabled(McpSecurityConfig.class); + } + + @Test + void categoryToolsAndDescribeOperation_doNotNeedOwnGate() { + // The tool beans are only wired into the gated controller; sanity-check their signatures. + Class[] tools = { + DescribeOperationTool.class, + StirlingConvertTool.class, + StirlingPagesTool.class, + StirlingMiscTool.class, + StirlingSecurityTool.class, + StirlingAiTool.class + }; + for (Class t : tools) { + assertTrue( + McpTool.class.isAssignableFrom(t), + t.getSimpleName() + " must implement McpTool"); + assertNotNull( + t.getAnnotation(org.springframework.stereotype.Component.class), + t.getSimpleName() + " must be @Component"); + } + } + + @Test + void mcpBeans_areNotSaasProfileRestricted() { + // Beans gate on mcp.enabled only; no @Profile, so MCP can run under the saas profile too. + Class[] beans = { + McpServerController.class, + McpSecurityConfig.class, + McpToolCatalog.class, + EngineCapabilityClient.class, + McpOperationExecutor.class, + DescribeOperationTool.class, + StirlingAiTool.class, + StirlingConvertTool.class, + StirlingMiscTool.class, + StirlingPagesTool.class, + StirlingSecurityTool.class, + StirlingUploadTool.class, + StirlingDownloadTool.class + }; + for (Class bean : beans) { + assertNull( + bean.getAnnotation(Profile.class), + bean.getSimpleName() + + " must not be @Profile-restricted so MCP can run under saas"); + } + } + + private static void assertGatedByEnabled(Class beanClass) { + ConditionalOnProperty conditional = beanClass.getAnnotation(ConditionalOnProperty.class); + assertNotNull(conditional, beanClass.getSimpleName() + " missing @ConditionalOnProperty"); + assertTrue( + Arrays.asList(conditional.name()).contains("mcp.enabled") + || Arrays.asList(conditional.value()).contains("mcp.enabled"), + beanClass.getSimpleName() + " must gate on mcp.enabled"); + assertEquals( + "true", + conditional.havingValue(), + beanClass.getSimpleName() + " must require mcp.enabled=true"); + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/mcp/McpServerControllerTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/McpServerControllerTest.java new file mode 100644 index 0000000000..0418295797 --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/McpServerControllerTest.java @@ -0,0 +1,238 @@ +package stirling.software.proprietary.mcp; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.util.Iterator; +import java.util.List; +import java.util.Set; +import java.util.function.Consumer; +import java.util.function.Supplier; + +import org.junit.jupiter.api.Test; +import org.springframework.beans.factory.ObjectProvider; +import org.springframework.http.HttpStatus; +import org.springframework.http.ResponseEntity; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.proprietary.mcp.catalog.McpToolCatalog; +import stirling.software.proprietary.mcp.tools.DescribeOperationTool; +import stirling.software.proprietary.mcp.tools.McpOperationExecutor; +import stirling.software.proprietary.mcp.tools.StirlingAiTool; +import stirling.software.proprietary.mcp.tools.StirlingConvertTool; +import stirling.software.proprietary.mcp.tools.StirlingMiscTool; +import stirling.software.proprietary.mcp.tools.StirlingPagesTool; +import stirling.software.proprietary.mcp.tools.StirlingSecurityTool; +import stirling.software.proprietary.service.AiEngineClient; + +import tools.jackson.databind.JsonNode; +import tools.jackson.databind.ObjectMapper; + +/** Unit test of the MCP server controller: JSON-RPC framing and the 6-tool contract. */ +class McpServerControllerTest { + + private final ObjectMapper mapper = new ObjectMapper(); + private final McpServerController controller = buildController(); + + private McpServerController buildController() { + ApplicationProperties props = new ApplicationProperties(); + props.getAutomaticallyGenerated().setAppVersion("test-version"); + ObjectProvider emptyCatalog = emptyProvider(); + ObjectProvider emptyEngine = emptyProvider(); + ObjectProvider emptyExecutor = emptyProvider(); + List tools = + List.of( + new DescribeOperationTool(mapper, emptyCatalog), + new StirlingConvertTool(mapper, emptyCatalog, emptyExecutor), + new StirlingPagesTool(mapper, emptyCatalog, emptyExecutor), + new StirlingMiscTool(mapper, emptyCatalog, emptyExecutor), + new StirlingSecurityTool(mapper, emptyCatalog, emptyExecutor), + new StirlingAiTool(mapper, emptyCatalog, emptyEngine)); + return new McpServerController(mapper, props, tools); + } + + private static ObjectProvider emptyProvider() { + return new ObjectProvider<>() { + @Override + public T getObject() { + throw new UnsupportedOperationException("no bean in unit tests"); + } + + @Override + public T getObject(Object... args) { + return getObject(); + } + + @Override + public T getIfAvailable() { + return null; + } + + @Override + public T getIfUnique() { + return null; + } + + @Override + public T getIfAvailable(Supplier defaultSupplier) { + return defaultSupplier == null ? null : defaultSupplier.get(); + } + + @Override + public void ifAvailable(Consumer dependencyConsumer) {} + + @Override + public Iterator iterator() { + return java.util.Collections.emptyIterator(); + } + }; + } + + @Test + void toolsList_returnsExactlySixTools() throws Exception { + JsonNode body = mapper.readTree("{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"tools/list\"}"); + + ResponseEntity response = controller.handle(body); + + assertEquals(HttpStatus.OK, response.getStatusCode()); + JsonNode tools = mapper.valueToTree(response.getBody()).get("result").get("tools"); + assertEquals(6, tools.size(), "tools/list must return exactly 6 tools"); + + Set names = + Set.of( + "stirling_describe_operation", + "stirling_convert", + "stirling_pages", + "stirling_misc", + "stirling_security", + "stirling_ai"); + Set seen = new java.util.HashSet<>(); + tools.forEach(t -> seen.add(t.get("name").asText())); + assertEquals(names, seen); + + for (JsonNode tool : tools) { + assertTrue(tool.get("description").asText().length() > 10, "description present"); + assertEquals("object", tool.get("inputSchema").get("type").asText()); + } + } + + @Test + void initialize_returnsServerInfoAndProtocolVersion() throws Exception { + JsonNode body = mapper.readTree("{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"initialize\"}"); + + ResponseEntity response = controller.handle(body); + + JsonNode result = mapper.valueToTree(response.getBody()).get("result"); + assertNotNull(result.get("protocolVersion")); + assertEquals("stirling-pdf-mcp", result.get("serverInfo").get("name").asText()); + assertEquals("test-version", result.get("serverInfo").get("version").asText()); + assertNotNull(result.get("capabilities").get("tools"), "tools capability advertised"); + } + + @Test + void ping_returnsEmptyResult() throws Exception { + JsonNode body = mapper.readTree("{\"jsonrpc\":\"2.0\",\"id\":7,\"method\":\"ping\"}"); + + ResponseEntity response = controller.handle(body); + + JsonNode out = mapper.valueToTree(response.getBody()); + assertEquals(7, out.get("id").asInt()); + assertNotNull(out.get("result")); + assertNull(out.get("error")); + } + + @Test + void notification_returnsNoContentWithEmptyBody() throws Exception { + // No id field: a JSON-RPC notification gets no response object. + JsonNode body = + mapper.readTree("{\"jsonrpc\":\"2.0\",\"method\":\"notifications/initialized\"}"); + + ResponseEntity response = controller.handle(body); + + assertEquals(HttpStatus.NO_CONTENT, response.getStatusCode()); + assertNull(response.getBody()); + } + + @Test + void unknownMethod_returnsMethodNotFoundError() throws Exception { + JsonNode body = + mapper.readTree("{\"jsonrpc\":\"2.0\",\"id\":3,\"method\":\"does/not/exist\"}"); + + ResponseEntity response = controller.handle(body); + + JsonNode error = mapper.valueToTree(response.getBody()).get("error"); + assertEquals(-32601, error.get("code").asInt()); + assertTrue(error.get("message").asText().contains("does/not/exist")); + } + + @Test + void toolsCall_unknownTool_returnsInvalidParams() throws Exception { + JsonNode body = + mapper.readTree( + "{\"jsonrpc\":\"2.0\",\"id\":4,\"method\":\"tools/call\"," + + "\"params\":{\"name\":\"stirling_does_not_exist\",\"arguments\":{}}}"); + + ResponseEntity response = controller.handle(body); + + JsonNode error = mapper.valueToTree(response.getBody()).get("error"); + assertEquals(-32602, error.get("code").asInt()); + } + + @Test + void toolsCall_describeOperation_withoutCatalog_returnsErrorContent() throws Exception { + // Null catalog: describe must surface an isError content block, not crash. + JsonNode body = + mapper.readTree( + "{\"jsonrpc\":\"2.0\",\"id\":5,\"method\":\"tools/call\"," + + "\"params\":{\"name\":\"stirling_describe_operation\"," + + "\"arguments\":{\"operation\":\"compress-pdf\"}}}"); + + ResponseEntity response = controller.handle(body); + + JsonNode result = mapper.valueToTree(response.getBody()).get("result"); + assertTrue(result.get("isError").asBoolean()); + String text = result.get("content").get(0).get("text").asText(); + assertTrue(text.toLowerCase().contains("catalog") || text.contains("compress-pdf")); + } + + @Test + void wrongShapeJson_returnsInvalidRequest() throws Exception { + // Valid JSON but not a JSON-RPC request object -> Invalid Request (-32600). + JsonNode body = mapper.readTree("{\"not\":\"a json-rpc frame\"}"); + + ResponseEntity response = controller.handle(body); + + assertEquals(HttpStatus.BAD_REQUEST, response.getStatusCode()); + JsonNode error = mapper.valueToTree(response.getBody()).get("error"); + assertEquals(-32600, error.get("code").asInt()); + } + + @Test + void initialize_echoesSupportedClientProtocolVersion() throws Exception { + // Older but supported revision -> server echoes it. + JsonNode body = + mapper.readTree( + "{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"initialize\"," + + "\"params\":{\"protocolVersion\":\"2025-03-26\"}}"); + + ResponseEntity response = controller.handle(body); + + JsonNode result = mapper.valueToTree(response.getBody()).get("result"); + assertEquals("2025-03-26", result.get("protocolVersion").asText()); + } + + @Test + void initialize_unknownClientProtocolVersion_fallsBackToPreferred() throws Exception { + JsonNode body = + mapper.readTree( + "{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"initialize\"," + + "\"params\":{\"protocolVersion\":\"1999-01-01\"}}"); + + ResponseEntity response = controller.handle(body); + + JsonNode result = mapper.valueToTree(response.getBody()).get("result"); + assertEquals("2025-06-18", result.get("protocolVersion").asText()); + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/mcp/catalog/McpToolCatalogTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/catalog/McpToolCatalogTest.java new file mode 100644 index 0000000000..71341cf6ce --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/catalog/McpToolCatalogTest.java @@ -0,0 +1,185 @@ +package stirling.software.proprietary.mcp.catalog; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; + +import java.lang.reflect.Field; +import java.util.List; +import java.util.Map; +import java.util.Set; + +import org.junit.jupiter.api.Test; +import org.springframework.context.ApplicationContext; +import org.springframework.web.bind.annotation.RequestMethod; +import org.springframework.web.servlet.mvc.method.annotation.RequestMappingHandlerMapping; + +import stirling.software.SPDF.config.EndpointConfiguration; +import stirling.software.common.model.ApplicationProperties; + +import tools.jackson.databind.ObjectMapper; +import tools.jackson.databind.node.ObjectNode; + +/** Catalog tests: DELETE/GET exclusion and disabled-PDF-op AI fall-through. */ +class McpToolCatalogTest { + + private final ObjectMapper mapper = new ObjectMapper(); + + @Test + void isInvocableMethod_excludesDeleteAndGet() { + assertTrue(McpToolCatalog.isInvocableMethod(Set.of(RequestMethod.POST))); + assertTrue(McpToolCatalog.isInvocableMethod(Set.of(RequestMethod.PUT))); + // DELETE/GET handlers must never be cataloged as runnable tools. + assertFalse(McpToolCatalog.isInvocableMethod(Set.of(RequestMethod.DELETE))); + assertFalse(McpToolCatalog.isInvocableMethod(Set.of(RequestMethod.GET))); + // Empty method set (matches all verbs) is not invocable. + assertFalse(McpToolCatalog.isInvocableMethod(Set.of())); + // Multi-verb mapping including POST stays invocable. + assertTrue( + McpToolCatalog.isInvocableMethod(Set.of(RequestMethod.POST, RequestMethod.DELETE))); + } + + @Test + void findByOperationId_disabledPdfOp_doesNotFallThroughToAi() throws Exception { + ApplicationContext ctx = mock(ApplicationContext.class); + when(ctx.getBeansOfType(RequestMappingHandlerMapping.class)).thenReturn(Map.of()); + EndpointConfiguration endpoints = mock(EndpointConfiguration.class); + // The PDF op's endpoint is disabled. + when(endpoints.isEndpointEnabledForUri(anyString())).thenReturn(false); + + McpToolCatalog catalog = + new McpToolCatalog(ctx, endpoints, new ApplicationProperties(), mapper); + + ObjectNode schema = mapper.createObjectNode(); + OperationMeta pdf = + new OperationMeta( + "collide", + OperationCategory.MISC, + "pdf op", + schema, + "mcp.tools.write", + OperationMeta.Target.JAVA_ENDPOINT, + "/api/v1/misc/collide", + null); + OperationMeta ai = + new OperationMeta( + "collide", + OperationCategory.AI, + "ai op", + schema, + "mcp.tools.write", + OperationMeta.Target.ENGINE_CAPABILITY, + "collide", + null); + seed(catalog, "pdfOps", "collide", pdf); + seed(catalog, "aiOps", "collide", ai); + + // A disabled PDF op must resolve to empty, not a colliding AI capability of the same id. + assertTrue( + catalog.findByOperationId("collide").isEmpty(), + "disabled PDF op must not resolve to a colliding AI capability"); + + // A genuine AI-only id still resolves. + seed( + catalog, + "aiOps", + "ai-only", + new OperationMeta( + "ai-only", + OperationCategory.AI, + "ai op", + schema, + "mcp.tools.write", + OperationMeta.Target.ENGINE_CAPABILITY, + "ai-only", + null)); + assertEquals("ai-only", catalog.findByOperationId("ai-only").orElseThrow().id()); + } + + @Test + void blockedOperations_hidesOp() throws Exception { + ApplicationProperties props = new ApplicationProperties(); + props.getMcp().setBlockedOperations(List.of("compress-pdf")); + McpToolCatalog catalog = catalogWithEndpointsEnabled(props); + seed(catalog, "pdfOps", "compress-pdf", miscOp("compress-pdf")); + seed(catalog, "pdfOps", "ocr-pdf", miscOp("ocr-pdf")); + + assertTrue( + catalog.findByOperationId("compress-pdf").isEmpty(), "blocked op must be hidden"); + assertEquals("ocr-pdf", catalog.findByOperationId("ocr-pdf").orElseThrow().id()); + assertFalse(idsOf(catalog).contains("compress-pdf")); + assertTrue(idsOf(catalog).contains("ocr-pdf")); + } + + @Test + void allowedOperations_isWhitelist() throws Exception { + ApplicationProperties props = new ApplicationProperties(); + props.getMcp().setAllowedOperations(List.of("compress-pdf")); + McpToolCatalog catalog = catalogWithEndpointsEnabled(props); + seed(catalog, "pdfOps", "compress-pdf", miscOp("compress-pdf")); + seed(catalog, "pdfOps", "ocr-pdf", miscOp("ocr-pdf")); + + assertEquals("compress-pdf", catalog.findByOperationId("compress-pdf").orElseThrow().id()); + assertTrue( + catalog.findByOperationId("ocr-pdf").isEmpty(), + "op not on the allow-list must be hidden"); + assertEquals(List.of("compress-pdf"), idsOf(catalog)); + } + + @Test + void blockedOperations_takePrecedenceOverAllowed() throws Exception { + ApplicationProperties props = new ApplicationProperties(); + props.getMcp().setAllowedOperations(List.of("compress-pdf")); + props.getMcp().setBlockedOperations(List.of("compress-pdf")); + McpToolCatalog catalog = catalogWithEndpointsEnabled(props); + seed(catalog, "pdfOps", "compress-pdf", miscOp("compress-pdf")); + + assertTrue( + catalog.findByOperationId("compress-pdf").isEmpty(), + "block-list must win over allow-list"); + } + + @Test + void emptyAllowAndBlockLists_exposeAllEnabledOps() throws Exception { + McpToolCatalog catalog = catalogWithEndpointsEnabled(new ApplicationProperties()); + seed(catalog, "pdfOps", "compress-pdf", miscOp("compress-pdf")); + + assertEquals("compress-pdf", catalog.findByOperationId("compress-pdf").orElseThrow().id()); + assertTrue(idsOf(catalog).contains("compress-pdf")); + } + + private McpToolCatalog catalogWithEndpointsEnabled(ApplicationProperties props) { + ApplicationContext ctx = mock(ApplicationContext.class); + when(ctx.getBeansOfType(RequestMappingHandlerMapping.class)).thenReturn(Map.of()); + EndpointConfiguration endpoints = mock(EndpointConfiguration.class); + when(endpoints.isEndpointEnabledForUri(anyString())).thenReturn(true); + return new McpToolCatalog(ctx, endpoints, props, mapper); + } + + private OperationMeta miscOp(String id) { + return new OperationMeta( + id, + OperationCategory.MISC, + id, + mapper.createObjectNode(), + "mcp.tools.write", + OperationMeta.Target.JAVA_ENDPOINT, + "/api/v1/misc/" + id, + null); + } + + private static List idsOf(McpToolCatalog catalog) { + return catalog.enabledOps(OperationCategory.MISC).stream().map(OperationMeta::id).toList(); + } + + @SuppressWarnings("unchecked") + private static void seed(McpToolCatalog catalog, String field, String id, OperationMeta meta) + throws Exception { + Field f = McpToolCatalog.class.getDeclaredField(field); + f.setAccessible(true); + ((Map) f.get(catalog)).put(id, meta); + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/mcp/catalog/SimpleSchemaGeneratorTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/catalog/SimpleSchemaGeneratorTest.java new file mode 100644 index 0000000000..601c4c11cd --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/catalog/SimpleSchemaGeneratorTest.java @@ -0,0 +1,59 @@ +package stirling.software.proprietary.mcp.catalog; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.util.HashSet; +import java.util.Set; + +import org.junit.jupiter.api.Test; + +import com.fasterxml.jackson.annotation.JsonIgnore; +import com.fasterxml.jackson.annotation.JsonProperty; + +import tools.jackson.databind.ObjectMapper; +import tools.jackson.databind.node.ObjectNode; + +/** Schema must describe the JSON wire contract, not raw Java field names. */ +class SimpleSchemaGeneratorTest { + + private final SimpleSchemaGenerator gen = new SimpleSchemaGenerator(new ObjectMapper()); + + @SuppressWarnings("unused") + static class SampleRequest { + @JsonProperty("file_name") + String fileName; + + @JsonProperty(required = true) + String mode; + + @JsonIgnore String internalSecret; + + boolean flag; + + @jakarta.validation.constraints.NotBlank String title; + } + + @Test + void usesJsonPropertyNames_skipsJsonIgnore_marksRequired() { + ObjectNode schema = gen.toSchema(SampleRequest.class); + ObjectNode props = (ObjectNode) schema.get("properties"); + + assertTrue(props.has("file_name"), "must use the @JsonProperty name"); + assertFalse(props.has("fileName"), "must not emit the raw field name"); + + assertFalse(props.has("internalSecret"), "@JsonIgnore field must be skipped"); + + assertTrue(props.has("flag")); + assertEquals("boolean", props.get("flag").get("type").asText()); + + assertNotNull(schema.get("required"), "required array expected"); + Set required = new HashSet<>(); + schema.get("required").forEach(n -> required.add(n.asText())); + assertTrue(required.contains("mode"), "@JsonProperty(required=true) -> required"); + assertTrue(required.contains("title"), "@NotBlank -> required"); + assertFalse(required.contains("file_name"), "optional field not required"); + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/mcp/engine/EngineCapabilityParseTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/engine/EngineCapabilityParseTest.java new file mode 100644 index 0000000000..9be2d9908d --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/engine/EngineCapabilityParseTest.java @@ -0,0 +1,138 @@ +package stirling.software.proprietary.mcp.engine; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.lang.reflect.Method; +import java.util.Map; + +import org.junit.jupiter.api.Test; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.proprietary.mcp.catalog.McpToolCatalog; +import stirling.software.proprietary.mcp.catalog.OperationCategory; +import stirling.software.proprietary.mcp.catalog.OperationMeta; + +import tools.jackson.databind.ObjectMapper; + +/** + * Verifies {@link EngineCapabilityClient#parseManifest} maps a manifest to {@link OperationMeta}. + */ +class EngineCapabilityParseTest { + + @Test + void manifest_maps_to_operation_meta() throws Exception { + ObjectMapper mapper = new ObjectMapper(); + ApplicationProperties props = new ApplicationProperties(); + // parseManifest doesn't touch the catalog, so a null catalog is fine. + EngineCapabilityClient client = + new EngineCapabilityClient(props, (McpToolCatalog) null, mapper); + + String body = + """ + { + "version": 1, + "capabilities": [ + { + "id": "pdf-question-answer", + "description": "Answer a question about a PDF.", + "input_schema": {"type": "object", "properties": {"question": {"type": "string"}}}, + "mode": "sync", + "required_scope": "mcp.tools.read", + "route": "/api/v1/pdf-question" + }, + { + "id": "pdf-edit-plan", + "description": "Produce an edit plan from natural language.", + "input_schema": {"type": "object"}, + "mode": "async", + "required_scope": "mcp.tools.write", + "route": "/api/v1/pdf-edit" + } + ] + } + """; + + Method parse = + EngineCapabilityClient.class.getDeclaredMethod("parseManifest", String.class); + parse.setAccessible(true); + @SuppressWarnings("unchecked") + Map parsed = (Map) parse.invoke(client, body); + + assertThat(parsed).hasSize(2); + + OperationMeta qa = parsed.get("pdf-question-answer"); + assertThat(qa).isNotNull(); + assertThat(qa.category()).isEqualTo(OperationCategory.AI); + assertThat(qa.requiredScope()).isEqualTo("mcp.tools.read"); + assertThat(qa.endpointPath()).isEqualTo("/api/v1/pdf-question"); + assertThat(qa.target()).isEqualTo(OperationMeta.Target.ENGINE_CAPABILITY); + assertThat(qa.paramSchema().get("type").asText()).isEqualTo("object"); + + OperationMeta edit = parsed.get("pdf-edit-plan"); + assertThat(edit).isNotNull(); + assertThat(edit.requiredScope()).isEqualTo("mcp.tools.write"); + assertThat(edit.endpointPath()).isEqualTo("/api/v1/pdf-edit"); + } + + @Test + void missing_capabilities_array_throws() { + ObjectMapper mapper = new ObjectMapper(); + EngineCapabilityClient client = + new EngineCapabilityClient( + new ApplicationProperties(), (McpToolCatalog) null, mapper); + + try { + Method parse = + EngineCapabilityClient.class.getDeclaredMethod("parseManifest", String.class); + parse.setAccessible(true); + parse.invoke(client, "{\"version\":1}"); + org.junit.jupiter.api.Assertions.fail("Expected IOException"); + } catch (java.lang.reflect.InvocationTargetException e) { + assertThat(e.getCause()).isInstanceOf(java.io.IOException.class); + } catch (Exception e) { + org.junit.jupiter.api.Assertions.fail("Unexpected exception: " + e); + } + } + + @Test + void unsafe_routes_are_skipped_and_blank_scope_fails_safe() throws Exception { + ObjectMapper mapper = new ObjectMapper(); + EngineCapabilityClient client = + new EngineCapabilityClient( + new ApplicationProperties(), (McpToolCatalog) null, mapper); + + String body = + """ + {"version":1,"capabilities":[ + {"id":"good","description":"ok","input_schema":{"type":"object"},"required_scope":"","route":"/api/v1/pdf-question"}, + {"id":"ssrf-at","description":"x","input_schema":{"type":"object"},"route":"@evil.com/steal"}, + {"id":"ssrf-proto","description":"x","input_schema":{"type":"object"},"route":"//evil.com/x"}, + {"id":"escape","description":"x","input_schema":{"type":"object"},"route":"/api/../../internal/secret"}, + {"id":"scheme","description":"x","input_schema":{"type":"object"},"route":"http://evil.com/x"}, + {"id":"non-api","description":"x","input_schema":{"type":"object"},"route":"/admin/settings"} + ]} + """; + + Method parse = + EngineCapabilityClient.class.getDeclaredMethod("parseManifest", String.class); + parse.setAccessible(true); + @SuppressWarnings("unchecked") + Map parsed = (Map) parse.invoke(client, body); + + // Only the safe, server-relative /api route survives. + assertThat(parsed.keySet()).containsExactly("good"); + // Blank required_scope fails safe to the stricter write scope. + assertThat(parsed.get("good").requiredScope()).isEqualTo("mcp.tools.write"); + } + + @Test + void isSafeRelativeRoute_acceptsOnlyServerRelativeApiPaths() { + assertThat(EngineCapabilityClient.isSafeRelativeRoute("/api/v1/pdf-question")).isTrue(); + assertThat(EngineCapabilityClient.isSafeRelativeRoute("@evil.com/x")).isFalse(); + assertThat(EngineCapabilityClient.isSafeRelativeRoute("//evil.com/x")).isFalse(); + assertThat(EngineCapabilityClient.isSafeRelativeRoute("/api/../secret")).isFalse(); + assertThat(EngineCapabilityClient.isSafeRelativeRoute("http://evil/x")).isFalse(); + assertThat(EngineCapabilityClient.isSafeRelativeRoute("/admin/x")).isFalse(); + assertThat(EngineCapabilityClient.isSafeRelativeRoute("/api/v1/x y")).isFalse(); + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/mcp/security/McpApiKeyIntegrationTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/security/McpApiKeyIntegrationTest.java new file mode 100644 index 0000000000..b8c1f953e9 --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/security/McpApiKeyIntegrationTest.java @@ -0,0 +1,148 @@ +package stirling.software.proprietary.mcp.security; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.net.URI; +import java.net.http.HttpClient; +import java.net.http.HttpRequest; +import java.net.http.HttpResponse; +import java.util.Optional; + +import org.junit.jupiter.api.Test; +import org.springframework.boot.SpringBootConfiguration; +import org.springframework.boot.autoconfigure.EnableAutoConfiguration; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.boot.test.web.server.LocalServerPort; +import org.springframework.context.annotation.Bean; +import org.springframework.context.annotation.Import; +import org.springframework.test.context.DynamicPropertyRegistry; +import org.springframework.test.context.DynamicPropertySource; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.proprietary.mcp.McpServerController; +import stirling.software.proprietary.mcp.tools.DescribeOperationTool; +import stirling.software.proprietary.mcp.tools.StirlingAiTool; +import stirling.software.proprietary.mcp.tools.StirlingConvertTool; +import stirling.software.proprietary.mcp.tools.StirlingMiscTool; +import stirling.software.proprietary.mcp.tools.StirlingPagesTool; +import stirling.software.proprietary.mcp.tools.StirlingSecurityTool; +import stirling.software.proprietary.security.model.User; +import stirling.software.proprietary.security.service.UserService; + +/** + * End-to-end test of {@code mcp.auth.mode=apikey} against the real security chain on live Jetty. + */ +@SpringBootTest( + classes = McpApiKeyIntegrationTest.TestApp.class, + webEnvironment = SpringBootTest.WebEnvironment.RANDOM_PORT) +class McpApiKeyIntegrationTest { + + private static final String VALID_KEY = "stirling-test-key-abc123"; + + @LocalServerPort private int port; + private final HttpClient http = HttpClient.newHttpClient(); + + @DynamicPropertySource + static void mcpProperties(DynamicPropertyRegistry registry) { + registry.add("mcp.enabled", () -> "true"); + registry.add("mcp.auth.mode", () -> "apikey"); + } + + @Test + void validApiKeyViaHeader_callsToolsList() throws Exception { + HttpResponse response = + postMcp( + b -> b.header("X-API-KEY", VALID_KEY), + "{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"tools/list\"}"); + assertThat(response.statusCode()).isEqualTo(200); + assertThat(response.body()).contains("stirling_describe_operation"); + } + + @Test + void validApiKeyViaBearer_callsToolsList() throws Exception { + HttpResponse response = + postMcp( + b -> b.header("Authorization", "Bearer " + VALID_KEY), + "{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"tools/list\"}"); + assertThat(response.statusCode()).isEqualTo(200); + } + + @Test + void noKey_isRejectedWith401() throws Exception { + HttpResponse response = + postMcp(b -> b, "{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"ping\"}"); + assertThat(response.statusCode()).isEqualTo(401); + } + + @Test + void wrongKey_isRejectedWith401() throws Exception { + HttpResponse response = + postMcp( + b -> b.header("X-API-KEY", "not-a-real-key"), + "{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"ping\"}"); + assertThat(response.statusCode()).isEqualTo(401); + } + + @Test + void noOAuthMetadataInApiKeyMode() throws Exception { + HttpRequest req = + HttpRequest.newBuilder() + .uri(URI.create(base() + "/.well-known/oauth-protected-resource")) + .GET() + .build(); + HttpResponse response = http.send(req, HttpResponse.BodyHandlers.ofString()); + assertThat(response.statusCode()).isNotEqualTo(200); + } + + private HttpResponse postMcp( + java.util.function.UnaryOperator headers, String body) + throws Exception { + HttpRequest.Builder builder = + HttpRequest.newBuilder() + .uri(URI.create(base() + "/mcp")) + .header("Content-Type", "application/json") + .POST(HttpRequest.BodyPublishers.ofString(body)); + return http.send(headers.apply(builder).build(), HttpResponse.BodyHandlers.ofString()); + } + + private String base() { + return "http://localhost:" + port; + } + + @SpringBootConfiguration + @EnableAutoConfiguration + @Import({ + McpSecurityConfig.class, + McpServerController.class, + DescribeOperationTool.class, + StirlingConvertTool.class, + StirlingPagesTool.class, + StirlingMiscTool.class, + StirlingSecurityTool.class, + StirlingAiTool.class + }) + static class TestApp { + + @Bean + ApplicationProperties applicationProperties() { + ApplicationProperties props = new ApplicationProperties(); + props.getMcp().setEnabled(true); + props.getMcp().getAuth().setMode("apikey"); + props.getAutomaticallyGenerated().setAppVersion("test"); + return props; + } + + @Bean + UserService userService() { + UserService mock = org.mockito.Mockito.mock(UserService.class); + User account = org.mockito.Mockito.mock(User.class); + org.mockito.Mockito.when(account.isEnabled()).thenReturn(true); + org.mockito.Mockito.when(account.getUsername()).thenReturn("alice"); + org.mockito.Mockito.when(mock.getUserByApiKey(org.mockito.ArgumentMatchers.anyString())) + .thenReturn(Optional.empty()); + org.mockito.Mockito.when(mock.getUserByApiKey(VALID_KEY)) + .thenReturn(Optional.of(account)); + return mock; + } + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/mcp/security/McpAudienceValidatorTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/security/McpAudienceValidatorTest.java new file mode 100644 index 0000000000..5bee580b7e --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/security/McpAudienceValidatorTest.java @@ -0,0 +1,88 @@ +package stirling.software.proprietary.mcp.security; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.time.Instant; +import java.util.List; +import java.util.Map; + +import org.junit.jupiter.api.Test; +import org.springframework.security.oauth2.core.OAuth2TokenValidatorResult; +import org.springframework.security.oauth2.jwt.Jwt; + +/** RFC 8707 audience validator: token {@code aud} must list the resource id; blank fails closed. */ +class McpAudienceValidatorTest { + + private static final String RESOURCE = "http://localhost:8080/mcp"; + + private final McpAudienceValidator validator = new McpAudienceValidator(RESOURCE); + + @Test + void matchingAudience_isAccepted() { + Jwt token = tokenWithAudience(List.of(RESOURCE)); + OAuth2TokenValidatorResult result = validator.validate(token); + assertThat(result.hasErrors()).isFalse(); + } + + @Test + void multiAudienceIncludingResource_isAccepted() { + Jwt token = tokenWithAudience(List.of("https://other.example.com", RESOURCE)); + OAuth2TokenValidatorResult result = validator.validate(token); + assertThat(result.hasErrors()).isFalse(); + } + + @Test + void wrongAudience_isRejected() { + Jwt token = tokenWithAudience(List.of("https://other.example.com")); + OAuth2TokenValidatorResult result = validator.validate(token); + assertThat(result.hasErrors()).isTrue(); + assertThat(result.getErrors()).anyMatch(e -> e.getErrorCode().equals("invalid_token")); + } + + @Test + void missingAudienceClaim_isRejected() { + Jwt token = tokenWithAudience(null); + OAuth2TokenValidatorResult result = validator.validate(token); + assertThat(result.hasErrors()).isTrue(); + } + + @Test + void blankResourceId_failsClosed_rejectingEvenMatchingTokens() { + McpAudienceValidator blank = new McpAudienceValidator(""); + OAuth2TokenValidatorResult result = blank.validate(tokenWithAudience(List.of(RESOURCE))); + assertThat(result.hasErrors()).isTrue(); + } + + @Test + void acceptedAudience_isAccepted_alongsideResourceId() { + // Supabase-style IdP: every token carries aud=authenticated, never the resource id. + McpAudienceValidator relaxed = new McpAudienceValidator(RESOURCE, List.of("authenticated")); + assertThat(relaxed.validate(tokenWithAudience(List.of("authenticated"))).hasErrors()) + .isFalse(); + assertThat(relaxed.validate(tokenWithAudience(List.of(RESOURCE))).hasErrors()).isFalse(); + assertThat(relaxed.validate(tokenWithAudience(List.of("something-else"))).hasErrors()) + .isTrue(); + } + + @Test + void blankAcceptedAudienceEntries_areIgnored() { + McpAudienceValidator relaxed = new McpAudienceValidator(RESOURCE, List.of("", " ")); + assertThat(relaxed.validate(tokenWithAudience(List.of(""))).hasErrors()).isTrue(); + assertThat(relaxed.validate(tokenWithAudience(List.of(RESOURCE))).hasErrors()).isFalse(); + } + + @Test + void blankResourceIdWithOnlyBlankAccepted_failsClosed() { + McpAudienceValidator blank = new McpAudienceValidator("", List.of(" ")); + assertThat(blank.validate(tokenWithAudience(List.of(RESOURCE))).hasErrors()).isTrue(); + } + + private static Jwt tokenWithAudience(List audience) { + return new Jwt( + "header.payload.signature", + Instant.now(), + Instant.now().plusSeconds(60), + Map.of("alg", "RS256"), + audience == null ? Map.of("sub", "u1") : Map.of("sub", "u1", "aud", audience)); + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/mcp/security/McpAuthenticationEntryPointTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/security/McpAuthenticationEntryPointTest.java new file mode 100644 index 0000000000..e098fe875f --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/security/McpAuthenticationEntryPointTest.java @@ -0,0 +1,89 @@ +package stirling.software.proprietary.mcp.security; + +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.ArgumentMatchers.anyInt; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import org.junit.jupiter.api.Test; +import org.mockito.ArgumentCaptor; +import org.springframework.security.oauth2.core.OAuth2AuthenticationException; +import org.springframework.security.oauth2.core.OAuth2Error; + +import jakarta.servlet.http.HttpServletRequest; +import jakarta.servlet.http.HttpServletResponse; + +/** The resource_metadata URL must reflect the public host behind a reverse proxy. */ +class McpAuthenticationEntryPointTest { + + private static final String META = "/.well-known/oauth-protected-resource"; + private final McpAuthenticationEntryPoint entryPoint = new McpAuthenticationEntryPoint(META); + + @Test + void usesForwardedHeadersForMetadataUrl() throws Exception { + HttpServletRequest req = mock(HttpServletRequest.class); + when(req.getScheme()).thenReturn("http"); + when(req.getServerName()).thenReturn("internal-host"); + when(req.getServerPort()).thenReturn(8080); + when(req.getHeader("X-Forwarded-Proto")).thenReturn("https"); + when(req.getHeader("X-Forwarded-Host")).thenReturn("mcp.example.com"); + when(req.getHeader("X-Forwarded-Port")).thenReturn("443"); + HttpServletResponse resp = mock(HttpServletResponse.class); + + entryPoint.commence(req, resp, null); + + ArgumentCaptor header = ArgumentCaptor.forClass(String.class); + verify(resp).setHeader(eq("WWW-Authenticate"), header.capture()); + String www = header.getValue(); + assertTrue( + www.contains("resource_metadata=\"https://mcp.example.com" + META + "\""), + "must use forwarded host/proto, got: " + www); + assertFalse(www.contains("internal-host"), "internal host must not leak"); + verify(resp).sendError(anyInt(), anyString()); + } + + @Test + void fallsBackToServletHostWithoutForwardedHeaders() throws Exception { + HttpServletRequest req = mock(HttpServletRequest.class); + when(req.getScheme()).thenReturn("http"); + when(req.getServerName()).thenReturn("localhost"); + when(req.getServerPort()).thenReturn(8080); + HttpServletResponse resp = mock(HttpServletResponse.class); + + entryPoint.commence(req, resp, null); + + ArgumentCaptor header = ArgumentCaptor.forClass(String.class); + verify(resp).setHeader(eq("WWW-Authenticate"), header.capture()); + assertTrue(header.getValue().contains("http://localhost:8080" + META), header.getValue()); + } + + @Test + void surfacesRejectionReasonWhenTokenPresented() throws Exception { + HttpServletRequest req = mock(HttpServletRequest.class); + when(req.getScheme()).thenReturn("https"); + when(req.getServerName()).thenReturn("mcp.example.com"); + when(req.getServerPort()).thenReturn(443); + when(req.getHeader("Authorization")).thenReturn("Bearer bad.token"); + HttpServletResponse resp = mock(HttpServletResponse.class); + + OAuth2Error error = + new OAuth2Error( + "invalid_token", + "Token audience does not include this server's resource id" + + " (https://mcp.example.com/mcp).", + null); + + entryPoint.commence(req, resp, new OAuth2AuthenticationException(error)); + + ArgumentCaptor header = ArgumentCaptor.forClass(String.class); + verify(resp).setHeader(eq("WWW-Authenticate"), header.capture()); + String www = header.getValue(); + assertTrue( + www.contains("error_description=\"invalid_token - Token audience does not include"), + "must surface the rejection reason, got: " + www); + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/mcp/security/McpConfigValidatorTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/security/McpConfigValidatorTest.java new file mode 100644 index 0000000000..a933ec37d2 --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/security/McpConfigValidatorTest.java @@ -0,0 +1,151 @@ +package stirling.software.proprietary.mcp.security; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.util.List; + +import org.junit.jupiter.api.Test; + +import stirling.software.common.model.ApplicationProperties; + +class McpConfigValidatorTest { + + private static ApplicationProperties.Mcp newMcp() { + return new ApplicationProperties.Mcp(); + } + + private static boolean hasWarn(List findings, String needle) { + return findings.stream() + .anyMatch( + f -> + f.severity() == McpConfigValidator.Severity.WARN + && f.message().contains(needle)); + } + + @Test + void apiKeyModeSkipsOAuthChecks() { + ApplicationProperties.Mcp mcp = newMcp(); + mcp.getAuth().setMode("apikey"); + + List findings = McpConfigValidator.validate(mcp); + + assertEquals(1, findings.size()); + assertEquals(McpConfigValidator.Severity.INFO, findings.get(0).severity()); + assertTrue(findings.get(0).message().contains("apikey")); + } + + @Test + void blankIssuerAndResourceProduceWarnings() { + // Defaults: oauth mode, blank issuer-uri and resource-id. + List findings = McpConfigValidator.validate(newMcp()); + + assertTrue(hasWarn(findings, "issuer-uri"), "blank issuer must warn"); + assertTrue(hasWarn(findings, "resource-id"), "blank resource-id must warn"); + } + + @Test + void subClaimWithRequireAccountWarns() { + ApplicationProperties.Mcp mcp = newMcp(); + mcp.getAuth().setIssuerUri("https://issuer.example.com"); + mcp.getAuth().setResourceId("https://host.example.com/mcp"); + // Defaults username-claim=sub, require-existing-account=true. + + assertTrue(hasWarn(McpConfigValidator.validate(mcp), "username-claim='sub'")); + } + + @Test + void completeConfigReportsReadyWithNoWarnings() { + ApplicationProperties.Mcp mcp = newMcp(); + mcp.getAuth().setIssuerUri("https://issuer.example.com"); + mcp.getAuth().setResourceId("https://host.example.com/mcp"); + mcp.getAuth().setUsernameClaim("email"); + mcp.setScopesEnabled(false); + + List findings = McpConfigValidator.validate(mcp); + + assertTrue( + findings.stream().noneMatch(f -> f.severity() == McpConfigValidator.Severity.WARN), + "complete config must have no warnings"); + assertTrue(findings.stream().anyMatch(f -> f.message().contains("look complete"))); + } + + @Test + void acceptedAudiencesCoverBlankResourceId() { + ApplicationProperties.Mcp mcp = newMcp(); + mcp.getAuth().setIssuerUri("https://issuer.example.com"); + mcp.getAuth().setResourceId(""); + mcp.getAuth().setAcceptedAudiences(List.of("authenticated")); + mcp.getAuth().setUsernameClaim("email"); + mcp.setScopesEnabled(false); + + List findings = McpConfigValidator.validate(mcp); + + assertFalse( + hasWarn(findings, "fails closed"), + "accepted-audiences must satisfy audience binding without a resource id"); + assertTrue( + findings.stream().anyMatch(f -> f.message().contains("accepted-audiences=")), + "configured accepted-audiences should be surfaced"); + } + + @Test + void strictAudienceHintsAtAcceptedAudiencesEscapeHatch() { + ApplicationProperties.Mcp mcp = newMcp(); + mcp.getAuth().setIssuerUri("https://issuer.example.com"); + mcp.getAuth().setResourceId("https://host.example.com/mcp"); + mcp.getAuth().setUsernameClaim("email"); + mcp.setScopesEnabled(false); + + List findings = McpConfigValidator.validate(mcp); + + assertTrue( + findings.stream().anyMatch(f -> f.message().contains("accepted-audiences")), + "should point coarse-audience IdPs at accepted-audiences"); + } + + @Test + void unrecognizedModeWarnsAboutOAuthFallback() { + ApplicationProperties.Mcp mcp = newMcp(); + mcp.getAuth().setMode("api-key"); // near-miss typo that silently runs the OAuth chain + + assertTrue(hasWarn(McpConfigValidator.validate(mcp), "is not recognized")); + } + + @Test + void requireExistingAccountFalseWarnsAboutOpenAccess() { + ApplicationProperties.Mcp mcp = newMcp(); + mcp.getAuth().setIssuerUri("https://issuer.example.com"); + mcp.getAuth().setResourceId("https://host.example.com/mcp"); + mcp.getAuth().setUsernameClaim("email"); + mcp.getAuth().setRequireExistingAccount(false); + + assertTrue(hasWarn(McpConfigValidator.validate(mcp), "require-existing-account=false")); + } + + @Test + void nonUrlResourceIdWarns() { + ApplicationProperties.Mcp mcp = newMcp(); + mcp.getAuth().setIssuerUri("https://issuer.example.com"); + mcp.getAuth().setResourceId("localhost:8080/mcp"); // missing scheme + + assertTrue(hasWarn(McpConfigValidator.validate(mcp), "is not an http(s) URL")); + } + + @Test + void allowListIsFlaggedAndOverlapWithBlockListWarns() { + ApplicationProperties.Mcp mcp = newMcp(); + mcp.getAuth().setIssuerUri("https://issuer.example.com"); + mcp.getAuth().setResourceId("https://host.example.com/mcp"); + mcp.setAllowedOperations(List.of("merge-pdfs", "split-pdf")); + mcp.setBlockedOperations(List.of("split-pdf")); + + List findings = McpConfigValidator.validate(mcp); + + assertTrue( + findings.stream().anyMatch(f -> f.message().contains("strict allow-list")), + "an allow-list should be surfaced"); + assertTrue(hasWarn(findings, "blocked wins"), "allowed+blocked overlap should warn"); + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/mcp/security/McpOAuthIntegrationTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/security/McpOAuthIntegrationTest.java new file mode 100644 index 0000000000..1e814ea580 --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/security/McpOAuthIntegrationTest.java @@ -0,0 +1,356 @@ +package stirling.software.proprietary.mcp.security; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.net.URI; +import java.net.http.HttpClient; +import java.net.http.HttpRequest; +import java.net.http.HttpResponse; +import java.security.KeyPair; +import java.security.KeyPairGenerator; +import java.security.interfaces.RSAPrivateKey; +import java.security.interfaces.RSAPublicKey; +import java.time.Instant; +import java.util.Date; +import java.util.List; + +import org.junit.jupiter.api.AfterAll; +import org.junit.jupiter.api.BeforeAll; +import org.junit.jupiter.api.Test; +import org.springframework.boot.SpringBootConfiguration; +import org.springframework.boot.autoconfigure.EnableAutoConfiguration; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.boot.test.web.server.LocalServerPort; +import org.springframework.context.annotation.Bean; +import org.springframework.context.annotation.Import; +import org.springframework.test.context.DynamicPropertyRegistry; +import org.springframework.test.context.DynamicPropertySource; + +import com.nimbusds.jose.JWSAlgorithm; +import com.nimbusds.jose.JWSHeader; +import com.nimbusds.jose.crypto.RSASSASigner; +import com.nimbusds.jose.jwk.JWKSet; +import com.nimbusds.jose.jwk.RSAKey; +import com.nimbusds.jwt.JWTClaimsSet; +import com.nimbusds.jwt.SignedJWT; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.proprietary.mcp.McpServerController; +import stirling.software.proprietary.mcp.tools.DescribeOperationTool; +import stirling.software.proprietary.mcp.tools.StirlingAiTool; +import stirling.software.proprietary.mcp.tools.StirlingConvertTool; +import stirling.software.proprietary.mcp.tools.StirlingMiscTool; +import stirling.software.proprietary.mcp.tools.StirlingPagesTool; +import stirling.software.proprietary.mcp.tools.StirlingSecurityTool; +import stirling.software.proprietary.security.service.UserService; + +import okhttp3.mockwebserver.Dispatcher; +import okhttp3.mockwebserver.MockResponse; +import okhttp3.mockwebserver.MockWebServer; +import okhttp3.mockwebserver.RecordedRequest; + +/** + * End-to-end OAuth test against the real {@link McpSecurityConfig} chain. A real RSA keypair signs + * JWTs; the public key is served as JWKS over HTTP (mockwebserver) and the resource server fetches + * and validates against it. The JDK HttpClient drives a live Jetty instance on a random port. + */ +@SpringBootTest( + classes = McpOAuthIntegrationTest.TestApp.class, + webEnvironment = SpringBootTest.WebEnvironment.RANDOM_PORT) +class McpOAuthIntegrationTest { + + private static final String ISSUER = "https://test-issuer.example.com"; + private static final String RESOURCE_ID = "http://localhost/mcp"; + + private static MockWebServer jwksServer; + private static RSAPrivateKey privateKey; + private static final String KEY_ID = "mcp-test-key"; + + @LocalServerPort private int port; + + private final HttpClient http = HttpClient.newHttpClient(); + + @BeforeAll + static void startJwks() throws Exception { + KeyPairGenerator gen = KeyPairGenerator.getInstance("RSA"); + gen.initialize(2048); + KeyPair kp = gen.generateKeyPair(); + privateKey = (RSAPrivateKey) kp.getPrivate(); + RSAPublicKey publicKey = (RSAPublicKey) kp.getPublic(); + + RSAKey jwk = new RSAKey.Builder(publicKey).keyID(KEY_ID).build(); + String jwksJson = new JWKSet(jwk).toString(); + + jwksServer = new MockWebServer(); + jwksServer.setDispatcher( + new Dispatcher() { + @Override + public MockResponse dispatch(RecordedRequest request) { + return new MockResponse() + .setHeader("Content-Type", "application/json") + .setBody(jwksJson); + } + }); + jwksServer.start(); + } + + @AfterAll + static void stopJwks() throws Exception { + if (jwksServer != null) { + jwksServer.shutdown(); + } + } + + @DynamicPropertySource + static void mcpProperties(DynamicPropertyRegistry registry) { + registry.add("mcp.enabled", () -> "true"); + registry.add("mcp.auth.issuer-uri", () -> ISSUER); + registry.add("mcp.auth.jwks-uri", () -> jwksServer.url("/jwks").toString()); + registry.add("mcp.auth.resource-id", () -> RESOURCE_ID); + } + + @Test + void noToken_returns401WithResourceMetadataHeader() throws Exception { + HttpResponse response = + postMcp(null, "{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"ping\"}"); + assertThat(response.statusCode()).isEqualTo(401); + String wwwAuth = response.headers().firstValue("WWW-Authenticate").orElse(""); + assertThat(wwwAuth).contains("resource_metadata="); + // The advertised URL must be the RFC 9728 path-inserted form for the /mcp resource. + assertThat(wwwAuth).contains("/.well-known/oauth-protected-resource/mcp"); + } + + @Test + void validToken_callsToolsListSuccessfully() throws Exception { + String token = + mintToken( + ISSUER, + List.of(RESOURCE_ID), + "mcp.tools.read mcp.tools.write", + Instant.now().plusSeconds(300)); + HttpResponse response = + postMcp(token, "{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"tools/list\"}"); + assertThat(response.statusCode()).isEqualTo(200); + assertThat(response.body()).contains("stirling_describe_operation"); + } + + @Test + void wrongAudience_isRejected() throws Exception { + String token = + mintToken( + ISSUER, + List.of("https://some-other-resource.example.com"), + "mcp.tools.read", + Instant.now().plusSeconds(300)); + HttpResponse response = + postMcp(token, "{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"ping\"}"); + assertThat(response.statusCode()).isEqualTo(401); + } + + @Test + void expiredToken_isRejected() throws Exception { + String token = + mintToken( + ISSUER, + List.of(RESOURCE_ID), + "mcp.tools.read", + Instant.now().minusSeconds(60)); + HttpResponse response = + postMcp(token, "{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"ping\"}"); + assertThat(response.statusCode()).isEqualTo(401); + } + + @Test + void validTokenButNoStirlingAccount_isRejectedWith403() throws Exception { + // 'ghost-user' is not provisioned, so account-binding rejects an otherwise-valid token. + String token = + mintToken( + "ghost-user", + ISSUER, + List.of(RESOURCE_ID), + "mcp.tools.read mcp.tools.write", + Instant.now().plusSeconds(300)); + HttpResponse response = + postMcp(token, "{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"tools/list\"}"); + assertThat(response.statusCode()).isEqualTo(403); + } + + @Test + void oversizedBody_isRejectedWith413() throws Exception { + String token = + mintToken( + ISSUER, + List.of(RESOURCE_ID), + "mcp.tools.read mcp.tools.write", + Instant.now().plusSeconds(300)); + StringBuilder big = + new StringBuilder("{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"x\",\"params\":\""); + big.append("A".repeat(300 * 1024)); + big.append("\"}"); + HttpResponse response = postMcp(token, big.toString()); + assertThat(response.statusCode()).isEqualTo(413); + } + + @Test + void oversizedChunkedBody_isRejectedWith413() throws Exception { + // No Content-Length (chunked transfer) exercises the streaming cap rather than the fast + // check. + String token = + mintToken( + ISSUER, + List.of(RESOURCE_ID), + "mcp.tools.read mcp.tools.write", + Instant.now().plusSeconds(300)); + byte[] big = + ("{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"x\",\"params\":\"" + + "A".repeat(300 * 1024) + + "\"}") + .getBytes(java.nio.charset.StandardCharsets.UTF_8); + HttpRequest request = + HttpRequest.newBuilder() + .uri(URI.create(base() + "/mcp")) + .header("Content-Type", "application/json") + .header("Authorization", "Bearer " + token) + .POST( + HttpRequest.BodyPublishers.ofInputStream( + () -> new java.io.ByteArrayInputStream(big))) + .build(); + HttpResponse response = http.send(request, HttpResponse.BodyHandlers.ofString()); + assertThat(response.statusCode()).isEqualTo(413); + } + + @Test + void malformedJson_returnsJsonRpcParseErrorEnvelope() throws Exception { + // Unparseable JSON must be wrapped as a JSON-RPC Parse error (-32700), not Spring's HTML + // 400. + String token = + mintToken( + ISSUER, + List.of(RESOURCE_ID), + "mcp.tools.read mcp.tools.write", + Instant.now().plusSeconds(300)); + HttpResponse response = postMcp(token, "{ this is not valid json "); + assertThat(response.statusCode()).isEqualTo(400); + assertThat(response.body()).contains("-32700"); + } + + @Test + void metadataEndpoint_isReachableWithoutToken() throws Exception { + HttpRequest request = + HttpRequest.newBuilder() + .uri(URI.create(base() + "/.well-known/oauth-protected-resource")) + .GET() + .build(); + HttpResponse response = http.send(request, HttpResponse.BodyHandlers.ofString()); + assertThat(response.statusCode()).isEqualTo(200); + assertThat(response.body()).contains(RESOURCE_ID); + assertThat(response.body()).contains(ISSUER); + assertThat(response.body()).contains("mcp.tools.read"); + } + + @Test + void pathInsertedMetadataEndpoint_servesCustomizedMetadata() throws Exception { + // RFC 9728 path-inserted form for the /mcp resource. Must be served by the MCP chain + // with authorization_servers populated; a default/uncustomized document here makes MCP + // clients fall back to treating this server as its own authorization server. + HttpRequest request = + HttpRequest.newBuilder() + .uri(URI.create(base() + "/.well-known/oauth-protected-resource/mcp")) + .GET() + .build(); + HttpResponse response = http.send(request, HttpResponse.BodyHandlers.ofString()); + assertThat(response.statusCode()).isEqualTo(200); + assertThat(response.body()).contains(RESOURCE_ID); + assertThat(response.body()).contains("authorization_servers"); + assertThat(response.body()).contains(ISSUER); + assertThat(response.body()).contains("mcp.tools.read"); + } + + private HttpResponse postMcp(String token, String body) throws Exception { + HttpRequest.Builder builder = + HttpRequest.newBuilder() + .uri(URI.create(base() + "/mcp")) + .header("Content-Type", "application/json") + .POST(HttpRequest.BodyPublishers.ofString(body)); + if (token != null) { + builder.header("Authorization", "Bearer " + token); + } + return http.send(builder.build(), HttpResponse.BodyHandlers.ofString()); + } + + private String base() { + return "http://localhost:" + port; + } + + private static String mintToken( + String issuer, List audience, String scope, Instant expiry) { + return mintToken("test-user", issuer, audience, scope, expiry); + } + + private static String mintToken( + String subject, String issuer, List audience, String scope, Instant expiry) { + try { + JWTClaimsSet claims = + new JWTClaimsSet.Builder() + .issuer(issuer) + .subject(subject) + .audience(audience) + .claim("scope", scope) + .issueTime(Date.from(Instant.now().minusSeconds(5))) + .expirationTime(Date.from(expiry)) + .build(); + SignedJWT jwt = + new SignedJWT( + new JWSHeader.Builder(JWSAlgorithm.RS256).keyID(KEY_ID).build(), + claims); + jwt.sign(new RSASSASigner(privateKey)); + return jwt.serialize(); + } catch (Exception e) { + throw new RuntimeException(e); + } + } + + @SpringBootConfiguration + @EnableAutoConfiguration + @Import({ + McpSecurityConfig.class, + McpServerController.class, + DescribeOperationTool.class, + StirlingConvertTool.class, + StirlingPagesTool.class, + StirlingMiscTool.class, + StirlingSecurityTool.class, + StirlingAiTool.class + }) + static class TestApp { + + @Bean + ApplicationProperties applicationProperties() { + ApplicationProperties props = new ApplicationProperties(); + props.getMcp().setEnabled(true); + props.getMcp().setMaxRequestBytes(256L * 1024); + props.getMcp().getAuth().setIssuerUri(ISSUER); + props.getMcp().getAuth().setJwksUri(jwksServer.url("/jwks").toString()); + props.getMcp().getAuth().setResourceId(RESOURCE_ID); + props.getAutomaticallyGenerated().setAppVersion("test"); + return props; + } + + /** Stub UserService: only 'test-user' is a provisioned, enabled account. */ + @Bean + UserService userService() { + UserService mock = org.mockito.Mockito.mock(UserService.class); + stirling.software.proprietary.security.model.User account = + org.mockito.Mockito.mock( + stirling.software.proprietary.security.model.User.class); + org.mockito.Mockito.when(account.isEnabled()).thenReturn(true); + org.mockito.Mockito.when(account.getUsername()).thenReturn("test-user"); + org.mockito.Mockito.when( + mock.findByUsernameIgnoreCase(org.mockito.ArgumentMatchers.anyString())) + .thenReturn(java.util.Optional.empty()); + org.mockito.Mockito.when(mock.findByUsernameIgnoreCase("test-user")) + .thenReturn(java.util.Optional.of(account)); + return mock; + } + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/mcp/security/McpScopeMetadataDisabledTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/security/McpScopeMetadataDisabledTest.java new file mode 100644 index 0000000000..17d15d1f11 --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/security/McpScopeMetadataDisabledTest.java @@ -0,0 +1,105 @@ +package stirling.software.proprietary.mcp.security; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.net.URI; +import java.net.http.HttpClient; +import java.net.http.HttpRequest; +import java.net.http.HttpResponse; + +import org.junit.jupiter.api.Test; +import org.springframework.boot.SpringBootConfiguration; +import org.springframework.boot.autoconfigure.EnableAutoConfiguration; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.boot.test.web.server.LocalServerPort; +import org.springframework.context.annotation.Bean; +import org.springframework.context.annotation.Import; +import org.springframework.test.context.DynamicPropertyRegistry; +import org.springframework.test.context.DynamicPropertySource; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.proprietary.mcp.McpServerController; +import stirling.software.proprietary.mcp.tools.DescribeOperationTool; +import stirling.software.proprietary.security.service.UserService; + +/** + * Regression test for the protected-resource metadata when {@code mcp.scopes-enabled=false} (the + * SaaS/Supabase setup, where the IdP cannot mint {@code mcp.tools.*} scopes). Advertising scopes + * the authorization server can't issue makes spec-compliant MCP clients request them and get + * bounced with {@code invalid_request}, so the metadata must omit them when scopes are not + * enforced. + */ +@SpringBootTest( + classes = McpScopeMetadataDisabledTest.TestApp.class, + webEnvironment = SpringBootTest.WebEnvironment.RANDOM_PORT) +class McpScopeMetadataDisabledTest { + + private static final String ISSUER = "https://test-issuer.example.com"; + private static final String RESOURCE_ID = "http://localhost/mcp"; + + @LocalServerPort private int port; + + private final HttpClient http = HttpClient.newHttpClient(); + + // McpSecurityConfig is @ConditionalOnProperty("mcp.enabled"); that condition reads the Spring + // Environment, so it must be set here (the ApplicationProperties bean alone is not enough to + // register the chain). + @DynamicPropertySource + static void mcpProperties(DynamicPropertyRegistry registry) { + registry.add("mcp.enabled", () -> "true"); + } + + @Test + void metadata_omitsToolScopes_whenScopesDisabled() throws Exception { + String body = getMetadata("/.well-known/oauth-protected-resource"); + assertThat(body).contains(RESOURCE_ID); + assertThat(body).contains(ISSUER); + assertThat(body).doesNotContain("mcp.tools.read"); + assertThat(body).doesNotContain("mcp.tools.write"); + } + + @Test + void pathInsertedMetadata_omitsToolScopes_whenScopesDisabled() throws Exception { + String body = getMetadata("/.well-known/oauth-protected-resource/mcp"); + assertThat(body).contains(RESOURCE_ID); + assertThat(body).contains("authorization_servers"); + assertThat(body).doesNotContain("mcp.tools.read"); + assertThat(body).doesNotContain("mcp.tools.write"); + } + + private String getMetadata(String path) throws Exception { + HttpRequest request = + HttpRequest.newBuilder() + .uri(URI.create("http://localhost:" + port + path)) + .GET() + .build(); + HttpResponse response = http.send(request, HttpResponse.BodyHandlers.ofString()); + assertThat(response.statusCode()).isEqualTo(200); + return response.body(); + } + + @SpringBootConfiguration + @EnableAutoConfiguration + @Import({McpSecurityConfig.class, McpServerController.class, DescribeOperationTool.class}) + static class TestApp { + + @Bean + ApplicationProperties applicationProperties() { + ApplicationProperties props = new ApplicationProperties(); + props.getMcp().setEnabled(true); + props.getMcp().setScopesEnabled(false); + props.getMcp().getAuth().setIssuerUri(ISSUER); + // No real JWKS fetch happens for the permitAll metadata endpoint; a placeholder URI is + // fine because NimbusJwtDecoder.withJwkSetUri(...) resolves the key set lazily. + props.getMcp().getAuth().setJwksUri(ISSUER + "/jwks"); + props.getMcp().getAuth().setResourceId(RESOURCE_ID); + props.getAutomaticallyGenerated().setAppVersion("test"); + return props; + } + + @Bean + UserService userService() { + return org.mockito.Mockito.mock(UserService.class); + } + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/mcp/tools/CategoryToolDispatchTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/tools/CategoryToolDispatchTest.java new file mode 100644 index 0000000000..de843ae423 --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/tools/CategoryToolDispatchTest.java @@ -0,0 +1,121 @@ +package stirling.software.proprietary.mcp.tools; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; + +import java.util.List; +import java.util.Optional; +import java.util.Set; + +import org.junit.jupiter.api.Test; +import org.springframework.beans.factory.ObjectProvider; + +import stirling.software.proprietary.mcp.McpCallContext; +import stirling.software.proprietary.mcp.catalog.McpToolCatalog; +import stirling.software.proprietary.mcp.catalog.OperationCategory; +import stirling.software.proprietary.mcp.catalog.OperationMeta; + +import tools.jackson.databind.ObjectMapper; +import tools.jackson.databind.node.ObjectNode; + +/** + * PDF category tools must not fake success: a bad/missing op returns the operation list, a valid + * scoped op delegates to the executor, and a missing scope is refused. + */ +class CategoryToolDispatchTest { + + private final ObjectMapper mapper = new ObjectMapper(); + + private OperationMeta miscOp() { + return new OperationMeta( + "compress-pdf", + OperationCategory.MISC, + "Compress a PDF", + mapper.createObjectNode(), + "mcp.tools.write", + OperationMeta.Target.JAVA_ENDPOINT, + "/api/v1/misc/compress-pdf", + null); + } + + private McpOperationExecutor executorReturning(ObjectNode sentinel) { + McpOperationExecutor executor = mock(McpOperationExecutor.class); + when(executor.execute(any(), any())).thenReturn(sentinel); + return executor; + } + + private StirlingMiscTool toolWith(McpOperationExecutor executor) { + OperationMeta meta = miscOp(); + McpToolCatalog catalog = mock(McpToolCatalog.class); + when(catalog.findByOperationId("compress-pdf")).thenReturn(Optional.of(meta)); + when(catalog.enabledOps(OperationCategory.MISC)).thenReturn(List.of(meta)); + @SuppressWarnings("unchecked") + ObjectProvider catalogProvider = mock(ObjectProvider.class); + when(catalogProvider.getIfAvailable()).thenReturn(catalog); + @SuppressWarnings("unchecked") + ObjectProvider executorProvider = mock(ObjectProvider.class); + when(executorProvider.getIfAvailable()).thenReturn(executor); + return new StirlingMiscTool(mapper, catalogProvider, executorProvider); + } + + private ObjectNode args(String op) { + ObjectNode a = mapper.createObjectNode(); + if (op != null) { + a.put("operation", op); + } + return a; + } + + private String textOf(ObjectNode result) { + return result.get("content").get(0).get("text").asText(); + } + + @Test + void validOpWithScope_delegatesToExecutor() { + ObjectNode sentinel = McpResponses.text(mapper, "EXECUTED"); + StirlingMiscTool tool = toolWith(executorReturning(sentinel)); + McpCallContext ctx = new McpCallContext("user", Set.of("mcp.tools.write"), true); + + ObjectNode result = tool.call(args("compress-pdf"), ctx); + + assertEquals("EXECUTED", textOf(result), "valid scoped op must run via the executor"); + } + + @Test + void unknownOperation_returnsAvailableOperationList() { + StirlingMiscTool tool = toolWith(executorReturning(mapper.createObjectNode())); + McpCallContext ctx = new McpCallContext("user", Set.of("mcp.tools.write"), true); + + ObjectNode result = tool.call(args("does-not-exist"), ctx); + + assertTrue(result.path("isError").asBoolean(false)); + String text = textOf(result); + assertTrue(text.contains("Available operations"), "should list available ops: " + text); + assertTrue(text.contains("compress-pdf"), "should include the valid op id: " + text); + } + + @Test + void missingOperation_returnsAvailableOperationList() { + StirlingMiscTool tool = toolWith(executorReturning(mapper.createObjectNode())); + McpCallContext ctx = new McpCallContext("user", Set.of("mcp.tools.write"), true); + + ObjectNode result = tool.call(args(null), ctx); + + assertTrue(result.path("isError").asBoolean(false)); + assertTrue(textOf(result).contains("Available operations")); + } + + @Test + void missingScope_returnsScopeError() { + StirlingMiscTool tool = toolWith(executorReturning(mapper.createObjectNode())); + McpCallContext ctx = new McpCallContext("user", Set.of("mcp.tools.read"), true); + + ObjectNode result = tool.call(args("compress-pdf"), ctx); + + assertTrue(result.path("isError").asBoolean(false)); + assertTrue(textOf(result).toLowerCase().contains("scope")); + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/mcp/tools/FileToolScopeTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/tools/FileToolScopeTest.java new file mode 100644 index 0000000000..966420548e --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/tools/FileToolScopeTest.java @@ -0,0 +1,55 @@ +package stirling.software.proprietary.mcp.tools; + +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.Mockito.mock; + +import java.util.Set; + +import org.junit.jupiter.api.Test; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.common.service.FileStorage; +import stirling.software.proprietary.mcp.McpCallContext; + +import tools.jackson.databind.ObjectMapper; +import tools.jackson.databind.node.ObjectNode; + +/** stirling_upload requires write scope; stirling_download requires read scope. */ +class FileToolScopeTest { + + private final ObjectMapper mapper = new ObjectMapper(); + + private McpCallContext noScopes() { + return new McpCallContext("user", Set.of(), true); + } + + private String text(ObjectNode result) { + return result.get("content").get(0).get("text").asText(); + } + + @Test + void upload_withoutWriteScope_isRefused() { + StirlingUploadTool tool = new StirlingUploadTool(mapper, mock(FileStorage.class)); + ObjectNode args = mapper.createObjectNode(); + args.put("file", "YWJj"); + + ObjectNode result = tool.call(args, noScopes()); + + assertTrue(result.path("isError").asBoolean(false)); + assertTrue(text(result).toLowerCase().contains("scope")); + } + + @Test + void download_withoutReadScope_isRefused() { + StirlingDownloadTool tool = + new StirlingDownloadTool( + mapper, mock(FileStorage.class), new ApplicationProperties()); + ObjectNode args = mapper.createObjectNode(); + args.put("fileId", "abc"); + + ObjectNode result = tool.call(args, noScopes()); + + assertTrue(result.path("isError").asBoolean(false)); + assertTrue(text(result).toLowerCase().contains("scope")); + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/mcp/tools/McpOperationExecutorTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/tools/McpOperationExecutorTest.java new file mode 100644 index 0000000000..f2d9387142 --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/mcp/tools/McpOperationExecutorTest.java @@ -0,0 +1,146 @@ +package stirling.software.proprietary.mcp.tools; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.nio.charset.StandardCharsets; +import java.util.Base64; + +import org.junit.jupiter.api.Test; +import org.mockito.ArgumentCaptor; +import org.springframework.core.io.ByteArrayResource; +import org.springframework.core.io.Resource; +import org.springframework.http.HttpHeaders; +import org.springframework.http.MediaType; +import org.springframework.http.ResponseEntity; +import org.springframework.util.MultiValueMap; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.common.service.FileStorage; +import stirling.software.common.service.InternalApiClient; +import stirling.software.proprietary.mcp.catalog.OperationCategory; +import stirling.software.proprietary.mcp.catalog.OperationMeta; + +import tools.jackson.databind.ObjectMapper; +import tools.jackson.databind.node.ObjectNode; + +class McpOperationExecutorTest { + + private final ObjectMapper mapper = new ObjectMapper(); + + private OperationMeta compressOp() { + return new OperationMeta( + "compress-pdf", + OperationCategory.MISC, + "Compress", + mapper.createObjectNode(), + "mcp.tools.write", + OperationMeta.Target.JAVA_ENDPOINT, + "/api/v1/misc/compress-pdf", + null); + } + + private ResponseEntity pdfResponse(byte[] bytes) { + HttpHeaders headers = new HttpHeaders(); + headers.setContentType(MediaType.APPLICATION_PDF); + Resource body = + new ByteArrayResource(bytes) { + @Override + public String getFilename() { + return "out.pdf"; + } + }; + return ResponseEntity.ok().headers(headers).body(body); + } + + @Test + @SuppressWarnings({"unchecked", "rawtypes"}) + void inlineBase64Input_runsAndReturnsResultInline() throws Exception { + InternalApiClient api = mock(InternalApiClient.class); + FileStorage storage = mock(FileStorage.class); + ApplicationProperties props = new ApplicationProperties(); + + when(api.post(eq("/api/v1/misc/compress-pdf"), any())) + .thenReturn(pdfResponse("RESULT".getBytes(StandardCharsets.UTF_8))); + when(storage.storeBytes(any(), anyString())).thenReturn("result-123"); + + McpOperationExecutor executor = new McpOperationExecutor(mapper, api, storage, props); + + ObjectNode args = mapper.createObjectNode(); + args.put("operation", "compress-pdf"); + args.put( + "file", + Base64.getEncoder().encodeToString("INPUT".getBytes(StandardCharsets.UTF_8))); + args.putObject("parameters").put("optimizeLevel", 2); + + ObjectNode result = executor.execute(compressOp(), args); + + assertFalse(result.path("isError").asBoolean(false)); + + ArgumentCaptor bodyCap = ArgumentCaptor.forClass(MultiValueMap.class); + verify(api).post(eq("/api/v1/misc/compress-pdf"), bodyCap.capture()); + MultiValueMap captured = bodyCap.getValue(); + assertTrue(captured.containsKey("fileInput"), "must send fileInput"); + assertEquals("2", String.valueOf(captured.getFirst("optimizeLevel")), "must pass params"); + + String text = result.get("content").get(0).get("text").asText(); + assertTrue(text.contains("result-123"), "must report the result fileId: " + text); + ObjectNode resBlock = (ObjectNode) result.get("content").get(1); + assertEquals("resource", resBlock.get("type").asText()); + String blob = resBlock.get("resource").get("blob").asText(); + assertEquals( + "RESULT", new String(Base64.getDecoder().decode(blob), StandardCharsets.UTF_8)); + } + + @Test + void missingFile_returnsError() { + McpOperationExecutor executor = + new McpOperationExecutor( + mapper, + mock(InternalApiClient.class), + mock(FileStorage.class), + new ApplicationProperties()); + ObjectNode args = mapper.createObjectNode(); + args.put("operation", "compress-pdf"); + + ObjectNode result = executor.execute(compressOp(), args); + + assertTrue(result.path("isError").asBoolean(false)); + assertTrue( + result.get("content") + .get(0) + .get("text") + .asText() + .toLowerCase() + .contains("input file")); + } + + @Test + void fileIdInput_retrievesStoredBytes() throws Exception { + InternalApiClient api = mock(InternalApiClient.class); + FileStorage storage = mock(FileStorage.class); + when(storage.fileExists("abc")).thenReturn(true); + when(storage.retrieveBytes("abc")).thenReturn("INPUT".getBytes(StandardCharsets.UTF_8)); + when(api.post(anyString(), any())) + .thenReturn(pdfResponse("OUT".getBytes(StandardCharsets.UTF_8))); + when(storage.storeBytes(any(), anyString())).thenReturn("res"); + + McpOperationExecutor executor = + new McpOperationExecutor(mapper, api, storage, new ApplicationProperties()); + ObjectNode args = mapper.createObjectNode(); + args.put("operation", "compress-pdf"); + args.put("fileId", "abc"); + + ObjectNode result = executor.execute(compressOp(), args); + + assertFalse(result.path("isError").asBoolean(false)); + verify(storage).retrieveBytes("abc"); + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/policy/config/AdminPolicyManagementAuthorityTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/policy/config/AdminPolicyManagementAuthorityTest.java new file mode 100644 index 0000000000..811aae6b4f --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/policy/config/AdminPolicyManagementAuthorityTest.java @@ -0,0 +1,58 @@ +package stirling.software.proprietary.policy.config; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.Mockito.when; + +import java.util.Optional; + +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; + +import stirling.software.proprietary.model.Team; +import stirling.software.proprietary.security.model.User; +import stirling.software.proprietary.security.service.UserService; + +/** Self-hosted policy context: a global admin may edit; scoping uses the current user's team. */ +@ExtendWith(MockitoExtension.class) +class AdminPolicyManagementAuthorityTest { + + @Mock private UserService userService; + + private AdminPolicyManagementAuthority authority() { + return new AdminPolicyManagementAuthority(userService); + } + + @Test + void adminMayEditPolicies() { + when(userService.isCurrentUserAdmin()).thenReturn(true); + assertTrue(authority().canEditPolicies()); + } + + @Test + void nonAdminMayNot() { + when(userService.isCurrentUserAdmin()).thenReturn(false); + assertFalse(authority().canEditPolicies()); + } + + @Test + void currentUserTeamIdResolvesFromTheCurrentUsersTeam() { + Team team = new Team(); + team.setId(42L); + User user = new User(); + user.setTeam(team); + when(userService.getCurrentUsername()).thenReturn("alice"); + when(userService.findByUsername("alice")).thenReturn(Optional.of(user)); + assertEquals(42L, authority().currentUserTeamId()); + } + + @Test + void currentUserTeamIdIsNullWhenNoCurrentUser() { + when(userService.getCurrentUsername()).thenReturn(null); + assertNull(authority().currentUserTeamId()); + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/policy/config/FolderAccessGuardTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/policy/config/FolderAccessGuardTest.java new file mode 100644 index 0000000000..ea33319a8b --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/policy/config/FolderAccessGuardTest.java @@ -0,0 +1,99 @@ +package stirling.software.proprietary.policy.config; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.nio.file.Path; +import java.util.List; + +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.io.TempDir; +import org.springframework.core.env.StandardEnvironment; + +import stirling.software.common.configuration.InstallationPathConfig; +import stirling.software.common.model.ApplicationProperties; +import stirling.software.proprietary.policy.model.InputSpec; +import stirling.software.proprietary.policy.model.OutputSpec; +import stirling.software.proprietary.policy.model.Policy; + +/** + * Tests for {@link FolderAccessGuard}: folder access is fail-closed, confined to the configured + * allowed roots, never reaches Stirling's own config directory, and is off entirely under SaaS. + */ +class FolderAccessGuardTest { + + @TempDir Path tempDir; + + private FolderAccessGuard guard(List allowedRoots, String... activeProfiles) { + ApplicationProperties properties = new ApplicationProperties(); + properties.getPolicies().setAllowedFolderRoots(allowedRoots); + StandardEnvironment environment = new StandardEnvironment(); + environment.setActiveProfiles(activeProfiles); + return new FolderAccessGuard(properties, environment); + } + + @Test + void permitsAndNormalisesADirectoryWithinAnAllowedRoot() { + FolderAccessGuard guard = guard(List.of(tempDir.toString())); + Path within = tempDir.resolve("inbox"); + + assertEquals(within.toAbsolutePath().normalize(), guard.requirePermitted(within)); + } + + @Test + void rejectsADirectoryOutsideEveryAllowedRoot() { + FolderAccessGuard guard = guard(List.of(tempDir.toString())); + assertThrows( + IllegalArgumentException.class, + () -> guard.requirePermitted(tempDir.resolveSibling("elsewhere"))); + } + + @Test + void rejectsTraversalThatWalksOutOfAnAllowedRoot() { + FolderAccessGuard guard = guard(List.of(tempDir.toString())); + assertThrows( + IllegalArgumentException.class, + () -> guard.requirePermitted(tempDir.resolve("..").resolve("escaped"))); + } + + @Test + void rejectsEverythingWhenNoRootsAreConfigured() { + FolderAccessGuard guard = guard(List.of()); + assertThrows(IllegalArgumentException.class, () -> guard.requirePermitted(tempDir)); + } + + @Test + void rejectsTheStirlingConfigDirectoryEvenWhenItWouldBeInsideAnAllowedRoot() { + Path configDir = + Path.of(InstallationPathConfig.getConfigPath()).toAbsolutePath().normalize(); + // Allow the config dir's parent, so only the protected-path rule can reject it. + FolderAccessGuard guard = guard(List.of(configDir.getParent().toString())); + + assertThrows( + IllegalArgumentException.class, + () -> guard.requirePermitted(configDir.resolve("settings.yml"))); + } + + @Test + void refusesAllFolderAccessUnderTheSaasProfile() { + FolderAccessGuard guard = guard(List.of(tempDir.toString()), "saas"); + assertThrows(IllegalArgumentException.class, () -> guard.requirePermitted(tempDir)); + } + + @Test + void usesFolderAccessDetectsFolderSourcesAndOutputs() { + FolderAccessGuard guard = guard(List.of(tempDir.toString())); + + assertTrue( + guard.usesFolderAccess( + policy(List.of(InputSpec.folder("/in")), OutputSpec.inline()))); + assertTrue(guard.usesFolderAccess(policy(List.of(), OutputSpec.folder("/out")))); + assertFalse(guard.usesFolderAccess(policy(List.of(), OutputSpec.inline()))); + } + + private static Policy policy(List sources, OutputSpec output) { + return new Policy("p1", "p", "owner", true, null, sources, List.of(), output); + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/policy/config/PolicyAccessGuardTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/policy/config/PolicyAccessGuardTest.java new file mode 100644 index 0000000000..9f86b87aec --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/policy/config/PolicyAccessGuardTest.java @@ -0,0 +1,84 @@ +package stirling.software.proprietary.policy.config; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.Mockito.when; + +import java.util.List; + +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.common.service.UserServiceInterface; +import stirling.software.proprietary.policy.model.OutputSpec; +import stirling.software.proprietary.policy.model.Policy; + +/** + * {@link PolicyAccessGuard}: policies are scoped to the caller's team. A user sees/accesses only + * their own team's policies (admins included — there is no cross-team escape). Login disabled + * (single-user) bypasses scoping. + */ +@ExtendWith(MockitoExtension.class) +class PolicyAccessGuardTest { + + @Mock private UserServiceInterface userService; + @Mock private PolicyManagementAuthority policyManagementAuthority; + + private PolicyAccessGuard guard(boolean loginEnabled) { + ApplicationProperties properties = new ApplicationProperties(); + properties.getSecurity().setEnableLogin(loginEnabled); + return new PolicyAccessGuard(userService, properties, policyManagementAuthority); + } + + @Test + void visibleFiltersToTheCallersTeam() { + when(policyManagementAuthority.currentUserTeamId()).thenReturn(1L); + List all = List.of(inTeam(1L), inTeam(2L), inTeam(1L), inTeam(null)); + List visible = guard(true).visible(all); + assertEquals(2, visible.size()); + assertTrue(visible.stream().allMatch(p -> Long.valueOf(1L).equals(p.teamId()))); + } + + @Test + void visibleReturnsEverythingWhenLoginDisabled() { + List all = List.of(inTeam(1L), inTeam(2L)); + assertEquals(all, guard(false).visible(all)); + } + + @Test + void canAccessOnlyOwnTeamsPolicy() { + when(policyManagementAuthority.currentUserTeamId()).thenReturn(1L); + assertTrue(guard(true).canAccess(inTeam(1L))); + assertFalse(guard(true).canAccess(inTeam(2L))); + assertFalse(guard(true).canAccess(inTeam(null))); + } + + @Test + void canAccessAnythingWhenLoginDisabled() { + assertTrue(guard(false).canAccess(inTeam(2L))); + } + + @Test + void ownerAndTeamForNewPolicyComeFromTheCurrentUserWhenLoginEnabled() { + when(userService.getCurrentUsername()).thenReturn("alice"); + when(policyManagementAuthority.currentUserTeamId()).thenReturn(7L); + assertEquals("alice", guard(true).ownerForNewPolicy()); + assertEquals(7L, guard(true).teamForNewPolicy()); + } + + @Test + void ownerAndTeamForNewPolicyAreNullWhenLoginDisabled() { + assertNull(guard(false).ownerForNewPolicy()); + assertNull(guard(false).teamForNewPolicy()); + } + + private static Policy inTeam(Long teamId) { + return new Policy( + "p1", "p", "owner", true, null, List.of(), List.of(), OutputSpec.inline(), teamId); + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/policy/engine/PolicyEngineTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/policy/engine/PolicyEngineTest.java new file mode 100644 index 0000000000..78348fa5c4 --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/policy/engine/PolicyEngineTest.java @@ -0,0 +1,388 @@ +package stirling.software.proprietary.policy.engine; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyInt; +import static org.mockito.ArgumentMatchers.anyLong; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.atLeastOnce; +import static org.mockito.Mockito.doReturn; +import static org.mockito.Mockito.lenient; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.io.InputStream; +import java.nio.file.Path; +import java.util.List; +import java.util.Map; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.TimeUnit; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.junit.jupiter.api.io.TempDir; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.slf4j.MDC; +import org.springframework.core.io.ByteArrayResource; +import org.springframework.core.io.Resource; +import org.springframework.http.HttpHeaders; +import org.springframework.http.HttpStatus; +import org.springframework.http.ResponseEntity; +import org.springframework.web.client.HttpClientErrorException; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.common.service.FileStorage; +import stirling.software.common.service.FileStorage.StoredFile; +import stirling.software.common.service.InternalApiClient; +import stirling.software.common.service.JobOwnershipService; +import stirling.software.common.service.JobQueue; +import stirling.software.common.service.ResourceMonitor; +import stirling.software.common.service.TaskManager; +import stirling.software.common.service.ToolMetadataService; +import stirling.software.common.util.TempFileManager; +import stirling.software.common.util.TempFileRegistry; +import stirling.software.proprietary.policy.model.OutputSpec; +import stirling.software.proprietary.policy.model.PipelineDefinition; +import stirling.software.proprietary.policy.model.PipelineStep; +import stirling.software.proprietary.policy.model.Policy; +import stirling.software.proprietary.policy.model.PolicyInputs; +import stirling.software.proprietary.policy.model.PolicyRun; +import stirling.software.proprietary.policy.model.PolicyRunStatus; +import stirling.software.proprietary.policy.output.InlineOutputSink; +import stirling.software.proprietary.policy.progress.PolicyProgressListener; + +import tools.jackson.databind.json.JsonMapper; + +/** + * Tests for {@link PolicyEngine}: async submission runs the pipeline on a virtual thread, registers + * outputs and progress with {@link TaskManager}, and surfaces terminal state via {@link + * PolicyRunRegistry}. The step executor and inline sink are real (with mocked collaborators) so the + * full run path is exercised. + */ +@ExtendWith(MockitoExtension.class) +class PolicyEngineTest { + + private static final String ROTATE = "/api/v1/general/rotate-pdf"; + private static final String COMPRESS = "/api/v1/misc/compress-pdf"; + + @Mock private InternalApiClient internalApiClient; + @Mock private ToolMetadataService toolMetadataService; + @Mock private TaskManager taskManager; + @Mock private FileStorage fileStorage; + @Mock private JobOwnershipService jobOwnershipService; + @Mock private ResourceMonitor resourceMonitor; + @Mock private JobQueue jobQueue; + + @TempDir Path tempDir; + + private PolicyRunRegistry registry; + private PolicyEngine engine; + + @BeforeEach + void setUp() { + ApplicationProperties props = new ApplicationProperties(); + props.getSystem().getTempFileManagement().setBaseTmpDir(tempDir.toString()); + props.getSystem().getTempFileManagement().setPrefix("policy-engine-test-"); + TempFileManager tempFileManager = new TempFileManager(new TempFileRegistry(), props); + PolicyExecutor executor = + new PolicyExecutor( + internalApiClient, + toolMetadataService, + tempFileManager, + JsonMapper.builder().build()); + registry = new PolicyRunRegistry(new ApplicationProperties()); + InlineOutputSink sink = new InlineOutputSink(fileStorage); + engine = + new PolicyEngine( + executor, + taskManager, + registry, + fileStorage, + jobOwnershipService, + List.of(sink), + resourceMonitor, + jobQueue); + + // Identity scoping: the run id is the generated UUID unchanged. Lenient because the + // resume/cancel tests do not submit a run. + lenient() + .when(jobOwnershipService.createScopedJobKey(anyString())) + .thenAnswer(inv -> inv.getArgument(0)); + // Default to running immediately; the queueing test overrides this. + lenient().when(resourceMonitor.shouldQueueJob(anyInt())).thenReturn(false); + } + + @Test + void submitRunsPipelineToCompletionAndRegistersOutputs() throws Exception { + when(toolMetadataService.isMultiInput(anyString())).thenReturn(false); + when(toolMetadataService.shouldUnpackZipResponse(anyString())).thenReturn(false); + stubEndpoint(ROTATE, pdf("rotated", "rotated.pdf")); + stubEndpoint(COMPRESS, pdf("compressed", "compressed.pdf")); + int[] counter = {0}; + when(fileStorage.storeInputStream(any(InputStream.class), anyString())) + .thenAnswer( + inv -> { + InputStream is = inv.getArgument(0); + long size = is.readAllBytes().length; + return new StoredFile("file-" + ++counter[0], size); + }); + + PolicyRunHandle handle = + engine.submit( + definition( + new PipelineStep(ROTATE, Map.of()), + new PipelineStep(COMPRESS, Map.of())), + PolicyInputs.of(List.of(pdf("input", "input.pdf"))), + PolicyProgressListener.NOOP); + + // The completion future resolves with the final run state, no polling needed. + String runId = handle.runId(); + PolicyRun run = handle.completion().get(10, TimeUnit.SECONDS); + assertEquals(PolicyRunStatus.COMPLETED, run.getStatus()); + assertEquals(1, run.getOutputs().size()); + assertEquals("compressed.pdf", run.getOutputs().get(0).getFileName()); + + // The run self-registers its results and completion with the job system. + verify(taskManager).createTask(runId); + verify(taskManager).setMultipleFileResults(eq(runId), any()); + verify(taskManager).setComplete(runId); + // Progress notes were written for each step. + verify(taskManager, atLeastOnce()).addNote(eq(runId), anyString()); + } + + @Test + void submitFailsRunWhenAToolErrors() throws Exception { + when(toolMetadataService.isMultiInput(ROTATE)).thenReturn(false); + when(internalApiClient.post(eq(ROTATE), any())).thenThrow(new RuntimeException("boom")); + + PolicyRunHandle handle = + engine.submit( + definition(new PipelineStep(ROTATE, Map.of())), + PolicyInputs.of(List.of(pdf("input", "input.pdf"))), + PolicyProgressListener.NOOP); + + String runId = handle.runId(); + PolicyRun run = handle.completion().get(10, TimeUnit.SECONDS); + assertEquals(PolicyRunStatus.FAILED, run.getStatus()); + verify(taskManager).setError(eq(runId), anyString()); + verify(taskManager, never()).setComplete(runId); + } + + @Test + void runBlockedByUsageLimit_surfacesErrorCodeAndSubscribed() throws Exception { + // A downstream tool call gets a 402 entitlement block. The run fails, but its errorCode + + // subscribed are taken from the 402 body so the client can pop the right usage-limit modal + // (the policy 402 happens server-side, out of reach of the apiClient interceptor). + when(toolMetadataService.isMultiInput(ROTATE)).thenReturn(false); + String body = "{\"error\":\"PAYG_LIMIT_REACHED\",\"subscribed\":true}"; + when(internalApiClient.post(eq(ROTATE), any())) + .thenThrow( + HttpClientErrorException.create( + HttpStatus.PAYMENT_REQUIRED, + "Payment Required", + HttpHeaders.EMPTY, + body.getBytes(java.nio.charset.StandardCharsets.UTF_8), + java.nio.charset.StandardCharsets.UTF_8)); + + PolicyRun run = + engine.submit( + definition(new PipelineStep(ROTATE, Map.of())), + PolicyInputs.of(List.of(pdf("input", "input.pdf"))), + PolicyProgressListener.NOOP) + .completion() + .get(10, TimeUnit.SECONDS); + + assertEquals(PolicyRunStatus.FAILED, run.getStatus()); + assertEquals("PAYG_LIMIT_REACHED", run.getErrorCode()); + assertEquals(Boolean.TRUE, run.getErrorSubscribed()); + } + + @Test + void runPolicyExecutesThePolicysPipeline() throws Exception { + when(toolMetadataService.isMultiInput(anyString())).thenReturn(false); + when(toolMetadataService.shouldUnpackZipResponse(anyString())).thenReturn(false); + stubEndpoint(ROTATE, pdf("rotated", "rotated.pdf")); + int[] counter = {0}; + when(fileStorage.storeInputStream(any(InputStream.class), anyString())) + .thenAnswer( + inv -> { + InputStream is = inv.getArgument(0); + return new StoredFile("file-" + ++counter[0], is.readAllBytes().length); + }); + + Policy policy = + new Policy( + "p1", + "rotate", + "owner", + true, + null, + List.of(new PipelineStep(ROTATE, Map.of())), + OutputSpec.inline()); + + PolicyRunHandle handle = + engine.runPolicy( + policy, + PolicyInputs.of(List.of(pdf("input", "input.pdf"))), + PolicyProgressListener.NOOP); + + PolicyRun run = handle.completion().get(10, TimeUnit.SECONDS); + assertEquals(PolicyRunStatus.COMPLETED, run.getStatus()); + verify(internalApiClient).post(eq(ROTATE), any()); + } + + @Test + void runPolicyDispatchesToolCallsAsTheOwner() throws Exception { + // Billing-attribution regression: the pipeline runs on a background worker thread, but the + // policy owner must be propagated as the audit principal so InternalApiClient (and thus + // PAYG) attributes each tool call to the owner — not the INTERNAL_API_USER fallback. + when(toolMetadataService.isMultiInput(anyString())).thenReturn(false); + when(toolMetadataService.shouldUnpackZipResponse(anyString())).thenReturn(false); + int[] counter = {0}; + when(fileStorage.storeInputStream(any(InputStream.class), anyString())) + .thenAnswer( + inv -> + new StoredFile( + "file-" + ++counter[0], + ((InputStream) inv.getArgument(0)).readAllBytes().length)); + + String[] principalAtDispatch = {""}; + when(internalApiClient.post(eq(ROTATE), any())) + .thenAnswer( + inv -> { + principalAtDispatch[0] = MDC.get("auditPrincipal"); + return ResponseEntity.ok(pdf("rotated", "rotated.pdf")); + }); + + Policy policy = + new Policy( + "p1", + "rotate", + "alice", + true, + null, + List.of(new PipelineStep(ROTATE, Map.of())), + OutputSpec.inline()); + + engine.runPolicy( + policy, + PolicyInputs.of(List.of(pdf("input", "input.pdf"))), + PolicyProgressListener.NOOP) + .completion() + .get(10, TimeUnit.SECONDS); + + assertEquals("alice", principalAtDispatch[0]); + } + + @Test + void adHocRunDispatchesToolCallsAsTheSubmittingUser() throws Exception { + // Ad-hoc runs (no stored policy) bill whoever kicked them off; the principal is captured on + // the request thread (here simulated via MDC) and re-established on the worker thread. + when(toolMetadataService.isMultiInput(anyString())).thenReturn(false); + when(toolMetadataService.shouldUnpackZipResponse(anyString())).thenReturn(false); + int[] counter = {0}; + when(fileStorage.storeInputStream(any(InputStream.class), anyString())) + .thenAnswer( + inv -> + new StoredFile( + "file-" + ++counter[0], + ((InputStream) inv.getArgument(0)).readAllBytes().length)); + + String[] principalAtDispatch = {""}; + when(internalApiClient.post(eq(ROTATE), any())) + .thenAnswer( + inv -> { + principalAtDispatch[0] = MDC.get("auditPrincipal"); + return ResponseEntity.ok(pdf("rotated", "rotated.pdf")); + }); + + MDC.put("auditPrincipal", "bob"); // the request thread's audit principal + try { + engine.submit( + definition(new PipelineStep(ROTATE, Map.of())), + PolicyInputs.of(List.of(pdf("input", "input.pdf"))), + PolicyProgressListener.NOOP) + .completion() + .get(10, TimeUnit.SECONDS); + } finally { + MDC.remove("auditPrincipal"); + } + + assertEquals("bob", principalAtDispatch[0]); + } + + @Test + void runIsQueuedUnderResourcePressure() { + when(resourceMonitor.shouldQueueJob(anyInt())).thenReturn(true); + // Returning an already-completed future keeps the run parked: the queued work (which would + // start the run) is never executed by this mock, so it stays PENDING. + doReturn(CompletableFuture.completedFuture(null)) + .when(jobQueue) + .queueJob(anyString(), anyInt(), any(), anyLong()); + + PolicyRunHandle handle = + engine.submit( + definition(new PipelineStep(ROTATE, Map.of())), + PolicyInputs.of(List.of(pdf("input", "input.pdf"))), + PolicyProgressListener.NOOP); + + verify(jobQueue).queueJob(eq(handle.runId()), anyInt(), any(), anyLong()); + assertEquals(PolicyRunStatus.PENDING, registry.get(handle.runId()).getStatus()); + } + + @Test + void runRejectedWhenQueueFullCarriesTransientErrorCode() { + when(resourceMonitor.shouldQueueJob(anyInt())).thenReturn(true); + // Admission rejected (queue full): the queued future completes exceptionally. + CompletableFuture rejected = new CompletableFuture<>(); + rejected.completeExceptionally( + new RuntimeException("Job queue full, please try again later")); + doReturn(rejected).when(jobQueue).queueJob(anyString(), anyInt(), any(), anyLong()); + + PolicyRunHandle handle = + engine.submit( + definition(new PipelineStep(ROTATE, Map.of())), + PolicyInputs.of(List.of(pdf("input", "input.pdf"))), + PolicyProgressListener.NOOP); + + PolicyRun run = registry.get(handle.runId()); + assertEquals(PolicyRunStatus.FAILED, run.getStatus()); + // Tagged transient so the client backs off and retries instead of hard-failing. + assertEquals("POLICY_QUEUE_FULL", run.getErrorCode()); + } + + @Test + void resumeIsNotYetImplemented() { + assertThrows(UnsupportedOperationException.class, () -> engine.resume("any", List.of())); + } + + @Test + void cancelUnknownRunReturnsFalse() { + assertFalse(engine.cancel("does-not-exist")); + } + + // --- helpers --- + + private static PipelineDefinition definition(PipelineStep... steps) { + return new PipelineDefinition("test", List.of(steps), OutputSpec.inline()); + } + + private void stubEndpoint(String endpoint, Resource body) { + when(internalApiClient.post(eq(endpoint), any())).thenReturn(ResponseEntity.ok(body)); + } + + private static ByteArrayResource pdf(String content, String filename) { + return new ByteArrayResource(content.getBytes()) { + @Override + public String getFilename() { + return filename; + } + }; + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/policy/engine/PolicyExecutorTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/policy/engine/PolicyExecutorTest.java new file mode 100644 index 0000000000..d682c35cd7 --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/policy/engine/PolicyExecutorTest.java @@ -0,0 +1,378 @@ +package stirling.software.proprietary.policy.engine; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.times; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.io.ByteArrayInputStream; +import java.io.ByteArrayOutputStream; +import java.io.IOException; +import java.io.InputStream; +import java.nio.file.Path; +import java.util.ArrayList; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import java.util.zip.ZipEntry; +import java.util.zip.ZipOutputStream; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.junit.jupiter.api.io.TempDir; +import org.mockito.ArgumentCaptor; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.springframework.core.io.ByteArrayResource; +import org.springframework.core.io.Resource; +import org.springframework.http.ResponseEntity; +import org.springframework.util.MultiValueMap; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.common.service.InternalApiClient; +import stirling.software.common.service.InternalApiTimeoutException; +import stirling.software.common.service.ToolMetadataService; +import stirling.software.common.util.TempFileManager; +import stirling.software.common.util.TempFileRegistry; +import stirling.software.proprietary.policy.model.OutputSpec; +import stirling.software.proprietary.policy.model.PipelineDefinition; +import stirling.software.proprietary.policy.model.PipelineStep; +import stirling.software.proprietary.policy.model.PolicyInputs; +import stirling.software.proprietary.policy.progress.PolicyProgressListener; + +import tools.jackson.databind.ObjectMapper; +import tools.jackson.databind.json.JsonMapper; + +/** + * Unit tests for {@link PolicyExecutor}, the shared pipeline step loop. Covers file chaining across + * steps, multi-input vs per-file dispatch, ZIP unpacking, structured-list parameter encoding, + * progress callbacks, and timeout propagation. External collaborators are mocked; {@link + * TempFileManager} is real so ZIP extraction exercises real code. + */ +@ExtendWith(MockitoExtension.class) +class PolicyExecutorTest { + + private static final String ROTATE = "/api/v1/general/rotate-pdf"; + private static final String COMPRESS = "/api/v1/misc/compress-pdf"; + private static final String SPLIT = "/api/v1/general/split-pages"; + private static final String MERGE = "/api/v1/general/merge-pdfs"; + + @Mock private InternalApiClient internalApiClient; + @Mock private ToolMetadataService toolMetadataService; + + @TempDir Path tempDir; + + private TempFileManager tempFileManager; + private PolicyExecutor executor; + + @BeforeEach + void setUp() { + ApplicationProperties props = new ApplicationProperties(); + props.getSystem().getTempFileManagement().setBaseTmpDir(tempDir.toString()); + props.getSystem().getTempFileManagement().setPrefix("policy-test-"); + tempFileManager = new TempFileManager(new TempFileRegistry(), props); + ObjectMapper objectMapper = JsonMapper.builder().build(); + executor = + new PolicyExecutor( + internalApiClient, toolMetadataService, tempFileManager, objectMapper); + } + + @Test + void executesStepsSequentiallyChainingOutputToInput() throws IOException { + when(toolMetadataService.isMultiInput(anyString())).thenReturn(false); + when(toolMetadataService.shouldUnpackZipResponse(anyString())).thenReturn(false); + stubEndpoint(ROTATE, pdf("rotated", "rotated.pdf")); + stubEndpoint(COMPRESS, pdf("compressed", "compressed.pdf")); + + List steps = new ArrayList<>(); + PolicyProgressListener listener = + new PolicyProgressListener() { + @Override + public void onStepStart(int stepIndex, int stepCount, String operation) { + steps.add(stepIndex); + } + }; + + PolicyExecutionResult result = + executor.execute( + definition( + new PipelineStep(ROTATE, Map.of()), + new PipelineStep(COMPRESS, Map.of())), + PolicyInputs.of(List.of(pdf("input", "input.pdf"))), + listener); + + assertEquals(1, result.files().size()); + assertEquals("compressed.pdf", result.files().get(0).getFilename()); + verify(internalApiClient, times(1)).post(eq(ROTATE), any()); + verify(internalApiClient, times(1)).post(eq(COMPRESS), any()); + // Progress fired once per step, in order. + assertEquals(List.of(1, 2), steps); + } + + @Test + void multiInputEndpointIsCalledOnceWithAllFiles() throws IOException { + when(toolMetadataService.isMultiInput(MERGE)).thenReturn(true); + when(toolMetadataService.shouldUnpackZipResponse(MERGE)).thenReturn(false); + stubEndpoint(MERGE, pdf("merged", "merged.pdf")); + + PolicyExecutionResult result = + executor.execute( + definition(new PipelineStep(MERGE, Map.of())), + PolicyInputs.of(List.of(pdf("a", "a.pdf"), pdf("b", "b.pdf"))), + PolicyProgressListener.NOOP); + + assertEquals(1, result.files().size()); + verify(internalApiClient, times(1)).post(eq(MERGE), any()); + } + + @Test + void singleInputEndpointIsCalledOncePerFile() throws IOException { + when(toolMetadataService.isMultiInput(ROTATE)).thenReturn(false); + when(toolMetadataService.shouldUnpackZipResponse(ROTATE)).thenReturn(false); + stubEndpoint(ROTATE, pdf("rotated", "rotated.pdf")); + + PolicyExecutionResult result = + executor.execute( + definition(new PipelineStep(ROTATE, Map.of())), + PolicyInputs.of(List.of(pdf("a", "a.pdf"), pdf("b", "b.pdf"))), + PolicyProgressListener.NOOP); + + assertEquals(2, result.files().size()); + verify(internalApiClient, times(2)).post(eq(ROTATE), any()); + } + + @Test + void zipResponseIsUnpackedIntoIndividualFiles() throws IOException { + when(toolMetadataService.isMultiInput(SPLIT)).thenReturn(false); + when(toolMetadataService.shouldUnpackZipResponse(SPLIT)).thenReturn(true); + stubEndpoint( + SPLIT, + zip( + "doc.zip", + List.of(new Entry("page-1.pdf", "one"), new Entry("page-2.pdf", "two")))); + + PolicyExecutionResult result = + executor.execute( + definition(new PipelineStep(SPLIT, Map.of())), + PolicyInputs.of(List.of(pdf("doc", "doc.pdf"))), + PolicyProgressListener.NOOP); + + assertEquals(2, result.files().size()); + assertEquals("page-1.pdf", result.files().get(0).getFilename()); + assertEquals("page-2.pdf", result.files().get(1).getFilename()); + } + + @Test + void structuredListParameterIsJsonEncodedAsSingleField() throws IOException { + String editText = "/api/v1/general/edit-text"; + when(toolMetadataService.isMultiInput(editText)).thenReturn(false); + when(toolMetadataService.shouldUnpackZipResponse(editText)).thenReturn(false); + stubEndpoint(editText, pdf("edited", "edited.pdf")); + + // LinkedHashMap so the serialized key order is deterministic for the assertion below. + Map edit = new LinkedHashMap<>(); + edit.put("find", "foo"); + edit.put("replace", "bar"); + Map params = new LinkedHashMap<>(); + params.put("edits", List.of(edit)); + params.put("useRegex", false); + + executor.execute( + definition(new PipelineStep(editText, params)), + PolicyInputs.of(List.of(pdf("in", "in.pdf"))), + PolicyProgressListener.NOOP); + + @SuppressWarnings("unchecked") + ArgumentCaptor> bodyCaptor = + ArgumentCaptor.forClass(MultiValueMap.class); + verify(internalApiClient).post(eq(editText), bodyCaptor.capture()); + MultiValueMap body = bodyCaptor.getValue(); + + List edits = body.get("edits"); + assertNotNull(edits); + assertEquals(1, edits.size()); + assertEquals("[{\"find\":\"foo\",\"replace\":\"bar\"}]", edits.get(0)); + } + + @Test + void supportingFilesAreBoundToTheirNamedFields() throws IOException { + String addStamp = "/api/v1/misc/add-stamp-to-pdf"; + when(toolMetadataService.isMultiInput(addStamp)).thenReturn(false); + when(toolMetadataService.shouldUnpackZipResponse(addStamp)).thenReturn(false); + stubEndpoint(addStamp, pdf("stamped", "stamped.pdf")); + + PipelineStep step = + new PipelineStep(addStamp, Map.of("opacity", 0.5), Map.of("stampImage", "logo")); + PolicyInputs inputs = + new PolicyInputs( + List.of(pdf("doc", "doc.pdf")), + Map.of("logo", List.of(pdf("logo-bytes", "logo.png")))); + + executor.execute( + new PipelineDefinition("stamp", List.of(step), OutputSpec.inline()), + inputs, + PolicyProgressListener.NOOP); + + @SuppressWarnings("unchecked") + ArgumentCaptor> bodyCaptor = + ArgumentCaptor.forClass(MultiValueMap.class); + verify(internalApiClient).post(eq(addStamp), bodyCaptor.capture()); + MultiValueMap body = bodyCaptor.getValue(); + // The document goes to fileInput; the supporting image is bound to its named field and is + // not part of the document stream. + assertEquals(1, body.get("fileInput").size()); + assertNotNull(body.get("stampImage")); + assertEquals(1, body.get("stampImage").size()); + } + + @Test + void missingSupportingFileFailsTheStep() { + String addStamp = "/api/v1/misc/add-stamp-to-pdf"; + when(toolMetadataService.isMultiInput(addStamp)).thenReturn(false); + PipelineStep step = new PipelineStep(addStamp, Map.of(), Map.of("stampImage", "logo")); + + IOException ex = + assertThrows( + IOException.class, + () -> + executor.execute( + new PipelineDefinition( + "stamp", List.of(step), OutputSpec.inline()), + PolicyInputs.of(List.of(pdf("doc", "doc.pdf"))), + PolicyProgressListener.NOOP)); + assertTrue(ex.getMessage().contains("logo")); + } + + @Test + void documentOfAnUnacceptedTypeFailsTheStep() { + String compress = "/api/v1/misc/compress-pdf"; + when(toolMetadataService.getExtensionTypes(false, compress)).thenReturn(List.of("pdf")); + + IOException ex = + assertThrows( + IOException.class, + () -> + executor.execute( + definition(new PipelineStep(compress, Map.of())), + PolicyInputs.of(List.of(pdf("img", "image.png"))), + PolicyProgressListener.NOOP)); + assertTrue(ex.getMessage().contains("image.png")); + // Type check happens before any dispatch. + verify(internalApiClient, never()).post(anyString(), any()); + } + + @Test + void documentOfAnAcceptedTypeProceeds() throws IOException { + String compress = "/api/v1/misc/compress-pdf"; + when(toolMetadataService.getExtensionTypes(false, compress)).thenReturn(List.of("pdf")); + when(toolMetadataService.isMultiInput(compress)).thenReturn(false); + when(toolMetadataService.shouldUnpackZipResponse(compress)).thenReturn(false); + stubEndpoint(compress, pdf("compressed", "compressed.pdf")); + + PolicyExecutionResult result = + executor.execute( + definition(new PipelineStep(compress, Map.of())), + PolicyInputs.of(List.of(pdf("doc", "doc.pdf"))), + PolicyProgressListener.NOOP); + + assertEquals(1, result.files().size()); + verify(internalApiClient, times(1)).post(eq(compress), any()); + } + + @Test + void filterOperationWithEmptyResultDropsTheFile() throws IOException { + String filter = "/api/v1/filter/filter-page-count"; + when(toolMetadataService.isMultiInput(filter)).thenReturn(false); + stubEndpoint(filter, pdf("", "filtered.pdf")); // empty body => filtered out + + PolicyExecutionResult result = + executor.execute( + definition(new PipelineStep(filter, Map.of())), + PolicyInputs.of(List.of(pdf("doc", "doc.pdf"))), + PolicyProgressListener.NOOP); + + assertEquals(0, result.files().size()); + } + + @Test + void timeoutFromAStepPropagates() { + when(toolMetadataService.isMultiInput(ROTATE)).thenReturn(false); + when(internalApiClient.post(eq(ROTATE), any())) + .thenThrow( + new InternalApiTimeoutException( + ROTATE, + java.time.Duration.ofSeconds(300), + new IOException("Read timed out"))); + + assertThrows( + InternalApiTimeoutException.class, + () -> + executor.execute( + definition(new PipelineStep(ROTATE, Map.of())), + PolicyInputs.of(List.of(pdf("in", "in.pdf"))), + PolicyProgressListener.NOOP)); + } + + @Test + void emptyPipelineIsRejected() { + assertThrows( + IllegalArgumentException.class, + () -> + executor.execute( + new PipelineDefinition("empty", List.of(), OutputSpec.inline()), + PolicyInputs.of(List.of(pdf("in", "in.pdf"))), + PolicyProgressListener.NOOP)); + } + + // --- helpers --- + + private static PipelineDefinition definition(PipelineStep... steps) { + return new PipelineDefinition("test", List.of(steps), OutputSpec.inline()); + } + + private void stubEndpoint(String endpoint, Resource body) { + when(internalApiClient.post(eq(endpoint), any())).thenReturn(ResponseEntity.ok(body)); + } + + private static ByteArrayResource pdf(String content, String filename) { + return new ByteArrayResource(content.getBytes()) { + @Override + public String getFilename() { + return filename; + } + }; + } + + private static ByteArrayResource zip(String filename, List entries) throws IOException { + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + try (ZipOutputStream zos = new ZipOutputStream(baos)) { + for (Entry entry : entries) { + zos.putNextEntry(new ZipEntry(entry.name())); + zos.write(entry.content().getBytes()); + zos.closeEntry(); + } + } + byte[] zipBytes = baos.toByteArray(); + return new ByteArrayResource(zipBytes) { + @Override + public String getFilename() { + return filename; + } + + @Override + public InputStream getInputStream() { + return new ByteArrayInputStream(zipBytes); + } + }; + } + + private record Entry(String name, String content) {} +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/policy/engine/PolicyRunRegistryTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/policy/engine/PolicyRunRegistryTest.java new file mode 100644 index 0000000000..1549ba62ac --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/policy/engine/PolicyRunRegistryTest.java @@ -0,0 +1,102 @@ +package stirling.software.proprietary.policy.engine; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.time.Instant; +import java.util.List; + +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.proprietary.policy.model.PipelineDefinition; +import stirling.software.proprietary.policy.model.PolicyRun; +import stirling.software.proprietary.policy.model.WaitState; + +/** + * Tests for {@link PolicyRunRegistry} eviction: terminal runs expire, active/paused runs persist. + */ +class PolicyRunRegistryTest { + + private PolicyRunRegistry registry; + + @BeforeEach + void setUp() { + registry = new PolicyRunRegistry(new ApplicationProperties()); + } + + @AfterEach + void tearDown() { + registry.shutdown(); + } + + @Test + void evictsTerminalRunsPastTheCutoff() { + PolicyRun completed = register("completed"); + completed.complete(List.of()); + PolicyRun failed = register("failed"); + failed.fail("boom"); + PolicyRun cancelled = register("cancelled"); + cancelled.cancel(); + + // A cutoff in the future means every terminal run finished "before" it. + int removed = registry.evictExpired(Instant.now().plusSeconds(60)); + + assertEquals(3, removed); + assertNull(registry.get("completed")); + assertNull(registry.get("failed")); + assertNull(registry.get("cancelled")); + } + + @Test + void retainsActiveAndPausedRunsRegardlessOfAge() { + register("pending"); // PENDING: never started + PolicyRun running = register("running"); + running.markRunning(); + PolicyRun waiting = register("waiting"); + waiting.waitForInput(new WaitState("needs a signature", 1, List.of())); + + int removed = registry.evictExpired(Instant.now().plusSeconds(60)); + + assertEquals(0, removed); + assertNotNull(registry.get("pending")); + assertNotNull(registry.get("running")); + assertNotNull(registry.get("waiting")); + } + + @Test + void keepsTerminalRunsStillWithinTheExpiryWindow() { + PolicyRun completed = register("recent"); + completed.complete(List.of()); + + // A cutoff in the past means the run was updated "after" it: too young to evict. + int removed = registry.evictExpired(Instant.now().minusSeconds(60)); + + assertEquals(0, removed); + assertNotNull(registry.get("recent")); + } + + @Test + void evictionLeavesUnrelatedRunsInPlace() { + PolicyRun completed = register("done"); + completed.complete(List.of()); + PolicyRun running = register("busy"); + running.markRunning(); + + registry.evictExpired(Instant.now().plusSeconds(60)); + + assertNull(registry.get("done")); + assertNotNull(registry.get("busy")); + assertTrue(registry.all().stream().anyMatch(r -> r.getRunId().equals("busy"))); + } + + private PolicyRun register(String runId) { + PolicyRun run = new PolicyRun(runId, null, new PipelineDefinition(runId, List.of(), null)); + registry.register(run); + return run; + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/policy/engine/PolicyRunnerTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/policy/engine/PolicyRunnerTest.java new file mode 100644 index 0000000000..2cefb1b27c --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/policy/engine/PolicyRunnerTest.java @@ -0,0 +1,159 @@ +package stirling.software.proprietary.policy.engine; + +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertSame; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.times; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.verifyNoInteractions; +import static org.mockito.Mockito.when; + +import java.util.List; +import java.util.Map; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.atomic.AtomicBoolean; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.ArgumentCaptor; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; + +import stirling.software.proprietary.policy.input.InputSource; +import stirling.software.proprietary.policy.input.ResolvedInput; +import stirling.software.proprietary.policy.model.InputSpec; +import stirling.software.proprietary.policy.model.OutputSpec; +import stirling.software.proprietary.policy.model.PipelineStep; +import stirling.software.proprietary.policy.model.Policy; +import stirling.software.proprietary.policy.model.PolicyInputs; +import stirling.software.proprietary.policy.model.PolicyRun; +import stirling.software.proprietary.policy.model.PolicyRunStatus; +import stirling.software.proprietary.policy.progress.PolicyProgressListener; + +/** + * Tests for {@link PolicyRunner}: the one place that turns a policy's sources into runs. Verifies + * it pulls every source, runs one job per unit of work, feeds each unit's completion hook the run + * outcome, and that a generator (no sources) still runs once. + */ +@ExtendWith(MockitoExtension.class) +class PolicyRunnerTest { + + @Mock private PolicyEngine policyEngine; + @Mock private InputSource folderSource; + + private PolicyRunner runner; + + @BeforeEach + void setUp() { + runner = new PolicyRunner(policyEngine, List.of(folderSource)); + } + + @Test + void runsOnceWithNoFilesWhenThePolicyHasNoSources() { + Policy policy = policy(List.of()); + when(policyEngine.runPolicy(eq(policy), any(), any())) + .thenReturn(new PolicyRunHandle("r", new CompletableFuture<>())); + + runner.run(policy); + + ArgumentCaptor inputs = ArgumentCaptor.forClass(PolicyInputs.class); + verify(policyEngine).runPolicy(eq(policy), inputs.capture(), any()); + assertTrue(inputs.getValue().primary().isEmpty()); + } + + @Test + void pullsEverySourceAndRunsOnePerUnitOfWork() throws Exception { + InputSpec spec = InputSpec.folder("/in"); + Policy policy = policy(List.of(spec)); + when(folderSource.supports(spec)).thenReturn(true); + when(folderSource.resolve(spec)) + .thenReturn( + List.of( + ResolvedInput.of(PolicyInputs.of(List.of())), + ResolvedInput.of(PolicyInputs.of(List.of())))); + when(policyEngine.runPolicy(any(), any(), any())) + .thenReturn(new PolicyRunHandle("r", new CompletableFuture<>())); + + runner.run(policy); + + verify(policyEngine, times(2)).runPolicy(eq(policy), any(), any()); + } + + @Test + void feedsEachUnitsCompletionHookTheRunOutcome() throws Exception { + InputSpec spec = InputSpec.folder("/in"); + Policy policy = policy(List.of(spec)); + AtomicBoolean outcome = new AtomicBoolean(false); + ResolvedInput unit = new ResolvedInput(PolicyInputs.of(List.of()), outcome::set); + when(folderSource.supports(spec)).thenReturn(true); + when(folderSource.resolve(spec)).thenReturn(List.of(unit)); + CompletableFuture completion = new CompletableFuture<>(); + when(policyEngine.runPolicy(any(), any(), any())) + .thenReturn(new PolicyRunHandle("r", completion)); + + runner.run(policy); + + PolicyRun run = mock(PolicyRun.class); + when(run.getStatus()).thenReturn(PolicyRunStatus.COMPLETED); + completion.complete(run); + + assertTrue(outcome.get()); + } + + @Test + void reportsFailureToTheCompletionHookWhenTheRunDoesNotComplete() throws Exception { + InputSpec spec = InputSpec.folder("/in"); + Policy policy = policy(List.of(spec)); + AtomicBoolean outcome = new AtomicBoolean(true); + ResolvedInput unit = new ResolvedInput(PolicyInputs.of(List.of()), outcome::set); + when(folderSource.supports(spec)).thenReturn(true); + when(folderSource.resolve(spec)).thenReturn(List.of(unit)); + CompletableFuture completion = new CompletableFuture<>(); + when(policyEngine.runPolicy(any(), any(), any())) + .thenReturn(new PolicyRunHandle("r", completion)); + + runner.run(policy); + completion.completeExceptionally(new RuntimeException("boom")); + + assertFalse(outcome.get()); + } + + @Test + void skipsSourcesWithNoMatchingBean() { + InputSpec spec = new InputSpec("s3", Map.of()); + Policy policy = policy(List.of(spec)); + when(folderSource.supports(spec)).thenReturn(false); + + runner.run(policy); + + verifyNoInteractions(policyEngine); + } + + @Test + void runWithSuppliedInputsBypassesSources() { + Policy policy = policy(List.of(InputSpec.folder("/in"))); + PolicyInputs inputs = PolicyInputs.of(List.of()); + PolicyRunHandle handle = new PolicyRunHandle("r", new CompletableFuture<>()); + when(policyEngine.runPolicy(policy, inputs, PolicyProgressListener.NOOP)) + .thenReturn(handle); + + assertSame(handle, runner.runWith(policy, inputs, PolicyProgressListener.NOOP)); + verifyNoInteractions(folderSource); + } + + private static Policy policy(List sources) { + return new Policy( + "p1", + "p", + "owner", + true, + null, + sources, + List.of(new PipelineStep("/api/v1/misc/compress-pdf", Map.of())), + OutputSpec.inline()); + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/policy/engine/PolicyValidatorTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/policy/engine/PolicyValidatorTest.java new file mode 100644 index 0000000000..edaa586511 --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/policy/engine/PolicyValidatorTest.java @@ -0,0 +1,114 @@ +package stirling.software.proprietary.policy.engine; + +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.doThrow; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.util.List; +import java.util.Map; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; + +import stirling.software.proprietary.policy.input.InputSource; +import stirling.software.proprietary.policy.model.InputSpec; +import stirling.software.proprietary.policy.model.OutputSpec; +import stirling.software.proprietary.policy.model.Policy; +import stirling.software.proprietary.policy.model.TriggerConfig; +import stirling.software.proprietary.policy.output.PolicyOutputSink; +import stirling.software.proprietary.policy.trigger.PolicyTrigger; + +/** Tests for {@link PolicyValidator}: routes each facet to its handler and surfaces failures. */ +@ExtendWith(MockitoExtension.class) +class PolicyValidatorTest { + + @Mock private PolicyTrigger trigger; + @Mock private InputSource inputSource; + @Mock private PolicyOutputSink outputSink; + + private PolicyValidator validator; + + @BeforeEach + void setUp() { + validator = + new PolicyValidator(List.of(trigger), List.of(inputSource), List.of(outputSink)); + } + + @Test + void delegatesEachFacetToItsHandler() { + when(trigger.type()).thenReturn("schedule"); + when(inputSource.supports(any())).thenReturn(true); + when(outputSink.supports(any())).thenReturn(true); + Policy policy = policy("schedule"); + + validator.validate(policy); + + verify(trigger).validate(policy); + verify(inputSource).validate(policy.sources().get(0)); + verify(outputSink).validate(policy.output()); + } + + @Test + void skipsTriggerValidationForAManualOnlyPolicy() { + when(inputSource.supports(any())).thenReturn(true); + when(outputSink.supports(any())).thenReturn(true); + + validator.validate(manualOnly()); + + verify(trigger, never()).validate(any()); + } + + @Test + void surfacesAnInvalidConfigFromAHandler() { + when(trigger.type()).thenReturn("schedule"); + doThrow(new IllegalArgumentException("invalid schedule")).when(trigger).validate(any()); + + IllegalArgumentException ex = + assertThrows( + IllegalArgumentException.class, + () -> validator.validate(policy("schedule"))); + assertTrue(ex.getMessage().contains("schedule")); + } + + @Test + void rejectsAnUnknownTriggerType() { + when(trigger.type()).thenReturn("schedule"); + + IllegalArgumentException ex = + assertThrows( + IllegalArgumentException.class, + () -> validator.validate(policy("mystery"))); + assertTrue(ex.getMessage().contains("unknown trigger type")); + } + + private static Policy policy(String triggerType) { + return new Policy( + "p1", + "p", + "owner", + true, + new TriggerConfig(triggerType, Map.of()), + List.of(InputSpec.folder("/in")), + List.of(), + OutputSpec.inline()); + } + + private static Policy manualOnly() { + return new Policy( + "p1", + "p", + "owner", + true, + null, + List.of(InputSpec.folder("/in")), + List.of(), + OutputSpec.inline()); + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/policy/input/FolderInputSourceTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/policy/input/FolderInputSourceTest.java new file mode 100644 index 0000000000..3e9d34ccf0 --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/policy/input/FolderInputSourceTest.java @@ -0,0 +1,137 @@ +package stirling.software.proprietary.policy.input; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.lenient; + +import java.io.IOException; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.List; +import java.util.Map; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.junit.jupiter.api.io.TempDir; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.springframework.core.env.StandardEnvironment; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.common.util.FileReadinessChecker; +import stirling.software.proprietary.policy.config.FolderAccessGuard; +import stirling.software.proprietary.policy.model.InputSpec; + +/** Tests for {@link FolderInputSource}: consume (claim + route) and snapshot (read-only) modes. */ +@ExtendWith(MockitoExtension.class) +class FolderInputSourceTest { + + @Mock private FileReadinessChecker readinessChecker; + + @TempDir Path tempDir; + + private FolderInputSource source; + + @BeforeEach + void setUp() { + ApplicationProperties properties = new ApplicationProperties(); + properties.getPolicies().setAllowedFolderRoots(List.of(tempDir.toString())); + FolderAccessGuard guard = new FolderAccessGuard(properties, new StandardEnvironment()); + source = new FolderInputSource(readinessChecker, guard); + // Lenient: the missing-dir / nonexistent-dir cases return before any readiness check. + lenient().when(readinessChecker.isReady(any())).thenReturn(true); + } + + @Test + void consumeClaimsFilesAndRoutesToDoneOnSuccess() throws IOException { + Path inputDir = Files.createDirectories(tempDir.resolve("in")); + Files.writeString(inputDir.resolve("doc.pdf"), "data"); + + List work = source.resolve(InputSpec.folder(inputDir.toString())); + + assertEquals(1, work.size()); + assertEquals(1, work.get(0).inputs().primary().size()); + // Claimed out of the input dir. + assertFalse(Files.exists(inputDir.resolve("doc.pdf"))); + assertTrue( + Files.exists( + inputDir.resolve(".stirling").resolve("processing").resolve("doc.pdf"))); + + work.get(0).onComplete().accept(true); + assertTrue(Files.exists(inputDir.resolve(".stirling").resolve("done").resolve("doc.pdf"))); + assertFalse( + Files.exists( + inputDir.resolve(".stirling").resolve("processing").resolve("doc.pdf"))); + } + + @Test + void consumeRoutesToErrorOnFailure() throws IOException { + Path inputDir = Files.createDirectories(tempDir.resolve("in")); + Files.writeString(inputDir.resolve("doc.pdf"), "data"); + + List work = source.resolve(InputSpec.folder(inputDir.toString())); + work.get(0).onComplete().accept(false); + + assertTrue(Files.exists(inputDir.resolve(".stirling").resolve("error").resolve("doc.pdf"))); + } + + @Test + void snapshotReadsWithoutClaiming() throws IOException { + Path inputDir = Files.createDirectories(tempDir.resolve("in")); + Files.writeString(inputDir.resolve("doc.pdf"), "data"); + + List work = + source.resolve( + new InputSpec( + "folder", + Map.of("directory", inputDir.toString(), "mode", "snapshot"))); + + assertEquals(1, work.size()); + // Not moved, and completing the run is a no-op. + assertTrue(Files.exists(inputDir.resolve("doc.pdf"))); + work.get(0).onComplete().accept(true); + assertTrue(Files.exists(inputDir.resolve("doc.pdf"))); + } + + @Test + void missingDirectoryOptionFails() { + assertThrows( + IllegalArgumentException.class, + () -> source.resolve(new InputSpec("folder", Map.of()))); + } + + @Test + void nonexistentDirectoryYieldsNoWork() throws IOException { + List work = + source.resolve(InputSpec.folder(tempDir.resolve("nope").toString())); + assertTrue(work.isEmpty()); + } + + @Test + void validateRejectsMissingDirectory() { + assertThrows( + IllegalArgumentException.class, + () -> source.validate(new InputSpec("folder", Map.of()))); + } + + @Test + void rejectsADirectoryOutsideTheAllowedRoots() { + Path outside = tempDir.resolveSibling("not-allowed"); + assertThrows( + IllegalArgumentException.class, + () -> source.resolve(InputSpec.folder(outside.toString()))); + assertThrows( + IllegalArgumentException.class, + () -> source.validate(InputSpec.folder(outside.toString()))); + } + + @Test + void watchTargetsIsTheConfiguredDirectory() { + Path inputDir = tempDir.resolve("in"); + assertEquals(List.of(inputDir), source.watchTargets(InputSpec.folder(inputDir.toString()))); + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/policy/output/FolderOutputSinkTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/policy/output/FolderOutputSinkTest.java new file mode 100644 index 0000000000..0d714162d7 --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/policy/output/FolderOutputSinkTest.java @@ -0,0 +1,105 @@ +package stirling.software.proprietary.policy.output; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.io.IOException; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.List; +import java.util.Map; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.io.TempDir; +import org.springframework.core.env.StandardEnvironment; +import org.springframework.core.io.ByteArrayResource; +import org.springframework.core.io.Resource; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.common.model.job.ResultFile; +import stirling.software.proprietary.policy.config.FolderAccessGuard; +import stirling.software.proprietary.policy.model.OutputSpec; + +/** Tests for {@link FolderOutputSink}: outputs are written to the configured directory on disk. */ +class FolderOutputSinkTest { + + @TempDir Path tempDir; + + private FolderOutputSink sink; + + @BeforeEach + void setUp() { + ApplicationProperties properties = new ApplicationProperties(); + properties.getPolicies().setAllowedFolderRoots(List.of(tempDir.toString())); + sink = new FolderOutputSink(new FolderAccessGuard(properties, new StandardEnvironment())); + } + + @Test + void writesEachOutputToTheDirectory() throws IOException { + Path out = tempDir.resolve("out"); + List outputs = List.of(named("a.pdf", "aaa"), named("b.pdf", "bb")); + + List results = + sink.deliver("run-1", outputs, OutputSpec.folder(out.toString())); + + assertEquals(2, results.size()); + assertTrue(Files.exists(out.resolve("a.pdf"))); + assertEquals("aaa", Files.readString(out.resolve("a.pdf"))); + assertEquals("bb", Files.readString(out.resolve("b.pdf"))); + } + + @Test + void collidingNamesGetAUniqueSuffix() throws IOException { + Path out = tempDir.resolve("out"); + List outputs = List.of(named("a.pdf", "first"), named("a.pdf", "second")); + + sink.deliver("run-1", outputs, OutputSpec.folder(out.toString())); + + assertTrue(Files.exists(out.resolve("a.pdf"))); + assertTrue(Files.exists(out.resolve("a (1).pdf"))); + } + + @Test + void missingDirectoryOptionIsRejected() { + OutputSpec noDir = new OutputSpec("folder", Map.of()); + assertThrows(IllegalArgumentException.class, () -> sink.validate(noDir)); + assertThrows( + IllegalArgumentException.class, + () -> sink.deliver("run-1", List.of(named("a.pdf", "x")), noDir)); + } + + @Test + void aDirectoryOutsideTheAllowedRootsIsRejected() { + OutputSpec outside = OutputSpec.folder(tempDir.resolveSibling("not-allowed").toString()); + assertThrows(IllegalArgumentException.class, () -> sink.validate(outside)); + assertThrows( + IllegalArgumentException.class, + () -> sink.deliver("run-1", List.of(named("a.pdf", "x")), outside)); + } + + @Test + void filenamesWithPathTraversalAreConfinedToTheDirectory() throws IOException { + Path out = tempDir.resolve("out"); + List outputs = + List.of(named("../escape.pdf", "x"), named("nested/deep.pdf", "y")); + + sink.deliver("run-1", outputs, OutputSpec.folder(out.toString())); + + // Each name is reduced to its bare form inside the target dir; nothing escapes. + assertTrue(Files.exists(out.resolve("escape.pdf"))); + assertTrue(Files.exists(out.resolve("deep.pdf"))); + assertFalse(Files.exists(tempDir.resolve("escape.pdf"))); + } + + private static ByteArrayResource named(String filename, String content) { + return new ByteArrayResource(content.getBytes()) { + @Override + public String getFilename() { + return filename; + } + }; + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/policy/store/InProcessPolicyStoreTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/policy/store/InProcessPolicyStoreTest.java new file mode 100644 index 0000000000..60f14be20d --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/policy/store/InProcessPolicyStoreTest.java @@ -0,0 +1,90 @@ +package stirling.software.proprietary.policy.store; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.util.List; +import java.util.Map; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; + +import stirling.software.proprietary.policy.model.OutputSpec; +import stirling.software.proprietary.policy.model.PipelineStep; +import stirling.software.proprietary.policy.model.Policy; +import stirling.software.proprietary.policy.model.TriggerConfig; + +/** Tests for {@link InProcessPolicyStore}: id assignment, upsert, trigger-type lookup, delete. */ +class InProcessPolicyStoreTest { + + private PolicyStore store; + + @BeforeEach + void setUp() { + store = new InProcessPolicyStore(); + } + + @Test + void savedPolicyGetsAnIdAndIsRetrievable() { + Policy saved = store.save(policy(null, "compress", null, true)); + + assertNotNull(saved.id()); + assertFalse(saved.id().isBlank()); + assertEquals(saved, store.get(saved.id()).orElseThrow()); + } + + @Test + void savingWithAnExistingIdUpdatesInPlace() { + Policy created = store.save(policy(null, "before", null, true)); + + store.save( + new Policy( + created.id(), + "after", + "owner", + true, + null, + List.of(), + OutputSpec.inline())); + + assertEquals(1, store.all().size()); + assertEquals("after", store.get(created.id()).orElseThrow().name()); + } + + @Test + void findByTriggerTypeReturnsOnlyEnabledMatches() { + store.save(policy(null, "nightly", "schedule", true)); + store.save(policy(null, "nightly-disabled", "schedule", false)); + store.save(policy(null, "hooked", "webhook", true)); + store.save(policy(null, "on-demand", null, true)); // manual-only: no trigger + + List scheduled = store.findByTriggerType("schedule"); + + assertEquals(1, scheduled.size()); + assertEquals("nightly", scheduled.get(0).name()); + } + + @Test + void deleteRemovesThePolicy() { + Policy saved = store.save(policy(null, "p", null, true)); + + assertTrue(store.delete(saved.id())); + assertTrue(store.get(saved.id()).isEmpty()); + assertFalse(store.delete(saved.id())); + } + + private static Policy policy(String id, String name, String triggerType, boolean enabled) { + TriggerConfig trigger = + triggerType == null ? null : new TriggerConfig(triggerType, Map.of()); + return new Policy( + id, + name, + "owner", + enabled, + trigger, + List.of(new PipelineStep("/api/v1/misc/compress-pdf", Map.of())), + OutputSpec.inline()); + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/policy/store/JpaPolicyStoreTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/policy/store/JpaPolicyStoreTest.java new file mode 100644 index 0000000000..2bdbc08a79 --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/policy/store/JpaPolicyStoreTest.java @@ -0,0 +1,131 @@ +package stirling.software.proprietary.policy.store; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.util.List; +import java.util.Map; +import java.util.Optional; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.ArgumentCaptor; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; + +import stirling.software.proprietary.policy.model.InputSpec; +import stirling.software.proprietary.policy.model.OutputSpec; +import stirling.software.proprietary.policy.model.PipelineStep; +import stirling.software.proprietary.policy.model.Policy; +import stirling.software.proprietary.policy.model.TriggerConfig; + +import tools.jackson.databind.ObjectMapper; +import tools.jackson.databind.json.JsonMapper; + +/** + * Tests for {@link JpaPolicyStore}'s entity mapping and query delegation. The repository is mocked; + * real Hibernate/H2 persistence is exercised at application boot (this module's convention is + * Mockito unit tests for store/service logic). + */ +@ExtendWith(MockitoExtension.class) +class JpaPolicyStoreTest { + + @Mock private PolicyRepository repository; + + private final ObjectMapper objectMapper = JsonMapper.builder().build(); + private JpaPolicyStore store; + + @BeforeEach + void setUp() { + store = new JpaPolicyStore(repository, objectMapper); + } + + @Test + void saveAssignsAnIdAndPersistsThePolicyAsJson() { + Policy saved = + store.save( + new Policy( + null, + "compress incoming", + "alice", + true, + new TriggerConfig("schedule", Map.of()), + List.of(InputSpec.folder("/in")), + List.of(new PipelineStep("/api/v1/misc/compress-pdf", Map.of())), + OutputSpec.inline())); + + assertNotNull(saved.id()); + ArgumentCaptor captor = ArgumentCaptor.forClass(PolicyEntity.class); + verify(repository).save(captor.capture()); + PolicyEntity entity = captor.getValue(); + assertEquals(saved.id(), entity.getId()); + assertEquals("schedule", entity.getTriggerType()); + assertTrue(entity.isEnabled()); + // The stored JSON round-trips back to an equal policy. + assertEquals(saved, objectMapper.readValue(entity.getPolicyJson(), Policy.class)); + } + + @Test + void getDeserializesThePolicyFromJson() { + Policy policy = + new Policy( + "p1", + "rotate", + "alice", + true, + null, // manual-only: no automatic trigger + List.of( + new PipelineStep( + "/api/v1/general/rotate-pdf", Map.of("angle", 90))), + OutputSpec.inline()); + when(repository.findById("p1")).thenReturn(Optional.of(entityFor(policy))); + + assertEquals(policy, store.get("p1").orElseThrow()); + } + + @Test + void findByTriggerTypeUsesTheEnabledQuery() { + Policy policy = + new Policy( + "p1", + "watch", + "alice", + true, + new TriggerConfig("schedule", Map.of()), + List.of(new PipelineStep("/api/v1/misc/compress-pdf", Map.of())), + OutputSpec.inline()); + when(repository.findByTriggerTypeAndEnabledTrue("schedule")) + .thenReturn(List.of(entityFor(policy))); + + List scheduled = store.findByTriggerType("schedule"); + + assertEquals(1, scheduled.size()); + assertEquals("p1", scheduled.get(0).id()); + } + + @Test + void deleteReturnsWhetherThePolicyExisted() { + when(repository.existsById("p1")).thenReturn(true); + assertTrue(store.delete("p1")); + verify(repository).deleteById("p1"); + + when(repository.existsById("missing")).thenReturn(false); + assertFalse(store.delete("missing")); + } + + private PolicyEntity entityFor(Policy policy) { + PolicyEntity entity = new PolicyEntity(); + entity.setId(policy.id()); + entity.setName(policy.name()); + entity.setOwner(policy.owner()); + entity.setEnabled(policy.enabled()); + entity.setTriggerType(policy.trigger() == null ? null : policy.trigger().type()); + entity.setPolicyJson(objectMapper.writeValueAsString(policy)); + return entity; + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/policy/trigger/FolderWatchTriggerTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/policy/trigger/FolderWatchTriggerTest.java new file mode 100644 index 0000000000..9b0b9f1d2a --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/policy/trigger/FolderWatchTriggerTest.java @@ -0,0 +1,178 @@ +package stirling.software.proprietary.policy.trigger; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.lenient; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.verifyNoInteractions; +import static org.mockito.Mockito.when; + +import java.nio.file.FileSystems; +import java.nio.file.Files; +import java.nio.file.Path; +import java.nio.file.WatchService; +import java.util.List; +import java.util.Map; +import java.util.Set; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.junit.jupiter.api.io.TempDir; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.proprietary.policy.engine.PolicyRunner; +import stirling.software.proprietary.policy.input.InputSource; +import stirling.software.proprietary.policy.model.InputSpec; +import stirling.software.proprietary.policy.model.OutputSpec; +import stirling.software.proprietary.policy.model.PipelineStep; +import stirling.software.proprietary.policy.model.Policy; +import stirling.software.proprietary.policy.model.TriggerConfig; +import stirling.software.proprietary.policy.store.PolicyStore; + +/** + * Tests for {@link FolderWatchTrigger}'s dispatch logic via the package-visible {@code + * runForChangedDirs}/{@code runAll}, plus its cross-facet validation. The OS watch loop and + * scheduled reconcile are thin glue around these and are not exercised here (a real {@code + * WatchService} is timing-dependent), mirroring how {@link ScheduleTriggerTest} drives {@code + * sweep} directly. The folder source is stubbed to mirror {@code FolderInputSource.watchTargets}. + */ +@ExtendWith(MockitoExtension.class) +class FolderWatchTriggerTest { + + @Mock private PolicyStore policyStore; + @Mock private PolicyRunner policyRunner; + @Mock private InputSource folderSource; + + @TempDir Path tempDir; + + private FolderWatchTrigger trigger; + + @BeforeEach + void setUp() { + trigger = + new FolderWatchTrigger( + policyStore, + policyRunner, + List.of(folderSource), + new ApplicationProperties()); + lenient().when(folderSource.supports(any())).thenReturn(true); + lenient() + .when(folderSource.watchTargets(any())) + .thenAnswer( + invocation -> { + InputSpec spec = invocation.getArgument(0); + Object dir = spec.options().get("directory"); + if (dir == null) { + throw new IllegalArgumentException( + "folder input requires a 'directory' option"); + } + return List.of(Path.of(dir.toString())); + }); + } + + @Test + void validateRejectsPolicyWithNoWatchableSource() { + assertThrows( + IllegalArgumentException.class, + () -> trigger.validate(folderWatch("p1", List.of()))); + } + + @Test + void validateAcceptsPolicyWithAFolderSource() { + trigger.validate(folderWatch("p1", List.of(InputSpec.folder("/in")))); + } + + @Test + void runsOnlyPoliciesDrawingFromTheChangedDirectory() { + Policy a = folderWatch("a", List.of(InputSpec.folder("/in/a"))); + Policy b = folderWatch("b", List.of(InputSpec.folder("/in/b"))); + when(policyStore.findByTriggerType("folder-watch")).thenReturn(List.of(a, b)); + + trigger.runForChangedDirs(Set.of(normalized("/in/a"))); + + verify(policyRunner).run(a); + verify(policyRunner, never()).run(b); + } + + @Test + void skipsAMisconfiguredPolicyButStillRunsTheOthers() { + Policy bad = folderWatch("bad", List.of(new InputSpec("folder", Map.of()))); + Policy good = folderWatch("good", List.of(InputSpec.folder("/in/a"))); + when(policyStore.findByTriggerType("folder-watch")).thenReturn(List.of(bad, good)); + + trigger.runForChangedDirs(Set.of(normalized("/in/a"))); + + verify(policyRunner).run(good); + verify(policyRunner, never()).run(bad); + } + + @Test + void anEmptyChangeSetDoesNothing() { + trigger.runForChangedDirs(Set.of()); + + verifyNoInteractions(policyStore, policyRunner); + } + + @Test + void reconcileRunsEveryFolderWatchPolicyAsASafetyNet() { + Policy a = folderWatch("a", List.of(InputSpec.folder("/in/a"))); + Policy b = folderWatch("b", List.of(InputSpec.folder("/in/b"))); + when(policyStore.findByTriggerType("folder-watch")).thenReturn(List.of(a, b)); + + trigger.runAll(); + + verify(policyRunner).run(a); + verify(policyRunner).run(b); + } + + @Test + void syncRegistrationsWatchesExistingDirsAndCancelsRemovedOnes() throws Exception { + Path dirA = Files.createDirectories(tempDir.resolve("a")); + Path dirB = Files.createDirectories(tempDir.resolve("b")); + Path missing = tempDir.resolve("missing"); // never created on disk + + Policy a = folderWatch("a", List.of(InputSpec.folder(dirA.toString()))); + Policy b = folderWatch("b", List.of(InputSpec.folder(dirB.toString()))); + Policy m = folderWatch("m", List.of(InputSpec.folder(missing.toString()))); + + WatchService service = FileSystems.getDefault().newWatchService(); + try { + trigger.watchService = service; + + when(policyStore.findByTriggerType("folder-watch")).thenReturn(List.of(a, b, m)); + trigger.syncRegistrations(); + // Existing dirs are watched; the non-existent one is skipped. + assertEquals( + Set.of(normalized(dirA.toString()), normalized(dirB.toString())), + trigger.watchedDirs()); + + // b's policy is removed: its registration is cancelled, a remains. + when(policyStore.findByTriggerType("folder-watch")).thenReturn(List.of(a)); + trigger.syncRegistrations(); + assertEquals(Set.of(normalized(dirA.toString())), trigger.watchedDirs()); + } finally { + service.close(); + } + } + + private static Path normalized(String dir) { + return Path.of(dir).toAbsolutePath().normalize(); + } + + private static Policy folderWatch(String id, List sources) { + return new Policy( + id, + "watcher", + "owner", + true, + new TriggerConfig("folder-watch", Map.of()), + sources, + List.of(new PipelineStep("/api/v1/misc/compress-pdf", Map.of())), + OutputSpec.inline()); + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/policy/trigger/PolicyTriggerManagerTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/policy/trigger/PolicyTriggerManagerTest.java new file mode 100644 index 0000000000..41504424a1 --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/policy/trigger/PolicyTriggerManagerTest.java @@ -0,0 +1,51 @@ +package stirling.software.proprietary.policy.trigger; + +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.Mockito.doThrow; +import static org.mockito.Mockito.verify; + +import java.util.List; + +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; + +/** + * Tests for {@link PolicyTriggerManager}: starts/stops every trigger, tolerating individual + * failures. + */ +@ExtendWith(MockitoExtension.class) +class PolicyTriggerManagerTest { + + @Mock private PolicyTrigger triggerA; + @Mock private PolicyTrigger triggerB; + + @Test + void startsAndStopsAllTriggers() { + PolicyTriggerManager manager = new PolicyTriggerManager(List.of(triggerA, triggerB)); + assertFalse(manager.isRunning()); + + manager.start(); + verify(triggerA).start(); + verify(triggerB).start(); + assertTrue(manager.isRunning()); + + manager.stop(); + verify(triggerA).stop(); + verify(triggerB).stop(); + assertFalse(manager.isRunning()); + } + + @Test + void oneTriggerFailingToStartDoesNotBlockTheOthers() { + doThrow(new RuntimeException("boom")).when(triggerA).start(); + PolicyTriggerManager manager = new PolicyTriggerManager(List.of(triggerA, triggerB)); + + manager.start(); + + verify(triggerB).start(); + assertTrue(manager.isRunning()); + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/policy/trigger/ScheduleTriggerTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/policy/trigger/ScheduleTriggerTest.java new file mode 100644 index 0000000000..889c50ee44 --- /dev/null +++ b/app/proprietary/src/test/java/stirling/software/proprietary/policy/trigger/ScheduleTriggerTest.java @@ -0,0 +1,147 @@ +package stirling.software.proprietary.policy.trigger; + +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.times; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.time.DayOfWeek; +import java.time.Instant; +import java.time.LocalTime; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import java.util.Set; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.proprietary.policy.engine.PolicyRunner; +import stirling.software.proprietary.policy.model.OutputSpec; +import stirling.software.proprietary.policy.model.PipelineStep; +import stirling.software.proprietary.policy.model.Policy; +import stirling.software.proprietary.policy.model.Schedule; +import stirling.software.proprietary.policy.model.TriggerConfig; +import stirling.software.proprietary.policy.store.PolicyStore; + +import tools.jackson.databind.json.JsonMapper; + +/** + * Tests for {@link ScheduleTrigger}'s due-firing logic via the package-visible {@code + * sweep(Instant)}. The trigger only decides when a policy is due; pulling sources and starting runs + * is the {@link PolicyRunner}'s job, so these assert it delegates to the runner. Schedules default + * to UTC, so explicit UTC instants make these deterministic. + */ +@ExtendWith(MockitoExtension.class) +class ScheduleTriggerTest { + + @Mock private PolicyStore policyStore; + @Mock private PolicyRunner policyRunner; + + private ScheduleTrigger trigger; + + @BeforeEach + void setUp() { + trigger = + new ScheduleTrigger( + policyStore, + policyRunner, + JsonMapper.builder().build(), + new ApplicationProperties()); + } + + @Test + void firesOncePerScheduleWhenItComesDue() { + Policy policy = scheduled("p1", new Schedule.Every(1, Schedule.Unit.MINUTES)); + when(policyStore.findByTriggerType("schedule")).thenReturn(List.of(policy)); + + Instant t0 = Instant.parse("2026-06-05T10:00:30Z"); + trigger.sweep(t0); // first sight: baseline, must not fire immediately + verify(policyRunner, never()).run(any()); + + trigger.sweep(t0.plusSeconds(120)); // the one-minute mark has passed + verify(policyRunner, times(1)).run(eq(policy)); + } + + @Test + void doesNotFireBeforeTheNextScheduledTime() { + Policy policy = scheduled("p1", new Schedule.Daily(LocalTime.of(3, 0))); // 03:00 UTC daily + when(policyStore.findByTriggerType("schedule")).thenReturn(List.of(policy)); + + Instant t0 = Instant.parse("2026-06-05T10:00:00Z"); + trigger.sweep(t0); + trigger.sweep(t0.plusSeconds(60)); // next 03:00 is far away + + verify(policyRunner, never()).run(any()); + } + + @Test + void firesWeeklyOnAChosenDay() { + // 2026-06-05 is a Friday; the next Monday 09:00 is the soonest firing. + Policy policy = + scheduled("p1", new Schedule.Weekly(Set.of(DayOfWeek.MONDAY), LocalTime.of(9, 0))); + when(policyStore.findByTriggerType("schedule")).thenReturn(List.of(policy)); + + Instant friday = Instant.parse("2026-06-05T10:00:00Z"); + trigger.sweep(friday); // baseline + trigger.sweep(Instant.parse("2026-06-08T09:00:00Z")); // Monday 09:00 + + verify(policyRunner, times(1)).run(eq(policy)); + } + + @Test + void skipsPoliciesWithAnInvalidSchedule() { + Policy policy = scheduledWithRawOptions("p1", Map.of()); // no schedule + when(policyStore.findByTriggerType("schedule")).thenReturn(List.of(policy)); + + trigger.sweep(Instant.parse("2026-06-05T10:00:00Z")); + + verify(policyRunner, never()).run(any()); + } + + @Test + void validateRejectsMissingSchedule() { + assertThrows( + IllegalArgumentException.class, + () -> trigger.validate(scheduledWithRawOptions("p1", Map.of()))); + } + + @Test + void validateRejectsAnInvalidSchedule() { + Map options = + Map.of("schedule", Map.of("type", "every", "count", -5, "unit", "MINUTES")); + assertThrows( + IllegalArgumentException.class, + () -> trigger.validate(scheduledWithRawOptions("p1", options))); + } + + @Test + void validateAcceptsAValidScheduleAndZone() { + Map options = new LinkedHashMap<>(); + options.put("schedule", new Schedule.Daily(LocalTime.of(2, 0))); + options.put("zone", "Europe/London"); + trigger.validate(scheduledWithRawOptions("p1", options)); + } + + private static Policy scheduled(String id, Schedule schedule) { + return scheduledWithRawOptions(id, Map.of("schedule", schedule)); + } + + private static Policy scheduledWithRawOptions(String id, Map options) { + return new Policy( + id, + "nightly", + "owner", + true, + new TriggerConfig("schedule", options), + List.of(new PipelineStep("/api/v1/misc/compress-pdf", Map.of())), + OutputSpec.inline()); + } +} diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/security/InitialSecuritySetupTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/security/InitialSecuritySetupTest.java index 12620b6e85..adf9486094 100644 --- a/app/proprietary/src/test/java/stirling/software/proprietary/security/InitialSecuritySetupTest.java +++ b/app/proprietary/src/test/java/stirling/software/proprietary/security/InitialSecuritySetupTest.java @@ -16,6 +16,7 @@ import org.junit.jupiter.api.extension.ExtendWith; import org.mockito.ArgumentCaptor; import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; +import org.springframework.core.env.Environment; import org.springframework.test.util.ReflectionTestUtils; import stirling.software.common.model.ApplicationProperties; @@ -36,6 +37,7 @@ class InitialSecuritySetupTest { @Mock private TeamService teamService; @Mock private DatabaseServiceInterface databaseService; @Mock private UserLicenseSettingsService licenseSettingsService; + @Mock private Environment environment; private ApplicationProperties applicationProperties; private InitialSecuritySetup initialSecuritySetup; @@ -53,13 +55,15 @@ class InitialSecuritySetupTest { when(userService.findByUsernameIgnoreCase(Role.INTERNAL_API_USER.getRoleId())) .thenReturn(Optional.of(internalUser)); when(teamService.getOrCreateInternalTeam()).thenReturn(internalTeam); + when(environment.getActiveProfiles()).thenReturn(new String[] {}); initialSecuritySetup = new InitialSecuritySetup( userService, teamService, applicationProperties, databaseService, - licenseSettingsService); + licenseSettingsService, + environment); } @Test diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/security/controller/api/UserControllerTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/security/controller/api/UserControllerTest.java index f5ea59f105..7bcc0da9f4 100644 --- a/app/proprietary/src/test/java/stirling/software/proprietary/security/controller/api/UserControllerTest.java +++ b/app/proprietary/src/test/java/stirling/software/proprietary/security/controller/api/UserControllerTest.java @@ -1,13 +1,16 @@ package stirling.software.proprietary.security.controller.api; +import static org.junit.jupiter.api.Assertions.assertEquals; import static org.mockito.ArgumentMatchers.any; import static org.mockito.Mockito.never; import static org.mockito.Mockito.verify; import static org.mockito.Mockito.when; +import static org.springframework.test.web.servlet.request.MockMvcRequestBuilders.get; import static org.springframework.test.web.servlet.request.MockMvcRequestBuilders.post; import static org.springframework.test.web.servlet.result.MockMvcResultMatchers.jsonPath; import static org.springframework.test.web.servlet.result.MockMvcResultMatchers.status; +import java.util.List; import java.util.Optional; import org.junit.jupiter.api.BeforeEach; @@ -24,6 +27,7 @@ import org.springframework.test.web.servlet.setup.MockMvcBuilders; import stirling.software.common.model.ApplicationProperties; import stirling.software.proprietary.model.Team; import stirling.software.proprietary.security.database.repository.UserRepository; +import stirling.software.proprietary.security.model.AuthenticationType; import stirling.software.proprietary.security.model.User; import stirling.software.proprietary.security.model.api.user.UsernameAndPass; import stirling.software.proprietary.security.repository.TeamRepository; @@ -150,4 +154,233 @@ class UserControllerTest { verify(loginAttemptService).resetAttempts("lockeduser"); } + + // --------------------------------------------------------------------- + // GET /api/v1/user/users - storage.signing.userListScope scoping + // --------------------------------------------------------------------- + + private static User user(long id, String username, boolean enabled, Team team) { + User u = new User(); + u.setId(id); + u.setUsername(username); + u.setEnabled(enabled); + u.setTeam(team); + return u; + } + + private static Team team(long id, String name) { + Team t = new Team(); + t.setId(id); + t.setName(name); + return t; + } + + private static Authentication auth(String username) { + return new UsernamePasswordAuthenticationToken(username, "pw"); + } + + @Test + void listUsersDefaultScopeIsOrgWide() throws Exception { + // Default "org" scope returns every enabled user via findAll(), no team lookup. + Team alpha = team(1L, "alpha"); + when(userRepository.findAll()) + .thenReturn( + List.of( + user(1L, "a@alpha.com", true, alpha), + user(2L, "b@alpha.com", true, alpha))); + + mockMvc.perform(get("/api/v1/user/users").principal(auth("a@alpha.com"))) + .andExpect(status().isOk()) + .andExpect(jsonPath("$.length()").value(2)) + .andExpect(jsonPath("$[0].username").value("a@alpha.com")) + .andExpect(jsonPath("$[1].username").value("b@alpha.com")); + + // Caller is resolved (for the anonymous-gate) but org scope still uses findAll, not team. + verify(userRepository, never()).findAllByTeamId(any()); + } + + @Test + void listUsersForbiddenForAnonymousCaller() throws Exception { + // Anonymous SaaS accounts must never enumerate users, regardless of scope. + User anon = user(1L, "anon_abc", true, team(1L, TeamService.DEFAULT_TEAM_NAME)); + anon.setAuthenticationType(AuthenticationType.ANONYMOUS); + when(userService.findByUsernameIgnoreCase("anon_abc")).thenReturn(Optional.of(anon)); + + mockMvc.perform(get("/api/v1/user/users").principal(auth("anon_abc"))) + .andExpect(status().isForbidden()); + + verify(userRepository, never()).findAll(); + verify(userRepository, never()).findAllByTeamId(any()); + } + + @Test + void listUsersOrgScopeFiltersDisabledUsers() throws Exception { + Team alpha = team(1L, "alpha"); + when(userRepository.findAll()) + .thenReturn( + List.of( + user(1L, "enabled@alpha.com", true, alpha), + user(2L, "disabled@alpha.com", false, alpha))); + + mockMvc.perform(get("/api/v1/user/users").principal(auth("enabled@alpha.com"))) + .andExpect(status().isOk()) + .andExpect(jsonPath("$.length()").value(1)) + .andExpect(jsonPath("$[0].username").value("enabled@alpha.com")); + } + + @Test + void listUsersTeamScopeReturnsOnlyCallerTeam() throws Exception { + applicationProperties.getStorage().getSigning().setUserListScope("team"); + Team alpha = team(7L, "alpha"); + User caller = user(1L, "caller@alpha.com", true, alpha); + when(userService.findByUsernameIgnoreCase("caller@alpha.com")) + .thenReturn(Optional.of(caller)); + when(userRepository.findAllByTeamId(7L)) + .thenReturn(List.of(caller, user(2L, "mate@alpha.com", true, alpha))); + + mockMvc.perform(get("/api/v1/user/users").principal(auth("caller@alpha.com"))) + .andExpect(status().isOk()) + .andExpect(jsonPath("$.length()").value(2)) + .andExpect(jsonPath("$[0].teamName").value("alpha")); + + verify(userRepository).findAllByTeamId(7L); + verify(userRepository, never()).findAll(); + } + + @Test + void listUsersTeamScopeWithMissingCallerReturnsEmpty() throws Exception { + applicationProperties.getStorage().getSigning().setUserListScope("team"); + when(userService.findByUsernameIgnoreCase("ghost@alpha.com")).thenReturn(Optional.empty()); + + mockMvc.perform(get("/api/v1/user/users").principal(auth("ghost@alpha.com"))) + .andExpect(status().isOk()) + .andExpect(jsonPath("$.length()").value(0)); + + verify(userRepository, never()).findAllByTeamId(any()); + verify(userRepository, never()).findAll(); + } + + @Test + void listUsersTeamScopeWithNullTeamReturnsSelfOnly() throws Exception { + applicationProperties.getStorage().getSigning().setUserListScope("team"); + User caller = user(1L, "solo@nowhere.com", true, null); + when(userService.findByUsernameIgnoreCase("solo@nowhere.com")) + .thenReturn(Optional.of(caller)); + + mockMvc.perform(get("/api/v1/user/users").principal(auth("solo@nowhere.com"))) + .andExpect(status().isOk()) + .andExpect(jsonPath("$.length()").value(1)) + .andExpect(jsonPath("$[0].username").value("solo@nowhere.com")); + + verify(userRepository, never()).findAllByTeamId(any()); + verify(userRepository, never()).findAll(); + } + + @Test + void listUsersTeamScopeOnDefaultTeamReturnsSelfOnly() throws Exception { + // A caller on a shared system team must not enumerate its members. + applicationProperties.getStorage().getSigning().setUserListScope("team"); + Team defaultTeam = team(1L, TeamService.DEFAULT_TEAM_NAME); + User caller = user(1L, "new@saas.com", true, defaultTeam); + when(userService.findByUsernameIgnoreCase("new@saas.com")).thenReturn(Optional.of(caller)); + + mockMvc.perform(get("/api/v1/user/users").principal(auth("new@saas.com"))) + .andExpect(status().isOk()) + .andExpect(jsonPath("$.length()").value(1)) + .andExpect(jsonPath("$[0].username").value("new@saas.com")); + + verify(userRepository, never()).findAllByTeamId(any()); + verify(userRepository, never()).findAll(); + } + + @Test + void listUsersTeamScopeOnInternalTeamReturnsSelfOnly() throws Exception { + applicationProperties.getStorage().getSigning().setUserListScope("team"); + Team internalTeam = team(2L, TeamService.INTERNAL_TEAM_NAME); + User caller = user(1L, "svc@saas.com", true, internalTeam); + when(userService.findByUsernameIgnoreCase("svc@saas.com")).thenReturn(Optional.of(caller)); + + mockMvc.perform(get("/api/v1/user/users").principal(auth("svc@saas.com"))) + .andExpect(status().isOk()) + .andExpect(jsonPath("$.length()").value(1)) + .andExpect(jsonPath("$[0].username").value("svc@saas.com")); + + verify(userRepository, never()).findAllByTeamId(any()); + verify(userRepository, never()).findAll(); + } + + @Test + void listUsersFailsClosedOnUnrecognisedScope() throws Exception { + // Any non-"org" value must restrict to the caller's team, not leak the instance. + applicationProperties.getStorage().getSigning().setUserListScope("tewm"); + Team alpha = team(3L, "alpha"); + when(userService.findByUsernameIgnoreCase("caller@alpha.com")) + .thenReturn(Optional.of(user(1L, "caller@alpha.com", true, alpha))); + when(userRepository.findAllByTeamId(3L)) + .thenReturn(List.of(user(1L, "caller@alpha.com", true, alpha))); + + mockMvc.perform(get("/api/v1/user/users").principal(auth("caller@alpha.com"))) + .andExpect(status().isOk()); + + verify(userRepository).findAllByTeamId(3L); + verify(userRepository, never()).findAll(); + } + + @Test + void listUsersFailsClosedOnBlankScope() throws Exception { + applicationProperties.getStorage().getSigning().setUserListScope(" "); + Team alpha = team(4L, "alpha"); + when(userService.findByUsernameIgnoreCase("caller@alpha.com")) + .thenReturn(Optional.of(user(1L, "caller@alpha.com", true, alpha))); + when(userRepository.findAllByTeamId(4L)).thenReturn(List.of()); + + mockMvc.perform(get("/api/v1/user/users").principal(auth("caller@alpha.com"))) + .andExpect(status().isOk()); + + verify(userRepository).findAllByTeamId(4L); + verify(userRepository, never()).findAll(); + } + + @Test + void listUsersFailsClosedOnNullScope() throws Exception { + // A null value must also fail closed to the caller's team. + applicationProperties.getStorage().getSigning().setUserListScope(null); + Team alpha = team(9L, "alpha"); + when(userService.findByUsernameIgnoreCase("caller@alpha.com")) + .thenReturn(Optional.of(user(1L, "caller@alpha.com", true, alpha))); + when(userRepository.findAllByTeamId(9L)).thenReturn(List.of()); + + mockMvc.perform(get("/api/v1/user/users").principal(auth("caller@alpha.com"))) + .andExpect(status().isOk()); + + verify(userRepository).findAllByTeamId(9L); + verify(userRepository, never()).findAll(); + } + + @Test + void listUsersOrgScopeIsCaseInsensitive() throws Exception { + applicationProperties.getStorage().getSigning().setUserListScope("ORG"); + when(userRepository.findAll()).thenReturn(List.of(user(1L, "a@alpha.com", true, null))); + + mockMvc.perform(get("/api/v1/user/users").principal(auth("a@alpha.com"))) + .andExpect(status().isOk()); + + verify(userRepository).findAll(); + verify(userRepository, never()).findAllByTeamId(any()); + } + + @Test + void listUsersRequiresAuthentication() throws Exception { + mockMvc.perform(get("/api/v1/user/users")).andExpect(status().isUnauthorized()); + + verify(userRepository, never()).findAll(); + verify(userRepository, never()).findAllByTeamId(any()); + } + + @Test + void signingUserListScopeDefaultsToOrg() { + // Self-host backward-compat: default must stay "org" (saas profile flips it to "team"). + assertEquals( + "org", new ApplicationProperties().getStorage().getSigning().getUserListScope()); + } } diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/service/AiEngineClientTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/service/AiEngineClientTest.java index f4c9af5920..259ae55a0b 100644 --- a/app/proprietary/src/test/java/stirling/software/proprietary/service/AiEngineClientTest.java +++ b/app/proprietary/src/test/java/stirling/software/proprietary/service/AiEngineClientTest.java @@ -84,4 +84,47 @@ class AiEngineClientTest { assertEquals(HttpStatus.SERVICE_UNAVAILABLE, ex.getStatusCode()); } + + @Test + @SuppressWarnings("unchecked") + void allVerbsSendEngineAuthHeaderWhenSecretConfigured() throws Exception { + // Regression for the engine shared-secret hardening: every verb that hits a non-public + // engine route (post/delete/get) must present X-Engine-Auth, or the route 401s once the + // secret is set. delete() backs the logout-time RAG purge, so a miss silently leaks data. + AiEngineClient secured = + new AiEngineClient(applicationProperties, httpClient, "top-secret"); + HttpResponse ok = mock(HttpResponse.class); + when(ok.statusCode()).thenReturn(200); + when(ok.body()).thenReturn("{}"); + org.mockito.ArgumentCaptor captor = + org.mockito.ArgumentCaptor.forClass(java.net.http.HttpRequest.class); + when(httpClient.send(captor.capture(), any(HttpResponse.BodyHandler.class))).thenReturn(ok); + + secured.post("/api/v1/pdf-question", "{}", "alice"); + secured.delete("/api/v1/documents/by-owner", "alice"); + secured.get("/api/v1/x", "alice"); + + for (java.net.http.HttpRequest req : captor.getAllValues()) { + assertEquals( + "top-secret", + req.headers().firstValue("X-Engine-Auth").orElse(null), + req.method() + " must carry the engine shared secret"); + } + } + + @Test + @SuppressWarnings("unchecked") + void noEngineAuthHeaderWhenSecretUnset() throws Exception { + AiEngineClient noSecret = new AiEngineClient(applicationProperties, httpClient, null); + HttpResponse ok = mock(HttpResponse.class); + when(ok.statusCode()).thenReturn(200); + when(ok.body()).thenReturn("{}"); + org.mockito.ArgumentCaptor captor = + org.mockito.ArgumentCaptor.forClass(java.net.http.HttpRequest.class); + when(httpClient.send(captor.capture(), any(HttpResponse.BodyHandler.class))).thenReturn(ok); + + noSecret.delete("/api/v1/documents/by-owner", "alice"); + + assertEquals(null, captor.getValue().headers().firstValue("X-Engine-Auth").orElse(null)); + } } diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/service/AiWorkflowServiceTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/service/AiWorkflowServiceTest.java index 5b73ea5cc6..9e611cf32e 100644 --- a/app/proprietary/src/test/java/stirling/software/proprietary/service/AiWorkflowServiceTest.java +++ b/app/proprietary/src/test/java/stirling/software/proprietary/service/AiWorkflowServiceTest.java @@ -57,6 +57,7 @@ import stirling.software.proprietary.model.api.ai.AiWorkflowFileInput; import stirling.software.proprietary.model.api.ai.AiWorkflowOutcome; import stirling.software.proprietary.model.api.ai.AiWorkflowRequest; import stirling.software.proprietary.model.api.ai.AiWorkflowResponse; +import stirling.software.proprietary.policy.engine.PolicyExecutor; import tools.jackson.databind.ObjectMapper; import tools.jackson.databind.json.JsonMapper; @@ -107,18 +108,20 @@ class AiWorkflowServiceTest { .when(fileIdStrategy.idFor(any(MultipartFile.class))) .thenAnswer(inv -> ((MultipartFile) inv.getArgument(0)).getOriginalFilename()); + PolicyExecutor policyExecutor = + new PolicyExecutor( + internalApiClient, toolMetadataService, tempFileManager, objectMapper); service = new AiWorkflowService( pdfDocumentFactory, aiEngineClient, pdfContentExtractor, objectMapper, - internalApiClient, fileStorage, - toolMetadataService, tempFileManager, fileIdStrategy, endpointResolver, + policyExecutor, null, new ApplicationProperties()); when(endpointResolver.getEnabledEndpointUrls()).thenReturn(List.of()); @@ -436,6 +439,35 @@ class AiWorkflowServiceTest { verify(internalApiClient, never()).post(anyString(), any()); } + @Test + void convertMarkdownRunsDeterministicConversionAndReturnsMdFile() throws IOException { + MockMultipartFile input = pdf("multi-column-test_lorem.pdf", "pdf-bytes"); + when(fileIdStrategy.idFor(any())).thenReturn("doc-1"); + stubOrchestrator( + """ + { + "outcome":"convert_markdown", + "reason":"PDF to Markdown requested.", + "filesToIngest":[{"id":"doc-1","name":"multi-column-test_lorem.pdf"}] + } + """); + when(toolMetadataService.shouldUnpackZipResponse("/api/v1/convert/pdf/markdown")) + .thenReturn(false); + stubEndpoint( + "/api/v1/convert/pdf/markdown", + pdfResource("# Title", "multi-column-test_lorem.md")); + AtomicInteger ids = stubFileStorage(); + + AiWorkflowResponse result = service.orchestrate(requestFor(input, "convert to markdown")); + + assertEquals(AiWorkflowOutcome.COMPLETED, result.getOutcome()); + assertEquals(1, result.getResultFiles().size()); + // Extension changes (pdf -> md), so the converter's response filename wins. + assertEquals("multi-column-test_lorem.md", result.getResultFiles().get(0).getFileName()); + assertEquals(1, ids.get()); + verify(internalApiClient, times(1)).post(eq("/api/v1/convert/pdf/markdown"), any()); + } + @Test void toolCallWithoutEndpointFallsBackToCannotContinue() throws IOException { MockMultipartFile input = pdf("input.pdf", "bytes"); diff --git a/app/proprietary/src/test/java/stirling/software/proprietary/service/PageLayoutArtifactContractTest.java b/app/proprietary/src/test/java/stirling/software/proprietary/service/PageLayoutArtifactContractTest.java deleted file mode 100644 index ae853b2e6a..0000000000 --- a/app/proprietary/src/test/java/stirling/software/proprietary/service/PageLayoutArtifactContractTest.java +++ /dev/null @@ -1,66 +0,0 @@ -package stirling.software.proprietary.service; - -import static org.junit.jupiter.api.Assertions.assertEquals; -import static org.junit.jupiter.api.Assertions.assertTrue; - -import java.util.List; - -import org.junit.jupiter.api.Test; - -import stirling.software.proprietary.service.PdfContentExtractor.LayoutFragment; -import stirling.software.proprietary.service.PdfContentExtractor.LayoutLine; -import stirling.software.proprietary.service.PdfContentExtractor.LayoutPage; -import stirling.software.proprietary.service.PdfContentExtractor.PageLayoutArtifact; -import stirling.software.proprietary.service.PdfContentExtractor.PageLayoutFileResult; - -import tools.jackson.databind.JsonNode; -import tools.jackson.databind.json.JsonMapper; - -/** - * Contract test: verifies that {@link PageLayoutArtifact} serializes to the JSON field names that - * the Python engine expects in {@code engine/src/stirling/contracts/pdf_to_markdown.py}. - * - *

The companion Python test in {@code tests/test_pdf_to_markdown.py} deserializes the same JSON - * literal and asserts field values. If either side renames a field, one of these tests fails. - */ -class PageLayoutArtifactContractTest { - - static final String CONTRACT_JSON = - """ - {"kind":"page_layout","files":[{"fileName":"test.pdf","pages":[{"pageNumber":1,"lines":[{"y":10.0,"fragments":[{"text":"Hello","x":1.0,"y":2.0,"width":30.0,"fontSize":12.0,"bold":true}]}]}]}]}"""; - - @Test - void pageLayoutArtifact_serialisesToExpectedJson() throws Exception { - LayoutFragment fragment = new LayoutFragment("Hello", 1.0f, 2.0f, 30.0f, 12.0f, true); - LayoutLine line = new LayoutLine(10.0f, List.of(fragment)); - LayoutPage page = new LayoutPage(1, List.of(line)); - - PageLayoutFileResult fileResult = new PageLayoutFileResult(); - fileResult.setFileName("test.pdf"); - fileResult.setPages(List.of(page)); - - PageLayoutArtifact artifact = new PageLayoutArtifact(); - artifact.setFiles(List.of(fileResult)); - - JsonNode json = new JsonMapper().valueToTree(artifact); - - assertEquals("page_layout", json.get("kind").asText()); - - JsonNode file = json.get("files").get(0); - assertEquals("test.pdf", file.get("fileName").asText()); - - JsonNode pg = file.get("pages").get(0); - assertEquals(1, pg.get("pageNumber").asInt()); - - JsonNode ln = pg.get("lines").get(0); - assertEquals(10.0, ln.get("y").asDouble(), 0.001); - - JsonNode frag = ln.get("fragments").get(0); - assertEquals("Hello", frag.get("text").asText()); - assertEquals(1.0, frag.get("x").asDouble(), 0.001); - assertEquals(2.0, frag.get("y").asDouble(), 0.001); - assertEquals(30.0, frag.get("width").asDouble(), 0.001); - assertEquals(12.0, frag.get("fontSize").asDouble(), 0.001); - assertTrue(frag.get("bold").asBoolean()); - } -} diff --git a/app/saas/src/main/java/stirling/software/saas/ai/controller/AiCreateController.java b/app/saas/src/main/java/stirling/software/saas/ai/controller/AiCreateController.java index 85520a239b..371060c032 100644 --- a/app/saas/src/main/java/stirling/software/saas/ai/controller/AiCreateController.java +++ b/app/saas/src/main/java/stirling/software/saas/ai/controller/AiCreateController.java @@ -29,36 +29,45 @@ import org.springframework.web.servlet.mvc.method.annotation.StreamingResponseBo import com.fasterxml.jackson.core.JsonProcessingException; import com.fasterxml.jackson.databind.ObjectMapper; +import io.swagger.v3.oas.annotations.Hidden; +import io.swagger.v3.oas.annotations.tags.Tag; + import jakarta.servlet.http.HttpServletRequest; import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; import stirling.software.proprietary.security.database.repository.UserRepository; +import stirling.software.proprietary.security.model.ApiKeyAuthenticationToken; import stirling.software.proprietary.security.model.User; import stirling.software.saas.ai.model.AiCreateSession; import stirling.software.saas.ai.repository.AiCreateSessionRepository; import stirling.software.saas.ai.service.AiCreateProxyService; import stirling.software.saas.ai.service.AiCreateSessionService; -import stirling.software.saas.service.CreditService; -import stirling.software.saas.service.TeamCreditService; +import stirling.software.saas.payg.cap.RequiresFeature; +import stirling.software.saas.payg.charge.ChargeContext; +import stirling.software.saas.payg.charge.JobChargeService; +import stirling.software.saas.payg.model.BillingCategory; +import stirling.software.saas.payg.model.FeatureGate; +import stirling.software.saas.payg.model.JobSource; +import stirling.software.saas.payg.model.ProcessType; import stirling.software.saas.util.AuthenticationUtils; -import stirling.software.saas.util.CreditHeaderUtils; @RestController @Profile("saas") @RequestMapping("/api/v1/ai/create") +@Tag(name = "AI") +@Hidden @RequiredArgsConstructor +@RequiresFeature(FeatureGate.AI_SUPPORT) @Slf4j public class AiCreateController { private final AiCreateSessionService sessionService; private final AiCreateProxyService proxyService; private final ObjectMapper objectMapper = new ObjectMapper(); - private final CreditService creditService; - private final TeamCreditService teamCreditService; private final UserRepository userRepository; - private final CreditHeaderUtils creditHeaderUtils; + private final JobChargeService jobChargeService; @PostMapping("/sessions") public ResponseEntity createSession( @@ -79,9 +88,45 @@ public class AiCreateController { session.getUserId(), session.getDocType(), session.getTemplateId()); + chargeForCreate(session); return ResponseEntity.ok(new CreateSessionResponse(session.getSessionId())); } + /** + * Bill one document for a new AI Create session — creating a document is the charge point; + * follow-up edits on the same session (outline / reprompt / draft / template / stream) carry no + * charge. AI usage is billable, so a JWT (web) session counts the same as an API-key one. + * + *

Best-effort: a charge failure must not block the user's session. Entitlement is already + * enforced upstream — this controller is {@code @RequiresFeature(AI_SUPPORT)}, so the + * EntitlementGuard 402s a team with no AI allowance before we ever get here; this call only + * does the accounting (free-grant draw + Stripe meter). + */ + private void chargeForCreate(AiCreateSession session) { + try { + Authentication auth = SecurityContextHolder.getContext().getAuthentication(); + User user = AuthenticationUtils.getCurrentUser(auth, userRepository); + if (user == null || user.getTeam() == null) { + return; + } + JobSource source = + auth instanceof ApiKeyAuthenticationToken ? JobSource.API : JobSource.WEB; + ChargeContext ctx = + new ChargeContext( + user.getId(), + user.getTeam().getId(), + source, + ProcessType.SINGLE_TOOL, + BillingCategory.AI); + jobChargeService.chargeStandalone(ctx, 1); + } catch (RuntimeException e) { + log.warn( + "AI create session {} charge failed; session proceeds unbilled: {}", + session.getSessionId(), + e.getMessage()); + } + } + @DeleteMapping("/sessions/{sessionId}") public ResponseEntity deleteSession(@PathVariable String sessionId) { sessionService.deleteSessionForCurrentUser(sessionId); @@ -182,8 +227,7 @@ public class AiCreateController { @PathVariable String sessionId, HttpServletRequest request) { sessionService.getSessionForCurrentUser(sessionId); log.info("AI create fillFields sessionId={}", sessionId); - return proxy( - "POST", "/api/create/sessions/" + sessionId + "/fields", request, false, false); + return proxy("POST", "/api/create/sessions/" + sessionId + "/fields", request, false); } @GetMapping( @@ -192,20 +236,11 @@ public class AiCreateController { public ResponseEntity stream( @PathVariable String sessionId, HttpServletRequest request) { sessionService.getSessionForCurrentUser(sessionId); - return proxy( - "GET", - "/api/create/sessions/" + sessionId + "/stream", - request, - true, - true); // Add credits header: frontend endpoint that triggers AI + return proxy("GET", "/api/create/sessions/" + sessionId + "/stream", request, true); } private ResponseEntity proxy( - String method, - String path, - HttpServletRequest request, - boolean acceptEventStream, - boolean includeCreditsHeader) { + String method, String path, HttpServletRequest request, boolean acceptEventStream) { try { HttpResponse response = proxyService.forward(method, path, request, acceptEventStream); @@ -219,11 +254,6 @@ public class AiCreateController { headers.set(HttpHeaders.CONTENT_TYPE, MediaType.TEXT_EVENT_STREAM_VALUE); } - // Add credit headers if requested - if (includeCreditsHeader) { - addCreditHeaders(headers); - } - StreamingResponseBody body = outputStream -> { try (InputStream inputStream = response.body()) { @@ -251,31 +281,6 @@ public class AiCreateController { .ifPresent(value -> headers.set(headerName, value)); } - /** - * Add credit headers to the response headers. - * - * @param headers The headers to add credit information to - */ - private void addCreditHeaders(HttpHeaders headers) { - try { - Authentication auth = SecurityContextHolder.getContext().getAuthentication(); - if (auth == null || !auth.isAuthenticated()) { - log.debug("[AI-CREATE] No authentication found, skipping credit header"); - return; - } - - User user = AuthenticationUtils.getCurrentUser(auth, userRepository); - int remainingCredits = - creditHeaderUtils.getRemainingCredits(user, creditService, teamCreditService); - if (remainingCredits >= 0) { - headers.set("X-Credits-Remaining", Integer.toString(remainingCredits)); - log.warn("[AI-CREATE] Added X-Credits-Remaining header: {}", remainingCredits); - } - } catch (Exception e) { - log.error("[AI-CREATE] Failed to add credit header: {}", e.getMessage(), e); - } - } - public record CreateSessionRequest( String prompt, String docType, diff --git a/app/saas/src/main/java/stirling/software/saas/ai/controller/AiCreateInternalController.java b/app/saas/src/main/java/stirling/software/saas/ai/controller/AiCreateInternalController.java index 34c44176e5..60c8dc4615 100644 --- a/app/saas/src/main/java/stirling/software/saas/ai/controller/AiCreateInternalController.java +++ b/app/saas/src/main/java/stirling/software/saas/ai/controller/AiCreateInternalController.java @@ -17,17 +17,25 @@ import org.springframework.web.server.ResponseStatusException; import com.fasterxml.jackson.core.JsonProcessingException; import com.fasterxml.jackson.databind.ObjectMapper; +import io.swagger.v3.oas.annotations.Hidden; +import io.swagger.v3.oas.annotations.tags.Tag; + import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; import stirling.software.saas.ai.model.AiCreateSession; import stirling.software.saas.ai.model.AiCreateSessionStatus; import stirling.software.saas.ai.service.AiCreateSessionService; +import stirling.software.saas.payg.cap.RequiresFeature; +import stirling.software.saas.payg.model.FeatureGate; @RestController @Profile("saas") @RequestMapping("/api/v1/ai/create/internal") +@Tag(name = "AI") +@Hidden @RequiredArgsConstructor +@RequiresFeature(FeatureGate.AI_SUPPORT) @Slf4j public class AiCreateInternalController { diff --git a/app/saas/src/main/java/stirling/software/saas/ai/controller/AiProxyController.java b/app/saas/src/main/java/stirling/software/saas/ai/controller/AiProxyController.java index 074da8707c..82f1678fed 100644 --- a/app/saas/src/main/java/stirling/software/saas/ai/controller/AiProxyController.java +++ b/app/saas/src/main/java/stirling/software/saas/ai/controller/AiProxyController.java @@ -9,8 +9,6 @@ import org.springframework.http.HttpHeaders; import org.springframework.http.HttpStatus; import org.springframework.http.MediaType; import org.springframework.http.ResponseEntity; -import org.springframework.security.core.Authentication; -import org.springframework.security.core.context.SecurityContextHolder; import org.springframework.web.bind.annotation.GetMapping; import org.springframework.web.bind.annotation.PathVariable; import org.springframework.web.bind.annotation.PostMapping; @@ -18,122 +16,110 @@ import org.springframework.web.bind.annotation.RequestMapping; import org.springframework.web.bind.annotation.RestController; import org.springframework.web.servlet.mvc.method.annotation.StreamingResponseBody; +import io.swagger.v3.oas.annotations.Hidden; +import io.swagger.v3.oas.annotations.tags.Tag; + import jakarta.servlet.http.HttpServletRequest; import lombok.extern.slf4j.Slf4j; -import stirling.software.proprietary.security.database.repository.UserRepository; -import stirling.software.proprietary.security.model.User; import stirling.software.saas.ai.service.AiProxyService; -import stirling.software.saas.service.CreditService; -import stirling.software.saas.service.TeamCreditService; -import stirling.software.saas.util.AuthenticationUtils; -import stirling.software.saas.util.CreditHeaderUtils; +import stirling.software.saas.payg.cap.RequiresFeature; +import stirling.software.saas.payg.model.FeatureGate; @RestController @Profile("saas") @RequestMapping("/api/v1/ai") +@RequiresFeature(FeatureGate.AI_SUPPORT) +@Tag(name = "AI") +@Hidden @Slf4j public class AiProxyController { private final AiProxyService aiProxyService; - private final CreditService creditService; - private final TeamCreditService teamCreditService; - private final UserRepository userRepository; - private final CreditHeaderUtils creditHeaderUtils; - public AiProxyController( - AiProxyService aiProxyService, - CreditService creditService, - TeamCreditService teamCreditService, - UserRepository userRepository, - CreditHeaderUtils creditHeaderUtils) { + public AiProxyController(AiProxyService aiProxyService) { this.aiProxyService = aiProxyService; - this.creditService = creditService; - this.teamCreditService = teamCreditService; - this.userRepository = userRepository; - this.creditHeaderUtils = creditHeaderUtils; } @PostMapping("/generate_section") public ResponseEntity generateSection(HttpServletRequest request) { - return proxy("POST", "/api/generate_section", request, false, false); + return proxy("POST", "/api/generate_section", request, false); } @PostMapping("/generate_all_sections") public ResponseEntity generateAllSections(HttpServletRequest request) { - return proxy("POST", "/api/generate_all_sections", request, false, false); + return proxy("POST", "/api/generate_all_sections", request, false); } @PostMapping("/intent/check") public ResponseEntity intentCheck(HttpServletRequest request) { - return proxy("POST", "/api/intent/check", request, false, false); + return proxy("POST", "/api/intent/check", request, false); } @PostMapping("/chat/route") public ResponseEntity chatRoute(HttpServletRequest request) { - return proxy("POST", "/api/chat/route", request, false, true); + return proxy("POST", "/api/chat/route", request, false); } @PostMapping("/chat/create-smart-folder") public ResponseEntity createSmartFolder(HttpServletRequest request) { - return proxy("POST", "/api/chat/create-smart-folder", request, false, true); + return proxy("POST", "/api/chat/create-smart-folder", request, false); } @PostMapping("/chat/info") public ResponseEntity chatInfo(HttpServletRequest request) { - return proxy("POST", "/api/chat/info", request, false, true); + return proxy("POST", "/api/chat/info", request, false); } @PostMapping("/pdf/answer") public ResponseEntity pdfAnswer(HttpServletRequest request) { - return proxy("POST", "/api/pdf/answer", request, false, false); + return proxy("POST", "/api/pdf/answer", request, false); } @PostMapping("/progressive_render") public ResponseEntity progressiveRender(HttpServletRequest request) { - return proxy("POST", "/api/progressive_render", request, false, false); + return proxy("POST", "/api/progressive_render", request, false); } @GetMapping("/versions/{userId}") public ResponseEntity versions( @PathVariable("userId") String userId, HttpServletRequest request) { - return proxy("GET", "/api/versions/" + userId, request, false, false); + return proxy("GET", "/api/versions/" + userId, request, false); } @GetMapping("/style/{userId}") public ResponseEntity style( @PathVariable("userId") String userId, HttpServletRequest request) { - return proxy("GET", "/api/style/" + userId, request, false, false); + return proxy("GET", "/api/style/" + userId, request, false); } @PostMapping("/style/{userId}") public ResponseEntity updateStyle( @PathVariable("userId") String userId, HttpServletRequest request) { - return proxy("POST", "/api/style/" + userId, request, false, false); + return proxy("POST", "/api/style/" + userId, request, false); } @PostMapping("/import_template") public ResponseEntity importTemplate(HttpServletRequest request) { - return proxy("POST", "/api/import_template", request, false, false); + return proxy("POST", "/api/import_template", request, false); } @PostMapping("/edit/sessions") public ResponseEntity createEditSession(HttpServletRequest request) { - return proxy("POST", "/api/edit/sessions", request, false, false); + return proxy("POST", "/api/edit/sessions", request, false); } @PostMapping("/edit/sessions/{sessionId}/messages") public ResponseEntity editSessionMessage( @PathVariable("sessionId") String sessionId, HttpServletRequest request) { - return proxy("POST", "/api/edit/sessions/" + sessionId + "/messages", request, false, true); + return proxy("POST", "/api/edit/sessions/" + sessionId + "/messages", request, false); } @PostMapping("/edit/sessions/{sessionId}/attachments") public ResponseEntity editSessionAttachment( @PathVariable("sessionId") String sessionId, HttpServletRequest request) { - return proxy( - "POST", "/api/edit/sessions/" + sessionId + "/attachments", request, false, false); + return proxy("POST", "/api/edit/sessions/" + sessionId + "/attachments", request, false); } @PostMapping( @@ -141,17 +127,17 @@ public class AiProxyController { produces = MediaType.TEXT_EVENT_STREAM_VALUE) public ResponseEntity runEditSession( @PathVariable("sessionId") String sessionId, HttpServletRequest request) { - return proxy("POST", "/api/edit/sessions/" + sessionId + "/run", request, true, false); + return proxy("POST", "/api/edit/sessions/" + sessionId + "/run", request, true); } @GetMapping("/pdf-editor/document") public ResponseEntity pdfEditorDocument(HttpServletRequest request) { - return proxy("GET", "/api/pdf-editor/document", request, false, false); + return proxy("GET", "/api/pdf-editor/document", request, false); } @PostMapping("/pdf-editor/upload") public ResponseEntity pdfEditorUpload(HttpServletRequest request) { - return proxy("POST", "/api/pdf-editor/upload", request, false, false); + return proxy("POST", "/api/pdf-editor/upload", request, false); } @GetMapping("/output/**") @@ -159,27 +145,22 @@ public class AiProxyController { String requestUri = request.getRequestURI(); String prefix = request.getContextPath() + "/api/v1/ai/output/"; String path = requestUri.startsWith(prefix) ? requestUri.substring(prefix.length()) : ""; - return proxy("GET", "/output/" + path, request, false, false); + return proxy("GET", "/output/" + path, request, false); } // Health endpoint at /api/v1/ai/health is owned by the proprietary AiEngineController; both // proxy to the same backing AI engine. No need for credit-aware wrapping on a health probe. /** - * Proxy method that optionally adds credit headers. + * Proxy method. * * @param method HTTP method * @param path API path * @param request The incoming request * @param acceptEventStream Whether to accept event stream responses - * @param includeCreditsHeader Whether to add credit balance header */ private ResponseEntity proxy( - String method, - String path, - HttpServletRequest request, - boolean acceptEventStream, - boolean includeCreditsHeader) { + String method, String path, HttpServletRequest request, boolean acceptEventStream) { try { // Forward to AI backend HttpResponse aiResponse = @@ -196,11 +177,6 @@ public class AiProxyController { headers.set(HttpHeaders.CONTENT_TYPE, MediaType.TEXT_EVENT_STREAM_VALUE); } - // Add credit headers if requested (after AI processing completes) - if (includeCreditsHeader) { - addCreditHeaders(headers); - } - StreamingResponseBody body = outputStream -> { try (InputStream inputStream = aiResponse.body()) { @@ -239,30 +215,4 @@ public class AiProxyController { } }); } - - /** - * Add credit headers to the response headers. - * - * @param headers The headers to add credit information to - */ - private void addCreditHeaders(HttpHeaders headers) { - try { - Authentication auth = SecurityContextHolder.getContext().getAuthentication(); - if (auth == null || !auth.isAuthenticated()) { - log.debug("[AI-PROXY] No authentication found, skipping credit header"); - return; - } - - User user = AuthenticationUtils.getCurrentUser(auth, userRepository); - int remainingCredits = - creditHeaderUtils.getRemainingCredits(user, creditService, teamCreditService); - if (remainingCredits >= 0) { - headers.set("X-Credits-Remaining", Integer.toString(remainingCredits)); - log.warn("[AI-PROXY] Added X-Credits-Remaining header: {}", remainingCredits); - } - headers.set("X-Credit-Source", "AI_TOOL_CALL"); - } catch (Exception e) { - log.error("[AI-PROXY] Failed to add credit header: {}", e.getMessage(), e); - } - } } diff --git a/app/saas/src/main/java/stirling/software/saas/config/CreditInterceptorConfig.java b/app/saas/src/main/java/stirling/software/saas/config/CreditInterceptorConfig.java deleted file mode 100644 index 978e9b34fa..0000000000 --- a/app/saas/src/main/java/stirling/software/saas/config/CreditInterceptorConfig.java +++ /dev/null @@ -1,29 +0,0 @@ -package stirling.software.saas.config; - -import org.springframework.context.annotation.Configuration; -import org.springframework.context.annotation.Profile; -import org.springframework.web.servlet.config.annotation.InterceptorRegistry; -import org.springframework.web.servlet.config.annotation.WebMvcConfigurer; - -import lombok.RequiredArgsConstructor; - -import stirling.software.saas.interceptor.UnifiedCreditInterceptor; - -@Configuration -@Profile("saas") -@RequiredArgsConstructor -public class CreditInterceptorConfig implements WebMvcConfigurer { - - private final UnifiedCreditInterceptor unifiedCreditInterceptor; - private final CreditsProperties creditsProperties; - - @Override - public void addInterceptors(InterceptorRegistry registry) { - if (creditsProperties.isEnabled()) { - registry.addInterceptor(unifiedCreditInterceptor) - .addPathPatterns("/api/**") - .excludePathPatterns( - "/api/v1/credits/**", "/api/v1/config/**", "/api/v1/info/**"); - } - } -} diff --git a/app/saas/src/main/java/stirling/software/saas/config/CreditsProperties.java b/app/saas/src/main/java/stirling/software/saas/config/CreditsProperties.java deleted file mode 100644 index 5b8df0d73b..0000000000 --- a/app/saas/src/main/java/stirling/software/saas/config/CreditsProperties.java +++ /dev/null @@ -1,75 +0,0 @@ -package stirling.software.saas.config; - -import java.util.Map; - -import org.springframework.boot.context.properties.ConfigurationProperties; -import org.springframework.context.annotation.Profile; -import org.springframework.stereotype.Component; - -import lombok.Data; - -@Data -@Component -@Profile("saas") -@ConfigurationProperties(prefix = "credits") -public class CreditsProperties { - - /** Whether the credits system is enabled */ - private boolean enabled = true; - - /** Credit allocations per billing cycle (monthly) */ - private CycleAllocations cycle = new CycleAllocations(); - - /** Reset configuration */ - private Reset reset = new Reset(); - - /** Error tracking configuration */ - private Errors errors = new Errors(); - - /** Cache configuration */ - private Cache cache = new Cache(); - - @Data - public static class CycleAllocations { - /** Whether admin role has unlimited credits */ - private boolean adminUnlimited = true; - - /** Credit allocations per billing cycle (monthly) per role */ - private Map allocations = - Map.of( - "ROLE_ADMIN", 1000, - "ROLE_PRO_USER", 500, - "ROLE_USER", 50, - "ROLE_LIMITED_API_USER", 10, - "ROLE_EXTRA_LIMITED_API_USER", 20, - "ROLE_WEB_ONLY_USER", 0, - "ROLE_DEMO_USER", 100); - } - - @Data - public static class Reset { - /** Cron expression for monthly reset (default: 1st of month 02:00 UTC) */ - private String cron = "0 0 2 1 * *"; - - /** Time zone for the reset schedule */ - private String zone = "UTC"; - } - - @Data - public static class Errors { - /** How long error counts are tracked (in minutes) */ - private int ttlMinutes = 60; - - /** Number of free processing errors before charging */ - private int freeProcessingErrors = 2; - } - - @Data - public static class Cache { - /** Enable local Caffeine cache for error counts */ - private boolean localEnabled = true; - - /** Enable Redis cache for multi-instance deployments */ - private boolean redisEnabled = false; - } -} diff --git a/app/saas/src/main/java/stirling/software/saas/controller/CreditController.java b/app/saas/src/main/java/stirling/software/saas/controller/CreditController.java deleted file mode 100644 index 9459c8bbb4..0000000000 --- a/app/saas/src/main/java/stirling/software/saas/controller/CreditController.java +++ /dev/null @@ -1,312 +0,0 @@ -package stirling.software.saas.controller; - -import java.util.Map; - -import org.springframework.context.annotation.Profile; -import org.springframework.http.ResponseEntity; -import org.springframework.security.access.prepost.PreAuthorize; -import org.springframework.security.core.Authentication; -import org.springframework.web.bind.annotation.*; - -import io.swagger.v3.oas.annotations.Hidden; -import io.swagger.v3.oas.annotations.Operation; -import io.swagger.v3.oas.annotations.media.Content; -import io.swagger.v3.oas.annotations.media.Schema; -import io.swagger.v3.oas.annotations.responses.ApiResponse; -import io.swagger.v3.oas.annotations.tags.Tag; - -import lombok.RequiredArgsConstructor; -import lombok.extern.slf4j.Slf4j; - -import stirling.software.proprietary.security.model.ApiKeyAuthenticationToken; -import stirling.software.proprietary.security.model.User; -import stirling.software.saas.security.EnhancedJwtAuthenticationToken; -import stirling.software.saas.service.CreditService; -import stirling.software.saas.service.CreditService.CreditSummary; -import stirling.software.saas.util.LogRedactionUtils; - -@RestController -@Profile("saas") -@RequestMapping("/api/v1/credits") -@Tag(name = "Credit Management", description = "Endpoints for managing user API credits") -@RequiredArgsConstructor -@Slf4j -public class CreditController { - - private final CreditService creditService; - - @GetMapping - @Operation( - summary = "Get user credit information", - description = - "Retrieve current credit balance and usage statistics for the authenticated user") - @ApiResponse( - responseCode = "200", - description = "Credit information retrieved successfully", - content = @Content(schema = @Schema(implementation = CreditSummary.class))) - public ResponseEntity getUserCredits(Authentication authentication) { - return ResponseEntity.ok(getCreditSummaryForAuthentication(authentication)); - } - - @PostMapping("/purchase") - @Hidden - @Operation( - summary = "Purchase additional credits", - description = "Add bought credits to user account (admin only)") - @PreAuthorize("hasRole('ADMIN')") - @ApiResponse(responseCode = "200", description = "Credits purchased successfully") - public ResponseEntity> purchaseCredits( - @RequestParam("username") String username, @RequestParam("credits") int credits) { - - if (credits <= 0) { - return ResponseEntity.badRequest().body(Map.of("error", "Credits must be positive")); - } - - try { - creditService.addBoughtCredits(username, credits); - log.info("Admin added {} credits to user: {}", credits, username); - return ResponseEntity.ok(Map.of("success", true, "creditsAdded", credits)); - } catch (IllegalArgumentException e) { - log.warn("purchaseCredits rejected: {}", e.getMessage()); - return ResponseEntity.badRequest().body(Map.of("error", "Invalid request")); - } - } - - @PostMapping("/purchase-by-supabase-id") - @Hidden - @Operation( - summary = "Purchase additional credits by Supabase ID", - description = "Add bought credits to user account using Supabase ID (admin only)") - @PreAuthorize("hasRole('ADMIN')") - @ApiResponse(responseCode = "200", description = "Credits purchased successfully") - public ResponseEntity> purchaseCreditsBySupabaseId( - @RequestParam("supabaseId") String supabaseId, @RequestParam("credits") int credits) { - - if (credits <= 0) { - return ResponseEntity.badRequest().body(Map.of("error", "Credits must be positive")); - } - - try { - creditService.addBoughtCreditsBySupabaseId(supabaseId, credits); - log.info( - "Admin added {} credits to user with Supabase ID: {}", - credits, - LogRedactionUtils.redactSupabaseId(supabaseId)); - return ResponseEntity.ok(Map.of("success", true, "creditsAdded", credits)); - } catch (IllegalArgumentException e) { - log.warn("purchaseCreditsBySupabaseId rejected: {}", e.getMessage()); - return ResponseEntity.badRequest().body(Map.of("error", "Invalid request")); - } - } - - @GetMapping("/user/{username}") - @Hidden - @Operation( - summary = "Get credit information for specific user", - description = "Retrieve credit information for a specific user (admin only)") - @PreAuthorize("hasRole('ADMIN')") - @ApiResponse( - responseCode = "200", - description = "User credit information retrieved successfully", - content = @Content(schema = @Schema(implementation = CreditSummary.class))) - public ResponseEntity getUserCreditsAdmin( - @PathVariable("username") String username) { - CreditSummary summary = creditService.getCreditSummary(username); - return ResponseEntity.ok(summary); - } - - @GetMapping("/user-by-supabase-id/{supabaseId}") - @Hidden - @Operation( - summary = "Get credit information for specific user by Supabase ID", - description = - "Retrieve credit information for a specific user using Supabase ID (admin only)") - @PreAuthorize("hasRole('ADMIN')") - @ApiResponse( - responseCode = "200", - description = "User credit information retrieved successfully", - content = @Content(schema = @Schema(implementation = CreditSummary.class))) - public ResponseEntity getUserCreditsAdminBySupabaseId( - @PathVariable("supabaseId") String supabaseId) { - CreditSummary summary = creditService.getCreditSummaryBySupabaseId(supabaseId); - return ResponseEntity.ok(summary); - } - - @PostMapping("/reset-cycle") - @Hidden - @Operation( - summary = "Reset cycle credits for all users", - description = "Manually trigger cycle credit reset for all users (admin only)") - @PreAuthorize("hasRole('ADMIN')") - @ApiResponse(responseCode = "200", description = "Cycle credits reset successfully") - public ResponseEntity resetCycleCredits() { - creditService.resetCycleCreditsForAllUsers(); - log.info("Manual cycle credit reset triggered by admin"); - return ResponseEntity.ok("Cycle credits reset successfully for all users"); - } - - @PostMapping("/set-bought-credits") - @Hidden - @Operation( - summary = "Set user's bought credits to a specific amount", - description = - "Hard set the bought credits balance for a specific user to an exact amount (admin only)") - @PreAuthorize("hasRole('ADMIN')") - @ApiResponse(responseCode = "200", description = "Bought credits set successfully") - public ResponseEntity> setBoughtCredits( - @RequestParam("username") String username, @RequestParam("credits") int credits) { - - if (credits < 0) { - return ResponseEntity.badRequest().body(Map.of("error", "Credits cannot be negative")); - } - - try { - creditService.setBoughtCredits(username, credits); - log.info("Admin set bought credits to {} for user: {}", credits, username); - return ResponseEntity.ok(Map.of("success", true, "boughtCredits", credits)); - } catch (IllegalArgumentException e) { - log.warn("setBoughtCredits rejected: {}", e.getMessage()); - return ResponseEntity.badRequest().body(Map.of("error", "Invalid request")); - } - } - - @PostMapping("/set-bought-credits-by-supabase-id") - @Hidden - @Operation( - summary = "Set user's bought credits to a specific amount by Supabase ID", - description = - "Hard set the bought credits balance for a specific user using Supabase ID to an exact amount (admin only)") - @PreAuthorize("hasRole('ADMIN')") - @ApiResponse(responseCode = "200", description = "Bought credits set successfully") - public ResponseEntity> setBoughtCreditsBySupabaseId( - @RequestParam("supabaseId") String supabaseId, @RequestParam("credits") int credits) { - - if (credits < 0) { - return ResponseEntity.badRequest().body(Map.of("error", "Credits cannot be negative")); - } - - try { - creditService.setBoughtCreditsBySupabaseId(supabaseId, credits); - log.info( - "Admin set bought credits to {} for user with Supabase ID: {}", - credits, - LogRedactionUtils.redactSupabaseId(supabaseId)); - return ResponseEntity.ok(Map.of("success", true, "boughtCredits", credits)); - } catch (IllegalArgumentException e) { - log.warn("setBoughtCreditsBySupabaseId rejected: {}", e.getMessage()); - return ResponseEntity.badRequest().body(Map.of("error", "Invalid request")); - } - } - - @PostMapping("/set-cycle-credits") - @Hidden - @Operation( - summary = "Set user's cycle credits remaining to a specific amount", - description = - "Hard set the cycle credits remaining balance for a specific user to an exact amount (admin only)") - @PreAuthorize("hasRole('ADMIN')") - @ApiResponse(responseCode = "200", description = "Cycle credits set successfully") - public ResponseEntity> setCycleCredits( - @RequestParam("username") String username, @RequestParam("credits") int credits) { - - if (credits < 0) { - return ResponseEntity.badRequest().body(Map.of("error", "Credits cannot be negative")); - } - - try { - creditService.setCycleCredits(username, credits); - log.info("Admin set cycle credits to {} for user: {}", credits, username); - return ResponseEntity.ok(Map.of("success", true, "cycleCredits", credits)); - } catch (IllegalArgumentException e) { - log.warn("setCycleCredits rejected: {}", e.getMessage()); - return ResponseEntity.badRequest().body(Map.of("error", "Invalid request")); - } - } - - @PostMapping("/set-cycle-credits-by-supabase-id") - @Hidden - @Operation( - summary = "Set user's cycle credits remaining to a specific amount by Supabase ID", - description = - "Hard set the cycle credits remaining balance for a specific user using Supabase ID to an exact amount (admin only)") - @PreAuthorize("hasRole('ADMIN')") - @ApiResponse(responseCode = "200", description = "Cycle credits set successfully") - public ResponseEntity> setCycleCreditsBySupabaseId( - @RequestParam("supabaseId") String supabaseId, @RequestParam("credits") int credits) { - - if (credits < 0) { - return ResponseEntity.badRequest().body(Map.of("error", "Credits cannot be negative")); - } - - try { - creditService.setCycleCreditsBySupabaseId(supabaseId, credits); - log.info( - "Admin set cycle credits to {} for user with Supabase ID: {}", - credits, - LogRedactionUtils.redactSupabaseId(supabaseId)); - return ResponseEntity.ok(Map.of("success", true, "cycleCredits", credits)); - } catch (IllegalArgumentException e) { - log.warn("setCycleCreditsBySupabaseId rejected: {}", e.getMessage()); - return ResponseEntity.badRequest().body(Map.of("error", "Invalid request")); - } - } - - @GetMapping("/usage") - @Operation( - summary = "Get credit usage summary", - description = "Get overview of credit usage (for authenticated user or admin view)") - public ResponseEntity getCreditUsage(Authentication authentication) { - CreditSummary summary = getCreditSummaryForAuthentication(authentication); - - // For unlimited users, don't show meaningless huge usage numbers - int cycleCreditsUsed = - summary.unlimited - ? 0 - : (summary.cycleCreditsAllocated - summary.cycleCreditsRemaining); - - UsageSummary usage = - new UsageSummary( - cycleCreditsUsed, - summary.totalBoughtCredits - summary.boughtCreditsRemaining, - summary.totalAvailableCredits, - summary.unlimited); - - return ResponseEntity.ok(usage); - } - - /** Resolves the current authentication to a credit summary, handling JWT and API-key auth. */ - private CreditSummary getCreditSummaryForAuthentication(Authentication authentication) { - if (authentication instanceof EnhancedJwtAuthenticationToken enhancedJwt) { - return creditService.getCreditSummaryBySupabaseId(enhancedJwt.getSupabaseId()); - } - if (authentication instanceof ApiKeyAuthenticationToken apiKeyToken) { - String apiKey = (String) apiKeyToken.getCredentials(); - // Principal is the resolved User entity (per SupabaseAuthenticationFilter). Prefer the - // linked Supabase ID; fall back to API-key-keyed credits if there's no supabase link - // or no User row (e.g. legacy API-key-only deployments). - if (apiKeyToken.getPrincipal() instanceof User user && user.getSupabaseId() != null) { - return creditService.getCreditSummaryBySupabaseId(user.getSupabaseId().toString()); - } - return creditService.getCreditSummaryByApiKey(apiKey); - } - return creditService.getCreditSummaryBySupabaseId(authentication.getName()); - } - - public static class UsageSummary { - public final int cycleCreditsUsed; - public final int boughtCreditsUsed; - public final int creditsRemaining; - public final boolean unlimited; - - public UsageSummary( - int cycleCreditsUsed, - int boughtCreditsUsed, - int creditsRemaining, - boolean unlimited) { - this.cycleCreditsUsed = cycleCreditsUsed; - this.boughtCreditsUsed = boughtCreditsUsed; - this.creditsRemaining = creditsRemaining; - this.unlimited = unlimited; - } - } -} diff --git a/app/saas/src/main/java/stirling/software/saas/interceptor/CreditErrorAdvice.java b/app/saas/src/main/java/stirling/software/saas/interceptor/CreditErrorAdvice.java deleted file mode 100644 index 280f8be6df..0000000000 --- a/app/saas/src/main/java/stirling/software/saas/interceptor/CreditErrorAdvice.java +++ /dev/null @@ -1,305 +0,0 @@ -package stirling.software.saas.interceptor; - -import java.util.HashMap; -import java.util.Map; -import java.util.Optional; - -import org.springframework.context.annotation.Profile; -import org.springframework.core.annotation.Order; -import org.springframework.http.HttpStatus; -import org.springframework.http.MediaType; -import org.springframework.http.ResponseEntity; -import org.springframework.security.core.Authentication; -import org.springframework.security.core.context.SecurityContextHolder; -import org.springframework.web.bind.annotation.ExceptionHandler; -import org.springframework.web.bind.annotation.RestControllerAdvice; - -import com.fasterxml.jackson.core.JsonProcessingException; -import com.fasterxml.jackson.databind.ObjectMapper; - -import io.micrometer.core.instrument.Counter; -import io.micrometer.core.instrument.MeterRegistry; - -import jakarta.servlet.http.HttpServletRequest; - -import lombok.extern.slf4j.Slf4j; - -import stirling.software.common.annotations.AutoJobPostMapping; -import stirling.software.proprietary.security.database.repository.UserRepository; -import stirling.software.proprietary.security.model.User; -import stirling.software.saas.model.CreditConsumptionResult; -import stirling.software.saas.service.CreditService; -import stirling.software.saas.service.ErrorTrackingService; -import stirling.software.saas.service.SaasTeamExtensionService; -import stirling.software.saas.service.TeamCreditService; -import stirling.software.saas.util.AuthenticationUtils; -import stirling.software.saas.util.CreditHeaderUtils; - -/** - * Scoped to controllers annotated with {@link AutoJobPostMapping} so it doesn't hijack the global - * exception flow. - */ -@RestControllerAdvice(annotations = AutoJobPostMapping.class) -@Profile("saas") -@Slf4j -@Order(1) -public class CreditErrorAdvice { - - private static final String ATTR_ELIGIBLE = "CREDIT_ELIGIBLE"; - private static final String ATTR_APIKEY = "CREDIT_API_KEY"; - private static final String ATTR_CHARGED = "CREDIT_CHARGED"; - private static final String ATTR_RESOURCE_WEIGHT = "CREDIT_RESOURCE_WEIGHT"; - - private final CreditService creditService; - private final TeamCreditService teamCreditService; - private final UserRepository userRepository; - private final ErrorTrackingService errorTrackingService; - private final SaasTeamExtensionService saasTeamExtensionService; - private final CreditHeaderUtils creditHeaderUtils; - private final Counter creditsConsumedCounter; - // Inlined: Stirling's parent build uses Jackson 3 (tools.jackson), no Jackson 2 ObjectMapper - // bean in the context. Stateless usage, so a fresh instance is fine. - private final ObjectMapper objectMapper = new ObjectMapper(); - - public CreditErrorAdvice( - CreditService creditService, - TeamCreditService teamCreditService, - UserRepository userRepository, - ErrorTrackingService errorTrackingService, - SaasTeamExtensionService saasTeamExtensionService, - CreditHeaderUtils creditHeaderUtils, - MeterRegistry meterRegistry) { - this.creditService = creditService; - this.teamCreditService = teamCreditService; - this.userRepository = userRepository; - this.errorTrackingService = errorTrackingService; - this.saasTeamExtensionService = saasTeamExtensionService; - this.creditHeaderUtils = creditHeaderUtils; - this.creditsConsumedCounter = - Counter.builder("credits.consumed") - .description("Number of credits actually consumed") - .tag("source", "error") - .register(meterRegistry); - } - - @ExceptionHandler(Throwable.class) - public ResponseEntity handleThrowable(HttpServletRequest request, Throwable ex) { - HttpStatus status = determineHttpStatus(ex); - log.debug( - "[CREDIT-DEBUG] CreditErrorAdvice: Handling exception: {} -> {}", - ex.getClass().getSimpleName(), - status); - - String message = Optional.ofNullable(ex.getMessage()).orElse("An error occurred"); - // Build error body - Map body = new HashMap<>(); - body.put("error", ex.getClass().getSimpleName()); - body.put("message", message); - body.put("status", status.value()); - - var builder = ResponseEntity.status(status); - - // Handle credit consumption for errors - if (Boolean.TRUE.equals(request.getAttribute(ATTR_ELIGIBLE)) - && request.getAttribute(ATTR_CHARGED) == null) { - - var apiKey = (String) request.getAttribute(ATTR_APIKEY); - var resourceWeight = (Integer) request.getAttribute(ATTR_RESOURCE_WEIGHT); - var isApiRequest = (Boolean) request.getAttribute("IS_API_REQUEST"); - int creditAmount = resourceWeight != null ? resourceWeight : 1; - - String identifierForErrorTracking = - apiKey; // Keep using apiKey/username for error tracking - if (apiKey != null - && errorTrackingService.recordErrorAndShouldConsumeCredit( - identifierForErrorTracking, - request.getRequestURI(), - ex, - status.value())) { - - // Get current user - Authentication auth = SecurityContextHolder.getContext().getAuthentication(); - User user = null; - try { - user = AuthenticationUtils.getCurrentUser(auth, userRepository); - } catch (Exception e) { - log.warn( - "[CREDIT-DEBUG] CreditErrorAdvice: Could not get user for team check: {}", - e.getMessage()); - } - - if (user == null) { - log.error( - "[CREDIT-DEBUG] CreditErrorAdvice: Unable to resolve user - skipping credit consumption"); - } else { - // Check if user is in a non-personal team (must match UnifiedCreditInterceptor - // logic) - Long targetTeamId = null; - if (user.getTeam() != null - && !saasTeamExtensionService.isPersonal(user.getTeam())) { - targetTeamId = user.getTeam().getId(); - } - - boolean consumed = false; - String creditSource = null; - - if (targetTeamId != null) { - // User is in a non-personal team - consume from team credit pool - consumed = teamCreditService.consumeCredit(targetTeamId, creditAmount); - creditSource = "TEAM_CREDITS"; - log.debug( - "[CREDIT-DEBUG] CreditErrorAdvice: Consumed {} credits from team {}", - creditAmount, - targetTeamId); - } else { - // No team - use waterfall logic for individual credits - boolean isApiRequestFlag = Boolean.TRUE.equals(isApiRequest); - CreditConsumptionResult result = - creditService.consumeCreditWithWaterfall( - user, creditAmount, isApiRequestFlag); - consumed = result.isSuccess(); - creditSource = result.getSource(); - - if (!consumed) { - log.error( - "[CREDIT-DEBUG] CreditErrorAdvice: Credit consumption failed for user: {} - {}", - user.getUsername(), - result.getMessage()); - } - } - - if (consumed) { - request.setAttribute(ATTR_CHARGED, Boolean.TRUE); - creditsConsumedCounter.increment(); - - // Set remaining credits header - int remainingCredits = - creditHeaderUtils.getRemainingCredits( - user, creditService, teamCreditService); - if (remainingCredits >= 0) { - builder.header( - "X-Credits-Remaining", Integer.toString(remainingCredits)); - log.warn( - "[CREDIT-HEADER] Added X-Credits-Remaining header: {}", - remainingCredits); - } - if (creditSource != null) { - builder.header("X-Credit-Source", creditSource); - } - - log.info( - "[CREDIT-DEBUG] CreditErrorAdvice: {} credits consumed from {} for user: {} (error case)", - creditAmount, - creditSource, - user.getUsername()); - } - } - } else { - log.debug( - "[CREDIT-DEBUG] CreditErrorAdvice: ErrorTrackingService says do NOT consume credit for this error"); - } - } else if (request.getAttribute(ATTR_CHARGED) != null) { - // Already charged, set header if user is authenticated - Authentication auth = SecurityContextHolder.getContext().getAuthentication(); - if (auth != null && auth.isAuthenticated()) { - try { - User user = AuthenticationUtils.getCurrentUser(auth, userRepository); - int remainingCredits = - creditHeaderUtils.getRemainingCredits( - user, creditService, teamCreditService); - if (remainingCredits >= 0) { - builder.header("X-Credits-Remaining", Integer.toString(remainingCredits)); - log.warn( - "[CREDIT-HEADER] Added X-Credits-Remaining header: {}", - remainingCredits); - } - } catch (Exception e) { - log.debug( - "[CREDIT-HEADER] Could not add credits header for already charged error: {}", - e.getMessage()); - } - } - log.debug("[CREDIT-DEBUG] CreditErrorAdvice: Header set for already charged error"); - } - - if (isSseRequest(request)) { - String payload = toJsonPayload(body); - String sseBody = "event: error\ndata: " + payload + "\n\n"; - return builder.contentType(MediaType.TEXT_EVENT_STREAM).body(sseBody); - } - - return builder.body(body); - } - - private String maskApiKey(String apiKey) { - if (apiKey == null || apiKey.length() < 8) { - return "***"; - } - return apiKey.substring(0, 4) + "***" + apiKey.substring(apiKey.length() - 4); - } - - private HttpStatus determineHttpStatus(Throwable throwable) { - // Map common exceptions to HTTP status codes - String exceptionClass = throwable.getClass().getSimpleName(); - switch (exceptionClass) { - case "IllegalArgumentException": - case "ValidationException": - case "MethodArgumentNotValidException": - return HttpStatus.BAD_REQUEST; - case "AccessDeniedException": - return HttpStatus.FORBIDDEN; - case "UsernameNotFoundException": - return HttpStatus.UNAUTHORIZED; - case "HttpMessageNotReadableException": - return HttpStatus.BAD_REQUEST; - case "MaxUploadSizeExceededException": - return HttpStatus.PAYLOAD_TOO_LARGE; - case "UnsupportedOperationException": - return HttpStatus.NOT_IMPLEMENTED; - default: - // Check error message for clues - String message = throwable.getMessage(); - if (message != null) { - if (message.toLowerCase().contains("validation") - || message.toLowerCase().contains("invalid parameter")) { - return HttpStatus.BAD_REQUEST; - } - if (message.toLowerCase().contains("not found")) { - return HttpStatus.NOT_FOUND; - } - } - return HttpStatus.INTERNAL_SERVER_ERROR; - } - } - - private boolean isSseRequest(HttpServletRequest request) { - String accept = request.getHeader("Accept"); - if (accept != null && accept.contains(MediaType.TEXT_EVENT_STREAM_VALUE)) { - return true; - } - String contentType = request.getContentType(); - return contentType != null && contentType.contains(MediaType.TEXT_EVENT_STREAM_VALUE); - } - - private String toJsonPayload(Map payload) { - try { - return objectMapper.writeValueAsString(payload); - } catch (JsonProcessingException exc) { - log.warn("Failed to serialize SSE error payload, falling back to string", exc); - String message = payload.getOrDefault("message", "An error occurred").toString(); - return "{\"error\":\"Error\",\"message\":\"" + message + "\",\"status\":500}"; - } - } - - public static class ErrorResponse { - public final String error; - public final String message; - public final int status; - - public ErrorResponse(String error, String message, int status) { - this.error = error; - this.message = message; - this.status = status; - } - } -} diff --git a/app/saas/src/main/java/stirling/software/saas/interceptor/CreditSuccessAdvice.java b/app/saas/src/main/java/stirling/software/saas/interceptor/CreditSuccessAdvice.java deleted file mode 100644 index 51e82b4db6..0000000000 --- a/app/saas/src/main/java/stirling/software/saas/interceptor/CreditSuccessAdvice.java +++ /dev/null @@ -1,226 +0,0 @@ -package stirling.software.saas.interceptor; - -import org.springframework.context.annotation.Profile; -import org.springframework.core.MethodParameter; -import org.springframework.http.MediaType; -import org.springframework.http.converter.HttpMessageConverter; -import org.springframework.http.server.ServerHttpRequest; -import org.springframework.http.server.ServerHttpResponse; -import org.springframework.http.server.ServletServerHttpRequest; -import org.springframework.http.server.ServletServerHttpResponse; -import org.springframework.security.core.Authentication; -import org.springframework.security.core.context.SecurityContextHolder; -import org.springframework.web.bind.annotation.RestControllerAdvice; -import org.springframework.web.servlet.mvc.method.annotation.ResponseBodyAdvice; - -import io.micrometer.core.instrument.Counter; -import io.micrometer.core.instrument.MeterRegistry; - -import lombok.extern.slf4j.Slf4j; - -import stirling.software.proprietary.security.database.repository.UserRepository; -import stirling.software.proprietary.security.model.User; -import stirling.software.saas.model.CreditConsumptionResult; -import stirling.software.saas.service.CreditService; -import stirling.software.saas.service.SaasTeamExtensionService; -import stirling.software.saas.service.TeamCreditService; -import stirling.software.saas.util.AuthenticationUtils; -import stirling.software.saas.util.CreditHeaderUtils; - -@RestControllerAdvice -@Profile("saas") -@Slf4j -public class CreditSuccessAdvice implements ResponseBodyAdvice { - - private static final String ATTR_ELIGIBLE = "CREDIT_ELIGIBLE"; - private static final String ATTR_APIKEY = "CREDIT_API_KEY"; - private static final String ATTR_CHARGED = "CREDIT_CHARGED"; - private static final String ATTR_RESOURCE_WEIGHT = "CREDIT_RESOURCE_WEIGHT"; - - private final CreditService creditService; - private final TeamCreditService teamCreditService; - private final UserRepository userRepository; - private final SaasTeamExtensionService saasTeamExtensionService; - private final CreditHeaderUtils creditHeaderUtils; - private final Counter creditsConsumedCounter; - - public CreditSuccessAdvice( - CreditService creditService, - TeamCreditService teamCreditService, - UserRepository userRepository, - SaasTeamExtensionService saasTeamExtensionService, - CreditHeaderUtils creditHeaderUtils, - MeterRegistry meterRegistry) { - this.creditService = creditService; - this.teamCreditService = teamCreditService; - this.userRepository = userRepository; - this.saasTeamExtensionService = saasTeamExtensionService; - this.creditHeaderUtils = creditHeaderUtils; - this.creditsConsumedCounter = - Counter.builder("credits.consumed") - .description("Number of credits actually consumed") - .tag("source", "success") - .register(meterRegistry); - } - - @Override - public boolean supports( - MethodParameter returnType, Class> converterType) { - // Only REST bodies; this covers @ResponseBody and ResponseEntity - return true; - } - - @Override - public Object beforeBodyWrite( - Object body, - MethodParameter returnType, - MediaType selectedContentType, - Class> selectedConverterType, - ServerHttpRequest request, - ServerHttpResponse response) { - - if (!(request instanceof ServletServerHttpRequest)) { - return body; - } - - var servletReq = ((ServletServerHttpRequest) request).getServletRequest(); - if (!Boolean.TRUE.equals(servletReq.getAttribute(ATTR_ELIGIBLE))) { - return body; - } - - if (servletReq.getAttribute(ATTR_CHARGED) != null) { - return body; - } - - // If the handler returned an error ResponseEntity (>=400) without throwing, - // don't spend here; the error advice will decide. - int status = 200; - if (response instanceof ServletServerHttpResponse) { - status = ((ServletServerHttpResponse) response).getServletResponse().getStatus(); - } - if (status >= 400) { - log.debug( - "[CREDIT-DEBUG] CreditSuccessAdvice: Error status {} detected, skipping credit consumption", - status); - return body; - } - - var apiKey = (String) servletReq.getAttribute(ATTR_APIKEY); - var resourceWeight = (Integer) servletReq.getAttribute(ATTR_RESOURCE_WEIGHT); - var isApiRequest = (Boolean) servletReq.getAttribute("IS_API_REQUEST"); - int creditAmount = resourceWeight != null ? resourceWeight : 1; - - if (apiKey != null) { - // Get current user - Authentication auth = SecurityContextHolder.getContext().getAuthentication(); - User user = null; - try { - user = AuthenticationUtils.getCurrentUser(auth, userRepository); - } catch (Exception e) { - log.warn( - "[CREDIT-DEBUG] CreditSuccessAdvice: Could not get user for team check: {}", - e.getMessage()); - } - - if (user == null) { - log.error( - "[CREDIT-DEBUG] CreditSuccessAdvice: Unable to resolve user - skipping credit consumption"); - return body; - } - - // Check if user is in a non-personal team (must match UnifiedCreditInterceptor logic) - // IMPORTANT: Limited API users (anonymous, extra limited) always use personal credits, - // never team credits - boolean isLimitedApiUser = - auth.getAuthorities().stream() - .anyMatch( - authority -> - "ROLE_LIMITED_API_USER".equals(authority.getAuthority()) - || "ROLE_EXTRA_LIMITED_API_USER" - .equals(authority.getAuthority())); - Long targetTeamId = null; - if (!isLimitedApiUser - && user.getTeam() != null - && !saasTeamExtensionService.isPersonal(user.getTeam())) { - targetTeamId = user.getTeam().getId(); - } - - final boolean consumed; - final String creditSource; - - if (targetTeamId != null) { - // User is in a non-personal team - use waterfall with leader overage - CreditConsumptionResult result = - teamCreditService.consumeCreditWithWaterfall(targetTeamId, creditAmount); - consumed = result.isSuccess(); - creditSource = result.getSource(); - - if (!consumed) { - log.error( - "[CREDIT-DEBUG] CreditSuccessAdvice: Team credit consumption failed:" - + " {}", - result.getMessage()); - } else { - log.debug( - "[CREDIT-DEBUG] CreditSuccessAdvice: Consumed {} credits from team {}" - + " via {}", - creditAmount, - targetTeamId, - creditSource); - } - } else { - // No team - use waterfall logic for individual credits - boolean isApiRequestFlag = Boolean.TRUE.equals(isApiRequest); - CreditConsumptionResult result = - creditService.consumeCreditWithWaterfall( - user, creditAmount, isApiRequestFlag); - consumed = result.isSuccess(); - creditSource = result.getSource(); - - if (!consumed) { - log.error( - "[CREDIT-DEBUG] CreditSuccessAdvice: Credit consumption failed for user: {} - {}", - user.getUsername(), - result.getMessage()); - } - } - - if (consumed) { - servletReq.setAttribute(ATTR_CHARGED, Boolean.TRUE); - creditsConsumedCounter.increment(); - - // Set remaining credits header - int remainingCredits = - creditHeaderUtils.getRemainingCredits( - user, creditService, teamCreditService); - if (remainingCredits >= 0) { - response.getHeaders() - .set("X-Credits-Remaining", Integer.toString(remainingCredits)); - log.warn( - "[CREDIT-HEADER] Added X-Credits-Remaining header: {}", - remainingCredits); - } - if (creditSource != null) { - response.getHeaders().set("X-Credit-Source", creditSource); - } - - log.info( - "[CREDIT-DEBUG] CreditSuccessAdvice: {} credits consumed from {} for user: {}", - creditAmount, - creditSource, - user.getUsername()); - } - } else { - log.warn("[CREDIT-DEBUG] CreditSuccessAdvice: No apiKey attribute found"); - } - - return body; - } - - private String maskApiKey(String apiKey) { - if (apiKey == null || apiKey.length() < 8) { - return "***"; - } - return apiKey.substring(0, 4) + "***" + apiKey.substring(apiKey.length() - 4); - } -} diff --git a/app/saas/src/main/java/stirling/software/saas/interceptor/UnifiedCreditInterceptor.java b/app/saas/src/main/java/stirling/software/saas/interceptor/UnifiedCreditInterceptor.java deleted file mode 100644 index 631217e1c0..0000000000 --- a/app/saas/src/main/java/stirling/software/saas/interceptor/UnifiedCreditInterceptor.java +++ /dev/null @@ -1,491 +0,0 @@ -package stirling.software.saas.interceptor; - -import org.springframework.context.annotation.Profile; -import org.springframework.security.core.Authentication; -import org.springframework.security.core.context.SecurityContextHolder; -import org.springframework.stereotype.Component; -import org.springframework.web.method.HandlerMethod; -import org.springframework.web.servlet.AsyncHandlerInterceptor; -import org.springframework.web.servlet.ModelAndView; - -import io.micrometer.core.instrument.Counter; -import io.micrometer.core.instrument.MeterRegistry; -import io.micrometer.core.instrument.Timer; - -import jakarta.servlet.http.HttpServletRequest; -import jakarta.servlet.http.HttpServletResponse; - -import lombok.extern.slf4j.Slf4j; - -import stirling.software.common.annotations.AutoJobPostMapping; -import stirling.software.proprietary.security.database.repository.UserRepository; -import stirling.software.proprietary.security.model.ApiKeyAuthenticationToken; -import stirling.software.proprietary.security.model.User; -import stirling.software.saas.config.CreditsProperties; -import stirling.software.saas.model.TeamCredit; -import stirling.software.saas.model.UserCredit; -import stirling.software.saas.repository.TeamMembershipRepository; -import stirling.software.saas.service.CreditService; -import stirling.software.saas.service.ErrorTrackingService; -import stirling.software.saas.service.SaasTeamExtensionService; -import stirling.software.saas.service.SaasUserExtensionService; -import stirling.software.saas.service.TeamCreditService; -import stirling.software.saas.util.AuthenticationUtils; - -@Component -@Profile("saas") -@Slf4j -public class UnifiedCreditInterceptor implements AsyncHandlerInterceptor { - - private final CreditService creditService; - private final ErrorTrackingService errorTrackingService; - private final CreditsProperties creditsProperties; - private final UserRepository userRepository; - private final TeamCreditService teamCreditService; - private final TeamMembershipRepository membershipRepository; - private final SaasUserExtensionService saasUserExtensionService; - private final SaasTeamExtensionService saasTeamExtensionService; - - private final Counter creditsCheckedCounter; - private final Counter creditsRejectedCounter; - private final Counter jwtBypassCounter; - private final Timer creditCheckTimer; - - private static final String ATTR_CREDIT_ELIGIBLE = "CREDIT_ELIGIBLE"; - private static final String ATTR_API_KEY = "CREDIT_API_KEY"; - private static final String ATTR_RESOURCE_WEIGHT = "CREDIT_RESOURCE_WEIGHT"; - private static final String ATTR_CHARGED = "CREDIT_CHARGED"; - - public UnifiedCreditInterceptor( - CreditService creditService, - ErrorTrackingService errorTrackingService, - CreditsProperties creditsProperties, - UserRepository userRepository, - TeamCreditService teamCreditService, - TeamMembershipRepository membershipRepository, - SaasUserExtensionService saasUserExtensionService, - SaasTeamExtensionService saasTeamExtensionService, - MeterRegistry meterRegistry) { - this.creditService = creditService; - this.errorTrackingService = errorTrackingService; - this.creditsProperties = creditsProperties; - this.userRepository = userRepository; - this.teamCreditService = teamCreditService; - this.membershipRepository = membershipRepository; - this.saasUserExtensionService = saasUserExtensionService; - this.saasTeamExtensionService = saasTeamExtensionService; - - this.creditsCheckedCounter = - Counter.builder("credits.validation.checked") - .description("Number of requests that had credit validation performed") - .register(meterRegistry); - this.creditsRejectedCounter = - Counter.builder("credits.validation.rejected") - .description("Number of requests rejected due to insufficient credits") - .register(meterRegistry); - this.jwtBypassCounter = - Counter.builder("credits.validation.jwt_bypass") - .description("Number of JWT requests that bypassed credit validation") - .register(meterRegistry); - this.creditCheckTimer = - Timer.builder("credits.validation.duration") - .description("Time taken to validate credits") - .register(meterRegistry); - } - - @Override - public boolean preHandle( - HttpServletRequest request, HttpServletResponse response, Object handler) - throws Exception { - - log.debug( - "[CREDIT-DEBUG] UnifiedCreditInterceptor.preHandle() - handler: {}", - handler.getClass().getSimpleName()); - - // Credits system disabled - allow all requests - if (!creditsProperties.isEnabled()) { - log.debug("[CREDIT-DEBUG] Credits system disabled - allowing request"); - return true; - } - - // Only apply to @AutoJobPostMapping endpoints and extract resource weight - if (!(handler instanceof HandlerMethod hm) - || !hm.getMethod().isAnnotationPresent(AutoJobPostMapping.class)) { - log.debug( - "[CREDIT-DEBUG] Handler not eligible for credit validation (no @AutoJobPostMapping)"); - return true; - } - - Authentication auth = SecurityContextHolder.getContext().getAuthentication(); - - log.debug( - "[CREDIT-DEBUG] Authentication: {}", - auth != null ? auth.getClass().getSimpleName() : "null"); - - User currentUser = null; - - // API key authentication always needs credit validation - if (auth instanceof ApiKeyAuthenticationToken) { - // API key users - proceed with normal credit validation - currentUser = (User) auth.getPrincipal(); - } else if (auth != null && auth.isAuthenticated()) { - // JWT users - get user from authentication details - // JwtAuthenticationToken.getPrincipal() might not be a User object - // so we need to look up the user by the Supabase ID from auth.getName() - - String supabaseId = AuthenticationUtils.extractSupabaseId(auth); - log.debug("[CREDIT-DEBUG] JWT authentication detected, Supabase ID: {}", supabaseId); - - // Look up the User object that should exist (authentication succeeded) - try { - java.util.UUID supabaseUuid = java.util.UUID.fromString(supabaseId); - java.util.Optional userOpt = userRepository.findBySupabaseId(supabaseUuid); - if (userOpt.isEmpty()) { - log.error( - "[CREDIT-DEBUG] JWT authenticated but no User found for Supabase ID: {}", - supabaseId); - response.setStatus(500); - response.setContentType("application/json"); - response.getWriter() - .write( - "{\"error\":\"USER_NOT_FOUND\",\"message\":\"Authenticated user not found in database\",\"status\":500}"); - return false; - } - currentUser = userOpt.get(); - - if (shouldApplyCreditsToJwtUser(currentUser)) { - // Anonymous users or other limited JWT users should consume credits - log.debug( - "[CREDIT-DEBUG] JWT user {} subject to credit validation due to limited role", - currentUser.getUsername()); - } else { - jwtBypassCounter.increment(); - log.debug( - "[CREDIT-DEBUG] JWT user {} bypassing credit validation (unlimited role)", - currentUser.getUsername()); - return true; - } - } catch (IllegalArgumentException e) { - log.error("[CREDIT-DEBUG] Invalid Supabase ID format: {}", supabaseId); - response.setStatus(400); - response.setContentType("application/json"); - response.getWriter() - .write( - "{\"error\":\"INVALID_USER_ID\",\"message\":\"Invalid user identifier format\",\"status\":400}"); - return false; - } - } else { - // SECURITY: Block all non-authenticated requests - log.warn( - "[CREDIT-DEBUG] Non-authenticated request blocked - authentication required for credit-controlled endpoints"); - response.setStatus(401); // 401 Unauthorized - response.setContentType("application/json"); - response.getWriter() - .write( - "{\"error\":\"AUTHENTICATION_REQUIRED\",\"message\":\"Authentication required to access this endpoint\",\"status\":401}"); - return false; - } - - // Extract resource weight from annotation - AutoJobPostMapping annotation = hm.getMethod().getAnnotation(AutoJobPostMapping.class); - int resourceWeight = - Math.max(1, Math.min(100, annotation.resourceWeight())); // Clamp to 1-100 - - String apiKey = getApiKeyForUser(auth, currentUser); - String maskedApiKey = maskApiKey(apiKey); - - log.debug( - "[CREDIT-DEBUG] Credit validation for user: {}, API key: {}, resource weight: {}", - currentUser.getUsername(), - maskedApiKey, - resourceWeight); - - // Track that we're performing credit validation - creditsCheckedCounter.increment(); - - // Check if user has SUFFICIENT credits for this operation (with timing) - Timer.Sample sample = Timer.start(); - boolean hasSufficientCredits; - int availableCredits = 0; - - // Check if user is a limited API user (anonymous, extra limited) - // Limited API users always use personal credits, never team credits - boolean isLimitedApiUser = - currentUser.getAuthorities().stream() - .anyMatch( - authority -> - "ROLE_LIMITED_API_USER".equals(authority.getAuthority()) - || "ROLE_EXTRA_LIMITED_API_USER" - .equals(authority.getAuthority())); - - if (auth instanceof ApiKeyAuthenticationToken) { - // API key auth - get credit balance - java.util.Optional userCreditsOpt = - creditService.getUserCreditsByApiKey(apiKey); - availableCredits = userCreditsOpt.map(UserCredit::getTotalAvailableCredits).orElse(0); - hasSufficientCredits = availableCredits >= resourceWeight; - } else { - // JWT user - check team credits if user is in a non-personal team, otherwise personal - // credits - Long teamId = null; - if (!isLimitedApiUser - && currentUser.getTeam() != null - && !saasTeamExtensionService.isPersonal(currentUser.getTeam())) { - teamId = currentUser.getTeam().getId(); - } - - if (teamId != null) { - // User is in a non-personal team - check team credits + leader overage billing - java.util.Optional teamCredits = - teamCreditService.getTeamCredits(teamId); - availableCredits = teamCredits.map(TeamCredit::getTotalAvailableCredits).orElse(0); - - // Check if sufficient credits OR team leader has metered billing - boolean hasTeamCredits = availableCredits >= resourceWeight; - boolean leaderHasMetered = checkTeamLeaderMeteredBilling(currentUser.getTeam()); - - hasSufficientCredits = hasTeamCredits || leaderHasMetered; - - log.debug( - "[CREDIT-DEBUG] Checking team {} credits for user {}: available={}" - + " required={} hasCredits={} leaderMetered={} sufficient={}", - teamId, - currentUser.getUsername(), - availableCredits, - resourceWeight, - hasTeamCredits, - leaderHasMetered, - hasSufficientCredits); - } else { - // Personal team or no team - check personal credits - UserCredit userCredits = creditService.getOrCreateUserCredits(currentUser); - availableCredits = userCredits.getTotalAvailableCredits(); - hasSufficientCredits = availableCredits >= resourceWeight; - log.debug( - "[CREDIT-DEBUG] Checking personal credits for user {}: available={} required={} sufficient={}", - currentUser.getUsername(), - availableCredits, - resourceWeight, - hasSufficientCredits); - } - } - sample.stop(creditCheckTimer); - - // Check if user has metered billing enabled (they can use overage credits even with - // insufficient free credits) - boolean hasMeteredBilling = saasUserExtensionService.isMeteredBillingEnabled(currentUser); - - if (!hasSufficientCredits && !hasMeteredBilling) { - creditsRejectedCounter.increment(); - - // Enhanced message for team members - // Note: Limited API users always use personal credits, so they get personal message - String message; - if (!isLimitedApiUser - && currentUser.getTeam() != null - && !saasTeamExtensionService.isPersonal(currentUser.getTeam())) { - message = - "Insufficient team credits. Team leader must enable overage billing for" - + " uninterrupted service."; - } else { - message = - "Insufficient API credits. Please purchase more credits or wait for your" - + " monthly cycle credits to reset."; - } - - log.warn( - "[CREDIT-DEBUG] Credit validation rejected - Method: {}, URI: {}, IP: {}," - + " User-Agent: {}, User: {}, Supabase ID: {}, Reason: {}", - request.getMethod(), - request.getRequestURI(), - getClientIpAddress(request), - request.getHeader("User-Agent"), - currentUser.getUsername(), - currentUser.getSupabaseId(), - message); - - response.setStatus(429); // 429 Too Many Requests - response.setContentType("application/json"); - response.getWriter() - .write( - String.format( - "{\"error\":\"INSUFFICIENT_CREDITS\",\"message\":\"%s\",\"status\":429}", - message)); - response.getWriter().flush(); - return false; - } - - // Log when metered billing users are using overage credits - if (!hasSufficientCredits && hasMeteredBilling) { - log.info( - "[CREDIT-DEBUG] Metered billing user {} proceeding with insufficient free credits (have: {}, need: {}) - will use overage credits (billed monthly)", - currentUser.getUsername(), - availableCredits, - resourceWeight); - } - - // Mark request as eligible for credit consumption - request.setAttribute(ATTR_CREDIT_ELIGIBLE, Boolean.TRUE); - request.setAttribute(ATTR_API_KEY, apiKey); - request.setAttribute(ATTR_RESOURCE_WEIGHT, resourceWeight); - - // Store whether this is API key or JWT authentication for advice classes - boolean isApiKeyAuth = auth instanceof ApiKeyAuthenticationToken; - request.setAttribute("IS_API_KEY_AUTH", isApiKeyAuth); - - // Store IS_API_REQUEST for waterfall logic (API key requests always consume credits) - request.setAttribute("IS_API_REQUEST", isApiKeyAuth); - - log.debug( - "[CREDIT-DEBUG] Credit validation passed - request marked as eligible for consumption (will consume after success/error)"); - - return true; - } - - @Override - public void postHandle( - HttpServletRequest request, - HttpServletResponse response, - Object handler, - ModelAndView modelAndView) - throws Exception { - // Success path now handled by CreditSuccessAdvice - no spending in postHandle anymore - log.debug("[CREDIT-DEBUG] postHandle: Success path will be handled by CreditSuccessAdvice"); - } - - @Override - public void afterCompletion( - HttpServletRequest request, HttpServletResponse response, Object handler, Exception ex) - throws Exception { - // Error path now handled by CreditErrorAdvice - no spending in afterCompletion anymore - if (ex != null) { - log.debug( - "[CREDIT-DEBUG] afterCompletion: Error path will be handled by CreditErrorAdvice: {}", - ex.getClass().getSimpleName()); - } else { - log.debug( - "[CREDIT-DEBUG] afterCompletion: Success path already handled by CreditSuccessAdvice"); - } - } - - @Override - public void afterConcurrentHandlingStarted( - HttpServletRequest request, HttpServletResponse response, Object handler) - throws Exception { - // For async requests (Callable, DeferredResult, etc.), prevent duplicate processing - // The actual postHandle/afterCompletion will be called when async processing completes - log.debug( - "[CREDIT-DEBUG] afterConcurrentHandlingStarted: Async processing started - skipping interceptor logic"); - } - - private String maskApiKey(String apiKey) { - if (apiKey == null || apiKey.length() < 8) { - return "***"; - } - return apiKey.substring(0, 4) + "***" + apiKey.substring(apiKey.length() - 4); - } - - private String getClientIpAddress(HttpServletRequest request) { - String xForwardedFor = request.getHeader("X-Forwarded-For"); - if (xForwardedFor != null && !xForwardedFor.isEmpty()) { - return xForwardedFor.split(",")[0].trim(); - } - - String xRealIp = request.getHeader("X-Real-IP"); - if (xRealIp != null && !xRealIp.isEmpty()) { - return xRealIp; - } - - return request.getRemoteAddr(); - } - - /** - * Determines if credit limits should apply to a JWT user. - * - *

Rules: - * - *

    - *
  • Metered billing users: always consume (free tier first, then report overage to Stripe) - *
  • Anonymous users: consume credits (web/API) - *
  • Regular users: consume credits (web/API) - *
  • Pro users: unlimited on web UI (waterfall logic handles this), but subject to checks - *
  • API users: always consume credits - *
  • Internal API users: unlimited everywhere - *
  • Admin users: unlimited everywhere - *
- */ - private boolean shouldApplyCreditsToJwtUser(User user) { - String roles = user.getRolesAsString(); - - // Internal API users are unlimited everywhere (for backend internal operations) - if (roles.contains("STIRLING-PDF-BACKEND-API-USER")) { - log.debug("[CREDIT-DEBUG] Internal API user {} - unlimited usage", user.getUsername()); - return false; - } - - // Pro users: Let them through to waterfall logic - // (Pro gets unlimited UI but API still consumes credits) - if (roles.contains("ROLE_PRO_USER")) { - log.debug( - "[CREDIT-DEBUG] Pro user {} - will be handled by waterfall logic", - user.getUsername()); - return true; // Changed from false - let waterfall handle Pro exemption - } - - // Admin users are unlimited everywhere - if (roles.contains("ROLE_ADMIN")) { - log.debug("[CREDIT-DEBUG] Admin user {} - unlimited usage", user.getUsername()); - return false; - } - - // All other users (anonymous, regular, limited API users, metered billing) consume credits - log.debug( - "[CREDIT-DEBUG] User {} with roles {} - subject to credit limits", - user.getUsername(), - roles); - return true; - } - - /** - * Gets the identifier for credit consumption. For API key users, use their actual API key. For - * JWT users, use the Supabase ID as identifier (auth.getName() returns Supabase ID). - */ - private String getApiKeyForUser(Authentication auth, User user) { - if (auth instanceof ApiKeyAuthenticationToken) { - return user.getApiKey(); - } else { - // For JWT users, return Supabase ID as the credit consumption identifier - return AuthenticationUtils.extractSupabaseId(auth); - } - } - - /** - * Check if team leader has metered billing enabled. This allows teams to use overage billing - * when team credits are exhausted. - * - * @param team the team to check - * @return true if team leader has metered billing enabled - */ - private boolean checkTeamLeaderMeteredBilling(stirling.software.proprietary.model.Team team) { - if (team == null || team.getId() == null) { - return false; - } - - try { - java.util.List leaders = - membershipRepository.findByTeamIdAndRole( - team.getId(), - stirling.software.common.model.enumeration.TeamRole.LEADER); - - if (leaders.isEmpty()) { - return false; - } - - User leader = leaders.get(0).getUser(); - return saasUserExtensionService.isMeteredBillingEnabled(leader); - } catch (Exception e) { - log.error("Error checking team leader metered billing: {}", e.getMessage()); - return false; - } - } -} diff --git a/app/saas/src/main/java/stirling/software/saas/model/CreditConsumptionResult.java b/app/saas/src/main/java/stirling/software/saas/model/CreditConsumptionResult.java deleted file mode 100644 index 14a449b042..0000000000 --- a/app/saas/src/main/java/stirling/software/saas/model/CreditConsumptionResult.java +++ /dev/null @@ -1,69 +0,0 @@ -package stirling.software.saas.model; - -import lombok.AllArgsConstructor; -import lombok.Data; - -/** - * Result of a credit consumption attempt with explicit waterfall logic. Indicates whether the - * operation succeeded and which credit source was used. - */ -@Data -@AllArgsConstructor -public class CreditConsumptionResult { - - /** Whether the credit consumption succeeded */ - private boolean success; - - /** - * The credit source used for this operation. Possible values: "PRO_PLAN" (Pro user with - * unlimited UI access, no credits consumed); "CYCLE_CREDITS" (free monthly cycle credit - * allocation); "BOUGHT_CREDITS" (one-time purchased credits); "METERED_SUBSCRIPTION" - * (pay-what-you-use metered billing, reported to Stripe); null (operation failed; see message - * for reason). - */ - private String source; - - /** Human-readable message about the result */ - private String message; - - /** - * Creates a successful result for unlimited access (Pro plan UI requests). - * - * @param source The credit source (typically "PRO_PLAN") - * @return CreditConsumptionResult indicating unlimited access - */ - public static CreditConsumptionResult unlimited(String source) { - return new CreditConsumptionResult(true, source, "Unlimited access"); - } - - /** - * Creates a successful result for credit consumption. - * - * @param source The credit source used - * @return CreditConsumptionResult indicating success - */ - public static CreditConsumptionResult success(String source) { - return new CreditConsumptionResult(true, source, "Credits consumed"); - } - - /** - * Creates a failure result. - * - * @param reason The reason for failure - * @return CreditConsumptionResult indicating failure - */ - public static CreditConsumptionResult failure(String reason) { - return new CreditConsumptionResult(false, null, reason); - } - - /** - * Creates a failure result with custom message. - * - * @param reason The reason code - * @param message Custom human-readable message - * @return CreditConsumptionResult indicating failure - */ - public static CreditConsumptionResult failure(String reason, String message) { - return new CreditConsumptionResult(false, null, message != null ? message : reason); - } -} diff --git a/app/saas/src/main/java/stirling/software/saas/model/ProcessingErrorType.java b/app/saas/src/main/java/stirling/software/saas/model/ProcessingErrorType.java deleted file mode 100644 index f30c9a5f84..0000000000 --- a/app/saas/src/main/java/stirling/software/saas/model/ProcessingErrorType.java +++ /dev/null @@ -1,139 +0,0 @@ -package stirling.software.saas.model; - -public enum ProcessingErrorType { - - /** - * Validation errors. should never cost credits. Examples: missing parameters, invalid file - * types, size limits exceeded, malformed requests, authentication failures - */ - VALIDATION_ERROR, - - /** - * Processing errors. should cost credits after 3rd attempt per user/endpoint. Examples: corrupt - * PDF files, unsupported PDF features, memory issues during processing, OCR failures on valid - * PDFs, conversion errors on valid files - */ - PROCESSING_ERROR, - - /** - * System errors. should not cost credits (our fault). Examples: database connection issues, - * filesystem problems, service unavailable, internal server errors - */ - SYSTEM_ERROR; - - /** Determine error type from exception and HTTP status */ - public static ProcessingErrorType classifyError( - Throwable throwable, int httpStatus, String endpoint) { - if (throwable == null) { - return classifyByHttpStatus(httpStatus); - } - - String errorMessage = throwable.getMessage(); - String exceptionClass = throwable.getClass().getSimpleName(); - - // Validation errors (client-side issues) - if (httpStatus == 400 || httpStatus == 422) { - if (isValidationError(errorMessage, exceptionClass)) { - return VALIDATION_ERROR; - } - } - - // Authentication/Authorization errors - if (httpStatus == 401 || httpStatus == 403) { - return VALIDATION_ERROR; - } - - // Rate limiting - if (httpStatus == 429) { - return VALIDATION_ERROR; - } - - // System errors (our fault) - if (httpStatus >= 500 || isSystemError(errorMessage, exceptionClass)) { - return SYSTEM_ERROR; - } - - // Processing errors (user's data issue but valid request) - if (isProcessingError(errorMessage, exceptionClass, endpoint)) { - return PROCESSING_ERROR; - } - - // Default to validation error to be safe - return VALIDATION_ERROR; - } - - private static ProcessingErrorType classifyByHttpStatus(int httpStatus) { - if (httpStatus >= 400 && httpStatus < 500) { - return VALIDATION_ERROR; - } else if (httpStatus >= 500) { - return SYSTEM_ERROR; - } - return VALIDATION_ERROR; - } - - private static boolean isValidationError(String errorMessage, String exceptionClass) { - if (errorMessage == null && exceptionClass == null) return false; - - String[] validationKeywords = { - "validation", "invalid parameter", "missing parameter", "malformed", - "bad request", "illegal argument", "file too large", "unsupported file type", - "empty file", "no file provided", "invalid format" - }; - - String[] validationExceptions = { - "IllegalArgumentException", - "ValidationException", - "BindException", - "MethodArgumentNotValidException", - "MissingServletRequestParameterException", - "HttpMessageNotReadableException", - "MaxUploadSizeExceededException" - }; - - return containsAny(errorMessage, validationKeywords) - || containsAny(exceptionClass, validationExceptions); - } - - private static boolean isSystemError(String errorMessage, String exceptionClass) { - if (errorMessage == null && exceptionClass == null) return false; - - String[] systemExceptions = { - "SQLException", - "IOException", - "OutOfMemoryError", - "TimeoutException", - "ConnectException", - "UnknownHostException", - "ServiceUnavailableException" - }; - - return containsAny(exceptionClass, systemExceptions); - } - - private static boolean isProcessingError( - String errorMessage, String exceptionClass, String endpoint) { - if (errorMessage == null && exceptionClass == null) return false; - - String[] processingExceptions = { - "PDFException", "COSVisitorException", "InvalidPDFException", - "ConversionException", "OCRException", "ParseException" - }; - - // If we're checking errors for an endpoint, it's already been identified as a tracked - // endpoint - // through @AutoJobPostMapping annotation, so we can assume it's a PDF processing endpoint - return containsAny(exceptionClass, processingExceptions) - || (endpoint != null && !isValidationError(errorMessage, exceptionClass)); - } - - private static boolean containsAny(String text, String[] keywords) { - if (text == null) return false; - String lowerText = text.toLowerCase(); - for (String keyword : keywords) { - if (lowerText.contains(keyword.toLowerCase())) { - return true; - } - } - return false; - } -} diff --git a/app/saas/src/main/java/stirling/software/saas/model/TeamCredit.java b/app/saas/src/main/java/stirling/software/saas/model/TeamCredit.java deleted file mode 100644 index 661cda1d4c..0000000000 --- a/app/saas/src/main/java/stirling/software/saas/model/TeamCredit.java +++ /dev/null @@ -1,131 +0,0 @@ -package stirling.software.saas.model; - -import java.io.Serializable; -import java.time.LocalDateTime; - -import org.hibernate.annotations.CreationTimestamp; -import org.hibernate.annotations.OnDelete; -import org.hibernate.annotations.OnDeleteAction; -import org.hibernate.annotations.UpdateTimestamp; - -import jakarta.persistence.Column; -import jakarta.persistence.Entity; -import jakarta.persistence.FetchType; -import jakarta.persistence.GeneratedValue; -import jakarta.persistence.GenerationType; -import jakarta.persistence.Id; -import jakarta.persistence.JoinColumn; -import jakarta.persistence.OneToOne; -import jakarta.persistence.Table; -import jakarta.persistence.Version; - -import lombok.Getter; -import lombok.NoArgsConstructor; -import lombok.Setter; - -import stirling.software.proprietary.model.Team; - -/** Shared credit pool for multi-member teams; see {@link UserCredit} for the per-user variant. */ -@Entity -@Table(name = "team_credits") -@NoArgsConstructor -@Getter -@Setter -public class TeamCredit implements Serializable { - - private static final long serialVersionUID = 1L; - - @Id - @GeneratedValue(strategy = GenerationType.IDENTITY) - @Column(name = "credit_id") - private Long id; - - @OneToOne(fetch = FetchType.LAZY) - @JoinColumn(name = "team_id", nullable = false, unique = true) - @OnDelete(action = OnDeleteAction.CASCADE) - private Team team; - - @Column(name = "cycle_credits_remaining") - private Integer cycleCreditsRemaining = 0; - - @Column(name = "cycle_credits_allocated") - private Integer cycleCreditsAllocated = 0; - - @Column(name = "bought_credits_remaining") - private Integer boughtCreditsRemaining = 0; - - @Column(name = "total_bought_credits") - private Integer totalBoughtCredits = 0; - - @Column(name = "last_cycle_reset_at") - private LocalDateTime lastCycleResetAt; - - @Column(name = "last_api_usage") - private LocalDateTime lastApiUsage; - - @Column(name = "total_api_calls_made") - private Long totalApiCallsMade = 0L; - - @CreationTimestamp - @Column(name = "created_at", updatable = false) - private LocalDateTime createdAt; - - @UpdateTimestamp - @Column(name = "updated_at") - private LocalDateTime updatedAt; - - @Version - @Column(name = "version") - private Long version; - - public TeamCredit(Team team) { - this.team = team; - } - - public int getTotalAvailableCredits() { - return (cycleCreditsRemaining != null ? cycleCreditsRemaining : 0) - + (boughtCreditsRemaining != null ? boughtCreditsRemaining : 0); - } - - public boolean hasCreditsAvailable() { - return getTotalAvailableCredits() > 0; - } - - /** - * Consume a credit from the team pool. Consumes cycle credits first, then bought credits. - * - * @return true if a credit was consumed, false if no credits available - */ - public boolean consumeCredit() { - if (cycleCreditsRemaining != null && cycleCreditsRemaining > 0) { - cycleCreditsRemaining--; - totalApiCallsMade++; - lastApiUsage = LocalDateTime.now(); - return true; - } else if (boughtCreditsRemaining != null && boughtCreditsRemaining > 0) { - boughtCreditsRemaining--; - totalApiCallsMade++; - lastApiUsage = LocalDateTime.now(); - return true; - } - return false; - } - - public void addBoughtCredits(int credits) { - if (credits > 0) { - boughtCreditsRemaining = - (boughtCreditsRemaining != null ? boughtCreditsRemaining : 0) + credits; - totalBoughtCredits = (totalBoughtCredits != null ? totalBoughtCredits : 0) + credits; - } - } - - public void resetCycleCredits(int cycleAllocation, LocalDateTime resetTime) { - this.cycleCreditsAllocated = cycleAllocation; - this.cycleCreditsRemaining = cycleAllocation; - this.lastCycleResetAt = resetTime; - } - - public boolean isCycleResetDue(LocalDateTime lastScheduledReset) { - return lastCycleResetAt == null || lastCycleResetAt.isBefore(lastScheduledReset); - } -} diff --git a/app/saas/src/main/java/stirling/software/saas/model/UserCredit.java b/app/saas/src/main/java/stirling/software/saas/model/UserCredit.java deleted file mode 100644 index 4a947df306..0000000000 --- a/app/saas/src/main/java/stirling/software/saas/model/UserCredit.java +++ /dev/null @@ -1,133 +0,0 @@ -package stirling.software.saas.model; - -import java.io.Serializable; -import java.time.LocalDateTime; - -import org.hibernate.annotations.CreationTimestamp; -import org.hibernate.annotations.OnDelete; -import org.hibernate.annotations.OnDeleteAction; -import org.hibernate.annotations.UpdateTimestamp; - -import jakarta.persistence.Column; -import jakarta.persistence.Entity; -import jakarta.persistence.FetchType; -import jakarta.persistence.GeneratedValue; -import jakarta.persistence.GenerationType; -import jakarta.persistence.Id; -import jakarta.persistence.JoinColumn; -import jakarta.persistence.ManyToOne; -import jakarta.persistence.Table; -import jakarta.persistence.Version; - -import lombok.Getter; -import lombok.NoArgsConstructor; -import lombok.Setter; - -import stirling.software.proprietary.security.model.User; - -/** - * Per-user credit pool. Layers a renewable monthly cycle pool ({@code cycleCreditsRemaining}) over - * a non-expiring purchased pool ({@code boughtCreditsRemaining}); cycle credits consume first. - */ -@Entity -@Table(name = "user_credits") -@NoArgsConstructor -@Getter -@Setter -public class UserCredit implements Serializable { - - private static final long serialVersionUID = 1L; - - @Id - @GeneratedValue(strategy = GenerationType.IDENTITY) - @Column(name = "credit_id") - private Long id; - - @ManyToOne(fetch = FetchType.LAZY) - @JoinColumn(name = "user_id", nullable = false) - @OnDelete(action = OnDeleteAction.CASCADE) - private User user; - - @Column(name = "cycle_credits_remaining") - private Integer cycleCreditsRemaining = 0; - - @Column(name = "cycle_credits_allocated") - private Integer cycleCreditsAllocated = 0; - - @Column(name = "bought_credits_remaining") - private Integer boughtCreditsRemaining = 0; - - @Column(name = "total_bought_credits") - private Integer totalBoughtCredits = 0; - - @Column(name = "last_cycle_reset_at") - private LocalDateTime lastCycleResetAt; - - @Column(name = "last_api_usage") - private LocalDateTime lastApiUsage; - - @Column(name = "total_api_calls_made") - private Long totalApiCallsMade = 0L; - - @CreationTimestamp - @Column(name = "created_at", updatable = false) - private LocalDateTime createdAt; - - @UpdateTimestamp - @Column(name = "updated_at") - private LocalDateTime updatedAt; - - @Version - @Column(name = "version") - private Long version; - - public UserCredit(User user) { - this.user = user; - // Cycle credits are initialized by CreditService after this object is created, - // typically during user registration or at the start of a new billing cycle, - // using values from the application configuration. - } - - public int getTotalAvailableCredits() { - return (cycleCreditsRemaining != null ? cycleCreditsRemaining : 0) - + (boughtCreditsRemaining != null ? boughtCreditsRemaining : 0); - } - - public boolean hasCreditsAvailable() { - return getTotalAvailableCredits() > 0; - } - - public boolean consumeCredit() { - // Consume cycle credits first, then bought credits. - if (cycleCreditsRemaining != null && cycleCreditsRemaining > 0) { - cycleCreditsRemaining--; - totalApiCallsMade++; - lastApiUsage = LocalDateTime.now(); - return true; - } else if (boughtCreditsRemaining != null && boughtCreditsRemaining > 0) { - boughtCreditsRemaining--; - totalApiCallsMade++; - lastApiUsage = LocalDateTime.now(); - return true; - } - return false; - } - - public void addBoughtCredits(int credits) { - if (credits > 0) { - boughtCreditsRemaining = - (boughtCreditsRemaining != null ? boughtCreditsRemaining : 0) + credits; - totalBoughtCredits = (totalBoughtCredits != null ? totalBoughtCredits : 0) + credits; - } - } - - public void resetCycleCredits(int cycleAllocation, LocalDateTime resetTime) { - this.cycleCreditsAllocated = cycleAllocation; - this.cycleCreditsRemaining = cycleAllocation; - this.lastCycleResetAt = resetTime; - } - - public boolean isCycleResetDue(LocalDateTime lastScheduledReset) { - return lastCycleResetAt == null || lastCycleResetAt.isBefore(lastScheduledReset); - } -} diff --git a/app/saas/src/main/java/stirling/software/saas/model/UserErrorTracker.java b/app/saas/src/main/java/stirling/software/saas/model/UserErrorTracker.java deleted file mode 100644 index 2ae97fffb7..0000000000 --- a/app/saas/src/main/java/stirling/software/saas/model/UserErrorTracker.java +++ /dev/null @@ -1,95 +0,0 @@ -package stirling.software.saas.model; - -import java.io.Serializable; -import java.time.LocalDateTime; - -import org.hibernate.annotations.CreationTimestamp; -import org.hibernate.annotations.UpdateTimestamp; - -import jakarta.persistence.Column; -import jakarta.persistence.Entity; -import jakarta.persistence.FetchType; -import jakarta.persistence.GeneratedValue; -import jakarta.persistence.GenerationType; -import jakarta.persistence.Id; -import jakarta.persistence.JoinColumn; -import jakarta.persistence.ManyToOne; -import jakarta.persistence.Table; - -import lombok.Getter; -import lombok.NoArgsConstructor; -import lombok.Setter; - -import stirling.software.proprietary.security.model.User; - -@Entity -@Table(name = "user_error_tracker") -@NoArgsConstructor -@Getter -@Setter -public class UserErrorTracker implements Serializable { - - private static final long serialVersionUID = 1L; - - @Id - @GeneratedValue(strategy = GenerationType.IDENTITY) - @Column(name = "error_tracker_id") - private Long id; - - @ManyToOne(fetch = FetchType.LAZY) - @JoinColumn(name = "user_id", nullable = false) - private User user; - - @Column(name = "endpoint") - private String endpoint; - - @Column(name = "processing_error_count") - private Integer processingErrorCount = 0; - - @Column(name = "last_processing_error") - private LocalDateTime lastProcessingError; - - @Column(name = "reset_after") - private LocalDateTime resetAfter; - - @CreationTimestamp - @Column(name = "created_at", updatable = false) - private LocalDateTime createdAt; - - @UpdateTimestamp - @Column(name = "updated_at") - private LocalDateTime updatedAt; - - public UserErrorTracker(User user, String endpoint, int ttlMinutes) { - this.user = user; - this.endpoint = endpoint; - this.resetAfter = LocalDateTime.now().plusMinutes(ttlMinutes); - } - - public boolean shouldChargeForProcessingError(int freeProcessingErrors) { - return processingErrorCount != null && processingErrorCount > freeProcessingErrors; - } - - public void recordProcessingError(int ttlMinutes) { - this.processingErrorCount = (processingErrorCount != null ? processingErrorCount : 0) + 1; - this.lastProcessingError = LocalDateTime.now(); - - // Refresh TTL on each error - this.resetAfter = LocalDateTime.now().plusMinutes(ttlMinutes); - } - - public void resetErrorCount(int ttlMinutes) { - this.processingErrorCount = 0; - this.lastProcessingError = null; - this.resetAfter = LocalDateTime.now().plusMinutes(ttlMinutes); - } - - public boolean isExpired() { - return resetAfter != null && LocalDateTime.now().isAfter(resetAfter); - } - - public int getErrorsUntilCharged(int freeProcessingErrors) { - int current = processingErrorCount != null ? processingErrorCount : 0; - return Math.max(0, freeProcessingErrors + 1 - current); - } -} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/api/CapMoneyUnits.java b/app/saas/src/main/java/stirling/software/saas/payg/api/CapMoneyUnits.java new file mode 100644 index 0000000000..7ecad6a8b2 --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/api/CapMoneyUnits.java @@ -0,0 +1,56 @@ +package stirling.software.saas.payg.api; + +/** + * Cap conversion between the dollar amount the leader edits in the UI and the unit count stored on + * {@code wallet_policy.cap_units}. + * + *

The cap is application-layer only — Stripe stays on a flat-priced single meter, so the + * conversion rate here doesn't need to match any Stripe price. It only needs to be stable: a leader + * who set "$25" should read back "$25" on the next page load. + * + *

V1 rate: {@value #UNITS_PER_USD} units = $1. This anchors the in-app cap representation to the + * same unit count the ledger writes, so the "X of Y units used" widget in the FE lines up against + * the cap. A future iteration can read this from {@code pricing_policy} once the per-policy money + * conversion lands. + * + *

Both directions floor: $24.50 → 2450 units, 2450 units → $24. The FE only sends whole-dollar + * inputs (the cap-edit field is an integer text box) so the floor on the read path is the only + * place rounding ever shows up, and only when an admin set a non-multiple via SQL. + */ +public final class CapMoneyUnits { + + /** + * Doc-units per USD. {@code 100} = "1 cent per unit" at the in-app display layer. Tied to the + * unit-count meter the engine writes to the ledger; not tied to Stripe pricing. + */ + public static final int UNITS_PER_USD = 100; + + /** Smallest currency unit per USD (always 100 cents in USD; explicit for clarity). */ + public static final int CENTS_PER_USD = 100; + + private CapMoneyUnits() {} + + /** Convert a dollar cap entered by the leader to doc-units for {@code cap_units}. */ + public static long usdToUnits(int capUsd) { + if (capUsd < 0) { + throw new IllegalArgumentException("capUsd must be >= 0"); + } + return (long) capUsd * UNITS_PER_USD; + } + + /** Convert {@code cap_units} back to dollars for the response payload. Floor on read. */ + public static int unitsToUsd(long capUnits) { + if (capUnits < 0L) { + return 0; + } + return (int) (capUnits / UNITS_PER_USD); + } + + /** Convert a dollar cap to smallest-currency-unit cents for {@code cap_source_money}. */ + public static long usdToCents(int capUsd) { + if (capUsd < 0) { + throw new IllegalArgumentException("capUsd must be >= 0"); + } + return (long) capUsd * CENTS_PER_USD; + } +} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/api/PaygWalletController.java b/app/saas/src/main/java/stirling/software/saas/payg/api/PaygWalletController.java new file mode 100644 index 0000000000..041eaa8f4a --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/api/PaygWalletController.java @@ -0,0 +1,424 @@ +package stirling.software.saas.payg.api; + +import java.time.LocalDateTime; +import java.time.format.DateTimeFormatter; +import java.util.ArrayList; +import java.util.Collections; +import java.util.HashMap; +import java.util.List; +import java.util.Locale; +import java.util.Map; +import java.util.Objects; +import java.util.Optional; + +import org.springframework.context.annotation.Profile; +import org.springframework.http.HttpStatus; +import org.springframework.http.ResponseEntity; +import org.springframework.security.access.prepost.PreAuthorize; +import org.springframework.security.core.Authentication; +import org.springframework.transaction.annotation.Transactional; +import org.springframework.web.bind.annotation.GetMapping; +import org.springframework.web.bind.annotation.PatchMapping; +import org.springframework.web.bind.annotation.RequestBody; +import org.springframework.web.bind.annotation.RequestMapping; +import org.springframework.web.bind.annotation.RestController; + +import io.swagger.v3.oas.annotations.Hidden; + +import jakarta.validation.Valid; +import jakarta.validation.constraints.Min; + +import lombok.extern.slf4j.Slf4j; + +import stirling.software.common.model.enumeration.TeamRole; +import stirling.software.proprietary.security.database.repository.UserRepository; +import stirling.software.proprietary.security.model.User; +import stirling.software.saas.model.TeamMembership; +import stirling.software.saas.payg.api.WalletSnapshotResponse.ActivityRow; +import stirling.software.saas.payg.api.WalletSnapshotResponse.CategoryBreakdown; +import stirling.software.saas.payg.api.WalletSnapshotResponse.MemberRow; +import stirling.software.saas.payg.billing.TeamBillingContext; +import stirling.software.saas.payg.billing.TeamBillingService; +import stirling.software.saas.payg.entitlement.EntitlementService; +import stirling.software.saas.payg.entitlement.EntitlementSnapshot; +import stirling.software.saas.payg.model.BillingCategory; +import stirling.software.saas.payg.model.LedgerEntryType; +import stirling.software.saas.payg.policy.PaygTeamExtensions; +import stirling.software.saas.payg.repository.PaygShadowChargeRepository; +import stirling.software.saas.payg.repository.PaygTeamExtensionsRepository; +import stirling.software.saas.payg.repository.WalletLedgerRepository; +import stirling.software.saas.payg.repository.WalletPolicyRepository; +import stirling.software.saas.payg.wallet.WalletLedgerEntry; +import stirling.software.saas.payg.wallet.WalletPolicy; +import stirling.software.saas.repository.TeamMembershipRepository; +import stirling.software.saas.util.AuthenticationUtils; + +/** + * Read + cap-mutation surface backing the FE PAYG Plan page. + * + *

{@code GET /api/v1/payg/wallet} is the single fetch the {@code useWallet} hook calls. Returns + * a fully-populated {@link WalletSnapshotResponse} — derived from {@link EntitlementService} (for + * spend / cap / period), {@link PaygTeamExtensions} (for subscription state), and {@link + * WalletLedgerRepository} (for the per-category breakdown widget). Leader callers also get a roster + * of team members + their per-member usage; member callers see an empty roster. + * + *

{@code PATCH /api/v1/payg/cap} updates {@code wallet_policy.cap_units} (no Stripe call — the + * cap is enforced application-side via the entitlement guard) and invalidates the team's snapshot + * cache so the next read reflects the change immediately. Only leaders may call this; the team is + * derived from the caller, so we authorise inside the method rather than via {@code @PreAuthorize} + * — the team id never appears on the path or query string. + * + *

Subscription state is sourced from {@code payg_team_extensions.payg_subscription_id} (added in + * V14): {@code stripeSubscriptionId} echoes it via {@link TeamBillingService}, and a team reads as + * {@link #STATUS_SUBSCRIBED} once {@code billing.subscribed()} is true — i.e. it has a subscription + * id, or a Stripe customer id as the pre-webhook bridge for a just-completed checkout whose + * subscription-created webhook hasn't landed yet (see {@code TeamBillingService.compute}). + */ +@Slf4j +@Hidden +@RestController +@RequestMapping("/api/v1/payg") +@Profile("saas") +public class PaygWalletController { + + static final String STATUS_FREE = "free"; + static final String STATUS_SUBSCRIBED = "subscribed"; + static final String ROLE_LEADER = "leader"; + static final String ROLE_MEMBER = "member"; + + /** + * Placeholder ceiling for the team-less empty snapshot only (authenticated caller without a + * membership — shouldn't happen post-migration). Teams always get the live {@code + * pricing_policy.free_tier_units} grant via {@link TeamBillingService}. + */ + private static final int FREE_TIER_LIMIT_UNITS_FALLBACK = 500; + + private static final DateTimeFormatter ISO_DATE = DateTimeFormatter.ISO_LOCAL_DATE; + + private final EntitlementService entitlementService; + private final TeamBillingService billingService; + private final TeamMembershipRepository memberRepo; + private final PaygTeamExtensionsRepository extRepo; + private final WalletPolicyRepository policyRepo; + private final WalletLedgerRepository ledgerRepo; + private final PaygShadowChargeRepository shadowRepo; + private final UserRepository userRepository; + + public PaygWalletController( + EntitlementService entitlementService, + TeamBillingService billingService, + TeamMembershipRepository memberRepo, + PaygTeamExtensionsRepository extRepo, + WalletPolicyRepository policyRepo, + WalletLedgerRepository ledgerRepo, + PaygShadowChargeRepository shadowRepo, + UserRepository userRepository) { + this.entitlementService = Objects.requireNonNull(entitlementService, "entitlementService"); + this.billingService = Objects.requireNonNull(billingService, "billingService"); + this.memberRepo = Objects.requireNonNull(memberRepo, "memberRepo"); + this.extRepo = Objects.requireNonNull(extRepo, "extRepo"); + this.policyRepo = Objects.requireNonNull(policyRepo, "policyRepo"); + this.ledgerRepo = Objects.requireNonNull(ledgerRepo, "ledgerRepo"); + this.shadowRepo = Objects.requireNonNull(shadowRepo, "shadowRepo"); + this.userRepository = Objects.requireNonNull(userRepository, "userRepository"); + } + + // --------------------------------------------------------------------------------------- + // GET /wallet — the single FE fetch + // --------------------------------------------------------------------------------------- + + @GetMapping("/wallet") + @PreAuthorize("isAuthenticated()") + @Transactional(readOnly = true) + public ResponseEntity getWallet(Authentication auth) { + User user; + try { + user = AuthenticationUtils.getCurrentUser(auth, userRepository); + } catch (SecurityException e) { + // SecurityException maps to 401 per the existing controller convention. + return ResponseEntity.status(HttpStatus.UNAUTHORIZED).build(); + } + + Optional primary = primaryMembership(user.getId()); + if (primary.isEmpty()) { + // Authenticated user without a team — shouldn't happen post-migration, but we don't + // want to 500. Return a free-tier-shaped empty snapshot so the FE renders the gated UI + // rather than blowing up on a null body. + return ResponseEntity.ok(emptySnapshot()); + } + + TeamMembership membership = primary.get(); + Long teamId = membership.getTeam().getId(); + boolean isLeader = membership.getRole() == TeamRole.LEADER; + + // Billing facts (window, free allowance, per-doc rate, doc cap) and the entitlement + // snapshot (period spend over that window) share the same composition service, so what + // the customer sees here is exactly what the 402 guard enforces. + TeamBillingContext billing = billingService.forTeam(teamId); + EntitlementSnapshot snap = entitlementService.getSnapshot(teamId); + + String status = billing.subscribed() ? STATUS_SUBSCRIBED : STATUS_FREE; + + boolean noCap = billing.subscribed() && billing.capMoneyMinor() == null; + Integer capMajor = + billing.capMoneyMinor() != null + ? Math.toIntExact(billing.capMoneyMinor() / 100) + : null; + + // Per-state by construction (see EntitlementService.computeSnapshot): free team → spend is + // lifetime free used, cap is the grant size; subscribed → spend is this month's net + // billable + // docs, cap is the monthly paid-doc ceiling (null = uncapped). + int spend = clampToInt(snap.periodSpendUnits()); + Integer limit = snap.periodCapUnits() != null ? clampToInt(snap.periodCapUnits()) : null; + + CategoryBreakdown breakdown = buildBreakdown(teamId, snap.periodStart(), snap.periodEnd()); + + // Estimated bill = paid (Stripe-metered) docs this period × rate — the free portion was + // already netted out at charge time, so this is the metered total, not spend − grant. + long periodPaid = shadowRepo.sumPaidUnits(teamId, snap.periodStart(), snap.periodEnd()); + Long estimatedBill = billingService.estimateBillMinor(billing, periodPaid).orElse(null); + + List members = + isLeader + ? buildMemberRows(teamId, snap.periodStart(), snap.periodEnd()) + : List.of(); + + WalletSnapshotResponse body = + new WalletSnapshotResponse( + teamId, + status, + isLeader ? ROLE_LEADER : ROLE_MEMBER, + ISO_DATE.format(snap.periodStart().toLocalDate()), + ISO_DATE.format(snap.periodEnd().toLocalDate()), + spend, + limit, + clampToInt(billing.freeGrantUnits()), + clampToInt(billing.freeRemainingUnits()), + billing.perDocMinor(), + billing.currency(), + estimatedBill, + capMajor, + noCap, + billing.subscriptionId(), + spend, + breakdown, + members, + buildActivity(teamId)); + return ResponseEntity.ok(body); + } + + private CategoryBreakdown buildBreakdown( + Long teamId, LocalDateTime periodStart, LocalDateTime periodEnd) { + Map byCategory = new HashMap<>(); + for (Object[] row : + ledgerRepo.sumPeriodAmountByCategory( + teamId, LedgerEntryType.DEBIT, periodStart, periodEnd)) { + if (row.length >= 2 + && row[0] instanceof BillingCategory cat + && row[1] instanceof Number n) { + byCategory.put(cat, n.longValue()); + } + } + return new CategoryBreakdown( + clampToInt(byCategory.getOrDefault(BillingCategory.API, 0L)), + clampToInt(byCategory.getOrDefault(BillingCategory.AI, 0L)), + clampToInt(byCategory.getOrDefault(BillingCategory.AUTOMATION, 0L))); + } + + /** + * Latest ledger entries shaped for the FE activity feed. DEBITs read as usage, REFUNDs as + * credits-back; system entries without a category render as {@code other}. + */ + private List buildActivity(Long teamId) { + List out = new ArrayList<>(); + for (WalletLedgerEntry e : ledgerRepo.findTop20ByTeamIdOrderByIdDesc(teamId)) { + BillingCategory category = e.getBillingCategory(); + String kind = category != null ? category.name().toLowerCase(Locale.ROOT) : "other"; + String categoryLabel = category != null ? categoryDisplayName(category) : "Document"; + String label = + e.getEntryType() == LedgerEntryType.REFUND + ? "Refund — " + categoryLabel + : categoryLabel + " usage"; + int docUnits = e.getAmountUnits() == null ? 0 : Math.abs(e.getAmountUnits()); + out.add( + new ActivityRow( + e.getId(), + kind, + label, + e.getOccurredAt() != null ? e.getOccurredAt().toString() : "", + docUnits)); + } + return out; + } + + private static String categoryDisplayName(BillingCategory category) { + return switch (category) { + case API -> "API"; + case AI -> "AI"; + case AUTOMATION -> "Automation"; + case BYPASSED -> "Manual"; + }; + } + + // --------------------------------------------------------------------------------------- + // PATCH /cap — leader-only, cap is application-layer, no Stripe call + // --------------------------------------------------------------------------------------- + + @PatchMapping("/cap") + @PreAuthorize("isAuthenticated()") + @Transactional + public ResponseEntity updateCap( + @Valid @RequestBody UpdateCapRequest req, Authentication auth) { + User user; + try { + user = AuthenticationUtils.getCurrentUser(auth, userRepository); + } catch (SecurityException e) { + return ResponseEntity.status(HttpStatus.UNAUTHORIZED).build(); + } + Optional primary = primaryMembership(user.getId()); + if (primary.isEmpty()) { + // No team → can't have a wallet to cap. + return ResponseEntity.status(HttpStatus.FORBIDDEN).build(); + } + TeamMembership membership = primary.get(); + if (membership.getRole() != TeamRole.LEADER) { + return ResponseEntity.status(HttpStatus.FORBIDDEN).build(); + } + Long teamId = membership.getTeam().getId(); + + WalletPolicy policy = + policyRepo + .findByTeamId(teamId) + .orElseGet( + () -> { + WalletPolicy created = new WalletPolicy(); + created.setTeamId(teamId); + return created; + }); + + if (req.noCap()) { + policy.setCapUnits(null); + policy.setCapSourceMoney(null); + } else { + long capMinor = CapMoneyUnits.usdToCents(req.capUsd()); + policy.setCapSourceMoney(capMinor); + // Derived document allowance: store both the money intent and the unit translation. + // The live snapshot recomputes from cap_source_money + current rate; this stored value + // is the enforcement fallback when the rate is unreachable. + TeamBillingContext billing = billingService.forTeam(teamId); + Optional docCap = billingService.docCapForMoney(billing, capMinor); + if (docCap.isPresent()) { + policy.setCapUnits(docCap.get()); + } else { + // Rate unknown (price-info fn unconfigured / Stripe blip): keep the legacy + // money-as-units conversion so the cap still binds rather than silently lifting. + log.warn( + "Per-document rate unavailable for team {}; storing legacy cap_units" + + " conversion.", + teamId); + policy.setCapUnits(CapMoneyUnits.usdToUnits(req.capUsd())); + } + } + policyRepo.save(policy); + entitlementService.invalidate(teamId); + return ResponseEntity.noContent().build(); + } + + /** Request body for {@link #updateCap}. */ + public record UpdateCapRequest(@Min(0) int capUsd, boolean noCap) {} + + // --------------------------------------------------------------------------------------- + // Helpers + // --------------------------------------------------------------------------------------- + + private Optional primaryMembership(Long userId) { + List rows = memberRepo.findPrimaryMembership(userId); + return rows.isEmpty() ? Optional.empty() : Optional.of(rows.get(0)); + } + + private List buildMemberRows( + Long teamId, LocalDateTime periodStart, LocalDateTime periodEnd) { + List all = memberRepo.findByTeamId(teamId); + if (all.isEmpty()) { + return List.of(); + } + LocalDateTime[] window = {periodStart, periodEnd}; + List out = new ArrayList<>(all.size()); + for (TeamMembership tm : all) { + User u = tm.getUser(); + if (u == null) { + continue; + } + // We could batch these; team sizes are small (FE design assumes ≤ ~20 members per + // team on the Plan page) so a per-member sum is fine. If teams grow we'd switch to + // a single GROUP BY actor_user_id query. + long spend = 0L; // sumPeriodAmountForMember stores signed debits (negative); negate. + try { + spend = -memberSpend(teamId, u.getId(), window[0], window[1]); + } catch (RuntimeException e) { + log.warn( + "buildMemberRows: per-member spend lookup failed for user {}", + u.getId(), + e); + } + String displayName = + Optional.ofNullable(u.getUsername()) + .orElse(Optional.ofNullable(u.getEmail()).orElse("")); + out.add( + new MemberRow( + Long.toString(u.getId()), + displayName, + Optional.ofNullable(u.getEmail()).orElse(""), + clampToInt(spend))); + } + return out; + } + + /** + * Per-member period spend in signed ledger units (debits are negative). Helper so the test + * slice can override without standing up a real database, and so the controller doesn't inline + * the negation arithmetic at every call site. + */ + long memberSpend(Long teamId, Long userId, LocalDateTime start, LocalDateTime end) { + return ledgerRepo.sumPeriodAmountForMember( + teamId, userId, LedgerEntryType.DEBIT, start, end); + } + + private static LocalDateTime[] currentMonthWindow() { + java.time.YearMonth ym = java.time.YearMonth.now(); + LocalDateTime start = ym.atDay(1).atStartOfDay(); + LocalDateTime end = ym.plusMonths(1).atDay(1).atStartOfDay(); + return new LocalDateTime[] {start, end}; + } + + private static int clampToInt(long v) { + if (v <= 0) return 0; + if (v >= Integer.MAX_VALUE) return Integer.MAX_VALUE; + return (int) v; + } + + private WalletSnapshotResponse emptySnapshot() { + LocalDateTime[] window = currentMonthWindow(); + return new WalletSnapshotResponse( + null, // teamId — unknown when the caller has no team membership + STATUS_FREE, + ROLE_MEMBER, + ISO_DATE.format(window[0].toLocalDate()), + ISO_DATE.format(window[1].toLocalDate()), + 0, + FREE_TIER_LIMIT_UNITS_FALLBACK, + FREE_TIER_LIMIT_UNITS_FALLBACK, + FREE_TIER_LIMIT_UNITS_FALLBACK, + null, + null, + null, + null, + false, + null, + 0, + new CategoryBreakdown(0, 0, 0), + List.of(), + Collections.emptyList()); + } +} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/api/WalletSnapshotResponse.java b/app/saas/src/main/java/stirling/software/saas/payg/api/WalletSnapshotResponse.java new file mode 100644 index 0000000000..0b29347c85 --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/api/WalletSnapshotResponse.java @@ -0,0 +1,100 @@ +package stirling.software.saas.payg.api; + +import java.math.BigDecimal; +import java.util.List; + +/** + * JSON payload returned by {@code GET /api/v1/payg/wallet}. Mirrors the {@code Wallet} type the + * frontend {@code useWallet} hook consumes, plus the leader-only fields ({@code members}, + * breakdowns, recent activity) used by the PAYG Plan page. + * + *

Every number is real: the billing window is the Stripe subscription's current period (via Sync + * Engine) for subscribed teams, the one-time free grant size comes from {@code + * pricing_policy.free_tier_units} (live balance from {@code + * payg_team_extensions.free_units_remaining}), and the per-document rate comes from the + * subscription's Stripe Price. Fields that can't be resolved are {@code null} and the FE renders + * "unknown" — never a substituted default. + * + * @param teamId the caller's primary team_id. Needed by the frontend so it can pass it to the + * Supabase edge functions that create Stripe Checkout / portal sessions — those run outside + * Spring Security and have no other way to resolve the caller's team. + * @param status {@code "free"} when the team has no Stripe subscription; {@code "subscribed"} once + * a card is on file and the engine bills meter events. + * @param role the current caller's role within their team — {@code "leader"} or {@code "member"}. + * Controls which UI variant the frontend renders. + * @param billingPeriodStart inclusive ISO date (yyyy-MM-dd) for the current cycle — the Stripe + * subscription period when subscribed, the calendar month otherwise. + * @param billingPeriodEnd exclusive ISO date (yyyy-MM-dd) for the current cycle. + * @param billableUsed alias of {@code spendUnitsThisPeriod} kept for clarity in the FE. For a free + * team this is the lifetime free documents used so far ({@code freeAllowance − freeRemaining}); + * for a subscribed team it's this month's net billable documents. + * @param billableLimit the team's document ceiling for the matching window: the one-time free grant + * ({@code freeAllowance}) for free teams; {@code floor(cap / perDocRate)} paid docs/month for + * capped subscribed teams; {@code null} when subscribed with no cap (uncapped). + * @param freeAllowance the team's one-time free document grant size (the "N" in "X of N free"). + * Never resets; survives subscribing. Applies to billable categories only. + * @param freeRemaining one-time free documents still available to the team ({@code + * payg_team_extensions.free_units_remaining}). 0 = grant exhausted. + * @param pricePerDocMinor paid per-document rate in minor units of {@code currency} (may be + * fractional — Stripe supports sub-cent rates); {@code null} when the rate can't be resolved. + * @param currency lower-case ISO 4217 currency of the subscription's Stripe Price; {@code null} + * when unknown (free teams, unresolved rate). + * @param estimatedBillMinor estimated charges so far this period in minor units of {@code + * currency}: paid (Stripe-metered) documents this period × {@code pricePerDocMinor}. The free + * portion was already netted out at charge time. Informational — the Stripe invoice is + * authoritative. {@code null} when the rate is unknown. + * @param capUsd the leader's monthly spending cap in major currency units; {@code null} when free + * or when the leader has opted into no-cap. (Field name predates multi-currency; the FE pairs + * it with {@code currency} for the symbol.) + * @param noCap {@code true} when the leader has explicitly disabled the cap. Only meaningful when + * subscribed. + * @param stripeSubscriptionId Stripe subscription id from {@code + * payg_team_extensions.payg_subscription_id}; {@code null} when status is free. + * @param spendUnitsThisPeriod documents debited this cycle across billable categories. + * @param categoryBreakdown per-category spend slice over the same billing window. + * @param members leader-only roster of team members + their per-member sub-caps. Empty for member + * callers. + * @param recent latest wallet-ledger entries (newest first) for the activity feed. + */ +public record WalletSnapshotResponse( + Long teamId, + String status, + String role, + String billingPeriodStart, + String billingPeriodEnd, + int billableUsed, + Integer billableLimit, + int freeAllowance, + int freeRemaining, + BigDecimal pricePerDocMinor, + String currency, + Long estimatedBillMinor, + Integer capUsd, + boolean noCap, + String stripeSubscriptionId, + int spendUnitsThisPeriod, + CategoryBreakdown categoryBreakdown, + List members, + List recent) { + + /** Per-category breakdown of {@code spendUnitsThisPeriod} for the in-app analytics widget. */ + public record CategoryBreakdown(int api, int ai, int automation) {} + + /** + * One row of the team-members table on the leader's Plan page — display-only per-member usage. + * (Per-member sub-caps aren't enforced yet. When they ship, a cap field returns here.) + */ + public record MemberRow(String userId, String name, String email, int spendUnits) {} + + /** + * One wallet-ledger entry shaped for the FE activity feed. + * + * @param id ledger entry id (stable React key) + * @param kind lower-case billing category ({@code api} / {@code ai} / {@code automation}) or + * {@code other} for system entries + * @param label human line, e.g. {@code "API usage"} or {@code "Refund — API"} + * @param ts ISO-8601 local timestamp of the entry + * @param docUnits absolute document count of the entry + */ + public record ActivityRow(long id, String kind, String label, String ts, int docUnits) {} +} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/billing/TeamBillingContext.java b/app/saas/src/main/java/stirling/software/saas/payg/billing/TeamBillingContext.java new file mode 100644 index 0000000000..55988fbe82 --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/billing/TeamBillingContext.java @@ -0,0 +1,47 @@ +package stirling.software.saas.payg.billing; + +import java.math.BigDecimal; +import java.time.LocalDateTime; + +/** + * One team's billing facts, composed by {@link TeamBillingService}. Two independent meters live + * here and must not be conflated: + * + *

    + *
  • the one-time lifetime free grant ({@link #freeGrantUnits} total, {@link + * #freeRemainingUnits} left) — gates an un-subscribed team and decides the free-vs-paid split + * of every job; never resets, survives subscribing; + *
  • the monthly billing window ({@link #periodStart}/{@link #periodEnd}) and the + * optional monthly spending cap ({@link #monthlyCapDocUnits}) — govern the subscribed invoice + * + cap only. + *
+ * + * @param subscribed team has a live PAYG subscription — i.e. {@code payg_subscription_id} is set. + * Cleared by {@code payg_unlink_subscription} on cancellation, so a cancelled team reads false. + * @param subscriptionId {@code payg_team_extensions.payg_subscription_id}; null when free + * @param periodStart inclusive start of the monthly billing window — the Stripe subscription's + * current period when subscribed, calendar month otherwise + * @param periodEnd exclusive end of the monthly billing window + * @param freeGrantUnits the team's one-time free grant size (policy {@code free_tier_units}); the + * denominator for "used X of N free". Never resets. + * @param freeRemainingUnits one-time free documents still available ({@code + * payg_team_extensions.free_units_remaining}). 0 = grant exhausted. + * @param perDocMinor paid per-document rate in minor units of {@link #currency()}; null when the + * rate can't be resolved (free team, price row unsynced) — display "unknown", never substitute + * @param currency lower-case ISO 4217 of the subscription's Price; null when unknown + * @param capMoneyMinor leader-set monthly spending cap in minor units ({@code + * wallet_policy.cap_source_money}); null = no cap configured + * @param monthlyCapDocUnits the subscribed monthly paid-document ceiling — {@code floor(capMoney / + * perDocRate)}; null = uncapped, or the team is not subscribed + */ +public record TeamBillingContext( + boolean subscribed, + String subscriptionId, + LocalDateTime periodStart, + LocalDateTime periodEnd, + long freeGrantUnits, + long freeRemainingUnits, + BigDecimal perDocMinor, + String currency, + Long capMoneyMinor, + Long monthlyCapDocUnits) {} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/billing/TeamBillingService.java b/app/saas/src/main/java/stirling/software/saas/payg/billing/TeamBillingService.java new file mode 100644 index 0000000000..260efbd9ab --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/billing/TeamBillingService.java @@ -0,0 +1,272 @@ +package stirling.software.saas.payg.billing; + +import java.math.BigDecimal; +import java.math.RoundingMode; +import java.time.Duration; +import java.time.LocalDateTime; +import java.util.Objects; +import java.util.Optional; + +import org.springframework.context.annotation.Profile; +import org.springframework.stereotype.Service; + +import com.github.benmanes.caffeine.cache.Cache; +import com.github.benmanes.caffeine.cache.Caffeine; + +import lombok.extern.slf4j.Slf4j; + +import stirling.software.saas.payg.policy.PaygTeamExtensions; +import stirling.software.saas.payg.policy.PricingPolicy; +import stirling.software.saas.payg.policy.PricingPolicyService; +import stirling.software.saas.payg.repository.PaygTeamExtensionsRepository; +import stirling.software.saas.payg.repository.WalletPolicyRepository; +import stirling.software.saas.payg.stripe.StripeSubscriptionDao; +import stirling.software.saas.payg.stripe.StripeSubscriptionDao.PriceRate; +import stirling.software.saas.payg.stripe.StripeSubscriptionDao.SubscriptionBilling; +import stirling.software.saas.payg.wallet.WalletPolicy; + +/** + * Single composition point for "what does billing look like for this team right now." Both the + * entitlement hot path and the wallet endpoint read from here, so what the customer sees is what + * the guard enforces. + * + *

Two independent meters (design 2026-06-11 — the free allowance is a one-time lifetime grant): + * + *

    + *
  • Free grant — one-time, per team. Size from {@code pricing_policy.free_tier_units}; + * live balance from the {@code payg_team_extensions.free_units_remaining} counter (maintained + * by the charge pipeline). Never resets, survives subscribing. Gates un-subscribed teams and + * drives the free-vs-paid split. + *
  • Monthly window + cap — the Stripe subscription period (calendar month otherwise) and + * the optional money cap. Govern the subscribed invoice + spending cap only. The per-document + * rate is the synced {@code stripe.prices.unit_amount} (PAYG prices are plain per-unit). + *
+ * + *

Cached per team for {@value #CACHE_TTL_SECONDS}s. {@code EntitlementService.invalidate} + * cascades into {@link #invalidate(Long)} so both caches drop together on cap edits / webhooks. + * Note the cached context's {@code freeRemainingUnits} is a 30s-stale read of the counter — the + * authoritative decrement happens in {@code JobChargeService} against the row directly; this cache + * is for display + the entitlement gate, where 30s staleness is the accepted cap-evaluation floor. + */ +@Slf4j +@Service +@Profile("saas") +public class TeamBillingService { + + static final int CACHE_TTL_SECONDS = 30; + private static final int CACHE_MAX_SIZE = 10_000; + + /** + * In-app display/estimate currency. The app prices in dollars; Stripe handles real currency + * selection at checkout. Used to pick the right Price for un-subscribed teams. + */ + private static final String DISPLAY_CURRENCY = "usd"; + + /** + * Stripe Price {@code lookup_key} for the PAYG per-document price. The stable handle we resolve + * an un-subscribed team's rate from (the default policy carries no price ids in the seed). + */ + private static final String PAYG_LOOKUP_KEY = "plan:processor"; + + private final PaygTeamExtensionsRepository extensionsRepository; + private final WalletPolicyRepository walletPolicyRepository; + private final PricingPolicyService pricingPolicyService; + private final StripeSubscriptionDao subscriptionDao; + + private final Cache cache; + + public TeamBillingService( + PaygTeamExtensionsRepository extensionsRepository, + WalletPolicyRepository walletPolicyRepository, + PricingPolicyService pricingPolicyService, + StripeSubscriptionDao subscriptionDao) { + this.extensionsRepository = + Objects.requireNonNull(extensionsRepository, "extensionsRepository"); + this.walletPolicyRepository = + Objects.requireNonNull(walletPolicyRepository, "walletPolicyRepository"); + this.pricingPolicyService = + Objects.requireNonNull(pricingPolicyService, "pricingPolicyService"); + this.subscriptionDao = Objects.requireNonNull(subscriptionDao, "subscriptionDao"); + this.cache = + Caffeine.newBuilder() + .maximumSize(CACHE_MAX_SIZE) + .expireAfterWrite(Duration.ofSeconds(CACHE_TTL_SECONDS)) + .build(); + } + + public TeamBillingContext forTeam(Long teamId) { + Objects.requireNonNull(teamId, "teamId"); + return cache.get(teamId, this::compute); + } + + /** Drop {@code teamId}'s entry after cap edits / subscription webhooks / grant consumption. */ + public void invalidate(Long teamId) { + if (teamId != null) { + cache.invalidate(teamId); + } + } + + private TeamBillingContext compute(Long teamId) { + Optional extOpt = extensionsRepository.findById(teamId); + Optional walletPolicyOpt = walletPolicyRepository.findByTeamId(teamId); + + String subscriptionId = extOpt.map(PaygTeamExtensions::getPaygSubscriptionId).orElse(null); + // payg_subscription_id is the single subscription switch. payg_link_subscription sets it + // (alongside stripe_customer_id, in the same write) on customer.subscription.created; + // payg_unlink_subscription nulls it on customer.subscription.deleted while deliberately + // keeping stripe_customer_id so a future re-subscribe can reuse the Stripe customer. So a + // cancelled team has a null subscription id and must read as free again. + // + // We deliberately do NOT fall back to stripe_customer_id presence. payg_link_subscription + // is the only writer of that column and it writes it together with the subscription id, so + // it can never be set "before the webhook lands" — there is no gap for it to bridge. A + // customer-id fallback would instead keep every team that ever subscribed pinned to + // subscribed forever (the customer outlives the subscription), which is the cancelled-team + // bug this guards against. + boolean subscribed = subscriptionId != null; + + long freeGrant = resolveGrant(teamId); + long freeRemaining = + extOpt.map(PaygTeamExtensions::getFreeUnitsRemaining) + .map(Long::longValue) + .orElse(0L); + + Optional billing = + subscriptionId != null + ? subscriptionDao.findBilling(subscriptionId) + : Optional.empty(); + + LocalDateTime[] window = + billing.map(b -> new LocalDateTime[] {b.periodStart(), b.periodEnd()}) + .orElseGet(TeamBillingService::calendarMonthWindow); + + BigDecimal perDocMinor = billing.map(SubscriptionBilling::perDocMinor).orElse(null); + String currency = billing.map(SubscriptionBilling::currency).orElse(null); + + // Un-subscribed teams have no Stripe subscription to read a rate from, but the cap + // estimate (the upgrade flow's "≈ N paid PDFs/month") still needs one. Resolve it from + // the default policy's USD Price — Stripe hasn't assigned the team a currency yet, and + // the whole app prices in dollars. Display-only: resolveMonthlyCap stays gated on + // `subscribed`, so this never starts enforcing a cap on a free team. + if (!subscribed && perDocMinor == null) { + Optional rate = + subscriptionDao.findRateByLookupKey(PAYG_LOOKUP_KEY, DISPLAY_CURRENCY); + if (rate.isPresent()) { + perDocMinor = rate.get().perDocMinor(); + currency = rate.get().currency(); + } + } + + Long capMoneyMinor = walletPolicyOpt.map(WalletPolicy::getCapSourceMoney).orElse(null); + Long legacyCapUnits = walletPolicyOpt.map(WalletPolicy::getCapUnits).orElse(null); + + Long monthlyCapDocUnits = + resolveMonthlyCap(subscribed, capMoneyMinor, legacyCapUnits, perDocMinor); + + return new TeamBillingContext( + subscribed, + subscriptionId, + window[0], + window[1], + freeGrant, + freeRemaining, + perDocMinor, + currency, + capMoneyMinor, + monthlyCapDocUnits); + } + + /** The policy grant size — the "N" denominator for display; the counter is the live balance. */ + private long resolveGrant(Long teamId) { + try { + PricingPolicy policy = pricingPolicyService.getEffectivePolicy(teamId); + Long grant = policy.getFreeTierUnits(); + return grant == null ? 0L : grant; + } catch (RuntimeException e) { + log.warn("No effective pricing policy for team {}: {}", teamId, e.getMessage()); + return 0L; + } + } + + /** + * The subscribed monthly paid-document ceiling; {@code null} = uncapped or not subscribed. The + * one-time free grant is NOT added here — it's a separate lifetime pool consumed at charge + * time. The cap purely limits how many paid documents the team will fund per billing period. + * + *

    + *
  • not subscribed → null (the free grant, not a money cap, is what bounds them); + *
  • subscribed, no money cap → uncapped (null), unless an admin set raw {@code cap_units}; + *
  • subscribed, money cap + known rate → {@code floor(capMoney / perDocRate)}; + *
  • subscribed, money cap but rate unknown → stored {@code cap_units} fallback (WARN). + *
+ */ + private Long resolveMonthlyCap( + boolean subscribed, Long capMoneyMinor, Long legacyCapUnits, BigDecimal perDocMinor) { + if (!subscribed) { + return null; + } + if (capMoneyMinor == null) { + return legacyCapUnits; // admin-set unit cap (source money null) still applies + } + if (perDocMinor != null && perDocMinor.signum() > 0) { + return BigDecimal.valueOf(capMoneyMinor) + .divide(perDocMinor, 0, RoundingMode.FLOOR) + .longValue(); + } + log.warn( + "Per-document rate unavailable; enforcing stored cap_units fallback ({}).", + legacyCapUnits); + return legacyCapUnits; + } + + /** + * Estimated charges for the current period in minor units of {@link + * TeamBillingContext#currency()}: the paid (metered) documents this period at the per-document + * rate. Informational — the Stripe invoice is authoritative. Empty when the rate is unknown. + * + * @param paidUnitsThisPeriod metered documents this period ({@code payg_units − + * free_units_consumed} summed over the period's charged jobs) + */ + public Optional estimateBillMinor(TeamBillingContext ctx, long paidUnitsThisPeriod) { + if (ctx.perDocMinor() == null) { + return Optional.empty(); + } + long paid = Math.max(0, paidUnitsThisPeriod); + BigDecimal bill = + ctx.perDocMinor() + .multiply(BigDecimal.valueOf(paid)) + .setScale(0, RoundingMode.HALF_UP); + return Optional.of(bill.longValue()); + } + + /** + * Documents a hypothetical monthly money cap would buy: {@code floor(capMinor / rate)}. Used by + * the cap editor's live preview and the {@code PATCH /cap} derived write. The free grant is NOT + * added — it's a separate one-time pool. Empty when the rate is unknown. + */ + public Optional docCapForMoney(TeamBillingContext ctx, long capMinor) { + if (ctx.perDocMinor() == null || ctx.perDocMinor().signum() <= 0) { + return Optional.empty(); + } + return Optional.of( + BigDecimal.valueOf(capMinor) + .divide(ctx.perDocMinor(), 0, RoundingMode.FLOOR) + .longValue()); + } + + /** + * Inclusive-start / exclusive-end window for the calendar month — the monthly billing window + * used when there's no Stripe subscription period to anchor on. + */ + static LocalDateTime[] calendarMonthWindow() { + return calendarMonthWindow(LocalDateTime.now()); + } + + /** Test seam — accepts a clock value so tests don't race the calendar boundary. */ + static LocalDateTime[] calendarMonthWindow(LocalDateTime now) { + java.time.YearMonth ym = java.time.YearMonth.from(now); + return new LocalDateTime[] { + ym.atDay(1).atStartOfDay(), ym.plusMonths(1).atDay(1).atStartOfDay() + }; + } +} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/cap/AiToolRoutes.java b/app/saas/src/main/java/stirling/software/saas/payg/cap/AiToolRoutes.java new file mode 100644 index 0000000000..c6a31f389a --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/cap/AiToolRoutes.java @@ -0,0 +1,42 @@ +package stirling.software.saas.payg.cap; + +import org.springframework.web.servlet.HandlerMapping; + +import jakarta.servlet.http.HttpServletRequest; + +/** + * Central definition of the AI document-tool route namespace ({@code /api/v1/ai/tools/**}). + * + *

These tools (e.g. {@code PdfCommentAgentController}, {@code MathAuditorAgentController}) live + * in the {@code proprietary} module, which does not depend on {@code saas} and therefore cannot + * carry the saas-only {@link RequiresFeature} annotation. Rather than weaken the layering, the PAYG + * hot-path components recognise the path prefix instead: + * + *

    + *
  • {@code PaygChargeInterceptor} brings these routes into scope and bills them as {@code + * BillingCategory.AI} on a direct call (an orchestrator-dispatched call still resolves to + * AUTOMATION first, via the {@code X-Stirling-Automation} header); + *
  • {@code EntitlementGuard} gates them on {@link + * stirling.software.saas.payg.model.FeatureGate#AI_SUPPORT}. + *
+ * + *

Kept as a single source of truth so the interceptor and the guard can never drift on what + * counts as an AI tool. + */ +public final class AiToolRoutes { + + /** Trailing slash so it matches the tool sub-paths, not a bare {@code /api/v1/ai/tools}. */ + public static final String PREFIX = "/api/v1/ai/tools/"; + + private AiToolRoutes() {} + + /** + * True when the request resolved to an AI document-tool endpoint. Prefers the matched route + * pattern (context-path independent, set by Spring MVC) and falls back to the raw request URI. + */ + public static boolean matches(HttpServletRequest request) { + Object pattern = request.getAttribute(HandlerMapping.BEST_MATCHING_PATTERN_ATTRIBUTE); + String path = pattern instanceof String s ? s : request.getRequestURI(); + return path != null && path.startsWith(PREFIX); + } +} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/cap/CapEvaluator.java b/app/saas/src/main/java/stirling/software/saas/payg/cap/CapEvaluator.java new file mode 100644 index 0000000000..d47b6e7365 --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/cap/CapEvaluator.java @@ -0,0 +1,101 @@ +package stirling.software.saas.payg.cap; + +import java.util.List; + +import stirling.software.saas.payg.model.EntitlementState; +import stirling.software.saas.payg.model.FeatureGate; +import stirling.software.saas.payg.model.FeatureSet; + +/** + * Pure-compute cap evaluation. Given a team's (or member's) spend, cap, and warn / degrade + * thresholds, returns the {@link EntitlementState}, the {@link FeatureSet} that should be in + * effect, and the corresponding enabled {@link FeatureGate}s. No DB access, no caches — the caller + * (entitlement service) supplies the inputs. + * + *

State transitions: + * + *

    + *
  • {@code capUnits == null} → {@code FULL} / {@link FeatureSet#FULL} unconditionally. + *
  • {@code spend / cap < warnPct} → {@code FULL}. + *
  • MINIMAL semantics: under DEGRADED+MINIMAL manual server-side tools (gated by {@link + * FeatureGate#OFFSITE_PROCESSING}) and client-side tools still work; only {@link + * FeatureGate#AUTOMATION} and {@link FeatureGate#AI_SUPPORT} are blocked. + *
  • {@code warnPct ≤ spend / cap < degradePct} → {@code WARNED}; feature set still {@link + * FeatureSet#FULL} — the warn band is a notification trigger, not a degradation. + *
  • {@code spend / cap ≥ degradePct} → {@code DEGRADED}; feature set drops to the policy's + * configured {@code degradedFeatureSet} (default {@link FeatureSet#MINIMAL}). + *
+ * + *

The percentage compare is integer math. We multiply spend by 100 before dividing — this keeps + * the precision and avoids floating-point on the hot path. Spend × 100 can overflow long at 9.2e16 + * units, which is not a realistic value (would represent quintillions of charged documents); we + * don't guard against it. + */ +public final class CapEvaluator { + + private CapEvaluator() {} + + /** + * Snapshot of one cap evaluation. The caller persists this into the appropriate {@code + * wallet_entitlement_snapshot} row (team-wide or per-member). + */ + public record Evaluation( + EntitlementState state, FeatureSet featureSet, List enabledGates) {} + + public static Evaluation evaluate( + long spendUnits, + Long capUnits, + int warnAtPct, + int degradeAtPct, + FeatureSet degradedFeatureSet) { + + if (capUnits == null || capUnits <= 0) { + return full(); + } + if (warnAtPct < 0 || degradeAtPct <= 0 || degradeAtPct < warnAtPct) { + // Defensive: misconfigured thresholds → treat as no-cap-effect to avoid surprise + // degradation. The admin endpoints that set the policy should validate; this + // protects the hot path from a bad row sneaking through. + return full(); + } + + // pct = floor((spend * 100) / cap). Integer arithmetic on the hot path. + long pct = (spendUnits * 100L) / capUnits; + + if (pct >= degradeAtPct) { + FeatureSet effective = + degradedFeatureSet != null ? degradedFeatureSet : FeatureSet.MINIMAL; + return new Evaluation(EntitlementState.DEGRADED, effective, gatesFor(effective)); + } + if (pct >= warnAtPct) { + // Warn band: still FULL feature set, but state flag is set so the FE can show a + // banner / send a notification. The wallet service emits a + // WalletEntitlementChanged event when state transitions; subscribers (email + // reminder, SSE to FE) act on that. + return new Evaluation( + EntitlementState.WARNED, FeatureSet.FULL, gatesFor(FeatureSet.FULL)); + } + return full(); + } + + /** Default enabled gates for a given feature set. */ + public static List gatesFor(FeatureSet set) { + if (set == null) { + return List.of(); + } + return switch (set) { + case FULL -> + List.of( + FeatureGate.OFFSITE_PROCESSING, + FeatureGate.AUTOMATION, + FeatureGate.AI_SUPPORT, + FeatureGate.CLIENT_SIDE); + case MINIMAL -> List.of(FeatureGate.OFFSITE_PROCESSING, FeatureGate.CLIENT_SIDE); + case CLIENT_ONLY -> List.of(FeatureGate.CLIENT_SIDE); + }; + } + + private static Evaluation full() { + return new Evaluation(EntitlementState.FULL, FeatureSet.FULL, gatesFor(FeatureSet.FULL)); + } +} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/cap/RequiresFeature.java b/app/saas/src/main/java/stirling/software/saas/payg/cap/RequiresFeature.java new file mode 100644 index 0000000000..97a61d0f30 --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/cap/RequiresFeature.java @@ -0,0 +1,46 @@ +package stirling.software.saas.payg.cap; + +import java.lang.annotation.ElementType; +import java.lang.annotation.Retention; +import java.lang.annotation.RetentionPolicy; +import java.lang.annotation.Target; + +import stirling.software.saas.payg.model.FeatureGate; + +/** + * Declares which {@link FeatureGate}(s) a controller method requires. Read at request time by + * {@code EntitlementGuard}; if any required gate is not in the team's currently-enabled gates the + * request is rejected with HTTP 402. + * + *

The annotation is not required on every endpoint. The guard's default rule is: + * + *

    + *
  • {@code @RequiresFeature} present → use exactly those gates. + *
  • No annotation, but the method has {@code @AutoJobPostMapping} → assume {@link + * FeatureGate#OFFSITE_PROCESSING}. + *
  • Neither → skip (admin endpoints, info, config — these don't accrue charges and shouldn't + * degrade). + *
+ * + *

So the only endpoints that need this annotation explicitly are those whose gate is + * different from the default {@code OFFSITE_PROCESSING} — chiefly {@code + * PipelineController} ({@link FeatureGate#AUTOMATION}) and the AI proxy layer ({@link + * FeatureGate#AI_SUPPORT}). Per-tool proliferation of the annotation is intentional non-goal. + * + *

Multiple gates declared = ALL must be enabled (AND, not OR). Realistic usage is single-gate; + * the array form is here for future combinations (e.g. an AI workflow inside a pipeline that needs + * both {@code AUTOMATION} and {@code AI_SUPPORT}). + * + *

{@code
+ * @RequiresFeature(FeatureGate.AUTOMATION)
+ * @AutoJobPostMapping("/pipeline")
+ * public ResponseEntity<...> runPipeline(@ModelAttribute PipelineRequest req) { ... }
+ * }
+ */ +@Target({ElementType.METHOD, ElementType.TYPE}) +@Retention(RetentionPolicy.RUNTIME) +public @interface RequiresFeature { + + /** One or more gates that must all be enabled for the request to proceed. */ + FeatureGate[] value(); +} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/charge/ChargeContext.java b/app/saas/src/main/java/stirling/software/saas/payg/charge/ChargeContext.java index ff4c94a138..f5914f8536 100644 --- a/app/saas/src/main/java/stirling/software/saas/payg/charge/ChargeContext.java +++ b/app/saas/src/main/java/stirling/software/saas/payg/charge/ChargeContext.java @@ -1,5 +1,6 @@ package stirling.software.saas.payg.charge; +import stirling.software.saas.payg.model.BillingCategory; import stirling.software.saas.payg.model.JobSource; import stirling.software.saas.payg.model.ProcessType; @@ -8,9 +9,18 @@ import stirling.software.saas.payg.model.ProcessType; * kind of process this is. Does NOT carry policy fields — the charge service resolves the effective * policy from {@code PricingPolicyService} so a stale snapshot from the caller can't desync from * the live policy. + * + *

{@code billingCategory} is the analytics axis for ledger + shadow rows and is determined by + * the interceptor before this context is built. Manual UI tools never reach {@code openProcess} + * (they short-circuit on {@link BillingCategory#BYPASSED}); any context constructed here therefore + * carries one of {@code API}, {@code AI}, or {@code AUTOMATION}. */ public record ChargeContext( - Long ownerUserId, Long ownerTeamId, JobSource source, ProcessType processType) { + Long ownerUserId, + Long ownerTeamId, + JobSource source, + ProcessType processType, + BillingCategory billingCategory) { public ChargeContext { if (ownerUserId == null) { @@ -22,5 +32,8 @@ public record ChargeContext( if (processType == null) { throw new IllegalArgumentException("processType is required"); } + if (billingCategory == null) { + throw new IllegalArgumentException("billingCategory is required"); + } } } diff --git a/app/saas/src/main/java/stirling/software/saas/payg/charge/JobChargeService.java b/app/saas/src/main/java/stirling/software/saas/payg/charge/JobChargeService.java index ee342f4e9d..ea3bb9cdab 100644 --- a/app/saas/src/main/java/stirling/software/saas/payg/charge/JobChargeService.java +++ b/app/saas/src/main/java/stirling/software/saas/payg/charge/JobChargeService.java @@ -2,12 +2,17 @@ package stirling.software.saas.payg.charge; import java.io.IOException; import java.nio.file.Path; +import java.time.LocalDateTime; import java.util.List; import java.util.Objects; +import java.util.Optional; +import java.util.UUID; import org.springframework.context.annotation.Profile; import org.springframework.stereotype.Service; import org.springframework.transaction.annotation.Transactional; +import org.springframework.transaction.support.TransactionSynchronization; +import org.springframework.transaction.support.TransactionSynchronizationManager; import org.springframework.web.multipart.MultipartFile; import lombok.extern.slf4j.Slf4j; @@ -17,11 +22,24 @@ import stirling.software.saas.payg.docs.DocumentMetrics; import stirling.software.saas.payg.job.JobContext; import stirling.software.saas.payg.job.JobService; import stirling.software.saas.payg.job.JoinOrOpenResult; +import stirling.software.saas.payg.job.ProcessingJob; +import stirling.software.saas.payg.meter.PaygMeterReportingService; +import stirling.software.saas.payg.model.BillingCategory; import stirling.software.saas.payg.model.JobSource; +import stirling.software.saas.payg.model.JobStatus; +import stirling.software.saas.payg.model.LedgerBucket; +import stirling.software.saas.payg.model.LedgerEntryType; +import stirling.software.saas.payg.model.ReferenceType; +import stirling.software.saas.payg.model.ShadowChargeStatus; +import stirling.software.saas.payg.policy.PaygTeamExtensions; import stirling.software.saas.payg.policy.PricingPolicy; import stirling.software.saas.payg.policy.PricingPolicyService; import stirling.software.saas.payg.repository.PaygShadowChargeRepository; +import stirling.software.saas.payg.repository.PaygTeamExtensionsRepository; +import stirling.software.saas.payg.repository.ProcessingJobRepository; +import stirling.software.saas.payg.repository.WalletLedgerRepository; import stirling.software.saas.payg.shadow.PaygShadowCharge; +import stirling.software.saas.payg.wallet.WalletLedgerEntry; /** * Orchestrates a tool call's open-process decision: look up the team's effective policy, resolve @@ -33,10 +51,9 @@ import stirling.software.saas.payg.shadow.PaygShadowCharge; * The real-charging path lives in a separate follow-up and reuses the same orchestration — only the * side-effect (shadow row vs ledger entry + Stripe call) differs. * - *

The {@code legacyCreditsCharged} field on the shadow row is set to {@code 0} here. When the - * legacy {@code CreditService} is wired to call this service (separate PR), the legacy debit amount - * becomes available and {@code diffPct} can be computed against it; until then the shadow row - * captures the PAYG units only. + *

The {@code legacyCreditsCharged} field on the shadow row is set to {@code 0}: the legacy + * credit engine has been removed, so there is no legacy debit to compare against and {@code + * diffPct} stays {@code 0}. The shadow row captures the PAYG units only. */ @Service @Profile("saas") @@ -47,16 +64,30 @@ public class JobChargeService { private final PricingPolicyService policyService; private final DocumentClassifier classifier; private final PaygShadowChargeRepository shadowRepository; + private final ProcessingJobRepository jobRepository; + private final PaygTeamExtensionsRepository teamExtensionsRepository; + private final PaygMeterReportingService meterReportingService; + private final WalletLedgerRepository ledgerRepository; public JobChargeService( JobService jobService, PricingPolicyService policyService, DocumentClassifier classifier, - PaygShadowChargeRepository shadowRepository) { + PaygShadowChargeRepository shadowRepository, + ProcessingJobRepository jobRepository, + PaygTeamExtensionsRepository teamExtensionsRepository, + PaygMeterReportingService meterReportingService, + WalletLedgerRepository ledgerRepository) { this.jobService = Objects.requireNonNull(jobService, "jobService"); this.policyService = Objects.requireNonNull(policyService, "policyService"); this.classifier = Objects.requireNonNull(classifier, "classifier"); this.shadowRepository = Objects.requireNonNull(shadowRepository, "shadowRepository"); + this.jobRepository = Objects.requireNonNull(jobRepository, "jobRepository"); + this.teamExtensionsRepository = + Objects.requireNonNull(teamExtensionsRepository, "teamExtensionsRepository"); + this.meterReportingService = + Objects.requireNonNull(meterReportingService, "meterReportingService"); + this.ledgerRepository = Objects.requireNonNull(ledgerRepository, "ledgerRepository"); } /** @@ -94,11 +125,115 @@ public class JobChargeService { int units = computeUnits(inputs, policy); result.job().setDocUnits(units); - recordShadowRow(ctx, result.job().getId(), policy.getId(), units); + int freeUsed = consumeFreeGrant(ctx, units); + recordShadowRow(ctx, result.job().getId(), policy.getId(), units, freeUsed); + recordLedgerDebit(ctx, result.job().getId(), policy.getId(), units); return new ChargeOutcome(result.job().getId(), units, ChargeOutcome.Disposition.OPENED); } + /** + * Charge a fixed number of units for a billable action that isn't file/lineage-driven — e.g. an + * AI Create session, billed once per document at session creation. Opens a standalone + * bookkeeping job (no lineage inputs, so follow-up calls never lineage-join it), draws the + * free-grant split, and writes the shadow + ledger rows exactly as {@link #openProcess} does, + * then closes the job so the paid portion meters to Stripe via the same {@code afterCommit} + * path and idempotency key ({@code process::close}). + * + *

Each call is independent: there is no join/dedup, so two sessions charge twice (correct — + * each is a distinct document). The caller passes the unit count; the policy {@code + * minChargeUnits} floor still applies. Must not be called for {@link BillingCategory#BYPASSED}. + * + * @return the bookkeeping job id (mostly useful for tests / tracing) + */ + @Transactional + public UUID chargeStandalone(ChargeContext ctx, int units) { + Objects.requireNonNull(ctx, "ctx"); + if (ctx.billingCategory() == BillingCategory.BYPASSED) { + throw new IllegalArgumentException("chargeStandalone must not be called for BYPASSED"); + } + + PricingPolicy policy = policyService.getEffectivePolicy(ctx.ownerTeamId()); + int chargeUnits = Math.max(units, policy.getMinChargeUnits()); + int stepLimit = resolveStepLimit(policy, ctx.source()); + + JobContext jobCtx = + new JobContext( + ctx.ownerUserId(), + ctx.ownerTeamId(), + ctx.source(), + ctx.processType(), + policy.getId(), + stepLimit); + ProcessingJob job = jobService.open(jobCtx, chargeUnits); + + int freeUsed = consumeFreeGrant(ctx, chargeUnits); + recordShadowRow(ctx, job.getId(), policy.getId(), chargeUnits, freeUsed); + recordLedgerDebit(ctx, job.getId(), policy.getId(), chargeUnits); + + // Close immediately — nothing will lineage-join a standalone job — so the paid portion + // meters via the same afterCommit hook + idempotency key as a normal process completion. + close(job.getId()); + return job.getId(); + } + + /** + * Draw this job's free portion from the team's one-time lifetime grant, atomically, and return + * the units taken (0..{@code units}); the remainder is the paid portion that will be metered to + * Stripe. Runs inside {@code openProcess}'s transaction with a pessimistic row lock so + * concurrent same-team charges split the grant exactly — no two jobs can both claim the last + * free unit. The grant is a soft floor: it never goes below 0, and the single job that crosses + * the boundary takes whatever's left (its remaining units bill). Skipped for non-billable / + * team-less calls (BYPASSED never reaches openProcess; guarded defensively). + */ + private int consumeFreeGrant(ChargeContext ctx, int units) { + BillingCategory category = ctx.billingCategory(); + if (category == null || category == BillingCategory.BYPASSED || ctx.ownerTeamId() == null) { + return 0; + } + Optional extOpt = + teamExtensionsRepository.findByIdForUpdate(ctx.ownerTeamId()); + if (extOpt.isEmpty()) { + return 0; + } + PaygTeamExtensions ext = extOpt.get(); + long remaining = ext.getFreeUnitsRemaining() == null ? 0L : ext.getFreeUnitsRemaining(); + int freeUsed = (int) Math.min(units, Math.max(0L, remaining)); + if (freeUsed > 0) { + ext.setFreeUnitsRemaining(remaining - freeUsed); + teamExtensionsRepository.save(ext); + } + return freeUsed; + } + + /** + * The live spend record. Everything the customer-facing side reads — the wallet endpoint's + * {@code spendUnitsThisPeriod}, the per-category breakdown ({@code wallet_category_summary} + * view), and the cap evaluator's period sum — derives from {@code wallet_ledger} DEBITs. Shadow + * rows are the comparison audit trail; this row is what actually counts. + * + *

Sign convention: debits are stored NEGATIVE (the entitlement snapshot negates the sum). + * Skipped for {@code BYPASSED} / uncategorised calls — manual UI work is never billed. + */ + private void recordLedgerDebit( + ChargeContext ctx, java.util.UUID jobId, Long policyId, int units) { + BillingCategory category = ctx.billingCategory(); + if (category == null || category == BillingCategory.BYPASSED) { + return; + } + WalletLedgerEntry entry = new WalletLedgerEntry(); + entry.setTeamId(ctx.ownerTeamId()); + entry.setActorUserId(ctx.ownerUserId()); + entry.setEntryType(LedgerEntryType.DEBIT); + entry.setBucket(LedgerBucket.CYCLE); + entry.setAmountUnits(-units); + entry.setReferenceType(ReferenceType.JOB); + entry.setReferenceId(jobId.toString()); + entry.setPolicyId(policyId); + entry.setBillingCategory(category); + ledgerRepository.save(entry); + } + private int resolveStepLimit(PricingPolicy policy, JobSource source) { Integer fromPolicy = policy.getStepLimits() == null ? null : policy.getStepLimits().get(source); @@ -117,10 +252,14 @@ public class JobChargeService { private int computeUnits(List inputs, PricingPolicy policy) { List multiparts = inputs.stream().map(JobInput::multipart).toList(); + // Reuse the temp file the caller already wrote (in PaygChargeInterceptor.preHandle for + // lineage hashing) instead of materialising the same bytes a second time inside the + // classifier. Saves one write + one read per PDF input. + List paths = inputs.stream().map(JobInput::path).toList(); DocumentMetrics metrics = multiparts.size() == 1 - ? classifier.classify(multiparts.get(0), policy) - : classifier.classify(multiparts, policy); + ? classifier.classify(multiparts.get(0), paths.get(0), policy) + : classifier.classify(multiparts, paths, policy); // Apply the policy-level minChargeUnits floor per design § 3.4. The classifier returns // raw docUnits with a "non-empty input → ≥1" floor; the charge formula's // max(min_charge_units, docUnits) layers on top. @@ -128,16 +267,264 @@ public class JobChargeService { } private void recordShadowRow( - ChargeContext ctx, java.util.UUID jobId, Long policyId, int units) { + ChargeContext ctx, + java.util.UUID jobId, + Long policyId, + int units, + int freeUnitsConsumed) { PaygShadowCharge row = new PaygShadowCharge(); row.setTeamId(ctx.ownerTeamId()); row.setJobId(jobId); row.setPolicyId(policyId); row.setPaygUnits(units); - // No legacy comparison yet — wired when the shadow path is connected to the legacy - // CreditService in the follow-up PR. Until then, diff stays at 0. + // Free-vs-paid split fixed at charge time: paid (metered) = paygUnits - freeUnitsConsumed, + // and a refund restores freeUnitsConsumed to the team's grant. + row.setFreeUnitsConsumed(freeUnitsConsumed); + // No legacy comparison: the legacy credit engine has been removed, so diff stays at 0. row.setLegacyCreditsCharged(0); row.setDiffPct(0); + row.setStatus(ShadowChargeStatus.CHARGED); + // PAYG analytics axis + caller surface — copied from ctx so the row stays self-describing + // after processing_job is pruned. Never affects what Stripe meters (single flat meter). + row.setBillingCategory(ctx.billingCategory()); + row.setJobSource(ctx.source()); shadowRepository.save(row); } + + /** + * First-step failure on a freshly-opened process: mimic a successful Stripe + * meter_event_adjustment(cancel) by flipping the shadow row to {@link + * ShadowChargeStatus#REFUNDED}, and close the process so a same-input retry can't lineage-join + * into a refunded chain for free work. + * + *

Idempotent: re-invoking on an already-REFUNDED row or already-CLOSED process is a silent + * no-op. + */ + @Transactional + public void markFirstStepFailed(UUID jobId, String refundReason) { + Objects.requireNonNull(jobId, "jobId"); + LocalDateTime now = LocalDateTime.now(); + + Optional rowOpt = shadowRepository.findFirstByJobIdOrderByIdAsc(jobId); + if (rowOpt.isEmpty()) { + log.debug( + "markFirstStepFailed: no shadow row for job {} (PAYG not active for team?)", + jobId); + } else { + PaygShadowCharge row = rowOpt.get(); + if (row.getStatus() != ShadowChargeStatus.REFUNDED) { + row.setStatus(ShadowChargeStatus.REFUNDED); + row.setRefundedAt(now); + row.setRefundReason(trimReason(refundReason)); + shadowRepository.save(row); + // Compensate the live ledger DEBIT written at openProcess so the period spend + // nets to zero for the failed work. Positive amount mirrors the negative debit; + // same JOB reference ties the pair together. The idempotency guard above (only + // on the CHARGED→REFUNDED transition) prevents double-credits on re-invocation. + BillingCategory category = row.getBillingCategory(); + if (category != null && category != BillingCategory.BYPASSED) { + WalletLedgerEntry refund = new WalletLedgerEntry(); + refund.setTeamId(row.getTeamId()); + refund.setEntryType(LedgerEntryType.REFUND); + refund.setBucket(LedgerBucket.CYCLE); + refund.setAmountUnits(row.getPaygUnits()); + refund.setReferenceType(ReferenceType.JOB); + refund.setReferenceId(jobId.toString()); + refund.setPolicyId(row.getPolicyId()); + refund.setBillingCategory(category); + ledgerRepository.save(refund); + // Hand back the free units this job consumed (first-step failures are + // pre-meter, so nothing was billed to Stripe — only the grant moved). Exactly + // what was taken at charge time, so the counter can't drift above the grant. + int freeConsumed = + row.getFreeUnitsConsumed() == null ? 0 : row.getFreeUnitsConsumed(); + if (freeConsumed > 0 && row.getTeamId() != null) { + teamExtensionsRepository.restoreFreeUnits(row.getTeamId(), freeConsumed); + } + } + } + } + + ProcessingJob job = jobRepository.findById(jobId).orElse(null); + if (job == null) { + log.warn("markFirstStepFailed: no ProcessingJob with id {}", jobId); + return; + } + if (job.getStatus() == JobStatus.OPEN) { + job.setStatus(JobStatus.CLOSED); + job.setClosedAt(now); + jobRepository.save(job); + } + } + + /** + * Closes a process and — as a fallback — meters its usage. The primary meter trigger is the + * charge interceptor's {@code afterCompletion} on a successful request (see {@link + * #meterJobUsage(UUID)}); this close-time meter exists to catch processes that were never + * cleanly completed (request thread died before {@code afterCompletion}) and are swept up later + * by {@code StaleJobCloser}. The deterministic idempotency key means a job already metered at + * completion is deduped here at Stripe, so the two paths never double-bill. + * + *

Idempotent w.r.t. process state (delegates to {@link JobService#close(UUID)}, which + * silently no-ops on an already-closed row). The meter POST runs in an {@code afterCommit} hook + * so a failed POST does not roll back the close; the reconciliation backfill (separate chunk) + * is the durability mechanism. + */ + @Transactional + public ProcessingJob close(UUID jobId) { + Objects.requireNonNull(jobId, "jobId"); + ProcessingJob closed = jobService.close(jobId); + + // The afterCommit hook only fires if there's an active transaction (Spring's + // @Transactional ensures that). If we're called outside one — e.g. a test using the raw + // bean — fall through with a debug log: the close() above already happened in a + // sub-transaction created by JobService, but the surrounding scope has no synchronization. + if (!TransactionSynchronizationManager.isSynchronizationActive()) { + log.debug("close({}): no active synchronization; skipping meter POST", jobId); + return closed; + } + + TransactionSynchronizationManager.registerSynchronization( + new TransactionSynchronization() { + @Override + public void afterCommit() { + try { + meterJobUsage(jobId); + } catch (RuntimeException e) { + // PaygMeterReportingService should already swallow; defence in depth so + // a thrown exception out of afterCommit doesn't leak past the + // synchronization boundary and bubble into the caller. + log.warn( + "afterCommit meter post for job {} threw unexpectedly: {}", + jobId, + e.getMessage()); + } + } + }); + + return closed; + } + + /** + * Post this job's billable usage to Stripe. The primary caller is the charge interceptor's + * {@code afterCompletion} on a successful OPENED request — i.e. the moment the work finishes — + * so the meter moves promptly. {@link #close(UUID)} also calls this from its {@code + * afterCommit} hook as the fallback for processes that were never cleanly completed (e.g. the + * request thread died); the deterministic idempotency key ({@code process::close}) makes + * the two paths dedup at Stripe, so a job metered at completion isn't billed again when it's + * later stale-closed. + * + *

Safe to call outside a transaction: it only reads (the job's openProcess DEBIT is already + * committed by the time either caller runs) and the POST is best-effort. Never throws — see + * {@link PaygMeterReportingService}. + * + *

Skips: no shadow row (not PAYG-tracked), REFUNDED row (first-step failure — never billed), + * BYPASSED/uncategorised, zero units, free-tier team (no Stripe customer), or usage still + * within the app-side free allowance. + */ + public void meterJobUsage(UUID jobId) { + Optional rowOpt = shadowRepository.findFirstByJobIdOrderByIdAsc(jobId); + if (rowOpt.isEmpty()) { + // No shadow row → not a PAYG-tracked job; nothing to meter. + return; + } + PaygShadowCharge row = rowOpt.get(); + if (row.getStatus() == ShadowChargeStatus.REFUNDED) { + // Refunded rows are zero-net charges; do not emit a meter event. + return; + } + BillingCategory category = row.getBillingCategory(); + if (category == null || category == BillingCategory.BYPASSED) { + // Defensive: BYPASSED rows shouldn't exist (interceptor short-circuits before + // openProcess), but tolerate if a future caller writes one. + log.debug("close({}): shadow row category={} → no meter event", jobId, category); + return; + } + Integer units = row.getPaygUnits(); + if (units == null || units <= 0) { + return; + } + Long teamId = row.getTeamId(); + if (teamId == null) { + return; + } + PaygTeamExtensions ext = teamExtensionsRepository.findById(teamId).orElse(null); + if (ext == null) { + return; + } + // payg_subscription_id is the single switch that says "this team is billed" (see + // PaygTeamExtensions). Gate on it directly now that V14 ships the column: a team with a + // Stripe customer but no live subscription — e.g. the brief window after checkout but + // before the subscription-created webhook lands — must not post meter events against a + // subscription that doesn't exist. A job finishing in that window is still metered later + // via the stale-close fallback, once the subscription has landed (same idempotency key). + String subscriptionId = ext.getPaygSubscriptionId(); + if (subscriptionId == null || subscriptionId.isBlank()) { + log.debug( + "close({}): team {} has no active subscription → no meter event", + jobId, + teamId); + return; + } + String stripeCustomerId = ext.getStripeCustomerId(); + if (stripeCustomerId == null || stripeCustomerId.isBlank()) { + // Subscribed but no customer id is a data inconsistency — we can't address the event. + log.warn( + "close({}): team {} has a subscription but no stripeCustomerId → cannot meter", + jobId, + teamId); + return; + } + + // Paid portion = units beyond the team's one-time free grant, fixed at charge time. The + // free grant is app-side only (Stripe's Prices are plain per-unit, no free tier), so the + // free units were already withheld when this row's free_units_consumed was set. + int freeConsumed = row.getFreeUnitsConsumed() == null ? 0 : row.getFreeUnitsConsumed(); + int paidUnits = units - freeConsumed; + if (paidUnits <= 0) { + log.debug( + "close({}): all {} units came from the free grant → no meter event", + jobId, + units); + return; + } + String idempotencyKey = "process:" + jobId + ":close"; + meterReportingService.recordUsage( + teamId, stripeCustomerId, paidUnits, category, idempotencyKey, jobId); + } + + /** + * Mid-chain 5xx on a JOINED step: return the step slot. The {@code lastStepAt} timestamp stays + * advanced (workflow window intentionally remains active for the next retry). No shadow-row + * change — only OPENED-disposition calls wrote a row. + * + *

Defensive lower bound: never drives {@code stepCount} below 1 (which would imply we'd + * decremented an already-decremented slot, or were called on a fresh process). + */ + @Transactional + public void decrementStepCount(UUID jobId) { + Objects.requireNonNull(jobId, "jobId"); + ProcessingJob job = jobRepository.findById(jobId).orElse(null); + if (job == null) { + log.warn("decrementStepCount: no ProcessingJob with id {}", jobId); + return; + } + int current = job.getStepCount() == null ? 0 : job.getStepCount(); + if (current <= 1) { + log.debug( + "decrementStepCount: stepCount already at {} for job {}; no-op", + current, + jobId); + return; + } + job.setStepCount(current - 1); + jobRepository.save(job); + } + + private static String trimReason(String reason) { + if (reason == null) { + return null; + } + return reason.length() > 128 ? reason.substring(0, 128) : reason; + } } diff --git a/app/saas/src/main/java/stirling/software/saas/payg/docs/DefaultDocumentClassifier.java b/app/saas/src/main/java/stirling/software/saas/payg/docs/DefaultDocumentClassifier.java index a46d975878..603f490a4f 100644 --- a/app/saas/src/main/java/stirling/software/saas/payg/docs/DefaultDocumentClassifier.java +++ b/app/saas/src/main/java/stirling/software/saas/payg/docs/DefaultDocumentClassifier.java @@ -4,6 +4,7 @@ import java.io.IOException; import java.io.InputStream; import java.io.OutputStream; import java.nio.file.Files; +import java.nio.file.Path; import java.util.List; import java.util.Objects; @@ -48,10 +49,16 @@ public class DefaultDocumentClassifier implements DocumentClassifier { @Override public DocumentMetrics classify(MultipartFile file, PricingPolicy policy) { + return classify(file, null, policy); + } + + @Override + public DocumentMetrics classify( + MultipartFile file, Path materialisedPath, PricingPolicy policy) { Objects.requireNonNull(file, "file"); Objects.requireNonNull(policy, "policy"); - FileFacts facts = inspect(file); + FileFacts facts = inspect(file, materialisedPath); long rawUnits = computeRawUnits(facts.pages, facts.bytes, policy); // toIntExact: fail loud on overflow rather than silently wrapping a billing number. int units = @@ -64,19 +71,35 @@ public class DefaultDocumentClassifier implements DocumentClassifier { @Override public DocumentMetrics classify(List files, PricingPolicy policy) { + return classify(files, null, policy); + } + + @Override + public DocumentMetrics classify( + List files, List materialisedPaths, PricingPolicy policy) { Objects.requireNonNull(files, "files"); Objects.requireNonNull(policy, "policy"); if (files.isEmpty()) { throw new IllegalArgumentException("files must not be empty"); } + if (materialisedPaths != null && materialisedPaths.size() != files.size()) { + throw new IllegalArgumentException( + "materialisedPaths size (" + + materialisedPaths.size() + + ") must equal files size (" + + files.size() + + ")"); + } int totalPages = 0; long totalBytes = 0; long rawUnitsSum = 0; String firstContentType = null; - for (MultipartFile file : files) { - FileFacts facts = inspect(file); + for (int i = 0; i < files.size(); i++) { + MultipartFile file = files.get(i); + Path path = materialisedPaths == null ? null : materialisedPaths.get(i); + FileFacts facts = inspect(file, path); // Sum the *raw* (unclamped) per-file units so the group cap below can actually bind. // Per-file clamping in this loop would make the group cap a no-op. rawUnitsSum = @@ -103,11 +126,17 @@ public class DefaultDocumentClassifier implements DocumentClassifier { totalUnits); } - private FileFacts inspect(MultipartFile file) { + private FileFacts inspect(MultipartFile file, Path materialisedPath) { long bytes = file.getSize(); String contentType = file.getContentType() != null ? file.getContentType() : DEFAULT_CONTENT_TYPE; - int pages = isPdf(contentType, file.getOriginalFilename()) ? readPageCount(file) : 0; + int pages = 0; + if (isPdf(contentType, file.getOriginalFilename())) { + pages = + materialisedPath != null + ? readPageCountFromPath(materialisedPath, file.getOriginalFilename()) + : readPageCount(file); + } return new FileFacts(pages, bytes, contentType); } @@ -141,9 +170,7 @@ public class DefaultDocumentClassifier implements DocumentClassifier { OutputStream out = Files.newOutputStream(temp.getPath())) { in.transferTo(out); } - try (PdfDocument doc = PdfDocument.open(temp.getPath())) { - return doc.pageCount(); - } + return readPageCountFromPath(temp.getPath(), file.getOriginalFilename()); } catch (IOException | RuntimeException e) { log.debug( "Could not read PDF page count for {} ({}); falling back to bytes-only units", @@ -153,6 +180,23 @@ public class DefaultDocumentClassifier implements DocumentClassifier { } } + /** + * Page-count read against an already-materialised file. Used by callers that already wrote the + * bytes to disk (the PAYG interceptor materialises every input for the lineage hash) so we + * avoid a second copy. + */ + private int readPageCountFromPath(Path path, String displayName) { + try (PdfDocument doc = PdfDocument.open(path)) { + return doc.pageCount(); + } catch (RuntimeException e) { + log.debug( + "Could not read PDF page count for {} ({}); falling back to bytes-only units", + displayName, + e.getClass().getSimpleName()); + return 0; + } + } + private static int saturatedAdd(int a, int b) { long sum = (long) a + b; if (sum > Integer.MAX_VALUE) { diff --git a/app/saas/src/main/java/stirling/software/saas/payg/docs/DocumentClassifier.java b/app/saas/src/main/java/stirling/software/saas/payg/docs/DocumentClassifier.java index c7af9deaf4..9c4966fde9 100644 --- a/app/saas/src/main/java/stirling/software/saas/payg/docs/DocumentClassifier.java +++ b/app/saas/src/main/java/stirling/software/saas/payg/docs/DocumentClassifier.java @@ -1,5 +1,6 @@ package stirling.software.saas.payg.docs; +import java.nio.file.Path; import java.util.List; import org.springframework.web.multipart.MultipartFile; @@ -11,6 +12,17 @@ import stirling.software.saas.payg.policy.PricingPolicy; * *

Returns {@code docUnits} with an absolute floor of 1 for non-empty input. {@code * policy.minChargeUnits} is applied at charge time, not here. + * + *

Each overload comes in two flavours: + * + *

    + *
  • {@code classify(MultipartFile, ...)} — classifier reads bytes via {@code getInputStream()} + * and writes its own temp file to feed jpdfium (for PDFs). Use when no on-disk copy exists. + *
  • {@code classify(MultipartFile, Path, ...)} — caller has already materialised the bytes to + * {@code Path}; classifier reads page count directly from there without re-writing. Hot-path + * callers (the PAYG interceptor) should use this form since they materialise inputs anyway + * for the lineage hash. + *
*/ public interface DocumentClassifier { @@ -22,4 +34,20 @@ public interface DocumentClassifier { * units, capped at {@code fileUnitCap × files.size()} and floored at 1. */ DocumentMetrics classify(List files, PricingPolicy policy); + + /** + * Same as {@link #classify(MultipartFile, PricingPolicy)} but uses {@code materialisedPath} for + * PDF page-count extraction, avoiding a second copy of the upload bytes to disk. Callers that + * hold an on-disk copy (e.g. the PAYG interceptor materialises every input for the lineage + * hash) should prefer this form. + */ + DocumentMetrics classify(MultipartFile file, Path materialisedPath, PricingPolicy policy); + + /** + * Multi-file variant of {@link #classify(MultipartFile, Path, PricingPolicy)}. {@code + * materialisedPaths} must align positionally with {@code files} — entry {@code i} is the + * on-disk copy of {@code files.get(i)}. + */ + DocumentMetrics classify( + List files, List materialisedPaths, PricingPolicy policy); } diff --git a/app/saas/src/main/java/stirling/software/saas/payg/entitlement/EntitlementGuard.java b/app/saas/src/main/java/stirling/software/saas/payg/entitlement/EntitlementGuard.java new file mode 100644 index 0000000000..54ebd24c04 --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/entitlement/EntitlementGuard.java @@ -0,0 +1,357 @@ +package stirling.software.saas.payg.entitlement; + +import java.io.IOException; +import java.util.Arrays; +import java.util.Collections; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import java.util.Optional; +import java.util.UUID; + +import org.springframework.context.annotation.Profile; +import org.springframework.core.annotation.AnnotationUtils; +import org.springframework.http.HttpHeaders; +import org.springframework.http.HttpStatus; +import org.springframework.http.MediaType; +import org.springframework.security.core.Authentication; +import org.springframework.security.core.context.SecurityContextHolder; +import org.springframework.stereotype.Component; +import org.springframework.web.method.HandlerMethod; +import org.springframework.web.servlet.HandlerInterceptor; + +import com.fasterxml.jackson.databind.ObjectMapper; + +import io.micrometer.core.instrument.Counter; +import io.micrometer.core.instrument.MeterRegistry; + +import jakarta.servlet.http.HttpServletRequest; +import jakarta.servlet.http.HttpServletResponse; + +import lombok.extern.slf4j.Slf4j; + +import stirling.software.common.annotations.AutoJobPostMapping; +import stirling.software.proprietary.security.database.repository.UserRepository; +import stirling.software.proprietary.security.model.ApiKeyAuthenticationToken; +import stirling.software.proprietary.security.model.User; +import stirling.software.saas.payg.cap.AiToolRoutes; +import stirling.software.saas.payg.cap.RequiresFeature; +import stirling.software.saas.payg.model.FeatureGate; +import stirling.software.saas.util.AuthenticationUtils; + +/** + * Hot-path entitlement check. Runs after {@code PaygChargeInterceptor} in the MVC chain and short- + * circuits the request before any handler work happens when the team's snapshot is missing one of + * the gates the route declared via {@link RequiresFeature}. + * + *

Scope: routes whose handler method (or bean type) carries either {@link AutoJobPostMapping} + * (multipart tool POSTs) or {@link RequiresFeature} (AI controllers, future non-multipart gated + * routes). Admin / info / config endpoints are excluded by the path-pattern in {@code + * PaygWebMvcConfig} and are additionally skipped here when they carry neither annotation, so non- + * billable infra never trips the guard. + * + *

Decision matrix: + * + *

+ * + * + * + * + * + *
authrequired gatessnapshot enabled?outcome
anonymousAUTOMATION or AI_SUPPORTn/a401 SIGNUP_REQUIRED
anonymousOFFSITE_PROCESSING / CLIENT_SIDEn/a200 (pass through)
authenticatedrequired ⊆ enabledyes200
authenticatedrequired ⊄ enabledno402 FEATURE_DEGRADED
+ * + *

Fail-open: any unexpected exception is logged at WARN and the request passes through. The cap + * pipeline must never block a customer because the guard tripped on a transient DB error. + */ +@Slf4j +@Component +@Profile("saas") +public class EntitlementGuard implements HandlerInterceptor { + + private static final FeatureGate[] DEFAULT_REQUIRED_GATES = {FeatureGate.OFFSITE_PROCESSING}; + + private final EntitlementService entitlementService; + private final UserRepository userRepository; + private final ObjectMapper objectMapper; + + private final Counter passCounter; + private final Counter deniedDegradedCounter; + private final Counter deniedPaygLimitCounter; + private final Counter deniedSignupRequiredCounter; + private final Counter errorsCounter; + private final Counter skippedNoAnnotationCounter; + + public EntitlementGuard( + EntitlementService entitlementService, + UserRepository userRepository, + MeterRegistry meterRegistry) { + this.entitlementService = entitlementService; + this.userRepository = userRepository; + this.objectMapper = new ObjectMapper(); + + this.passCounter = + Counter.builder("payg.entitlement.guard") + .tag("outcome", "pass") + .register(meterRegistry); + this.deniedDegradedCounter = + Counter.builder("payg.entitlement.guard") + .tag("outcome", "denied_degraded") + .register(meterRegistry); + this.deniedPaygLimitCounter = + Counter.builder("payg.entitlement.guard") + .tag("outcome", "denied_payg_limit") + .register(meterRegistry); + this.deniedSignupRequiredCounter = + Counter.builder("payg.entitlement.guard") + .tag("outcome", "denied_signup_required") + .register(meterRegistry); + this.skippedNoAnnotationCounter = + Counter.builder("payg.entitlement.guard") + .tag("outcome", "skipped") + .register(meterRegistry); + this.errorsCounter = + Counter.builder("payg.entitlement.guard.errors") + .description("EntitlementGuard internal failures (fail-open)") + .register(meterRegistry); + } + + @Override + public boolean preHandle( + HttpServletRequest request, HttpServletResponse response, Object handler) { + if (!(handler instanceof HandlerMethod hm)) { + return true; + } + // Scope: AutoJobPostMapping routes (multipart tool POSTs) OR routes that explicitly + // declare @RequiresFeature (e.g. AI controllers — JSON-bodied, no AutoJobPostMapping). + // Admin / info / config endpoints carry neither annotation and never trip the guard. + boolean hasAutoJobPostMapping = + AnnotationUtils.findAnnotation(hm.getMethod(), AutoJobPostMapping.class) != null + || AnnotationUtils.findAnnotation( + hm.getBeanType(), AutoJobPostMapping.class) + != null; + boolean hasRequiresFeature = + AnnotationUtils.findAnnotation(hm.getMethod(), RequiresFeature.class) != null + || AnnotationUtils.findAnnotation(hm.getBeanType(), RequiresFeature.class) + != null; + // AI document tools (/api/v1/ai/tools/**) live in the proprietary module and can't carry + // @RequiresFeature; recognise them by path so they're gated on AI_SUPPORT — see + // AiToolRoutes and PaygChargeInterceptor, which classify the same routes as AI. + boolean aiToolRoute = AiToolRoutes.matches(request); + if (!hasAutoJobPostMapping && !hasRequiresFeature && !aiToolRoute) { + skippedNoAnnotationCounter.increment(); + return true; + } + + FeatureGate[] required = + aiToolRoute ? new FeatureGate[] {FeatureGate.AI_SUPPORT} : resolveRequiredGates(hm); + Authentication auth = SecurityContextHolder.getContext().getAuthentication(); + + boolean anonymous = isAnonymous(auth); + boolean billable = isBillable(required); + + if (anonymous) { + if (billable) { + return write401SignupRequired(response, required); + } + // Anonymous user calling a manual / OFFSITE-only tool — let it through; PAYG only + // charges authenticated requests. + passCounter.increment(); + return true; + } + + Long teamId; + try { + teamId = resolveTeamId(auth); + } catch (RuntimeException e) { + log.warn("EntitlementGuard resolveTeamId failed; passing through", e); + errorsCounter.increment(); + return true; + } + if (teamId == null) { + // Defensive: authenticated principal with no team — shouldn't happen post-migration, + // but we don't want to lock those users out. PaygChargeInterceptor short-circuits the + // same shape upstream. + passCounter.increment(); + return true; + } + + EntitlementSnapshot snapshot; + try { + snapshot = entitlementService.getSnapshot(teamId); + } catch (RuntimeException e) { + log.warn("EntitlementGuard getSnapshot failed for team {}; passing through", teamId, e); + errorsCounter.increment(); + return true; + } + + // API-key calls are always billable usage (BillingCategory.API) — there is no "free + // manual" path for a programmatic client the way there is for a JWT/web user, whose + // everyday tool calls are BYPASSED and never reach a gate. So once the team is over its + // free allowance / spending cap (DEGRADED), every API-key call hard-stops, regardless of + // which gate the route declares. The gate loop below would otherwise wave through an API + // call to a plain server tool (it needs only OFFSITE_PROCESSING, which survives DEGRADED), + // letting an unsubscribed team keep consuming the API for free past its allowance. + if (auth instanceof ApiKeyAuthenticationToken && snapshot.isDegraded()) { + return write402PaygLimitReached(response, snapshot); + } + + List enabled = snapshot.enabledGates(); + for (FeatureGate gate : required) { + if (enabled == null || !enabled.contains(gate)) { + return write402FeatureDegraded(response, required, snapshot); + } + } + passCounter.increment(); + return true; + } + + static FeatureGate[] resolveRequiredGates(HandlerMethod hm) { + RequiresFeature ann = AnnotationUtils.findAnnotation(hm.getMethod(), RequiresFeature.class); + if (ann == null) { + ann = AnnotationUtils.findAnnotation(hm.getBeanType(), RequiresFeature.class); + } + if (ann != null && ann.value().length > 0) { + return ann.value(); + } + return DEFAULT_REQUIRED_GATES; + } + + private static boolean isAnonymous(Authentication auth) { + if (auth == null || !auth.isAuthenticated()) { + return true; + } + // Spring's anonymous filter installs a token whose name is "anonymousUser". + return "anonymousUser".equals(auth.getName()); + } + + private static boolean isBillable(FeatureGate[] required) { + for (FeatureGate g : required) { + if (g == FeatureGate.AUTOMATION || g == FeatureGate.AI_SUPPORT) { + return true; + } + } + return false; + } + + private Long resolveTeamId(Authentication auth) { + if (auth instanceof ApiKeyAuthenticationToken + && auth.getPrincipal() instanceof User apiUser) { + return apiUser.getTeam() == null ? null : apiUser.getTeam().getId(); + } + String supabaseId = AuthenticationUtils.extractSupabaseId(auth); + if (supabaseId == null) { + return null; + } + UUID supabaseUuid; + try { + supabaseUuid = UUID.fromString(supabaseId); + } catch (IllegalArgumentException e) { + // Username-style principals (legacy local accounts) — no Supabase ID to look up. Skip. + return null; + } + return userRepository + .findBySupabaseId(supabaseUuid) + .map(u -> u.getTeam() == null ? null : u.getTeam().getId()) + .orElse(null); + } + + private boolean write401SignupRequired(HttpServletResponse response, FeatureGate[] required) { + deniedSignupRequiredCounter.increment(); + Map body = new LinkedHashMap<>(); + body.put("error", "SIGNUP_REQUIRED"); + body.put("category", inferCategory(required)); + writeJson(response, HttpStatus.UNAUTHORIZED, body); + return false; + } + + /** + * 402 for a billable API-key call once the team is over its allowance / cap. The message is + * tailored by subscription state: an un-subscribed team is told to subscribe (their free + * allowance is spent); a subscribed team is told it hit its own spending cap. Programmatic + * clients get a stable {@code error} code plus the spend/cap numbers so they can surface + * something actionable. + */ + private boolean write402PaygLimitReached( + HttpServletResponse response, EntitlementSnapshot snapshot) { + deniedPaygLimitCounter.increment(); + Map body = new LinkedHashMap<>(); + body.put("error", "PAYG_LIMIT_REACHED"); + body.put("subscribed", snapshot.subscribed()); + body.put( + "message", + snapshot.subscribed() + ? "Your team has reached its monthly spending cap. Raise the cap to" + + " continue, or wait for it to reset next billing period." + : "Your team has used its free document allowance." + + " Subscribe to continue using the API."); + body.put("state", snapshot.state().name()); + body.put("spendUnits", snapshot.periodSpendUnits()); + body.put("capUnits", snapshot.periodCapUnits()); + body.put( + "periodEnd", + Optional.ofNullable(snapshot.periodEnd()).map(Object::toString).orElse(null)); + writeJson(response, HttpStatus.PAYMENT_REQUIRED, body); + return false; + } + + private boolean write402FeatureDegraded( + HttpServletResponse response, FeatureGate[] required, EntitlementSnapshot snapshot) { + deniedDegradedCounter.increment(); + Map body = new LinkedHashMap<>(); + body.put("error", "FEATURE_DEGRADED"); + // subscribed tells the client which usage-limit modal to show: a subscribed team is over + // its spending cap; an un-subscribed one has spent its free allowance. (PAYG_LIMIT_REACHED + // already carries this; mirror it here so the JWT/web path can pick the right modal too.) + body.put("subscribed", snapshot.subscribed()); + body.put("missingGates", missingGates(required, snapshot.enabledGates())); + body.put("state", snapshot.state().name()); + body.put( + "periodEnd", + Optional.ofNullable(snapshot.periodEnd()).map(Object::toString).orElse(null)); + body.put("capUnits", snapshot.periodCapUnits()); + body.put("spendUnits", snapshot.periodSpendUnits()); + writeJson(response, HttpStatus.PAYMENT_REQUIRED, body); + return false; + } + + private static List missingGates(FeatureGate[] required, List enabled) { + List enabledOrEmpty = enabled == null ? Collections.emptyList() : enabled; + return Arrays.stream(required) + .filter(g -> !enabledOrEmpty.contains(g)) + .map(Enum::name) + .toList(); + } + + private static String inferCategory(FeatureGate[] required) { + // Mirrors PaygChargeInterceptor.determineCategory precedence: AUTOMATION dominates AI. + for (FeatureGate g : required) { + if (g == FeatureGate.AUTOMATION) { + return "AUTOMATION"; + } + } + for (FeatureGate g : required) { + if (g == FeatureGate.AI_SUPPORT) { + return "AI"; + } + } + return "OFFSITE_PROCESSING"; + } + + private void writeJson( + HttpServletResponse response, HttpStatus status, Map body) { + response.setStatus(status.value()); + response.setContentType(MediaType.APPLICATION_JSON_VALUE); + response.setCharacterEncoding("UTF-8"); + try { + byte[] payload = objectMapper.writeValueAsBytes(body); + response.setHeader(HttpHeaders.CONTENT_LENGTH, Integer.toString(payload.length)); + response.getOutputStream().write(payload); + response.getOutputStream().flush(); + } catch (IOException e) { + // Container will fall back to its default error page — we did set the status code, + // so the client still sees the right HTTP code even if the body fails to write. + log.warn("EntitlementGuard write response body failed", e); + errorsCounter.increment(); + } + } +} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/entitlement/EntitlementService.java b/app/saas/src/main/java/stirling/software/saas/payg/entitlement/EntitlementService.java new file mode 100644 index 0000000000..397a714bfa --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/entitlement/EntitlementService.java @@ -0,0 +1,182 @@ +package stirling.software.saas.payg.entitlement; + +import java.time.Duration; +import java.time.LocalDateTime; +import java.time.YearMonth; +import java.util.List; +import java.util.Objects; +import java.util.Optional; + +import org.springframework.context.annotation.Profile; +import org.springframework.stereotype.Service; +import org.springframework.transaction.annotation.Transactional; + +import com.github.benmanes.caffeine.cache.Cache; +import com.github.benmanes.caffeine.cache.Caffeine; + +import lombok.extern.slf4j.Slf4j; + +import stirling.software.saas.payg.billing.TeamBillingContext; +import stirling.software.saas.payg.billing.TeamBillingService; +import stirling.software.saas.payg.cap.CapEvaluator; +import stirling.software.saas.payg.cap.CapEvaluator.Evaluation; +import stirling.software.saas.payg.model.EntitlementState; +import stirling.software.saas.payg.model.FeatureSet; +import stirling.software.saas.payg.repository.WalletLedgerRepository; +import stirling.software.saas.payg.repository.WalletPolicyRepository; +import stirling.software.saas.payg.wallet.WalletPolicy; + +/** + * Hot-path entitlement lookup. Returns the {@link EntitlementSnapshot} for a team: the billing + * facts (window, free allowance, document cap) come from {@link TeamBillingService}; this service + * layers the period spend (ledger SUM over that window) and the warn/degrade evaluation on top. + * + *

Backed by a per-team Caffeine cache with {@value #CACHE_TTL_SECONDS}s TTL and {@value + * #CACHE_MAX_SIZE}-entry cap. The TTL is the correctness floor — a cap change becomes visible on + * every instance within that window without coordination. Mutators (wallet policy admin updates, + * subscription webhook handlers) call {@link #invalidate(Long)} to drop a single team's entry + * immediately on the originating instance. + */ +@Slf4j +@Service +@Profile("saas") +public class EntitlementService { + + static final int CACHE_TTL_SECONDS = 30; + private static final int CACHE_MAX_SIZE = 10_000; + + private static final int WARN_AT_PCT = 80; + private static final int DEGRADE_AT_PCT = 100; + + private final TeamBillingService teamBillingService; + private final WalletPolicyRepository walletPolicyRepository; + private final WalletLedgerRepository ledgerRepository; + + private final Cache snapshotCache; + + public EntitlementService( + TeamBillingService teamBillingService, + WalletPolicyRepository walletPolicyRepository, + WalletLedgerRepository ledgerRepository) { + this.teamBillingService = Objects.requireNonNull(teamBillingService, "teamBillingService"); + this.walletPolicyRepository = + Objects.requireNonNull(walletPolicyRepository, "walletPolicyRepository"); + this.ledgerRepository = Objects.requireNonNull(ledgerRepository, "ledgerRepository"); + this.snapshotCache = + Caffeine.newBuilder() + .maximumSize(CACHE_MAX_SIZE) + .expireAfterWrite(Duration.ofSeconds(CACHE_TTL_SECONDS)) + .recordStats() + .build(); + } + + /** + * Returns the entitlement snapshot for {@code teamId}. Caches per-team for {@value + * #CACHE_TTL_SECONDS}s — burst requests share a single SUM query against the ledger. + * + *

{@code null} teamId throws — the guard short-circuits team-less requests upstream so a + * null reach here is a programming error. + */ + public EntitlementSnapshot getSnapshot(Long teamId) { + Objects.requireNonNull(teamId, "teamId"); + return snapshotCache.get(teamId, this::computeSnapshot); + } + + /** + * Drops {@code teamId}'s cache entry. Call after subscription state changes (webhook handlers), + * cap edits, or manual ledger adjustments so the next read recomputes immediately rather than + * waiting out the TTL. Also drops the underlying billing context so window/cap facts recompute + * together with the spend. + */ + public void invalidate(Long teamId) { + if (teamId != null) { + snapshotCache.invalidate(teamId); + teamBillingService.invalidate(teamId); + } + } + + /** Visible for tests. */ + long cacheSize() { + return snapshotCache.estimatedSize(); + } + + @Transactional(readOnly = true) + EntitlementSnapshot computeSnapshot(Long teamId) { + TeamBillingContext billing = teamBillingService.forTeam(teamId); + Optional walletPolicyOpt = walletPolicyRepository.findByTeamId(teamId); + + FeatureSet degradedSet = + walletPolicyOpt.map(WalletPolicy::getDegradedFeatureSet).orElse(FeatureSet.MINIMAL); + int warnAtPct = + walletPolicyOpt + .map(WalletPolicy::getWarnAtPct) + .filter(Objects::nonNull) + .orElse(WARN_AT_PCT); + int degradeAtPct = + walletPolicyOpt + .map(WalletPolicy::getDegradeAtPct) + .filter(Objects::nonNull) + .orElse(DEGRADE_AT_PCT); + + // Subscription-anchored window when subscribed; calendar month otherwise. Used for the + // subscribed monthly cap + the displayed billing period. + LocalDateTime periodStart = billing.periodStart(); + LocalDateTime periodEnd = billing.periodEnd(); + + Evaluation eval; + long snapshotSpend; + Long snapshotCap; + + if (billing.subscribed()) { + // Subscribed: gate on the monthly spending cap. Spend = this period's net billable + // documents (DEBIT minus REFUND so a refunded job doesn't read as spent). The one-time + // free grant doesn't gate a paying team — it only reduced what they were metered. + long signedNet = ledgerRepository.sumPeriodNetBillable(teamId, periodStart, periodEnd); + long periodSpend = signedNet < 0 ? -signedNet : 0L; + Long cap = billing.monthlyCapDocUnits(); + eval = CapEvaluator.evaluate(periodSpend, cap, warnAtPct, degradeAtPct, degradedSet); + snapshotSpend = periodSpend; + snapshotCap = cap; + } else { + // Unsubscribed: gate on the one-time lifetime free grant. Exhausted (remaining ≤ 0, or + // no grant configured) → DEGRADED so billable categories hard-stop; otherwise evaluate + // the warn/degrade band on used-of-grant. + long grant = billing.freeGrantUnits(); + long remaining = billing.freeRemainingUnits(); + long used = Math.max(0L, grant - remaining); + if (remaining <= 0L) { + eval = + new Evaluation( + EntitlementState.DEGRADED, + degradedSet, + CapEvaluator.gatesFor(degradedSet)); + } else { + eval = CapEvaluator.evaluate(used, grant, warnAtPct, degradeAtPct, degradedSet); + } + snapshotSpend = used; + snapshotCap = grant; + } + + return new EntitlementSnapshot( + eval.state(), + eval.featureSet(), + List.copyOf(eval.enabledGates()), + snapshotSpend, + snapshotCap, + periodStart, + periodEnd, + billing.subscribed()); + } + + /** + * Inclusive-start / exclusive-end window for the calendar-month period. Test seam — takes a + * clock value so tests don't race the calendar boundary. The live snapshot window comes from + * {@link TeamBillingService}; this remains for the forthcoming {@code BILLING_CYCLE} work. + */ + static LocalDateTime[] currentMonthWindow(LocalDateTime now) { + YearMonth ym = YearMonth.from(now); + LocalDateTime start = ym.atDay(1).atStartOfDay(); + LocalDateTime end = ym.plusMonths(1).atDay(1).atStartOfDay(); + return new LocalDateTime[] {start, end}; + } +} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/entitlement/EntitlementSnapshot.java b/app/saas/src/main/java/stirling/software/saas/payg/entitlement/EntitlementSnapshot.java new file mode 100644 index 0000000000..b3d44f3c71 --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/entitlement/EntitlementSnapshot.java @@ -0,0 +1,46 @@ +package stirling.software.saas.payg.entitlement; + +import java.time.LocalDateTime; +import java.util.List; + +import stirling.software.saas.payg.model.EntitlementState; +import stirling.software.saas.payg.model.FeatureGate; +import stirling.software.saas.payg.model.FeatureSet; + +/** + * Immutable snapshot of a team's entitlement state as of a single point in time. Returned by {@link + * EntitlementService#getSnapshot(Long)} and consumed by {@code EntitlementGuard}. + * + *

Contrast with {@link WalletEntitlementSnapshot}: the JPA entity is the persisted + * snapshot that the recompute path writes (one row per team, optionally per member). This record is + * the computed-now view the hot-path guard reads — backed by a 30s Caffeine cache so a + * request burst doesn't hammer the ledger SUM. + * + * @param state aggregate state — FULL, WARNED, or DEGRADED. + * @param featureSet bundle name in effect (FULL on no-cap / warn band; degraded set on DEGRADED). + * @param enabledGates the gates the guard checks against — request proceeds only if every required + * gate is in this list. + * @param periodSpendUnits sum of debited units in {@code [periodStart, periodEnd)}, in canonical + * doc-units (positive). + * @param periodCapUnits the cap applied — free-tier units for un-subscribed teams, {@code + * wallet_policy.cap_units} for subscribed teams. {@code null} means uncapped. + * @param periodStart inclusive start of the current cap period. + * @param periodEnd exclusive end of the current cap period. + * @param subscribed whether the team has an active PAYG subscription. Drives the messaging when a + * billable call is hard-stopped: an un-subscribed team is told to subscribe; a subscribed team + * that hit its self-set spending cap is told to raise it. + */ +public record EntitlementSnapshot( + EntitlementState state, + FeatureSet featureSet, + List enabledGates, + long periodSpendUnits, + Long periodCapUnits, + LocalDateTime periodStart, + LocalDateTime periodEnd, + boolean subscribed) { + + public boolean isDegraded() { + return state == EntitlementState.DEGRADED; + } +} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/filter/PaygChargeInterceptor.java b/app/saas/src/main/java/stirling/software/saas/payg/filter/PaygChargeInterceptor.java new file mode 100644 index 0000000000..eda04a4789 --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/filter/PaygChargeInterceptor.java @@ -0,0 +1,592 @@ +package stirling.software.saas.payg.filter; + +import java.io.IOException; +import java.io.InputStream; +import java.io.OutputStream; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.ArrayList; +import java.util.Collections; +import java.util.List; +import java.util.Optional; +import java.util.UUID; + +import org.springframework.context.annotation.Profile; +import org.springframework.core.annotation.AnnotationUtils; +import org.springframework.security.core.Authentication; +import org.springframework.security.core.context.SecurityContextHolder; +import org.springframework.stereotype.Component; +import org.springframework.util.MultiValueMap; +import org.springframework.web.method.HandlerMethod; +import org.springframework.web.multipart.MultipartFile; +import org.springframework.web.multipart.MultipartHttpServletRequest; +import org.springframework.web.servlet.AsyncHandlerInterceptor; +import org.springframework.web.servlet.HandlerMapping; + +import io.micrometer.core.instrument.Counter; +import io.micrometer.core.instrument.MeterRegistry; +import io.micrometer.core.instrument.Timer; + +import jakarta.servlet.http.HttpServletRequest; +import jakarta.servlet.http.HttpServletResponse; + +import lombok.extern.slf4j.Slf4j; + +import stirling.software.common.annotations.AutoJobPostMapping; +import stirling.software.common.util.TempFile; +import stirling.software.common.util.TempFileManager; +import stirling.software.proprietary.security.database.repository.UserRepository; +import stirling.software.proprietary.security.model.ApiKeyAuthenticationToken; +import stirling.software.proprietary.security.model.User; +import stirling.software.saas.payg.cap.AiToolRoutes; +import stirling.software.saas.payg.cap.RequiresFeature; +import stirling.software.saas.payg.charge.ChargeContext; +import stirling.software.saas.payg.charge.ChargeOutcome; +import stirling.software.saas.payg.charge.JobChargeService; +import stirling.software.saas.payg.charge.JobInput; +import stirling.software.saas.payg.job.JobService; +import stirling.software.saas.payg.model.BillingCategory; +import stirling.software.saas.payg.model.FeatureGate; +import stirling.software.saas.payg.model.JobSource; +import stirling.software.saas.payg.model.JobStepStatus; +import stirling.software.saas.payg.model.ProcessType; +import stirling.software.saas.util.AuthenticationUtils; + +/** + * The hot-path PAYG interceptor, registered in {@code PaygWebMvcConfig}. + * + *

{@code preHandle}: gates on {@code @AutoJobPostMapping} OR {@code @RequiresFeature} (the + * latter lets AI controllers — JSON-bodied, no AutoJobPostMapping — bill correctly), reads the + * parsed multipart parts, materialises each input to a {@code TempFile}, and asks {@link + * JobChargeService#openProcess} to open (or join) a process. The resulting {@link ChargeOutcome} + * plus input temp-files are stashed as request attributes for {@code afterCompletion}. Routes + * without multipart inputs short-circuit inside {@code doPreHandle} without touching the charge + * service. + * + *

{@code afterCompletion}: branches on HTTP status — 2xx hashes the response body for OUTPUT + * lineage; 4xx records a step append for audit; 5xx triggers refund-and-close (OPENED) or + * step-quota return (JOINED). Closes all input temp files and the response wrapper at the end. + * + *

Fail-open everywhere: any unexpected {@link RuntimeException} is swallowed, logged at WARN, + * and counted on {@code payg.filter.errors}. The customer's tool call always proceeds. + * + *

Async controllers ({@code DeferredResult}, {@code CompletableFuture}) are handled + * transparently — {@link AsyncHandlerInterceptor#afterConcurrentHandlingStarted} is a no-op; the + * normal {@code afterCompletion} fires when the async work resolves. + */ +@Slf4j +@Component +@Profile("saas") +public class PaygChargeInterceptor implements AsyncHandlerInterceptor { + + static final String ATTR_JOB_ID = PaygChargeInterceptor.class.getName() + ".JOB_ID"; + static final String ATTR_DISPOSITION = PaygChargeInterceptor.class.getName() + ".DISPOSITION"; + static final String ATTR_INPUT_TEMP_FILES = + PaygChargeInterceptor.class.getName() + ".INPUT_TEMP_FILES"; + static final String ATTR_INPUT_BYTES = PaygChargeInterceptor.class.getName() + ".INPUT_BYTES"; + static final String ATTR_FAILED = PaygChargeInterceptor.class.getName() + ".FAILED"; + static final String ATTR_TOOL_ID = PaygChargeInterceptor.class.getName() + ".TOOL_ID"; + + private static final String AUTOMATION_HEADER = "X-Stirling-Automation"; + + /** + * Optional header the Tauri desktop shell sets so saas-side traffic from the embedded client + * can be classified as {@link JobSource#DESKTOP_APP} instead of {@code WEB}. No anti-spoof — + * V12 step limits for DESKTOP_APP and WEB are identical, so the worst-case abuse value is zero + * today. Tighten if/when their limits diverge. + */ + private static final String DESKTOP_CLIENT_HEADER = "X-Stirling-Client"; + + /** Matches {@code processing_job_step.tool_id} column width (VARCHAR(128)). */ + private static final int TOOL_ID_MAX_LENGTH = 128; + + private final JobChargeService chargeService; + private final JobService jobService; + private final UserRepository userRepository; + private final TempFileManager tempFileManager; + private final PaygOutputExtractor outputExtractor; + private final PaygFilterProperties properties; + + private final Counter errorsCounter; + private final Counter callsOpened; + private final Counter callsJoined; + private final Counter callsShortCircuit; + private final Counter callsBypassed; + private final Counter refundsCounter; + + /** preHandle wall-clock per request. Separate from afterCompletion — different populations. */ + private final Timer preHandleTimer; + + /** afterCompletion wall-clock per request. Includes response hashing + step append + refund. */ + private final Timer afterCompletionTimer; + + public PaygChargeInterceptor( + JobChargeService chargeService, + JobService jobService, + UserRepository userRepository, + TempFileManager tempFileManager, + PaygOutputExtractor outputExtractor, + PaygFilterProperties properties, + MeterRegistry meterRegistry) { + this.chargeService = chargeService; + this.jobService = jobService; + this.userRepository = userRepository; + this.tempFileManager = tempFileManager; + this.outputExtractor = outputExtractor; + this.properties = properties; + + this.errorsCounter = + Counter.builder("payg.filter.errors") + .description("PAYG interceptor / filter internal failures") + .register(meterRegistry); + this.callsOpened = + Counter.builder("payg.filter.calls") + .tag("disposition", "OPENED") + .register(meterRegistry); + this.callsJoined = + Counter.builder("payg.filter.calls") + .tag("disposition", "JOINED") + .register(meterRegistry); + this.callsShortCircuit = + Counter.builder("payg.filter.calls") + .tag("disposition", "SHORT_CIRCUIT") + .register(meterRegistry); + this.callsBypassed = + Counter.builder("payg.filter.bypassed") + .description( + "Manual UI tool calls that skipped openProcess (BillingCategory.BYPASSED)") + .register(meterRegistry); + this.refundsCounter = + Counter.builder("payg.filter.refunds") + .description("First-step 5xx refunds applied to shadow rows") + .register(meterRegistry); + this.preHandleTimer = + Timer.builder("payg.filter.duration") + .tag("phase", "preHandle") + .description("preHandle wall-clock per request") + .register(meterRegistry); + this.afterCompletionTimer = + Timer.builder("payg.filter.duration") + .tag("phase", "afterCompletion") + .description("afterCompletion wall-clock per request") + .register(meterRegistry); + } + + @Override + public boolean preHandle( + HttpServletRequest request, HttpServletResponse response, Object handler) { + Timer.Sample sample = Timer.start(); + try { + if (!properties.isEnabled()) { + return true; + } + if (!(handler instanceof HandlerMethod hm)) { + callsShortCircuit.increment(); + return true; + } + // In-scope when the handler carries @AutoJobPostMapping (multipart tool POSTs) OR + // @RequiresFeature (AI controllers, future non-multipart gated routes). Without one of + // these the interceptor short-circuits — admin / info / static routes never run + // determineCategory. + boolean hasAutoJobPostMapping = + AnnotationUtils.findAnnotation(hm.getMethod(), AutoJobPostMapping.class) != null + || AnnotationUtils.findAnnotation( + hm.getBeanType(), AutoJobPostMapping.class) + != null; + boolean hasRequiresFeature = + AnnotationUtils.findAnnotation(hm.getMethod(), RequiresFeature.class) != null + || AnnotationUtils.findAnnotation( + hm.getBeanType(), RequiresFeature.class) + != null; + // AI document tools (/api/v1/ai/tools/**) live in the proprietary module and can't + // carry @RequiresFeature, so they're recognised by path — see AiToolRoutes. + boolean aiToolRoute = AiToolRoutes.matches(request); + if (!hasAutoJobPostMapping && !hasRequiresFeature && !aiToolRoute) { + callsShortCircuit.increment(); + return true; + } + // Bypass fast-path: determine the BillingCategory BEFORE any multipart + // materialisation or openProcess call. Manual UI tool calls (BYPASSED) skip the + // entire ledger/shadow pipeline — no temp files, no DB writes. + Authentication auth = SecurityContextHolder.getContext().getAuthentication(); + BillingCategory category = determineCategory(hm, request, auth); + if (category == BillingCategory.BYPASSED) { + callsBypassed.increment(); + return true; + } + try { + doPreHandle(request, auth, category); + } catch (RuntimeException e) { + log.warn("PAYG preHandle failed; passing through unbilled", e); + errorsCounter.increment(); + request.setAttribute(ATTR_FAILED, Boolean.TRUE); + cleanupInputs(request); + } + return true; + } finally { + sample.stop(preHandleTimer); + } + } + + private void doPreHandle( + HttpServletRequest request, Authentication auth, BillingCategory category) { + User currentUser = resolveUser(auth); + if (currentUser == null) { + callsShortCircuit.increment(); + return; + } + + if (!(request instanceof MultipartHttpServletRequest mreq)) { + callsShortCircuit.increment(); + return; + } + + MultiValueMap map = mreq.getMultiFileMap(); + List nonEmpty = new ArrayList<>(); + for (List bucket : map.values()) { + for (MultipartFile mp : bucket) { + if (mp.getSize() > 0) { + nonEmpty.add(mp); + } + } + } + if (nonEmpty.isEmpty()) { + callsShortCircuit.increment(); + return; + } + + List tempFiles = new ArrayList<>(nonEmpty.size()); + List inputs = new ArrayList<>(nonEmpty.size()); + long totalInputBytes = 0L; + try { + for (MultipartFile mp : nonEmpty) { + TempFile tf = tempFileManager.createManagedTempFile(".upload"); + tempFiles.add(tf); + try (InputStream in = mp.getInputStream(); + OutputStream out = Files.newOutputStream(tf.getPath())) { + in.transferTo(out); + } + inputs.add(new JobInput(mp, tf.getPath())); + totalInputBytes += mp.getSize(); + } + } catch (IOException e) { + for (TempFile tf : tempFiles) { + tf.close(); + } + throw new RuntimeException("Failed to materialise multipart input", e); + } + // Publish as an immutable list. preHandle fully populates `tempFiles` before this + // setAttribute, and cleanupInputs / afterCompletion are read-only consumers — exposing + // an unmodifiable view (a) makes the safe-publication guarantee from the container's + // attribute-map synchronization unambiguous to readers on different threads (sync vs + // async-dispatch), and (b) hard-fails any future caller that tries to mutate it. + request.setAttribute(ATTR_INPUT_TEMP_FILES, Collections.unmodifiableList(tempFiles)); + request.setAttribute(ATTR_INPUT_BYTES, totalInputBytes); + request.setAttribute(ATTR_TOOL_ID, resolveToolId(request)); + + ChargeContext ctx = + new ChargeContext( + currentUser.getId(), + currentUser.getTeam() == null ? null : currentUser.getTeam().getId(), + determineSource(request, auth), + ProcessType.SINGLE_TOOL, + category); + + ChargeOutcome outcome; + try { + outcome = chargeService.openProcess(ctx, inputs); + } catch (IOException e) { + throw new RuntimeException("openProcess IO failure", e); + } + request.setAttribute(ATTR_JOB_ID, outcome.processId()); + request.setAttribute(ATTR_DISPOSITION, outcome.disposition()); + + if (outcome.disposition() == ChargeOutcome.Disposition.OPENED) { + callsOpened.increment(); + } else { + callsJoined.increment(); + } + } + + @Override + public void afterCompletion( + HttpServletRequest request, + HttpServletResponse response, + Object handler, + Exception ex) { + Timer.Sample sample = Timer.start(); + try { + if (Boolean.TRUE.equals(request.getAttribute(ATTR_FAILED))) { + cleanupInputs(request); + closeWrapper(request); + return; + } + UUID jobId = (UUID) request.getAttribute(ATTR_JOB_ID); + if (jobId == null) { + closeWrapper(request); + return; + } + try { + doAfterCompletion(request, response, jobId); + } catch (RuntimeException e) { + log.warn( + "PAYG afterCompletion failed for job {}; lineage may be incomplete", + jobId, + e); + errorsCounter.increment(); + } finally { + cleanupInputs(request); + closeWrapper(request); + } + } finally { + sample.stop(afterCompletionTimer); + } + } + + private void doAfterCompletion( + HttpServletRequest request, HttpServletResponse response, UUID jobId) { + int status = response.getStatus(); + ChargeOutcome.Disposition disposition = + (ChargeOutcome.Disposition) request.getAttribute(ATTR_DISPOSITION); + String toolId = + Optional.ofNullable((String) request.getAttribute(ATTR_TOOL_ID)).orElse("unknown"); + Long inputBytes = (Long) request.getAttribute(ATTR_INPUT_BYTES); + + // Step audit row — appended for every disposition + every outcome class. Done first so a + // refund-and-close still has the failure recorded against the now-CLOSED process. + JobStepStatus stepStatus = status < 400 ? JobStepStatus.OK : JobStepStatus.FAILED; + String errorCode = status >= 400 ? String.valueOf(status) : null; + try { + jobService.appendStep(jobId, toolId, stepStatus, null, inputBytes, errorCode); + } catch (RuntimeException e) { + log.debug("appendStep failed for job {}: {}", jobId, e.getMessage()); + } + + if (status >= 500) { + if (disposition == ChargeOutcome.Disposition.OPENED) { + chargeService.markFirstStepFailed(jobId, "first-step-5xx:" + status); + refundsCounter.increment(); + } else { + chargeService.decrementStepCount(jobId); + } + return; + } + if (status >= 400) { + // 4xx: customer paid for the attempt. No OUTPUT recording, no refund. + // Still a successful-from-billing-standpoint OPENED process — meter it below. + meterIfOpened(jobId, disposition); + return; + } + + // Success: this is the moment the billable work finished, so this is when we tell Stripe. + // Only the OPENED request meters — JOINED follow-up steps (chained tools on the same + // document) added no units and must not re-meter. The process stays OPEN for further + // lineage joins; StaleJobCloser closing it later is a no-op at Stripe thanks to the shared + // idempotency key. metering is best-effort and must never break the response teardown. + meterIfOpened(jobId, disposition); + recordOutputs(request, response, jobId); + } + + /** + * Fire the Stripe meter for a just-finished process, but only when this request OPENED it. Runs + * on the request-teardown thread (the response is already flushed to the client); {@code + * meterJobUsage} is best-effort and swallows its own failures, but we still guard here so a + * meter hiccup can't disturb lineage/cleanup that follows. + */ + private void meterIfOpened(UUID jobId, ChargeOutcome.Disposition disposition) { + if (disposition != ChargeOutcome.Disposition.OPENED) { + return; + } + try { + chargeService.meterJobUsage(jobId); + } catch (RuntimeException e) { + log.warn("Meter-on-completion failed for job {}: {}", jobId, e.getMessage()); + errorsCounter.increment(); + } + } + + private void recordOutputs( + HttpServletRequest request, HttpServletResponse response, UUID jobId) { + PaygResponseBodyWrapper wrapper = + (PaygResponseBodyWrapper) + request.getAttribute(PaygResponseBodyWrapperFilter.REQUEST_ATTRIBUTE); + if (wrapper == null) { + return; + } + Long maxBytes = properties.getResponse().getMaxBytes(); + if (maxBytes != null && wrapper.bytesWritten() > maxBytes) { + log.debug( + "Response size {} exceeds payg.filter.response.max-bytes={}; skipping OUTPUT recording", + wrapper.bytesWritten(), + maxBytes); + return; + } + Path bodyPath; + try { + bodyPath = wrapper.materialisedPath(); + } catch (IOException e) { + log.debug("materialisedPath failed for job {}: {}", jobId, e.getMessage()); + return; + } + if (bodyPath == null) { + return; + } + List pdfs = + outputExtractor.extract(response.getContentType(), bodyPath); + try { + for (PaygOutputExtractor.ExtractedPdf pdf : pdfs) { + try { + jobService.recordOutput(jobId, pdf.path()); + } catch (IOException e) { + log.debug( + "recordOutput failed for job {} path {}: {}", + jobId, + pdf.path(), + e.getMessage()); + } + } + } finally { + for (PaygOutputExtractor.ExtractedPdf pdf : pdfs) { + pdf.close(); + } + } + } + + @Override + public void afterConcurrentHandlingStarted( + HttpServletRequest request, HttpServletResponse response, Object handler) { + // Async handoff: don't touch state. afterCompletion will fire when the async work resolves. + } + + @SuppressWarnings("unchecked") + private void cleanupInputs(HttpServletRequest request) { + Object raw = request.getAttribute(ATTR_INPUT_TEMP_FILES); + if (raw instanceof List) { + for (Object o : (List) raw) { + if (o instanceof TempFile tf) { + tf.close(); + } + } + } + request.removeAttribute(ATTR_INPUT_TEMP_FILES); + } + + private void closeWrapper(HttpServletRequest request) { + Object wrapper = request.getAttribute(PaygResponseBodyWrapperFilter.REQUEST_ATTRIBUTE); + if (wrapper instanceof PaygResponseBodyWrapper w) { + w.close(); + } + } + + private User resolveUser(Authentication auth) { + if (auth == null || !auth.isAuthenticated()) { + return null; + } + if (auth instanceof ApiKeyAuthenticationToken && auth.getPrincipal() instanceof User u) { + return u; + } + try { + String supabaseId = AuthenticationUtils.extractSupabaseId(auth); + if (supabaseId == null) { + return null; + } + UUID supabaseUuid = UUID.fromString(supabaseId); + return userRepository.findBySupabaseId(supabaseUuid).orElse(null); + } catch (RuntimeException e) { + log.debug("PAYG resolveUser failed: {}", e.getMessage()); + return null; + } + } + + private static JobSource determineSource(HttpServletRequest request, Authentication auth) { + String automationHeader = request.getHeader(AUTOMATION_HEADER); + if (automationHeader != null && "true".equalsIgnoreCase(automationHeader.trim())) { + return JobSource.PIPELINE; + } + String desktopHeader = request.getHeader(DESKTOP_CLIENT_HEADER); + if (desktopHeader != null && "desktop".equalsIgnoreCase(desktopHeader.trim())) { + return JobSource.DESKTOP_APP; + } + if (auth instanceof ApiKeyAuthenticationToken) { + return JobSource.API; + } + return JobSource.WEB; + } + + /** + * Resolve the {@link BillingCategory} for this request. Precedence: {@code + * X-Stirling-Automation: true} or {@code @RequiresFeature(AUTOMATION)} → AUTOMATION; + * {@code @RequiresFeature(AI_SUPPORT)} → AI; an AI document-tool route ({@link AiToolRoutes}) → + * AI; API-key auth → API; otherwise BYPASSED (manual UI tool — short-circuited in {@link + * #preHandle}). + * + *

Method-level {@code @RequiresFeature} wins over class-level. Multiple gates: AUTOMATION + * dominates AI within a single annotation. The AI-tool path check sits below the automation + * header on purpose: an AI tool dispatched inside a policy / AI workflow bills as AUTOMATION, + * while a direct call to it bills as AI. + */ + private static BillingCategory determineCategory( + HandlerMethod handler, HttpServletRequest request, Authentication auth) { + String automationHeader = request.getHeader(AUTOMATION_HEADER); + if (automationHeader != null && "true".equalsIgnoreCase(automationHeader.trim())) { + return BillingCategory.AUTOMATION; + } + RequiresFeature ann = + AnnotationUtils.findAnnotation(handler.getMethod(), RequiresFeature.class); + if (ann == null) { + ann = AnnotationUtils.findAnnotation(handler.getBeanType(), RequiresFeature.class); + } + if (ann != null) { + boolean ai = false; + for (FeatureGate gate : ann.value()) { + if (gate == FeatureGate.AUTOMATION) { + return BillingCategory.AUTOMATION; + } + if (gate == FeatureGate.AI_SUPPORT) { + ai = true; + } + } + if (ai) { + return BillingCategory.AI; + } + } + // AI document tools (proprietary module, recognised by path). A direct call bills as AI; an + // orchestrator-dispatched call already returned AUTOMATION above via the automation header. + if (AiToolRoutes.matches(request)) { + return BillingCategory.AI; + } + if (auth instanceof ApiKeyAuthenticationToken) { + return BillingCategory.API; + } + return BillingCategory.BYPASSED; + } + + /** + * Resolves the {@code tool_id} value stored on {@code processing_job_step}. Prefers the route + * pattern (e.g. {@code /api/v1/security/add-password}) over the raw URI so audit rollups + * aggregate by endpoint rather than by request — path variables, query strings, and matrix + * params don't pollute the column. Falls back to the raw URI when the pattern isn't available + * (non-Spring-MVC dispatches, async re-dispatch edges). + * + *

Truncates to {@link #TOOL_ID_MAX_LENGTH} to match the column's {@code VARCHAR(128)} width. + * Logs at WARN + increments {@link #errorsCounter} when truncation actually happens so support + * notices the {@code tool_id} they expected isn't what we stored. + */ + private String resolveToolId(HttpServletRequest request) { + Object pattern = request.getAttribute(HandlerMapping.BEST_MATCHING_PATTERN_ATTRIBUTE); + String value = pattern instanceof String s ? s : request.getRequestURI(); + if (value == null) { + return "unknown"; + } + if (value.length() <= TOOL_ID_MAX_LENGTH) { + return value; + } + log.warn( + "tool_id length {} exceeds column max {}; truncating. value='{}'", + value.length(), + TOOL_ID_MAX_LENGTH, + value); + errorsCounter.increment(); + return value.substring(0, TOOL_ID_MAX_LENGTH); + } +} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/filter/PaygFilterProperties.java b/app/saas/src/main/java/stirling/software/saas/payg/filter/PaygFilterProperties.java new file mode 100644 index 0000000000..89a24e3881 --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/filter/PaygFilterProperties.java @@ -0,0 +1,50 @@ +package stirling.software.saas.payg.filter; + +import org.springframework.boot.context.properties.ConfigurationProperties; +import org.springframework.context.annotation.Profile; +import org.springframework.stereotype.Component; + +import lombok.Getter; +import lombok.Setter; + +/** + * Configuration knobs for the PAYG filter + interceptor stack. See {@code PAYG_FILTER_DESIGN.md} + * §16 for the rationale on each default. + * + *

    + *
  • {@code payg.filter.enabled} — master kill switch. Restart-required; no + * {@code @RefreshScope}. When {@code false}, the wrapper filter passes through and the + * interceptor short-circuits in {@code preHandle}. + *
  • {@code payg.filter.response.in-memory-threshold-bytes} — the wrapper buffers below this in + * a {@link java.io.ByteArrayOutputStream}; above it spills to a {@code TempFile}. + *
  • {@code payg.filter.response.max-bytes} — optional ceiling. Responses exceeding this skip + * OUTPUT recording in {@code afterCompletion}. {@code null} = unbounded. + *
+ */ +@Component +@Profile("saas") +@ConfigurationProperties(prefix = "payg.filter") +@Getter +@Setter +public class PaygFilterProperties { + + private boolean enabled = true; + + private final Response response = new Response(); + + @Getter + @Setter + public static class Response { + /** 10 MiB. Tiny responses stay in RAM; large responses spill. */ + private long inMemoryThresholdBytes = 10L * 1024L * 1024L; + + /** + * Ceiling for OUTPUT recording. Responses larger than this skip the per-PDF hash + ZIP + * unpack — the bytes still flowed through to the client unmodified, only lineage capture is + * dropped. Default 500 MiB is generous for the largest realistic Stirling responses (full + * split-to-ZIP on a 1000-page document) while preventing pathological cases from tying up + * the interceptor for minutes. Set to {@code null} for "no ceiling at all". + */ + private Long maxBytes = 500L * 1024L * 1024L; + } +} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/filter/PaygOutputExtractor.java b/app/saas/src/main/java/stirling/software/saas/payg/filter/PaygOutputExtractor.java new file mode 100644 index 0000000000..e413d6c73d --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/filter/PaygOutputExtractor.java @@ -0,0 +1,241 @@ +package stirling.software.saas.payg.filter; + +import java.io.IOException; +import java.io.InputStream; +import java.nio.file.Files; +import java.nio.file.Path; +import java.nio.file.StandardCopyOption; +import java.util.ArrayList; +import java.util.List; +import java.util.Locale; +import java.util.Objects; +import java.util.zip.ZipEntry; +import java.util.zip.ZipInputStream; + +import org.springframework.context.annotation.Profile; +import org.springframework.stereotype.Component; + +import lombok.extern.slf4j.Slf4j; + +import stirling.software.common.util.TempFile; +import stirling.software.common.util.TempFileManager; + +/** + * Extracts the PDF artefacts from a controller's response body for lineage OUTPUT recording. Two + * content-type paths: + * + *
    + *
  • {@code application/pdf} — return the response body path verbatim (the wrapper already + * materialised it to disk). + *
  • {@code application/zip} — iterate entries; each that has a {@code .pdf} extension AND + * starts with the {@code %PDF-} magic-byte sequence becomes one extracted {@link Path}. + *
+ * + *

Anything else yields an empty list — non-PDF responses don't drive lineage matching. Encrypted + * ZIPs, malformed entries, and IO failures log at debug and yield an empty list (per design §15 + * fail-open). Nested ZIP-of-ZIPs is intentionally out of scope: only outer-level PDFs are + * extracted. + * + *

Each extracted ZIP entry is written to its own {@link TempFile} returned in the result list. + * Callers are responsible for closing those temp files — typically the interceptor closes them in + * the same {@code afterCompletion} that calls this extractor. The whole-response body path is owned + * by the wrapper and is NOT in the returned list when extraction takes the direct-PDF path. + */ +@Slf4j +@Component +@Profile("saas") +public class PaygOutputExtractor { + + /** {@code %PDF-} in ASCII. */ + private static final byte[] PDF_MAGIC = {0x25, 0x50, 0x44, 0x46, 0x2D}; + + /** {@code PK\x03\x04} — local file header for any non-empty ZIP. */ + private static final byte[] ZIP_MAGIC = {0x50, 0x4B, 0x03, 0x04}; + + private static final String ZIP_CONTENT_TYPE = "application/zip"; + private static final String PDF_CONTENT_TYPE = "application/pdf"; + private static final String OCTET_STREAM_CONTENT_TYPE = "application/octet-stream"; + + private final TempFileManager tempFileManager; + + public PaygOutputExtractor(TempFileManager tempFileManager) { + this.tempFileManager = Objects.requireNonNull(tempFileManager, "tempFileManager"); + } + + /** + * Extract PDF paths from {@code bodyPath} according to {@code contentType}. The returned list + * carries {@link ExtractedPdf} records — each wraps a {@link Path} and an indicator of whether + * the caller owns the temp file lifecycle. + * + * @param contentType the response Content-Type header (may be null / parametrised — only the + * base media type is inspected) + * @param bodyPath the on-disk full response body (from the wrapper's {@code materialisedPath}) + * @return zero-or-more PDFs extracted from the body. Empty when content type is not PDF/ZIP, + * when extraction failed, when no entries matched, or when {@code bodyPath} is null. + */ + public List extract(String contentType, Path bodyPath) { + if (bodyPath == null) { + return List.of(); + } + String mediaType = stripParameters(contentType); + + if (PDF_CONTENT_TYPE.equalsIgnoreCase(mediaType)) { + // Wrapper-owned path. Don't claim ownership. Magic-byte check protects against tools + // that emit application/pdf for non-PDF payloads (mirrors the ZIP-entry path below). + if (!isPdfMagic(bodyPath)) { + log.debug( + "Response advertised application/pdf but content does not start with" + + " %PDF- magic bytes; skipping OUTPUT recording. body={}", + bodyPath); + return List.of(); + } + return List.of(new ExtractedPdf(bodyPath, null)); + } + if (ZIP_CONTENT_TYPE.equalsIgnoreCase(mediaType)) { + return extractZip(bodyPath); + } + // Stirling-PDF tool endpoints sometimes set Content-Type to application/octet-stream (or + // no header at all) even when the body is a real PDF or ZIP — Spring's default + // StreamingResponseBody path doesn't always negotiate content type. When the declared + // Content-Type is missing or generic, sniff magic bytes in a single head-read so we don't + // open the body twice on the common negative path. + if (mediaType == null || OCTET_STREAM_CONTENT_TYPE.equalsIgnoreCase(mediaType)) { + BodyMagic magic = sniffMagic(bodyPath); + if (magic == BodyMagic.PDF) { + return List.of(new ExtractedPdf(bodyPath, null)); + } + if (magic == BodyMagic.ZIP) { + return extractZip(bodyPath); + } + } + return List.of(); + } + + /** Discriminator returned by {@link #sniffMagic(Path)}. */ + private enum BodyMagic { + PDF, + ZIP, + NEITHER + } + + /** + * Single-pass magic-byte sniff used by the generic Content-Type branch. Opens {@code path} + * once, reads enough bytes to compare against both PDF and ZIP magics, returns the first match + * (or {@link BodyMagic#NEITHER} if neither matched / read failed). Replaces two consecutive + * {@link #isMagic} calls that would have opened the file twice on the common negative path. + */ + private BodyMagic sniffMagic(Path path) { + int needed = Math.max(PDF_MAGIC.length, ZIP_MAGIC.length); + byte[] head = new byte[needed]; + int read; + try (InputStream in = Files.newInputStream(path)) { + read = in.read(head); + } catch (IOException e) { + log.debug("Magic-byte sniff failed for {}", path, e); + return BodyMagic.NEITHER; + } + if (read >= PDF_MAGIC.length && startsWith(head, PDF_MAGIC)) { + return BodyMagic.PDF; + } + if (read >= ZIP_MAGIC.length && startsWith(head, ZIP_MAGIC)) { + return BodyMagic.ZIP; + } + return BodyMagic.NEITHER; + } + + private static boolean startsWith(byte[] buf, byte[] prefix) { + for (int i = 0; i < prefix.length; i++) { + if (buf[i] != prefix[i]) { + return false; + } + } + return true; + } + + private List extractZip(Path bodyPath) { + List results = new ArrayList<>(); + try (InputStream rawIn = Files.newInputStream(bodyPath); + ZipInputStream zin = new ZipInputStream(rawIn)) { + ZipEntry entry; + while ((entry = zin.getNextEntry()) != null) { + try { + if (entry.isDirectory()) { + continue; + } + String name = entry.getName(); + if (name == null || !name.toLowerCase(Locale.ROOT).endsWith(".pdf")) { + continue; + } + TempFile temp = tempFileManager.createManagedTempFile(".pdf"); + Files.copy(zin, temp.getPath(), StandardCopyOption.REPLACE_EXISTING); + if (isPdfMagic(temp.getPath())) { + results.add(new ExtractedPdf(temp.getPath(), temp)); + } else { + temp.close(); + } + } finally { + zin.closeEntry(); + } + } + } catch (IOException | IllegalArgumentException e) { + // ZipException extends IOException; IllegalArgumentException covers ZipEntry + // "MALFORMED" surfaces from zlib. Fail-open: caller still serves the response. + log.debug( + "ZIP unpack failed for response body {} ({}); skipping per-PDF OUTPUT recording", + bodyPath, + e.getClass().getSimpleName()); + // Close anything we already opened. + for (ExtractedPdf p : results) { + p.close(); + } + return List.of(); + } + return results; + } + + private boolean isPdfMagic(Path path) { + return isMagic(path, PDF_MAGIC); + } + + /** Returns true if the first {@code magic.length} bytes of {@code path} equal {@code magic}. */ + private boolean isMagic(Path path, byte[] magic) { + byte[] head = new byte[magic.length]; + try (InputStream in = Files.newInputStream(path)) { + int read = in.read(head); + if (read != magic.length) { + return false; + } + for (int i = 0; i < magic.length; i++) { + if (head[i] != magic[i]) { + return false; + } + } + return true; + } catch (IOException e) { + log.debug("Magic-byte check failed for {}", path, e); + return false; + } + } + + private static String stripParameters(String contentType) { + if (contentType == null) { + return null; + } + int semi = contentType.indexOf(';'); + return (semi < 0 ? contentType : contentType.substring(0, semi)).trim(); + } + + /** + * One PDF extracted from the response body. {@link #ownedTempFile} is non-null only when this + * extractor created the temp file (ZIP entries); the {@link #path} for direct-PDF responses + * remains owned by the wrapper. + */ + public record ExtractedPdf(Path path, TempFile ownedTempFile) implements AutoCloseable { + @Override + public void close() { + if (ownedTempFile != null) { + ownedTempFile.close(); + } + } + } +} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/filter/PaygResponseBodyWrapper.java b/app/saas/src/main/java/stirling/software/saas/payg/filter/PaygResponseBodyWrapper.java new file mode 100644 index 0000000000..25450c54f1 --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/filter/PaygResponseBodyWrapper.java @@ -0,0 +1,297 @@ +package stirling.software.saas.payg.filter; + +import java.io.BufferedOutputStream; +import java.io.ByteArrayOutputStream; +import java.io.IOException; +import java.io.OutputStream; +import java.io.OutputStreamWriter; +import java.io.PrintWriter; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.Objects; + +import jakarta.servlet.ServletOutputStream; +import jakarta.servlet.WriteListener; +import jakarta.servlet.http.HttpServletResponse; +import jakarta.servlet.http.HttpServletResponseWrapper; + +import lombok.extern.slf4j.Slf4j; + +import stirling.software.common.util.TempFile; +import stirling.software.common.util.TempFileManager; + +/** + * Tees the controller's response body so the PAYG interceptor can hash it for OUTPUT lineage + * recording. Writes still flow through to the real client output stream unmodified — the wrapper + * just keeps a parallel copy. + * + *

Memory model: bytes accumulate in an in-memory {@link ByteArrayOutputStream} until the + * configured {@code in-memory-threshold-bytes} is crossed; after that the wrapper spills the + * existing buffer plus all subsequent writes to a {@link TempFile} owned by {@link + * TempFileManager}. Tiny responses (error JSONs, small PDFs) stay entirely in RAM with zero disk + * IO; large responses (split-to-ZIP, big compressed PDFs) spill cleanly. See design doc §8. + * + *

The interceptor calls {@link #materialisedPath()} from {@code afterCompletion} to get a {@code + * Path} suitable for the lineage detector + ZIP unpack. The wrapper guarantees a Path even if the + * response stayed in memory — it materialises the buffer to a {@link TempFile} on demand so callers + * have a uniform file-based interface. + * + *

{@link #close()} closes any {@link TempFile} the wrapper created. Callers MUST invoke close in + * a finally — typically the interceptor's {@code afterCompletion} after it's done hashing. + * + *

Thread safety: all mutating methods (record paths, materialisedPath, close, + * resetBuffer) synchronize on the wrapper instance. The Servlet spec serialises controller writes + * onto a single dispatch thread, but the {@link jakarta.servlet.AsyncListener#onComplete} callback + * that closes the wrapper for async controllers runs on a container thread distinct from the + * dispatch thread that produced the body. The synchronization makes that handoff safe and also + * guards against future callers (e.g. tests) that might invoke {@link #materialisedPath} or {@link + * #close} from a non-dispatch thread. + */ +@Slf4j +public class PaygResponseBodyWrapper extends HttpServletResponseWrapper implements AutoCloseable { + + private final TempFileManager tempFileManager; + private final long inMemoryThresholdBytes; + + /** Lazily-created on the first getOutputStream() / getWriter() call. */ + private TeeingServletOutputStream teeOut; + + private PrintWriter writer; + + /** + * In-memory accumulator until the threshold is crossed. Becomes null after spill (helps GC of + * potentially large buffers). + */ + private ByteArrayOutputStream memoryBuffer = new ByteArrayOutputStream(); + + /** Non-null once we've spilled. Owns the {@link TempFile} below. */ + private OutputStream spillStream; + + /** Non-null once we've spilled, OR once {@link #materialisedPath()} forced materialisation. */ + private TempFile spillFile; + + private long bytesWritten; + private boolean spilled; + + public PaygResponseBodyWrapper( + HttpServletResponse response, + TempFileManager tempFileManager, + long inMemoryThresholdBytes) { + super(response); + this.tempFileManager = Objects.requireNonNull(tempFileManager, "tempFileManager"); + if (inMemoryThresholdBytes < 0) { + throw new IllegalArgumentException( + "inMemoryThresholdBytes must be >= 0, got " + inMemoryThresholdBytes); + } + this.inMemoryThresholdBytes = inMemoryThresholdBytes; + } + + @Override + public ServletOutputStream getOutputStream() throws IOException { + if (writer != null) { + // Servlet spec: getOutputStream() and getWriter() are mutually exclusive per request. + throw new IllegalStateException( + "getWriter() was already called on this response; cannot switch to getOutputStream()"); + } + if (teeOut == null) { + teeOut = new TeeingServletOutputStream(super.getOutputStream()); + } + return teeOut; + } + + @Override + public PrintWriter getWriter() throws IOException { + if (teeOut != null) { + throw new IllegalStateException( + "getOutputStream() was already called on this response; cannot switch to getWriter()"); + } + if (writer == null) { + String encoding = getCharacterEncoding() != null ? getCharacterEncoding() : "UTF-8"; + // Wrap super.getOutputStream() directly with our TeeingServletOutputStream so writes + // through the Writer path also get tee'd. The Writer just adds character→byte encoding. + teeOut = new TeeingServletOutputStream(super.getOutputStream()); + writer = new PrintWriter(new OutputStreamWriter(teeOut, encoding)); + } + return writer; + } + + @Override + public synchronized void resetBuffer() { + super.resetBuffer(); + if (memoryBuffer != null) { + memoryBuffer.reset(); + } + // If we'd already spilled, the only safe move is to abandon the spill file: the client + // hasn't seen the body yet (otherwise resetBuffer would be illegal), but our tee captured + // bytes we now want to forget. + if (spilled) { + closeSpillQuietly(); + spillFile = null; + spillStream = null; + spilled = false; + memoryBuffer = new ByteArrayOutputStream(); + } + bytesWritten = 0; + } + + /** + * Returns a {@link Path} containing the full response body, or {@code null} if no bytes were + * written. The returned path is owned by this wrapper — do NOT delete or modify it. Use {@link + * #close()} to release. + * + *

If the response stayed under the threshold, this materialises the in-memory buffer to a + * {@link TempFile} on demand so the caller always gets a file-based handle (uniform with the + * spilled path). + */ + public synchronized Path materialisedPath() throws IOException { + if (bytesWritten == 0) { + return null; + } + // Flush the writer so any character data lands in the underlying byte stream first. + if (writer != null) { + writer.flush(); + } + if (spilled) { + spillStream.flush(); + return spillFile.getPath(); + } + // Stayed in memory — materialise on demand for a uniform Path-based interface. + if (spillFile == null) { + spillFile = tempFileManager.createManagedTempFile(".body"); + Files.write(spillFile.getPath(), memoryBuffer.toByteArray()); + } + return spillFile.getPath(); + } + + public synchronized long bytesWritten() { + return bytesWritten; + } + + @Override + public synchronized void close() { + closeSpillQuietly(); + } + + private void closeSpillQuietly() { + if (spillStream != null) { + try { + spillStream.close(); + } catch (IOException e) { + log.debug("Ignoring close error on spill stream: {}", e.getMessage()); + } + spillStream = null; + } + if (spillFile != null) { + spillFile.close(); // TempFile.close() deletes the file + spillFile = null; + } + } + + /** + * Routes bytes both to the real client output stream AND to our buffer (in-memory or spilled). + * Writes are not re-batched — each call to the delegate corresponds exactly to one call here. + */ + private final class TeeingServletOutputStream extends ServletOutputStream { + + private final ServletOutputStream delegate; + + TeeingServletOutputStream(ServletOutputStream delegate) { + this.delegate = delegate; + } + + @Override + public void write(int b) throws IOException { + delegate.write(b); + recordSingleByte((byte) b); + } + + @Override + public void write(byte[] b, int off, int len) throws IOException { + delegate.write(b, off, len); + recordRange(b, off, len); + } + + @Override + public void flush() throws IOException { + delegate.flush(); + // Read spilled / spillStream under the outer monitor so we observe the publication + // written by spillToDisk() on another thread. The writers (recordSingleByte, + // recordRange, spillToDisk) already mutate these fields under + // synchronized(PaygResponseBodyWrapper.this); without matching locking here the + // JMM permits stale reads (false `spilled`, null `spillStream`) and Aikido AI + // flagged that gap. + synchronized (PaygResponseBodyWrapper.this) { + if (spilled && spillStream != null) { + spillStream.flush(); + } + } + } + + @Override + public void close() throws IOException { + delegate.close(); + synchronized (PaygResponseBodyWrapper.this) { + if (spilled && spillStream != null) { + spillStream.flush(); + } + } + } + + @Override + public boolean isReady() { + return delegate.isReady(); + } + + @Override + public void setWriteListener(WriteListener writeListener) { + delegate.setWriteListener(writeListener); + } + } + + private synchronized void recordSingleByte(byte b) throws IOException { + if (spilled) { + spillStream.write(b & 0xFF); + } else if (bytesWritten + 1 > inMemoryThresholdBytes) { + spillToDisk(); + spillStream.write(b & 0xFF); + } else { + memoryBuffer.write(b & 0xFF); + } + bytesWritten++; + } + + private synchronized void recordRange(byte[] b, int off, int len) throws IOException { + if (spilled) { + spillStream.write(b, off, len); + } else if (bytesWritten + len > inMemoryThresholdBytes) { + // This write crosses the threshold. Spill the existing in-memory buffer, then write + // this entire chunk to disk too — we don't bother splitting it for the sake of staying + // exactly at the threshold. Going over by at most one chunk is fine. + spillToDisk(); + spillStream.write(b, off, len); + } else { + memoryBuffer.write(b, off, len); + } + bytesWritten += len; + } + + /** Spill stream buffer size — coalesces Tomcat's per-chunk syscalls into 64 KiB writes. */ + private static final int SPILL_BUFFER_SIZE = 64 * 1024; + + private void spillToDisk() throws IOException { + spillFile = tempFileManager.createManagedTempFile(".body"); + // Wrap in BufferedOutputStream — without this every Tomcat chunk (default 8 KiB) was a + // separate syscall to the temp file, which dominates wall-clock on big spilled responses. + spillStream = + new BufferedOutputStream( + Files.newOutputStream(spillFile.getPath()), SPILL_BUFFER_SIZE); + memoryBuffer.writeTo(spillStream); + memoryBuffer = null; // help GC of potentially large buffer + spilled = true; + log.debug( + "PaygResponseBodyWrapper spilled to {} after {} bytes (threshold {})", + spillFile.getPath(), + bytesWritten, + inMemoryThresholdBytes); + } +} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/filter/PaygResponseBodyWrapperFilter.java b/app/saas/src/main/java/stirling/software/saas/payg/filter/PaygResponseBodyWrapperFilter.java new file mode 100644 index 0000000000..8c6876fb63 --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/filter/PaygResponseBodyWrapperFilter.java @@ -0,0 +1,138 @@ +package stirling.software.saas.payg.filter; + +import java.io.IOException; + +import org.springframework.context.annotation.Profile; +import org.springframework.stereotype.Component; +import org.springframework.web.filter.OncePerRequestFilter; + +import jakarta.servlet.AsyncEvent; +import jakarta.servlet.AsyncListener; +import jakarta.servlet.FilterChain; +import jakarta.servlet.ServletException; +import jakarta.servlet.http.HttpServletRequest; +import jakarta.servlet.http.HttpServletResponse; + +import lombok.extern.slf4j.Slf4j; + +import stirling.software.common.util.TempFileManager; + +/** + * Wraps the {@link HttpServletResponse} with a {@link PaygResponseBodyWrapper} for every request, + * making the response body retrievable by the downstream {@code PaygChargeInterceptor} in {@code + * afterCompletion}. The wrapper is stashed as a request attribute under {@link #REQUEST_ATTRIBUTE} + * so the interceptor can find it. + * + *

Pure plumbing — no business logic. When {@code payg.filter.enabled=false} the filter passes + * through unchanged. + * + *

Lifecycle: the wrapper is closed in a {@code finally} after the chain returns for sync + * requests; for async controllers ({@code DeferredResult}, {@code CompletableFuture}), close is + * deferred to an {@link AsyncListener} so the wrapper survives the async window. Close is + * idempotent so a defensive call by the interceptor's {@code afterCompletion} is harmless. + */ +@Slf4j +@Component +@Profile("saas") +public class PaygResponseBodyWrapperFilter extends OncePerRequestFilter { + + /** Request-attribute key under which the wrapper is exposed to the interceptor. */ + public static final String REQUEST_ATTRIBUTE = + PaygResponseBodyWrapperFilter.class.getName() + ".WRAPPER"; + + private final TempFileManager tempFileManager; + private final PaygFilterProperties properties; + + public PaygResponseBodyWrapperFilter( + TempFileManager tempFileManager, PaygFilterProperties properties) { + this.tempFileManager = tempFileManager; + this.properties = properties; + } + + @Override + protected void doFilterInternal( + HttpServletRequest request, HttpServletResponse response, FilterChain chain) + throws ServletException, IOException { + + if (!properties.isEnabled()) { + chain.doFilter(request, response); + return; + } + + PaygResponseBodyWrapper wrapper; + try { + wrapper = + new PaygResponseBodyWrapper( + response, + tempFileManager, + properties.getResponse().getInMemoryThresholdBytes()); + } catch (RuntimeException e) { + // Wrapper construction failure: fail-open. Pass through unwrapped — OUTPUT recording + // is lost but the customer's tool call still runs. + log.warn("PaygResponseBodyWrapper construction failed; passing through unwrapped", e); + chain.doFilter(request, response); + return; + } + + request.setAttribute(REQUEST_ATTRIBUTE, wrapper); + boolean asyncStarted = false; + try { + chain.doFilter(request, wrapper); + asyncStarted = request.isAsyncStarted(); + if (asyncStarted) { + // Async controller: defer attribute removal + close to async dispatch completion. + // The interceptor's afterCompletion fires on the async dispatch and needs the + // wrapper attribute still present at that point. close() is idempotent so a + // defensive call by the interceptor is harmless. + request.getAsyncContext().addListener(new ReleaseOnAsyncComplete(request, wrapper)); + } + } finally { + if (!asyncStarted) { + // Sync path: interceptor.afterCompletion has already run inside chain.doFilter. + request.removeAttribute(REQUEST_ATTRIBUTE); + wrapper.close(); + } + } + } + + /** + * For async dispatches, removes the wrapper attribute and closes the wrapper after the async + * dispatch completes. The Servlet container fires exactly one of {@code onComplete} / {@code + * onError} / {@code onTimeout} per async context lifecycle. + */ + private static final class ReleaseOnAsyncComplete implements AsyncListener { + + private final HttpServletRequest request; + private final PaygResponseBodyWrapper wrapper; + + ReleaseOnAsyncComplete(HttpServletRequest request, PaygResponseBodyWrapper wrapper) { + this.request = request; + this.wrapper = wrapper; + } + + private void release() { + request.removeAttribute(REQUEST_ATTRIBUTE); + wrapper.close(); + } + + @Override + public void onComplete(AsyncEvent event) { + release(); + } + + @Override + public void onTimeout(AsyncEvent event) { + release(); + } + + @Override + public void onError(AsyncEvent event) { + release(); + } + + @Override + public void onStartAsync(AsyncEvent event) { + // re-dispatch retains the listener — no-op + } + } +} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/filter/PaygWebMvcConfig.java b/app/saas/src/main/java/stirling/software/saas/payg/filter/PaygWebMvcConfig.java new file mode 100644 index 0000000000..72c502fcff --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/filter/PaygWebMvcConfig.java @@ -0,0 +1,73 @@ +package stirling.software.saas.payg.filter; + +import org.springframework.boot.web.servlet.FilterRegistrationBean; +import org.springframework.context.annotation.Bean; +import org.springframework.context.annotation.Configuration; +import org.springframework.context.annotation.Profile; +import org.springframework.web.servlet.config.annotation.InterceptorRegistry; +import org.springframework.web.servlet.config.annotation.WebMvcConfigurer; + +import lombok.RequiredArgsConstructor; + +import stirling.software.saas.payg.entitlement.EntitlementGuard; + +/** + * Wires the PAYG filter + interceptor into Spring MVC. Two registrations: + * + *

    + *
  • {@link PaygResponseBodyWrapperFilter} as a Servlet filter — registered with no explicit + * order so it sits at the end of the Spring filter chain (after all security filters). Pure + * response-wrapping plumbing. + *
  • {@link PaygChargeInterceptor} as a Spring MVC interceptor — intercepts {@code /api/**} with + * admin/info/health exclusions. + *
+ */ +@Configuration +@Profile("saas") +@RequiredArgsConstructor +public class PaygWebMvcConfig implements WebMvcConfigurer { + + private final PaygChargeInterceptor paygChargeInterceptor; + private final EntitlementGuard entitlementGuard; + + @Bean + public FilterRegistrationBean + paygResponseBodyWrapperFilterRegistration(PaygResponseBodyWrapperFilter filter) { + FilterRegistrationBean reg = + new FilterRegistrationBean<>(filter); + reg.addUrlPatterns("/api/*"); + return reg; + } + + /** + * The {@code PaygChargeInterceptor} runs after the {@link #ENTITLEMENT_GUARD_ORDER guard}, so + * {@code openProcess} only fires for requests the guard has admitted. See {@link + * #ENTITLEMENT_GUARD_ORDER} for the full ordering rationale. + */ + public static final int INTERCEPTOR_ORDER = 1000; + + /** + * The {@code EntitlementGuard} runs BEFORE the charge interceptor. Spring runs interceptors in + * ascending order on the way in and skips a later interceptor's {@code preHandle} (and its + * {@code afterCompletion}) entirely once an earlier one returns {@code false} — so a request + * the guard refuses (over its free allowance / spending cap, or with no subscription to bill) + * short-circuits with its 402 before the charge interceptor ever runs. A blocked request + * therefore never opens a process, materialises inputs, or writes a charge: a refused operation + * must not bill, and running the guard first guarantees that structurally rather than by + * compensating after the fact. + */ + public static final int ENTITLEMENT_GUARD_ORDER = 900; + + @Override + public void addInterceptors(InterceptorRegistry registry) { + registry.addInterceptor(paygChargeInterceptor) + .addPathPatterns("/api/**") + .excludePathPatterns("/api/v1/config/**", "/api/v1/info/**", "/api/v1/admin/**") + .order(INTERCEPTOR_ORDER); + + registry.addInterceptor(entitlementGuard) + .addPathPatterns("/api/**") + .excludePathPatterns("/api/v1/config/**", "/api/v1/info/**", "/api/v1/admin/**") + .order(ENTITLEMENT_GUARD_ORDER); + } +} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/job/JobService.java b/app/saas/src/main/java/stirling/software/saas/payg/job/JobService.java index 005833c70f..5addb5b76d 100644 --- a/app/saas/src/main/java/stirling/software/saas/payg/job/JobService.java +++ b/app/saas/src/main/java/stirling/software/saas/payg/job/JobService.java @@ -6,9 +6,12 @@ import java.time.Duration; import java.time.LocalDateTime; import java.util.ArrayList; import java.util.Comparator; +import java.util.HashMap; import java.util.List; +import java.util.Map; import java.util.Objects; import java.util.Optional; +import java.util.Set; import java.util.UUID; import org.springframework.beans.factory.annotation.Value; @@ -20,6 +23,7 @@ import lombok.extern.slf4j.Slf4j; import stirling.software.saas.payg.lineage.HashLineageDetector; import stirling.software.saas.payg.lineage.LineageMatch; +import stirling.software.saas.payg.lineage.LineageSignature; import stirling.software.saas.payg.model.ArtifactKind; import stirling.software.saas.payg.model.JobStatus; import stirling.software.saas.payg.model.JobStepStatus; @@ -86,7 +90,15 @@ public class JobService { throw new IllegalArgumentException("inputs must not be empty"); } - Optional bestMatch = findBestMatch(ctx.ownerUserId(), inputs); + // Extract signatures ONCE per input, then reuse for both the lineage lookup and the + // post-decision record() call. Avoids hashing every input twice on the hot path. + Map> signaturesByInput = new HashMap<>(inputs.size()); + for (Path input : inputs) { + signaturesByInput.put(input, detector.extractSignatures(input)); + } + + Optional bestMatch = + findBestMatch(ctx.ownerUserId(), inputs, signaturesByInput); if (bestMatch.isPresent()) { ProcessingJob existing = @@ -100,7 +112,7 @@ public class JobService { + " but no such ProcessingJob row" + " exists (stale signature?)")); if (existing.getStepCount() < ctx.stepLimit()) { - return joinExisting(existing, inputs); + return joinExisting(existing, signaturesByInput); } // Step-limit hit: spawn a fresh job. The new job will share input signatures with // the existing chain so future tool calls still lineage-match into the workflow, @@ -111,7 +123,7 @@ public class JobService { ctx.stepLimit()); } - return openFresh(ctx, inputs); + return openFresh(ctx, signaturesByInput); } /** @@ -199,25 +211,40 @@ public class JobService { return stale.size(); } - private Optional findBestMatch(Long userId, List inputs) - throws IOException { + private Optional findBestMatch( + Long userId, List inputs, Map> signaturesByInput) { List matches = new ArrayList<>(inputs.size()); for (Path input : inputs) { - detector.detect(userId, input).ifPresent(matches::add); + detector.detect(userId, signaturesByInput.get(input)).ifPresent(matches::add); } return matches.stream().max(Comparator.comparing(LineageMatch::jobLastStepAt)); } - private JoinOrOpenResult joinExisting(ProcessingJob existing, List inputs) - throws IOException { + private JoinOrOpenResult joinExisting( + ProcessingJob existing, Map> signaturesByInput) { existing.setStepCount(existing.getStepCount() + 1); existing.setLastStepAt(LocalDateTime.now()); ProcessingJob saved = jobRepository.save(existing); - recordAllInputs(saved.getId(), inputs); + recordAllInputs(saved.getId(), signaturesByInput); return new JoinOrOpenResult(saved, JoinOrOpenResult.Disposition.JOINED); } - private JoinOrOpenResult openFresh(JobContext ctx, List inputs) throws IOException { + /** + * Open a standalone process with no lineage inputs, for a billable action that isn't + * file/lineage-driven (e.g. an AI Create session). Because no input signatures are recorded, + * nothing downstream can lineage-join it — each such charge stands alone. {@code docUnits} is + * persisted so the charge service's shadow + ledger rows agree with the job. + */ + @Transactional + public ProcessingJob open(JobContext ctx, int docUnits) { + Objects.requireNonNull(ctx, "ctx"); + ProcessingJob job = openFresh(ctx, Map.of()).job(); + job.setDocUnits(docUnits); + return jobRepository.save(job); + } + + private JoinOrOpenResult openFresh( + JobContext ctx, Map> signaturesByInput) { ProcessingJob fresh = new ProcessingJob(); fresh.setId(UUID.randomUUID()); fresh.setOwnerUserId(ctx.ownerUserId()); @@ -231,13 +258,13 @@ public class JobService { fresh.setLastStepAt(now); fresh.setStatus(JobStatus.OPEN); ProcessingJob saved = jobRepository.save(fresh); - recordAllInputs(saved.getId(), inputs); + recordAllInputs(saved.getId(), signaturesByInput); return new JoinOrOpenResult(saved, JoinOrOpenResult.Disposition.OPENED); } - private void recordAllInputs(UUID jobId, List inputs) throws IOException { - for (Path input : inputs) { - detector.record(jobId, input, ArtifactKind.INPUT); + private void recordAllInputs(UUID jobId, Map> signaturesByInput) { + for (Set signatures : signaturesByInput.values()) { + detector.record(jobId, signatures, ArtifactKind.INPUT); } } } diff --git a/app/saas/src/main/java/stirling/software/saas/payg/job/StaleJobCloser.java b/app/saas/src/main/java/stirling/software/saas/payg/job/StaleJobCloser.java index 6c48139889..432965aa7c 100644 --- a/app/saas/src/main/java/stirling/software/saas/payg/job/StaleJobCloser.java +++ b/app/saas/src/main/java/stirling/software/saas/payg/job/StaleJobCloser.java @@ -1,34 +1,69 @@ package stirling.software.saas.payg.job; +import java.util.List; + import org.springframework.context.annotation.Profile; import org.springframework.scheduling.annotation.Scheduled; import org.springframework.stereotype.Component; -import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; +import stirling.software.saas.payg.charge.JobChargeService; + /** * Auto-closes {@code OPEN} jobs whose {@code last_step_at} is older than the workflow window. Runs * every minute. API users never have to call {@code close()} explicitly — this scheduler is the - * safety net. + * safety net, and (for metered teams) the point at which the Stripe meter event is posted. + * + *

Each stale job is closed individually through {@link JobChargeService#close(java.util.UUID)} + * rather than {@code JobService.closeStale()} (a bulk status flip). That routing matters: {@code + * JobChargeService.close} registers the {@code afterCommit} hook that posts the billable usage to + * Stripe via {@code PaygMeterReportingService}. A bulk flip would close the rows but never meter + * them — usage would accrue in the wallet ledger yet never reach the customer's invoice. + * + *

Per-job transactions + failure isolation: each {@code chargeService.close(id)} runs in its own + * transaction (cross-bean proxied call from this non-transactional scheduled method), so the + * afterCommit meter POST fires once per job and one job's failure can't abort the rest of the + * sweep. The meter event's idempotency key ({@code process::close}) makes a re-run on the next + * tick safe even if a close half-completed. * *

Single-fire only at V1: not {@code @SchedulerLock}'d, consistent with the other - * {@code @Scheduled} tasks in {@code :saas} (none of them are guarded against multi-pod - * double-fires today either). Multi-pod cluster-correctness for all schedulers is tracked in design - * § 9 as a separate cleanup. The underlying {@code closeStale()} call is idempotent — duplicate - * firings read an empty stale set on the second pod, no data corruption risk. + * {@code @Scheduled} tasks in {@code :saas}. Multi-pod cluster-correctness for all schedulers is + * tracked in design § 9 as a separate cleanup; the per-job close + meter idempotency key mean a + * double-fire across pods reads a shrinking stale set and never double-bills. */ @Component @Profile("saas") -@RequiredArgsConstructor @Slf4j public class StaleJobCloser { private final JobService jobService; + private final JobChargeService chargeService; + + public StaleJobCloser(JobService jobService, JobChargeService chargeService) { + this.jobService = jobService; + this.chargeService = chargeService; + } @Scheduled(fixedRateString = "${payg.job.stale-close-interval-ms:60000}") public void closeStale() { - int closed = jobService.closeStale(); + List stale = jobService.findStale(); + if (stale.isEmpty()) { + return; + } + int closed = 0; + for (ProcessingJob job : stale) { + try { + // Routes through the charge service so the afterCommit meter hook fires for + // metered teams. Idempotent: a job already closed by a racing tick no-ops. + chargeService.close(job.getId()); + closed++; + } catch (RuntimeException e) { + // Isolate per job — a single bad row (or a transient meter-path issue) must not + // strand the rest of the stale set open. Next tick retries. + log.warn("StaleJobCloser failed to close job {}: {}", job.getId(), e.getMessage()); + } + } if (closed > 0) { log.info("StaleJobCloser closed {} job(s) idle past the workflow window.", closed); } diff --git a/app/saas/src/main/java/stirling/software/saas/payg/lineage/DefaultHashLineageDetector.java b/app/saas/src/main/java/stirling/software/saas/payg/lineage/DefaultHashLineageDetector.java index 2d261ee02b..082054a09f 100644 --- a/app/saas/src/main/java/stirling/software/saas/payg/lineage/DefaultHashLineageDetector.java +++ b/app/saas/src/main/java/stirling/software/saas/payg/lineage/DefaultHashLineageDetector.java @@ -58,37 +58,45 @@ public class DefaultHashLineageDetector implements HashLineageDetector { @Override public Optional detect(Long userId, Path inputFile) throws IOException { - Objects.requireNonNull(userId, "userId"); Objects.requireNonNull(inputFile, "inputFile"); + return detect(userId, extractSignatures(inputFile)); + } - Set signatures = extractAll(inputFile); + @Override + public Optional detect(Long userId, Set signatures) { + Objects.requireNonNull(userId, "userId"); + Objects.requireNonNull(signatures, "signatures"); if (signatures.isEmpty()) { // No extractor produced anything for this content. Treat as no-match. - log.debug("No signatures extracted from {}; lineage check returns empty.", inputFile); return Optional.empty(); } - return store.findOpenJobForSignatures(userId, signatures, workflowWindow); } @Override public void record(UUID jobId, Path file, ArtifactKind kind) throws IOException { - Objects.requireNonNull(jobId, "jobId"); Objects.requireNonNull(file, "file"); - Objects.requireNonNull(kind, "kind"); + record(jobId, extractSignatures(file), kind); + } - Set signatures = extractAll(file); + @Override + public void record(UUID jobId, Set signatures, ArtifactKind kind) { + Objects.requireNonNull(jobId, "jobId"); + Objects.requireNonNull(signatures, "signatures"); + Objects.requireNonNull(kind, "kind"); if (signatures.isEmpty()) { - log.debug( - "No signatures extracted from {} for job {} ({}); nothing recorded.", - file, - jobId, - kind); + log.debug("No signatures to record for job {} ({}); skipping.", jobId, kind); return; } store.record(jobId, signatures, kind); } + @Override + public Set extractSignatures(Path file) { + Objects.requireNonNull(file, "file"); + return extractAll(file); + } + private Set extractAll(Path file) { Set union = new HashSet<>(); for (LineageSignatureExtractor extractor : extractors) { diff --git a/app/saas/src/main/java/stirling/software/saas/payg/lineage/HashLineageDetector.java b/app/saas/src/main/java/stirling/software/saas/payg/lineage/HashLineageDetector.java index 60aa589bfd..66559eaf40 100644 --- a/app/saas/src/main/java/stirling/software/saas/payg/lineage/HashLineageDetector.java +++ b/app/saas/src/main/java/stirling/software/saas/payg/lineage/HashLineageDetector.java @@ -3,6 +3,7 @@ package stirling.software.saas.payg.lineage; import java.io.IOException; import java.nio.file.Path; import java.util.Optional; +import java.util.Set; import java.util.UUID; import stirling.software.saas.payg.model.ArtifactKind; @@ -33,4 +34,25 @@ public interface HashLineageDetector { /** Records the file's signatures against the given job as either INPUT or OUTPUT. */ void record(UUID jobId, Path file, ArtifactKind kind) throws IOException; + + /** + * Pre-computes the signature set for {@code file} so callers can avoid hashing the same bytes + * twice when they need both {@link #detect} and {@link #record} for the same file. Returned set + * may be empty (no extractor recognised the content) — treat as "no signatures to match or + * record" by both consumers. + */ + Set extractSignatures(Path file); + + /** + * Same as {@link #detect(Long, Path)} but operating on pre-computed signatures. Useful when the + * caller has already extracted them (e.g. via {@link #extractSignatures}) and wants to avoid + * re-hashing. Empty {@code signatures} short-circuits to {@link Optional#empty}. + */ + Optional detect(Long userId, Set signatures); + + /** + * Same as {@link #record(UUID, Path, ArtifactKind)} but operating on pre-computed signatures. + * Empty {@code signatures} is a no-op. + */ + void record(UUID jobId, Set signatures, ArtifactKind kind); } diff --git a/app/saas/src/main/java/stirling/software/saas/payg/meter/PaygMeterEventLog.java b/app/saas/src/main/java/stirling/software/saas/payg/meter/PaygMeterEventLog.java new file mode 100644 index 0000000000..e19cf522e6 --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/meter/PaygMeterEventLog.java @@ -0,0 +1,69 @@ +package stirling.software.saas.payg.meter; + +import java.time.LocalDateTime; +import java.util.UUID; + +import jakarta.persistence.Column; +import jakarta.persistence.Entity; +import jakarta.persistence.GeneratedValue; +import jakarta.persistence.GenerationType; +import jakarta.persistence.Id; +import jakarta.persistence.Table; + +import lombok.Getter; +import lombok.NoArgsConstructor; +import lombok.Setter; + +/** + * Backend-side audit row for one Stripe meter-event POST attempt ({@code payg_meter_event_log}, + * V15). A row is written pending ({@code posted_to_stripe_at} NULL) just before the POST + * and stamped on success; a failed POST leaves it unposted with the Stripe error captured. Rows + * still unposted after a short delay are retried by {@link PaygMeterReconcileScheduler} — this is + * the durability mechanism behind the fail-open meter path, so a Stripe blip never silently + * under-bills. + * + *

{@code idempotency_key} is UNIQUE and identical to the key sent to Stripe ({@code + * process::close}); the unique constraint gives safe at-least-once semantics across the dual + * meter triggers (completion + stale-close) and reconcile retries. + */ +@Entity +@Table(name = "payg_meter_event_log") +@Getter +@Setter +@NoArgsConstructor +public class PaygMeterEventLog { + + @Id + @GeneratedValue(strategy = GenerationType.IDENTITY) + @Column(name = "event_id") + private Long eventId; + + @Column(name = "team_id", nullable = false) + private Long teamId; + + @Column(name = "job_id") + private UUID jobId; + + @Column(name = "idempotency_key", nullable = false, unique = true, length = 128) + private String idempotencyKey; + + @Column(name = "units", nullable = false) + private Integer units; + + /** + * Insert time; the DB column defaults to {@code CURRENT_TIMESTAMP} (set by {@code + * insertPending}). + */ + @Column(name = "occurred_at", nullable = false, insertable = false, updatable = false) + private LocalDateTime occurredAt; + + /** NULL while pending; stamped when the meter-payg-units edge fn returns success. */ + @Column(name = "posted_to_stripe_at") + private LocalDateTime postedToStripeAt; + + @Column(name = "stripe_error_code", length = 64) + private String stripeErrorCode; + + @Column(name = "stripe_error_body", columnDefinition = "text") + private String stripeErrorBody; +} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/meter/PaygMeterReconcileScheduler.java b/app/saas/src/main/java/stirling/software/saas/payg/meter/PaygMeterReconcileScheduler.java new file mode 100644 index 0000000000..81cde551fb --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/meter/PaygMeterReconcileScheduler.java @@ -0,0 +1,139 @@ +package stirling.software.saas.payg.meter; + +import java.time.Duration; +import java.time.LocalDateTime; +import java.util.List; +import java.util.Map; +import java.util.Objects; +import java.util.stream.Collectors; + +import org.springframework.beans.factory.annotation.Value; +import org.springframework.context.annotation.Profile; +import org.springframework.data.domain.PageRequest; +import org.springframework.scheduling.annotation.Scheduled; +import org.springframework.stereotype.Component; + +import io.micrometer.core.instrument.Counter; +import io.micrometer.core.instrument.MeterRegistry; + +import lombok.extern.slf4j.Slf4j; + +import stirling.software.saas.payg.policy.PaygTeamExtensions; +import stirling.software.saas.payg.repository.PaygMeterEventLogRepository; +import stirling.software.saas.payg.repository.PaygTeamExtensionsRepository; + +/** + * Retries PAYG meter events that were logged but never confirmed posted to Stripe — the durability + * half of the fail-open meter path. {@link PaygMeterReportingService} writes a pending {@code + * payg_meter_event_log} row before each POST and stamps it on success; anything left unposted + * (Stripe blip, pod crash between POST and stamp, edge-fn outage) is picked up here and re-sent + * under the same idempotency key, so Stripe dedups rather than double-charging. + * + *

Only retries rows inside Stripe's 24h idempotency window — past that a same-key retry is no + * longer guaranteed to dedup, so stuck rows are logged for manual reconciliation rather than + * risking a double charge. Skips teams that have since unsubscribed (nothing to bill). Like {@link + * stirling.software.saas.payg.lineage.LineagePruneScheduler} it is not {@code @SchedulerLock}'d: a + * duplicate firing on a multi-pod deploy re-sends the same keys, which dedup at Stripe — + * idempotent, wasted IO at worst. + */ +@Component +@Profile("saas") +@Slf4j +public class PaygMeterReconcileScheduler { + + /** Stripe's meter-event idempotency window — a same-key retry past this may double-charge. */ + private static final Duration STRIPE_IDEMPOTENCY_WINDOW = Duration.ofHours(24); + + private final PaygMeterEventLogRepository eventLogRepository; + private final PaygTeamExtensionsRepository teamExtensionsRepository; + private final PaygMeterReportingService meterReportingService; + private final boolean enabled; + private final Duration retryDelay; + private final int batchSize; + private final Counter retriedCounter; + + public PaygMeterReconcileScheduler( + PaygMeterEventLogRepository eventLogRepository, + PaygTeamExtensionsRepository teamExtensionsRepository, + PaygMeterReportingService meterReportingService, + @Value("${payg.meter.reconcile.enabled:true}") boolean enabled, + @Value("${payg.meter.reconcile.retry-delay:PT5M}") Duration retryDelay, + @Value("${payg.meter.reconcile.batch-size:100}") int batchSize, + MeterRegistry meterRegistry) { + this.eventLogRepository = Objects.requireNonNull(eventLogRepository, "eventLogRepository"); + this.teamExtensionsRepository = + Objects.requireNonNull(teamExtensionsRepository, "teamExtensionsRepository"); + this.meterReportingService = + Objects.requireNonNull(meterReportingService, "meterReportingService"); + this.enabled = enabled; + this.retryDelay = Objects.requireNonNull(retryDelay, "retryDelay"); + this.batchSize = batchSize > 0 ? batchSize : 100; + this.retriedCounter = + Counter.builder("payg.meter.reconcile.retried") + .description("PAYG meter events re-posted to Stripe by the reconcile job") + .register(meterRegistry); + } + + @Scheduled(cron = "${payg.meter.reconcile-cron:0 */15 * * * *}", zone = "UTC") + public void reconcile() { + if (!enabled) { + return; + } + LocalDateTime now = LocalDateTime.now(); + // Give the live POST a moment to land before retrying; stay inside the 24h dedup window. + LocalDateTime cutoff = now.minus(retryDelay); + LocalDateTime floor = now.minus(STRIPE_IDEMPOTENCY_WINDOW); + + List retryable = + eventLogRepository.findRetryable(cutoff, floor, PageRequest.of(0, batchSize)); + + // Batch-fetch this page's team extensions in one query (keyed by team id) rather than a + // findById per row — avoids an N+1 when the page spans several teams. + List teamIds = + retryable.stream().map(PaygMeterEventLog::getTeamId).distinct().toList(); + Map extById = + teamExtensionsRepository.findAllById(teamIds).stream() + .collect(Collectors.toMap(PaygTeamExtensions::getTeamId, ext -> ext)); + + int retried = 0; + for (PaygMeterEventLog row : retryable) { + PaygTeamExtensions ext = extById.get(row.getTeamId()); + if (ext == null) { + continue; + } + String subscriptionId = ext.getPaygSubscriptionId(); + String stripeCustomerId = ext.getStripeCustomerId(); + if (subscriptionId == null + || subscriptionId.isBlank() + || stripeCustomerId == null + || stripeCustomerId.isBlank()) { + // Team unsubscribed since the event was logged — nothing to bill; leave the row. + continue; + } + // Same idempotency key → Stripe dedups if the original actually landed. recordUsage + // re-inserts pending as a no-op, re-POSTs, and stamps the row on success. Category is + // not re-derived (analytics metadata only); units + key are what bill. + meterReportingService.recordUsage( + row.getTeamId(), + stripeCustomerId, + row.getUnits() == null ? 0 : row.getUnits(), + null, + row.getIdempotencyKey(), + row.getJobId()); + retried++; + } + if (retried > 0) { + retriedCounter.increment(retried); + log.info("PaygMeterReconcileScheduler retried {} unposted meter event(s).", retried); + } + + long stuck = eventLogRepository.countStuck(floor); + if (stuck > 0) { + log.warn( + "{} PAYG meter event(s) stuck unposted past Stripe's {}h idempotency window —" + + " manual reconciliation needed.", + stuck, + STRIPE_IDEMPOTENCY_WINDOW.toHours()); + } + } +} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/meter/PaygMeterReportingService.java b/app/saas/src/main/java/stirling/software/saas/payg/meter/PaygMeterReportingService.java new file mode 100644 index 0000000000..e578edd15d --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/meter/PaygMeterReportingService.java @@ -0,0 +1,221 @@ +package stirling.software.saas.payg.meter; + +import java.util.Map; +import java.util.UUID; + +import org.springframework.beans.factory.annotation.Value; +import org.springframework.context.annotation.Profile; +import org.springframework.http.HttpEntity; +import org.springframework.http.HttpHeaders; +import org.springframework.http.HttpMethod; +import org.springframework.http.MediaType; +import org.springframework.http.ResponseEntity; +import org.springframework.stereotype.Service; +import org.springframework.web.client.RestTemplate; + +import io.micrometer.core.instrument.Counter; +import io.micrometer.core.instrument.MeterRegistry; + +import lombok.extern.slf4j.Slf4j; + +import stirling.software.saas.payg.model.BillingCategory; +import stirling.software.saas.payg.repository.PaygMeterEventLogRepository; + +/** + * POSTs PAYG billable usage to the Supabase {@code meter-payg-units} edge function. Called from + * {@code JobChargeService.close()} in an {@code afterCommit} hook, so the wallet ledger DEBIT (the + * customer's authoritative bill) is already durable before we tell Stripe about it. + * + *

Stripe is on a single flat-priced meter forever. {@link BillingCategory} ships as metadata for + * analytics — pricing never reads it. Free-tier teams (no Stripe subscription) skip this call + * entirely; the ledger entry is the only record needed. + * + *

Failure mode: we owe Stripe an event but the customer's bill via the ledger is correct. Log + * WARN, bump {@code payg.meter.errors}, and swallow — durability comes from the {@code + * payg_meter_event_log} row written around every attempt (pending → posted/failed) and {@link + * PaygMeterReconcileScheduler}, which retries unposted rows, not from retries here. Caller's {@code + * close()} must not roll back because Stripe wobbled. + * + *

Both config keys default to empty so unit tests / local dev never crash on missing env. When + * blank, this service no-ops at WARN-debug level — useful for SaaS smoke tests that don't want to + * touch the real edge function. + */ +@Service +@Profile("saas") +@Slf4j +public class PaygMeterReportingService { + + private final String endpoint; + private final String authToken; + private final RestTemplate restTemplate; + private final PaygMeterEventLogRepository eventLogRepository; + private final Counter errorsCounter; + + /** Stripe error bodies can be large; the column is TEXT but we cap to keep rows sane. */ + private static final int MAX_ERROR_BODY = 4000; + + public PaygMeterReportingService( + @Value("${payg.meter.endpoint:}") String endpoint, + @Value("${payg.meter.auth-token:}") String authToken, + RestTemplate saasRestTemplate, + PaygMeterEventLogRepository eventLogRepository, + MeterRegistry meterRegistry) { + this.endpoint = endpoint; + this.authToken = authToken; + this.restTemplate = saasRestTemplate; + this.eventLogRepository = eventLogRepository; + this.errorsCounter = + Counter.builder("payg.meter.errors") + .description("Failures POSTing PAYG meter events to Supabase edge function") + .register(meterRegistry); + } + + /** + * Best-effort POST of a single billable event, wrapped in a durable audit row. Idempotency on + * the Supabase side is keyed on {@code idempotency_key} — supply a deterministic value (e.g. + * {@code "process::close"}) so a retry, a reconcile replay, or a double-fire from two + * pods never charges twice. + * + *

Flow: write a pending {@code payg_meter_event_log} row (idempotent), POST to the edge fn, + * then stamp the row posted or record the Stripe error. The row is what {@link + * PaygMeterReconcileScheduler} retries, so a failed POST is recoverable rather than silently + * dropped. + * + *

Never throws. The wallet ledger entry is the source of truth for what the customer is + * billed; if the POST fails the only loss is that Stripe doesn't see this event until reconcile + * retries it. + */ + public void recordUsage( + Long teamId, + String stripeCustomerId, + int units, + BillingCategory category, + String idempotencyKey, + UUID jobId) { + if (endpoint == null || endpoint.isBlank()) { + log.debug( + "payg.meter.endpoint not configured; skipping meter event for team {} key {}", + teamId, + idempotencyKey); + return; + } + if (units <= 0) { + // Zero-unit events would inflate event count without changing the bill — defensive. + log.debug( + "Skipping meter event with units={} for team {} key {}", + units, + teamId, + idempotencyKey); + return; + } + + // Durable pending row before the POST so a failure leaves a record the reconcile scheduler + // can retry. Idempotent insert (ON CONFLICT DO NOTHING) — the completion + stale-close + // triggers and reconcile retries all share the key. Best-effort: a logging failure must + // never stop us from actually metering. + try { + eventLogRepository.insertPending(teamId, jobId, idempotencyKey, units); + } catch (Exception e) { + log.warn( + "payg_meter_event_log pending insert failed for key {} (still metering): {}", + idempotencyKey, + e.getMessage()); + } + + PostOutcome outcome = + postToStripe(teamId, stripeCustomerId, units, category, idempotencyKey); + + try { + if (outcome.success()) { + eventLogRepository.markPosted(idempotencyKey); + } else { + eventLogRepository.markFailed( + idempotencyKey, outcome.errorCode(), outcome.errorBody()); + } + } catch (Exception e) { + log.warn( + "payg_meter_event_log result update failed for key {}: {}", + idempotencyKey, + e.getMessage()); + } + } + + /** + * POST the event to the edge fn. Never throws; returns success / the captured Stripe error. + * Increments {@code payg.meter.errors} on any non-2xx or exception (unchanged metric contract). + */ + private PostOutcome postToStripe( + Long teamId, + String stripeCustomerId, + int units, + BillingCategory category, + String idempotencyKey) { + try { + HttpHeaders headers = new HttpHeaders(); + if (authToken != null && !authToken.isBlank()) { + headers.setBearerAuth(authToken); + } + headers.setContentType(MediaType.APPLICATION_JSON); + + Map body = + Map.of( + "team_id", + // JSON number — the edge fn type-checks and ignores strings. + teamId == null ? -1L : teamId, + "stripe_customer_id", + stripeCustomerId == null ? "" : stripeCustomerId, + "units", + units, + "idempotency_key", + idempotencyKey, + "metadata", + Map.of("category", category == null ? "UNKNOWN" : category.name())); + + ResponseEntity response = + restTemplate.exchange( + endpoint, + HttpMethod.POST, + new HttpEntity<>(body, headers), + String.class); + + if (response.getStatusCode().is2xxSuccessful()) { + return PostOutcome.ok(); + } + log.warn( + "Meter event POST returned {} for team {} key {}: {}", + response.getStatusCode(), + teamId, + idempotencyKey, + response.getBody()); + errorsCounter.increment(); + return PostOutcome.error( + String.valueOf(response.getStatusCode().value()), response.getBody()); + } catch (Exception e) { + // Catch-all by design: this method MUST NOT propagate. The customer's bill via the + // ledger is correct; we just owe Stripe an event the reconcile scheduler will retry. + log.warn( + "Meter event POST failed for team {} key {}: {}", + teamId, + idempotencyKey, + e.getMessage()); + errorsCounter.increment(); + return PostOutcome.error("exception", e.getMessage()); + } + } + + /** Outcome of one edge-fn POST attempt. */ + private record PostOutcome(boolean success, String errorCode, String errorBody) { + static PostOutcome ok() { + return new PostOutcome(true, null, null); + } + + static PostOutcome error(String code, String body) { + String trimmedCode = code != null && code.length() > 64 ? code.substring(0, 64) : code; + String trimmedBody = + body != null && body.length() > MAX_ERROR_BODY + ? body.substring(0, MAX_ERROR_BODY) + : body; + return new PostOutcome(false, trimmedCode, trimmedBody); + } + } +} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/model/BillingCategory.java b/app/saas/src/main/java/stirling/software/saas/payg/model/BillingCategory.java new file mode 100644 index 0000000000..2b1c0fee74 --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/model/BillingCategory.java @@ -0,0 +1,18 @@ +package stirling.software.saas.payg.model; + +/** + * Analytics / in-app breakdown axis stamped on every billable ledger entry and shadow charge. PAYG + * stays on a single flat-priced Stripe meter — category is metadata only, never affects pricing. + * + *

Precedence at translation time (interceptor): AUTOMATION → AI → API → BYPASSED. {@link + * #BYPASSED} is the default for manual UI tool calls that never hit a billable code path. + * + *

Listing order matters only as the default sentinel ({@link #BYPASSED} first); no downstream + * relies on {@code ordinal()}. + */ +public enum BillingCategory { + BYPASSED, + API, + AI, + AUTOMATION +} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/model/ShadowChargeStatus.java b/app/saas/src/main/java/stirling/software/saas/payg/model/ShadowChargeStatus.java new file mode 100644 index 0000000000..74653712e2 --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/model/ShadowChargeStatus.java @@ -0,0 +1,18 @@ +package stirling.software.saas.payg.model; + +/** + * Lifecycle state for a {@code payg_shadow_charge} row. Mimics the Stripe meter_event_adjustment + * (type=cancel) mechanism that real-mode will invoke when a freshly-opened process fails on its + * first step. The reconciliation report's "true net" query is {@code SUM(payg_units) WHERE status = + * 'CHARGED'}. + */ +public enum ShadowChargeStatus { + /** Default at process open; the would-be charge stands. */ + CHARGED, + /** + * Set when a 5xx first-step failure was observed in the same request's {@code afterCompletion}. + * Real-mode equivalent is a successful Stripe {@code meter_event_adjustment} posted in the same + * transaction-after-commit hook. + */ + REFUNDED +} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/policy/PaygTeamExtensions.java b/app/saas/src/main/java/stirling/software/saas/payg/policy/PaygTeamExtensions.java index 2c0603fc80..c564f85562 100644 --- a/app/saas/src/main/java/stirling/software/saas/payg/policy/PaygTeamExtensions.java +++ b/app/saas/src/main/java/stirling/software/saas/payg/policy/PaygTeamExtensions.java @@ -58,6 +58,29 @@ public class PaygTeamExtensions implements Serializable { @Column(name = "stripe_customer_id", unique = true, length = 128) private String stripeCustomerId; + /** + * Stripe subscription id (sub_xxx) for this team's PAYG metered subscription. {@code null} + * means the team has not added a card yet — engine writes shadow rows only and free-tier gating + * applies. {@code non-null} means engine posts meter events to Stripe on every billable tool + * call. + * + *

This is the single switch that determines whether a team is billed. Mutated only by the + * {@code payg_link_subscription} / {@code payg_unlink_subscription} Postgres functions (V14) — + * never directly via JPA writes. Treat as read-only from Java. + */ + @Column(name = "payg_subscription_id", unique = true, length = 128) + private String paygSubscriptionId; + + /** + * Remaining one-time free documents for this team (the lifetime grant). Seeded from the + * effective pricing policy's {@code free_tier_units} when this row is created (V14 trigger, + * updated in V19); decremented by the charge pipeline when a billable charge is written and + * restored on a first-step refund. Never replenishes; survives subscribing. This counter — not + * the wallet ledger — is the source of truth for the grant, so old ledger rows can be pruned. + */ + @Column(name = "free_units_remaining", nullable = false) + private Long freeUnitsRemaining = 0L; + @CreationTimestamp @Column(name = "created_at", updatable = false) private LocalDateTime createdAt; diff --git a/app/saas/src/main/java/stirling/software/saas/payg/policy/PricingPolicy.java b/app/saas/src/main/java/stirling/software/saas/payg/policy/PricingPolicy.java index f1443664cc..9581177369 100644 --- a/app/saas/src/main/java/stirling/software/saas/payg/policy/PricingPolicy.java +++ b/app/saas/src/main/java/stirling/software/saas/payg/policy/PricingPolicy.java @@ -73,6 +73,16 @@ public class PricingPolicy implements Serializable { @Column(name = "file_unit_cap", nullable = false) private Integer fileUnitCap = 1000; + /** + * One-time lifetime free document grant handed to a team on creation. {@code 0} (default) means + * no free grant. NOT per-cycle: it never replenishes and a team keeps any unused portion after + * subscribing. The value is copied into {@code payg_team_extensions.free_units_remaining} when + * the team's sidecar row is created (V14 trigger, updated in V19); from then on the per-team + * counter is authoritative and this column is only the seed for new teams. + */ + @Column(name = "free_tier_units", nullable = false) + private Long freeTierUnits = 0L; + /** * Max tool steps allowed in one process before it splits, keyed by the caller's {@link * JobSource}. Self-hosted teams typically get a higher limit via a per-team policy override. diff --git a/app/saas/src/main/java/stirling/software/saas/payg/policy/PricingPolicyService.java b/app/saas/src/main/java/stirling/software/saas/payg/policy/PricingPolicyService.java index e273a391d9..0fd7bb45d1 100644 --- a/app/saas/src/main/java/stirling/software/saas/payg/policy/PricingPolicyService.java +++ b/app/saas/src/main/java/stirling/software/saas/payg/policy/PricingPolicyService.java @@ -79,19 +79,28 @@ public class PricingPolicyService { * {@code is_default = TRUE}. Throws {@link IllegalStateException} if no default exists — the * seed migration is expected to put one there. * + *

A {@code null} {@code teamId} (admin-created user, deleted team, pre-team-migration + * account) is valid and returns the default policy directly. Throwing here would silently + * exclude the team-less cohort from PAYG via the filter's fail-open path, biasing the + * shadow-vs-legacy reconciliation. + * *

{@link Transactional}({@code readOnly = true}) so the eager-loaded {@code stepLimits} and * {@code stripePriceIds} collections initialize inside the same session. */ @Transactional(readOnly = true) public PricingPolicy getEffectivePolicy(Long teamId) { - Objects.requireNonNull(teamId, "teamId"); + if (teamId == null) { + return loadDefaultPolicy(); + } return byTeamCache.get(teamId, this::loadEffectivePolicy); } /** Bypasses the cache. Useful for admin endpoints that want a fresh read after a mutation. */ @Transactional(readOnly = true) public PricingPolicy getEffectivePolicyUncached(Long teamId) { - Objects.requireNonNull(teamId, "teamId"); + if (teamId == null) { + return loadDefaultPolicy(); + } return loadEffectivePolicy(teamId); } @@ -243,6 +252,10 @@ public class PricingPolicyService { teamId, id); } + return loadDefaultPolicy(); + } + + private PricingPolicy loadDefaultPolicy() { return policyRepository .findFirstByIsDefaultTrue() .orElseThrow( diff --git a/app/saas/src/main/java/stirling/software/saas/payg/repository/PaygMeterEventLogRepository.java b/app/saas/src/main/java/stirling/software/saas/payg/repository/PaygMeterEventLogRepository.java new file mode 100644 index 0000000000..7d752bd14f --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/repository/PaygMeterEventLogRepository.java @@ -0,0 +1,81 @@ +package stirling.software.saas.payg.repository; + +import java.time.LocalDateTime; +import java.util.List; +import java.util.UUID; + +import org.springframework.data.domain.Pageable; +import org.springframework.data.jpa.repository.JpaRepository; +import org.springframework.data.jpa.repository.Modifying; +import org.springframework.data.jpa.repository.Query; +import org.springframework.data.repository.query.Param; +import org.springframework.stereotype.Repository; +import org.springframework.transaction.annotation.Transactional; + +import stirling.software.saas.payg.meter.PaygMeterEventLog; + +@Repository +public interface PaygMeterEventLogRepository extends JpaRepository { + + /** + * Idempotent pending insert. The two meter triggers (charge-completion + stale-close) share the + * idempotency key, and a reconcile retry re-runs the same path, so {@code ON CONFLICT DO + * NOTHING} keeps the audit row at exactly one per key. {@code occurred_at} defaults to {@code + * CURRENT_TIMESTAMP} at the DB; {@code posted_to_stripe_at} stays NULL until success. + * Transactional because it's called from the non-transactional meter-reporting path. + */ + @Transactional + @Modifying + @Query( + value = + "INSERT INTO payg_meter_event_log (team_id, job_id, idempotency_key, units)" + + " VALUES (:teamId, :jobId, :key, :units)" + + " ON CONFLICT (idempotency_key) DO NOTHING", + nativeQuery = true) + void insertPending( + @Param("teamId") Long teamId, + @Param("jobId") UUID jobId, + @Param("key") String key, + @Param("units") int units); + + /** Stamp an event posted; clears any prior error. No-op if already posted. */ + @Transactional + @Modifying + @Query( + "UPDATE PaygMeterEventLog e" + + " SET e.postedToStripeAt = CURRENT_TIMESTAMP," + + " e.stripeErrorCode = NULL, e.stripeErrorBody = NULL" + + " WHERE e.idempotencyKey = :key AND e.postedToStripeAt IS NULL") + int markPosted(@Param("key") String key); + + /** Record the latest Stripe error against a still-pending event. */ + @Transactional + @Modifying + @Query( + "UPDATE PaygMeterEventLog e" + + " SET e.stripeErrorCode = :code, e.stripeErrorBody = :body" + + " WHERE e.idempotencyKey = :key AND e.postedToStripeAt IS NULL") + int markFailed( + @Param("key") String key, @Param("code") String code, @Param("body") String body); + + /** + * Events still unposted whose attempt is older than {@code cutoff} (give the live POST time to + * land first) but within {@code floor} — Stripe's 24h idempotency window — so a retry under the + * same key safely dedups rather than risking a double charge. Oldest first. + */ + @Query( + "SELECT e FROM PaygMeterEventLog e" + + " WHERE e.postedToStripeAt IS NULL" + + " AND e.occurredAt < :cutoff AND e.occurredAt >= :floor" + + " ORDER BY e.occurredAt ASC") + List findRetryable( + @Param("cutoff") LocalDateTime cutoff, + @Param("floor") LocalDateTime floor, + Pageable pageable); + + /** Events stuck unposted past the safe retry window — surfaced for manual reconciliation. */ + @Query( + "SELECT COUNT(e) FROM PaygMeterEventLog e" + + " WHERE e.postedToStripeAt IS NULL AND e.occurredAt < :floor") + long countStuck(@Param("floor") LocalDateTime floor); +} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/repository/PaygShadowChargeRepository.java b/app/saas/src/main/java/stirling/software/saas/payg/repository/PaygShadowChargeRepository.java index ec1a51963a..8a2911f9c4 100644 --- a/app/saas/src/main/java/stirling/software/saas/payg/repository/PaygShadowChargeRepository.java +++ b/app/saas/src/main/java/stirling/software/saas/payg/repository/PaygShadowChargeRepository.java @@ -2,6 +2,8 @@ package stirling.software.saas.payg.repository; import java.time.LocalDateTime; import java.util.List; +import java.util.Optional; +import java.util.UUID; import org.springframework.data.jpa.repository.JpaRepository; import org.springframework.data.jpa.repository.Query; @@ -19,4 +21,28 @@ public interface PaygShadowChargeRepository extends JpaRepository findInWindow( @Param("from") LocalDateTime from, @Param("to") LocalDateTime to); + + /** + * The shadow row written when the given process was opened. At most one row per {@code jobId} + * exists by construction — {@code openProcess} writes exactly one row on OPENED and zero on + * JOINED — so callers can treat the result as a single optional. Returns the first row by id + * defensively if a duplicate ever appears. + */ + Optional findFirstByJobIdOrderByIdAsc(UUID jobId); + + /** + * Paid (Stripe-metered) documents for a team in a period: {@code SUM(payg_units − + * free_units_consumed)} over CHARGED rows. This is exactly what was reported to Stripe in the + * window, so the wallet's "estimated bill so far" is the metered total × rate. REFUNDED rows + * are excluded. + */ + @Query( + "SELECT COALESCE(SUM(s.paygUnits - s.freeUnitsConsumed), 0) FROM PaygShadowCharge s" + + " WHERE s.teamId = :teamId" + + " AND s.status = stirling.software.saas.payg.model.ShadowChargeStatus.CHARGED" + + " AND s.occurredAt >= :from AND s.occurredAt < :to") + long sumPaidUnits( + @Param("teamId") Long teamId, + @Param("from") LocalDateTime from, + @Param("to") LocalDateTime to); } diff --git a/app/saas/src/main/java/stirling/software/saas/payg/repository/PaygTeamExtensionsRepository.java b/app/saas/src/main/java/stirling/software/saas/payg/repository/PaygTeamExtensionsRepository.java index 3473eeeff7..35fb6bbf7a 100644 --- a/app/saas/src/main/java/stirling/software/saas/payg/repository/PaygTeamExtensionsRepository.java +++ b/app/saas/src/main/java/stirling/software/saas/payg/repository/PaygTeamExtensionsRepository.java @@ -3,12 +3,40 @@ package stirling.software.saas.payg.repository; import java.util.Optional; import org.springframework.data.jpa.repository.JpaRepository; +import org.springframework.data.jpa.repository.Lock; +import org.springframework.data.jpa.repository.Modifying; +import org.springframework.data.jpa.repository.Query; +import org.springframework.data.repository.query.Param; import org.springframework.stereotype.Repository; +import jakarta.persistence.LockModeType; + import stirling.software.saas.payg.policy.PaygTeamExtensions; @Repository public interface PaygTeamExtensionsRepository extends JpaRepository { Optional findByStripeCustomerId(String stripeCustomerId); + + /** + * Pessimistic-write load of the sidecar row, used by the charge pipeline to deduct the one-time + * free grant atomically. The lock serialises concurrent charges for the same team so + * the per-job {@code free_units_consumed} split (and therefore the metered paid portion) is + * exact — two simultaneous jobs can't both believe they drew from the same remaining unit. + * Different teams never contend; the lock is held only for the {@code openProcess} transaction. + */ + @Lock(LockModeType.PESSIMISTIC_WRITE) + @Query("SELECT e FROM PaygTeamExtensions e WHERE e.teamId = :teamId") + Optional findByIdForUpdate(@Param("teamId") Long teamId); + + /** + * Atomically returns {@code freeUnitsConsumed} to the team's grant on a refund. Increment is + * commutative so no lock is needed; the amount restored is exactly what the job consumed, so it + * can never exceed the original grant. + */ + @Modifying + @Query( + "UPDATE PaygTeamExtensions e SET e.freeUnitsRemaining = e.freeUnitsRemaining + :units" + + " WHERE e.teamId = :teamId") + int restoreFreeUnits(@Param("teamId") Long teamId, @Param("units") long units); } diff --git a/app/saas/src/main/java/stirling/software/saas/payg/repository/WalletLedgerRepository.java b/app/saas/src/main/java/stirling/software/saas/payg/repository/WalletLedgerRepository.java index 2a2a003429..095f6cb924 100644 --- a/app/saas/src/main/java/stirling/software/saas/payg/repository/WalletLedgerRepository.java +++ b/app/saas/src/main/java/stirling/software/saas/payg/repository/WalletLedgerRepository.java @@ -14,7 +14,30 @@ import stirling.software.saas.payg.wallet.WalletLedgerEntry; @Repository public interface WalletLedgerRepository extends JpaRepository { - List findByTeamIdOrderByOccurredAtDesc(Long teamId); + /** Most recent entries for the Plan page activity feed. */ + List findTop20ByTeamIdOrderByIdDesc(Long teamId); + + /** + * Per-category debit totals over an arbitrary window, as positive units. Replaces the + * calendar-month {@code wallet_category_summary} view on the wallet endpoint — subscribed + * teams' billing windows are anchored to the Stripe subscription period, not month starts. Rows + * with {@code NULL} category (system entries) are excluded; BYPASSED never reaches the ledger + * by construction. + */ + @Query( + "SELECT e.billingCategory AS category, COALESCE(SUM(-e.amountUnits), 0) AS units" + + " FROM WalletLedgerEntry e" + + " WHERE e.teamId = :teamId" + + " AND e.entryType = :entryType" + + " AND e.billingCategory IS NOT NULL" + + " AND e.occurredAt >= :periodStart" + + " AND e.occurredAt < :periodEnd" + + " GROUP BY e.billingCategory") + List sumPeriodAmountByCategory( + @Param("teamId") Long teamId, + @Param("entryType") LedgerEntryType entryType, + @Param("periodStart") LocalDateTime periodStart, + @Param("periodEnd") LocalDateTime periodEnd); /** Sum of signed amounts over a team's entries — the wallet's current balance in units. */ @Query( @@ -34,6 +57,24 @@ public interface WalletLedgerRepository extends JpaRepository= :periodStart" + + " AND e.occurredAt < :periodEnd") + long sumPeriodNetBillable( + @Param("teamId") Long teamId, + @Param("periodStart") LocalDateTime periodStart, + @Param("periodEnd") LocalDateTime periodEnd); + /** Per-member period spend (only when the member has a sub-cap configured). */ @Query( "SELECT COALESCE(SUM(e.amountUnits), 0) FROM WalletLedgerEntry e" diff --git a/app/saas/src/main/java/stirling/software/saas/payg/shadow/PaygShadowCharge.java b/app/saas/src/main/java/stirling/software/saas/payg/shadow/PaygShadowCharge.java index 495930c08c..57aea1a0e4 100644 --- a/app/saas/src/main/java/stirling/software/saas/payg/shadow/PaygShadowCharge.java +++ b/app/saas/src/main/java/stirling/software/saas/payg/shadow/PaygShadowCharge.java @@ -8,6 +8,8 @@ import org.hibernate.annotations.CreationTimestamp; import jakarta.persistence.Column; import jakarta.persistence.Entity; +import jakarta.persistence.EnumType; +import jakarta.persistence.Enumerated; import jakarta.persistence.GeneratedValue; import jakarta.persistence.GenerationType; import jakarta.persistence.Id; @@ -17,10 +19,18 @@ import lombok.Getter; import lombok.NoArgsConstructor; import lombok.Setter; +import stirling.software.saas.payg.model.BillingCategory; +import stirling.software.saas.payg.model.JobSource; +import stirling.software.saas.payg.model.ShadowChargeStatus; + /** * Per-job comparison row written while a team is in {@code PAYG_SHADOW} mode: what the legacy * engine actually charged vs. what the PAYG engine would have charged. Aggregated daily by the * shadow-reconciliation report; deletable after promotion. + * + *

The {@link #status} column mirrors the eventual Stripe meter_event_adjustment(cancel) flow: + * rows are written {@code CHARGED}; a freshly-opened process whose first step 5xx's is flipped to + * {@code REFUNDED} in the same request's {@code afterCompletion}. */ @Entity @Table(name = "payg_shadow_charge") @@ -48,6 +58,15 @@ public class PaygShadowCharge implements Serializable { @Column(name = "payg_units", nullable = false) private Integer paygUnits; + /** + * How many of {@link #paygUnits} were drawn from the team's one-time free grant at charge time. + * The paid (Stripe-metered) portion is {@code paygUnits - freeUnitsConsumed}; a refund restores + * this many units to {@code payg_team_extensions.free_units_remaining}. {@code 0} for pre-V19 + * rows and for jobs that consumed no free units (team's grant already exhausted). + */ + @Column(name = "free_units_consumed", nullable = false) + private Integer freeUnitsConsumed = 0; + @Column(name = "legacy_credits_charged", nullable = false) private Integer legacyCreditsCharged; @@ -55,6 +74,36 @@ public class PaygShadowCharge implements Serializable { @Column(name = "diff_pct", nullable = false) private Integer diffPct; + @Enumerated(EnumType.STRING) + @Column(name = "status", nullable = false, length = 16) + private ShadowChargeStatus status = ShadowChargeStatus.CHARGED; + + /** Set when {@link #status} flips to {@link ShadowChargeStatus#REFUNDED}. */ + @Column(name = "refunded_at") + private LocalDateTime refundedAt; + + /** Free-form reason, e.g. {@code "first-step-5xx:503"}. */ + @Column(name = "refund_reason", length = 128) + private String refundReason; + + /** + * PAYG analytics axis copied from the request that produced this shadow row. {@code null} for + * pre-V16 rows; populated by the charge interceptor going forward. Never affects what Stripe + * would meter — the flat-meter assumption holds, this is breakdown metadata. + */ + @Enumerated(EnumType.STRING) + @Column(name = "billing_category", length = 16) + private BillingCategory billingCategory; + + /** + * Caller surface that originated the job (copied from {@code processing_job.source} at write + * time so the row stays self-describing after the job table is pruned). {@code null} for + * pre-V16 rows backfilled when {@code processing_job} is no longer present. + */ + @Enumerated(EnumType.STRING) + @Column(name = "job_source", length = 32) + private JobSource jobSource; + @CreationTimestamp @Column(name = "occurred_at", nullable = false, updatable = false) private LocalDateTime occurredAt; diff --git a/app/saas/src/main/java/stirling/software/saas/payg/stripe/StripeSubscriptionDao.java b/app/saas/src/main/java/stirling/software/saas/payg/stripe/StripeSubscriptionDao.java new file mode 100644 index 0000000000..5859c2099c --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/stripe/StripeSubscriptionDao.java @@ -0,0 +1,215 @@ +package stirling.software.saas.payg.stripe; + +import java.math.BigDecimal; +import java.sql.ResultSet; +import java.sql.SQLException; +import java.time.Instant; +import java.time.LocalDateTime; +import java.time.ZoneId; +import java.util.List; +import java.util.Objects; +import java.util.Optional; + +import org.springframework.context.annotation.Profile; +import org.springframework.dao.DataAccessException; +import org.springframework.jdbc.core.JdbcTemplate; +import org.springframework.stereotype.Repository; + +import lombok.extern.slf4j.Slf4j; + +/** + * Read-only accessor for the Stripe Sync Engine schema ({@code stripe.*}). Gives the PAYG layer the + * team's real billing window and the per-document rate of the Price its subscription bills against + * - both live in Stripe, mirrored into Postgres by the sync engine, and are NOT duplicated in + * {@code stirling_pdf} (money lives in Stripe). + * + *

PAYG prices are plain {@code per_unit} metered prices, so {@code stripe.prices.unit_amount} + * carries the rate directly. The free grant is deliberately NOT in Stripe - it's the one-time + * {@code pricing_policy.free_tier_units} pool, applied app-side (free units are never metered), + * because un-subscribed teams get the same grant and have no Stripe Price at all. + * + *

Defensive by construction: the {@code stripe} schema only exists where the sync engine has run + * (dev + prod Supabase, not unit-test H2). Any {@link DataAccessException} - missing schema, + * missing row, connectivity blip - degrades to {@link Optional#empty()} with a WARN so callers fall + * back to calendar-month windows rather than 500ing the wallet endpoint. + */ +@Slf4j +@Repository +@Profile("saas") +public class StripeSubscriptionDao { + + /** + * Billing window + per-document rate of one subscription. Epoch seconds come back from sync + * engine as INTEGER columns; converted to {@link LocalDateTime} in the system zone so they + * compare cleanly against {@code wallet_ledger.occurred_at} (written by + * {@code @CreationTimestamp} with JVM-local semantics). + * + * @param priceId the Stripe Price the subscription's (sole) item bills against; null if the + * item row hasn't synced yet + * @param currency lower-case ISO 4217 of that Price; null when the price row is missing + * @param perDocMinor per-document rate in minor units (may be fractional via {@code + * unit_amount_decimal}); null when the price row is missing or carries no usable amount + * (e.g. a tiered price, which PAYG doesn't use) + */ + public record SubscriptionBilling( + LocalDateTime periodStart, + LocalDateTime periodEnd, + String priceId, + String status, + String currency, + BigDecimal perDocMinor) {} + + /** + * The per-document rate of a Price looked up directly (not via a subscription) - used to price + * the cap estimate for un-subscribed teams, whose default policy points at Stripe Prices that + * carry the same {@code unit_amount} they'd be billed at on subscribing. + * + * @param priceId the resolved Stripe Price id + * @param currency lower-case ISO 4217 of that Price + * @param perDocMinor per-document rate in minor units (may be fractional); never null - a row + * with no usable amount is filtered out rather than returned with a null rate + */ + public record PriceRate(String priceId, String currency, BigDecimal perDocMinor) {} + + private static final String QUERY = + "SELECT s.current_period_start, s.current_period_end, s.status::text AS status," + + " p.id AS price_id, p.currency AS currency," + + " p.unit_amount AS unit_amount, p.unit_amount_decimal AS unit_amount_decimal" + + " FROM stripe.subscriptions s" + + " LEFT JOIN LATERAL (" + + " SELECT si.price FROM stripe.subscription_items si" + + " WHERE si.subscription = s.id AND COALESCE(si.deleted, false) = false" + + " ORDER BY si.created DESC NULLS LAST LIMIT 1" + + " ) item ON true" + + " LEFT JOIN stripe.prices p ON p.id = item.price" + + " WHERE s.id = ?"; + + private final JdbcTemplate jdbcTemplate; + + public StripeSubscriptionDao(JdbcTemplate jdbcTemplate) { + this.jdbcTemplate = Objects.requireNonNull(jdbcTemplate, "jdbcTemplate"); + } + + /** Billing window + rate for {@code subscriptionId}; empty if unsynced or schema absent. */ + public Optional findBilling(String subscriptionId) { + if (subscriptionId == null || subscriptionId.isBlank()) { + return Optional.empty(); + } + try { + List rows = + jdbcTemplate.query( + QUERY, + (rs, i) -> { + long startEpoch = rs.getLong("current_period_start"); + boolean startNull = rs.wasNull(); + long endEpoch = rs.getLong("current_period_end"); + boolean endNull = rs.wasNull(); + if (startNull || endNull) { + return null; + } + return new SubscriptionBilling( + toLocal(startEpoch), + toLocal(endEpoch), + rs.getString("price_id"), + rs.getString("status"), + rs.getString("currency"), + extractRate(rs)); + }, + subscriptionId); + return rows.stream().filter(Objects::nonNull).findFirst(); + } catch (DataAccessException e) { + // Missing stripe schema (sync engine not provisioned) or transient DB issue. The + // caller falls back to a calendar-month window; usage still accrues correctly. + log.warn( + "stripe.subscriptions lookup failed for {}: {}", + subscriptionId, + e.getMessage()); + return Optional.empty(); + } + } + + /** + * Per-document rate of the active {@code stripe.prices} row with the given {@code lookupKey} in + * {@code currency} - the elegant mirror of {@link #findBilling}, reading the same synced table + * by Stripe Price {@code lookup_key} instead of via a subscription. The PAYG layer uses this to + * price the cap estimate for an un-subscribed team: there's no subscription to read a rate off, + * but the PAYG Price (lookup key {@code plan:processor}) carries the very rate they'd be billed + * at. We resolve by lookup_key rather than the default policy's price ids because those aren't + * seeded - the lookup key is the stable, env-agnostic handle (same one the price-lookup edge + * function uses). + * + *

Empty when the {@code stripe} schema is absent (H2 unit tests), no active matching row + * exists, or the row carries no usable per-unit amount (e.g. a tiered price, which PAYG doesn't + * use). Callers degrade to "no estimate" exactly as the subscribed path degrades to a + * calendar-month window. + */ + public Optional findRateByLookupKey(String lookupKey, String currency) { + if (lookupKey == null || lookupKey.isBlank() || currency == null || currency.isBlank()) { + return Optional.empty(); + } + String sql = + "SELECT p.id AS price_id, p.currency AS currency," + + " p.unit_amount AS unit_amount," + + " p.unit_amount_decimal AS unit_amount_decimal" + + " FROM stripe.prices p" + + " WHERE p.lookup_key = ? AND p.currency = ?" + + " AND COALESCE(p.active, true) = true" + + " ORDER BY p.created DESC NULLS LAST LIMIT 1"; + try { + List rows = + jdbcTemplate.query( + sql, + (rs, i) -> { + BigDecimal rate = extractRate(rs); + if (rate == null) { + return null; + } + return new PriceRate( + rs.getString("price_id"), rs.getString("currency"), rate); + }, + lookupKey, + currency.toLowerCase()); + return rows.stream().filter(Objects::nonNull).findFirst(); + } catch (DataAccessException e) { + log.warn( + "stripe.prices rate lookup failed for lookup_key {} / currency {}: {}", + lookupKey, + currency, + e.getMessage()); + return Optional.empty(); + } + } + + /** + * Reads the per-document rate off a {@code stripe.prices} row. Prefers {@code + * unit_amount_decimal} (sub-minor-unit precision, e.g. half-cent rates), falls back to the + * integer {@code unit_amount}, and returns null when neither is usable or the amount is ≤ 0 + * (tiered/zero prices PAYG doesn't bill on). Both queries alias the columns identically so this + * is shared verbatim. + */ + private static BigDecimal extractRate(ResultSet rs) throws SQLException { + BigDecimal rate = null; + String decimal = rs.getString("unit_amount_decimal"); + if (decimal != null && !decimal.isBlank()) { + try { + rate = new BigDecimal(decimal); + } catch (NumberFormatException ignore) { + rate = null; + } + } + if (rate == null) { + long unitAmount = rs.getLong("unit_amount"); + if (!rs.wasNull()) { + rate = BigDecimal.valueOf(unitAmount); + } + } + if (rate != null && rate.signum() <= 0) { + rate = null; + } + return rate; + } + + private static LocalDateTime toLocal(long epochSeconds) { + return LocalDateTime.ofInstant(Instant.ofEpochSecond(epochSeconds), ZoneId.systemDefault()); + } +} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/test/PaygCucumberThrowController.java b/app/saas/src/main/java/stirling/software/saas/payg/test/PaygCucumberThrowController.java new file mode 100644 index 0000000000..8a9082e2c8 --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/payg/test/PaygCucumberThrowController.java @@ -0,0 +1,78 @@ +package stirling.software.saas.payg.test; + +import org.springframework.context.annotation.Profile; +import org.springframework.http.MediaType; +import org.springframework.http.ResponseEntity; +import org.springframework.web.bind.annotation.RequestMapping; +import org.springframework.web.bind.annotation.RequestParam; +import org.springframework.web.bind.annotation.RestController; +import org.springframework.web.multipart.MultipartFile; + +import io.swagger.v3.oas.annotations.Hidden; + +import lombok.extern.slf4j.Slf4j; + +import stirling.software.common.annotations.AutoJobPostMapping; +import stirling.software.common.enumeration.ResourceWeight; + +/** + * Cucumber-only force-5xx endpoint. Gated behind the {@code payg-cucumber} Spring profile so the + * bean never registers in production — only the PAYG cucumber compose stack activates the profile + * (via {@code SPRING_PROFILES_ACTIVE=saas,payg-cucumber} in {@code docker-compose-saas.yml}). + * + *

Purpose: drive the PAYG filter+interceptor's 5xx-first-step branch end-to-end. No reliably- + * 5xx-ing real tool endpoint exists in current Stirling — every malformed input is caught as 4xx by + * {@code GlobalExceptionHandler}. Without this stub the only way to exercise the refund path was a + * manual procedure (a temporary throw endpoint added, run, removed) documented in {@code + * notes/PAYG_DESIGN.md} §7.5.2 M1. This controller replaces that procedure with a profile- gated + * automated scenario. + * + *

{@link AutoJobPostMapping} consumes {@code multipart/form-data} so the filter's {@code + * MultipartHttpServletRequest} cast runs and the input lineage hash is computed before the + * controller throws. The PAYG filter chain therefore sees: preHandle → openProcess (CHARGED row + * written) → controller throws → afterCompletion observes status 500 → markFirstStepFailed (row → + * REFUNDED, job → CLOSED). + * + *

The thrown {@link IllegalStateException} is unwrapped by {@code GlobalExceptionHandler}'s + * RuntimeException handler — it has no IOException / IllegalArgument / BaseAppException cause, so + * it falls through to the "Unexpected RuntimeException" branch which sets HTTP 500. Don't use + * {@code ResponseStatusException} here — its dedicated handler is reached before the unwrapping + * branch but the upstream test still passes; {@link IllegalStateException} is the more honest mimic + * of a real bug-driven 500. + * + *

{@code @Hidden} keeps this endpoint out of OpenAPI / Swagger output even when the profile is + * active, so no docs leak to anyone running the cucumber stack interactively. + * + *

Return type matters: the method declares {@link ResponseEntity}{@code }, not + * {@code void}. Spring's {@code InvocableHandlerMethod} picks a return-value handler based on the + * declared return type; {@code void} matches the "no body, status 200" handler regardless of + * what the {@code @Around} advice on {@link AutoJobPostMapping} actually returns at runtime. The + * first version of this stub declared {@code void} and the 500 ResponseEntity from {@code + * JobExecutorService}'s catch block was silently discarded — the client saw a 200 with an empty + * body. Declaring {@code ResponseEntity} matches the real runtime type and lets the advice's + * 500 reach the wire. + */ +@Slf4j +@RestController +@Profile("payg-cucumber") +@RequestMapping("/api/v1/payg-cucumber") +@Hidden +public class PaygCucumberThrowController { + + @AutoJobPostMapping( + value = "/throw-500", + consumes = MediaType.MULTIPART_FORM_DATA_VALUE, + resourceWeight = ResourceWeight.SMALL_WEIGHT) + public ResponseEntity throw500( + @RequestParam(value = "fileInput", required = false) MultipartFile fileInput) { + // The file is read by the PAYG filter via getMultiFileMap() before we get here; the + // controller param is just to keep Spring's multipart binding happy. We don't touch it. + log.warn( + "PAYG cucumber forced 500 (fileInput name='{}' size={} bytes)", + fileInput != null ? fileInput.getOriginalFilename() : null, + fileInput != null ? fileInput.getSize() : 0); + throw new IllegalStateException("PAYG cucumber forced 500"); + // unreachable — kept as a type signature so AutoJobAspect's @Around return value (the + // 500 ResponseEntity from JobExecutorService) actually reaches the wire. + } +} diff --git a/app/saas/src/main/java/stirling/software/saas/payg/wallet/WalletLedgerEntry.java b/app/saas/src/main/java/stirling/software/saas/payg/wallet/WalletLedgerEntry.java index 90aed63cd3..4971e296f3 100644 --- a/app/saas/src/main/java/stirling/software/saas/payg/wallet/WalletLedgerEntry.java +++ b/app/saas/src/main/java/stirling/software/saas/payg/wallet/WalletLedgerEntry.java @@ -22,6 +22,7 @@ import lombok.Getter; import lombok.NoArgsConstructor; import lombok.Setter; +import stirling.software.saas.payg.model.BillingCategory; import stirling.software.saas.payg.model.LedgerBucket; import stirling.software.saas.payg.model.LedgerEntryType; import stirling.software.saas.payg.model.ReferenceType; @@ -76,6 +77,15 @@ public class WalletLedgerEntry implements Serializable { @Column(name = "stripe_event_id", length = 128) private String stripeEventId; + /** + * PAYG analytics axis. {@code null} for system entries (grants, resets) and pre-V16 rows; set + * by the charge interceptor for billable debits. Stripe pricing never reads this — it's a + * single flat meter — so the column stays soft-typed (no NOT NULL, no FK). + */ + @Enumerated(EnumType.STRING) + @Column(name = "billing_category", length = 16) + private BillingCategory billingCategory; + @CreationTimestamp @Column(name = "occurred_at", nullable = false, updatable = false) private LocalDateTime occurredAt; diff --git a/app/saas/src/main/java/stirling/software/saas/repository/TeamCreditRepository.java b/app/saas/src/main/java/stirling/software/saas/repository/TeamCreditRepository.java deleted file mode 100644 index 3757ce60d9..0000000000 --- a/app/saas/src/main/java/stirling/software/saas/repository/TeamCreditRepository.java +++ /dev/null @@ -1,68 +0,0 @@ -package stirling.software.saas.repository; - -import java.time.LocalDateTime; -import java.util.List; -import java.util.Optional; - -import org.springframework.data.jpa.repository.JpaRepository; -import org.springframework.data.jpa.repository.Modifying; -import org.springframework.data.jpa.repository.Query; -import org.springframework.data.repository.query.Param; -import org.springframework.stereotype.Repository; - -import stirling.software.saas.model.TeamCredit; - -@Repository -public interface TeamCreditRepository extends JpaRepository { - - /** Find team credits by team ID. */ - @Query("SELECT tc FROM TeamCredit tc WHERE tc.team.id = :teamId") - Optional findByTeamId(@Param("teamId") Long teamId); - - /** - * Atomically consume credits from the team pool. Uses the {@code @Version} column on {@link - * TeamCredit} for optimistic locking - concurrent attempts will fail-fast rather than - * over-deduct. Returns 1 on success, 0 if insufficient balance or version conflict. - */ - @Modifying - @Query( - value = - """ - UPDATE team_credits - SET cycle_credits_remaining = CASE - WHEN cycle_credits_remaining >= :amount THEN cycle_credits_remaining - :amount - WHEN cycle_credits_remaining > 0 AND bought_credits_remaining >= (:amount - cycle_credits_remaining) - THEN 0 - ELSE cycle_credits_remaining - END, - bought_credits_remaining = CASE - WHEN cycle_credits_remaining >= :amount THEN bought_credits_remaining - WHEN cycle_credits_remaining > 0 AND bought_credits_remaining >= (:amount - cycle_credits_remaining) - THEN bought_credits_remaining - (:amount - cycle_credits_remaining) - WHEN cycle_credits_remaining = 0 AND bought_credits_remaining >= :amount - THEN bought_credits_remaining - :amount - ELSE bought_credits_remaining - END, - total_api_calls_made = total_api_calls_made + :amount, - last_api_usage = CURRENT_TIMESTAMP, - updated_at = CURRENT_TIMESTAMP, - version = version + 1 - WHERE team_id = :teamId - AND (cycle_credits_remaining + bought_credits_remaining) >= :amount - """, - nativeQuery = true) - int consumeCredit(@Param("teamId") Long teamId, @Param("amount") int amount); - - @Query( - "SELECT CASE WHEN COUNT(tc) > 0 THEN true ELSE false END FROM TeamCredit tc WHERE tc.team.id = :teamId") - boolean existsByTeamId(@Param("teamId") Long teamId); - - @Modifying - @Query("DELETE FROM TeamCredit tc WHERE tc.team.id = :teamId") - void deleteByTeamId(@Param("teamId") Long teamId); - - @Query( - "SELECT tc FROM TeamCredit tc WHERE tc.lastCycleResetAt IS NULL OR tc.lastCycleResetAt < :lastScheduledReset") - List findCreditsNeedingCycleReset( - @Param("lastScheduledReset") LocalDateTime lastScheduledReset); -} diff --git a/app/saas/src/main/java/stirling/software/saas/repository/TeamMembershipRepository.java b/app/saas/src/main/java/stirling/software/saas/repository/TeamMembershipRepository.java index b3231224d2..36c04f09c3 100644 --- a/app/saas/src/main/java/stirling/software/saas/repository/TeamMembershipRepository.java +++ b/app/saas/src/main/java/stirling/software/saas/repository/TeamMembershipRepository.java @@ -42,6 +42,24 @@ public interface TeamMembershipRepository extends JpaRepository findByUserId(@Param("userId") Long userId); + /** + * Resolve the single membership a user belongs to. In the PAYG design every user is owned by + * exactly one team — personal-team-for-new-signups, then optionally migrated when they accept + * an invite (the old personal team is deleted on accept). For diagnostic safety this picks the + * earliest-created row if multiple exist, but in steady state there is exactly one. + * + *

Returns both the team and its role so the PAYG wallet endpoint can answer "what does this + * user see?" in a single query rather than a list-then-filter dance. + * + * @param userId the user ID + * @return the user's primary membership, if any + */ + @Query( + "SELECT tm FROM TeamMembership tm JOIN FETCH tm.team" + + " WHERE tm.user.id = :userId" + + " ORDER BY tm.createdAt ASC") + List findPrimaryMembership(@Param("userId") Long userId); + /** * Find all members with a specific role in a team * @@ -54,6 +72,16 @@ public interface TeamMembershipRepository extends JpaRepository findByTeamIdAndRole( @Param("teamId") Long teamId, @Param("role") TeamRole role); + /** + * Count members with a specific role in a team. Lighter than {@link #findByTeamIdAndRole} when + * only the tally is needed (e.g. last-leader checks) — avoids fetching and join-loading rows. + * + * @param teamId the team ID + * @param role the team role (LEADER or MEMBER) + * @return number of members with that role + */ + long countByTeamIdAndRole(Long teamId, TeamRole role); + /** * Check if a user is a member of a team * diff --git a/app/saas/src/main/java/stirling/software/saas/repository/UserCreditRepository.java b/app/saas/src/main/java/stirling/software/saas/repository/UserCreditRepository.java deleted file mode 100644 index 49a0042b19..0000000000 --- a/app/saas/src/main/java/stirling/software/saas/repository/UserCreditRepository.java +++ /dev/null @@ -1,148 +0,0 @@ -package stirling.software.saas.repository; - -import java.time.LocalDateTime; -import java.util.List; -import java.util.Optional; -import java.util.UUID; - -import org.springframework.data.jpa.repository.JpaRepository; -import org.springframework.data.jpa.repository.Modifying; -import org.springframework.data.jpa.repository.Query; -import org.springframework.data.repository.query.Param; -import org.springframework.stereotype.Repository; - -import stirling.software.proprietary.security.model.User; -import stirling.software.saas.model.UserCredit; - -/** - * JPA repository for {@link UserCredit}. Includes JPQL queries for the common read paths and native - * SQL for atomic credit-consumption updates (avoids select-then-update races). - * - *

Native queries reference {@code user_credits} and {@code users} unqualified — they pick up - * Hibernate's {@code default_schema} (set to {@code stirling_pdf} in {@code - * application-saas.properties}). Keeping the schema out of the SQL means a future schema rename is - * a one-property change instead of a sweep of native SQL. - */ -@Repository -public interface UserCreditRepository extends JpaRepository { - - Optional findByUser(User user); - - Optional findByUserId(Long userId); - - @Query( - "SELECT uc FROM UserCredit uc WHERE uc.lastCycleResetAt IS NULL OR uc.lastCycleResetAt < :lastScheduledReset") - List findCreditsNeedingCycleReset( - @Param("lastScheduledReset") LocalDateTime lastScheduledReset); - - @Query("SELECT SUM(uc.totalApiCallsMade) FROM UserCredit uc") - Long getTotalApiCallsAcrossAllUsers(); - - @Query("SELECT SUM(uc.cycleCreditsRemaining + uc.boughtCreditsRemaining) FROM UserCredit uc") - Long getTotalAvailableCreditsAcrossAllUsers(); - - @Query("SELECT uc FROM UserCredit uc WHERE uc.user.apiKey = :apiKey") - Optional findByUserApiKey(@Param("apiKey") String apiKey); - - @Query("SELECT uc FROM UserCredit uc WHERE uc.user.supabaseId = :supabaseId") - Optional findBySupabaseId(@Param("supabaseId") UUID supabaseId); - - @Query("SELECT COUNT(uc) FROM UserCredit uc WHERE uc.lastApiUsage >= :since") - Long countActiveUsersInPeriod(@Param("since") LocalDateTime since); - - @Modifying - @Query( - value = - "UPDATE user_credits " - + "SET " - + " cycle_credits_remaining = " - + " CASE " - + " WHEN cycle_credits_remaining >= :creditAmount THEN cycle_credits_remaining - :creditAmount " - + " ELSE 0 " - + " END, " - + " bought_credits_remaining = " - + " CASE " - + " WHEN cycle_credits_remaining < :creditAmount " - + " THEN GREATEST(0, bought_credits_remaining - (:creditAmount - cycle_credits_remaining)) " - + " ELSE bought_credits_remaining " - + " END, " - + " total_api_calls_made = total_api_calls_made + 1, " - + " last_api_usage = now() " - + "WHERE user_id = (SELECT user_id FROM users WHERE api_key = :apiKey) " - + " AND (cycle_credits_remaining + bought_credits_remaining >= :creditAmount)", - nativeQuery = true) - int consumeCredit(@Param("apiKey") String apiKey, @Param("creditAmount") int creditAmount); - - @Modifying - @Query( - value = - "UPDATE user_credits " - + "SET " - + " cycle_credits_remaining = " - + " CASE " - + " WHEN cycle_credits_remaining >= :creditAmount THEN cycle_credits_remaining - :creditAmount " - + " ELSE 0 " - + " END, " - + " bought_credits_remaining = " - + " CASE " - + " WHEN cycle_credits_remaining < :creditAmount " - + " THEN GREATEST(0, bought_credits_remaining - (:creditAmount - cycle_credits_remaining)) " - + " ELSE bought_credits_remaining " - + " END, " - + " total_api_calls_made = total_api_calls_made + 1, " - + " last_api_usage = now() " - + "WHERE user_id = (SELECT u.user_id FROM users u WHERE u.supabase_id = :supabaseId) " - + " AND (cycle_credits_remaining + bought_credits_remaining >= :creditAmount)", - nativeQuery = true) - int consumeCreditBySupabaseId( - @Param("supabaseId") UUID supabaseId, @Param("creditAmount") int creditAmount); - - /** - * Consumes ONLY cycle credits (does not touch bought credits). Used in explicit waterfall - * logic. - */ - @Modifying - @Query( - value = - "UPDATE user_credits " - + "SET " - + " cycle_credits_remaining = cycle_credits_remaining - :amount, " - + " total_api_calls_made = total_api_calls_made + 1, " - + " last_api_usage = now() " - + "WHERE user_id = (SELECT u.user_id FROM users u WHERE u.supabase_id = :supabaseId) " - + " AND cycle_credits_remaining >= :amount", - nativeQuery = true) - int consumeCycleCredits(@Param("supabaseId") UUID supabaseId, @Param("amount") int amount); - - /** Consumes ONLY bought credits (does not touch cycle credits). */ - @Modifying - @Query( - value = - "UPDATE user_credits " - + "SET " - + " bought_credits_remaining = bought_credits_remaining - :amount, " - + " total_api_calls_made = total_api_calls_made + 1, " - + " last_api_usage = now() " - + "WHERE user_id = (SELECT u.user_id FROM users u WHERE u.supabase_id = :supabaseId) " - + " AND bought_credits_remaining >= :amount", - nativeQuery = true) - int consumeBoughtCredits(@Param("supabaseId") UUID supabaseId, @Param("amount") int amount); - - /** Checks if user has sufficient cycle credits (does NOT consume them). */ - @Query( - value = - "SELECT CASE WHEN uc.cycle_credits_remaining >= :amount THEN TRUE ELSE FALSE END " - + "FROM user_credits uc " - + "WHERE uc.user_id = (SELECT u.user_id FROM users u WHERE u.supabase_id = :supabaseId)", - nativeQuery = true) - Boolean hasCycleCredits(@Param("supabaseId") UUID supabaseId, @Param("amount") int amount); - - /** Checks if user has sufficient bought credits (does NOT consume them). */ - @Query( - value = - "SELECT CASE WHEN uc.bought_credits_remaining >= :amount THEN TRUE ELSE FALSE END " - + "FROM user_credits uc " - + "WHERE uc.user_id = (SELECT u.user_id FROM users u WHERE u.supabase_id = :supabaseId)", - nativeQuery = true) - Boolean hasBoughtCredits(@Param("supabaseId") UUID supabaseId, @Param("amount") int amount); -} diff --git a/app/saas/src/main/java/stirling/software/saas/repository/UserErrorTrackerRepository.java b/app/saas/src/main/java/stirling/software/saas/repository/UserErrorTrackerRepository.java deleted file mode 100644 index c67174b85a..0000000000 --- a/app/saas/src/main/java/stirling/software/saas/repository/UserErrorTrackerRepository.java +++ /dev/null @@ -1,41 +0,0 @@ -package stirling.software.saas.repository; - -import java.time.LocalDateTime; -import java.util.List; -import java.util.Optional; - -import org.springframework.data.jpa.repository.JpaRepository; -import org.springframework.data.jpa.repository.Modifying; -import org.springframework.data.jpa.repository.Query; -import org.springframework.data.repository.query.Param; - -import stirling.software.proprietary.security.model.User; -import stirling.software.saas.model.UserErrorTracker; - -public interface UserErrorTrackerRepository extends JpaRepository { - - Optional findByUserAndEndpoint(User user, String endpoint); - - Optional findByUserIdAndEndpoint(Long userId, String endpoint); - - @Query( - "SELECT uet FROM UserErrorTracker uet WHERE uet.user.apiKey = :apiKey AND uet.endpoint = :endpoint") - Optional findByUserApiKeyAndEndpoint( - @Param("apiKey") String apiKey, @Param("endpoint") String endpoint); - - @Query("SELECT uet FROM UserErrorTracker uet WHERE uet.resetAfter <= :currentDateTime") - List findExpiredErrorTrackers( - @Param("currentDateTime") LocalDateTime currentDateTime); - - @Modifying - @Query("DELETE FROM UserErrorTracker uet WHERE uet.resetAfter <= :currentDateTime") - int deleteExpiredErrorTrackers(@Param("currentDateTime") LocalDateTime currentDateTime); - - @Query( - "SELECT uet FROM UserErrorTracker uet WHERE uet.user = :user AND uet.processingErrorCount >= 3") - List findHighErrorCountForUser(@Param("user") User user); - - @Query( - "SELECT COUNT(uet) FROM UserErrorTracker uet WHERE uet.processingErrorCount >= :threshold") - Long countUsersWithHighErrorCount(@Param("threshold") int threshold); -} diff --git a/app/saas/src/main/java/stirling/software/saas/security/EnhancedJwtAuthenticationToken.java b/app/saas/src/main/java/stirling/software/saas/security/EnhancedJwtAuthenticationToken.java index c678453191..6991c2bfae 100644 --- a/app/saas/src/main/java/stirling/software/saas/security/EnhancedJwtAuthenticationToken.java +++ b/app/saas/src/main/java/stirling/software/saas/security/EnhancedJwtAuthenticationToken.java @@ -6,6 +6,8 @@ import org.springframework.security.core.GrantedAuthority; import org.springframework.security.oauth2.jwt.Jwt; import org.springframework.security.oauth2.server.resource.authentication.JwtAuthenticationToken; +import stirling.software.proprietary.security.model.User; + /** * JWT auth token that exposes the Supabase subject UUID and email alongside the standard claims, so * downstream code (audit, credit accounting) can avoid re-parsing the JWT every request. @@ -14,15 +16,35 @@ public class EnhancedJwtAuthenticationToken extends JwtAuthenticationToken { private final String supabaseId; private final String email; + private final User user; public EnhancedJwtAuthenticationToken( Jwt jwt, Collection authorities, String email, String supabaseId) { + this(jwt, authorities, email, supabaseId, null); + } + + public EnhancedJwtAuthenticationToken( + Jwt jwt, + Collection authorities, + String email, + String supabaseId, + User user) { super(jwt, authorities, email); this.email = email; this.supabaseId = supabaseId; + this.user = user; + } + + /** + * Returns the resolved local {@link User} when available so shared {@code principal instanceof + * User} authorization works under JWT auth; falls back to the decoded Jwt. + */ + @Override + public Object getPrincipal() { + return user != null ? user : super.getPrincipal(); } public String getSupabaseId() { diff --git a/app/saas/src/main/java/stirling/software/saas/security/SupabaseAuthenticationFilter.java b/app/saas/src/main/java/stirling/software/saas/security/SupabaseAuthenticationFilter.java index 08cdfcf572..aa678f7cbf 100644 --- a/app/saas/src/main/java/stirling/software/saas/security/SupabaseAuthenticationFilter.java +++ b/app/saas/src/main/java/stirling/software/saas/security/SupabaseAuthenticationFilter.java @@ -63,7 +63,6 @@ public class SupabaseAuthenticationFilter extends OncePerRequestFilter { private final TeamService teamService; private final UserService userService; private final SupabaseUserService supabaseUserService; - private final stirling.software.saas.service.CreditService creditService; private final SaasTeamService saasTeamService; private final JwtDecoder jwtDecoder; private final AuthenticationEntryPoint authenticationEntryPoint = @@ -73,13 +72,11 @@ public class SupabaseAuthenticationFilter extends OncePerRequestFilter { TeamService teamService, UserService userService, SupabaseUserService supabaseUserService, - stirling.software.saas.service.CreditService creditService, SaasTeamService saasTeamService, JwtDecoder jwtDecoder) { this.teamService = teamService; this.userService = userService; this.supabaseUserService = supabaseUserService; - this.creditService = creditService; this.saasTeamService = saasTeamService; this.jwtDecoder = jwtDecoder; } @@ -155,9 +152,15 @@ public class SupabaseAuthenticationFilter extends OncePerRequestFilter { User user = getOrCreateUser(jwt); + // Full accounts carry the resolved User as principal for shared + // instanceof-User authorization; anonymous sessions keep the raw Jwt. EnhancedJwtAuthenticationToken authToken = new EnhancedJwtAuthenticationToken( - jwt, user.getAuthorities(), user.getUsername(), supabaseId); + jwt, + user.getAuthorities(), + user.getUsername(), + supabaseId, + isAnonymous(jwt) ? null : user); SecurityContextHolder.getContext().setAuthentication(authToken); // Hot path: runs on every authenticated request (>10 per page on a typical SPA), @@ -259,7 +262,10 @@ public class SupabaseAuthenticationFilter extends OncePerRequestFilter { user.setUsername(supabaseUser.getEmail()); } try { - return userService.saveUser(user); + User saved = userService.saveUser(user); + // Give the account its own team rather than the shared Default team. + saved.setTeam(saasTeamService.ensurePersonalTeam(saved)); + return saved; } catch (DataIntegrityViolationException e) { log.warn( "Email collision upgrading anonymous user {} to {}: {}", @@ -341,7 +347,8 @@ public class SupabaseAuthenticationFilter extends OncePerRequestFilter { newUser.setEnabled(true); newUser.setFirstLogin(true); newUser.setRoleName(roleId); - newUser.setTeam(teamService.getOrCreateDefaultTeam()); + // No shared Default team; a per-user personal team is assigned after save (team_id + // nullable). newUser.setAuthenticationType(authenticationType); newUser.setSupabaseId(supabaseId); newUser.addAuthority(new Authority(roleId, newUser)); @@ -376,18 +383,7 @@ public class SupabaseAuthenticationFilter extends OncePerRequestFilter { // Only the DB-race winner runs first-time init; the losers skip it. if (weCreatedThisUser) { try { - creditService.getOrCreateUserCredits(savedUser); - } catch (Exception e) { - log.warn( - "Failed to initialize credits for new user {} ({}): {}", - LogRedactionUtils.redactSupabaseId(supabaseId), - LogRedactionUtils.redactEmail(savedUser.getUsername()), - e.getMessage()); - } - - try { - saasTeamService.createPersonalTeam(savedUser); - savedUser = userService.findBySupabaseId(supabaseId).orElse(savedUser); + savedUser.setTeam(saasTeamService.ensurePersonalTeam(savedUser)); } catch (Exception e) { log.warn( "Failed to create personal team for new user {} ({}): {}", diff --git a/app/saas/src/main/java/stirling/software/saas/security/SupabaseSecurityConfig.java b/app/saas/src/main/java/stirling/software/saas/security/SupabaseSecurityConfig.java index feaa83a39c..c0a78d56e0 100644 --- a/app/saas/src/main/java/stirling/software/saas/security/SupabaseSecurityConfig.java +++ b/app/saas/src/main/java/stirling/software/saas/security/SupabaseSecurityConfig.java @@ -23,8 +23,10 @@ import org.springframework.security.config.annotation.web.builders.HttpSecurity; import org.springframework.security.config.annotation.web.configuration.EnableWebSecurity; import org.springframework.security.config.annotation.web.configurers.AbstractHttpConfigurer; import org.springframework.security.config.http.SessionCreationPolicy; +import org.springframework.security.core.Authentication; import org.springframework.security.core.GrantedAuthority; import org.springframework.security.core.authority.SimpleGrantedAuthority; +import org.springframework.security.core.context.SecurityContextHolder; import org.springframework.security.oauth2.core.OAuth2Error; import org.springframework.security.oauth2.core.OAuth2TokenValidator; import org.springframework.security.oauth2.core.OAuth2TokenValidatorResult; @@ -44,9 +46,9 @@ import lombok.extern.slf4j.Slf4j; import stirling.software.common.model.ApplicationProperties; import stirling.software.common.util.RequestUriUtils; +import stirling.software.proprietary.security.model.User; import stirling.software.proprietary.security.service.TeamService; import stirling.software.proprietary.security.service.UserService; -import stirling.software.saas.service.CreditService; import stirling.software.saas.service.SaasTeamService; import stirling.software.saas.service.SupabaseUserService; @@ -63,7 +65,6 @@ public class SupabaseSecurityConfig { private final UserService userService; private final TeamService teamService; private final SupabaseUserService supabaseUserService; - private final CreditService creditService; private final SaasTeamService saasTeamService; private final ApplicationProperties applicationProperties; @@ -118,7 +119,6 @@ public class SupabaseSecurityConfig { teamService, userService, supabaseUserService, - creditService, saasTeamService, jwtDecoder), BearerTokenAuthenticationFilter.class) @@ -257,7 +257,7 @@ public class SupabaseSecurityConfig { applicationProperties.getSystem() != null && applicationProperties.getSystem().getCorsAllowedOrigins() != null && !applicationProperties.getSystem().getCorsAllowedOrigins().isEmpty(); - List origins = + List configuredOrigins = operatorOverride ? applicationProperties.getSystem().getCorsAllowedOrigins() : List.of( @@ -267,6 +267,18 @@ public class SupabaseSecurityConfig { "https://stirling.com", "https://app.stirling.com", "https://api.stirling.com"); + // Always allow the desktop (Tauri) app's webview origins so the bundled + // desktop client can reach the cloud backend regardless of the operator's + // configured web origins. A browser can never present a tauri:// (or + // tauri.localhost) origin, so these are desktop-app identities — safe to + // allow alongside allowCredentials=true. Mirrors core WebMvcConfig. + List origins = new ArrayList<>(configuredOrigins); + for (String desktopOrigin : + List.of("tauri://localhost", "http://tauri.localhost", "https://tauri.localhost")) { + if (!origins.contains(desktopOrigin)) { + origins.add(desktopOrigin); + } + } if (origins.stream().anyMatch(o -> o.contains("*"))) { log.warn( "CORS origins contain a wildcard paired with allowCredentials=true: {}." @@ -285,7 +297,7 @@ public class SupabaseSecurityConfig { "Accept", "Origin", "X-API-KEY")); - cfg.setExposedHeaders(List.of("WWW-Authenticate", "X-Credits-Remaining")); + cfg.setExposedHeaders(List.of("WWW-Authenticate")); cfg.setAllowCredentials(true); cfg.setMaxAge(3600L); UrlBasedCorsConfigurationSource source = new UrlBasedCorsConfigurationSource(); @@ -333,6 +345,16 @@ public class SupabaseSecurityConfig { .map(GrantedAuthority::getAuthority) .collect(Collectors.joining(","))); } - return new EnhancedJwtAuthenticationToken(jwt, authorities, email, supabaseId); + // BearerTokenAuthenticationFilter overwrites the context SupabaseAuthenticationFilter + // built; carry its resolved User across so instanceof-User authorization keeps working. + User user = null; + Authentication existing = SecurityContextHolder.getContext().getAuthentication(); + if (existing instanceof EnhancedJwtAuthenticationToken enhanced + && supabaseId != null + && supabaseId.equals(enhanced.getSupabaseId()) + && enhanced.getPrincipal() instanceof User existingUser) { + user = existingUser; + } + return new EnhancedJwtAuthenticationToken(jwt, authorities, email, supabaseId, user); } } diff --git a/app/saas/src/main/java/stirling/software/saas/security/TeamLeaderPolicyManagementAuthority.java b/app/saas/src/main/java/stirling/software/saas/security/TeamLeaderPolicyManagementAuthority.java new file mode 100644 index 0000000000..e2f5b65b47 --- /dev/null +++ b/app/saas/src/main/java/stirling/software/saas/security/TeamLeaderPolicyManagementAuthority.java @@ -0,0 +1,32 @@ +package stirling.software.saas.security; + +import org.springframework.context.annotation.Profile; +import org.springframework.stereotype.Component; + +import lombok.RequiredArgsConstructor; + +import stirling.software.proprietary.policy.config.PolicyManagementAuthority; + +/** + * SaaS policy context: only the LEADER of the current user's team may edit policies, and every user + * is scoped to their own team. Replaces the self-hosted global-admin check, which is meaningless on + * SaaS (a single global admin exists for the whole deployment, never per-org) — and the admin gets + * no cross-team escape: scoping binds them like everyone else. + */ +@Component +@Profile("saas") +@RequiredArgsConstructor +public class TeamLeaderPolicyManagementAuthority implements PolicyManagementAuthority { + + private final TeamSecurityExpressions teamSecurity; + + @Override + public boolean canEditPolicies() { + return teamSecurity.isCurrentUserTeamLeader(); + } + + @Override + public Long currentUserTeamId() { + return teamSecurity.currentUserTeamId(); + } +} diff --git a/app/saas/src/main/java/stirling/software/saas/security/TeamSecurityExpressions.java b/app/saas/src/main/java/stirling/software/saas/security/TeamSecurityExpressions.java index 0cc91991c3..3dfdf541cc 100644 --- a/app/saas/src/main/java/stirling/software/saas/security/TeamSecurityExpressions.java +++ b/app/saas/src/main/java/stirling/software/saas/security/TeamSecurityExpressions.java @@ -44,6 +44,27 @@ public class TeamSecurityExpressions { .orElse(false); } + /** Whether the current authenticated user is a {@code LEADER} of their own team. */ + public boolean isCurrentUserTeamLeader() { + User currentUser = getCurrentUser(); + if (currentUser == null || currentUser.getTeam() == null) { + return false; + } + return membershipRepository + .findByTeamIdAndUserId(currentUser.getTeam().getId(), currentUser.getId()) + .map(membership -> membership.getRole() == TeamRole.LEADER) + .orElse(false); + } + + /** The current authenticated user's team id, or {@code null} if unauthenticated / teamless. */ + public Long currentUserTeamId() { + User currentUser = getCurrentUser(); + if (currentUser == null || currentUser.getTeam() == null) { + return null; + } + return currentUser.getTeam().getId(); + } + /** Whether the current authenticated user is any kind of member of the given team. */ public boolean isTeamMember(Long teamId) { User currentUser = getCurrentUser(); diff --git a/app/saas/src/main/java/stirling/software/saas/service/CreditBackfillRunner.java b/app/saas/src/main/java/stirling/software/saas/service/CreditBackfillRunner.java deleted file mode 100644 index 008a7e9373..0000000000 --- a/app/saas/src/main/java/stirling/software/saas/service/CreditBackfillRunner.java +++ /dev/null @@ -1,78 +0,0 @@ -package stirling.software.saas.service; - -import java.util.List; - -import org.springframework.boot.ApplicationArguments; -import org.springframework.boot.ApplicationRunner; -import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty; -import org.springframework.context.annotation.Profile; -import org.springframework.stereotype.Component; -import org.springframework.transaction.annotation.Transactional; - -import lombok.RequiredArgsConstructor; -import lombok.extern.slf4j.Slf4j; - -import stirling.software.proprietary.security.database.repository.UserRepository; -import stirling.software.proprietary.security.model.User; - -/** - * ApplicationRunner that backfills user_credits table for existing users who don't have credit rows - * yet. This prevents existing users from being hard-blocked when the credit system is enabled. - * - *

This runs once at application startup after the database schema is ready. - */ -@Component -@Profile("saas") -@ConditionalOnProperty(name = "credits.enabled", havingValue = "true", matchIfMissing = true) -@RequiredArgsConstructor -@Slf4j -public class CreditBackfillRunner implements ApplicationRunner { - - private final UserRepository userRepository; - private final CreditService creditService; - - @Override - @Transactional - public void run(ApplicationArguments args) { - try { - backfillUserCredits(); - } catch (Exception e) { - log.error("Failed to backfill user credits", e); - // Don't throw; this shouldn't prevent app startup - } - } - - private void backfillUserCredits() { - log.info("Starting user credits backfill for existing users..."); - - List usersNeedingCredits = userRepository.findUsersWithApiKeyButNoCredits(); - - if (usersNeedingCredits.isEmpty()) { - log.info( - "No users need credit backfill; all users with API keys already have credit rows"); - return; - } - - log.info("Found {} users with API keys that need credit rows", usersNeedingCredits.size()); - - int backfilled = 0; - for (User user : usersNeedingCredits) { - try { - // Use the existing getOrCreateUserCredits method which handles proper allocation - creditService.getOrCreateUserCredits(user); - backfilled++; - - if (backfilled % 100 == 0) { - log.info("Backfilled credits for {} users so far...", backfilled); - } - } catch (Exception e) { - log.warn( - "Failed to create credits for user {}: {}", - user.getUsername(), - e.getMessage()); - } - } - - log.info("Successfully backfilled user_credits for {} existing users", backfilled); - } -} diff --git a/app/saas/src/main/java/stirling/software/saas/service/CreditResetScheduler.java b/app/saas/src/main/java/stirling/software/saas/service/CreditResetScheduler.java deleted file mode 100644 index 075491bb37..0000000000 --- a/app/saas/src/main/java/stirling/software/saas/service/CreditResetScheduler.java +++ /dev/null @@ -1,107 +0,0 @@ -package stirling.software.saas.service; - -import java.time.LocalDateTime; -import java.time.ZoneId; -import java.time.ZonedDateTime; -import java.time.temporal.TemporalAdjusters; - -import org.springframework.boot.context.event.ApplicationReadyEvent; -import org.springframework.context.annotation.Profile; -import org.springframework.context.event.EventListener; -import org.springframework.scheduling.annotation.Scheduled; -import org.springframework.stereotype.Service; - -import lombok.RequiredArgsConstructor; -import lombok.extern.slf4j.Slf4j; - -import stirling.software.saas.config.CreditsProperties; - -@Service -@Profile("saas") -@Slf4j -@RequiredArgsConstructor -public class CreditResetScheduler { - - private final CreditService creditService; - private final CreditsProperties creditsProperties; - - /** - * Reset cycle credits for all users and teams on the 1st of each month at 2 AM UTC This runs - * monthly, resetting credits based on user roles and team seats - */ - @Scheduled(cron = "${credits.reset.cron:0 0 2 1 * *}", zone = "${credits.reset.zone:UTC}") - public void resetCycleCredits() { - log.info( - "Starting monthly credit reset for all users and teams (schedule: {}, zone: {})", - creditsProperties.getReset().getCron(), - creditsProperties.getReset().getZone()); - - try { - ZoneId configuredZone = ZoneId.of(creditsProperties.getReset().getZone()); - LocalDateTime resetTime = LocalDateTime.now(configuredZone); - creditService.resetCycleCreditsForAllUsers(resetTime); - creditService.resetCycleCreditsForAllTeams(resetTime); - log.info("Monthly credit reset completed successfully at {}", resetTime); - } catch (Exception e) { - log.error("Error during monthly credit reset", e); - } - } - - /** Check for missed resets on application startup */ - @EventListener(ApplicationReadyEvent.class) - public void onApplicationReady() { - try { - ZoneId configuredZone = ZoneId.of(creditsProperties.getReset().getZone()); - LocalDateTime now = LocalDateTime.now(configuredZone); - LocalDateTime lastScheduledReset = getMostRecentScheduledReset(now, configuredZone); - - log.info( - "Checking for missed cycle credit resets. Last scheduled: {}, Current: {}", - lastScheduledReset, - now); - - creditService.resetCycleCreditsForAllUsers(lastScheduledReset); - creditService.resetCycleCreditsForAllTeams(lastScheduledReset); - log.info("Catch-up cycle credit reset completed"); - } catch (Exception e) { - log.error("Error during catch-up credit reset", e); - } - } - - /** Get the most recent scheduled reset time based on configured schedule and zone */ - private LocalDateTime getMostRecentScheduledReset(LocalDateTime now, ZoneId configuredZone) { - ZonedDateTime zonedNow = now.atZone(configuredZone); - - // Find the 1st of the current month at the configured time (default 02:00) - ZonedDateTime firstOfMonth = - zonedNow.with(TemporalAdjusters.firstDayOfMonth()) - .withHour(2) - .withMinute(0) - .withSecond(0) - .withNano(0); - - // If it's the 1st and before the reset hour, or if current time is before the 1st at 2 AM, - // go to previous month's 1st - if (zonedNow.isBefore(firstOfMonth)) { - firstOfMonth = firstOfMonth.minusMonths(1); - } - - return firstOfMonth.toLocalDateTime(); - } - - /** - * Cleanup and maintenance task; runs daily at 3 AM UTC. Performs maintenance tasks like - * cleaning up old data. - */ - @Scheduled(cron = "0 0 3 * * *", zone = "UTC") - public void performDailyMaintenance() { - log.debug("Starting daily credit system maintenance"); - - try { - // API call history cleanup is no longer needed; audit system handles this - log.debug("Daily credit system maintenance completed"); - } catch (Exception e) { - log.error("Error during daily credit system maintenance", e); - } - } -} diff --git a/app/saas/src/main/java/stirling/software/saas/service/CreditService.java b/app/saas/src/main/java/stirling/software/saas/service/CreditService.java deleted file mode 100644 index 775c97d8d4..0000000000 --- a/app/saas/src/main/java/stirling/software/saas/service/CreditService.java +++ /dev/null @@ -1,1179 +0,0 @@ -package stirling.software.saas.service; - -import java.time.LocalDateTime; -import java.time.ZoneId; -import java.time.ZonedDateTime; -import java.util.List; -import java.util.Map; -import java.util.Optional; -import java.util.UUID; - -import org.slf4j.MDC; -import org.springframework.context.annotation.Profile; -import org.springframework.security.authentication.UsernamePasswordAuthenticationToken; -import org.springframework.security.core.Authentication; -import org.springframework.security.core.context.SecurityContextHolder; -import org.springframework.stereotype.Service; -import org.springframework.transaction.annotation.Transactional; -import org.springframework.transaction.support.TransactionSynchronization; -import org.springframework.transaction.support.TransactionSynchronizationManager; - -import io.micrometer.core.instrument.Counter; -import io.micrometer.core.instrument.Gauge; -import io.micrometer.core.instrument.MeterRegistry; - -import lombok.extern.slf4j.Slf4j; - -import stirling.software.proprietary.security.database.repository.UserRepository; -import stirling.software.proprietary.security.model.ApiKeyAuthenticationToken; -import stirling.software.proprietary.security.model.User; -import stirling.software.proprietary.security.service.UserService; -import stirling.software.saas.billing.service.StripeUsageReportingService; -import stirling.software.saas.config.CreditsProperties; -import stirling.software.saas.model.CreditConsumptionResult; -import stirling.software.saas.model.TeamCredit; -import stirling.software.saas.model.UserCredit; -import stirling.software.saas.repository.TeamCreditRepository; -import stirling.software.saas.repository.UserCreditRepository; -import stirling.software.saas.util.LogRedactionUtils; - -@Service -@Profile("saas") -@Slf4j -@Transactional -public class CreditService { - - private final UserCreditRepository userCreditRepository; - private final TeamCreditRepository teamCreditRepository; - private final UserRepository userRepository; - private final UserService userService; - private final CreditsProperties creditsProperties; - private final TeamCreditService teamCreditService; - private final StripeUsageReportingService stripeUsageReportingService; - private final SaasUserExtensionService saasUserExtensionService; - private final SaasTeamExtensionService saasTeamExtensionService; - - // Telemetry metrics - private final Counter creditsConsumedCounter; - private final Counter creditConsumptionFailuresCounter; - private final Counter cycleResetCounter; - private final Counter stripeReportFailuresCounter; - - public CreditService( - UserCreditRepository userCreditRepository, - TeamCreditRepository teamCreditRepository, - UserRepository userRepository, - UserService userService, - CreditsProperties creditsProperties, - TeamCreditService teamCreditService, - StripeUsageReportingService stripeUsageReportingService, - SaasUserExtensionService saasUserExtensionService, - SaasTeamExtensionService saasTeamExtensionService, - MeterRegistry meterRegistry) { - this.userCreditRepository = userCreditRepository; - this.teamCreditRepository = teamCreditRepository; - this.userRepository = userRepository; - this.userService = userService; - this.creditsProperties = creditsProperties; - this.teamCreditService = teamCreditService; - this.stripeUsageReportingService = stripeUsageReportingService; - this.saasUserExtensionService = saasUserExtensionService; - this.saasTeamExtensionService = saasTeamExtensionService; - - // Initialize metrics - this.creditsConsumedCounter = - Counter.builder("credits.consumed") - .description("Number of credits consumed") - .register(meterRegistry); - this.creditConsumptionFailuresCounter = - Counter.builder("credits.consumption.failures") - .description("Number of failed credit consumption attempts") - .register(meterRegistry); - this.cycleResetCounter = - Counter.builder("credits.cycle_reset") - .description("Number of credit cycle resets performed") - .register(meterRegistry); - this.stripeReportFailuresCounter = - Counter.builder("credits.stripe_report.failures") - .description("Stripe meter post failed after the DB debit committed") - .register(meterRegistry); - - // Active gauges for current credit levels - Gauge.builder("credits.total_available", this, CreditService::getTotalAvailableCredits) - .description("Total credits available across all users") - .register(meterRegistry); - Gauge.builder("credits.total_api_calls", this, CreditService::getTotalApiCalls) - .description("Total API calls made across all users") - .register(meterRegistry); - } - - public Optional getUserCreditsByApiKey(String apiKey) { - return userCreditRepository.findByUserApiKey(apiKey); - } - - public Optional getUserCreditsBySupabaseId(String supabaseId) { - try { - UUID supabaseUuid = UUID.fromString(supabaseId); - return userCreditRepository.findBySupabaseId(supabaseUuid); - } catch (IllegalArgumentException e) { - log.warn("Invalid Supabase ID format: {}", supabaseId); - return Optional.empty(); - } - } - - public Optional getUserBySupabaseId(UUID supabaseId) { - return userService.findBySupabaseId(supabaseId); - } - - public Optional getUserCreditsFromAuthentication() { - Authentication authentication = SecurityContextHolder.getContext().getAuthentication(); - if (authentication == null) { - return Optional.empty(); - } - - if (authentication instanceof ApiKeyAuthenticationToken) { - // API Key authentication: credit limits apply - User user = (User) authentication.getPrincipal(); - return getUserCreditsByUserId(user.getId()); - } else if (authentication instanceof UsernamePasswordAuthenticationToken) { - // JWT/Session authentication: unlimited for frontend users - User user = (User) authentication.getPrincipal(); - return getUserCreditsByUserId(user.getId()); - } - - return Optional.empty(); - } - - public Optional getUserCreditsByUserId(Long userId) { - return userCreditRepository.findByUserId(userId); - } - - public UserCredit getOrCreateUserCredits(User user) { - Optional existing = userCreditRepository.findByUser(user); - if (existing.isPresent()) { - UserCredit credits = existing.get(); - // Check if cycle reset is needed based on last scheduled reset - LocalDateTime lastScheduledReset = getMostRecentScheduledReset(); - if (credits.isCycleResetDue(lastScheduledReset)) { - int allocation = getCycleAllocationForUser(user); - credits.resetCycleCredits(allocation, lastScheduledReset); - return userCreditRepository.save(credits); - } - return credits; - } - - // Create new credits for user with proper allocation - UserCredit newCredits = new UserCredit(user); - int allocation = getCycleAllocationForUser(user); - newCredits.resetCycleCredits(allocation, LocalDateTime.now()); - return userCreditRepository.save(newCredits); - } - - private LocalDateTime getMostRecentScheduledReset() { - ZoneId configuredZone = ZoneId.of(creditsProperties.getReset().getZone()); - LocalDateTime now = LocalDateTime.now(configuredZone); - ZonedDateTime zonedNow = now.atZone(configuredZone); - - // Extract hour from cron expression (format: "0 0 2 1 * *" -> hour is 2) - String cronExpression = creditsProperties.getReset().getCron(); - int resetHour = extractHourFromCron(cronExpression); - - // Find first day of current month at the configured time - ZonedDateTime firstOfMonth = - zonedNow.withDayOfMonth(1) - .withHour(resetHour) - .withMinute(0) - .withSecond(0) - .withNano(0); - - // If we're on the 1st but before the reset hour, use previous month's first day - if (zonedNow.getDayOfMonth() == 1 && zonedNow.getHour() < resetHour) { - firstOfMonth = firstOfMonth.minusMonths(1); - } - // If we're before the 1st of this month, use previous month's first day - else if (zonedNow.isBefore(firstOfMonth)) { - firstOfMonth = firstOfMonth.minusMonths(1); - } - - return firstOfMonth.toLocalDateTime(); - } - - private int extractHourFromCron(String cronExpression) { - try { - // Cron format: "second minute hour day month weekday" - String[] parts = cronExpression.split("\\s+"); - if (parts.length >= 3) { - return Integer.parseInt(parts[2]); - } - } catch (NumberFormatException e) { - log.warn( - "Failed to parse hour from cron expression '{}', using default 2", - cronExpression); - } - return 2; // Default to 2 AM - } - - public boolean hasCreditsAvailable(String apiKey) { - Optional credits = getUserCreditsByApiKey(apiKey); - if (credits.isPresent()) { - return credits.get().hasCreditsAvailable(); - } - - // Lazy create UserCredit for existing users who don't have rows yet - Optional userOpt = userRepository.findByApiKey(apiKey); - if (userOpt.isPresent()) { - User user = userOpt.get(); - UserCredit newCredits = getOrCreateUserCredits(user); - return newCredits.hasCreditsAvailable(); - } - - // No user found with this API key - return false; - } - - public boolean consumeCredit(String apiKey, int creditAmount) { - int rowsUpdated = userCreditRepository.consumeCredit(apiKey, creditAmount); - - if (rowsUpdated == 1) { - creditsConsumedCounter.increment(creditAmount); - if (log.isTraceEnabled()) { - log.trace("{} credits consumed for API key: {}", creditAmount, maskApiKey(apiKey)); - } - return true; - } - - creditConsumptionFailuresCounter.increment(); - log.warn( - "Credit consumption failed for API key: {} - insufficient credits (requested: {})", - maskApiKey(apiKey), - creditAmount); - return false; - } - - /** Consume credits for a user identified by Supabase ID; metered overage bills to Stripe. */ - public boolean consumeCreditBySupabaseId(String supabaseId, int creditAmount) { - try { - UUID supabaseUuid = UUID.fromString(supabaseId); - log.debug( - "[CREDIT-CONSUME] Starting credit consumption for Supabase ID: {}, amount: {}", - supabaseId, - creditAmount); - - // Check if user is usage-based - Optional userOpt = userService.findBySupabaseId(supabaseUuid); - - if (userOpt.isEmpty()) { - log.error("[CREDIT-CONSUME] User not found for Supabase ID: {}", supabaseId); - creditConsumptionFailuresCounter.increment(); - return false; - } - - User user = userOpt.get(); - boolean isUsageBased = hasMeteredBillingEnabled(user); - log.info( - "[CREDIT-CONSUME] User {} - Metered billing enabled: {}, Roles: {}", - user.getUsername(), - isUsageBased, - user.getRolesAsString()); - - if (isUsageBased) { - // Metered billing: Try to consume free credits first, then report overage to Stripe - UserCredit userCredits = getOrCreateUserCredits(user); - - log.info( - "[CREDIT-CONSUME] Metered billing user detected: {} - Cycle credits remaining: {}, Amount needed: {}", - user.getUsername(), - userCredits.getCycleCreditsRemaining(), - creditAmount); - - if (userCredits.getCycleCreditsRemaining() >= creditAmount) { - // Covered by free tier; consume normally - int rowsUpdated = - userCreditRepository.consumeCreditBySupabaseId( - supabaseUuid, creditAmount); - - if (rowsUpdated == 1) { - creditsConsumedCounter.increment(creditAmount); - if (log.isTraceEnabled()) { - log.trace( - "[USAGE-BASED] {} credits consumed from free tier for user: {}", - creditAmount, - supabaseId); - } - return true; - } - } else { - // Partial or full overage: consume free credits in this tx, report the overage - // to Stripe after commit (see scheduleStripeReportAfterCommit). - int freeCreditsUsed = - userCredits.getCycleCreditsRemaining() != null - ? userCredits.getCycleCreditsRemaining() - : 0; - int overageCredits = creditAmount - freeCreditsUsed; - - log.warn( - "[CREDIT-CONSUME] OVERAGE DETECTED for user: {} - Free credits available: {}, Credits needed: {}, Overage: {}", - user.getUsername(), - freeCreditsUsed, - creditAmount, - overageCredits); - - // Consume available free credits (if any) - if (freeCreditsUsed > 0) { - log.debug( - "[CREDIT-CONSUME] Consuming {} free credits first", - freeCreditsUsed); - int rowsUpdated = - userCreditRepository.consumeCreditBySupabaseId( - supabaseUuid, freeCreditsUsed); - if (rowsUpdated != 1) { - log.warn( - "[USAGE-BASED] Failed to consume {} free credits for user: {}", - freeCreditsUsed, - supabaseId); - creditConsumptionFailuresCounter.increment(); - return false; - } - } - - String operationId = MDC.get("requestId"); - String idempotencyKey = - stripeUsageReportingService.generateIdempotencyKey( - supabaseId, overageCredits, operationId); - - scheduleStripeReportAfterCommit( - supabaseId, - overageCredits, - idempotencyKey, - creditAmount, - freeCreditsUsed); - return true; - } - - // Lost a concurrent-debit race: the in-memory balance check passed but the atomic - // UPDATE found insufficient credits. Surface the failure so the caller can retry. - log.warn( - "[USAGE-BASED] Concurrent-debit race lost the free-tier consumption for" - + " user {}; caller should retry.", - supabaseId); - creditConsumptionFailuresCounter.increment(); - return false; - } - - // Existing prepaid logic for Pro/Credit-Based users - log.debug( - "[CREDIT-CONSUME] Non-usage-based user; using standard prepaid credit consumption"); - int rowsUpdated = - userCreditRepository.consumeCreditBySupabaseId(supabaseUuid, creditAmount); - - if (rowsUpdated == 1) { - creditsConsumedCounter.increment(creditAmount); - if (log.isTraceEnabled()) { - log.trace("{} credits consumed for Supabase ID: {}", creditAmount, supabaseId); - } - log.debug( - "[CREDIT-CONSUME] Standard credit consumption successful for user: {}", - supabaseId); - return true; - } - - creditConsumptionFailuresCounter.increment(); - log.warn( - "[CREDIT-CONSUME] Credit consumption failed for Supabase ID: {} - insufficient credits (requested: {})", - supabaseId, - creditAmount); - return false; - } catch (IllegalArgumentException e) { - log.error( - "[CREDIT-CONSUME] Invalid Supabase ID format: {} - cannot consume credits", - supabaseId, - e); - creditConsumptionFailuresCounter.increment(); - return false; - } catch (RuntimeException e) { - log.error( - "[CREDIT-CONSUME] Unexpected runtime error consuming credits for user: {} - {}", - supabaseId, - e.getMessage(), - e); - creditConsumptionFailuresCounter.increment(); - return false; - } catch (Exception e) { - log.error( - "[CREDIT-CONSUME] Unexpected error consuming credits for user: {} - {}", - supabaseId, - e.getMessage(), - e); - creditConsumptionFailuresCounter.increment(); - return false; - } - } - - /** - * Checks if a user has metered billing enabled. For users with metered billing, credits are - * billed on usage through Stripe metering rather than being allocated on a monthly cycle basis. - * - * @param user User to check - * @return true if the user has metered billing enabled, false otherwise - */ - private boolean hasMeteredBillingEnabled(User user) { - return saasUserExtensionService.isMeteredBillingEnabled(user); - } - - /** - * Posts the Stripe meter event for an overage debit in a {@code TransactionSynchronization} - * afterCommit hook, so the DB row lock is released before the HTTP call to Stripe. - * - *

If no transaction is active (e.g. a test calling consume directly) the report runs - * synchronously instead, so the meter event still fires. - */ - private void scheduleStripeReportAfterCommit( - String supabaseId, - int overageCredits, - String idempotencyKey, - int creditAmount, - int freeCreditsUsed) { - - Runnable reportToStripe = - () -> { - log.info( - "[CREDIT-CONSUME] Posting Stripe meter event - User: {}, Overage: {}," - + " Idempotency: {}", - supabaseId, - overageCredits, - idempotencyKey); - - boolean reported; - try { - reported = - stripeUsageReportingService.reportUsageToStripe( - supabaseId, overageCredits, idempotencyKey); - } catch (RuntimeException e) { - // Don't let a Stripe exception unwind the afterCommit chain — the DB - // debit has already committed. - log.error( - "[CREDIT-CONSUME] Stripe meter post threw for user {} (overage {});" - + " usage owed-but-unbilled until a retry succeeds", - supabaseId, - overageCredits, - e); - stripeReportFailuresCounter.increment(); - return; - } - - if (reported) { - creditsConsumedCounter.increment(creditAmount); - log.info( - "[USAGE-BASED] User {} consumed {} free + {} overage credits" - + " (total: {}); Stripe meter posted.", - supabaseId, - freeCreditsUsed, - overageCredits, - creditAmount); - } else { - // DB has the debit, Stripe doesn't. The idempotency key is stable, so a - // replay with the same key recovers the meter event without - // double-charging. - stripeReportFailuresCounter.increment(); - log.error( - "[USAGE-BASED] Failed to post Stripe meter event for user {}" - + " (overage {}); usage owed-but-unbilled. Idempotency key" - + " is stable: replay with key '{}' to recover.", - supabaseId, - overageCredits, - idempotencyKey); - } - }; - - if (TransactionSynchronizationManager.isSynchronizationActive()) { - TransactionSynchronizationManager.registerSynchronization( - new TransactionSynchronization() { - @Override - public void afterCommit() { - reportToStripe.run(); - } - }); - } else { - log.warn( - "[CREDIT-CONSUME] No active transaction; reporting Stripe usage synchronously." - + " Expected only in tests."); - reportToStripe.run(); - } - } - - /** Check if a user has credits available by Supabase ID (unified approach). */ - public boolean hasCreditsAvailableBySupabaseId(String supabaseId) { - Optional credits = getUserCreditsBySupabaseId(supabaseId); - if (credits.isPresent()) { - return credits.get().hasCreditsAvailable(); - } - - // Lazy create UserCredit for existing users who don't have rows yet - try { - UUID supabaseUuid = UUID.fromString(supabaseId); - Optional userOpt = userService.findBySupabaseId(supabaseUuid); - if (userOpt.isPresent()) { - User user = userOpt.get(); - UserCredit newCredits = getOrCreateUserCredits(user); - return newCredits.hasCreditsAvailable(); - } - } catch (IllegalArgumentException e) { - log.warn("Invalid Supabase ID format: {}", supabaseId); - } - - // No user found with this Supabase ID - return false; - } - - public boolean isApiKeyAuthenticated() { - Authentication authentication = SecurityContextHolder.getContext().getAuthentication(); - return authentication instanceof ApiKeyAuthenticationToken; - } - - public boolean isJwtAuthenticated() { - Authentication authentication = SecurityContextHolder.getContext().getAuthentication(); - return authentication instanceof UsernamePasswordAuthenticationToken; - } - - public void addBoughtCredits(String username, int credits) { - Optional userOpt = userRepository.findByUsername(username); - if (userOpt.isEmpty()) { - throw new IllegalArgumentException("User not found: " + username); - } - - User user = userOpt.get(); - UserCredit userCredits = getOrCreateUserCredits(user); - userCredits.addBoughtCredits(credits); - userCreditRepository.save(userCredits); - - log.info( - "Added {} bought credits to user: {}. Total available: {}", - credits, - username, - userCredits.getTotalAvailableCredits()); - } - - public void setBoughtCredits(String username, int credits) { - Optional userOpt = userRepository.findByUsername(username); - if (userOpt.isEmpty()) { - throw new IllegalArgumentException("User not found: " + username); - } - User user = userOpt.get(); - UserCredit userCredits = getOrCreateUserCredits(user); - - int previousBought = userCredits.getBoughtCreditsRemaining(); - userCredits.setBoughtCreditsRemaining(credits); - userCredits.setTotalBoughtCredits(credits); // Also update total bought to match - - userCreditRepository.save(userCredits); - log.info( - "Set bought credits for user: {} from {} to {}. Total available: {}", - username, - previousBought, - credits, - userCredits.getTotalAvailableCredits()); - } - - public void setCycleCredits(String username, int credits) { - Optional userOpt = userRepository.findByUsername(username); - if (userOpt.isEmpty()) { - throw new IllegalArgumentException("User not found: " + username); - } - User user = userOpt.get(); - UserCredit userCredits = getOrCreateUserCredits(user); - - int previousCycle = userCredits.getCycleCreditsRemaining(); - userCredits.setCycleCreditsRemaining(credits); - - userCreditRepository.save(userCredits); - log.info( - "Set cycle credits for user: {} from {} to {}. Total available: {}", - username, - previousCycle, - credits, - userCredits.getTotalAvailableCredits()); - } - - public void addBoughtCreditsBySupabaseId(String supabaseId, int credits) { - UserCredit userCredits = getUserCreditsBySupabaseIdWithValidation(supabaseId); - userCredits.addBoughtCredits(credits); - userCreditRepository.save(userCredits); - - log.info( - "Added {} bought credits to user with Supabase ID: {}. Total available: {}", - credits, - supabaseId, - userCredits.getTotalAvailableCredits()); - } - - public void setBoughtCreditsBySupabaseId(String supabaseId, int credits) { - UserCredit userCredits = getUserCreditsBySupabaseIdWithValidation(supabaseId); - - int previousBought = userCredits.getBoughtCreditsRemaining(); - userCredits.setBoughtCreditsRemaining(credits); - userCredits.setTotalBoughtCredits(credits); // Also update total bought to match - - userCreditRepository.save(userCredits); - log.info( - "Set bought credits for user with Supabase ID: {} from {} to {}. Total available: {}", - supabaseId, - previousBought, - credits, - userCredits.getTotalAvailableCredits()); - } - - public void setCycleCreditsBySupabaseId(String supabaseId, int credits) { - UserCredit userCredits = getUserCreditsBySupabaseIdWithValidation(supabaseId); - - int previousCycle = userCredits.getCycleCreditsRemaining(); - userCredits.setCycleCreditsRemaining(credits); - - userCreditRepository.save(userCredits); - log.info( - "Set cycle credits for user with Supabase ID: {} from {} to {}. Total available: {}", - supabaseId, - previousCycle, - credits, - userCredits.getTotalAvailableCredits()); - } - - public void resetCycleCreditsForAllUsers(LocalDateTime lastScheduledReset) { - List creditsNeedingReset = - userCreditRepository.findCreditsNeedingCycleReset(lastScheduledReset); - - for (UserCredit credit : creditsNeedingReset) { - int allocation = getCycleAllocationForUser(credit.getUser()); - credit.resetCycleCredits(allocation, lastScheduledReset); - userCreditRepository.save(credit); - cycleResetCounter.increment(); - } - - log.info( - "Reset cycle credits for {} users based on scheduled reset time: {}", - creditsNeedingReset.size(), - lastScheduledReset); - } - - // Backward compatibility method - public void resetCycleCreditsForAllUsers() { - LocalDateTime now = LocalDateTime.now(); - resetCycleCreditsForAllUsers(now); - } - - public void resetCycleCreditsForAllTeams(LocalDateTime lastScheduledReset) { - List creditsNeedingReset = - teamCreditRepository.findCreditsNeedingCycleReset(lastScheduledReset); - - int proAllocation = - creditsProperties.getCycle().getAllocations().getOrDefault("ROLE_PRO_USER", 500); - int totalCycleAllocation = proAllocation; - - for (TeamCredit credit : creditsNeedingReset) { - credit.resetCycleCredits(totalCycleAllocation, lastScheduledReset); - teamCreditRepository.save(credit); - cycleResetCounter.increment(); - - log.info( - "Reset cycle credits for team {} to {} (fixed PRO amount)", - credit.getTeam().getId(), - totalCycleAllocation); - } - - log.info( - "Reset cycle credits for {} teams based on scheduled reset time: {}", - creditsNeedingReset.size(), - lastScheduledReset); - } - - // Backward compatibility method - public void resetCycleCreditsForAllTeams() { - LocalDateTime now = LocalDateTime.now(); - resetCycleCreditsForAllTeams(now); - } - - /** Credit summary keyed by API key; for API-key-only users without a linked Supabase ID. */ - public CreditSummary getCreditSummaryByApiKey(String apiKey) { - Optional creditsOpt = getUserCreditsByApiKey(apiKey); - if (creditsOpt.isEmpty()) { - return new CreditSummary(); - } - UserCredit credits = creditsOpt.get(); - boolean isUnlimited = credits.getCycleCreditsAllocated() == Integer.MAX_VALUE; - return new CreditSummary( - credits.getCycleCreditsRemaining(), - credits.getCycleCreditsAllocated(), - credits.getBoughtCreditsRemaining(), - credits.getTotalBoughtCredits(), - credits.getTotalAvailableCredits(), - credits.getLastCycleResetAt(), - credits.getLastApiUsage(), - isUnlimited); - } - - public CreditSummary getCreditSummary(String username) { - // Note: This method is kept for admin functions that need to lookup users by username. - Optional userOpt = userRepository.findByUsername(username); - if (userOpt.isEmpty()) { - log.warn("No user found with username: {}", username); - return new CreditSummary(); - } - - User user = userOpt.get(); - UserCredit credits = getOrCreateUserCredits(user); - boolean isUnlimited = credits.getCycleCreditsAllocated() == Integer.MAX_VALUE; - return new CreditSummary( - credits.getCycleCreditsRemaining(), - credits.getCycleCreditsAllocated(), - credits.getBoughtCreditsRemaining(), - credits.getTotalBoughtCredits(), - credits.getTotalAvailableCredits(), - credits.getLastCycleResetAt(), - credits.getLastApiUsage(), - isUnlimited); - } - - /** - * Credit summary for a user. Non-personal team members use the shared team pool; personal-team - * or teamless users use individual credits. - */ - public CreditSummary getCreditSummaryBySupabaseId(String supabaseId) { - // First, look up the user to check for team membership - UUID supabaseUuid; - try { - supabaseUuid = UUID.fromString(supabaseId); - } catch (IllegalArgumentException e) { - log.warn("Invalid Supabase ID format: {}", supabaseId); - return new CreditSummary(); - } - - Optional userOpt = userService.findBySupabaseId(supabaseUuid); - if (userOpt.isEmpty()) { - log.warn("No user found with Supabase ID: {}", supabaseId); - return new CreditSummary(); - } - - User user = userOpt.get(); - - // Check if user has LIMITED_API_USER role (anonymous/guest users). - // Limited API users always use personal credits, never team credits. - boolean isLimitedApiUser = - user.getAuthorities().stream() - .anyMatch( - authority -> - "ROLE_LIMITED_API_USER".equals(authority.getAuthority()) - || "ROLE_EXTRA_LIMITED_API_USER" - .equals(authority.getAuthority())); - - if (isLimitedApiUser) { - log.debug("User {} is limited API user; using personal credits", user.getUsername()); - } - // Check if user is on a non-personal team; if so, return team credits. - // Skip this check for limited API users. - else if (user.getTeam() != null && !saasTeamExtensionService.isPersonal(user.getTeam())) { - Long teamId = user.getTeam().getId(); - log.debug( - "User {} is on team {} - returning team credits instead of personal credits", - user.getUsername(), - teamId); - - Optional teamCreditsOpt = teamCreditService.getTeamCredits(teamId); - if (teamCreditsOpt.isPresent()) { - TeamCredit tc = teamCreditsOpt.get(); - return new CreditSummary( - tc.getCycleCreditsRemaining() != null ? tc.getCycleCreditsRemaining() : 0, - tc.getCycleCreditsAllocated() != null ? tc.getCycleCreditsAllocated() : 0, - tc.getBoughtCreditsRemaining() != null ? tc.getBoughtCreditsRemaining() : 0, - tc.getTotalBoughtCredits() != null ? tc.getTotalBoughtCredits() : 0, - tc.getTotalAvailableCredits(), - tc.getLastCycleResetAt(), - tc.getLastApiUsage(), - false // teams never have unlimited credits - ); - } else { - log.warn("Team {} exists but has no credit record; returning empty", teamId); - return new CreditSummary(); - } - } - - // User is limited API user, not on a team, or on a personal team; return personal credits - log.debug("User {} using personal credits", user.getUsername()); - Optional creditsOpt = getUserCreditsBySupabaseId(supabaseId); - if (creditsOpt.isEmpty()) { - // Lazy initialization: try to create UserCredit for existing user - log.info( - "UserCredit missing for Supabase ID {}, creating now", - LogRedactionUtils.redactSupabaseId(supabaseId)); - UserCredit newCredits = initializeCreditsForUser(user); - boolean isUnlimited = newCredits.getCycleCreditsAllocated() == Integer.MAX_VALUE; - return new CreditSummary( - newCredits.getCycleCreditsRemaining(), - newCredits.getCycleCreditsAllocated(), - newCredits.getBoughtCreditsRemaining(), - newCredits.getTotalBoughtCredits(), - newCredits.getTotalAvailableCredits(), - newCredits.getLastCycleResetAt(), - newCredits.getLastApiUsage(), - isUnlimited); - } - - UserCredit credits = creditsOpt.get(); - boolean isUnlimited = credits.getCycleCreditsAllocated() == Integer.MAX_VALUE; - return new CreditSummary( - credits.getCycleCreditsRemaining(), - credits.getCycleCreditsAllocated(), - credits.getBoughtCreditsRemaining(), - credits.getTotalBoughtCredits(), - credits.getTotalAvailableCredits(), - credits.getLastCycleResetAt(), - credits.getLastApiUsage(), - isUnlimited); - } - - /** Helper method to lookup user by Supabase ID with proper error handling */ - private UserCredit getUserCreditsBySupabaseIdWithValidation(String supabaseId) { - try { - UUID supabaseUuid = UUID.fromString(supabaseId); - Optional userOpt = userService.findBySupabaseId(supabaseUuid); - if (userOpt.isEmpty()) { - throw new IllegalArgumentException( - "User not found with Supabase ID: " + supabaseId); - } - - User user = userOpt.get(); - return getOrCreateUserCredits(user); - } catch (IllegalArgumentException e) { - if (e.getMessage().startsWith("Invalid UUID")) { - throw new IllegalArgumentException("Invalid Supabase ID format: " + supabaseId); - } - throw e; - } - } - - private String maskApiKey(String apiKey) { - if (apiKey == null || apiKey.length() < 8) { - return "***"; - } - return apiKey.substring(0, 4) + "***" + apiKey.substring(apiKey.length() - 4); - } - - /** Get total available credits across all users (for metrics gauge) */ - private Double getTotalAvailableCredits() { - try { - Long total = userCreditRepository.getTotalAvailableCreditsAcrossAllUsers(); - return total != null ? total.doubleValue() : 0.0; - } catch (Exception e) { - log.error("Error calculating total available credits for metrics", e); - return 0.0; - } - } - - /** Get total API calls across all users (for metrics gauge) */ - private Double getTotalApiCalls() { - try { - Long total = userCreditRepository.getTotalApiCallsAcrossAllUsers(); - return total != null ? total.doubleValue() : 0.0; - } catch (Exception e) { - log.error("Error calculating total API calls for metrics", e); - return 0.0; - } - } - - /** Get cycle credit allocation for a user based on configuration */ - private int getCycleAllocationForUser(User user) { - if (user == null || user.getRolesAsString() == null) { - log.warn("User or roles is null, returning 0 credits"); - return 0; - } - - String rolesString = user.getRolesAsString(); - Map allocations = creditsProperties.getCycle().getAllocations(); - - log.debug( - "Getting credit allocation for user {} with roles: {}", - user.getUsername(), - rolesString); - log.debug("Available credit allocations: {}", allocations); - - // Check roles in priority order - if (rolesString.contains("ROLE_ADMIN") && creditsProperties.getCycle().isAdminUnlimited()) { - log.debug("User {} has admin unlimited credits", user.getUsername()); - return Integer.MAX_VALUE; - } - - // Internal API users (including test API key) get unlimited credits - if (rolesString.contains("ROLE_INTERNAL_API_USER")) { - log.debug("User {} has internal API unlimited credits", user.getUsername()); - return Integer.MAX_VALUE; - } - - for (Map.Entry entry : allocations.entrySet()) { - if (rolesString.contains(entry.getKey())) { - log.debug( - "User {} matched role {} with {} credits", - user.getUsername(), - entry.getKey(), - entry.getValue()); - return entry.getValue(); - } - } - - // Default allocation - int defaultCredits = allocations.getOrDefault("ROLE_USER", 50); - log.debug( - "User {} using default ROLE_USER allocation: {} credits", - user.getUsername(), - defaultCredits); - return defaultCredits; - } - - /** Initialize credits for a new user */ - public UserCredit initializeCreditsForUser(User user) { - log.info( - "Initializing credits for user: {} (id: {})", - LogRedactionUtils.redactEmail(user.getUsername()), - user.getId()); - UserCredit credits = new UserCredit(user); - int allocation = getCycleAllocationForUser(user); - log.info( - "Allocated {} credits to user: {}", - allocation, - LogRedactionUtils.redactEmail(user.getUsername())); - credits.resetCycleCredits(allocation, LocalDateTime.now()); - UserCredit saved = userCreditRepository.save(credits); - log.info( - "Successfully saved UserCredit for user: {} with allocation: {}", - LogRedactionUtils.redactEmail(user.getUsername()), - allocation); - return saved; - } - - /** - * Refresh cycle credits after a role change. Resets {@code cycleCreditsRemaining} to the new - * allocation; preserves {@code boughtCreditsRemaining}. - */ - public void refreshCreditsAfterRoleChange(User user) { - log.info( - "Refreshing credits for user: {} after role change", - LogRedactionUtils.redactEmail(user.getUsername())); - - Optional creditsOpt = userCreditRepository.findByUserId(user.getId()); - if (creditsOpt.isEmpty()) { - log.warn( - "No credits found for user {}, initializing", - LogRedactionUtils.redactEmail(user.getUsername())); - initializeCreditsForUser(user); - return; - } - - UserCredit credits = creditsOpt.get(); - int oldAllocation = credits.getCycleCreditsAllocated(); - int oldRemaining = credits.getCycleCreditsRemaining(); - int newAllocation = getCycleAllocationForUser(user); - - log.info( - "Updating credits for user {} from {}/{} to {}/{} cycle credits", - user.getUsername(), - oldRemaining, - oldAllocation, - newAllocation, - newAllocation); - - // Full reset: sets both allocation and remaining to the new amount. - // This gives full credits on upgrade, but removes excess on downgrade. - credits.resetCycleCredits(newAllocation, LocalDateTime.now()); - userCreditRepository.save(credits); - - log.info( - "Successfully refreshed credits for user {}: {} cycle credits available", - user.getUsername(), - newAllocation); - } - - /** - * Resets cycle credit allocation after a role change. More efficient version when the caller - * already knows the target allocation. Performs a FULL RESET: updates allocation, remaining, - * and timestamp. - * - *

Different from setCycleCredits() which only adjusts remaining credits. - * - * @param userId The user ID - * @param newAllocation The new cycle credit allocation amount - */ - public void resetCycleAllocationForRoleChange(Long userId, int newAllocation) { - log.info("Resetting cycle allocation for user ID {} to {}", userId, newAllocation); - - Optional creditsOpt = userCreditRepository.findByUserId(userId); - if (creditsOpt.isEmpty()) { - log.warn("No credits found for user ID {}, cannot reset allocation", userId); - throw new IllegalStateException("User credits not found for user ID: " + userId); - } - - UserCredit credits = creditsOpt.get(); - int oldAllocation = credits.getCycleCreditsAllocated(); - int oldRemaining = credits.getCycleCreditsRemaining(); - - log.info( - "Resetting allocation for user ID {} from {}/{} to {}/{} cycle credits", - userId, - oldRemaining, - oldAllocation, - newAllocation, - newAllocation); - - // Full reset: sets both allocation and remaining to the new amount - credits.resetCycleCredits(newAllocation, LocalDateTime.now()); - userCreditRepository.save(credits); - - log.info( - "Successfully reset cycle allocation for user ID {}: {} credits available", - userId, - newAllocation); - } - - /** - * Consumes credits with explicit waterfall logic for Pro billing model. Implements the - * following priority order: - * - *

    - *
  1. Cycle Credits: Try consuming from cycle credit allocation (100/month for Pro, - * 25/month for Free) - *
  2. Bought Credits: Try consuming from purchased credits - *
  3. Metered Billing: If user.has_metered_billing_enabled, report to Stripe and allow - *
  4. Reject: No available credit source (Pro users without metered billing get - * helpful message about enabling overage billing) - *
- * - * @param user User consuming credits - * @param creditAmount Amount to consume - * @param isApiRequest True if API key request, false if UI request (both consume credits for - * Pro users) - * @return CreditConsumptionResult with source and success status - */ - public CreditConsumptionResult consumeCreditWithWaterfall( - User user, int creditAmount, boolean isApiRequest) { - - log.debug( - "[WATERFALL] Starting credit consumption for user {} - Amount: {}, API request: {}", - user.getUsername(), - creditAmount, - isApiRequest); - - // Internal API users (e.g., CUSTOM_API_USER) get unlimited credits, no Supabase ID needed - if (user.getRolesAsString().contains("STIRLING-PDF-BACKEND-API-USER")) { - log.debug( - "[WATERFALL] Internal API user {} - unlimited credits, skipping consumption", - user.getUsername()); - creditsConsumedCounter.increment(creditAmount); - return CreditConsumptionResult.success("INTERNAL_API_UNLIMITED"); - } - - UUID supabaseId = user.getSupabaseId(); - if (supabaseId == null) { - log.error("[WATERFALL] User {} has no Supabase ID", user.getUsername()); - return CreditConsumptionResult.failure("User has no Supabase ID"); - } - - // STEP 2: Try cycle free credits - Boolean hasCycle = userCreditRepository.hasCycleCredits(supabaseId, creditAmount); - if (Boolean.TRUE.equals(hasCycle)) { - int rowsUpdated = userCreditRepository.consumeCycleCredits(supabaseId, creditAmount); - if (rowsUpdated == 1) { - creditsConsumedCounter.increment(creditAmount); - log.info( - "[WATERFALL] Consumed {} cycle credits for user: {}", - creditAmount, - user.getUsername()); - return CreditConsumptionResult.success("CYCLE_CREDITS"); - } - } - - // STEP 3: Try purchased credits - Boolean hasBought = userCreditRepository.hasBoughtCredits(supabaseId, creditAmount); - if (Boolean.TRUE.equals(hasBought)) { - int rowsUpdated = userCreditRepository.consumeBoughtCredits(supabaseId, creditAmount); - if (rowsUpdated == 1) { - creditsConsumedCounter.increment(creditAmount); - log.info( - "[WATERFALL] Consumed {} bought credits for user: {}", - creditAmount, - user.getUsername()); - return CreditConsumptionResult.success("BOUGHT_CREDITS"); - } - } - - // STEP 4: Try metered billing (check flag, not role) - if (saasUserExtensionService.isMeteredBillingEnabled(user)) { - log.info( - "[WATERFALL] User {} has metered billing enabled; scheduling {} credits for" - + " Stripe report (after commit)", - user.getUsername(), - creditAmount); - - String operationId = MDC.get("requestId"); - String idempotencyKey = - stripeUsageReportingService.generateIdempotencyKey( - supabaseId.toString(), creditAmount, operationId); - - scheduleStripeReportAfterCommit( - supabaseId.toString(), - creditAmount, - idempotencyKey, - creditAmount, - /* freeCreditsUsed= */ 0); - return CreditConsumptionResult.success("METERED_SUBSCRIPTION"); - } else if (user.getRolesAsString().contains("ROLE_PRO_USER")) { - // Pro user without metered billing enabled; reject with helpful message - log.warn( - "[WATERFALL] Pro user {} has exhausted credits but metered billing not enabled. Rejecting request.", - user.getUsername()); - log.info( - "[WATERFALL] User should set up overage billing via UI to enable uninterrupted service."); - - creditConsumptionFailuresCounter.increment(); - return CreditConsumptionResult.failure( - "Credits exhausted. Please enable overage billing in settings for uninterrupted service."); - } - - // STEP 5: Reject; no available credit source - log.warn( - "[WATERFALL] No credit source available for user: {} (needed: {} credits)", - user.getUsername(), - creditAmount); - creditConsumptionFailuresCounter.increment(); - return CreditConsumptionResult.failure("INSUFFICIENT_CREDITS"); - } - - public static class CreditSummary { - public final int cycleCreditsRemaining; - public final int cycleCreditsAllocated; - public final int boughtCreditsRemaining; - public final int totalBoughtCredits; - public final int totalAvailableCredits; - public final LocalDateTime cycleResetDate; - public final LocalDateTime lastApiUsage; - public final boolean unlimited; - - public CreditSummary() { - this(0, 0, 0, 0, 0, null, null, false); - } - - public CreditSummary( - int cycleCreditsRemaining, - int cycleCreditsAllocated, - int boughtCreditsRemaining, - int totalBoughtCredits, - int totalAvailableCredits, - LocalDateTime cycleResetDate, - LocalDateTime lastApiUsage, - boolean unlimited) { - this.cycleCreditsRemaining = cycleCreditsRemaining; - this.cycleCreditsAllocated = cycleCreditsAllocated; - this.boughtCreditsRemaining = boughtCreditsRemaining; - this.totalBoughtCredits = totalBoughtCredits; - this.totalAvailableCredits = totalAvailableCredits; - this.cycleResetDate = cycleResetDate; - this.lastApiUsage = lastApiUsage; - this.unlimited = unlimited; - } - } -} diff --git a/app/saas/src/main/java/stirling/software/saas/service/ErrorTrackingService.java b/app/saas/src/main/java/stirling/software/saas/service/ErrorTrackingService.java deleted file mode 100644 index 91ffa23186..0000000000 --- a/app/saas/src/main/java/stirling/software/saas/service/ErrorTrackingService.java +++ /dev/null @@ -1,315 +0,0 @@ -package stirling.software.saas.service; - -import java.time.LocalDateTime; -import java.util.Optional; -import java.util.concurrent.TimeUnit; - -import org.springframework.context.annotation.Profile; -import org.springframework.scheduling.annotation.Scheduled; -import org.springframework.stereotype.Service; -import org.springframework.transaction.annotation.Transactional; - -import com.github.benmanes.caffeine.cache.Cache; -import com.github.benmanes.caffeine.cache.Caffeine; - -import lombok.extern.slf4j.Slf4j; - -import stirling.software.proprietary.security.database.repository.UserRepository; -import stirling.software.proprietary.security.model.User; -import stirling.software.saas.config.CreditsProperties; -import stirling.software.saas.model.ProcessingErrorType; -import stirling.software.saas.model.UserErrorTracker; -import stirling.software.saas.repository.UserErrorTrackerRepository; - -@Service -@Profile("saas") -@Slf4j -@Transactional -public class ErrorTrackingService { - - private final UserErrorTrackerRepository errorTrackerRepository; - private final UserRepository userRepository; - private final CreditsProperties creditsProperties; - - /** - * Local cache for error counts to reduce database chatter. - * - *

This cache is used to temporarily store error counts for each API key and endpoint, - * reducing the frequency of database writes and lookups. - * - *

Nullability: This field may be {@code null} if local caching is disabled via {@link - * CreditsProperties#getCache()#isLocalEnabled()}. All usages must check for null before - * accessing or invoking methods on this cache. - * - *

Lifecycle: The cache is initialized in the constructor based on configuration and - * remains unchanged for the lifetime of this service instance. - * - *

Thread-safety: The underlying Caffeine cache is thread-safe. - */ - private final Cache errorCountCache; - - public ErrorTrackingService( - UserErrorTrackerRepository errorTrackerRepository, - UserRepository userRepository, - CreditsProperties creditsProperties) { - this.errorTrackerRepository = errorTrackerRepository; - this.userRepository = userRepository; - this.creditsProperties = creditsProperties; - - // Initialize cache based on configuration - this.errorCountCache = - creditsProperties.getCache().isLocalEnabled() - ? Caffeine.newBuilder() - .maximumSize(10000) - .expireAfterWrite( - creditsProperties.getErrors().getTtlMinutes(), - TimeUnit.MINUTES) - .build() - : null; - } - - /** - * Record an error and determine if credits should be consumed - * - * @param apiKey User's API key - * @param endpoint The endpoint that failed - * @param throwable The exception that occurred - * @param httpStatus HTTP response status - * @return true if credits should be consumed for this error - */ - public boolean recordErrorAndShouldConsumeCredit( - String apiKey, String endpoint, Throwable throwable, int httpStatus) { - ProcessingErrorType errorType = - ProcessingErrorType.classifyError(throwable, httpStatus, endpoint); - - // Never charge for validation errors or system errors - if (errorType != ProcessingErrorType.PROCESSING_ERROR) { - log.debug( - "Error classified as {}, no credit consumption for API key: {}, endpoint: {}", - errorType, - maskApiKey(apiKey), - endpoint); - return false; - } - - String cacheKey = apiKey + "|" + endpoint; - - if (errorCountCache != null) { - // Use cache for fast tracking - ErrorCountCache cachedCount = errorCountCache.get(cacheKey, k -> new ErrorCountCache()); - cachedCount.incrementErrorCount(); - - boolean shouldCharge = - cachedCount.getErrorCount() - > creditsProperties.getErrors().getFreeProcessingErrors(); - - // Persist to DB when crossing the charging threshold or on first error - if (shouldCharge - && cachedCount.getErrorCount() - == creditsProperties.getErrors().getFreeProcessingErrors() + 1) { - persistErrorToDatabase(apiKey, endpoint); - } - - log.info( - "Processing error recorded (cached) for API key: {}, endpoint: {}, error count: {}, will charge: {}", - maskApiKey(apiKey), - endpoint, - cachedCount.getErrorCount(), - shouldCharge); - - return shouldCharge; - } else { - // Fallback to direct DB tracking - return recordErrorDirectToDatabase(apiKey, endpoint); - } - } - - private boolean recordErrorDirectToDatabase(String apiKey, String endpoint) { - Optional userOpt = userRepository.findByApiKey(apiKey); - if (userOpt.isEmpty()) { - log.warn("User not found for API key: {}", maskApiKey(apiKey)); - return false; - } - - User user = userOpt.get(); - UserErrorTracker tracker = getOrCreateErrorTracker(user, endpoint); - - tracker.recordProcessingError(creditsProperties.getErrors().getTtlMinutes()); - errorTrackerRepository.save(tracker); - - boolean shouldCharge = - tracker.shouldChargeForProcessingError( - creditsProperties.getErrors().getFreeProcessingErrors()); - - log.info( - "Processing error recorded (DB) for user: {}, endpoint: {}, error count: {}, will charge: {}", - user.getUsername(), - endpoint, - tracker.getProcessingErrorCount(), - shouldCharge); - - return shouldCharge; - } - - private void persistErrorToDatabase(String apiKey, String endpoint) { - try { - Optional userOpt = userRepository.findByApiKey(apiKey); - if (userOpt.isPresent()) { - User user = userOpt.get(); - UserErrorTracker tracker = getOrCreateErrorTracker(user, endpoint); - // Set to threshold + 1 to indicate charging has started - tracker.setProcessingErrorCount( - creditsProperties.getErrors().getFreeProcessingErrors() + 1); - tracker.setLastProcessingError(LocalDateTime.now()); - tracker.setResetAfter( - LocalDateTime.now() - .plusMinutes(creditsProperties.getErrors().getTtlMinutes())); - errorTrackerRepository.save(tracker); - log.debug( - "Persisted error threshold crossing to DB for API key: {}, endpoint: {}", - maskApiKey(apiKey), - endpoint); - } - } catch (Exception e) { - log.error( - "Failed to persist error to database for API key: {}, endpoint: {}", - maskApiKey(apiKey), - endpoint, - e); - } - } - - /** Check if a user has high error counts that might indicate abuse */ - public boolean hasHighErrorCount(String apiKey, String endpoint) { - Optional trackerOpt = - errorTrackerRepository.findByUserApiKeyAndEndpoint(apiKey, endpoint); - return trackerOpt - .map( - t -> - t.shouldChargeForProcessingError( - creditsProperties.getErrors().getFreeProcessingErrors())) - .orElse(false); - } - - /** Get error information for a user and endpoint */ - public ErrorInfo getErrorInfo(String apiKey, String endpoint) { - String cacheKey = apiKey + "|" + endpoint; - - if (errorCountCache != null) { - // Check cache first - ErrorCountCache cachedCount = errorCountCache.getIfPresent(cacheKey); - if (cachedCount != null) { - int currentCount = cachedCount.getErrorCount(); - int freeErrors = creditsProperties.getErrors().getFreeProcessingErrors(); - return new ErrorInfo( - currentCount, - Math.max(0, freeErrors - currentCount), - currentCount > freeErrors, - cachedCount.getLastErrorTime()); - } - } - - // Fallback to DB - Optional trackerOpt = - errorTrackerRepository.findByUserApiKeyAndEndpoint(apiKey, endpoint); - if (trackerOpt.isEmpty()) { - return new ErrorInfo( - 0, creditsProperties.getErrors().getFreeProcessingErrors(), false, null); - } - - UserErrorTracker tracker = trackerOpt.get(); - - // Reset if expired - if (tracker.isExpired()) { - tracker.resetErrorCount(creditsProperties.getErrors().getTtlMinutes()); - errorTrackerRepository.save(tracker); - return new ErrorInfo( - 0, creditsProperties.getErrors().getFreeProcessingErrors(), false, null); - } - - return new ErrorInfo( - tracker.getProcessingErrorCount(), - tracker.getErrorsUntilCharged( - creditsProperties.getErrors().getFreeProcessingErrors()), - tracker.shouldChargeForProcessingError( - creditsProperties.getErrors().getFreeProcessingErrors()), - tracker.getLastProcessingError()); - } - - private UserErrorTracker getOrCreateErrorTracker(User user, String endpoint) { - Optional existing = - errorTrackerRepository.findByUserAndEndpoint(user, endpoint); - - if (existing.isPresent()) { - UserErrorTracker tracker = existing.get(); - - // Reset if expired - if (tracker.isExpired()) { - tracker.resetErrorCount(creditsProperties.getErrors().getTtlMinutes()); - } - - return tracker; - } - - // Create new tracker - return new UserErrorTracker(user, endpoint, creditsProperties.getErrors().getTtlMinutes()); - } - - /** Clean up expired error trackers every hour */ - @Scheduled(cron = "0 0 * * * *") - public void cleanupExpiredErrorTrackers() { - try { - int deleted = errorTrackerRepository.deleteExpiredErrorTrackers(LocalDateTime.now()); - if (deleted > 0) { - log.debug("Cleaned up {} expired error trackers", deleted); - } - } catch (Exception e) { - log.error("Error cleaning up expired error trackers", e); - } - } - - private String maskApiKey(String apiKey) { - if (apiKey == null || apiKey.length() < 8) { - return "***"; - } - return apiKey.substring(0, 4) + "***" + apiKey.substring(apiKey.length() - 4); - } - - /** Information about user's error status for an endpoint */ - public static class ErrorInfo { - public final int currentErrorCount; - public final int errorsUntilCharged; - public final boolean isChargingForErrors; - public final LocalDateTime lastError; - - public ErrorInfo( - int currentErrorCount, - int errorsUntilCharged, - boolean isChargingForErrors, - LocalDateTime lastError) { - this.currentErrorCount = currentErrorCount; - this.errorsUntilCharged = errorsUntilCharged; - this.isChargingForErrors = isChargingForErrors; - this.lastError = lastError; - } - } - - /** Cache entry for tracking error counts in memory */ - private static class ErrorCountCache { - private int errorCount = 0; - private LocalDateTime lastErrorTime = LocalDateTime.now(); - - public void incrementErrorCount() { - errorCount++; - lastErrorTime = LocalDateTime.now(); - } - - public int getErrorCount() { - return errorCount; - } - - public LocalDateTime getLastErrorTime() { - return lastErrorTime; - } - } -} diff --git a/app/saas/src/main/java/stirling/software/saas/service/SaasTeamService.java b/app/saas/src/main/java/stirling/software/saas/service/SaasTeamService.java index 02bf87e270..e54aa00df2 100644 --- a/app/saas/src/main/java/stirling/software/saas/service/SaasTeamService.java +++ b/app/saas/src/main/java/stirling/software/saas/service/SaasTeamService.java @@ -2,7 +2,6 @@ package stirling.software.saas.service; import java.time.LocalDateTime; import java.util.List; -import java.util.Optional; import java.util.UUID; import org.springframework.context.annotation.Profile; @@ -21,16 +20,12 @@ import stirling.software.proprietary.security.database.repository.UserRepository import stirling.software.proprietary.security.model.User; import stirling.software.proprietary.security.repository.TeamRepository; import stirling.software.saas.billing.repository.BillingSubscriptionRepository; -import stirling.software.saas.config.CreditsProperties; import stirling.software.saas.config.SupabaseConfigurationProperties; -import stirling.software.saas.model.TeamCredit; import stirling.software.saas.model.TeamInvitation; import stirling.software.saas.model.TeamMembership; import stirling.software.saas.repository.SaasTeamExtensionsRepository; -import stirling.software.saas.repository.TeamCreditRepository; import stirling.software.saas.repository.TeamInvitationRepository; import stirling.software.saas.repository.TeamMembershipRepository; -import stirling.software.saas.repository.UserCreditRepository; /** SaaS-only team management: invitations, personal teams, seat caps, paid-subscription gating. */ @Service @@ -43,11 +38,7 @@ public class SaasTeamService { private final TeamMembershipRepository membershipRepository; private final TeamInvitationRepository invitationRepository; private final UserRepository userRepository; - private final UserCreditRepository userCreditRepository; private final BillingSubscriptionRepository billingSubscriptionRepository; - private final TeamCreditService teamCreditService; - private final TeamCreditRepository teamCreditRepository; - private final CreditsProperties creditsProperties; private final RestTemplate restTemplate; private final RateLimitService rateLimitService; private final SupabaseConfigurationProperties supabaseConfig; @@ -59,6 +50,16 @@ public class SaasTeamService { public static final String DEFAULT_TEAM_NAME = "Default"; public static final String INTERNAL_TEAM_NAME = "Internal"; + /** Returns the user's personal team, creating one if they have none. Idempotent. */ + @Transactional + public Team ensurePersonalTeam(User user) { + Team existing = user.getTeam(); + if (existing != null && saasTeamExtensionService.isPersonal(existing)) { + return existing; + } + return createPersonalTeam(user); + } + /** * Create personal team for new user during signup or migrate existing user from Default team * @@ -100,9 +101,6 @@ public class SaasTeamService { user.setTeam(savedTeam); userRepository.save(user); - // Initialize team credits - teamCreditService.initializeTeamCredits(savedTeam, user); - // Clean up old Default/Internal team membership if (oldTeam != null && (DEFAULT_TEAM_NAME.equals(oldTeam.getName()) @@ -303,6 +301,11 @@ public class SaasTeamService { throw new IllegalStateException("Team has no available seats"); } + // Validate: accepting won't orphan a team the user leads or that has a paid plan. + // Accepting moves the user off their current team; leaveTeam already blocks the + // last leader of a team from walking away, so accept must enforce the same rule. + assertCanLeaveCurrentTeamsToJoinAnother(acceptingUser); + // User can only be in one team . leave existing teams before joining new one List existingMemberships = membershipRepository.findByUserId(acceptingUser.getId()); @@ -442,6 +445,46 @@ public class SaasTeamService { teamId); } + /** + * Guard against silently orphaning a team when a user accepts an invite to another one. + * + *

{@link #acceptInvitation} moves a user to the inviting team by first leaving their current + * team(s). Personal teams are disposable (they get deleted on accept), but a non-personal team + * must not be left memberless while still billing. {@link #leaveTeam} already refuses to let + * the last leader walk away; accept took a shortcut around that check, which let a paid team's + * leader join another team and orphan their old team together with its live subscription. + * + *

So: for each non-personal team the user leads as its last leader, block the + * accept. The message points them at the right remedy — cancel the plan if the team is paid, + * otherwise transfer leadership first. + * + * @param user the user attempting to accept an invitation + * @throws IllegalStateException if accepting would orphan a team the user leads + */ + private void assertCanLeaveCurrentTeamsToJoinAnother(User user) { + for (TeamMembership membership : membershipRepository.findByUserId(user.getId())) { + Team team = membership.getTeam(); + if (saasTeamExtensionService.isPersonal(team) || !membership.isLeader()) { + // Personal teams are deleted on accept; non-leaders leaving never orphans a team. + continue; + } + // Only reached for a non-personal team the user leads — at most one such team in the + // one-team-per-user model — so this count runs ~once, not per membership. + if (membershipRepository.countByTeamIdAndRole(team.getId(), TeamRole.LEADER) > 1) { + // Another leader remains, so the team keeps an owner. + continue; + } + if (hasActivePaidSubscription(team)) { + throw new IllegalStateException( + "Your team has an active plan and you are its last leader. Cancel the plan" + + " or transfer leadership before joining another team."); + } + throw new IllegalStateException( + "You are the last leader of your team. Transfer leadership before joining" + + " another team."); + } + } + /** * Leave team (self-removal) * @@ -738,64 +781,6 @@ public class SaasTeamService { teamRepository.save(team); - Optional creditOpt = teamCreditRepository.findByTeamId(teamId); - - int fixedAllocation = - creditsProperties.getCycle().getAllocations().getOrDefault("ROLE_PRO_USER", 500); - - if (creditOpt.isPresent()) { - TeamCredit credit = creditOpt.get(); - - int oldAllocation = - credit.getCycleCreditsAllocated() != null - ? credit.getCycleCreditsAllocated() - : 0; - - if (oldAllocation != fixedAllocation) { - int currentRemaining = - credit.getCycleCreditsRemaining() != null - ? credit.getCycleCreditsRemaining() - : 0; - int allocationDifference = fixedAllocation - oldAllocation; - - credit.setCycleCreditsAllocated(fixedAllocation); - - int newRemaining = Math.max(0, currentRemaining + allocationDifference); - credit.setCycleCreditsRemaining(newRemaining); - - teamCreditRepository.save(credit); - - log.info( - "Updated team {} credit allocation: {} -> {} (fixed PRO amount). Remaining: {} -> {}", - teamId, - oldAllocation, - fixedAllocation, - currentRemaining, - newRemaining); - } else { - log.debug( - "Team {} already has fixed allocation of {} credits, no update needed", - teamId, - fixedAllocation); - } - } else { - log.warn("Team {} missing credit record; creating with fixed allocation", teamId); - TeamCredit credit = new TeamCredit(team); - - credit.setCycleCreditsAllocated(fixedAllocation); - credit.setCycleCreditsRemaining(fixedAllocation); - credit.setBoughtCreditsRemaining(0); - credit.setTotalBoughtCredits(0); - credit.setTotalApiCallsMade(0L); - credit.setLastCycleResetAt(LocalDateTime.now()); - teamCreditRepository.save(credit); - - log.info( - "Created team_credits record for team {} with {} fixed credits (unlimited seats model)", - teamId, - fixedAllocation); - } - log.info( "Team {} seat allocation updated: maxSeats={}, seatsUsed={}, isPersonal={}", teamId, diff --git a/app/saas/src/main/java/stirling/software/saas/service/SaasUserAccountService.java b/app/saas/src/main/java/stirling/software/saas/service/SaasUserAccountService.java index c4b0ad769d..51ad9cab17 100644 --- a/app/saas/src/main/java/stirling/software/saas/service/SaasUserAccountService.java +++ b/app/saas/src/main/java/stirling/software/saas/service/SaasUserAccountService.java @@ -31,6 +31,7 @@ public class SaasUserAccountService { private final SupabaseUserService supabaseUserService; private final SaasUserExtensionService saasUserExtensionService; private final SaasTeamExtensionService saasTeamExtensionService; + private final SaasTeamService saasTeamService; /** * Resolve a local {@link User} from a Supabase UUID string. Throws if the ID format is invalid @@ -173,6 +174,8 @@ public class SaasUserAccountService { user.setUsername(email); } user = userService.saveUser(user); + // Give the upgraded user their own team rather than the shared Default team. + user.setTeam(saasTeamService.ensurePersonalTeam(user)); log.info( "Upgraded anonymous user {} to {} ({})", user.getId(), diff --git a/app/saas/src/main/java/stirling/software/saas/service/TeamCreditService.java b/app/saas/src/main/java/stirling/software/saas/service/TeamCreditService.java deleted file mode 100644 index 5d1ca00484..0000000000 --- a/app/saas/src/main/java/stirling/software/saas/service/TeamCreditService.java +++ /dev/null @@ -1,279 +0,0 @@ -package stirling.software.saas.service; - -import java.time.LocalDateTime; -import java.util.List; -import java.util.Optional; - -import org.springframework.context.annotation.Profile; -import org.springframework.stereotype.Service; -import org.springframework.transaction.annotation.Transactional; - -import lombok.RequiredArgsConstructor; -import lombok.extern.slf4j.Slf4j; - -import stirling.software.common.model.enumeration.TeamRole; -import stirling.software.proprietary.model.Team; -import stirling.software.proprietary.security.model.User; -import stirling.software.saas.billing.service.StripeUsageReportingService; -import stirling.software.saas.config.CreditsProperties; -import stirling.software.saas.model.CreditConsumptionResult; -import stirling.software.saas.model.TeamCredit; -import stirling.software.saas.model.TeamMembership; -import stirling.software.saas.repository.TeamCreditRepository; -import stirling.software.saas.repository.TeamMembershipRepository; - -/** - * Service for managing team credit pools. Handles credit initialization, consumption, and cycle - * resets for teams. - */ -@Service -@Profile("saas") -@RequiredArgsConstructor -@Slf4j -public class TeamCreditService { - - private final TeamCreditRepository teamCreditRepository; - private final TeamMembershipRepository membershipRepository; - private final CreditsProperties creditsProperties; - private final StripeUsageReportingService stripeUsageReportingService; - private final SaasUserExtensionService saasUserExtensionService; - - /** Initialise a fixed PRO credit allocation for a new team. */ - @Transactional - public TeamCredit initializeTeamCredits(Team team, User primaryUser) { - Optional existing = teamCreditRepository.findByTeamId(team.getId()); - if (existing.isPresent()) { - log.debug("Team credits already exist for team {}", team.getId()); - return existing.get(); - } - - TeamCredit credits = new TeamCredit(team); - - // Fixed PRO allocation; seat-independent. - int proAllocation = - creditsProperties.getCycle().getAllocations().getOrDefault("ROLE_PRO_USER", 500); - int totalCycleAllocation = proAllocation; - - credits.setCycleCreditsAllocated(totalCycleAllocation); - credits.setCycleCreditsRemaining(totalCycleAllocation); - credits.setLastCycleResetAt(LocalDateTime.now()); - - TeamCredit saved = teamCreditRepository.save(credits); - log.info( - "Initialized team credits for team {} with {} cycle credits (fixed PRO amount)", - team.getId(), - totalCycleAllocation); - return saved; - } - - /** - * Check if team has credits available - * - * @param teamId the team ID - * @return true if team has credits available - */ - public boolean hasCreditsAvailable(Long teamId) { - return teamCreditRepository - .findByTeamId(teamId) - .map(TeamCredit::hasCreditsAvailable) - .orElse(false); - } - - /** - * Atomically consume credits from team pool - * - * @param teamId the team ID - * @param amount number of credits to consume - * @return true if credits were consumed, false if insufficient credits or version conflict - */ - @Transactional - public boolean consumeCredit(Long teamId, int amount) { - int rowsUpdated = teamCreditRepository.consumeCredit(teamId, amount); - if (rowsUpdated == 0) { - log.warn( - "Failed to consume {} credits for team {} (insufficient credits or version conflict)", - amount, - teamId); - return false; - } - log.debug("Consumed {} credits for team {}", amount, teamId); - return true; - } - - /** - * Get team credit summary for a user's team. - * - * @param user the user - * @return Optional of TeamCredit for the user's team - */ - public Optional getCreditSummaryForUser(User user) { - if (user.getTeam() == null) { - log.warn("User {} has no team assigned", user.getId()); - return Optional.empty(); - } - - Long teamId = user.getTeam().getId(); - log.debug("Using user's team {} for credit summary", teamId); - return teamCreditRepository.findByTeamId(teamId); - } - - /** - * Get team credits by team ID - * - * @param teamId the team ID - * @return Optional of TeamCredit - */ - public Optional getTeamCredits(Long teamId) { - return teamCreditRepository.findByTeamId(teamId); - } - - /** - * Add bought credits to team pool - * - * @param teamId the team ID - * @param credits number of credits to add - */ - @Transactional - public void addBoughtCredits(Long teamId, int credits) { - TeamCredit teamCredit = - teamCreditRepository - .findByTeamId(teamId) - .orElseThrow(() -> new IllegalArgumentException("Team credits not found")); - - teamCredit.addBoughtCredits(credits); - teamCreditRepository.save(teamCredit); - log.info("Added {} bought credits to team {}", credits, teamId); - } - - /** - * Reset cycle credits for team - * - * @param teamId the team ID - * @param cycleAllocation new cycle allocation - * @param resetTime reset timestamp - */ - @Transactional - public void resetCycleCredits(Long teamId, int cycleAllocation, LocalDateTime resetTime) { - TeamCredit teamCredit = - teamCreditRepository - .findByTeamId(teamId) - .orElseThrow(() -> new IllegalArgumentException("Team credits not found")); - - teamCredit.resetCycleCredits(cycleAllocation, resetTime); - teamCreditRepository.save(teamCredit); - log.info("Reset cycle credits for team {} to {}", teamId, cycleAllocation); - } - - /** - * Consume from the team credit pool; falls through to the team leader's metered Stripe billing - * when the pool is exhausted. - */ - @Transactional - public CreditConsumptionResult consumeCreditWithWaterfall(Long teamId, int amount) { - log.debug("[TEAM-CREDIT] Starting consumption for team {} - amount: {}", teamId, amount); - - // Step 1: Try consuming from team credit pool - int rowsUpdated = teamCreditRepository.consumeCredit(teamId, amount); - if (rowsUpdated == 1) { - log.info("[TEAM-CREDIT] Consumed {} credits from team {} pool", amount, teamId); - return CreditConsumptionResult.success("TEAM_CREDITS"); - } - - log.warn("[TEAM-CREDIT] Team {} credit pool exhausted; checking leader overage", teamId); - - // Step 2: Get team leader - Optional leaderOpt = getTeamLeader(teamId); - if (leaderOpt.isEmpty()) { - log.error("[TEAM-CREDIT] Team {} has no leader; cannot use overage billing", teamId); - return CreditConsumptionResult.failure("NO_TEAM_LEADER"); - } - - User teamLeader = leaderOpt.get(); - - // Step 3: Check if team leader has metered billing enabled - if (!saasUserExtensionService.isMeteredBillingEnabled(teamLeader)) { - log.warn( - "[TEAM-CREDIT] Team {} leader {} does not have metered billing enabled", - teamId, - teamLeader.getUsername()); - return CreditConsumptionResult.failure( - "TEAM_CREDITS_EXHAUSTED_NO_OVERAGE", - "Team credits exhausted. Team leader must enable overage billing for" - + " uninterrupted service."); - } - - // Step 4: Report overage to Stripe via team leader's metered billing - String leaderSupabaseId = - teamLeader.getSupabaseId() != null ? teamLeader.getSupabaseId().toString() : null; - - if (leaderSupabaseId == null) { - log.error("[TEAM-CREDIT] Team leader {} has no Supabase ID", teamLeader.getUsername()); - return CreditConsumptionResult.failure("LEADER_NO_SUPABASE_ID"); - } - - try { - String operationId = org.slf4j.MDC.get("requestId"); - if (operationId == null || operationId.isBlank()) { - operationId = java.util.UUID.randomUUID().toString(); - } - String idempotencyKey = - stripeUsageReportingService.generateIdempotencyKey( - leaderSupabaseId, amount, operationId); - - log.info( - "[TEAM-CREDIT] Reporting {} overage credits to Stripe for team {} leader {}", - amount, - teamId, - teamLeader.getUsername()); - - boolean reported = - stripeUsageReportingService.reportUsageToStripe( - leaderSupabaseId, amount, idempotencyKey); - - if (reported) { - log.info( - "[TEAM-CREDIT] Successfully reported {} overage credits for team {} via" - + " leader {}", - amount, - teamId, - teamLeader.getUsername()); - return CreditConsumptionResult.success("TEAM_LEADER_METERED"); - } else { - log.error("[TEAM-CREDIT] Failed to report overage to Stripe for team {}", teamId); - return CreditConsumptionResult.failure( - "STRIPE_REPORTING_FAILED", - "Unable to report usage to Stripe. Please try again."); - } - } catch (Exception e) { - log.error( - "[TEAM-CREDIT] Exception reporting overage for team {}: {}", - teamId, - e.getMessage(), - e); - return CreditConsumptionResult.failure( - "STRIPE_REPORTING_ERROR", "Error reporting usage: " + e.getMessage()); - } - } - - /** Returns the team's LEADER (first one if multiple exist) for overage-billing routing. */ - private Optional getTeamLeader(Long teamId) { - List leaders = - membershipRepository.findByTeamIdAndRole(teamId, TeamRole.LEADER); - - if (leaders.isEmpty()) { - log.warn("Team {} has no leaders", teamId); - return Optional.empty(); - } - - // Return first leader (typically only one leader per team) - TeamMembership leader = leaders.get(0); - User leaderUser = leader.getUser(); - log.debug( - "Found team {} leader: {} (user ID: {})", - teamId, - leaderUser.getUsername(), - leaderUser.getId()); - - return Optional.of(leaderUser); - } -} diff --git a/app/saas/src/main/java/stirling/software/saas/service/UserRoleService.java b/app/saas/src/main/java/stirling/software/saas/service/UserRoleService.java index 7a6452cdb9..e54758221c 100644 --- a/app/saas/src/main/java/stirling/software/saas/service/UserRoleService.java +++ b/app/saas/src/main/java/stirling/software/saas/service/UserRoleService.java @@ -12,10 +12,9 @@ import stirling.software.proprietary.security.database.repository.AuthorityRepos import stirling.software.proprietary.security.database.repository.UserRepository; import stirling.software.proprietary.security.model.Authority; import stirling.software.proprietary.security.model.User; -import stirling.software.saas.config.CreditsProperties; import stirling.software.saas.util.LogRedactionUtils; -/** Changes user roles and refreshes their credit allocation. */ +/** Changes user roles (and the matching authority grant/revoke). */ @Service @Profile("saas") @RequiredArgsConstructor @@ -24,8 +23,6 @@ public class UserRoleService { private final UserRepository userRepository; private final AuthorityRepository authorityRepository; - private final CreditService creditService; - private final CreditsProperties creditsProperties; /** * Change a user's role @@ -58,7 +55,7 @@ public class UserRoleService { /** * Downgrade a user to FREE tier (ROLE_USER) * - *

Changes role from PRO_USER to USER and resets cycle credit allocation to FREE tier. + *

Revokes ROLE_PRO_USER by changing the role/authority from PRO_USER to USER. * * @param user the user to downgrade */ @@ -70,24 +67,15 @@ public class UserRoleService { changeRole(user, Role.USER.getRoleId()); - // Reset credits to FREE tier allocation - int freeAllocation = - creditsProperties - .getCycle() - .getAllocations() - .getOrDefault(Role.USER.getRoleId(), 25); - creditService.resetCycleAllocationForRoleChange(user.getId(), freeAllocation); - log.info( - "Successfully downgraded user {} to FREE with {} cycle credits", - LogRedactionUtils.redactEmail(user.getUsername()), - freeAllocation); + "Successfully downgraded user {} to FREE", + LogRedactionUtils.redactEmail(user.getUsername())); } /** * Upgrade a user to PRO tier (ROLE_PRO_USER) * - *

Changes role from USER to PRO_USER and resets cycle credit allocation to PRO tier. + *

Grants ROLE_PRO_USER by changing the role/authority from USER to PRO_USER. * * @param user the user to upgrade */ @@ -98,30 +86,8 @@ public class UserRoleService { changeRole(user, Role.PRO_USER.getRoleId()); - // Reset credits to PRO tier allocation - int proAllocation = - creditsProperties - .getCycle() - .getAllocations() - .getOrDefault(Role.PRO_USER.getRoleId(), 100); - creditService.resetCycleAllocationForRoleChange(user.getId(), proAllocation); - log.info( - "Successfully upgraded user {} to PRO with {} cycle credits", - LogRedactionUtils.redactEmail(user.getUsername()), - proAllocation); - } - - /** - * Get credit allocation for a specific role - * - * @param roleId the role ID (e.g., "ROLE_USER", "ROLE_PRO_USER") - * @return the cycle credit allocation for that role - */ - public int getCreditAllocationForRole(String roleId) { - return creditsProperties - .getCycle() - .getAllocations() - .getOrDefault(roleId, Role.USER.getRoleId().equals(roleId) ? 25 : 100); + "Successfully upgraded user {} to PRO", + LogRedactionUtils.redactEmail(user.getUsername())); } } diff --git a/app/saas/src/main/java/stirling/software/saas/util/CreditHeaderUtils.java b/app/saas/src/main/java/stirling/software/saas/util/CreditHeaderUtils.java deleted file mode 100644 index 9e432b0682..0000000000 --- a/app/saas/src/main/java/stirling/software/saas/util/CreditHeaderUtils.java +++ /dev/null @@ -1,97 +0,0 @@ -package stirling.software.saas.util; - -import java.util.Optional; - -import org.springframework.context.annotation.Profile; -import org.springframework.stereotype.Component; - -import lombok.RequiredArgsConstructor; -import lombok.extern.slf4j.Slf4j; - -import stirling.software.proprietary.security.model.User; -import stirling.software.saas.model.TeamCredit; -import stirling.software.saas.model.UserCredit; -import stirling.software.saas.service.CreditService; -import stirling.software.saas.service.SaasTeamExtensionService; -import stirling.software.saas.service.TeamCreditService; - -/** - * Resolves the user's remaining credit balance. Uses the team pool for non-personal team members, - * otherwise the user's individual credits (looked up by Supabase ID or API key). - */ -@Component -@Profile("saas") -@RequiredArgsConstructor -@Slf4j -public class CreditHeaderUtils { - - private final SaasTeamExtensionService saasTeamExtensionService; - - /** - * Get the remaining credits for a user, checking team credits first (non-personal teams only). - * - * @param user The user whose credits to check - * @param creditService The credit service to fetch user credits - * @param teamCreditService The team credit service to fetch team credits - * @return The remaining credit balance, or -1 if credits cannot be determined - */ - public int getRemainingCredits( - User user, CreditService creditService, TeamCreditService teamCreditService) { - try { - // Limited-API users always read personal credits. - boolean isLimitedApiUser = - user.getAuthorities().stream() - .anyMatch( - authority -> - "ROLE_LIMITED_API_USER".equals(authority.getAuthority()) - || "ROLE_EXTRA_LIMITED_API_USER" - .equals(authority.getAuthority())); - - Long targetTeamId = null; - if (!isLimitedApiUser - && user.getTeam() != null - && !saasTeamExtensionService.isPersonal(user.getTeam())) { - targetTeamId = user.getTeam().getId(); - } - - if (targetTeamId != null) { - return teamCreditService - .getTeamCredits(targetTeamId) - .map(TeamCredit::getTotalAvailableCredits) - .orElse(-1); - } else { - log.debug( - "[CREDIT-HEADER] Getting personal credits - SupabaseId: {}, ApiKey: {}, Username: {}", - user.getSupabaseId(), - user.getApiKey() != null ? "present" : "null", - user.getUsername()); - - Optional credits; - if (user.getSupabaseId() != null) { - credits = - creditService.getUserCreditsBySupabaseId( - user.getSupabaseId().toString()); - log.debug( - "[CREDIT-HEADER] Looked up by SupabaseId - Found: {}", - credits.isPresent()); - } else if (user.getApiKey() != null) { - credits = creditService.getUserCreditsByApiKey(user.getApiKey()); - log.debug( - "[CREDIT-HEADER] Looked up by ApiKey - Found: {}", credits.isPresent()); - } else { - log.warn( - "[CREDIT-HEADER] No SupabaseId or ApiKey for user: {}", - user.getUsername()); - return -1; - } - - int remaining = credits.map(UserCredit::getTotalAvailableCredits).orElse(-1); - log.debug("[CREDIT-HEADER] Returning credits: {}", remaining); - return remaining; - } - } catch (Exception e) { - log.warn("[CREDIT-HEADER] Could not get remaining credits: {}", e.getMessage(), e); - return -1; - } - } -} diff --git a/app/saas/src/main/resources/application-dev.properties b/app/saas/src/main/resources/application-dev.properties index 725fe79fee..ee8bf80ff2 100644 --- a/app/saas/src/main/resources/application-dev.properties +++ b/app/saas/src/main/resources/application-dev.properties @@ -23,3 +23,10 @@ spring.datasource.hikari.data-source-properties.ApplicationName=stirling-consoli logging.level.stirling.software.saas=DEBUG logging.level.org.springframework.security.oauth2.jwt=WARN logging.level.org.springframework.security.oauth2.server.resource=WARN + +# Supabase meter edge fn the Java backend calls (server-to-server, on job close). +# URL is not a secret; auth rides the existing SUPABASE_EDGE_FUNCTION_SECRET (same +# shared secret the team-invitation flow uses — no service-role key in the Java env). +# Blank secret → the meter service no-ops with a WARN, so the app still boots. +# The billing portal is NOT here — the FE calls create-customer-portal-session directly. +payg.meter.endpoint=https://qacaivhsjtftfwtgjvva.supabase.co/functions/v1/meter-payg-units diff --git a/app/saas/src/main/resources/application-saas.properties b/app/saas/src/main/resources/application-saas.properties index 3676ad7d82..63ba016950 100644 --- a/app/saas/src/main/resources/application-saas.properties +++ b/app/saas/src/main/resources/application-saas.properties @@ -35,7 +35,37 @@ app.supabase.clock-skew-seconds=${app.jwt.clock-skew-seconds:120} app.supabase.edge-function-url=https://${app.supabase.project-ref}.supabase.co/functions/v1 app.supabase.edge-function-secret=${SUPABASE_EDGE_FUNCTION_SECRET:} +# ---------- PAYG meter reporting ---------- +# Posts billable usage to the Supabase `meter-payg-units` edge function in the JobChargeService +# close() afterCommit hook. Defaults to empty so unit tests / local dev are no-ops; set +# PAYG_METER_ENDPOINT in deployed envs to enable. Compose with the Supabase functions base via +# env (e.g. PAYG_METER_ENDPOINT=$SUPABASE_FUNCTIONS_URL/meter-payg-units) rather than templating +# here — concatenating an empty default would produce a half-valid URL. +# Auth rides the existing backend<->edge-fn shared secret (SUPABASE_EDGE_FUNCTION_SECRET, same +# one team-invitation-email uses) — the backend never holds the RLS-bypassing service-role key. +payg.meter.endpoint=${PAYG_METER_ENDPOINT:} +payg.meter.auth-token=${app.supabase.edge-function-secret:} + +# Reconcile job: retries meter events logged to payg_meter_event_log but never confirmed posted +# (Stripe blip, edge-fn outage, pod crash mid-POST). Re-sends under the same idempotency key so +# Stripe dedups; only within Stripe's 24h window. Defaults are sensible; tune/disable via env. +payg.meter.reconcile.enabled=${PAYG_METER_RECONCILE_ENABLED:true} +payg.meter.reconcile.retry-delay=${PAYG_METER_RECONCILE_RETRY_DELAY:PT5M} +payg.meter.reconcile.batch-size=${PAYG_METER_RECONCILE_BATCH_SIZE:100} +payg.meter.reconcile-cron=${PAYG_METER_RECONCILE_CRON:0 */15 * * * *} + +# ---------- PAYG customer portal ---------- +# The Stripe billing-portal session is minted FE-direct: the Plan page calls the Supabase +# `create-customer-portal-session` edge function with the user's JWT (same pattern as checkout), +# and the function's SECURITY DEFINER RPC enforces team membership. No Java proxy, no backend +# config — return_url host allowlisting lives in the edge fn (PORTAL_ALLOWED_RETURN_HOSTS). + supabase.url=https://${app.supabase.project-ref}.supabase.co spring.security.oauth2.resourceserver.jwt.jwk-set-uri=https://${app.supabase.project-ref}.supabase.co/auth/v1/.well-known/jwks.json spring.security.oauth2.resourceserver.jwt.audiences=${app.supabase.expected-aud} + +# ---------- Multi-tenant scoping ---------- +# Restrict the signing user picker to the caller's team; SaaS must be 'team' +# or unrelated tenants leak emails to each other. +storage.signing.userListScope=team diff --git a/app/saas/src/main/resources/db/migration/saas/V13__payg_shadow_charge_status.sql b/app/saas/src/main/resources/db/migration/saas/V13__payg_shadow_charge_status.sql new file mode 100644 index 0000000000..66da92a108 --- /dev/null +++ b/app/saas/src/main/resources/db/migration/saas/V13__payg_shadow_charge_status.sql @@ -0,0 +1,20 @@ +-- Refund tracking on shadow rows. Shadow models the eventual Stripe meter_event_adjustment(cancel) +-- by flipping status from CHARGED to REFUNDED in the same request's afterCompletion when a +-- freshly-opened process fails with 5xx on its first step. +-- +-- Reconciliation report selects SUM(payg_units) WHERE status = 'CHARGED' to get the true net +-- Stripe would bill. + +ALTER TABLE payg_shadow_charge + ADD COLUMN IF NOT EXISTS status VARCHAR(16) NOT NULL DEFAULT 'CHARGED', + ADD COLUMN IF NOT EXISTS refunded_at TIMESTAMP, + ADD COLUMN IF NOT EXISTS refund_reason VARCHAR(128); + +CREATE INDEX IF NOT EXISTS idx_payg_shadow_status_time + ON payg_shadow_charge (status, occurred_at); + +-- Hot-path index for findFirstByJobIdOrderByIdAsc: hit on every 5xx-first-step refund to flip +-- the row to REFUNDED. UNIQUE because at most one shadow row exists per processing_job by +-- construction (openProcess writes exactly one on OPENED, zero on JOINED). +CREATE UNIQUE INDEX IF NOT EXISTS uq_payg_shadow_job_id + ON payg_shadow_charge (job_id); diff --git a/app/saas/src/main/resources/db/migration/saas/V14__payg_subscription_state.sql b/app/saas/src/main/resources/db/migration/saas/V14__payg_subscription_state.sql new file mode 100644 index 0000000000..eca2665de7 --- /dev/null +++ b/app/saas/src/main/resources/db/migration/saas/V14__payg_subscription_state.sql @@ -0,0 +1,189 @@ +-- PAYG subscription state — the column + functions that let new customers reach Stripe billing. +-- +-- This migration is half of the Stripe/Supabase wire-up (PR-SB-1 in `notes/PAYG_DESIGN.md` +-- revision note + `payg-stripe-supabase-plan.html`). It's strictly additive: +-- * one new column on payg_team_extensions (payg_subscription_id) +-- * one new column on pricing_policy (free_tier_units_per_cycle) +-- * two RPC functions (payg_link_subscription, payg_unlink_subscription) — the only writers +-- of subscription state, called by stripe-webhook + create-payg-team-subscription edge fns +-- * an AFTER-INSERT trigger on teams that auto-creates the payg_team_extensions sidecar row +-- so every new signup is PAYG-by-default +-- * an RLS policy that lets team LEADERs (and the service role) link subscriptions +-- +-- No behaviour change for the running app until PR-SB-4 wires PaygMeterReportingService and +-- the free-tier gate into JobChargeService. Until then this just exposes new state for the +-- edge functions in PR-SB-2 to write through to. +-- +-- Design references: +-- * notes/PAYG_DESIGN.md (revision note 2026-06-03 — "subscription presence is the gate") +-- * payg-stripe-supabase-plan.html §3.1 — RPC functions; §3.5 — RLS policy + +-- --------------------------------------------------------------------------------------------- +-- 1. New columns +-- --------------------------------------------------------------------------------------------- + +ALTER TABLE stirling_pdf.payg_team_extensions + ADD COLUMN IF NOT EXISTS payg_subscription_id VARCHAR(128) UNIQUE; + +COMMENT ON COLUMN stirling_pdf.payg_team_extensions.payg_subscription_id IS + 'Stripe subscription id (sub_xxx) for this team''s PAYG metered subscription. ' + 'NULL = team has not added a card yet; engine writes shadow rows only. ' + 'NOT NULL = engine posts meter events to Stripe on every billable tool call. ' + 'Mutated exclusively by payg_link_subscription / payg_unlink_subscription RPC functions.'; + +ALTER TABLE stirling_pdf.pricing_policy + ADD COLUMN IF NOT EXISTS free_tier_units_per_cycle BIGINT NOT NULL DEFAULT 0; + +COMMENT ON COLUMN stirling_pdf.pricing_policy.free_tier_units_per_cycle IS + 'Doc units a team on this policy can consume per cycle before they must add a card. ' + 'Default 0 = no free tier (block immediately). The seeded default policy will set this ' + 'to the launch free-tier size; the special "launch" policy used by the day-1 legacy ' + 'migration script (see PAYG_DESIGN.md §3.10 revised) can override.'; + +-- --------------------------------------------------------------------------------------------- +-- 2. RPC: payg_link_subscription +-- +-- Called by: +-- * supabase/functions/create-payg-team-subscription/index.ts (post-Stripe-Checkout, with +-- either user JWT [normal path, RLS-enforced] or service-role [day-1 migration script]) +-- * supabase/functions/stripe-webhook/handlers/payg-subscription.ts on +-- customer.subscription.created (idempotent — second invocation with same args is a no-op) +-- --------------------------------------------------------------------------------------------- + +CREATE OR REPLACE FUNCTION stirling_pdf.payg_link_subscription( + p_team_id BIGINT, + p_customer_id TEXT, + p_subscription_id TEXT +) RETURNS VOID +LANGUAGE plpgsql +SECURITY INVOKER +AS $$ +BEGIN + UPDATE stirling_pdf.payg_team_extensions + SET stripe_customer_id = p_customer_id, + payg_subscription_id = p_subscription_id, + updated_at = now() + WHERE team_id = p_team_id; + + IF NOT FOUND THEN + RAISE EXCEPTION 'payg_team_extensions row missing for team %', p_team_id + USING ERRCODE = 'foreign_key_violation'; + END IF; + + INSERT INTO stirling_pdf.payg_subscription_change_log(team_id, action, subscription_id) + VALUES (p_team_id, 'LINKED', p_subscription_id); +END $$; + +COMMENT ON FUNCTION stirling_pdf.payg_link_subscription(BIGINT, TEXT, TEXT) IS + 'Idempotent link of a Stripe subscription to a team. SECURITY INVOKER means RLS applies — ' + 'the caller must be a LEADER of the team (or hold the service-role bypass). ' + 'Writes an audit row to payg_subscription_change_log.'; + +-- --------------------------------------------------------------------------------------------- +-- 3. RPC: payg_unlink_subscription +-- +-- Called by stripe-webhook handlers/payg-subscription.ts on customer.subscription.deleted +-- (after Stripe's own retries have given up). Drops the team back to free-tier-then-block. +-- --------------------------------------------------------------------------------------------- + +CREATE OR REPLACE FUNCTION stirling_pdf.payg_unlink_subscription( + p_team_id BIGINT, + p_reason TEXT +) RETURNS VOID +LANGUAGE plpgsql +SECURITY INVOKER +AS $$ +BEGIN + UPDATE stirling_pdf.payg_team_extensions + SET payg_subscription_id = NULL, + updated_at = now() + WHERE team_id = p_team_id; + -- We deliberately keep stripe_customer_id — the team may add a new card later and we'd + -- like to reuse the existing Stripe customer record rather than create a duplicate. + + INSERT INTO stirling_pdf.payg_subscription_change_log(team_id, action, reason) + VALUES (p_team_id, 'UNLINKED', p_reason); +END $$; + +COMMENT ON FUNCTION stirling_pdf.payg_unlink_subscription(BIGINT, TEXT) IS + 'Drops the team back to free-tier-then-block derived state. Reason is logged for audit ' + '(typically subscription_deleted | admin | card_removed).'; + +-- --------------------------------------------------------------------------------------------- +-- 4. Auto-create payg_team_extensions row when a team is created +-- +-- Every new signup gets a payg_team_extensions row with NULL pricing_policy_id (which the +-- backend's PricingPolicyService resolves to the default policy). The free-tier gate kicks in +-- from the very first tool call. +-- --------------------------------------------------------------------------------------------- + +CREATE OR REPLACE FUNCTION stirling_pdf.payg_create_team_extensions_trigger() +RETURNS TRIGGER +LANGUAGE plpgsql +AS $$ +BEGIN + INSERT INTO stirling_pdf.payg_team_extensions(team_id) + VALUES (NEW.team_id) + ON CONFLICT (team_id) DO NOTHING; + RETURN NEW; +END $$; + +DROP TRIGGER IF EXISTS trg_payg_create_team_extensions ON stirling_pdf.teams; +CREATE TRIGGER trg_payg_create_team_extensions + AFTER INSERT ON stirling_pdf.teams + FOR EACH ROW + EXECUTE FUNCTION stirling_pdf.payg_create_team_extensions_trigger(); + +COMMENT ON TRIGGER trg_payg_create_team_extensions ON stirling_pdf.teams IS + 'Ensures every team has a payg_team_extensions sidecar row from creation. New customers ' + 'are PAYG-default from minute one — they consume free-tier units until they add a card.'; + +-- --------------------------------------------------------------------------------------------- +-- 5. Backfill: any existing team without a sidecar row gets one now +-- --------------------------------------------------------------------------------------------- + +INSERT INTO stirling_pdf.payg_team_extensions(team_id) +SELECT t.team_id + FROM stirling_pdf.teams t + WHERE NOT EXISTS ( + SELECT 1 FROM stirling_pdf.payg_team_extensions x WHERE x.team_id = t.team_id + ); + +-- --------------------------------------------------------------------------------------------- +-- 6. RLS policy +-- +-- Service-role bypasses RLS (backend reads + day-1 migration script writes via the service-role +-- key). For user-initiated writes via the frontend Add-Card flow, only team LEADERs can link a +-- subscription. SELECT remains permissive — anyone in the team can see the row. +-- --------------------------------------------------------------------------------------------- + +ALTER TABLE stirling_pdf.payg_team_extensions ENABLE ROW LEVEL SECURITY; + +-- Read: any team member can see their team's payg row. +DROP POLICY IF EXISTS payg_team_ext_select ON stirling_pdf.payg_team_extensions; +CREATE POLICY payg_team_ext_select + ON stirling_pdf.payg_team_extensions + FOR SELECT + USING ( + team_id IN ( + SELECT tm.team_id + FROM stirling_pdf.team_memberships tm + JOIN stirling_pdf.users u ON u.user_id = tm.user_id + WHERE u.supabase_auth_id = auth.uid() + ) + ); + +-- Update: only LEADERs of the team can update (i.e. link / unlink a subscription). +DROP POLICY IF EXISTS payg_team_ext_leader_update ON stirling_pdf.payg_team_extensions; +CREATE POLICY payg_team_ext_leader_update + ON stirling_pdf.payg_team_extensions + FOR UPDATE + USING ( + team_id IN ( + SELECT tm.team_id + FROM stirling_pdf.team_memberships tm + JOIN stirling_pdf.users u ON u.user_id = tm.user_id + WHERE u.supabase_auth_id = auth.uid() + AND tm.role = 'LEADER' + ) + ); diff --git a/app/saas/src/main/resources/db/migration/saas/V15__payg_audit_logs.sql b/app/saas/src/main/resources/db/migration/saas/V15__payg_audit_logs.sql new file mode 100644 index 0000000000..1b7e61e9c1 --- /dev/null +++ b/app/saas/src/main/resources/db/migration/saas/V15__payg_audit_logs.sql @@ -0,0 +1,76 @@ +-- PAYG audit-log tables. Two append-only logs: +-- +-- * payg_meter_event_log — written by the backend's PaygMeterReportingService +-- on every Stripe meter event POST attempt. Gives us a +-- record independent of Stripe's own logs so we can +-- replay-after-24h-window (Stripe's idempotency window) +-- and run nightly reconciliation against Stripe's +-- meter-event list. +-- +-- * payg_subscription_change_log — written by V14's two RPC functions on every +-- subscription link / unlink. Independent of Stripe's +-- webhook log; lets us diagnose "why is this team in +-- free-tier-block when their Stripe sub is active?" +-- without leaving our DB. +-- +-- Both are pure additive; nothing reads them yet. PR-SB-5 (nightly reconcile) wires the +-- meter-event log; the subscription change log is queried only from admin tooling. +-- +-- Design references: +-- * payg-stripe-supabase-plan.html §3.10 — twin migrations +-- * payg-stripe-supabase-plan.html §8 H5 — 24h idempotency window mitigation + +-- --------------------------------------------------------------------------------------------- +-- 1. payg_meter_event_log — backend-side audit of every Stripe meter event we tried to post. +-- --------------------------------------------------------------------------------------------- + +CREATE TABLE IF NOT EXISTS stirling_pdf.payg_meter_event_log ( + event_id BIGSERIAL PRIMARY KEY, + team_id BIGINT NOT NULL REFERENCES stirling_pdf.teams(team_id) ON DELETE CASCADE, + job_id UUID, + idempotency_key VARCHAR(128) NOT NULL UNIQUE, + units INTEGER NOT NULL, + occurred_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP, + posted_to_stripe_at TIMESTAMP, + -- NULL while pending; set when the meter-payg-units edge fn returns success. NULL after + -- 24h means the event never made it to Stripe — nightly reconcile retries with a fresh + -- idempotency-key suffix (see §8 H5 mitigation). + stripe_error_code VARCHAR(64), + stripe_error_body TEXT, + metadata JSONB +); + +CREATE INDEX IF NOT EXISTS idx_payg_meter_event_team_time + ON stirling_pdf.payg_meter_event_log (team_id, occurred_at); + +CREATE INDEX IF NOT EXISTS idx_payg_meter_event_unposted + ON stirling_pdf.payg_meter_event_log (occurred_at) + WHERE posted_to_stripe_at IS NULL; + +COMMENT ON TABLE stirling_pdf.payg_meter_event_log IS + 'Backend audit of every Stripe meter event POST attempt. Independent of Stripe meter ' + 'history. idempotency_key is the same one passed to Stripe; the UNIQUE constraint here ' + 'gives us safe at-least-once semantics even on backend retry. Rows older than 24h with ' + 'posted_to_stripe_at IS NULL are stuck and retried by the nightly reconcile job.'; + +-- --------------------------------------------------------------------------------------------- +-- 2. payg_subscription_change_log — written by V14's RPC functions on every link / unlink. +-- --------------------------------------------------------------------------------------------- + +CREATE TABLE IF NOT EXISTS stirling_pdf.payg_subscription_change_log ( + change_id BIGSERIAL PRIMARY KEY, + team_id BIGINT NOT NULL REFERENCES stirling_pdf.teams(team_id) ON DELETE CASCADE, + action VARCHAR(32) NOT NULL, + -- LINKED — payg_link_subscription written subscription_id + -- UNLINKED — payg_unlink_subscription cleared the subscription + subscription_id VARCHAR(128), + reason TEXT, + changed_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP +); + +CREATE INDEX IF NOT EXISTS idx_payg_sub_change_team_time + ON stirling_pdf.payg_subscription_change_log (team_id, changed_at); + +COMMENT ON TABLE stirling_pdf.payg_subscription_change_log IS + 'Append-only log of every subscription link / unlink. Written by V14 RPC functions; ' + 'never updated. Diagnostic value when reconciling against Stripe webhook history.'; diff --git a/app/saas/src/main/resources/db/migration/saas/V16__payg_billing_category.sql b/app/saas/src/main/resources/db/migration/saas/V16__payg_billing_category.sql new file mode 100644 index 0000000000..f78458a1b6 --- /dev/null +++ b/app/saas/src/main/resources/db/migration/saas/V16__payg_billing_category.sql @@ -0,0 +1,58 @@ +-- PAYG analytics axis: stamp every billable ledger entry / shadow row with the category that +-- produced it (API | AI | AUTOMATION | BYPASSED). PAYG stays on a single flat-priced Stripe meter +-- forever — this column is for in-app breakdowns and analytics, never for Stripe pricing. +-- +-- All adds are nullable: pre-V16 rows have no category and stay NULL; the interceptor populates +-- it for new rows going forward. + +-- --------------------------------------------------------------------------------------------- +-- 1. wallet_ledger.billing_category +-- --------------------------------------------------------------------------------------------- +ALTER TABLE wallet_ledger ADD COLUMN IF NOT EXISTS billing_category VARCHAR(16) NULL; +COMMENT ON COLUMN wallet_ledger.billing_category IS + 'API | AI | AUTOMATION | BYPASSED. NULL = system entry or pre-V16 backfill.'; + +-- Partial index — only billable rows ever read this column, and NULLs would just bloat the tree. +CREATE INDEX IF NOT EXISTS idx_wallet_ledger_team_category_period + ON wallet_ledger (team_id, billing_category, occurred_at) + WHERE billing_category IS NOT NULL; + +-- --------------------------------------------------------------------------------------------- +-- 2. payg_shadow_charge.billing_category + job_source +-- --------------------------------------------------------------------------------------------- +ALTER TABLE payg_shadow_charge + ADD COLUMN IF NOT EXISTS billing_category VARCHAR(16) NULL, + ADD COLUMN IF NOT EXISTS job_source VARCHAR(32) NULL; + +-- Backfill job_source from processing_job (best-effort — rows whose job has already been pruned +-- stay NULL, which is fine: the shadow row is self-describing post-V16 and only legacy ones lack +-- the column.) +UPDATE payg_shadow_charge sc + SET job_source = pj.source + FROM processing_job pj + WHERE pj.job_id = sc.job_id + AND sc.job_source IS NULL; + +-- --------------------------------------------------------------------------------------------- +-- 3. pricing_policy_stripe_price.stripe_product_id +-- Operator populates this manually per row when seeding new policies. Nullable for backward +-- compatibility with existing rows that don't carry a Product reference. +-- --------------------------------------------------------------------------------------------- +ALTER TABLE pricing_policy_stripe_price + ADD COLUMN IF NOT EXISTS stripe_product_id VARCHAR(128) NULL; + +-- --------------------------------------------------------------------------------------------- +-- 4. wallet_category_summary view — pre-grouped per-team, per-month, per-category aggregate that +-- the in-app breakdown widget reads. Recomputed live on every SELECT; cheap thanks to the +-- partial index above. +-- --------------------------------------------------------------------------------------------- +CREATE OR REPLACE VIEW wallet_category_summary AS +SELECT + team_id, + date_trunc('month', occurred_at) AS period_start, + billing_category, + SUM(CASE WHEN amount_units < 0 THEN -amount_units ELSE 0 END) AS units_debited, + COUNT(*) FILTER (WHERE entry_type = 'DEBIT') AS debit_count +FROM wallet_ledger +WHERE billing_category IS NOT NULL +GROUP BY team_id, date_trunc('month', occurred_at), billing_category; diff --git a/app/saas/src/main/resources/db/migration/saas/V17__merge_supabase_id_into_auth_id.sql b/app/saas/src/main/resources/db/migration/saas/V17__merge_supabase_id_into_auth_id.sql new file mode 100644 index 0000000000..cc6b5d033e --- /dev/null +++ b/app/saas/src/main/resources/db/migration/saas/V17__merge_supabase_id_into_auth_id.sql @@ -0,0 +1,34 @@ +-- Consolidate users.supabase_id into users.supabase_auth_id. +-- +-- The `supabase_auth_id` column is the canonical link to Supabase Auth — it was +-- created by the initial Supabase schema migration (Sep 2025) and is referenced +-- by every RLS policy in the Supabase side of the world (V14's +-- payg_team_ext_select / payg_team_ext_leader_update, the public.payg_* +-- SECURITY DEFINER RPCs, etc.). +-- +-- PR #6384 ("SaaS Consolidation") accidentally added a parallel `supabase_id` +-- column via Flyway V2 — same purpose, different name. Java's User entity then +-- mapped to this new column. The result was a split-brain: +-- * Pre-#6384 users had supabase_auth_id populated, supabase_id NULL. +-- * Post-#6384 users had supabase_id populated, supabase_auth_id NULL. +-- * RLS policies + RPCs always check supabase_auth_id, so post-#6384 users +-- failed every membership check. +-- +-- This migration: +-- 1. Backfills supabase_auth_id from supabase_id where the former is NULL. +-- 2. Drops the supabase_id column and its unique index. +-- +-- The Java User entity has been switched to @Column(name = "supabase_auth_id") +-- in the same change-set; this migration assumes the new code is already +-- deployed (or will be deployed together with this migration). + +-- 1. Backfill the canonical column from the duplicate, where needed. +UPDATE users + SET supabase_auth_id = supabase_id + WHERE supabase_auth_id IS NULL + AND supabase_id IS NOT NULL; + +-- 2. Drop the duplicate column. IF EXISTS guards against environments where +-- the column was already removed manually. +DROP INDEX IF EXISTS uk_users_supabase_id; +ALTER TABLE users DROP COLUMN IF EXISTS supabase_id; diff --git a/app/saas/src/main/resources/db/migration/saas/V19__payg_lifetime_free_grant.sql b/app/saas/src/main/resources/db/migration/saas/V19__payg_lifetime_free_grant.sql new file mode 100644 index 0000000000..ed62b34e02 --- /dev/null +++ b/app/saas/src/main/resources/db/migration/saas/V19__payg_lifetime_free_grant.sql @@ -0,0 +1,93 @@ +-- PAYG free allowance: monthly per-cycle allowance → one-time LIFETIME grant. +-- +-- Product decision (2026-06-11): every team gets a one-time free document grant. It does NOT +-- replenish monthly and is NOT lost when the team subscribes — they keep whatever is unused. +-- +-- Mechanics: the grant is tracked as a running counter on the team sidecar +-- (payg_team_extensions.free_units_remaining), seeded once from the team's effective pricing +-- policy and maintained by the charge pipeline (deducted when a billable DEBIT is written, +-- restored on a first-step refund). Because the counter is authoritative, the wallet_ledger is +-- no longer the source of truth for the grant and its old rows can be pruned after a retention +-- window (separate future job). + +-- --------------------------------------------------------------------------------------------- +-- 1. Rename the policy column — it is no longer "per cycle", it's the one-time grant size. +-- --------------------------------------------------------------------------------------------- + +ALTER TABLE stirling_pdf.pricing_policy + RENAME COLUMN free_tier_units_per_cycle TO free_tier_units; + +COMMENT ON COLUMN stirling_pdf.pricing_policy.free_tier_units IS + 'One-time lifetime free document grant handed to a team on creation (copied into ' + 'payg_team_extensions.free_units_remaining). NOT per-cycle: it never replenishes and ' + 'survives subscribing. 0 = no free grant (block / meter from the first document).'; + +-- --------------------------------------------------------------------------------------------- +-- 2. The running counter on the team sidecar. +-- --------------------------------------------------------------------------------------------- + +ALTER TABLE stirling_pdf.payg_team_extensions + ADD COLUMN IF NOT EXISTS free_units_remaining BIGINT NOT NULL DEFAULT 0; + +COMMENT ON COLUMN stirling_pdf.payg_team_extensions.free_units_remaining IS + 'Remaining one-time free documents for this team. Seeded from the effective pricing ' + 'policy''s free_tier_units at row creation; decremented by min(jobUnits, remaining) when a ' + 'billable charge is written; restored on a first-step refund. Lifetime — never resets. ' + 'Authoritative source for the free grant (independent of wallet_ledger retention).'; + +-- --------------------------------------------------------------------------------------------- +-- 3. Per-job free/paid split on the shadow row — makes metering + refunds exact and removes +-- any need to SUM the ledger over a team's lifetime. +-- --------------------------------------------------------------------------------------------- + +ALTER TABLE stirling_pdf.payg_shadow_charge + ADD COLUMN IF NOT EXISTS free_units_consumed INT NOT NULL DEFAULT 0; + +COMMENT ON COLUMN stirling_pdf.payg_shadow_charge.free_units_consumed IS + 'How many of this job''s payg_units came out of the team''s free grant at charge time. ' + 'Paid (metered) units = payg_units - free_units_consumed. A refund restores this many ' + 'units to payg_team_extensions.free_units_remaining.'; + +-- --------------------------------------------------------------------------------------------- +-- 4. Seed the counter at team creation. Replace the V14 trigger function so new teams get the +-- default policy's grant from minute one. (The trigger itself still points at this function.) +-- --------------------------------------------------------------------------------------------- + +CREATE OR REPLACE FUNCTION stirling_pdf.payg_create_team_extensions_trigger() +RETURNS TRIGGER +LANGUAGE plpgsql +AS $$ +BEGIN + INSERT INTO stirling_pdf.payg_team_extensions(team_id, free_units_remaining) + VALUES ( + NEW.team_id, + COALESCE( + (SELECT pp.free_tier_units FROM stirling_pdf.pricing_policy pp + WHERE pp.is_default = TRUE LIMIT 1), + 0) + ) + ON CONFLICT (team_id) DO NOTHING; + RETURN NEW; +END $$; + +-- --------------------------------------------------------------------------------------------- +-- 5. Backfill existing teams. remaining = max(0, grant - lifetime_consumed). lifetime_consumed +-- is -SUM(amount_units) over the team's DEBIT+REFUND ledger entries (debits negative, refunds +-- positive), so grant + SUM(amount_units) collapses to grant - consumed. One-time read of the +-- ledger; after this the counter stands alone. Grant = team override policy, else the default. +-- --------------------------------------------------------------------------------------------- + +UPDATE stirling_pdf.payg_team_extensions ext + SET free_units_remaining = GREATEST( + 0, + COALESCE( + (SELECT pp.free_tier_units FROM stirling_pdf.pricing_policy pp + WHERE pp.policy_id = ext.pricing_policy_id), + (SELECT pp.free_tier_units FROM stirling_pdf.pricing_policy pp + WHERE pp.is_default = TRUE LIMIT 1), + 0) + + COALESCE( + (SELECT SUM(wl.amount_units) FROM stirling_pdf.wallet_ledger wl + WHERE wl.team_id = ext.team_id + AND wl.entry_type IN ('DEBIT', 'REFUND')), + 0)); diff --git a/app/saas/src/main/resources/db/migration/saas/V20__payg_launch_free_grant.sql b/app/saas/src/main/resources/db/migration/saas/V20__payg_launch_free_grant.sql new file mode 100644 index 0000000000..d9dc756e31 --- /dev/null +++ b/app/saas/src/main/resources/db/migration/saas/V20__payg_launch_free_grant.sql @@ -0,0 +1,42 @@ +-- PAYG launch free grant: give the default pricing policy a real one-time grant. +-- +-- V14 added pricing_policy.free_tier_units with DEFAULT 0, and the default policy seeded in V12 +-- predates the column — so on a fresh deploy every team's free_units_remaining seeds to 0 and +-- V19's "every team gets a one-time free grant" intent ships dead (teams are gated / metered from +-- the very first billable document). This migration sets the launch grant on the default policy +-- and re-seeds existing teams that V19 left at 0 (V19 ran its backfill while the grant was still +-- 0, so every then-existing team computed to 0). +-- +-- The launch value lives on the default policy row; tune it there (or via a future admin surface). +-- Both updates are guarded so a deliberately-tuned value — e.g. a smaller test grant — is never +-- clobbered. + +-- --------------------------------------------------------------------------------------------- +-- 1. Launch grant on the default policy, only where it's still the accidental 0. +-- --------------------------------------------------------------------------------------------- +UPDATE stirling_pdf.pricing_policy + SET free_tier_units = 500 + WHERE is_default = TRUE + AND free_tier_units = 0; + +-- --------------------------------------------------------------------------------------------- +-- 2. Re-seed existing teams V19 left at 0. Same recompute as V19's backfill — remaining = +-- max(0, grant + net signed DEBIT/REFUND) — now that the grant is non-zero. Guarded to +-- free_units_remaining = 0: a team with a deliberately-set positive balance is left alone, and +-- a team that genuinely exhausted a real grant also recomputes to 0, so the guard is safe. +-- --------------------------------------------------------------------------------------------- +UPDATE stirling_pdf.payg_team_extensions ext + SET free_units_remaining = GREATEST( + 0, + COALESCE( + (SELECT pp.free_tier_units FROM stirling_pdf.pricing_policy pp + WHERE pp.policy_id = ext.pricing_policy_id), + (SELECT pp.free_tier_units FROM stirling_pdf.pricing_policy pp + WHERE pp.is_default = TRUE LIMIT 1), + 0) + + COALESCE( + (SELECT SUM(wl.amount_units) FROM stirling_pdf.wallet_ledger wl + WHERE wl.team_id = ext.team_id + AND wl.entry_type IN ('DEBIT', 'REFUND')), + 0)) + WHERE ext.free_units_remaining = 0; diff --git a/app/saas/src/main/resources/db/migration/saas/V21__drop_wallet_category_summary_view.sql b/app/saas/src/main/resources/db/migration/saas/V21__drop_wallet_category_summary_view.sql new file mode 100644 index 0000000000..857e0bfa65 --- /dev/null +++ b/app/saas/src/main/resources/db/migration/saas/V21__drop_wallet_category_summary_view.sql @@ -0,0 +1,11 @@ +-- Drop the unused wallet_category_summary view. +-- +-- V16 created this view to back the wallet's per-category spend breakdown via +-- WalletCategorySummaryDao. That DAO was never wired up — the breakdown is built from the JPA +-- repository (WalletLedgerRepository.sumPeriodAmountByCategory) instead — so both the DAO and this +-- view have zero readers. The DAO is deleted in the same change; this drops the dead view. +-- +-- Done as a new migration (not by editing V16) so Flyway's checksum validation doesn't fail on +-- databases that already applied V16. IF EXISTS keeps it safe on DBs where V16 hasn't run. + +DROP VIEW IF EXISTS stirling_pdf.wallet_category_summary; diff --git a/app/saas/src/main/resources/static/modern-logo.svg b/app/saas/src/main/resources/static/modern-logo.svg new file mode 100644 index 0000000000..a4a1a1f87e --- /dev/null +++ b/app/saas/src/main/resources/static/modern-logo.svg @@ -0,0 +1,4 @@ + + + + diff --git a/app/saas/src/main/resources/static/saas-landing.html b/app/saas/src/main/resources/static/saas-landing.html new file mode 100644 index 0000000000..3ffef1d7bf --- /dev/null +++ b/app/saas/src/main/resources/static/saas-landing.html @@ -0,0 +1,250 @@ + + + + + + + Stirling PDF - Cloud API + + + + + + + + + + + +

+
+ + +
+

+ You've reached the Stirling PDF Cloud API. This endpoint serves the REST API that powers our apps and your integrations. +

+
+ + Open API Documentation + Read the Developer Docs + +
+ +
+

Just looking to edit a PDF?

+

+ This page is the API endpoint - it doesn't have an editor. If you want the full interface: +

+ +
+ +
+ +
+

Want to run your own?

+

+ Stirling PDF is open source. You can self-host the same toolkit on your own hardware or run it inside your network: +

+ +
+ +
+ +
+

Support

+ +
+
+
+ + + + diff --git a/app/saas/src/test/java/stirling/software/saas/ai/controller/AiCreateControllerTest.java b/app/saas/src/test/java/stirling/software/saas/ai/controller/AiCreateControllerTest.java new file mode 100644 index 0000000000..0614869702 --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/ai/controller/AiCreateControllerTest.java @@ -0,0 +1,946 @@ +package stirling.software.saas.ai.controller; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyBoolean; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.verifyNoInteractions; +import static org.mockito.Mockito.when; + +import java.io.ByteArrayInputStream; +import java.io.ByteArrayOutputStream; +import java.io.InputStream; +import java.net.http.HttpResponse; +import java.nio.charset.StandardCharsets; +import java.time.Instant; +import java.util.List; +import java.util.Map; + +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.ArgumentCaptor; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.http.HttpHeaders; +import org.springframework.http.HttpStatus; +import org.springframework.http.MediaType; +import org.springframework.http.ResponseEntity; +import org.springframework.security.authentication.UsernamePasswordAuthenticationToken; +import org.springframework.security.core.context.SecurityContextHolder; +import org.springframework.web.server.ResponseStatusException; +import org.springframework.web.servlet.mvc.method.annotation.StreamingResponseBody; + +import jakarta.servlet.http.HttpServletRequest; + +import stirling.software.proprietary.model.Team; +import stirling.software.proprietary.security.database.repository.UserRepository; +import stirling.software.proprietary.security.model.ApiKeyAuthenticationToken; +import stirling.software.proprietary.security.model.User; +import stirling.software.saas.ai.controller.AiCreateController.AiCreateSessionResponse; +import stirling.software.saas.ai.controller.AiCreateController.AiCreateSessionSummary; +import stirling.software.saas.ai.controller.AiCreateController.CreateSessionRequest; +import stirling.software.saas.ai.controller.AiCreateController.CreateSessionResponse; +import stirling.software.saas.ai.controller.AiCreateController.DraftRequest; +import stirling.software.saas.ai.controller.AiCreateController.DraftSection; +import stirling.software.saas.ai.controller.AiCreateController.OutlineRequest; +import stirling.software.saas.ai.controller.AiCreateController.RepromptRequest; +import stirling.software.saas.ai.controller.AiCreateController.TemplateRequest; +import stirling.software.saas.ai.model.AiCreateSession; +import stirling.software.saas.ai.model.AiCreateSessionStatus; +import stirling.software.saas.ai.repository.AiCreateSessionRepository.AiCreateSessionSummaryProjection; +import stirling.software.saas.ai.service.AiCreateProxyService; +import stirling.software.saas.ai.service.AiCreateSessionService; +import stirling.software.saas.payg.charge.ChargeContext; +import stirling.software.saas.payg.charge.JobChargeService; +import stirling.software.saas.payg.model.BillingCategory; +import stirling.software.saas.payg.model.JobSource; +import stirling.software.saas.payg.model.ProcessType; + +/** + * Pure unit tests for {@link AiCreateController}. All collaborators are mocked; the controller's + * handler methods are invoked directly and asserted via {@link ResponseEntity} / {@code verify}. + * + *

The controller reads {@code SecurityContextHolder} in the charge path, so each relevant test + * seeds an authentication and {@link #clearSecurityContext()} resets it afterwards. + */ +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class AiCreateControllerTest { + + @Mock private AiCreateSessionService sessionService; + @Mock private AiCreateProxyService proxyService; + @Mock private UserRepository userRepository; + @Mock private JobChargeService jobChargeService; + + private AiCreateController controller; + + @org.junit.jupiter.api.BeforeEach + void setUp() { + controller = + new AiCreateController( + sessionService, proxyService, userRepository, jobChargeService); + } + + @AfterEach + void clearSecurityContext() { + SecurityContextHolder.clearContext(); + } + + // ---------------------------------------------------------------------------------------------- + // createSession + chargeForCreate + // ---------------------------------------------------------------------------------------------- + + @Nested + @DisplayName("createSession") + class CreateSession { + + @Test + @DisplayName("happy path returns 200 with the new sessionId and charges one AI unit") + void createSession_happyPath_returnsIdAndCharges() { + authenticateWeb(userWithTeam(7L, 100L)); + AiCreateSession created = session("sess-1", "user-x"); + when(sessionService.createSession( + "write a report", "letter", "tmpl-1", "tex", "preview")) + .thenReturn(created); + + CreateSessionRequest req = + new CreateSessionRequest( + "write a report", "letter", "tmpl-1", "tex", "preview"); + ResponseEntity resp = controller.createSession(req); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(resp.getBody()).isNotNull(); + assertThat(resp.getBody().sessionId()).isEqualTo("sess-1"); + verify(sessionService) + .createSession("write a report", "letter", "tmpl-1", "tex", "preview"); + } + + @Test + @DisplayName("WEB auth charges a single AI unit with WEB source and the user's team") + void createSession_webAuth_chargesAiUnitWithWebSource() { + authenticateWeb(userWithTeam(7L, 100L)); + when(sessionService.createSession(any(), any(), any(), any(), any())) + .thenReturn(session("sess-1", "user-x")); + + controller.createSession(new CreateSessionRequest("p", null, null, null, null)); + + ArgumentCaptor ctx = ArgumentCaptor.forClass(ChargeContext.class); + verify(jobChargeService).chargeStandalone(ctx.capture(), eq(1)); + ChargeContext c = ctx.getValue(); + assertThat(c.ownerUserId()).isEqualTo(7L); + assertThat(c.ownerTeamId()).isEqualTo(100L); + assertThat(c.source()).isEqualTo(JobSource.WEB); + assertThat(c.processType()).isEqualTo(ProcessType.SINGLE_TOOL); + assertThat(c.billingCategory()).isEqualTo(BillingCategory.AI); + } + + @Test + @DisplayName("API-key auth charges with API source (AI usage billed the same as web)") + void createSession_apiKeyAuth_chargesAiUnitWithApiSource() { + User user = userWithTeam(7L, 100L); + authenticateApiKey(user); + when(sessionService.createSession(any(), any(), any(), any(), any())) + .thenReturn(session("sess-1", "user-x")); + + controller.createSession(new CreateSessionRequest("p", null, null, null, null)); + + ArgumentCaptor ctx = ArgumentCaptor.forClass(ChargeContext.class); + verify(jobChargeService).chargeStandalone(ctx.capture(), eq(1)); + assertThat(ctx.getValue().source()).isEqualTo(JobSource.API); + assertThat(ctx.getValue().billingCategory()).isEqualTo(BillingCategory.AI); + } + + @Test + @DisplayName("null prompt is rejected with 400 and never reaches the service") + void createSession_nullPrompt_throwsBadRequest() { + CreateSessionRequest req = new CreateSessionRequest(null, null, null, null, null); + + assertThatThrownBy(() -> controller.createSession(req)) + .isInstanceOf(ResponseStatusException.class) + .satisfies( + e -> + assertThat(((ResponseStatusException) e).getStatusCode()) + .isEqualTo(HttpStatus.BAD_REQUEST)); + verifyNoInteractions(sessionService); + verifyNoInteractions(jobChargeService); + } + + @Test + @DisplayName("blank/whitespace prompt is rejected with 400") + void createSession_blankPrompt_throwsBadRequest() { + CreateSessionRequest req = new CreateSessionRequest(" ", null, null, null, null); + + assertThatThrownBy(() -> controller.createSession(req)) + .isInstanceOf(ResponseStatusException.class); + verifyNoInteractions(sessionService); + verifyNoInteractions(jobChargeService); + } + + @Test + @DisplayName("no authentication: session still created, charge is skipped (no NPE)") + void createSession_noAuth_skipsChargeButCreatesSession() { + SecurityContextHolder.clearContext(); + when(sessionService.createSession(any(), any(), any(), any(), any())) + .thenReturn(session("sess-1", "user-x")); + + ResponseEntity resp = + controller.createSession(new CreateSessionRequest("p", null, null, null, null)); + + assertThat(resp.getBody().sessionId()).isEqualTo("sess-1"); + verify(jobChargeService, never()) + .chargeStandalone(any(), org.mockito.ArgumentMatchers.anyInt()); + } + + @Test + @DisplayName("user has no team: charge is skipped (free-grant accounting needs a team)") + void createSession_userWithoutTeam_skipsCharge() { + User user = new User(); + user.setId(7L); + // No team. + authenticateWeb(user); + when(sessionService.createSession(any(), any(), any(), any(), any())) + .thenReturn(session("sess-1", "user-x")); + + controller.createSession(new CreateSessionRequest("p", null, null, null, null)); + + verify(jobChargeService, never()) + .chargeStandalone(any(), org.mockito.ArgumentMatchers.anyInt()); + } + + @Test + @DisplayName("charge failure is best-effort: session is still returned to the caller") + void createSession_chargeThrows_sessionStillSucceeds() { + authenticateWeb(userWithTeam(7L, 100L)); + when(sessionService.createSession(any(), any(), any(), any(), any())) + .thenReturn(session("sess-1", "user-x")); + when(jobChargeService.chargeStandalone(any(), org.mockito.ArgumentMatchers.anyInt())) + .thenThrow(new IllegalStateException("stripe down")); + + ResponseEntity resp = + controller.createSession(new CreateSessionRequest("p", null, null, null, null)); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(resp.getBody().sessionId()).isEqualTo("sess-1"); + } + } + + // ---------------------------------------------------------------------------------------------- + // deleteSession + // ---------------------------------------------------------------------------------------------- + + @Test + @DisplayName("deleteSession returns 204 and delegates to the service") + void deleteSession_returnsNoContentAndDelegates() { + ResponseEntity resp = controller.deleteSession("sess-1"); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.NO_CONTENT); + assertThat(resp.getBody()).isNull(); + verify(sessionService).deleteSessionForCurrentUser("sess-1"); + } + + @Test + @DisplayName("deleteSession propagates a not-found from the service") + void deleteSession_propagatesServiceError() { + org.mockito.Mockito.doThrow(new ResponseStatusException(HttpStatus.NOT_FOUND)) + .when(sessionService) + .deleteSessionForCurrentUser("missing"); + + assertThatThrownBy(() -> controller.deleteSession("missing")) + .isInstanceOf(ResponseStatusException.class); + } + + // ---------------------------------------------------------------------------------------------- + // getSession + toResponse mapping + // ---------------------------------------------------------------------------------------------- + + @Test + @DisplayName("getSession maps every entity field onto the response record") + void getSession_mapsAllFields() { + AiCreateSession s = session("sess-1", "user-x"); + s.setDocType("letter"); + s.setTemplateId("tmpl-1"); + s.setTemplateTex("\\documentclass{}"); + s.setPreviewTex("preview-tex"); + s.setPromptInitial("first prompt"); + s.setPromptLatest("latest prompt"); + s.setOutlineText("- one\n- two"); + s.setOutlineFilename("outline.txt"); + s.setOutlineApproved(true); + s.setOutlineConstraints("{\"tone\":\"formal\"}"); + s.setDraftSections("[{\"label\":\"Intro\",\"value\":\"hi\"}]"); + s.setPolishedLatex("\\section{Intro}"); + s.setPdfUrl("https://signed/url.pdf"); + Instant created = Instant.parse("2024-01-01T00:00:00Z"); + Instant updated = Instant.parse("2024-01-02T00:00:00Z"); + s.setCreatedAt(created); + s.setUpdatedAt(updated); + s.setStatus(AiCreateSessionStatus.DRAFT_READY); + when(sessionService.getSessionForCurrentUser("sess-1")).thenReturn(s); + + ResponseEntity resp = controller.getSession("sess-1"); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.OK); + AiCreateSessionResponse body = resp.getBody(); + assertThat(body).isNotNull(); + assertThat(body.sessionId()).isEqualTo("sess-1"); + assertThat(body.userId()).isEqualTo("user-x"); + assertThat(body.docType()).isEqualTo("letter"); + assertThat(body.templateId()).isEqualTo("tmpl-1"); + assertThat(body.templateTex()).isEqualTo("\\documentclass{}"); + assertThat(body.previewTex()).isEqualTo("preview-tex"); + assertThat(body.promptInitial()).isEqualTo("first prompt"); + assertThat(body.promptLatest()).isEqualTo("latest prompt"); + assertThat(body.outlineText()).isEqualTo("- one\n- two"); + assertThat(body.outlineFilename()).isEqualTo("outline.txt"); + assertThat(body.outlineApproved()).isTrue(); + assertThat(body.outlineConstraints()).containsEntry("tone", "formal"); + assertThat(body.draftSections()).containsExactly(new DraftSection("Intro", "hi")); + assertThat(body.polishedLatex()).isEqualTo("\\section{Intro}"); + assertThat(body.pdfUrl()).isEqualTo("https://signed/url.pdf"); + assertThat(body.createdAt()).isEqualTo(created); + assertThat(body.updatedAt()).isEqualTo(updated); + assertThat(body.status()).isEqualTo("DRAFT_READY"); + } + + @Test + @DisplayName("getSession with null status maps status to null and null payloads to null") + void getSession_nullStatusAndPayloads_mapToNull() { + AiCreateSession s = session("sess-1", "user-x"); + s.setStatus(null); + s.setOutlineConstraints(null); + s.setDraftSections(null); + when(sessionService.getSessionForCurrentUser("sess-1")).thenReturn(s); + + AiCreateSessionResponse body = controller.getSession("sess-1").getBody(); + + assertThat(body).isNotNull(); + assertThat(body.status()).isNull(); + assertThat(body.outlineConstraints()).isNull(); + assertThat(body.draftSections()).isNull(); + } + + @Test + @DisplayName("getSession with blank payloads parses to null rather than throwing") + void getSession_blankPayloads_mapToNull() { + AiCreateSession s = session("sess-1", "user-x"); + s.setOutlineConstraints(" "); + s.setDraftSections(""); + when(sessionService.getSessionForCurrentUser("sess-1")).thenReturn(s); + + AiCreateSessionResponse body = controller.getSession("sess-1").getBody(); + + assertThat(body.outlineConstraints()).isNull(); + assertThat(body.draftSections()).isNull(); + } + + @Test + @DisplayName("getSession with malformed JSON payloads degrades to null (logged, not thrown)") + void getSession_malformedPayloads_mapToNull() { + AiCreateSession s = session("sess-1", "user-x"); + s.setOutlineConstraints("{not-valid-json"); + s.setDraftSections("[oops"); + when(sessionService.getSessionForCurrentUser("sess-1")).thenReturn(s); + + AiCreateSessionResponse body = controller.getSession("sess-1").getBody(); + + assertThat(body.outlineConstraints()).isNull(); + assertThat(body.draftSections()).isNull(); + } + + @Test + @DisplayName("getSession propagates a not-found from the service") + void getSession_notFound_propagates() { + when(sessionService.getSessionForCurrentUser("missing")) + .thenThrow( + new ResponseStatusException(HttpStatus.NOT_FOUND, "AI session not found")); + + assertThatThrownBy(() -> controller.getSession("missing")) + .isInstanceOf(ResponseStatusException.class); + } + + // ---------------------------------------------------------------------------------------------- + // listSessions + toSummary mapping + page/size clamping + // ---------------------------------------------------------------------------------------------- + + @Nested + @DisplayName("listSessions") + class ListSessions { + + @Test + @DisplayName("maps projections to summaries and forwards includeDrafts") + void listSessions_mapsProjections() { + AiCreateSessionSummaryProjection p = + projection( + "sess-1", + "letter", + "tmpl-1", + "latest", + "initial", + AiCreateSessionStatus.SAVED, + "https://pdf", + Instant.parse("2024-01-01T00:00:00Z"), + Instant.parse("2024-01-02T00:00:00Z")); + when(sessionService.listSessionSummariesForCurrentUser(any(), eq(true))) + .thenReturn(List.of(p)); + + ResponseEntity> resp = + controller.listSessions(0, 10, true); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(resp.getBody()).hasSize(1); + AiCreateSessionSummary summary = resp.getBody().get(0); + assertThat(summary.sessionId()).isEqualTo("sess-1"); + assertThat(summary.docType()).isEqualTo("letter"); + assertThat(summary.templateId()).isEqualTo("tmpl-1"); + assertThat(summary.promptLatest()).isEqualTo("latest"); + assertThat(summary.promptInitial()).isEqualTo("initial"); + assertThat(summary.status()).isEqualTo("SAVED"); + assertThat(summary.pdfUrl()).isEqualTo("https://pdf"); + } + + @Test + @DisplayName("projection with null status maps to a null status string") + void listSessions_nullStatus_mapsToNull() { + AiCreateSessionSummaryProjection p = + projection("s", null, null, null, null, null, null, null, null); + when(sessionService.listSessionSummariesForCurrentUser(any(), eq(false))) + .thenReturn(List.of(p)); + + AiCreateSessionSummary summary = controller.listSessions(0, 10, false).getBody().get(0); + + assertThat(summary.status()).isNull(); + } + + @Test + @DisplayName("negative page is clamped to 0") + void listSessions_negativePage_clampedToZero() { + when(sessionService.listSessionSummariesForCurrentUser(any(), anyBoolean())) + .thenReturn(List.of()); + + controller.listSessions(-5, 10, false); + + ArgumentCaptor pr = + ArgumentCaptor.forClass(org.springframework.data.domain.PageRequest.class); + verify(sessionService).listSessionSummariesForCurrentUser(pr.capture(), eq(false)); + assertThat(pr.getValue().getPageNumber()).isZero(); + } + + @Test + @DisplayName("size above 50 is capped at 50") + void listSessions_oversizeSize_cappedAt50() { + when(sessionService.listSessionSummariesForCurrentUser(any(), anyBoolean())) + .thenReturn(List.of()); + + controller.listSessions(0, 9999, false); + + ArgumentCaptor pr = + ArgumentCaptor.forClass(org.springframework.data.domain.PageRequest.class); + verify(sessionService).listSessionSummariesForCurrentUser(pr.capture(), eq(false)); + assertThat(pr.getValue().getPageSize()).isEqualTo(50); + } + + @Test + @DisplayName("size below 1 is floored to 1") + void listSessions_zeroSize_flooredToOne() { + when(sessionService.listSessionSummariesForCurrentUser(any(), anyBoolean())) + .thenReturn(List.of()); + + controller.listSessions(0, 0, false); + + ArgumentCaptor pr = + ArgumentCaptor.forClass(org.springframework.data.domain.PageRequest.class); + verify(sessionService).listSessionSummariesForCurrentUser(pr.capture(), eq(false)); + assertThat(pr.getValue().getPageSize()).isEqualTo(1); + } + + @Test + @DisplayName("empty result yields an empty list, not null") + void listSessions_empty_returnsEmptyList() { + when(sessionService.listSessionSummariesForCurrentUser(any(), anyBoolean())) + .thenReturn(List.of()); + + ResponseEntity> resp = + controller.listSessions(0, 10, false); + + assertThat(resp.getBody()).isEmpty(); + } + } + + // ---------------------------------------------------------------------------------------------- + // updateOutline + // ---------------------------------------------------------------------------------------------- + + @Nested + @DisplayName("updateOutline") + class UpdateOutline { + + @Test + @DisplayName("serializes constraints to JSON and forwards text + filename") + void updateOutline_serializesConstraints() { + AiCreateSession updated = session("sess-1", "user-x"); + when(sessionService.updateOutline(eq("sess-1"), eq("the outline"), eq("o.txt"), any())) + .thenReturn(updated); + + OutlineRequest req = + new OutlineRequest("the outline", "o.txt", Map.of("tone", "formal")); + ResponseEntity resp = controller.updateOutline("sess-1", req); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.OK); + ArgumentCaptor payload = ArgumentCaptor.forClass(String.class); + verify(sessionService) + .updateOutline(eq("sess-1"), eq("the outline"), eq("o.txt"), payload.capture()); + assertThat(payload.getValue()).contains("\"tone\":\"formal\""); + } + + @Test + @DisplayName("null constraints forward a null payload (use AI-generated outline)") + void updateOutline_nullConstraints_forwardsNullPayload() { + when(sessionService.updateOutline(any(), any(), any(), any())) + .thenReturn(session("sess-1", "user-x")); + + controller.updateOutline("sess-1", new OutlineRequest("text", null, null)); + + verify(sessionService).updateOutline("sess-1", "text", null, null); + } + + @Test + @DisplayName("empty outline string is allowed (signals AI-generated outline)") + void updateOutline_emptyOutlineText_isAllowed() { + when(sessionService.updateOutline(any(), any(), any(), any())) + .thenReturn(session("sess-1", "user-x")); + + controller.updateOutline("sess-1", new OutlineRequest("", null, null)); + + verify(sessionService).updateOutline("sess-1", "", null, null); + } + + @Test + @DisplayName("null outline text is rejected with 400 before touching the service") + void updateOutline_nullText_throwsBadRequest() { + OutlineRequest req = new OutlineRequest(null, "o.txt", Map.of("a", "b")); + + assertThatThrownBy(() -> controller.updateOutline("sess-1", req)) + .isInstanceOf(ResponseStatusException.class) + .satisfies( + e -> + assertThat(((ResponseStatusException) e).getStatusCode()) + .isEqualTo(HttpStatus.BAD_REQUEST)); + verifyNoInteractions(sessionService); + } + } + + // ---------------------------------------------------------------------------------------------- + // reprompt + // ---------------------------------------------------------------------------------------------- + + @Test + @DisplayName("reprompt forwards the prompt and returns the mapped session") + void reprompt_forwardsPrompt() { + AiCreateSession s = session("sess-1", "user-x"); + s.setStatus(AiCreateSessionStatus.OUTLINE_PENDING); + when(sessionService.reprompt("sess-1", "new prompt")).thenReturn(s); + + ResponseEntity resp = + controller.reprompt("sess-1", new RepromptRequest("new prompt")); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(resp.getBody().sessionId()).isEqualTo("sess-1"); + assertThat(resp.getBody().status()).isEqualTo("OUTLINE_PENDING"); + verify(sessionService).reprompt("sess-1", "new prompt"); + } + + @Test + @DisplayName("reprompt with null prompt is rejected with 400") + void reprompt_nullPrompt_throwsBadRequest() { + assertThatThrownBy(() -> controller.reprompt("sess-1", new RepromptRequest(null))) + .isInstanceOf(ResponseStatusException.class); + verifyNoInteractions(sessionService); + } + + @Test + @DisplayName("reprompt with blank prompt is rejected with 400") + void reprompt_blankPrompt_throwsBadRequest() { + assertThatThrownBy(() -> controller.reprompt("sess-1", new RepromptRequest(" "))) + .isInstanceOf(ResponseStatusException.class); + verifyNoInteractions(sessionService); + } + + // ---------------------------------------------------------------------------------------------- + // updateDraft + // ---------------------------------------------------------------------------------------------- + + @Nested + @DisplayName("updateDraft") + class UpdateDraft { + + @Test + @DisplayName("serializes draft sections to a JSON array and forwards it") + void updateDraft_serializesSections() { + when(sessionService.updateDraftSections(eq("sess-1"), any())) + .thenReturn(session("sess-1", "user-x")); + + DraftRequest req = + new DraftRequest( + List.of( + new DraftSection("Intro", "hi"), + new DraftSection("Body", "x"))); + controller.updateDraft("sess-1", req); + + ArgumentCaptor payload = ArgumentCaptor.forClass(String.class); + verify(sessionService).updateDraftSections(eq("sess-1"), payload.capture()); + assertThat(payload.getValue()).contains("\"label\":\"Intro\""); + assertThat(payload.getValue()).contains("\"value\":\"hi\""); + } + + @Test + @DisplayName("empty list is allowed and serialized to []") + void updateDraft_emptyList_serializesToEmptyArray() { + when(sessionService.updateDraftSections(eq("sess-1"), any())) + .thenReturn(session("sess-1", "user-x")); + + controller.updateDraft("sess-1", new DraftRequest(List.of())); + + verify(sessionService).updateDraftSections("sess-1", "[]"); + } + + @Test + @DisplayName("null draft sections is rejected with 400") + void updateDraft_nullSections_throwsBadRequest() { + assertThatThrownBy(() -> controller.updateDraft("sess-1", new DraftRequest(null))) + .isInstanceOf(ResponseStatusException.class) + .satisfies( + e -> + assertThat(((ResponseStatusException) e).getStatusCode()) + .isEqualTo(HttpStatus.BAD_REQUEST)); + verifyNoInteractions(sessionService); + } + } + + // ---------------------------------------------------------------------------------------------- + // updateTemplate + // ---------------------------------------------------------------------------------------------- + + @Nested + @DisplayName("updateTemplate") + class UpdateTemplate { + + @Test + @DisplayName("docType only is accepted and forwarded") + void updateTemplate_docTypeOnly() { + when(sessionService.updateTemplate("sess-1", "letter", null)) + .thenReturn(session("sess-1", "user-x")); + + ResponseEntity resp = + controller.updateTemplate("sess-1", new TemplateRequest("letter", null)); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.OK); + verify(sessionService).updateTemplate("sess-1", "letter", null); + } + + @Test + @DisplayName("templateId only is accepted and forwarded") + void updateTemplate_templateIdOnly() { + when(sessionService.updateTemplate("sess-1", null, "tmpl-9")) + .thenReturn(session("sess-1", "user-x")); + + controller.updateTemplate("sess-1", new TemplateRequest(null, "tmpl-9")); + + verify(sessionService).updateTemplate("sess-1", null, "tmpl-9"); + } + + @Test + @DisplayName("both docType and templateId null is rejected with 400") + void updateTemplate_bothNull_throwsBadRequest() { + assertThatThrownBy( + () -> + controller.updateTemplate( + "sess-1", new TemplateRequest(null, null))) + .isInstanceOf(ResponseStatusException.class) + .satisfies( + e -> + assertThat(((ResponseStatusException) e).getStatusCode()) + .isEqualTo(HttpStatus.BAD_REQUEST)); + verifyNoInteractions(sessionService); + } + + @Test + @DisplayName("both docType and templateId blank is rejected with 400") + void updateTemplate_bothBlank_throwsBadRequest() { + assertThatThrownBy( + () -> + controller.updateTemplate( + "sess-1", new TemplateRequest(" ", ""))) + .isInstanceOf(ResponseStatusException.class); + verifyNoInteractions(sessionService); + } + } + + // ---------------------------------------------------------------------------------------------- + // fillFields (proxy, no credit header) + // ---------------------------------------------------------------------------------------------- + + @Nested + @DisplayName("fillFields") + class FillFields { + + @Test + @DisplayName("checks ownership, proxies POST, copies headers, and streams the body through") + void fillFields_proxiesAndStreams() throws Exception { + HttpServletRequest req = mock(HttpServletRequest.class); + HttpResponse upstream = + upstreamResponse( + 200, + "section data", + httpHeaders(Map.of(HttpHeaders.CONTENT_TYPE, "application/json"))); + when(proxyService.forward( + eq("POST"), + eq("/api/create/sessions/sess-1/fields"), + eq(req), + eq(false))) + .thenReturn(upstream); + + ResponseEntity resp = controller.fillFields("sess-1", req); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(resp.getHeaders().getFirst(HttpHeaders.CONTENT_TYPE)) + .isEqualTo("application/json"); + // Ownership guard runs before the proxy. + verify(sessionService).getSessionForCurrentUser("sess-1"); + assertThat(drain(resp.getBody())).isEqualTo("section data"); + } + + @Test + @DisplayName("ownership failure short-circuits before proxying") + void fillFields_ownershipFailure_doesNotProxy() throws Exception { + HttpServletRequest req = mock(HttpServletRequest.class); + when(sessionService.getSessionForCurrentUser("sess-1")) + .thenThrow(new ResponseStatusException(HttpStatus.NOT_FOUND)); + + assertThatThrownBy(() -> controller.fillFields("sess-1", req)) + .isInstanceOf(ResponseStatusException.class); + verify(proxyService, never()).forward(any(), any(), any(), anyBoolean()); + } + + @Test + @DisplayName("upstream error: returns 503 with a JSON error body, never throws") + void fillFields_proxyThrows_returns503() throws Exception { + HttpServletRequest req = mock(HttpServletRequest.class); + when(proxyService.forward(any(), any(), any(), anyBoolean())) + .thenThrow(new java.io.IOException("backend down")); + + ResponseEntity resp = controller.fillFields("sess-1", req); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.SERVICE_UNAVAILABLE); + assertThat(resp.getHeaders().getContentType()).isEqualTo(MediaType.APPLICATION_JSON); + assertThat(drain(resp.getBody())).contains("AI backend unavailable"); + } + } + + // ---------------------------------------------------------------------------------------------- + // stream (proxy, accept event-stream) + // ---------------------------------------------------------------------------------------------- + + @Nested + @DisplayName("stream") + class Stream { + + @Test + @DisplayName( + "checks ownership, proxies GET as event-stream, defaults Content-Type, and streams") + void stream_proxiesEventStreamAndStreams() throws Exception { + HttpServletRequest req = mock(HttpServletRequest.class); + HttpResponse upstream = + upstreamResponse(200, "data: hi\n\n", httpHeaders(Map.of())); + when(proxyService.forward( + eq("GET"), eq("/api/create/sessions/sess-1/stream"), eq(req), eq(true))) + .thenReturn(upstream); + + ResponseEntity resp = controller.stream("sess-1", req); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.OK); + // No upstream Content-Type → defaulted to text/event-stream. + assertThat(resp.getHeaders().getFirst(HttpHeaders.CONTENT_TYPE)) + .isEqualTo(MediaType.TEXT_EVENT_STREAM_VALUE); + // Ownership guard runs before the proxy. + verify(sessionService).getSessionForCurrentUser("sess-1"); + assertThat(drain(resp.getBody())).isEqualTo("data: hi\n\n"); + } + + @Test + @DisplayName("upstream non-2xx status is passed through; explicit Content-Type wins") + void stream_upstreamStatusAndExplicitContentTypePassedThrough() throws Exception { + SecurityContextHolder.clearContext(); + HttpServletRequest req = mock(HttpServletRequest.class); + HttpResponse upstream = + upstreamResponse( + 404, + "not found", + httpHeaders( + Map.of( + HttpHeaders.CONTENT_TYPE, + "text/plain", + HttpHeaders.CACHE_CONTROL, + "no-cache"))); + when(proxyService.forward(any(), any(), any(), eq(true))).thenReturn(upstream); + + ResponseEntity resp = controller.stream("sess-1", req); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.NOT_FOUND); + assertThat(resp.getHeaders().getFirst(HttpHeaders.CONTENT_TYPE)) + .isEqualTo("text/plain"); + assertThat(resp.getHeaders().getFirst(HttpHeaders.CACHE_CONTROL)).isEqualTo("no-cache"); + } + + @Test + @DisplayName("unmappable upstream status code falls back to 502 Bad Gateway") + void stream_unresolvableStatus_fallsBackToBadGateway() throws Exception { + SecurityContextHolder.clearContext(); + HttpServletRequest req = mock(HttpServletRequest.class); + // 299 is not a defined HttpStatus enum constant → HttpStatus.resolve returns null. + HttpResponse upstream = + upstreamResponse(299, "weird", httpHeaders(Map.of())); + when(proxyService.forward(any(), any(), any(), eq(true))).thenReturn(upstream); + + ResponseEntity resp = controller.stream("sess-1", req); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.BAD_GATEWAY); + } + + @Test + @DisplayName("ownership failure short-circuits before proxying") + void stream_ownershipFailure_doesNotProxy() throws Exception { + HttpServletRequest req = mock(HttpServletRequest.class); + when(sessionService.getSessionForCurrentUser("sess-1")) + .thenThrow(new ResponseStatusException(HttpStatus.NOT_FOUND)); + + assertThatThrownBy(() -> controller.stream("sess-1", req)) + .isInstanceOf(ResponseStatusException.class); + verify(proxyService, never()).forward(any(), any(), any(), anyBoolean()); + } + } + + // ---------------------------------------------------------------------------------------------- + // helpers + // ---------------------------------------------------------------------------------------------- + + private static AiCreateSession session(String sessionId, String userId) { + AiCreateSession s = new AiCreateSession(); + s.setSessionId(sessionId); + s.setUserId(userId); + return s; + } + + private static User userWithTeam(long userId, long teamId) { + User user = new User(); + user.setId(userId); + Team team = new Team(); + team.setId(teamId); + user.setTeam(team); + return user; + } + + /** + * Authenticated WEB principal: 3-arg ctor so isAuthenticated()==true, principal is the User. + */ + private static void authenticateWeb(User user) { + UsernamePasswordAuthenticationToken auth = + new UsernamePasswordAuthenticationToken(user, null, List.of()); + SecurityContextHolder.getContext().setAuthentication(auth); + } + + private static void authenticateApiKey(User user) { + ApiKeyAuthenticationToken auth = + new ApiKeyAuthenticationToken(user, "the-api-key", List.of()); + SecurityContextHolder.getContext().setAuthentication(auth); + } + + private static java.net.http.HttpHeaders httpHeaders(Map single) { + Map> multi = new java.util.HashMap<>(); + single.forEach((k, v) -> multi.put(k, List.of(v))); + return java.net.http.HttpHeaders.of(multi, (k, v) -> true); + } + + @SuppressWarnings("unchecked") + private static HttpResponse upstreamResponse( + int status, String body, java.net.http.HttpHeaders headers) { + HttpResponse response = mock(HttpResponse.class); + when(response.statusCode()).thenReturn(status); + when(response.headers()).thenReturn(headers); + when(response.body()) + .thenReturn(new ByteArrayInputStream(body.getBytes(StandardCharsets.UTF_8))); + return response; + } + + private static String drain(StreamingResponseBody body) throws Exception { + ByteArrayOutputStream out = new ByteArrayOutputStream(); + body.writeTo(out); + return out.toString(StandardCharsets.UTF_8); + } + + private static AiCreateSessionSummaryProjection projection( + String sessionId, + String docType, + String templateId, + String promptLatest, + String promptInitial, + AiCreateSessionStatus status, + String pdfUrl, + Instant createdAt, + Instant updatedAt) { + return new AiCreateSessionSummaryProjection() { + @Override + public String getSessionId() { + return sessionId; + } + + @Override + public String getDocType() { + return docType; + } + + @Override + public String getTemplateId() { + return templateId; + } + + @Override + public String getPromptLatest() { + return promptLatest; + } + + @Override + public String getPromptInitial() { + return promptInitial; + } + + @Override + public AiCreateSessionStatus getStatus() { + return status; + } + + @Override + public String getPdfUrl() { + return pdfUrl; + } + + @Override + public Instant getCreatedAt() { + return createdAt; + } + + @Override + public Instant getUpdatedAt() { + return updatedAt; + } + }; + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/ai/controller/AiCreateInternalControllerTest.java b/app/saas/src/test/java/stirling/software/saas/ai/controller/AiCreateInternalControllerTest.java new file mode 100644 index 0000000000..f5af6ce483 --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/ai/controller/AiCreateInternalControllerTest.java @@ -0,0 +1,457 @@ +package stirling.software.saas.ai.controller; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.ArgumentMatchers.isNull; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.util.List; +import java.util.Map; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.ArgumentCaptor; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.springframework.http.HttpStatus; +import org.springframework.http.ResponseEntity; +import org.springframework.web.server.ResponseStatusException; + +import stirling.software.saas.ai.model.AiCreateSession; +import stirling.software.saas.ai.model.AiCreateSessionStatus; +import stirling.software.saas.ai.service.AiCreateSessionService; + +/** + * Unit tests for {@link AiCreateInternalController}. The controller is a thin internal facade over + * {@link AiCreateSessionService}: it maps an entity to a response record, and on update serialises + * the JSON-shaped fields (outline constraints / draft sections) before delegating. We mock the + * service and assert the {@link ResponseEntity} plus the exact arguments forwarded. + */ +@ExtendWith(MockitoExtension.class) +class AiCreateInternalControllerTest { + + @Mock private AiCreateSessionService sessionService; + + // The controller's @RequiredArgsConstructor only takes sessionService; the ObjectMapper field + // is an inline initializer, so a real Jackson instance is exercised by these tests. + private AiCreateInternalController controller; + + @BeforeEach + void setUp() { + controller = new AiCreateInternalController(sessionService); + } + + // --- helpers ------------------------------------------------------------------------------- + + private static AiCreateSession session(String sessionId) { + AiCreateSession s = new AiCreateSession(); + s.setSessionId(sessionId); + s.setUserId("user-1"); + s.setDocType("report"); + s.setTemplateId("tmpl-1"); + s.setTemplateTex("\\documentclass{article}"); + s.setPreviewTex("preview"); + s.setPromptInitial("initial prompt"); + s.setPromptLatest("latest prompt"); + s.setOutlineText("outline body"); + s.setOutlineFilename("outline.txt"); + s.setOutlineApproved(true); + s.setPolishedLatex("\\section{x}"); + s.setPdfUrl("https://example.com/doc.pdf"); + s.setStatus(AiCreateSessionStatus.DRAFT_READY); + return s; + } + + @Nested + @DisplayName("getSession") + class GetSession { + + @Test + @DisplayName("returns 200 with the session mapped to a response record") + void getSession_mapsAllScalarFields() { + AiCreateSession s = session("sess-abc"); + when(sessionService.getSession("sess-abc")).thenReturn(s); + + ResponseEntity resp = + controller.getSession("sess-abc"); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.OK); + AiCreateController.AiCreateSessionResponse body = resp.getBody(); + assertThat(body).isNotNull(); + assertThat(body.sessionId()).isEqualTo("sess-abc"); + assertThat(body.userId()).isEqualTo("user-1"); + assertThat(body.docType()).isEqualTo("report"); + assertThat(body.templateId()).isEqualTo("tmpl-1"); + assertThat(body.templateTex()).isEqualTo("\\documentclass{article}"); + assertThat(body.previewTex()).isEqualTo("preview"); + assertThat(body.promptInitial()).isEqualTo("initial prompt"); + assertThat(body.promptLatest()).isEqualTo("latest prompt"); + assertThat(body.outlineText()).isEqualTo("outline body"); + assertThat(body.outlineFilename()).isEqualTo("outline.txt"); + assertThat(body.outlineApproved()).isTrue(); + assertThat(body.polishedLatex()).isEqualTo("\\section{x}"); + assertThat(body.pdfUrl()).isEqualTo("https://example.com/doc.pdf"); + assertThat(body.status()).isEqualTo("DRAFT_READY"); + + verify(sessionService).getSession("sess-abc"); + } + + @Test + @DisplayName("maps a null status to a null status string, not an NPE") + void getSession_nullStatus_mapsToNull() { + AiCreateSession s = session("sess-null-status"); + s.setStatus(null); + when(sessionService.getSession("sess-null-status")).thenReturn(s); + + ResponseEntity resp = + controller.getSession("sess-null-status"); + + assertThat(resp.getBody()).isNotNull(); + assertThat(resp.getBody().status()).isNull(); + } + + @Test + @DisplayName("leaves outlineConstraints/draftSections null when the entity stored none") + void getSession_noJsonPayloads_yieldsNullCollections() { + AiCreateSession s = session("sess-empty-json"); + s.setOutlineConstraints(null); + s.setDraftSections(" "); // blank string is treated as absent + when(sessionService.getSession("sess-empty-json")).thenReturn(s); + + AiCreateController.AiCreateSessionResponse body = + controller.getSession("sess-empty-json").getBody(); + + assertThat(body).isNotNull(); + assertThat(body.outlineConstraints()).isNull(); + assertThat(body.draftSections()).isNull(); + } + + @Test + @DisplayName("parses stored outlineConstraints/draftSections JSON back into the response") + void getSession_parsesStoredJsonPayloads() { + AiCreateSession s = session("sess-json"); + s.setOutlineConstraints("{\"tone\":\"formal\",\"pages\":3}"); + s.setDraftSections("[{\"label\":\"Intro\",\"value\":\"hello\"}]"); + when(sessionService.getSession("sess-json")).thenReturn(s); + + AiCreateController.AiCreateSessionResponse body = + controller.getSession("sess-json").getBody(); + + assertThat(body).isNotNull(); + assertThat(body.outlineConstraints()) + .containsEntry("tone", "formal") + .containsEntry("pages", 3); + assertThat(body.draftSections()).hasSize(1); + assertThat(body.draftSections().get(0).label()).isEqualTo("Intro"); + assertThat(body.draftSections().get(0).value()).isEqualTo("hello"); + } + + @Test + @DisplayName("malformed stored JSON is swallowed and surfaces as null, not a 500") + void getSession_malformedJson_returnsNullCollections() { + AiCreateSession s = session("sess-bad-json"); + s.setOutlineConstraints("{not valid json"); + s.setDraftSections("[also not valid"); + when(sessionService.getSession("sess-bad-json")).thenReturn(s); + + AiCreateController.AiCreateSessionResponse body = + controller.getSession("sess-bad-json").getBody(); + + assertThat(body).isNotNull(); + assertThat(body.outlineConstraints()).isNull(); + assertThat(body.draftSections()).isNull(); + } + + @Test + @DisplayName("propagates a 404 ResponseStatusException from the service") + void getSession_notFound_propagates() { + when(sessionService.getSession("missing")) + .thenThrow( + new ResponseStatusException( + HttpStatus.NOT_FOUND, "AI session not found")); + + assertThatThrownBy(() -> controller.getSession("missing")) + .isInstanceOf(ResponseStatusException.class) + .satisfies( + ex -> + assertThat(((ResponseStatusException) ex).getStatusCode()) + .isEqualTo(HttpStatus.NOT_FOUND)); + } + } + + @Nested + @DisplayName("updateSession") + class UpdateSession { + + @Test + @DisplayName("forwards every scalar field and serialises JSON payloads to the service") + void updateSession_serialisesPayloadsAndForwardsAllFields() { + AiCreateSession updated = session("sess-1"); + when(sessionService.applyInternalUpdate( + eq("sess-1"), + any(), + any(), + any(), + any(), + any(), + any(), + any(), + any(), + any(), + any())) + .thenReturn(updated); + + AiCreateInternalController.UpdateSessionRequest req = + new AiCreateInternalController.UpdateSessionRequest( + "new outline", + "new.txt", + Boolean.TRUE, + Map.of("tone", "casual"), + List.of(new AiCreateController.DraftSection("Body", "content")), + "\\section{polished}", + "https://example.com/out.pdf", + "letter", + "tmpl-9", + AiCreateSessionStatus.POLISHED_READY); + + ResponseEntity resp = + controller.updateSession("sess-1", req); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(resp.getBody()).isNotNull(); + assertThat(resp.getBody().sessionId()).isEqualTo("sess-1"); + + ArgumentCaptor constraints = ArgumentCaptor.forClass(String.class); + ArgumentCaptor sections = ArgumentCaptor.forClass(String.class); + verify(sessionService) + .applyInternalUpdate( + eq("sess-1"), + eq("new outline"), + eq("new.txt"), + eq(Boolean.TRUE), + constraints.capture(), + sections.capture(), + eq("\\section{polished}"), + eq("https://example.com/out.pdf"), + eq("letter"), + eq("tmpl-9"), + eq(AiCreateSessionStatus.POLISHED_READY)); + + // Constraints serialised to a JSON object string carrying the map entry. + assertThat(constraints.getValue()).contains("\"tone\"").contains("\"casual\""); + // Draft sections serialised to a JSON array string carrying the record fields. + assertThat(sections.getValue()) + .startsWith("[") + .contains("\"label\":\"Body\"") + .contains("\"value\":\"content\""); + } + + @Test + @DisplayName( + "passes null payloads through when outline constraints / draft sections absent") + void updateSession_nullCollections_forwardsNullPayloads() { + AiCreateSession updated = session("sess-2"); + when(sessionService.applyInternalUpdate( + eq("sess-2"), + isNull(), + isNull(), + isNull(), + isNull(), + isNull(), + isNull(), + isNull(), + isNull(), + isNull(), + isNull())) + .thenReturn(updated); + + AiCreateInternalController.UpdateSessionRequest req = + new AiCreateInternalController.UpdateSessionRequest( + null, null, null, null, null, null, null, null, null, null); + + ResponseEntity resp = + controller.updateSession("sess-2", req); + + assertThat(resp.getBody()).isNotNull(); + // Both JSON-shaped fields forwarded as null (not "null" strings) since they were + // absent. + verify(sessionService) + .applyInternalUpdate( + eq("sess-2"), + isNull(), + isNull(), + isNull(), + isNull(), + isNull(), + isNull(), + isNull(), + isNull(), + isNull(), + isNull()); + } + + @Test + @DisplayName("serialises an empty constraints map / sections list to '{}' and '[]'") + void updateSession_emptyCollections_serialiseToEmptyJson() { + when(sessionService.applyInternalUpdate( + eq("sess-3"), + any(), + any(), + any(), + any(), + any(), + any(), + any(), + any(), + any(), + any())) + .thenReturn(session("sess-3")); + + AiCreateInternalController.UpdateSessionRequest req = + new AiCreateInternalController.UpdateSessionRequest( + null, null, null, Map.of(), List.of(), null, null, null, null, null); + + controller.updateSession("sess-3", req); + + ArgumentCaptor constraints = ArgumentCaptor.forClass(String.class); + ArgumentCaptor sections = ArgumentCaptor.forClass(String.class); + verify(sessionService) + .applyInternalUpdate( + eq("sess-3"), + isNull(), + isNull(), + isNull(), + constraints.capture(), + sections.capture(), + isNull(), + isNull(), + isNull(), + isNull(), + isNull()); + // Empty (but present) collections still serialise: distinguishes "absent" from "empty". + assertThat(constraints.getValue()).isEqualTo("{}"); + assertThat(sections.getValue()).isEqualTo("[]"); + } + + @Test + @DisplayName("response reflects the entity the service returns after the update") + void updateSession_responseReflectsReturnedEntity() { + AiCreateSession returned = session("sess-4"); + returned.setStatus(AiCreateSessionStatus.SAVED); + returned.setPdfUrl("https://example.com/final.pdf"); + when(sessionService.applyInternalUpdate( + eq("sess-4"), + any(), + any(), + any(), + any(), + any(), + any(), + any(), + any(), + any(), + any())) + .thenReturn(returned); + + AiCreateInternalController.UpdateSessionRequest req = + new AiCreateInternalController.UpdateSessionRequest( + "x", null, null, null, null, null, null, null, null, null); + + AiCreateController.AiCreateSessionResponse body = + controller.updateSession("sess-4", req).getBody(); + + assertThat(body).isNotNull(); + assertThat(body.status()).isEqualTo("SAVED"); + assertThat(body.pdfUrl()).isEqualTo("https://example.com/final.pdf"); + } + + @Test + @DisplayName("propagates a 404 when the service cannot find the session to update") + void updateSession_notFound_propagates() { + when(sessionService.applyInternalUpdate( + eq("missing"), + any(), + any(), + any(), + any(), + any(), + any(), + any(), + any(), + any(), + any())) + .thenThrow( + new ResponseStatusException( + HttpStatus.NOT_FOUND, "AI session not found")); + + AiCreateInternalController.UpdateSessionRequest req = + new AiCreateInternalController.UpdateSessionRequest( + "x", null, null, null, null, null, null, null, null, null); + + assertThatThrownBy(() -> controller.updateSession("missing", req)) + .isInstanceOf(ResponseStatusException.class) + .satisfies( + ex -> + assertThat(((ResponseStatusException) ex).getStatusCode()) + .isEqualTo(HttpStatus.NOT_FOUND)); + } + + @Test + @DisplayName("round-trip: serialised draft sections parse back identically in the response") + void updateSession_draftSectionsRoundTrip() { + // Service echoes back the payload it was handed so we can confirm serialise -> store -> + // parse is lossless for the DraftSection shape. + when(sessionService.applyInternalUpdate( + eq("sess-rt"), + any(), + any(), + any(), + any(), + any(), + any(), + any(), + any(), + any(), + any())) + .thenAnswer( + inv -> { + AiCreateSession s = session("sess-rt"); + s.setOutlineConstraints((String) inv.getArgument(4)); + s.setDraftSections((String) inv.getArgument(5)); + return s; + }); + + AiCreateInternalController.UpdateSessionRequest req = + new AiCreateInternalController.UpdateSessionRequest( + null, + null, + null, + Map.of("depth", "deep"), + List.of( + new AiCreateController.DraftSection("A", "1"), + new AiCreateController.DraftSection("B", "2")), + null, + null, + null, + null, + null); + + AiCreateController.AiCreateSessionResponse body = + controller.updateSession("sess-rt", req).getBody(); + + assertThat(body).isNotNull(); + assertThat(body.outlineConstraints()).containsEntry("depth", "deep"); + assertThat(body.draftSections()).hasSize(2); + assertThat(body.draftSections().get(0).label()).isEqualTo("A"); + assertThat(body.draftSections().get(0).value()).isEqualTo("1"); + assertThat(body.draftSections().get(1).label()).isEqualTo("B"); + assertThat(body.draftSections().get(1).value()).isEqualTo("2"); + } + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/ai/controller/AiProxyControllerTest.java b/app/saas/src/test/java/stirling/software/saas/ai/controller/AiProxyControllerTest.java new file mode 100644 index 0000000000..021c04351d --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/ai/controller/AiProxyControllerTest.java @@ -0,0 +1,569 @@ +package stirling.software.saas.ai.controller; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyBoolean; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.io.ByteArrayInputStream; +import java.io.ByteArrayOutputStream; +import java.io.InputStream; +import java.net.http.HttpResponse; +import java.nio.charset.StandardCharsets; +import java.util.List; +import java.util.Map; + +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.http.HttpHeaders; +import org.springframework.http.HttpStatus; +import org.springframework.http.MediaType; +import org.springframework.http.ResponseEntity; +import org.springframework.web.servlet.mvc.method.annotation.StreamingResponseBody; + +import jakarta.servlet.http.HttpServletRequest; + +import stirling.software.saas.ai.service.AiProxyService; + +/** + * Pure unit tests for {@link AiProxyController}. Every collaborator is mocked; each handler is + * invoked directly and asserted via {@link ResponseEntity} / {@code verify}. + * + *

All endpoints funnel through one private {@code proxy(method, path, request, + * acceptEventStream)} helper, so the suite has two halves: + * + *

    + *
  1. per-endpoint tests that pin the exact {@code (method, path, acceptEventStream)} contract a + * given handler forwards (the path-mapping surface), and + *
  2. behavioural tests around the single shared {@code proxy} body: header copy, status + * resolution, and the 503 error fallback. + *
+ */ +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class AiProxyControllerTest { + + @Mock private AiProxyService aiProxyService; + + private AiProxyController controller; + + @org.junit.jupiter.api.BeforeEach + void setUp() { + controller = new AiProxyController(aiProxyService); + } + + // ---------------------------------------------------------------------------------------------- + // Endpoint path/method mapping — each handler pins the exact upstream contract it forwards. + // ---------------------------------------------------------------------------------------------- + + @Nested + @DisplayName("endpoint → upstream (method, path, acceptEventStream) mapping") + class EndpointMapping { + + @Test + @DisplayName("generateSection POSTs to /api/generate_section, non-stream") + void generateSection() throws Exception { + HttpServletRequest req = req(); + stubForward("POST", "/api/generate_section", req, false, ok("body")); + + ResponseEntity resp = controller.generateSection(req); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.OK); + verify(aiProxyService).forward("POST", "/api/generate_section", req, false); + } + + @Test + @DisplayName("generateAllSections POSTs to /api/generate_all_sections") + void generateAllSections() throws Exception { + HttpServletRequest req = req(); + stubForward("POST", "/api/generate_all_sections", req, false, ok("body")); + + controller.generateAllSections(req); + + verify(aiProxyService).forward("POST", "/api/generate_all_sections", req, false); + } + + @Test + @DisplayName("intentCheck POSTs to /api/intent/check") + void intentCheck() throws Exception { + HttpServletRequest req = req(); + stubForward("POST", "/api/intent/check", req, false, ok("body")); + + controller.intentCheck(req); + + verify(aiProxyService).forward("POST", "/api/intent/check", req, false); + } + + @Test + @DisplayName("chatRoute POSTs to /api/chat/route") + void chatRoute() throws Exception { + HttpServletRequest req = req(); + stubForward("POST", "/api/chat/route", req, false, ok("body")); + + controller.chatRoute(req); + + verify(aiProxyService).forward("POST", "/api/chat/route", req, false); + } + + @Test + @DisplayName("createSmartFolder POSTs to /api/chat/create-smart-folder") + void createSmartFolder() throws Exception { + HttpServletRequest req = req(); + stubForward("POST", "/api/chat/create-smart-folder", req, false, ok("body")); + + controller.createSmartFolder(req); + + verify(aiProxyService).forward("POST", "/api/chat/create-smart-folder", req, false); + } + + @Test + @DisplayName("chatInfo POSTs to /api/chat/info") + void chatInfo() throws Exception { + HttpServletRequest req = req(); + stubForward("POST", "/api/chat/info", req, false, ok("body")); + + controller.chatInfo(req); + + verify(aiProxyService).forward("POST", "/api/chat/info", req, false); + } + + @Test + @DisplayName("pdfAnswer POSTs to /api/pdf/answer") + void pdfAnswer() throws Exception { + HttpServletRequest req = req(); + stubForward("POST", "/api/pdf/answer", req, false, ok("body")); + + controller.pdfAnswer(req); + + verify(aiProxyService).forward("POST", "/api/pdf/answer", req, false); + } + + @Test + @DisplayName("progressiveRender POSTs to /api/progressive_render") + void progressiveRender() throws Exception { + HttpServletRequest req = req(); + stubForward("POST", "/api/progressive_render", req, false, ok("body")); + + controller.progressiveRender(req); + + verify(aiProxyService).forward("POST", "/api/progressive_render", req, false); + } + + @Test + @DisplayName("versions GETs /api/versions/{userId} with the path variable interpolated") + void versions() throws Exception { + HttpServletRequest req = req(); + stubForward("GET", "/api/versions/user-42", req, false, ok("body")); + + controller.versions("user-42", req); + + verify(aiProxyService).forward("GET", "/api/versions/user-42", req, false); + } + + @Test + @DisplayName("style (GET) GETs /api/style/{userId}") + void styleGet() throws Exception { + HttpServletRequest req = req(); + stubForward("GET", "/api/style/user-42", req, false, ok("body")); + + controller.style("user-42", req); + + verify(aiProxyService).forward("GET", "/api/style/user-42", req, false); + } + + @Test + @DisplayName("updateStyle POSTs /api/style/{userId}") + void updateStyle() throws Exception { + HttpServletRequest req = req(); + stubForward("POST", "/api/style/user-42", req, false, ok("body")); + + controller.updateStyle("user-42", req); + + verify(aiProxyService).forward("POST", "/api/style/user-42", req, false); + } + + @Test + @DisplayName("importTemplate POSTs /api/import_template") + void importTemplate() throws Exception { + HttpServletRequest req = req(); + stubForward("POST", "/api/import_template", req, false, ok("body")); + + controller.importTemplate(req); + + verify(aiProxyService).forward("POST", "/api/import_template", req, false); + } + + @Test + @DisplayName("createEditSession POSTs /api/edit/sessions") + void createEditSession() throws Exception { + HttpServletRequest req = req(); + stubForward("POST", "/api/edit/sessions", req, false, ok("body")); + + controller.createEditSession(req); + + verify(aiProxyService).forward("POST", "/api/edit/sessions", req, false); + } + + @Test + @DisplayName("editSessionMessage POSTs /api/edit/sessions/{id}/messages") + void editSessionMessage() throws Exception { + HttpServletRequest req = req(); + stubForward("POST", "/api/edit/sessions/sess-9/messages", req, false, ok("body")); + + controller.editSessionMessage("sess-9", req); + + verify(aiProxyService) + .forward("POST", "/api/edit/sessions/sess-9/messages", req, false); + } + + @Test + @DisplayName("editSessionAttachment POSTs /api/edit/sessions/{id}/attachments") + void editSessionAttachment() throws Exception { + HttpServletRequest req = req(); + stubForward("POST", "/api/edit/sessions/sess-9/attachments", req, false, ok("body")); + + controller.editSessionAttachment("sess-9", req); + + verify(aiProxyService) + .forward("POST", "/api/edit/sessions/sess-9/attachments", req, false); + } + + @Test + @DisplayName("runEditSession POSTs /api/edit/sessions/{id}/run as an event stream") + void runEditSession() throws Exception { + HttpServletRequest req = req(); + // acceptEventStream == true here. + stubForward("POST", "/api/edit/sessions/sess-9/run", req, true, ok("data: x\n\n")); + + controller.runEditSession("sess-9", req); + + verify(aiProxyService).forward("POST", "/api/edit/sessions/sess-9/run", req, true); + } + + @Test + @DisplayName("pdfEditorDocument GETs /api/pdf-editor/document") + void pdfEditorDocument() throws Exception { + HttpServletRequest req = req(); + stubForward("GET", "/api/pdf-editor/document", req, false, ok("body")); + + controller.pdfEditorDocument(req); + + verify(aiProxyService).forward("GET", "/api/pdf-editor/document", req, false); + } + + @Test + @DisplayName("pdfEditorUpload POSTs /api/pdf-editor/upload") + void pdfEditorUpload() throws Exception { + HttpServletRequest req = req(); + stubForward("POST", "/api/pdf-editor/upload", req, false, ok("body")); + + controller.pdfEditorUpload(req); + + verify(aiProxyService).forward("POST", "/api/pdf-editor/upload", req, false); + } + } + + // ---------------------------------------------------------------------------------------------- + // output(**) — derives the upstream path from the raw request URI minus the proxy prefix. + // ---------------------------------------------------------------------------------------------- + + @Nested + @DisplayName("output(**) wildcard path derivation") + class OutputPathDerivation { + + @Test + @DisplayName( + "strips the contextPath + /api/v1/ai/output/ prefix and forwards the remainder") + void output_stripsPrefixAndForwardsRemainder() throws Exception { + HttpServletRequest req = mock(HttpServletRequest.class); + when(req.getContextPath()).thenReturn(""); + when(req.getRequestURI()).thenReturn("/api/v1/ai/output/foo/bar.png"); + stubForward("GET", "/output/foo/bar.png", req, false, ok("img-bytes")); + + ResponseEntity resp = controller.output(req); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.OK); + verify(aiProxyService).forward("GET", "/output/foo/bar.png", req, false); + } + + @Test + @DisplayName("honours a non-empty servlet contextPath when computing the prefix") + void output_honoursContextPath() throws Exception { + HttpServletRequest req = mock(HttpServletRequest.class); + when(req.getContextPath()).thenReturn("/stirling"); + when(req.getRequestURI()).thenReturn("/stirling/api/v1/ai/output/nested/file.pdf"); + stubForward("GET", "/output/nested/file.pdf", req, false, ok("pdf")); + + controller.output(req); + + verify(aiProxyService).forward("GET", "/output/nested/file.pdf", req, false); + } + + @Test + @DisplayName("URI not under the expected prefix yields an empty remainder (path /output/)") + void output_uriOutsidePrefix_emptyRemainder() throws Exception { + HttpServletRequest req = mock(HttpServletRequest.class); + when(req.getContextPath()).thenReturn(""); + // Does not start with /api/v1/ai/output/ → substring branch skipped, path stays "". + when(req.getRequestURI()).thenReturn("/totally/different"); + stubForward("GET", "/output/", req, false, ok("body")); + + controller.output(req); + + verify(aiProxyService).forward("GET", "/output/", req, false); + } + } + + // ---------------------------------------------------------------------------------------------- + // Shared proxy body — header copy + status resolution + streaming. + // ---------------------------------------------------------------------------------------------- + + @Nested + @DisplayName("proxy() header copy, status resolution and streaming") + class ProxyBody { + + @Test + @DisplayName("copies the whitelisted upstream headers onto the response") + void copiesWhitelistedHeaders() throws Exception { + HttpServletRequest req = req(); + HttpResponse upstream = + upstreamResponse( + 200, + "payload", + httpHeaders( + Map.of( + HttpHeaders.CONTENT_TYPE, + "application/json", + HttpHeaders.CACHE_CONTROL, + "no-cache", + "X-Accel-Buffering", + "no", + HttpHeaders.CONTENT_DISPOSITION, + "attachment; filename=a.pdf", + HttpHeaders.CONTENT_LENGTH, + "7"))); + when(aiProxyService.forward(any(), any(), any(), anyBoolean())).thenReturn(upstream); + + ResponseEntity resp = controller.generateSection(req); + + HttpHeaders h = resp.getHeaders(); + assertThat(h.getFirst(HttpHeaders.CONTENT_TYPE)).isEqualTo("application/json"); + assertThat(h.getFirst(HttpHeaders.CACHE_CONTROL)).isEqualTo("no-cache"); + assertThat(h.getFirst("X-Accel-Buffering")).isEqualTo("no"); + assertThat(h.getFirst(HttpHeaders.CONTENT_DISPOSITION)) + .isEqualTo("attachment; filename=a.pdf"); + assertThat(h.getFirst(HttpHeaders.CONTENT_LENGTH)).isEqualTo("7"); + } + + @Test + @DisplayName("streams the upstream body straight through to the output stream") + void streamsBodyThrough() throws Exception { + HttpServletRequest req = req(); + stubForward("POST", "/api/generate_section", req, false, ok("hello-stream")); + + ResponseEntity resp = controller.generateSection(req); + + assertThat(drain(resp.getBody())).isEqualTo("hello-stream"); + } + + @Test + @DisplayName("upstream non-2xx status is passed through verbatim") + void passesThroughUpstreamStatus() throws Exception { + HttpServletRequest req = req(); + HttpResponse upstream = + upstreamResponse(404, "nope", httpHeaders(Map.of())); + when(aiProxyService.forward(any(), any(), any(), anyBoolean())).thenReturn(upstream); + + ResponseEntity resp = controller.generateSection(req); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.NOT_FOUND); + } + + @Test + @DisplayName("unmappable upstream status (299) falls back to 502 Bad Gateway") + void unmappableStatus_fallsBackToBadGateway() throws Exception { + HttpServletRequest req = req(); + // 299 is not a defined HttpStatus enum constant → HttpStatus.resolve returns null. + HttpResponse upstream = + upstreamResponse(299, "weird", httpHeaders(Map.of())); + when(aiProxyService.forward(any(), any(), any(), anyBoolean())).thenReturn(upstream); + + ResponseEntity resp = controller.generateSection(req); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.BAD_GATEWAY); + } + + @Test + @DisplayName( + "event-stream endpoint with no upstream Content-Type defaults to text/event-stream") + void eventStreamDefaultsContentType() throws Exception { + HttpServletRequest req = req(); + // runEditSession uses acceptEventStream == true. + HttpResponse upstream = + upstreamResponse(200, "data: x\n\n", httpHeaders(Map.of())); + when(aiProxyService.forward( + eq("POST"), eq("/api/edit/sessions/s/run"), eq(req), eq(true))) + .thenReturn(upstream); + + ResponseEntity resp = controller.runEditSession("s", req); + + assertThat(resp.getHeaders().getFirst(HttpHeaders.CONTENT_TYPE)) + .isEqualTo(MediaType.TEXT_EVENT_STREAM_VALUE); + } + + @Test + @DisplayName("event-stream endpoint keeps an explicit upstream Content-Type (no override)") + void eventStreamKeepsExplicitContentType() throws Exception { + HttpServletRequest req = req(); + HttpResponse upstream = + upstreamResponse( + 200, + "data: x\n\n", + httpHeaders(Map.of(HttpHeaders.CONTENT_TYPE, "text/plain"))); + when(aiProxyService.forward( + eq("POST"), eq("/api/edit/sessions/s/run"), eq(req), eq(true))) + .thenReturn(upstream); + + ResponseEntity resp = controller.runEditSession("s", req); + + assertThat(resp.getHeaders().getFirst(HttpHeaders.CONTENT_TYPE)) + .isEqualTo("text/plain"); + } + + @Test + @DisplayName("non-event-stream endpoint with no upstream Content-Type leaves it unset") + void nonEventStream_noContentType_leavesUnset() throws Exception { + HttpServletRequest req = req(); + stubForward("POST", "/api/generate_section", req, false, ok("body")); + + ResponseEntity resp = controller.generateSection(req); + + assertThat(resp.getHeaders().containsHeader(HttpHeaders.CONTENT_TYPE)).isFalse(); + } + + @Test + @DisplayName("a header carrying CR/LF injection is dropped, not copied") + void crlfInjectionHeaderDropped() throws Exception { + HttpServletRequest req = req(); + HttpResponse upstream = + upstreamResponse( + 200, + "body", + httpHeaders( + Map.of( + HttpHeaders.CACHE_CONTROL, + "no-cache\r\nX-Injected: evil"))); + when(aiProxyService.forward(any(), any(), any(), anyBoolean())).thenReturn(upstream); + + ResponseEntity resp = controller.generateSection(req); + + assertThat(resp.getHeaders().containsHeader(HttpHeaders.CACHE_CONTROL)).isFalse(); + assertThat(resp.getHeaders().containsHeader("X-Injected")).isFalse(); + } + + @Test + @DisplayName("absent upstream headers are simply omitted (no blank values set)") + void absentHeadersOmitted() throws Exception { + HttpServletRequest req = req(); + stubForward("POST", "/api/generate_section", req, false, ok("body")); + + ResponseEntity resp = controller.generateSection(req); + + assertThat(resp.getHeaders().containsHeader(HttpHeaders.CACHE_CONTROL)).isFalse(); + assertThat(resp.getHeaders().containsHeader(HttpHeaders.CONTENT_DISPOSITION)).isFalse(); + assertThat(resp.getHeaders().containsHeader("X-Accel-Buffering")).isFalse(); + } + } + + // ---------------------------------------------------------------------------------------------- + // Error fallback — any forward() failure becomes a 503 with a JSON error body, never throws. + // ---------------------------------------------------------------------------------------------- + + @Nested + @DisplayName("error fallback") + class ErrorFallback { + + @Test + @DisplayName("IOException from forward yields 503 + JSON error body") + void ioException_returns503() throws Exception { + HttpServletRequest req = req(); + when(aiProxyService.forward(any(), any(), any(), anyBoolean())) + .thenThrow(new java.io.IOException("backend down")); + + ResponseEntity resp = controller.generateSection(req); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.SERVICE_UNAVAILABLE); + assertThat(resp.getHeaders().getContentType()).isEqualTo(MediaType.APPLICATION_JSON); + assertThat(drain(resp.getBody())).contains("AI backend unavailable"); + } + + @Test + @DisplayName("InterruptedException from forward also degrades to a 503, never propagates") + void interruptedException_returns503() throws Exception { + HttpServletRequest req = req(); + when(aiProxyService.forward(any(), any(), any(), anyBoolean())) + .thenThrow(new InterruptedException("interrupted")); + + ResponseEntity resp = controller.generateSection(req); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.SERVICE_UNAVAILABLE); + assertThat(drain(resp.getBody())).contains("AI backend unavailable"); + } + } + + // ---------------------------------------------------------------------------------------------- + // helpers + // ---------------------------------------------------------------------------------------------- + + /** A bare mocked request — these handlers never read from it directly (the service does). */ + private static HttpServletRequest req() { + return mock(HttpServletRequest.class); + } + + private void stubForward( + String method, + String path, + HttpServletRequest req, + boolean acceptEventStream, + HttpResponse response) + throws Exception { + when(aiProxyService.forward(eq(method), eq(path), eq(req), eq(acceptEventStream))) + .thenReturn(response); + } + + private static HttpResponse ok(String body) { + return upstreamResponse(200, body, httpHeaders(Map.of())); + } + + private static java.net.http.HttpHeaders httpHeaders(Map single) { + Map> multi = new java.util.HashMap<>(); + single.forEach((k, v) -> multi.put(k, List.of(v))); + return java.net.http.HttpHeaders.of(multi, (k, v) -> true); + } + + @SuppressWarnings("unchecked") + private static HttpResponse upstreamResponse( + int status, String body, java.net.http.HttpHeaders headers) { + HttpResponse response = mock(HttpResponse.class); + when(response.statusCode()).thenReturn(status); + when(response.headers()).thenReturn(headers); + when(response.body()) + .thenReturn(new ByteArrayInputStream(body.getBytes(StandardCharsets.UTF_8))); + return response; + } + + private static String drain(StreamingResponseBody body) throws Exception { + ByteArrayOutputStream out = new ByteArrayOutputStream(); + body.writeTo(out); + return out.toString(StandardCharsets.UTF_8); + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/ai/service/AiCreateProxyServiceTest.java b/app/saas/src/test/java/stirling/software/saas/ai/service/AiCreateProxyServiceTest.java new file mode 100644 index 0000000000..39e949505d --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/ai/service/AiCreateProxyServiceTest.java @@ -0,0 +1,676 @@ +package stirling.software.saas.ai.service; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.verifyNoInteractions; +import static org.mockito.Mockito.when; + +import java.io.ByteArrayInputStream; +import java.io.IOException; +import java.io.InputStream; +import java.io.UncheckedIOException; +import java.net.URI; +import java.net.http.HttpClient; +import java.net.http.HttpRequest; +import java.net.http.HttpResponse; +import java.nio.charset.StandardCharsets; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.ArgumentCaptor; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.test.util.ReflectionTestUtils; + +import jakarta.servlet.ServletInputStream; +import jakarta.servlet.http.HttpServletRequest; + +import stirling.software.proprietary.security.database.repository.UserRepository; +import stirling.software.proprietary.security.service.UserService; + +/** + * Unit tests for {@link AiCreateProxyService}. + * + *

The service forwards an inbound HTTP request to the AI "create" backend. It builds its own + * {@link HttpClient} in the constructor (no injection), so each test swaps in a mocked client via + * {@link ReflectionTestUtils} and captures the outgoing {@link HttpRequest} to assert on the URL, + * method and headers. All collaborators ({@link HttpServletRequest}, {@link UserService}, {@link + * UserRepository}) are mocked; no Spring context, DB or real network is involved. + * + *

Header semantics under test: Content-Type and Authorization are forwarded only when present + * and non-blank; X-API-KEY is taken from the inbound header first and otherwise resolved from the + * authenticated user (any lookup failure is swallowed); Accept is overridden to {@code + * text/event-stream} when SSE is requested. GET/DELETE send no body; other methods stream the + * request input stream lazily. + */ +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class AiCreateProxyServiceTest { + + private static final String BASE_URL = "http://ai-backend:5001"; + + @Mock private UserRepository userRepository; + @Mock private UserService userService; + @Mock private HttpServletRequest request; + @Mock private HttpClient httpClient; + + @SuppressWarnings("unchecked") + private final HttpResponse response = + (HttpResponse) org.mockito.Mockito.mock(HttpResponse.class); + + private AiCreateProxyService service; + + @BeforeEach + void setUp() throws Exception { + service = new AiCreateProxyService(BASE_URL, userRepository, userService); + // Swap the internally-built client for our mock so no real network call happens. + ReflectionTestUtils.setField(service, "httpClient", httpClient); + // Default: the mocked client returns our stub response for any send(). + when(httpClient.send(any(HttpRequest.class), any(HttpResponse.BodyHandler.class))) + .thenReturn(response); + } + + /** Capture the single HttpRequest the service hands to the client. */ + private HttpRequest captureSentRequest() throws Exception { + ArgumentCaptor captor = ArgumentCaptor.forClass(HttpRequest.class); + verify(httpClient).send(captor.capture(), any(HttpResponse.BodyHandler.class)); + return captor.getValue(); + } + + private static String header(HttpRequest req, String name) { + return req.headers().firstValue(name).orElse(null); + } + + /** Build a fresh service backed by the shared mock client for base-URL variations. */ + private AiCreateProxyService serviceWithBase(String base) { + AiCreateProxyService svc = new AiCreateProxyService(base, userRepository, userService); + ReflectionTestUtils.setField(svc, "httpClient", httpClient); + return svc; + } + + @Nested + @DisplayName("target URL assembly") + class UrlAssembly { + + @Test + @DisplayName("joins base + leading-slash path + query string") + void joinsBasePathAndQuery() throws Exception { + when(request.getQueryString()).thenReturn("model=foo&n=2"); + + service.forward("GET", "/v1/chat", request, false); + + assertThat(captureSentRequest().uri()) + .isEqualTo(URI.create("http://ai-backend:5001/v1/chat?model=foo&n=2")); + } + + @Test + @DisplayName("prepends a slash when the path lacks one") + void prependsMissingSlash() throws Exception { + when(request.getQueryString()).thenReturn(null); + + service.forward("GET", "v1/health", request, false); + + assertThat(captureSentRequest().uri()) + .isEqualTo(URI.create("http://ai-backend:5001/v1/health")); + } + + @Test + @DisplayName("null query string is ignored") + void nullQueryIgnored() throws Exception { + when(request.getQueryString()).thenReturn(null); + + service.forward("GET", "/v1/ping", request, false); + + assertThat(captureSentRequest().uri()) + .isEqualTo(URI.create("http://ai-backend:5001/v1/ping")); + } + + @Test + @DisplayName("blank query string is ignored") + void blankQueryIgnored() throws Exception { + when(request.getQueryString()).thenReturn(" "); + + service.forward("GET", "/v1/ping", request, false); + + assertThat(captureSentRequest().uri()) + .isEqualTo(URI.create("http://ai-backend:5001/v1/ping")); + } + + @Test + @DisplayName("trailing slash on the configured base URL is trimmed") + void trimsTrailingSlashOnBase() throws Exception { + AiCreateProxyService svc = serviceWithBase("http://ai-backend:5001/"); + when(request.getQueryString()).thenReturn(null); + + svc.forward("GET", "/v1/x", request, false); + + assertThat(captureSentRequest().uri()) + .isEqualTo(URI.create("http://ai-backend:5001/v1/x")); + } + + @Test + @DisplayName("surrounding whitespace on the base URL is trimmed") + void trimsWhitespaceOnBase() throws Exception { + AiCreateProxyService svc = serviceWithBase(" http://ai-backend:5001 "); + when(request.getQueryString()).thenReturn(null); + + svc.forward("GET", "/v1/y", request, false); + + assertThat(captureSentRequest().uri()) + .isEqualTo(URI.create("http://ai-backend:5001/v1/y")); + } + + @Test + @DisplayName("blank base URL falls back to the localhost default") + void blankBaseUrlFallsBackToDefault() throws Exception { + AiCreateProxyService svc = serviceWithBase(" "); + when(request.getQueryString()).thenReturn(null); + + svc.forward("GET", "/v1/z", request, false); + + assertThat(captureSentRequest().uri()) + .isEqualTo(URI.create("http://localhost:5001/v1/z")); + } + + @Test + @DisplayName("null base URL falls back to the localhost default") + void nullBaseUrlFallsBackToDefault() throws Exception { + AiCreateProxyService svc = serviceWithBase(null); + when(request.getQueryString()).thenReturn(null); + + svc.forward("GET", "/v1/q", request, false); + + assertThat(captureSentRequest().uri()) + .isEqualTo(URI.create("http://localhost:5001/v1/q")); + } + } + + @Nested + @DisplayName("Content-Type header forwarding") + class ContentTypeForwarding { + + @Test + @DisplayName("forwards a present inbound Content-Type on a body-bearing request") + void forwardsPresentContentType() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getContentType()).thenReturn("application/json"); + when(request.getInputStream()) + .thenReturn(servletInputStream("{}".getBytes(StandardCharsets.UTF_8))); + + service.forward("POST", "/v1/chat", request, false); + + assertThat(header(captureSentRequest(), "Content-Type")).isEqualTo("application/json"); + } + + @Test + @DisplayName("omits Content-Type when the inbound value is null") + void omitsWhenNull() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getContentType()).thenReturn(null); + + service.forward("GET", "/v1/x", request, false); + + assertThat(captureSentRequest().headers().firstValue("Content-Type")).isEmpty(); + } + + @Test + @DisplayName("omits Content-Type when the inbound value is blank") + void omitsWhenBlank() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getContentType()).thenReturn(" "); + + service.forward("GET", "/v1/x", request, false); + + assertThat(captureSentRequest().headers().firstValue("Content-Type")).isEmpty(); + } + } + + @Nested + @DisplayName("Authorization header forwarding") + class AuthorizationForwarding { + + @Test + @DisplayName("forwards a present Authorization header verbatim") + void forwardsPresentAuth() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("Authorization")).thenReturn("Bearer abc.def"); + + service.forward("GET", "/v1/x", request, false); + + assertThat(header(captureSentRequest(), "Authorization")).isEqualTo("Bearer abc.def"); + } + + @Test + @DisplayName("omits the header when Authorization is null") + void omitsWhenNull() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("Authorization")).thenReturn(null); + + service.forward("GET", "/v1/x", request, false); + + assertThat(captureSentRequest().headers().firstValue("Authorization")).isEmpty(); + } + + @Test + @DisplayName("omits the header when Authorization is blank") + void omitsWhenBlank() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("Authorization")).thenReturn(" "); + + service.forward("GET", "/v1/x", request, false); + + assertThat(captureSentRequest().headers().firstValue("Authorization")).isEmpty(); + } + } + + @Nested + @DisplayName("X-API-KEY resolution") + class ApiKeyResolution { + + @Test + @DisplayName("uses the X-API-KEY header from the request when present (no user lookup)") + void usesRequestHeaderAndSkipsUserLookup() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("X-API-KEY")).thenReturn("req-key-123"); + + service.forward("GET", "/v1/x", request, false); + + assertThat(header(captureSentRequest(), "X-API-KEY")).isEqualTo("req-key-123"); + // Header short-circuits the authenticated-user fallback entirely. + verifyNoInteractions(userService); + } + + @Test + @DisplayName("falls back to the authenticated user's API key when the header is absent") + void fallsBackToUserApiKey() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("X-API-KEY")).thenReturn(null); + when(userService.getCurrentUsername()).thenReturn("alice"); + when(userService.getApiKeyForUser("alice")).thenReturn("user-key-xyz"); + + service.forward("GET", "/v1/x", request, false); + + assertThat(header(captureSentRequest(), "X-API-KEY")).isEqualTo("user-key-xyz"); + verify(userService).getApiKeyForUser("alice"); + } + + @Test + @DisplayName("falls back to the user key when the inbound header is blank") + void blankHeaderTriggersFallback() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("X-API-KEY")).thenReturn(" "); + when(userService.getCurrentUsername()).thenReturn("bob"); + when(userService.getApiKeyForUser("bob")).thenReturn("bob-key"); + + service.forward("GET", "/v1/x", request, false); + + assertThat(header(captureSentRequest(), "X-API-KEY")).isEqualTo("bob-key"); + } + + @Test + @DisplayName("no X-API-KEY header is set when there is no authenticated user") + void noUser_noHeader() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("X-API-KEY")).thenReturn(null); + when(userService.getCurrentUsername()).thenReturn(null); + + service.forward("GET", "/v1/x", request, false); + + assertThat(captureSentRequest().headers().firstValue("X-API-KEY")).isEmpty(); + // Username was null/blank, so the key lookup is never attempted. + verify(userService, never()).getApiKeyForUser(any()); + } + + @Test + @DisplayName("blank username from the security context yields no header") + void blankUsername_noHeader() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("X-API-KEY")).thenReturn(null); + when(userService.getCurrentUsername()).thenReturn(" "); + + service.forward("GET", "/v1/x", request, false); + + assertThat(captureSentRequest().headers().firstValue("X-API-KEY")).isEmpty(); + verify(userService, never()).getApiKeyForUser(any()); + } + + @Test + @DisplayName("a resolved-but-blank user key is not forwarded") + void blankResolvedKey_noHeader() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("X-API-KEY")).thenReturn(null); + when(userService.getCurrentUsername()).thenReturn("carol"); + when(userService.getApiKeyForUser("carol")).thenReturn(""); + + service.forward("GET", "/v1/x", request, false); + + assertThat(captureSentRequest().headers().firstValue("X-API-KEY")).isEmpty(); + } + + @Test + @DisplayName("a null resolved user key is not forwarded") + void nullResolvedKey_noHeader() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("X-API-KEY")).thenReturn(null); + when(userService.getCurrentUsername()).thenReturn("dan"); + when(userService.getApiKeyForUser("dan")).thenReturn(null); + + service.forward("GET", "/v1/x", request, false); + + assertThat(captureSentRequest().headers().firstValue("X-API-KEY")).isEmpty(); + } + + @Test + @DisplayName("an exception while resolving the user key is swallowed; no header forwarded") + void userKeyLookupThrows_isSwallowed() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("X-API-KEY")).thenReturn(null); + when(userService.getCurrentUsername()).thenReturn("erin"); + when(userService.getApiKeyForUser("erin")) + .thenThrow(new RuntimeException("key store offline")); + + // Must not propagate: extractUserApiKey() catches and returns null. + service.forward("GET", "/v1/x", request, false); + + assertThat(captureSentRequest().headers().firstValue("X-API-KEY")).isEmpty(); + } + } + + @Nested + @DisplayName("Accept header handling") + class AcceptHandling { + + @Test + @DisplayName("acceptEventStream overrides any inbound Accept with text/event-stream") + void eventStreamOverrides() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("Accept")).thenReturn("application/json"); + + service.forward("GET", "/v1/stream", request, true); + + assertThat(header(captureSentRequest(), "Accept")).isEqualTo("text/event-stream"); + } + + @Test + @DisplayName("event stream is requested even with no inbound Accept header") + void eventStreamWithoutInboundAccept() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("Accept")).thenReturn(null); + + service.forward("GET", "/v1/stream", request, true); + + assertThat(header(captureSentRequest(), "Accept")).isEqualTo("text/event-stream"); + } + + @Test + @DisplayName("passes a non-stream Accept header through unchanged") + void passesInboundAcceptThrough() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("Accept")).thenReturn("application/json"); + + service.forward("GET", "/v1/x", request, false); + + assertThat(header(captureSentRequest(), "Accept")).isEqualTo("application/json"); + } + + @Test + @DisplayName("no Accept header set when inbound Accept is absent and SSE not requested") + void noAcceptWhenAbsentAndNotStreaming() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("Accept")).thenReturn(null); + + service.forward("GET", "/v1/x", request, false); + + assertThat(captureSentRequest().headers().firstValue("Accept")).isEmpty(); + } + + @Test + @DisplayName("blank inbound Accept is not forwarded when SSE not requested") + void blankAcceptNotForwarded() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("Accept")).thenReturn(" "); + + service.forward("GET", "/v1/x", request, false); + + assertThat(captureSentRequest().headers().firstValue("Accept")).isEmpty(); + } + } + + @Nested + @DisplayName("HTTP method and body publisher selection") + class MethodAndBody { + + @Test + @DisplayName("GET sends an empty body and never reads the request input stream") + void getHasNoBody() throws Exception { + when(request.getQueryString()).thenReturn(null); + + service.forward("GET", "/v1/x", request, false); + + HttpRequest sent = captureSentRequest(); + assertThat(sent.method()).isEqualTo("GET"); + assertThat(sent.bodyPublisher()).isPresent(); + assertThat(sent.bodyPublisher().get().contentLength()).isZero(); + // GET/DELETE short-circuit before touching the body. + verify(request, never()).getInputStream(); + } + + @Test + @DisplayName("DELETE sends an empty body and never reads the request input stream") + void deleteHasNoBody() throws Exception { + when(request.getQueryString()).thenReturn(null); + + service.forward("DELETE", "/v1/item/9", request, false); + + HttpRequest sent = captureSentRequest(); + assertThat(sent.method()).isEqualTo("DELETE"); + assertThat(sent.bodyPublisher().get().contentLength()).isZero(); + verify(request, never()).getInputStream(); + } + + @Test + @DisplayName("method name matching is case-insensitive for the no-body branch") + void lowercaseGetStillNoBody() throws Exception { + when(request.getQueryString()).thenReturn(null); + + service.forward("get", "/v1/x", request, false); + + HttpRequest sent = captureSentRequest(); + assertThat(sent.method()).isEqualToIgnoringCase("get"); + assertThat(sent.bodyPublisher().get().contentLength()).isZero(); + verify(request, never()).getInputStream(); + } + + @Test + @DisplayName("lowercase delete also routes through the no-body branch") + void lowercaseDeleteNoBody() throws Exception { + when(request.getQueryString()).thenReturn(null); + + service.forward("delete", "/v1/item/1", request, false); + + HttpRequest sent = captureSentRequest(); + assertThat(sent.bodyPublisher().get().contentLength()).isZero(); + verify(request, never()).getInputStream(); + } + + @Test + @DisplayName("POST streams the request input stream as an unknown-length body") + void postStreamsInputStream() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getContentType()).thenReturn("application/json"); + byte[] payload = "{\"x\":1}".getBytes(StandardCharsets.UTF_8); + when(request.getInputStream()).thenReturn(servletInputStream(payload)); + + service.forward("POST", "/v1/chat", request, false); + + HttpRequest sent = captureSentRequest(); + assertThat(sent.method()).isEqualTo("POST"); + assertThat(sent.bodyPublisher()).isPresent(); + // ofInputStream publishes with an unknown content length (-1). + assertThat(sent.bodyPublisher().get().contentLength()).isEqualTo(-1L); + } + + @Test + @DisplayName("PUT also streams the request body via the input-stream publisher") + void putStreamsInputStream() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getContentType()).thenReturn(null); + when(request.getInputStream()) + .thenReturn(servletInputStream("raw".getBytes(StandardCharsets.UTF_8))); + + service.forward("PUT", "/v1/item/3", request, false); + + HttpRequest sent = captureSentRequest(); + assertThat(sent.method()).isEqualTo("PUT"); + assertThat(sent.bodyPublisher().get().contentLength()).isEqualTo(-1L); + } + + @Test + @DisplayName( + "the streamed body publisher lazily emits the exact request bytes when drained") + void streamedBodyContainsRequestBytes() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getContentType()).thenReturn("application/json"); + when(request.getInputStream()) + .thenReturn(servletInputStream("hello-body".getBytes(StandardCharsets.UTF_8))); + + service.forward("POST", "/v1/chat", request, false); + + HttpRequest sent = captureSentRequest(); + // The supplier is lazy: getInputStream() is only invoked once the body is consumed. + String body = drainBody(sent.bodyPublisher().get()); + assertThat(body).isEqualTo("hello-body"); + } + + @Test + @DisplayName( + "an IOException while opening the request stream surfaces as UncheckedIOException") + void inputStreamFailureBecomesUnchecked() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getContentType()).thenReturn("application/json"); + when(request.getInputStream()).thenThrow(new IOException("stream gone")); + + service.forward("POST", "/v1/chat", request, false); + + // The failure only triggers when the lazy supplier runs at body-drain time. + HttpRequest sent = captureSentRequest(); + assertThatThrownBy(() -> drainBody(sent.bodyPublisher().get())) + .isInstanceOf(UncheckedIOException.class) + .hasRootCauseInstanceOf(IOException.class); + } + } + + @Nested + @DisplayName("response propagation and send delegation") + class SendDelegation { + + @Test + @DisplayName("returns exactly the response produced by the underlying client") + void returnsClientResponse() throws Exception { + when(request.getQueryString()).thenReturn(null); + + HttpResponse result = service.forward("GET", "/v1/x", request, false); + + assertThat(result).isSameAs(response); + } + + @Test + @DisplayName("an IOException from the client propagates to the caller") + void clientIoExceptionPropagates() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(httpClient.send(any(HttpRequest.class), any(HttpResponse.BodyHandler.class))) + .thenThrow(new IOException("connection refused")); + + assertThatThrownBy(() -> service.forward("GET", "/v1/x", request, false)) + .isInstanceOf(IOException.class) + .hasMessage("connection refused"); + } + + @Test + @DisplayName("an InterruptedException from the client propagates to the caller") + void clientInterruptedExceptionPropagates() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(httpClient.send(any(HttpRequest.class), any(HttpResponse.BodyHandler.class))) + .thenThrow(new InterruptedException("interrupted")); + + assertThatThrownBy(() -> service.forward("GET", "/v1/x", request, false)) + .isInstanceOf(InterruptedException.class); + // Clear the interrupt flag the thrown InterruptedException may have left. + Thread.interrupted(); + } + } + + // --- helpers ------------------------------------------------------------------------------ + + /** Drain a BodyPublisher to a UTF-8 string, propagating any error the supplier throws. */ + private static String drainBody(HttpRequest.BodyPublisher publisher) { + java.io.ByteArrayOutputStream out = new java.io.ByteArrayOutputStream(); + java.util.concurrent.atomic.AtomicReference error = + new java.util.concurrent.atomic.AtomicReference<>(); + java.util.concurrent.Flow.Subscriber subscriber = + new java.util.concurrent.Flow.Subscriber<>() { + @Override + public void onSubscribe(java.util.concurrent.Flow.Subscription s) { + s.request(Long.MAX_VALUE); + } + + @Override + public void onNext(java.nio.ByteBuffer item) { + byte[] chunk = new byte[item.remaining()]; + item.get(chunk); + out.write(chunk, 0, chunk.length); + } + + @Override + public void onError(Throwable t) { + error.set(t); + } + + @Override + public void onComplete() {} + }; + publisher.subscribe(subscriber); + Throwable t = error.get(); + if (t instanceof RuntimeException re) { + throw re; + } + if (t != null) { + throw new RuntimeException(t); + } + return out.toString(StandardCharsets.UTF_8); + } + + /** Minimal ServletInputStream over a fixed byte array for streaming-body tests. */ + private static ServletInputStream servletInputStream(byte[] data) { + ByteArrayInputStream delegate = new ByteArrayInputStream(data); + return new ServletInputStream() { + @Override + public int read() { + return delegate.read(); + } + + @Override + public boolean isFinished() { + return delegate.available() == 0; + } + + @Override + public boolean isReady() { + return true; + } + + @Override + public void setReadListener(jakarta.servlet.ReadListener readListener) { + // no-op: synchronous reads only in tests + } + }; + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/ai/service/AiCreateSessionServiceTest.java b/app/saas/src/test/java/stirling/software/saas/ai/service/AiCreateSessionServiceTest.java new file mode 100644 index 0000000000..dd410cca59 --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/ai/service/AiCreateSessionServiceTest.java @@ -0,0 +1,783 @@ +package stirling.software.saas.ai.service; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.time.Instant; +import java.util.HashMap; +import java.util.List; +import java.util.Map; +import java.util.Optional; +import java.util.UUID; + +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.ArgumentCaptor; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.data.domain.PageRequest; +import org.springframework.data.domain.Pageable; +import org.springframework.http.HttpStatus; +import org.springframework.mock.web.MockHttpServletRequest; +import org.springframework.mock.web.MockHttpSession; +import org.springframework.security.authentication.UsernamePasswordAuthenticationToken; +import org.springframework.security.core.authority.SimpleGrantedAuthority; +import org.springframework.security.core.context.SecurityContextHolder; +import org.springframework.security.oauth2.jwt.Jwt; +import org.springframework.web.context.request.RequestContextHolder; +import org.springframework.web.context.request.ServletRequestAttributes; +import org.springframework.web.server.ResponseStatusException; + +import stirling.software.common.service.UserServiceInterface; +import stirling.software.saas.ai.model.AiCreateSession; +import stirling.software.saas.ai.model.AiCreateSessionStatus; +import stirling.software.saas.ai.repository.AiCreateSessionRepository; +import stirling.software.saas.security.EnhancedJwtAuthenticationToken; + +/** + * Unit tests for {@link AiCreateSessionService}. + * + *

The service is a thin persistence orchestrator over {@link AiCreateSessionRepository} plus a + * three-tier user-id resolution chain: {@code UserServiceInterface.getCurrentUsername()} -> + * Supabase id from the {@link SecurityContextHolder} authentication -> servlet session-scoped id -> + * the {@code "default_user"} fallback. Repository.save is stubbed to echo its argument so the + * field-mutation assertions can read back what the service set. SecurityContext and + * RequestContextHolder are reset after every test to keep the static thread-local state isolated. + */ +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class AiCreateSessionServiceTest { + + @Mock private AiCreateSessionRepository repository; + @Mock private UserServiceInterface userService; + + private static final String SUPABASE_ID = "11111111-2222-3333-4444-555555555555"; + private static final String DEFAULT_USER_ID = "default_user"; + + @BeforeEach + void echoSave() { + // Persistence is a no-op for these unit tests; save() returns the same managed entity so + // mutation assertions can read it back. + when(repository.save(any(AiCreateSession.class))).thenAnswer(inv -> inv.getArgument(0)); + } + + @AfterEach + void clearStatics() { + SecurityContextHolder.clearContext(); + RequestContextHolder.resetRequestAttributes(); + } + + /** Service with a present (but unstubbed-by-default) UserServiceInterface. */ + private AiCreateSessionService serviceWithUserService() { + return new AiCreateSessionService(repository, Optional.of(userService)); + } + + /** Service with no UserServiceInterface bean wired (Optional.empty). */ + private AiCreateSessionService serviceWithoutUserService() { + return new AiCreateSessionService(repository, Optional.empty()); + } + + /** Authenticate the SecurityContext with a Supabase-id-bearing JWT token. */ + private static void authenticateJwt(String supabaseId) { + Map headers = new HashMap<>(); + headers.put("alg", "RS256"); + Map claims = new HashMap<>(); + claims.put("sub", supabaseId); + claims.put("email", "user@example.com"); + Jwt jwt = new Jwt("token", Instant.now(), Instant.now().plusSeconds(3600), headers, claims); + EnhancedJwtAuthenticationToken auth = + new EnhancedJwtAuthenticationToken( + jwt, + List.of(new SimpleGrantedAuthority("ROLE_USER")), + "user@example.com", + supabaseId); + SecurityContextHolder.getContext().setAuthentication(auth); + } + + /** Bind a servlet request (optionally with a live HttpSession) to the current thread. */ + private static MockHttpServletRequest bindRequest(boolean withSession, String sessionId) { + MockHttpServletRequest request = new MockHttpServletRequest(); + if (withSession) { + // (ServletContext, id) ctor fixes the session id so "session:" is deterministic. + request.setSession(new MockHttpSession(null, sessionId)); + } + RequestContextHolder.setRequestAttributes(new ServletRequestAttributes(request)); + return request; + } + + /** A persisted AiCreateSession owned by the given user. */ + private static AiCreateSession existingSession(String sessionId, String userId) { + AiCreateSession session = new AiCreateSession(); + session.setSessionId(sessionId); + session.setUserId(userId); + session.setStatus(AiCreateSessionStatus.OUTLINE_PENDING); + return session; + } + + // ------------------------------------------------------------------------------------------- + // resolveUserId() — three-tier precedence chain + // ------------------------------------------------------------------------------------------- + + @Nested + @DisplayName("resolveUserId precedence") + class ResolveUserId { + + @Test + @DisplayName("UserServiceInterface username wins over everything else") + void userServiceUsernameWins() { + // Even with a JWT auth present, the username from the user service takes priority. + authenticateJwt(SUPABASE_ID); + when(userService.getCurrentUsername()).thenReturn("alice@corp.com"); + + assertThat(serviceWithUserService().resolveUserId()).isEqualTo("alice@corp.com"); + } + + @Test + @DisplayName("blank username is ignored and the chain falls through to the JWT id") + void blankUsernameFallsThroughToJwt() { + authenticateJwt(SUPABASE_ID); + when(userService.getCurrentUsername()).thenReturn(" "); + + assertThat(serviceWithUserService().resolveUserId()).isEqualTo(SUPABASE_ID); + } + + @Test + @DisplayName("anonymousUser username is ignored and the chain falls through") + void anonymousUsernameFallsThrough() { + authenticateJwt(SUPABASE_ID); + when(userService.getCurrentUsername()).thenReturn("anonymousUser"); + + assertThat(serviceWithUserService().resolveUserId()).isEqualTo(SUPABASE_ID); + } + + @Test + @DisplayName("null username is ignored and the chain falls through") + void nullUsernameFallsThrough() { + authenticateJwt(SUPABASE_ID); + when(userService.getCurrentUsername()).thenReturn(null); + + assertThat(serviceWithUserService().resolveUserId()).isEqualTo(SUPABASE_ID); + } + + @Test + @DisplayName("a throwing user service is swallowed and the chain falls through") + void throwingUserServiceSwallowedAndFallsThrough() { + authenticateJwt(SUPABASE_ID); + when(userService.getCurrentUsername()).thenThrow(new RuntimeException("boom")); + + assertThat(serviceWithUserService().resolveUserId()).isEqualTo(SUPABASE_ID); + } + + @Test + @DisplayName("absent user service bean skips tier 1 and uses the JWT id") + void absentUserServiceUsesJwt() { + authenticateJwt(SUPABASE_ID); + + assertThat(serviceWithoutUserService().resolveUserId()).isEqualTo(SUPABASE_ID); + } + + @Test + @DisplayName("unauthenticated 2-arg token (isAuthenticated=false) is skipped") + void unauthenticatedTokenSkipped() { + // The 2-arg UsernamePasswordAuthenticationToken ctor leaves isAuthenticated()=false, + // so the JWT branch is bypassed and we fall through to the default. + SecurityContextHolder.getContext() + .setAuthentication(new UsernamePasswordAuthenticationToken("bob", "creds")); + + assertThat(serviceWithoutUserService().resolveUserId()).isEqualTo(DEFAULT_USER_ID); + } + + @Test + @DisplayName("authenticated principal id is used via the generic getName() fallback") + void authenticatedPrincipalNameUsed() { + // A non-JWT authenticated token: extractSupabaseId falls back to getName(). + SecurityContextHolder.getContext() + .setAuthentication( + new UsernamePasswordAuthenticationToken( + "carol", + "creds", + List.of(new SimpleGrantedAuthority("ROLE_USER")))); + + assertThat(serviceWithoutUserService().resolveUserId()).isEqualTo("carol"); + } + + @Test + @DisplayName("authenticated 'anonymousUser' name is rejected, chain falls through") + void anonymousAuthNameRejected() { + SecurityContextHolder.getContext() + .setAuthentication( + new UsernamePasswordAuthenticationToken( + "anonymousUser", + "creds", + List.of(new SimpleGrantedAuthority("ROLE_ANONYMOUS")))); + + assertThat(serviceWithoutUserService().resolveUserId()).isEqualTo(DEFAULT_USER_ID); + } + + @Test + @DisplayName("no auth + a live HttpSession yields a session-scoped id") + void sessionScopedIdWhenNoAuth() { + bindRequest(true, "sess-abc"); + + assertThat(serviceWithoutUserService().resolveUserId()).isEqualTo("session:sess-abc"); + } + + @Test + @DisplayName("no auth + a request without a session falls through to default") + void noSessionFallsThroughToDefault() { + bindRequest(false, null); + + assertThat(serviceWithoutUserService().resolveUserId()).isEqualTo(DEFAULT_USER_ID); + } + + @Test + @DisplayName("no user service, no auth, no request context -> default_user") + void defaultUserWhenNothingResolves() { + assertThat(serviceWithoutUserService().resolveUserId()).isEqualTo(DEFAULT_USER_ID); + } + + @Test + @DisplayName("JWT id is preferred over an available session-scoped id") + void jwtPreferredOverSession() { + authenticateJwt(SUPABASE_ID); + bindRequest(true, "sess-xyz"); + + assertThat(serviceWithoutUserService().resolveUserId()).isEqualTo(SUPABASE_ID); + } + } + + // ------------------------------------------------------------------------------------------- + // createSession + // ------------------------------------------------------------------------------------------- + + @Nested + @DisplayName("createSession") + class CreateSession { + + @Test + @DisplayName("populates every field, generates a session id, and persists once") + void populatesAndSaves() { + when(userService.getCurrentUsername()).thenReturn("owner"); + AiCreateSessionService service = serviceWithUserService(); + + AiCreateSession out = + service.createSession( + "my prompt", "report", "tmpl-1", "\\documentclass{}", "preview"); + + assertThat(out.getUserId()).isEqualTo("owner"); + assertThat(out.getDocType()).isEqualTo("report"); + assertThat(out.getTemplateId()).isEqualTo("tmpl-1"); + assertThat(out.getTemplateTex()).isEqualTo("\\documentclass{}"); + assertThat(out.getPreviewTex()).isEqualTo("preview"); + assertThat(out.getPromptInitial()).isEqualTo("my prompt"); + assertThat(out.getPromptLatest()).isEqualTo("my prompt"); + assertThat(out.isOutlineApproved()).isFalse(); + assertThat(out.getStatus()).isEqualTo(AiCreateSessionStatus.OUTLINE_PENDING); + // A random UUID session id was generated. + assertThat(out.getSessionId()).isNotBlank(); + assertThat(UUID.fromString(out.getSessionId())).isNotNull(); + verify(repository).save(out); + } + + @Test + @DisplayName("two sessions get distinct generated ids") + void distinctSessionIds() { + when(userService.getCurrentUsername()).thenReturn("owner"); + AiCreateSessionService service = serviceWithUserService(); + + AiCreateSession a = service.createSession("p", null, null, null, null); + AiCreateSession b = service.createSession("p", null, null, null, null); + + assertThat(a.getSessionId()).isNotEqualTo(b.getSessionId()); + } + + @Test + @DisplayName("uses default_user when nothing else resolves the identity") + void usesDefaultUser() { + AiCreateSession out = + serviceWithoutUserService().createSession("p", "doc", "t", "tex", "prev"); + + assertThat(out.getUserId()).isEqualTo(DEFAULT_USER_ID); + } + } + + // ------------------------------------------------------------------------------------------- + // getSession / getSessionForCurrentUser + // ------------------------------------------------------------------------------------------- + + @Nested + @DisplayName("getSession / getSessionForCurrentUser") + class GetSession { + + @Test + @DisplayName("getSession returns the persisted row") + void getSessionReturnsRow() { + AiCreateSession row = existingSession("s1", "owner"); + when(repository.findById("s1")).thenReturn(Optional.of(row)); + + assertThat(serviceWithoutUserService().getSession("s1")).isSameAs(row); + } + + @Test + @DisplayName("getSession throws 404 when the row is missing") + void getSessionMissingThrows404() { + when(repository.findById("nope")).thenReturn(Optional.empty()); + + assertThatThrownBy(() -> serviceWithoutUserService().getSession("nope")) + .isInstanceOf(ResponseStatusException.class) + .satisfies( + ex -> + assertThat(((ResponseStatusException) ex).getStatusCode()) + .isEqualTo(HttpStatus.NOT_FOUND)); + } + + @Test + @DisplayName("getSessionForCurrentUser returns the row when the owner matches") + void ownerMatchReturnsRow() { + when(userService.getCurrentUsername()).thenReturn("owner"); + AiCreateSession row = existingSession("s1", "owner"); + when(repository.findById("s1")).thenReturn(Optional.of(row)); + + assertThat(serviceWithUserService().getSessionForCurrentUser("s1")).isSameAs(row); + } + + @Test + @DisplayName("getSessionForCurrentUser hides another user's session behind a 404") + void foreignOwnerThrows404() { + when(userService.getCurrentUsername()).thenReturn("intruder"); + AiCreateSession row = existingSession("s1", "owner"); + when(repository.findById("s1")).thenReturn(Optional.of(row)); + + assertThatThrownBy(() -> serviceWithUserService().getSessionForCurrentUser("s1")) + .isInstanceOf(ResponseStatusException.class) + .satisfies( + ex -> + assertThat(((ResponseStatusException) ex).getStatusCode()) + .isEqualTo(HttpStatus.NOT_FOUND)); + } + } + + // ------------------------------------------------------------------------------------------- + // updateOutline + // ------------------------------------------------------------------------------------------- + + @Nested + @DisplayName("updateOutline") + class UpdateOutline { + + @Test + @DisplayName("sets outline text, filename, constraints, approval flag and APPROVED status") + void fullUpdate() { + when(userService.getCurrentUsername()).thenReturn("owner"); + AiCreateSession row = existingSession("s1", "owner"); + when(repository.findById("s1")).thenReturn(Optional.of(row)); + + AiCreateSession out = + serviceWithUserService() + .updateOutline("s1", "the outline", "outline.tex", "be brief"); + + assertThat(out.getOutlineText()).isEqualTo("the outline"); + assertThat(out.getOutlineFilename()).isEqualTo("outline.tex"); + assertThat(out.getOutlineConstraints()).isEqualTo("be brief"); + assertThat(out.isOutlineApproved()).isTrue(); + assertThat(out.getStatus()).isEqualTo(AiCreateSessionStatus.OUTLINE_APPROVED); + verify(repository).save(row); + } + + @Test + @DisplayName("blank filename is not applied; null constraints are left untouched") + void blankFilenameAndNullConstraintsIgnored() { + when(userService.getCurrentUsername()).thenReturn("owner"); + AiCreateSession row = existingSession("s1", "owner"); + row.setOutlineFilename("keep.tex"); + row.setOutlineConstraints("keep-constraints"); + when(repository.findById("s1")).thenReturn(Optional.of(row)); + + AiCreateSession out = serviceWithUserService().updateOutline("s1", "txt", " ", null); + + assertThat(out.getOutlineFilename()).isEqualTo("keep.tex"); + assertThat(out.getOutlineConstraints()).isEqualTo("keep-constraints"); + // Still approved + status flipped even with skipped optional fields. + assertThat(out.isOutlineApproved()).isTrue(); + assertThat(out.getStatus()).isEqualTo(AiCreateSessionStatus.OUTLINE_APPROVED); + } + + @Test + @DisplayName("empty-string constraints ARE applied (only null is skipped)") + void emptyConstraintsApplied() { + when(userService.getCurrentUsername()).thenReturn("owner"); + AiCreateSession row = existingSession("s1", "owner"); + row.setOutlineConstraints("old"); + when(repository.findById("s1")).thenReturn(Optional.of(row)); + + AiCreateSession out = serviceWithUserService().updateOutline("s1", "t", "f", ""); + + assertThat(out.getOutlineConstraints()).isEmpty(); + } + + @Test + @DisplayName("a foreign session 404s before any mutation or save") + void foreignSessionBlocked() { + when(userService.getCurrentUsername()).thenReturn("intruder"); + AiCreateSession row = existingSession("s1", "owner"); + when(repository.findById("s1")).thenReturn(Optional.of(row)); + + assertThatThrownBy(() -> serviceWithUserService().updateOutline("s1", "t", "f", "c")) + .isInstanceOf(ResponseStatusException.class); + assertThat(row.getOutlineText()).isNull(); + verify(repository, never()).save(any()); + } + } + + // ------------------------------------------------------------------------------------------- + // updateDraftSections + // ------------------------------------------------------------------------------------------- + + @Test + @DisplayName("updateDraftSections stores sections and flips status to DRAFT_READY") + void updateDraftSections() { + when(userService.getCurrentUsername()).thenReturn("owner"); + AiCreateSession row = existingSession("s1", "owner"); + when(repository.findById("s1")).thenReturn(Optional.of(row)); + + AiCreateSession out = serviceWithUserService().updateDraftSections("s1", "section json"); + + assertThat(out.getDraftSections()).isEqualTo("section json"); + assertThat(out.getStatus()).isEqualTo(AiCreateSessionStatus.DRAFT_READY); + verify(repository).save(row); + } + + // ------------------------------------------------------------------------------------------- + // updateTemplate + // ------------------------------------------------------------------------------------------- + + @Nested + @DisplayName("updateTemplate") + class UpdateTemplate { + + @Test + @DisplayName("updates docType and templateId when both are non-blank") + void updatesBoth() { + when(userService.getCurrentUsername()).thenReturn("owner"); + AiCreateSession row = existingSession("s1", "owner"); + row.setDocType("old-doc"); + row.setTemplateId("old-tmpl"); + when(repository.findById("s1")).thenReturn(Optional.of(row)); + + AiCreateSession out = + serviceWithUserService().updateTemplate("s1", "new-doc", "new-tmpl"); + + assertThat(out.getDocType()).isEqualTo("new-doc"); + assertThat(out.getTemplateId()).isEqualTo("new-tmpl"); + verify(repository).save(row); + } + + @Test + @DisplayName("null/blank inputs leave the existing template untouched") + void blankInputsKeepExisting() { + when(userService.getCurrentUsername()).thenReturn("owner"); + AiCreateSession row = existingSession("s1", "owner"); + row.setDocType("old-doc"); + row.setTemplateId("old-tmpl"); + when(repository.findById("s1")).thenReturn(Optional.of(row)); + + AiCreateSession out = serviceWithUserService().updateTemplate("s1", null, " "); + + assertThat(out.getDocType()).isEqualTo("old-doc"); + assertThat(out.getTemplateId()).isEqualTo("old-tmpl"); + // Still persists (no-op save) — the method always saves. + verify(repository).save(row); + } + } + + // ------------------------------------------------------------------------------------------- + // reprompt + // ------------------------------------------------------------------------------------------- + + @Test + @DisplayName("reprompt resets all derived artifacts and re-enters OUTLINE_PENDING") + void reprompt() { + when(userService.getCurrentUsername()).thenReturn("owner"); + AiCreateSession row = existingSession("s1", "owner"); + row.setPromptLatest("old prompt"); + row.setOutlineText("old outline"); + row.setOutlineFilename("old.tex"); + row.setOutlineApproved(true); + row.setOutlineConstraints("old constraints"); + row.setDraftSections("old draft"); + row.setPolishedLatex("old latex"); + row.setPdfUrl("https://old/url"); + row.setStatus(AiCreateSessionStatus.POLISHED_READY); + when(repository.findById("s1")).thenReturn(Optional.of(row)); + + AiCreateSession out = serviceWithUserService().reprompt("s1", "fresh prompt"); + + assertThat(out.getPromptLatest()).isEqualTo("fresh prompt"); + assertThat(out.getOutlineText()).isNull(); + assertThat(out.getOutlineFilename()).isNull(); + assertThat(out.isOutlineApproved()).isFalse(); + assertThat(out.getOutlineConstraints()).isNull(); + assertThat(out.getDraftSections()).isNull(); + assertThat(out.getPolishedLatex()).isNull(); + assertThat(out.getPdfUrl()).isNull(); + assertThat(out.getStatus()).isEqualTo(AiCreateSessionStatus.OUTLINE_PENDING); + verify(repository).save(row); + } + + // ------------------------------------------------------------------------------------------- + // deleteSessionForCurrentUser + // ------------------------------------------------------------------------------------------- + + @Nested + @DisplayName("deleteSessionForCurrentUser") + class DeleteSession { + + @Test + @DisplayName("deletes the owner's session") + void deletesOwnerSession() { + when(userService.getCurrentUsername()).thenReturn("owner"); + AiCreateSession row = existingSession("s1", "owner"); + when(repository.findById("s1")).thenReturn(Optional.of(row)); + + serviceWithUserService().deleteSessionForCurrentUser("s1"); + + verify(repository).delete(row); + } + + @Test + @DisplayName("a foreign session 404s and is never deleted") + void foreignSessionNotDeleted() { + when(userService.getCurrentUsername()).thenReturn("intruder"); + AiCreateSession row = existingSession("s1", "owner"); + when(repository.findById("s1")).thenReturn(Optional.of(row)); + + assertThatThrownBy(() -> serviceWithUserService().deleteSessionForCurrentUser("s1")) + .isInstanceOf(ResponseStatusException.class); + verify(repository, never()).delete(any()); + } + } + + // ------------------------------------------------------------------------------------------- + // applyInternalUpdate — null-coalescing partial update, no ownership check + // ------------------------------------------------------------------------------------------- + + @Nested + @DisplayName("applyInternalUpdate") + class ApplyInternalUpdate { + + @Test + @DisplayName("applies every non-null field including outlineApproved=false") + void appliesAllFields() { + // No ownership check on the internal path: it uses getSession, not the per-user guard. + AiCreateSession row = existingSession("s1", "owner"); + row.setOutlineApproved(true); + when(repository.findById("s1")).thenReturn(Optional.of(row)); + + AiCreateSession out = + serviceWithoutUserService() + .applyInternalUpdate( + "s1", + "outline", + "o.tex", + Boolean.FALSE, + "constraints", + "draft", + "latex", + "https://pdf/url", + "doc", + "tmpl", + AiCreateSessionStatus.POLISHED_READY); + + assertThat(out.getOutlineText()).isEqualTo("outline"); + assertThat(out.getOutlineFilename()).isEqualTo("o.tex"); + // Boolean.FALSE is non-null so it IS applied, flipping the prior true. + assertThat(out.isOutlineApproved()).isFalse(); + assertThat(out.getOutlineConstraints()).isEqualTo("constraints"); + assertThat(out.getDraftSections()).isEqualTo("draft"); + assertThat(out.getPolishedLatex()).isEqualTo("latex"); + assertThat(out.getPdfUrl()).isEqualTo("https://pdf/url"); + assertThat(out.getDocType()).isEqualTo("doc"); + assertThat(out.getTemplateId()).isEqualTo("tmpl"); + assertThat(out.getStatus()).isEqualTo(AiCreateSessionStatus.POLISHED_READY); + verify(repository).save(row); + } + + @Test + @DisplayName("all-null arguments leave the row untouched but still persist") + void allNullLeavesUntouched() { + AiCreateSession row = existingSession("s1", "owner"); + row.setOutlineText("keep-outline"); + row.setDocType("keep-doc"); + row.setStatus(AiCreateSessionStatus.DRAFT_READY); + when(repository.findById("s1")).thenReturn(Optional.of(row)); + + AiCreateSession out = + serviceWithoutUserService() + .applyInternalUpdate( + "s1", null, null, null, null, null, null, null, null, null, + null); + + assertThat(out.getOutlineText()).isEqualTo("keep-outline"); + assertThat(out.getDocType()).isEqualTo("keep-doc"); + assertThat(out.getStatus()).isEqualTo(AiCreateSessionStatus.DRAFT_READY); + verify(repository).save(row); + } + + @Test + @DisplayName("internal path 404s when the session does not exist") + void missingSession404() { + when(repository.findById("ghost")).thenReturn(Optional.empty()); + + assertThatThrownBy( + () -> + serviceWithoutUserService() + .applyInternalUpdate( + "ghost", "x", null, null, null, null, null, + null, null, null, null)) + .isInstanceOf(ResponseStatusException.class); + } + + @Test + @DisplayName("internal path ignores ownership — updates a row owned by another user") + void ignoresOwnership() { + // applyInternalUpdate uses getSession (not getSessionForCurrentUser); current identity + // is irrelevant. Confirm a non-matching identity still updates the row. + authenticateJwt(SUPABASE_ID); + AiCreateSession row = existingSession("s1", "someone-else"); + when(repository.findById("s1")).thenReturn(Optional.of(row)); + + AiCreateSession out = + serviceWithoutUserService() + .applyInternalUpdate( + "s1", + null, + null, + null, + null, + null, + null, + "https://done/pdf", + null, + null, + AiCreateSessionStatus.SAVED); + + assertThat(out.getPdfUrl()).isEqualTo("https://done/pdf"); + assertThat(out.getStatus()).isEqualTo(AiCreateSessionStatus.SAVED); + } + } + + // ------------------------------------------------------------------------------------------- + // list* — delegation to the right repository finder for the resolved user + // ------------------------------------------------------------------------------------------- + + @Nested + @DisplayName("listing methods") + class Listing { + + @Test + @DisplayName("no-arg list delegates to findByUserIdOrderByUpdatedAtDesc(userId)") + void listNoArg() { + when(userService.getCurrentUsername()).thenReturn("owner"); + List expected = List.of(existingSession("s1", "owner")); + when(repository.findByUserIdOrderByUpdatedAtDesc("owner")).thenReturn(expected); + + assertThat(serviceWithUserService().listSessionsForCurrentUser()).isSameAs(expected); + } + + @Test + @DisplayName("paged list delegates with the pageable for the resolved user") + void listPaged() { + when(userService.getCurrentUsername()).thenReturn("owner"); + Pageable pageable = PageRequest.of(0, 20); + List expected = List.of(existingSession("s1", "owner")); + when(repository.findByUserIdOrderByUpdatedAtDesc("owner", pageable)) + .thenReturn(expected); + + assertThat(serviceWithUserService().listSessionsForCurrentUser(pageable)) + .isSameAs(expected); + } + + @Test + @DisplayName("includeDrafts=true returns the all-sessions finder") + void listIncludeDraftsTrue() { + when(userService.getCurrentUsername()).thenReturn("owner"); + Pageable pageable = PageRequest.of(0, 10); + List expected = List.of(existingSession("s1", "owner")); + when(repository.findByUserIdOrderByUpdatedAtDesc("owner", pageable)) + .thenReturn(expected); + + assertThat(serviceWithUserService().listSessionsForCurrentUser(pageable, true)) + .isSameAs(expected); + verify(repository, never()) + .findByUserIdAndPdfUrlIsNotNullOrderByUpdatedAtDesc(any(), any()); + } + + @Test + @DisplayName("includeDrafts=false returns only sessions with a non-null pdfUrl") + void listIncludeDraftsFalse() { + when(userService.getCurrentUsername()).thenReturn("owner"); + Pageable pageable = PageRequest.of(0, 10); + List expected = List.of(existingSession("s1", "owner")); + when(repository.findByUserIdAndPdfUrlIsNotNullOrderByUpdatedAtDesc("owner", pageable)) + .thenReturn(expected); + + assertThat(serviceWithUserService().listSessionsForCurrentUser(pageable, false)) + .isSameAs(expected); + verify(repository, never()) + .findByUserIdOrderByUpdatedAtDesc(eq("owner"), any(Pageable.class)); + } + + @Test + @DisplayName("summary list includeDrafts=true uses the all-summaries projection finder") + void summariesIncludeDraftsTrue() { + when(userService.getCurrentUsername()).thenReturn("owner"); + Pageable pageable = PageRequest.of(0, 10); + List expected = List.of(); + when(repository.findSummariesByUserIdOrderByUpdatedAtDesc("owner", pageable)) + .thenReturn(expected); + + assertThat(serviceWithUserService().listSessionSummariesForCurrentUser(pageable, true)) + .isSameAs(expected); + verify(repository, never()) + .findSummariesByUserIdAndPdfUrlIsNotNullOrderByUpdatedAtDesc(any(), any()); + } + + @Test + @DisplayName("summary list includeDrafts=false uses the pdf-only summaries finder") + void summariesIncludeDraftsFalse() { + when(userService.getCurrentUsername()).thenReturn("owner"); + Pageable pageable = PageRequest.of(0, 10); + List expected = List.of(); + when(repository.findSummariesByUserIdAndPdfUrlIsNotNullOrderByUpdatedAtDesc( + "owner", pageable)) + .thenReturn(expected); + + assertThat(serviceWithUserService().listSessionSummariesForCurrentUser(pageable, false)) + .isSameAs(expected); + verify(repository, never()).findSummariesByUserIdOrderByUpdatedAtDesc(any(), any()); + } + + @Test + @DisplayName("listing for an unidentified caller queries the default_user partition") + void listForDefaultUser() { + List expected = List.of(); + when(repository.findByUserIdOrderByUpdatedAtDesc(DEFAULT_USER_ID)).thenReturn(expected); + + assertThat(serviceWithoutUserService().listSessionsForCurrentUser()).isSameAs(expected); + ArgumentCaptor userIdCaptor = ArgumentCaptor.forClass(String.class); + verify(repository).findByUserIdOrderByUpdatedAtDesc(userIdCaptor.capture()); + assertThat(userIdCaptor.getValue()).isEqualTo(DEFAULT_USER_ID); + } + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/ai/service/AiProxyServiceTest.java b/app/saas/src/test/java/stirling/software/saas/ai/service/AiProxyServiceTest.java new file mode 100644 index 0000000000..817958bacd --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/ai/service/AiProxyServiceTest.java @@ -0,0 +1,678 @@ +package stirling.software.saas.ai.service; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.verifyNoInteractions; +import static org.mockito.Mockito.when; + +import java.io.ByteArrayInputStream; +import java.io.IOException; +import java.io.InputStream; +import java.net.URI; +import java.net.http.HttpClient; +import java.net.http.HttpRequest; +import java.net.http.HttpResponse; +import java.nio.charset.StandardCharsets; +import java.util.List; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.ArgumentCaptor; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.test.util.ReflectionTestUtils; + +import jakarta.servlet.ServletInputStream; +import jakarta.servlet.http.HttpServletRequest; +import jakarta.servlet.http.Part; + +import stirling.software.proprietary.security.database.repository.UserRepository; +import stirling.software.proprietary.security.service.UserService; + +/** + * Unit tests for {@link AiProxyService}. + * + *

The service forwards HTTP requests to an AI backend. It constructs its own {@link HttpClient} + * internally (no constructor injection), so each test swaps in a mocked client via {@link + * ReflectionTestUtils} and captures the outgoing {@link HttpRequest} to assert on URL, method and + * headers. All collaborators ({@link HttpServletRequest}, {@link UserService}, {@link + * UserRepository}) are mocked; no Spring context, DB or real network is involved. + * + *

Header semantics under test: Authorization is forwarded when present/non-blank; X-API-KEY is + * taken from the request header first and otherwise resolved from the authenticated user; Accept is + * overridden to {@code text/event-stream} when SSE is requested; the target URL is assembled from + * the configured base URL, the path and the query string. + */ +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class AiProxyServiceTest { + + private static final String BASE_URL = "http://ai-backend:5001"; + + @Mock private UserRepository userRepository; + @Mock private UserService userService; + @Mock private HttpServletRequest request; + @Mock private HttpClient httpClient; + + @SuppressWarnings("unchecked") + private final HttpResponse response = + (HttpResponse) org.mockito.Mockito.mock(HttpResponse.class); + + private AiProxyService service; + + @BeforeEach + void setUp() throws Exception { + service = new AiProxyService(BASE_URL, userRepository, userService); + // Swap the internally-built client for our mock so no real network call happens. + ReflectionTestUtils.setField(service, "httpClient", httpClient); + // Default: the mocked client returns our stub response for any send(). + when(httpClient.send(any(HttpRequest.class), any(HttpResponse.BodyHandler.class))) + .thenReturn(response); + } + + /** Capture the single HttpRequest the service hands to the client. */ + private HttpRequest captureSentRequest() throws Exception { + ArgumentCaptor captor = ArgumentCaptor.forClass(HttpRequest.class); + verify(httpClient).send(captor.capture(), any(HttpResponse.BodyHandler.class)); + return captor.getValue(); + } + + private static String header(HttpRequest req, String name) { + return req.headers().firstValue(name).orElse(null); + } + + @Nested + @DisplayName("target URL assembly") + class UrlAssembly { + + @Test + @DisplayName("joins base + leading-slash path + query string") + void joinsBasePathAndQuery() throws Exception { + when(request.getQueryString()).thenReturn("model=foo&n=2"); + + service.forward("GET", "/v1/chat", request, false); + + HttpRequest sent = captureSentRequest(); + assertThat(sent.uri()) + .isEqualTo(URI.create("http://ai-backend:5001/v1/chat?model=foo&n=2")); + } + + @Test + @DisplayName("prepends a slash when the path lacks one") + void prependsMissingSlash() throws Exception { + when(request.getQueryString()).thenReturn(null); + + service.forward("GET", "v1/health", request, false); + + assertThat(captureSentRequest().uri()) + .isEqualTo(URI.create("http://ai-backend:5001/v1/health")); + } + + @Test + @DisplayName("blank query string is ignored") + void blankQueryIgnored() throws Exception { + when(request.getQueryString()).thenReturn(" "); + + service.forward("GET", "/v1/ping", request, false); + + assertThat(captureSentRequest().uri()) + .isEqualTo(URI.create("http://ai-backend:5001/v1/ping")); + } + + @Test + @DisplayName("trailing slash on the configured base URL is trimmed") + void trimsTrailingSlashOnBase() throws Exception { + AiProxyService svc = + new AiProxyService("http://ai-backend:5001/", userRepository, userService); + ReflectionTestUtils.setField(svc, "httpClient", httpClient); + when(request.getQueryString()).thenReturn(null); + + svc.forward("GET", "/v1/x", request, false); + + assertThat(captureSentRequest().uri()) + .isEqualTo(URI.create("http://ai-backend:5001/v1/x")); + } + + @Test + @DisplayName("surrounding whitespace on the base URL is trimmed") + void trimsWhitespaceOnBase() throws Exception { + AiProxyService svc = + new AiProxyService(" http://ai-backend:5001 ", userRepository, userService); + ReflectionTestUtils.setField(svc, "httpClient", httpClient); + when(request.getQueryString()).thenReturn(null); + + svc.forward("GET", "/v1/y", request, false); + + assertThat(captureSentRequest().uri()) + .isEqualTo(URI.create("http://ai-backend:5001/v1/y")); + } + + @Test + @DisplayName("blank base URL falls back to the localhost default") + void blankBaseUrlFallsBackToDefault() throws Exception { + AiProxyService svc = new AiProxyService(" ", userRepository, userService); + ReflectionTestUtils.setField(svc, "httpClient", httpClient); + when(request.getQueryString()).thenReturn(null); + + svc.forward("GET", "/v1/z", request, false); + + assertThat(captureSentRequest().uri()) + .isEqualTo(URI.create("http://localhost:5001/v1/z")); + } + + @Test + @DisplayName("null base URL falls back to the localhost default") + void nullBaseUrlFallsBackToDefault() throws Exception { + AiProxyService svc = new AiProxyService(null, userRepository, userService); + ReflectionTestUtils.setField(svc, "httpClient", httpClient); + when(request.getQueryString()).thenReturn(null); + + svc.forward("GET", "/v1/q", request, false); + + assertThat(captureSentRequest().uri()) + .isEqualTo(URI.create("http://localhost:5001/v1/q")); + } + } + + @Nested + @DisplayName("Authorization header forwarding") + class AuthorizationForwarding { + + @Test + @DisplayName("forwards a present Authorization header verbatim") + void forwardsPresentAuth() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("Authorization")).thenReturn("Bearer abc.def"); + + service.forward("GET", "/v1/x", request, false); + + assertThat(header(captureSentRequest(), "Authorization")).isEqualTo("Bearer abc.def"); + } + + @Test + @DisplayName("omits the header when Authorization is null") + void omitsWhenNull() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("Authorization")).thenReturn(null); + + service.forward("GET", "/v1/x", request, false); + + assertThat(captureSentRequest().headers().firstValue("Authorization")).isEmpty(); + } + + @Test + @DisplayName("omits the header when Authorization is blank") + void omitsWhenBlank() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("Authorization")).thenReturn(" "); + + service.forward("GET", "/v1/x", request, false); + + assertThat(captureSentRequest().headers().firstValue("Authorization")).isEmpty(); + } + } + + @Nested + @DisplayName("X-API-KEY resolution") + class ApiKeyResolution { + + @Test + @DisplayName("uses the X-API-KEY header from the request when present (no user lookup)") + void usesRequestHeaderAndSkipsUserLookup() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("X-API-KEY")).thenReturn("req-key-123"); + + service.forward("GET", "/v1/x", request, false); + + assertThat(header(captureSentRequest(), "X-API-KEY")).isEqualTo("req-key-123"); + // Header short-circuits the authenticated-user fallback entirely. + verifyNoInteractions(userService); + } + + @Test + @DisplayName("falls back to the authenticated user's API key when the header is absent") + void fallsBackToUserApiKey() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("X-API-KEY")).thenReturn(null); + when(userService.getCurrentUsername()).thenReturn("alice"); + when(userService.getApiKeyForUser("alice")).thenReturn("user-key-xyz"); + + service.forward("GET", "/v1/x", request, false); + + assertThat(header(captureSentRequest(), "X-API-KEY")).isEqualTo("user-key-xyz"); + verify(userService).getApiKeyForUser("alice"); + } + + @Test + @DisplayName("falls back to the user key when the header is blank") + void blankHeaderTriggersFallback() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("X-API-KEY")).thenReturn(" "); + when(userService.getCurrentUsername()).thenReturn("bob"); + when(userService.getApiKeyForUser("bob")).thenReturn("bob-key"); + + service.forward("GET", "/v1/x", request, false); + + assertThat(header(captureSentRequest(), "X-API-KEY")).isEqualTo("bob-key"); + } + + @Test + @DisplayName("no X-API-KEY header set when there is no authenticated user") + void noUser_noHeader() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("X-API-KEY")).thenReturn(null); + when(userService.getCurrentUsername()).thenReturn(null); + + service.forward("GET", "/v1/x", request, false); + + assertThat(captureSentRequest().headers().firstValue("X-API-KEY")).isEmpty(); + // Username was null/blank, so the key lookup is never attempted. + verify(userService, never()).getApiKeyForUser(any()); + } + + @Test + @DisplayName("blank username from the security context yields no header") + void blankUsername_noHeader() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("X-API-KEY")).thenReturn(null); + when(userService.getCurrentUsername()).thenReturn(" "); + + service.forward("GET", "/v1/x", request, false); + + assertThat(captureSentRequest().headers().firstValue("X-API-KEY")).isEmpty(); + verify(userService, never()).getApiKeyForUser(any()); + } + + @Test + @DisplayName("a resolved-but-blank user key is not forwarded") + void blankResolvedKey_noHeader() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("X-API-KEY")).thenReturn(null); + when(userService.getCurrentUsername()).thenReturn("carol"); + when(userService.getApiKeyForUser("carol")).thenReturn(""); + + service.forward("GET", "/v1/x", request, false); + + assertThat(captureSentRequest().headers().firstValue("X-API-KEY")).isEmpty(); + } + + @Test + @DisplayName("an exception while resolving the user key is swallowed; no header forwarded") + void userKeyLookupThrows_isSwallowed() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("X-API-KEY")).thenReturn(null); + when(userService.getCurrentUsername()).thenReturn("dave"); + when(userService.getApiKeyForUser("dave")) + .thenThrow(new RuntimeException("key store offline")); + + // Must not propagate: extractUserApiKey() catches and returns null. + service.forward("GET", "/v1/x", request, false); + + assertThat(captureSentRequest().headers().firstValue("X-API-KEY")).isEmpty(); + } + } + + @Nested + @DisplayName("Accept header handling") + class AcceptHandling { + + @Test + @DisplayName("acceptEventStream overrides any inbound Accept with text/event-stream") + void eventStreamOverrides() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("Accept")).thenReturn("application/json"); + + service.forward("GET", "/v1/stream", request, true); + + assertThat(header(captureSentRequest(), "Accept")).isEqualTo("text/event-stream"); + } + + @Test + @DisplayName("event stream is requested even with no inbound Accept header") + void eventStreamWithoutInboundAccept() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("Accept")).thenReturn(null); + + service.forward("GET", "/v1/stream", request, true); + + assertThat(header(captureSentRequest(), "Accept")).isEqualTo("text/event-stream"); + } + + @Test + @DisplayName("passes a non-stream Accept header through unchanged") + void passesInboundAcceptThrough() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("Accept")).thenReturn("application/json"); + + service.forward("GET", "/v1/x", request, false); + + assertThat(header(captureSentRequest(), "Accept")).isEqualTo("application/json"); + } + + @Test + @DisplayName("no Accept header set when inbound Accept is absent and SSE not requested") + void noAcceptWhenAbsentAndNotStreaming() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("Accept")).thenReturn(null); + + service.forward("GET", "/v1/x", request, false); + + assertThat(captureSentRequest().headers().firstValue("Accept")).isEmpty(); + } + + @Test + @DisplayName("blank inbound Accept is not forwarded when SSE not requested") + void blankAcceptNotForwarded() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getHeader("Accept")).thenReturn(" "); + + service.forward("GET", "/v1/x", request, false); + + assertThat(captureSentRequest().headers().firstValue("Accept")).isEmpty(); + } + } + + @Nested + @DisplayName("HTTP method and body publisher selection") + class MethodAndBody { + + @Test + @DisplayName("GET sends no body and never reads the request input stream") + void getHasNoBody() throws Exception { + when(request.getQueryString()).thenReturn(null); + + service.forward("GET", "/v1/x", request, false); + + HttpRequest sent = captureSentRequest(); + assertThat(sent.method()).isEqualTo("GET"); + assertThat(sent.bodyPublisher()).isPresent(); + assertThat(sent.bodyPublisher().get().contentLength()).isZero(); + // GET/DELETE short-circuit before touching the body. + verify(request, never()).getInputStream(); + verify(request, never()).getParts(); + } + + @Test + @DisplayName("DELETE sends no body and never reads the request input stream") + void deleteHasNoBody() throws Exception { + when(request.getQueryString()).thenReturn(null); + + service.forward("DELETE", "/v1/item/9", request, false); + + HttpRequest sent = captureSentRequest(); + assertThat(sent.method()).isEqualTo("DELETE"); + assertThat(sent.bodyPublisher().get().contentLength()).isZero(); + verify(request, never()).getInputStream(); + } + + @Test + @DisplayName("method name matching is case-insensitive for the no-body branch") + void lowercaseGetStillNoBody() throws Exception { + when(request.getQueryString()).thenReturn(null); + + service.forward("get", "/v1/x", request, false); + + HttpRequest sent = captureSentRequest(); + // Lowercase still routes through the GET/DELETE no-body branch. + assertThat(sent.method()).isEqualToIgnoringCase("get"); + assertThat(sent.bodyPublisher().get().contentLength()).isZero(); + verify(request, never()).getInputStream(); + } + + @Test + @DisplayName("POST with a plain content type streams the request input stream as the body") + void postStreamsInputStream() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getContentType()).thenReturn("application/json"); + ServletInputStream sis = + servletInputStream("{\"x\":1}".getBytes(StandardCharsets.UTF_8)); + when(request.getInputStream()).thenReturn(sis); + + service.forward("POST", "/v1/chat", request, false); + + HttpRequest sent = captureSentRequest(); + assertThat(sent.method()).isEqualTo("POST"); + // ofInputStream publishes with an unknown length (-1). + assertThat(sent.bodyPublisher()).isPresent(); + // Inbound Content-Type is propagated since the body publisher provides none. + assertThat(header(sent, "Content-Type")).isEqualTo("application/json"); + } + + @Test + @DisplayName("POST with no inbound content type sets no Content-Type header") + void postWithoutContentType() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getContentType()).thenReturn(null); + when(request.getInputStream()) + .thenReturn(servletInputStream("raw".getBytes(StandardCharsets.UTF_8))); + + service.forward("POST", "/v1/chat", request, false); + + assertThat(captureSentRequest().headers().firstValue("Content-Type")).isEmpty(); + } + } + + @Nested + @DisplayName("multipart/form-data re-encoding") + class Multipart { + + @Test + @DisplayName("re-encodes parts and sets a generated multipart boundary Content-Type") + void reencodesMultipartBody() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getContentType()).thenReturn("multipart/form-data; boundary=inbound"); + + Part field = textPart("prompt", "hello world"); + Part file = filePart("file", "doc.pdf", "application/pdf", "PDF-BYTES"); + when(request.getParts()).thenReturn(List.of(field, file)); + + service.forward("POST", "/v1/upload", request, false); + + HttpRequest sent = captureSentRequest(); + String contentType = header(sent, "Content-Type"); + assertThat(contentType).startsWith("multipart/form-data; boundary=----spdf-"); + // A fresh boundary is generated rather than reusing the inbound one. + assertThat(contentType).doesNotContain("inbound"); + // Body has a known length (ofByteArray), unlike the streamed-input branch. + assertThat(sent.bodyPublisher()).isPresent(); + assertThat(sent.bodyPublisher().get().contentLength()).isPositive(); + } + + @Test + @DisplayName("the generated boundary in the header matches the one used in the body bytes") + void boundaryHeaderMatchesBody() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getContentType()).thenReturn("multipart/form-data"); + Part textPart = textPart("k", "v"); + when(request.getParts()).thenReturn(List.of(textPart)); + + service.forward("POST", "/v1/upload", request, false); + + HttpRequest sent = captureSentRequest(); + String contentType = header(sent, "Content-Type"); + String boundary = + contentType.substring(contentType.indexOf("boundary=") + "boundary=".length()); + + String body = drainBody(sent.bodyPublisher().get()); + assertThat(body).contains("--" + boundary); + assertThat(body).contains("--" + boundary + "--"); + assertThat(body).contains("Content-Disposition: form-data; name=\"k\"").contains("v"); + } + + @Test + @DisplayName("a file part renders a filename in its Content-Disposition") + void filePartRendersFilename() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getContentType()).thenReturn("multipart/form-data"); + Part filePart = filePart("file", "a.pdf", "application/pdf", "DATA"); + when(request.getParts()).thenReturn(List.of(filePart)); + + service.forward("POST", "/v1/upload", request, false); + + String body = drainBody(captureSentRequest().bodyPublisher().get()); + assertThat(body) + .contains("Content-Disposition: form-data; name=\"file\"; filename=\"a.pdf\"") + .contains("Content-Type: application/pdf") + .contains("DATA"); + } + + @Test + @DisplayName("an empty parts collection still produces a valid closing boundary") + void emptyPartsClosesBoundary() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getContentType()).thenReturn("multipart/form-data"); + when(request.getParts()).thenReturn(List.of()); + + service.forward("POST", "/v1/upload", request, false); + + HttpRequest sent = captureSentRequest(); + String contentType = header(sent, "Content-Type"); + String boundary = + contentType.substring(contentType.indexOf("boundary=") + "boundary=".length()); + // Closing delimiter line + the trailing empty writeLine each append CRLF. + assertThat(drainBody(sent.bodyPublisher().get())) + .isEqualTo("--" + boundary + "--\r\n\r\n"); + } + + @Test + @DisplayName("a getParts() failure is surfaced as IOException") + void getPartsFailureBecomesIoException() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(request.getContentType()).thenReturn("multipart/form-data"); + when(request.getParts()) + .thenThrow(new jakarta.servlet.ServletException("bad multipart")); + + assertThatThrownBy(() -> service.forward("POST", "/v1/upload", request, false)) + .isInstanceOf(IOException.class) + .hasMessageContaining("Failed to proxy multipart request"); + // Failed before reaching the client: send() is never invoked. + verify(httpClient, never()) + .send(any(HttpRequest.class), any(HttpResponse.BodyHandler.class)); + } + } + + @Nested + @DisplayName("response propagation and send delegation") + class SendDelegation { + + @Test + @DisplayName("returns exactly the response produced by the underlying client") + void returnsClientResponse() throws Exception { + when(request.getQueryString()).thenReturn(null); + + HttpResponse result = service.forward("GET", "/v1/x", request, false); + + assertThat(result).isSameAs(response); + } + + @Test + @DisplayName("an IOException from the client propagates to the caller") + void clientIoExceptionPropagates() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(httpClient.send(any(HttpRequest.class), any(HttpResponse.BodyHandler.class))) + .thenThrow(new IOException("connection refused")); + + assertThatThrownBy(() -> service.forward("GET", "/v1/x", request, false)) + .isInstanceOf(IOException.class) + .hasMessage("connection refused"); + } + + @Test + @DisplayName("an InterruptedException from the client propagates to the caller") + void clientInterruptedExceptionPropagates() throws Exception { + when(request.getQueryString()).thenReturn(null); + when(httpClient.send(any(HttpRequest.class), any(HttpResponse.BodyHandler.class))) + .thenThrow(new InterruptedException("interrupted")); + + assertThatThrownBy(() -> service.forward("GET", "/v1/x", request, false)) + .isInstanceOf(InterruptedException.class); + // Clear the interrupt flag the thrown InterruptedException may have left. + Thread.interrupted(); + } + } + + // --- helpers ------------------------------------------------------------------------------ + + private static Part textPart(String name, String value) throws IOException { + Part p = org.mockito.Mockito.mock(Part.class); + when(p.getName()).thenReturn(name); + when(p.getSubmittedFileName()).thenReturn(null); + when(p.getContentType()).thenReturn(null); + when(p.getInputStream()) + .thenReturn(new ByteArrayInputStream(value.getBytes(StandardCharsets.UTF_8))); + return p; + } + + private static Part filePart(String name, String filename, String contentType, String value) + throws IOException { + Part p = org.mockito.Mockito.mock(Part.class); + when(p.getName()).thenReturn(name); + when(p.getSubmittedFileName()).thenReturn(filename); + when(p.getContentType()).thenReturn(contentType); + when(p.getInputStream()) + .thenReturn(new ByteArrayInputStream(value.getBytes(StandardCharsets.UTF_8))); + return p; + } + + private static String drainBody(HttpRequest.BodyPublisher publisher) { + java.io.ByteArrayOutputStream out = new java.io.ByteArrayOutputStream(); + java.util.concurrent.Flow.Subscriber subscriber = + new java.util.concurrent.Flow.Subscriber<>() { + @Override + public void onSubscribe(java.util.concurrent.Flow.Subscription s) { + s.request(Long.MAX_VALUE); + } + + @Override + public void onNext(java.nio.ByteBuffer item) { + byte[] chunk = new byte[item.remaining()]; + item.get(chunk); + out.write(chunk, 0, chunk.length); + } + + @Override + public void onError(Throwable t) { + throw new RuntimeException(t); + } + + @Override + public void onComplete() {} + }; + publisher.subscribe(subscriber); + return out.toString(StandardCharsets.UTF_8); + } + + /** Minimal ServletInputStream over a fixed byte array for streaming-body tests. */ + private static ServletInputStream servletInputStream(byte[] data) { + ByteArrayInputStream delegate = new ByteArrayInputStream(data); + return new ServletInputStream() { + @Override + public int read() { + return delegate.read(); + } + + @Override + public boolean isFinished() { + return delegate.available() == 0; + } + + @Override + public boolean isReady() { + return true; + } + + @Override + public void setReadListener(jakarta.servlet.ReadListener readListener) { + // no-op: synchronous reads only in tests + } + }; + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/controller/CreditControllerApiKeyTest.java b/app/saas/src/test/java/stirling/software/saas/controller/CreditControllerApiKeyTest.java deleted file mode 100644 index 89ec2473d0..0000000000 --- a/app/saas/src/test/java/stirling/software/saas/controller/CreditControllerApiKeyTest.java +++ /dev/null @@ -1,89 +0,0 @@ -package stirling.software.saas.controller; - -import static org.assertj.core.api.Assertions.assertThat; -import static org.mockito.ArgumentMatchers.eq; -import static org.mockito.Mockito.verify; -import static org.mockito.Mockito.when; - -import java.util.UUID; - -import org.junit.jupiter.api.Test; -import org.junit.jupiter.api.extension.ExtendWith; -import org.mockito.Mock; -import org.mockito.junit.jupiter.MockitoExtension; -import org.springframework.http.ResponseEntity; - -import stirling.software.proprietary.security.model.ApiKeyAuthenticationToken; -import stirling.software.proprietary.security.model.User; -import stirling.software.saas.service.CreditService; -import stirling.software.saas.service.CreditService.CreditSummary; - -/** - * Regression coverage for finding #15: API-key users used to always see empty credits because the - * controller blindly passed the API key string through to {@code getCreditSummaryBySupabaseId}, - * which then blew up on {@code UUID.fromString}. The new code reads the User from the principal and - * prefers the linked Supabase ID, falling back to API-key-keyed credits. - */ -@ExtendWith(MockitoExtension.class) -class CreditControllerApiKeyTest { - - @Mock private CreditService creditService; - - @Test - void apiKeyUserWithSupabaseIdGetsResolvedToSupabaseLookup() { - UUID supabaseId = UUID.randomUUID(); - User u = new User(); - u.setSupabaseId(supabaseId); - - CreditSummary expected = creditSummary(42, 100); - when(creditService.getCreditSummaryBySupabaseId(supabaseId.toString())) - .thenReturn(expected); - - CreditController controller = new CreditController(creditService); - ApiKeyAuthenticationToken token = - new ApiKeyAuthenticationToken(u, "the-api-key", java.util.List.of()); - - ResponseEntity resp = controller.getUserCredits(token); - - assertThat(resp.getBody()).isSameAs(expected); - verify(creditService).getCreditSummaryBySupabaseId(supabaseId.toString()); - } - - @Test - void apiKeyUserWithoutSupabaseIdFallsBackToApiKeyLookup() { - User u = new User(); - // No supabaseId set — covers self-hosted / OSS-style API-only users. - CreditSummary expected = creditSummary(7, 25); - when(creditService.getCreditSummaryByApiKey("apikey-no-supabase")).thenReturn(expected); - - CreditController controller = new CreditController(creditService); - ApiKeyAuthenticationToken token = - new ApiKeyAuthenticationToken(u, "apikey-no-supabase", java.util.List.of()); - - ResponseEntity resp = controller.getUserCredits(token); - - assertThat(resp.getBody()).isSameAs(expected); - verify(creditService).getCreditSummaryByApiKey(eq("apikey-no-supabase")); - } - - @Test - void apiKeyTokenWithoutUserPrincipalFallsBackToApiKeyLookup() { - // Edge: token wasn't constructed with a User principal. Should still attempt API-key - // lookup rather than throw. - CreditSummary expected = creditSummary(0, 0); - when(creditService.getCreditSummaryByApiKey("orphan-key")).thenReturn(expected); - - CreditController controller = new CreditController(creditService); - ApiKeyAuthenticationToken token = - new ApiKeyAuthenticationToken("not-a-user", "orphan-key", java.util.List.of()); - - ResponseEntity resp = controller.getUserCredits(token); - - assertThat(resp.getBody()).isNotNull(); - verify(creditService).getCreditSummaryByApiKey("orphan-key"); - } - - private static CreditSummary creditSummary(int remaining, int allocated) { - return new CreditSummary(remaining, allocated, 0, 0, remaining, null, null, false); - } -} diff --git a/app/saas/src/test/java/stirling/software/saas/controller/SaasTeamControllerTest.java b/app/saas/src/test/java/stirling/software/saas/controller/SaasTeamControllerTest.java new file mode 100644 index 0000000000..ea8721a4c3 --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/controller/SaasTeamControllerTest.java @@ -0,0 +1,1355 @@ +package stirling.software.saas.controller; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyLong; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.doThrow; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.time.LocalDateTime; +import java.util.List; +import java.util.Map; +import java.util.Optional; +import java.util.UUID; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.InjectMocks; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.http.HttpStatus; +import org.springframework.http.ResponseEntity; +import org.springframework.transaction.TransactionStatus; +import org.springframework.transaction.interceptor.TransactionAspectSupport; + +import stirling.software.common.model.enumeration.InvitationStatus; +import stirling.software.common.model.enumeration.Role; +import stirling.software.common.model.enumeration.TeamRole; +import stirling.software.proprietary.model.Team; +import stirling.software.proprietary.security.database.repository.UserRepository; +import stirling.software.proprietary.security.model.User; +import stirling.software.proprietary.security.repository.TeamRepository; +import stirling.software.proprietary.security.service.TeamService; +import stirling.software.proprietary.security.service.UserService; +import stirling.software.saas.controller.SaasTeamController.InviteUserRequest; +import stirling.software.saas.controller.SaasTeamController.RenameTeamRequest; +import stirling.software.saas.controller.SaasTeamController.UpdateSeatsRequest; +import stirling.software.saas.model.TeamInvitation; +import stirling.software.saas.model.TeamMembership; +import stirling.software.saas.repository.TeamInvitationRepository; +import stirling.software.saas.repository.TeamMembershipRepository; +import stirling.software.saas.security.TeamSecurityExpressions; +import stirling.software.saas.service.SaasTeamExtensionService; +import stirling.software.saas.service.SaasTeamService; + +/** + * Pure-Mockito unit tests for {@link SaasTeamController}. + * + *

Each handler is invoked directly with mocked collaborators and the returned {@link + * ResponseEntity} (status + body) is asserted, alongside repository/service interaction + * verification. The controller uses {@code @RequiredArgsConstructor}, so {@link InjectMocks} wires + * the mocks by type into the field-injection constructor. + * + *

The {@code getCurrentUser()} helper resolves the principal via {@code + * userService.getCurrentUsername()} then {@code userService.findByUsername(...)}; tests that reach + * a handler body stub that pair. Several handlers are {@code @Transactional} and call {@link + * TransactionAspectSupport#currentTransactionStatus()} on their error paths, which requires a + * thread-bound transaction info; {@code TransactionSupport} installs a no-op one and clears it. + */ +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class SaasTeamControllerTest { + + @Mock private TeamRepository teamRepository; + @Mock private UserRepository userRepository; + @Mock private TeamService teamService; + @Mock private SaasTeamService saasTeamService; + @Mock private SaasTeamExtensionService saasTeamExtensionService; + @Mock private TeamMembershipRepository membershipRepository; + @Mock private TeamInvitationRepository invitationRepository; + @Mock private UserService userService; + @Mock private TeamSecurityExpressions teamSecurityExpressions; + + @InjectMocks private SaasTeamController controller; + + private static final String CURRENT_USERNAME = "alice"; + private static final String CURRENT_EMAIL = "alice@example.com"; + + private User currentUser; + + @BeforeEach + void setUp() { + currentUser = user(7L, CURRENT_USERNAME, CURRENT_EMAIL); + } + + // ===== helpers ===== + + private static User user(Long id, String username, String email) { + User u = new User(); + u.setId(id); + u.setUsername(username); + u.setEmail(email); + return u; + } + + private static Team team(Long id, String name) { + Team t = new Team(); + t.setId(id); + t.setName(name); + return t; + } + + private TeamInvitation invitation( + Long id, Team team, User inviter, String inviteeEmail, InvitationStatus status) { + TeamInvitation inv = new TeamInvitation(); + inv.setInvitationId(id); + inv.setTeam(team); + inv.setInviter(inviter); + inv.setInviteeEmail(inviteeEmail); + inv.setStatus(status); + inv.setInvitationToken("tok-" + id); + inv.setExpiresAt(LocalDateTime.now().plusDays(3)); + return inv; + } + + private TeamMembership membership(Team team, User member, TeamRole role) { + TeamMembership m = new TeamMembership(); + m.setTeam(team); + m.setUser(member); + m.setRole(role); + m.setAcceptedAt(LocalDateTime.now()); + return m; + } + + /** Make {@code getCurrentUser()} resolve to {@link #currentUser}. */ + private void stubCurrentUser() { + when(userService.getCurrentUsername()).thenReturn(CURRENT_USERNAME); + when(userService.findByUsername(CURRENT_USERNAME)).thenReturn(Optional.of(currentUser)); + } + + @SuppressWarnings("unchecked") + private static Map body(ResponseEntity response) { + return (Map) response.getBody(); + } + + @Nested + @DisplayName("inviteUser") + class InviteUser { + + private InviteUserRequest request(Long teamId, String email) { + InviteUserRequest r = new InviteUserRequest(); + r.setTeamId(teamId); + r.setEmail(email); + return r; + } + + @Test + @DisplayName("happy path returns 200 with the invitation DTO") + void happyPath() { + stubCurrentUser(); + when(teamSecurityExpressions.isTeamLeader(10L)).thenReturn(true); + Team team = team(10L, "Acme"); + TeamInvitation inv = + invitation(99L, team, currentUser, "bob@example.com", InvitationStatus.PENDING); + when(saasTeamService.inviteUserToTeam(10L, "bob@example.com", currentUser)) + .thenReturn(inv); + + ResponseEntity response = controller.inviteUser(request(10L, "bob@example.com")); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + SaasTeamController.InvitationDTO dto = + (SaasTeamController.InvitationDTO) response.getBody(); + assertThat(dto.getInvitationId()).isEqualTo(99L); + assertThat(dto.getTeamName()).isEqualTo("Acme"); + assertThat(dto.getInviteeEmail()).isEqualTo("bob@example.com"); + assertThat(dto.getInviterEmail()).isEqualTo(CURRENT_EMAIL); + assertThat(dto.getStatus()).isEqualTo("PENDING"); + } + + @Test + @DisplayName("non-leader is rejected with 403 before any service call") + void nonLeaderForbidden() { + stubCurrentUser(); + when(teamSecurityExpressions.isTeamLeader(10L)).thenReturn(false); + + ResponseEntity response = controller.inviteUser(request(10L, "bob@example.com")); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.FORBIDDEN); + assertThat(body(response)) + .containsEntry("error", "Only team leaders can invite members"); + verify(saasTeamService, never()).inviteUserToTeam(anyLong(), anyString(), any()); + } + + @Test + @DisplayName("IllegalArgumentException from service maps to 400 with its message") + void serviceIllegalArgument_isBadRequest() { + stubCurrentUser(); + when(teamSecurityExpressions.isTeamLeader(10L)).thenReturn(true); + when(saasTeamService.inviteUserToTeam(eq(10L), eq("bob@example.com"), any())) + .thenThrow(new IllegalArgumentException("User is already a team member")); + + ResponseEntity response = controller.inviteUser(request(10L, "bob@example.com")); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + assertThat(body(response)).containsEntry("error", "User is already a team member"); + } + + @Test + @DisplayName("SecurityException from service maps to 400 with its message") + void serviceSecurityException_isBadRequest() { + stubCurrentUser(); + when(teamSecurityExpressions.isTeamLeader(10L)).thenReturn(true); + when(saasTeamService.inviteUserToTeam(eq(10L), eq("bob@example.com"), any())) + .thenThrow(new SecurityException("Only team leaders can invite members")); + + ResponseEntity response = controller.inviteUser(request(10L, "bob@example.com")); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + assertThat(body(response)) + .containsEntry("error", "Only team leaders can invite members"); + } + + @Test + @DisplayName("unexpected RuntimeException maps to 500 with a generic message") + void unexpectedError_isServerError() { + stubCurrentUser(); + when(teamSecurityExpressions.isTeamLeader(10L)).thenReturn(true); + when(saasTeamService.inviteUserToTeam(eq(10L), eq("bob@example.com"), any())) + .thenThrow(new RuntimeException("db down")); + + ResponseEntity response = controller.inviteUser(request(10L, "bob@example.com")); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.INTERNAL_SERVER_ERROR); + assertThat(body(response)).containsEntry("error", "Failed to send invitation"); + } + + @Test + @DisplayName("getCurrentUser failure (user not found) is caught as 400 SecurityException") + void currentUserNotFound_isBadRequest() { + when(userService.getCurrentUsername()).thenReturn(CURRENT_USERNAME); + when(userService.findByUsername(CURRENT_USERNAME)).thenReturn(Optional.empty()); + + ResponseEntity response = controller.inviteUser(request(10L, "bob@example.com")); + + // getCurrentUser throws SecurityException, caught by the (SecurityException|IAE) + // branch. + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + assertThat(body(response)) + .containsEntry("error", "User not found: " + CURRENT_USERNAME); + verify(teamSecurityExpressions, never()).isTeamLeader(anyLong()); + } + } + + @Nested + @DisplayName("acceptInvitation") + class AcceptInvitation { + + @Test + @DisplayName("happy path returns 200 success message") + void happyPath() throws Exception { + stubCurrentUser(); + + ResponseEntity response = controller.acceptInvitation("tok-1"); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(body(response)).containsEntry("message", "Invitation accepted"); + assertThat(body(response)).containsEntry("success", true); + verify(saasTeamService).acceptInvitationAndGrantRole("tok-1", currentUser); + } + + @Test + @DisplayName("IllegalStateException (expired/already-accepted) maps to 400 and rolls back") + void callerFixableFailure_isBadRequest() throws Exception { + stubCurrentUser(); + TransactionSupport tx = TransactionSupport.bind(); + try { + doThrow(new IllegalStateException("Invitation expired")) + .when(saasTeamService) + .acceptInvitationAndGrantRole("tok-1", currentUser); + + ResponseEntity response = controller.acceptInvitation("tok-1"); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + assertThat(body(response)).containsEntry("error", "Invitation expired"); + verify(tx.status()).setRollbackOnly(); + } finally { + tx.unbind(); + } + } + + @Test + @DisplayName("unexpected error maps to 500 and rolls back") + void unexpectedError_isServerError() throws Exception { + stubCurrentUser(); + TransactionSupport tx = TransactionSupport.bind(); + try { + doThrow(new RuntimeException("boom")) + .when(saasTeamService) + .acceptInvitationAndGrantRole("tok-1", currentUser); + + ResponseEntity response = controller.acceptInvitation("tok-1"); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.INTERNAL_SERVER_ERROR); + assertThat(body(response)).containsEntry("error", "Failed to accept invitation"); + verify(tx.status()).setRollbackOnly(); + } finally { + tx.unbind(); + } + } + } + + @Nested + @DisplayName("rejectInvitation") + class RejectInvitation { + + @Test + @DisplayName( + "happy path: pending invitation for the current user is set REJECTED and saved") + void happyPath() { + stubCurrentUser(); + Team team = team(10L, "Acme"); + TeamInvitation inv = + invitation( + 5L, + team, + user(2L, "leader", "lead@x.com"), + CURRENT_EMAIL, + InvitationStatus.PENDING); + when(invitationRepository.findByInvitationToken("tok-5")).thenReturn(Optional.of(inv)); + + ResponseEntity response = controller.rejectInvitation("tok-5"); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(body(response)).containsEntry("message", "Invitation rejected"); + assertThat(inv.getStatus()).isEqualTo(InvitationStatus.REJECTED); + verify(invitationRepository).save(inv); + } + + @Test + @DisplayName("invitation matched by username (not email) is also accepted") + void matchedByUsername() { + stubCurrentUser(); + TeamInvitation inv = + invitation( + 6L, + team(10L, "Acme"), + user(2L, "leader", "lead@x.com"), + CURRENT_USERNAME, + InvitationStatus.PENDING); + when(invitationRepository.findByInvitationToken("tok-6")).thenReturn(Optional.of(inv)); + + ResponseEntity response = controller.rejectInvitation("tok-6"); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(inv.getStatus()).isEqualTo(InvitationStatus.REJECTED); + } + + @Test + @DisplayName("missing invitation maps to 404") + void notFound() { + stubCurrentUser(); + when(invitationRepository.findByInvitationToken("nope")).thenReturn(Optional.empty()); + + ResponseEntity response = controller.rejectInvitation("nope"); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.NOT_FOUND); + assertThat(body(response)).containsEntry("error", "Invitation not found"); + verify(invitationRepository, never()).save(any()); + } + + @Test + @DisplayName("invitation addressed to someone else maps to 403 (security)") + void wrongRecipient_forbidden() { + stubCurrentUser(); + TeamInvitation inv = + invitation( + 7L, + team(10L, "Acme"), + user(2L, "leader", "lead@x.com"), + "someone-else@x.com", + InvitationStatus.PENDING); + when(invitationRepository.findByInvitationToken("tok-7")).thenReturn(Optional.of(inv)); + + ResponseEntity response = controller.rejectInvitation("tok-7"); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.FORBIDDEN); + assertThat(body(response)) + .containsEntry( + "error", "You cannot reject an invitation that was not sent to you"); + verify(invitationRepository, never()).save(any()); + } + + @Test + @DisplayName("non-pending invitation maps to 403 (illegal state)") + void nonPending_forbidden() { + stubCurrentUser(); + TeamInvitation inv = + invitation( + 8L, + team(10L, "Acme"), + user(2L, "leader", "lead@x.com"), + CURRENT_EMAIL, + InvitationStatus.ACCEPTED); + when(invitationRepository.findByInvitationToken("tok-8")).thenReturn(Optional.of(inv)); + + ResponseEntity response = controller.rejectInvitation("tok-8"); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.FORBIDDEN); + assertThat(body(response)) + .containsEntry("error", "Can only reject pending invitations"); + verify(invitationRepository, never()).save(any()); + } + } + + @Nested + @DisplayName("cancelInvitation") + class CancelInvitation { + + @Test + @DisplayName("leader cancels a pending invitation -> 200 and status CANCELLED") + void happyPath() { + stubCurrentUser(); + Team team = team(10L, "Acme"); + TeamInvitation inv = + invitation(11L, team, currentUser, "bob@x.com", InvitationStatus.PENDING); + when(invitationRepository.findById(11L)).thenReturn(Optional.of(inv)); + TeamMembership leaderMembership = membership(team, currentUser, TeamRole.LEADER); + when(membershipRepository.findByTeamIdAndUserId(10L, currentUser.getId())) + .thenReturn(Optional.of(leaderMembership)); + + ResponseEntity response = controller.cancelInvitation(11L); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(body(response)).containsEntry("message", "Invitation cancelled"); + assertThat(inv.getStatus()).isEqualTo(InvitationStatus.CANCELLED); + verify(invitationRepository).save(inv); + } + + @Test + @DisplayName("missing invitation -> 404") + void notFound() { + stubCurrentUser(); + when(invitationRepository.findById(11L)).thenReturn(Optional.empty()); + + ResponseEntity response = controller.cancelInvitation(11L); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.NOT_FOUND); + assertThat(body(response)).containsEntry("error", "Invitation not found"); + } + + @Test + @DisplayName("caller not a member of the team -> 403") + void notAMember_forbidden() { + stubCurrentUser(); + Team team = team(10L, "Acme"); + TeamInvitation inv = + invitation(11L, team, currentUser, "bob@x.com", InvitationStatus.PENDING); + when(invitationRepository.findById(11L)).thenReturn(Optional.of(inv)); + when(membershipRepository.findByTeamIdAndUserId(10L, currentUser.getId())) + .thenReturn(Optional.empty()); + + ResponseEntity response = controller.cancelInvitation(11L); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.FORBIDDEN); + assertThat(body(response)).containsEntry("error", "You are not a member of this team"); + verify(invitationRepository, never()).save(any()); + } + + @Test + @DisplayName("member but not leader -> 403") + void memberNotLeader_forbidden() { + stubCurrentUser(); + Team team = team(10L, "Acme"); + TeamInvitation inv = + invitation(11L, team, currentUser, "bob@x.com", InvitationStatus.PENDING); + when(invitationRepository.findById(11L)).thenReturn(Optional.of(inv)); + when(membershipRepository.findByTeamIdAndUserId(10L, currentUser.getId())) + .thenReturn(Optional.of(membership(team, currentUser, TeamRole.MEMBER))); + + ResponseEntity response = controller.cancelInvitation(11L); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.FORBIDDEN); + assertThat(body(response)) + .containsEntry("error", "Only team leaders can cancel invitations"); + verify(invitationRepository, never()).save(any()); + } + + @Test + @DisplayName("non-pending invitation -> 403 (illegal state)") + void nonPending_forbidden() { + stubCurrentUser(); + Team team = team(10L, "Acme"); + TeamInvitation inv = + invitation(11L, team, currentUser, "bob@x.com", InvitationStatus.CANCELLED); + when(invitationRepository.findById(11L)).thenReturn(Optional.of(inv)); + when(membershipRepository.findByTeamIdAndUserId(10L, currentUser.getId())) + .thenReturn(Optional.of(membership(team, currentUser, TeamRole.LEADER))); + + ResponseEntity response = controller.cancelInvitation(11L); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.FORBIDDEN); + assertThat(body(response)) + .containsEntry("error", "Can only cancel pending invitations"); + verify(invitationRepository, never()).save(any()); + } + } + + @Nested + @DisplayName("getPendingInvitations") + class GetPendingInvitations { + + @Test + @DisplayName("returns DTOs for the current user's pending invitations") + void happyPath() { + stubCurrentUser(); + Team team = team(10L, "Acme"); + TeamInvitation inv = + invitation( + 20L, + team, + user(2L, "leader", "lead@x.com"), + CURRENT_EMAIL, + InvitationStatus.PENDING); + when(invitationRepository.findPendingInvitationsByEmail(eq(CURRENT_EMAIL), any())) + .thenReturn(List.of(inv)); + + ResponseEntity response = controller.getPendingInvitations(); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + @SuppressWarnings("unchecked") + List dtos = + (List) response.getBody(); + assertThat(dtos).hasSize(1); + assertThat(dtos.get(0).getInvitationId()).isEqualTo(20L); + assertThat(dtos.get(0).getInviteeEmail()).isEqualTo(CURRENT_EMAIL); + } + + @Test + @DisplayName("empty list returns 200 with an empty body") + void empty() { + stubCurrentUser(); + when(invitationRepository.findPendingInvitationsByEmail(eq(CURRENT_EMAIL), any())) + .thenReturn(List.of()); + + ResponseEntity response = controller.getPendingInvitations(); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat((List) response.getBody()).isEmpty(); + } + + @Test + @DisplayName("repository failure maps to 500") + void repoFailure_isServerError() { + stubCurrentUser(); + when(invitationRepository.findPendingInvitationsByEmail(anyString(), any())) + .thenThrow(new RuntimeException("db down")); + + ResponseEntity response = controller.getPendingInvitations(); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.INTERNAL_SERVER_ERROR); + assertThat(body(response)).containsEntry("error", "Failed to fetch invitations"); + } + } + + @Nested + @DisplayName("getMyTeams") + class GetMyTeams { + + @Test + @DisplayName( + "no memberships -> personal team created, then memberships re-fetched and returned") + void noMemberships_createsPersonalTeam() { + stubCurrentUser(); + Team personal = team(1L, "My Team"); + when(membershipRepository.findByUserId(currentUser.getId())) + .thenReturn(List.of()) // first call: empty + .thenReturn(List.of(membership(personal, currentUser, TeamRole.LEADER))); + when(saasTeamExtensionService.isPersonal(personal)).thenReturn(true); + when(saasTeamExtensionService.getTeamType(personal)).thenReturn("PERSONAL"); + when(membershipRepository.countByTeamId(1L)).thenReturn(1L); + when(saasTeamExtensionService.getMaxSeats(personal)).thenReturn(1); + when(saasTeamExtensionService.getSeatsUsed(personal)).thenReturn(1); + + ResponseEntity response = controller.getMyTeams(); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + verify(saasTeamService).createPersonalTeam(currentUser); + @SuppressWarnings("unchecked") + List dtos = + (List) response.getBody(); + assertThat(dtos).hasSize(1); + assertThat(dtos.get(0).getTeamId()).isEqualTo(1L); + assertThat(dtos.get(0).getIsPersonal()).isTrue(); + assertThat(dtos.get(0).getIsLeader()).isTrue(); + assertThat(dtos.get(0).getMemberCount()).isEqualTo(1); + assertThat(dtos.get(0).getMaxSeats()).isEqualTo(1); + } + + @Test + @DisplayName("already has a personal team -> no migration, returns existing teams") + void existingPersonalTeam_noMigration() { + stubCurrentUser(); + Team personal = team(1L, "My Team"); + when(membershipRepository.findByUserId(currentUser.getId())) + .thenReturn(List.of(membership(personal, currentUser, TeamRole.LEADER))); + when(saasTeamExtensionService.isPersonal(personal)).thenReturn(true); + when(saasTeamExtensionService.getTeamType(personal)).thenReturn("PERSONAL"); + when(membershipRepository.countByTeamId(1L)).thenReturn(1L); + when(saasTeamExtensionService.getMaxSeats(personal)).thenReturn(1); + when(saasTeamExtensionService.getSeatsUsed(personal)).thenReturn(1); + + ResponseEntity response = controller.getMyTeams(); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + verify(saasTeamService, never()).createPersonalTeam(any()); + } + + @Test + @DisplayName("only on legacy Default team -> migrates to a personal team") + void onlyOnDefaultTeam_migrates() { + stubCurrentUser(); + Team legacy = team(2L, "Default"); + Team personal = team(1L, "My Team"); + when(membershipRepository.findByUserId(currentUser.getId())) + .thenReturn(List.of(membership(legacy, currentUser, TeamRole.MEMBER))) + .thenReturn(List.of(membership(personal, currentUser, TeamRole.LEADER))); + // Team equals() (Lombok onlyExplicitlyIncluded with no fields) treats all Team + // instances as equal, so isPersonal cannot be stubbed per-instance. The legacy team + // being non-personal plus the "Default" name is what triggers migration here. + when(saasTeamExtensionService.isPersonal(any(Team.class))).thenReturn(false); + when(saasTeamExtensionService.getTeamType(any(Team.class))).thenReturn("STANDARD"); + when(membershipRepository.countByTeamId(anyLong())).thenReturn(1L); + when(saasTeamExtensionService.getMaxSeats(any(Team.class))).thenReturn(1); + when(saasTeamExtensionService.getSeatsUsed(any(Team.class))).thenReturn(1); + + ResponseEntity response = controller.getMyTeams(); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + verify(saasTeamService).createPersonalTeam(currentUser); + } + + @Test + @DisplayName("on a real (non-system, non-personal) team -> no migration") + void onRealTeam_noMigration() { + stubCurrentUser(); + Team realTeam = team(3L, "Engineering"); + when(membershipRepository.findByUserId(currentUser.getId())) + .thenReturn(List.of(membership(realTeam, currentUser, TeamRole.MEMBER))); + when(saasTeamExtensionService.isPersonal(realTeam)).thenReturn(false); + when(saasTeamExtensionService.getTeamType(realTeam)).thenReturn("STANDARD"); + when(membershipRepository.countByTeamId(3L)).thenReturn(4L); + when(saasTeamExtensionService.getMaxSeats(realTeam)).thenReturn(10); + when(saasTeamExtensionService.getSeatsUsed(realTeam)).thenReturn(4); + + ResponseEntity response = controller.getMyTeams(); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + verify(saasTeamService, never()).createPersonalTeam(any()); + @SuppressWarnings("unchecked") + List dtos = + (List) response.getBody(); + assertThat(dtos.get(0).getIsLeader()).isFalse(); + assertThat(dtos.get(0).getMemberCount()).isEqualTo(4); + assertThat(dtos.get(0).getMaxSeats()).isEqualTo(10); + assertThat(dtos.get(0).getSeatsUsed()).isEqualTo(4); + } + + @Test + @DisplayName("personal-team creation failure surfaces as 500") + void createPersonalTeamFails_isServerError() { + stubCurrentUser(); + when(membershipRepository.findByUserId(currentUser.getId())).thenReturn(List.of()); + when(saasTeamService.createPersonalTeam(currentUser)) + .thenThrow(new RuntimeException("insert failed")); + + ResponseEntity response = controller.getMyTeams(); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.INTERNAL_SERVER_ERROR); + assertThat(body(response)).containsEntry("error", "Failed to fetch teams"); + } + } + + @Nested + @DisplayName("getTeamMembers") + class GetTeamMembers { + + @Test + @DisplayName("returns member DTOs for the team") + void happyPath() { + Team team = team(10L, "Acme"); + User bob = user(2L, "bob", "bob@x.com"); + when(membershipRepository.findByTeamId(10L)) + .thenReturn(List.of(membership(team, bob, TeamRole.MEMBER))); + + ResponseEntity response = controller.getTeamMembers(10L); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + @SuppressWarnings("unchecked") + List dtos = + (List) response.getBody(); + assertThat(dtos).hasSize(1); + assertThat(dtos.get(0).getId()).isEqualTo(2L); + assertThat(dtos.get(0).getUsername()).isEqualTo("bob"); + assertThat(dtos.get(0).getRole()).isEqualTo("MEMBER"); + } + + @Test + @DisplayName("repository failure maps to 500") + void repoFailure_isServerError() { + when(membershipRepository.findByTeamId(10L)).thenThrow(new RuntimeException("db")); + + ResponseEntity response = controller.getTeamMembers(10L); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.INTERNAL_SERVER_ERROR); + assertThat(body(response)).containsEntry("error", "Failed to fetch team members"); + } + } + + @Nested + @DisplayName("getTeamInvitations") + class GetTeamInvitations { + + @Test + @DisplayName("returns invitation DTOs for the team") + void happyPath() { + Team team = team(10L, "Acme"); + TeamInvitation inv = + invitation(30L, team, currentUser, "bob@x.com", InvitationStatus.PENDING); + when(invitationRepository.findByTeamId(10L)).thenReturn(List.of(inv)); + + ResponseEntity response = controller.getTeamInvitations(10L); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + @SuppressWarnings("unchecked") + List dtos = + (List) response.getBody(); + assertThat(dtos).hasSize(1); + assertThat(dtos.get(0).getInvitationId()).isEqualTo(30L); + } + + @Test + @DisplayName("repository failure maps to 500") + void repoFailure_isServerError() { + when(invitationRepository.findByTeamId(10L)).thenThrow(new RuntimeException("db")); + + ResponseEntity response = controller.getTeamInvitations(10L); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.INTERNAL_SERVER_ERROR); + assertThat(body(response)).containsEntry("error", "Failed to fetch invitations"); + } + } + + @Nested + @DisplayName("removeTeamMember") + class RemoveTeamMember { + + @Test + @DisplayName("removes member and revokes PRO role when the removed user was PRO") + void happyPath_revokesProRole() throws Exception { + stubCurrentUser(); + User proMember = user(2L, "bob", "bob@x.com"); + addRole(proMember, Role.PRO_USER.getRoleId()); + when(userRepository.findById(2L)).thenReturn(Optional.of(proMember)); + Team team = team(10L, "Acme"); + when(teamRepository.findById(10L)).thenReturn(Optional.of(team)); + + ResponseEntity response = controller.removeTeamMember(10L, 2L); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(body(response)).containsEntry("message", "Member removed successfully"); + verify(saasTeamService).removeTeamMember(10L, 2L, currentUser); + verify(userService).changeRole(proMember, Role.USER.getRoleId()); + } + + @Test + @DisplayName("non-PRO removed user is not downgraded") + void happyPath_nonProUntouched() throws Exception { + stubCurrentUser(); + User member = user(2L, "bob", "bob@x.com"); + addRole(member, Role.USER.getRoleId()); + when(userRepository.findById(2L)).thenReturn(Optional.of(member)); + when(teamRepository.findById(10L)).thenReturn(Optional.of(team(10L, "Acme"))); + + ResponseEntity response = controller.removeTeamMember(10L, 2L); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + verify(userService, never()).changeRole(any(), anyString()); + } + + @Test + @DisplayName("member not found -> 400 and rollback") + void memberNotFound_isBadRequest() { + stubCurrentUser(); + TransactionSupport tx = TransactionSupport.bind(); + try { + when(userRepository.findById(2L)).thenReturn(Optional.empty()); + + ResponseEntity response = controller.removeTeamMember(10L, 2L); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + assertThat(body(response)).containsEntry("error", "Member not found"); + verify(tx.status()).setRollbackOnly(); + } finally { + tx.unbind(); + } + } + + @Test + @DisplayName("service SecurityException -> 400 and rollback") + void serviceSecurityException_isBadRequest() { + stubCurrentUser(); + TransactionSupport tx = TransactionSupport.bind(); + try { + when(userRepository.findById(2L)) + .thenReturn(Optional.of(user(2L, "bob", "b@x.com"))); + when(teamRepository.findById(10L)).thenReturn(Optional.of(team(10L, "Acme"))); + doThrow(new SecurityException("Only team leaders can remove members")) + .when(saasTeamService) + .removeTeamMember(10L, 2L, currentUser); + + ResponseEntity response = controller.removeTeamMember(10L, 2L); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + assertThat(body(response)) + .containsEntry("error", "Only team leaders can remove members"); + verify(tx.status()).setRollbackOnly(); + } finally { + tx.unbind(); + } + } + + @Test + @DisplayName("unexpected error -> 500 and rollback") + void unexpectedError_isServerError() { + stubCurrentUser(); + TransactionSupport tx = TransactionSupport.bind(); + try { + when(userRepository.findById(2L)) + .thenReturn(Optional.of(user(2L, "bob", "b@x.com"))); + when(teamRepository.findById(10L)).thenReturn(Optional.of(team(10L, "Acme"))); + doThrow(new RuntimeException("boom")) + .when(saasTeamService) + .removeTeamMember(10L, 2L, currentUser); + + ResponseEntity response = controller.removeTeamMember(10L, 2L); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.INTERNAL_SERVER_ERROR); + assertThat(body(response)).containsEntry("error", "Failed to remove member"); + verify(tx.status()).setRollbackOnly(); + } finally { + tx.unbind(); + } + } + } + + @Nested + @DisplayName("leaveTeam") + class LeaveTeam { + + @Test + @DisplayName("PRO user leaving has PRO role revoked") + void happyPath_revokesProRole() throws Exception { + stubCurrentUser(); + addRole(currentUser, Role.PRO_USER.getRoleId()); + when(teamRepository.findById(10L)).thenReturn(Optional.of(team(10L, "Acme"))); + + ResponseEntity response = controller.leaveTeam(10L); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(body(response)).containsEntry("message", "Left team successfully"); + verify(saasTeamService).leaveTeam(10L, currentUser); + verify(userService).changeRole(currentUser, Role.USER.getRoleId()); + } + + @Test + @DisplayName("non-PRO user leaving is not downgraded") + void happyPath_nonProUntouched() throws Exception { + stubCurrentUser(); + addRole(currentUser, Role.USER.getRoleId()); + when(teamRepository.findById(10L)).thenReturn(Optional.of(team(10L, "Acme"))); + + ResponseEntity response = controller.leaveTeam(10L); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + verify(userService, never()).changeRole(any(), anyString()); + } + + @Test + @DisplayName("last-leader IllegalStateException -> 400 and rollback") + void lastLeader_isBadRequest() { + stubCurrentUser(); + TransactionSupport tx = TransactionSupport.bind(); + try { + when(teamRepository.findById(10L)).thenReturn(Optional.of(team(10L, "Acme"))); + doThrow(new IllegalStateException("Cannot leave as the last team leader.")) + .when(saasTeamService) + .leaveTeam(10L, currentUser); + + ResponseEntity response = controller.leaveTeam(10L); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + assertThat(body(response)) + .containsEntry("error", "Cannot leave as the last team leader."); + verify(tx.status()).setRollbackOnly(); + } finally { + tx.unbind(); + } + } + + @Test + @DisplayName("unexpected error -> 500 and rollback") + void unexpectedError_isServerError() { + stubCurrentUser(); + TransactionSupport tx = TransactionSupport.bind(); + try { + doThrow(new RuntimeException("boom")).when(teamRepository).findById(10L); + + ResponseEntity response = controller.leaveTeam(10L); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.INTERNAL_SERVER_ERROR); + assertThat(body(response)).containsEntry("error", "Failed to leave team"); + verify(tx.status()).setRollbackOnly(); + } finally { + tx.unbind(); + } + } + } + + @Nested + @DisplayName("renameTeamByLeader") + class RenameTeam { + + private RenameTeamRequest req(String name) { + RenameTeamRequest r = new RenameTeamRequest(); + r.setNewName(name); + return r; + } + + @Test + @DisplayName("happy path renames a standard team and trims the name") + void happyPath() { + stubCurrentUser(); + Team team = team(10L, "Old Name"); + when(teamRepository.findById(10L)).thenReturn(Optional.of(team)); + when(saasTeamExtensionService.isPersonal(team)).thenReturn(false); + + ResponseEntity response = controller.renameTeamByLeader(10L, req(" New Name ")); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(body(response)).containsEntry("message", "Team renamed successfully"); + assertThat(body(response)).containsEntry("newName", "New Name"); + assertThat(team.getName()).isEqualTo("New Name"); + verify(teamRepository).save(team); + } + + @Test + @DisplayName("blank name -> 400 before any lookup") + void blankName_isBadRequest() { + ResponseEntity response = controller.renameTeamByLeader(10L, req(" ")); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + assertThat(body(response)).containsEntry("error", "Team name cannot be empty"); + verify(teamRepository, never()).findById(anyLong()); + } + + @Test + @DisplayName("null name -> 400") + void nullName_isBadRequest() { + ResponseEntity response = controller.renameTeamByLeader(10L, req(null)); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + assertThat(body(response)).containsEntry("error", "Team name cannot be empty"); + } + + @Test + @DisplayName("team not found -> 400 with message") + void teamNotFound_isBadRequest() { + when(teamRepository.findById(10L)).thenReturn(Optional.empty()); + + ResponseEntity response = controller.renameTeamByLeader(10L, req("New")); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + assertThat(body(response)).containsEntry("error", "Team not found"); + } + + @Test + @DisplayName("personal team cannot be renamed -> 400") + void personalTeam_isBadRequest() { + Team team = team(10L, "My Team"); + when(teamRepository.findById(10L)).thenReturn(Optional.of(team)); + when(saasTeamExtensionService.isPersonal(team)).thenReturn(true); + + ResponseEntity response = controller.renameTeamByLeader(10L, req("New")); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + assertThat(body(response)).containsEntry("error", "Cannot rename personal team"); + verify(teamRepository, never()).save(any()); + } + + @Test + @DisplayName("Internal team cannot be renamed -> 400") + void internalTeam_isBadRequest() { + Team team = team(10L, TeamService.INTERNAL_TEAM_NAME); + when(teamRepository.findById(10L)).thenReturn(Optional.of(team)); + when(saasTeamExtensionService.isPersonal(team)).thenReturn(false); + + ResponseEntity response = controller.renameTeamByLeader(10L, req("New")); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + assertThat(body(response)).containsEntry("error", "Cannot rename Internal team"); + verify(teamRepository, never()).save(any()); + } + + @Test + @DisplayName("persistence failure -> 500") + void saveFailure_isServerError() { + stubCurrentUser(); + Team team = team(10L, "Old"); + when(teamRepository.findById(10L)).thenReturn(Optional.of(team)); + when(saasTeamExtensionService.isPersonal(team)).thenReturn(false); + when(teamRepository.save(team)).thenThrow(new RuntimeException("db")); + + ResponseEntity response = controller.renameTeamByLeader(10L, req("New")); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.INTERNAL_SERVER_ERROR); + assertThat(body(response)).containsEntry("error", "Failed to rename team"); + } + } + + @Nested + @DisplayName("updateTeamSeats") + class UpdateTeamSeats { + + private UpdateSeatsRequest req(Integer maxSeats) { + UpdateSeatsRequest r = new UpdateSeatsRequest(); + r.setMaxSeats(maxSeats); + return r; + } + + @Test + @DisplayName("happy path returns seat math (available = max - used)") + void happyPath() { + Team team = team(10L, "Acme"); + when(teamRepository.findById(10L)).thenReturn(Optional.of(team)); + when(saasTeamExtensionService.getMaxSeats(team)).thenReturn(10); + when(saasTeamExtensionService.getSeatsUsed(team)).thenReturn(3); + + ResponseEntity response = controller.updateTeamSeats(10L, req(10)); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + verify(saasTeamService).updateTeamSeats(10L, 10); + assertThat(body(response)).containsEntry("success", true); + assertThat(body(response)).containsEntry("teamId", 10L); + assertThat(body(response)).containsEntry("maxSeats", 10); + assertThat(body(response)).containsEntry("seatsUsed", 3); + assertThat(body(response)).containsEntry("availableSeats", 7); + } + + @Test + @DisplayName("invalid seats (service IllegalArgumentException) -> 400") + void invalidSeats_isBadRequest() { + doThrow(new IllegalArgumentException("maxSeats must be at least 1")) + .when(saasTeamService) + .updateTeamSeats(10L, 0); + + ResponseEntity response = controller.updateTeamSeats(10L, req(0)); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + assertThat(body(response)).containsEntry("error", "maxSeats must be at least 1"); + } + + @Test + @DisplayName("unexpected error -> 500") + void unexpectedError_isServerError() { + doThrow(new RuntimeException("boom")).when(saasTeamService).updateTeamSeats(10L, 5); + + ResponseEntity response = controller.updateTeamSeats(10L, req(5)); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.INTERNAL_SERVER_ERROR); + assertThat(body(response)).containsEntry("error", "Failed to update team seats"); + } + } + + @Nested + @DisplayName("getUserPrimaryTeamBySupabaseId") + class GetUserPrimaryTeam { + + @Test + @DisplayName("happy path returns the user's primary team payload") + void happyPath() { + UUID uuid = UUID.randomUUID(); + User user = user(2L, "bob", "bob@x.com"); + user.setSupabaseId(uuid); + Team primary = team(10L, "Acme"); + user.setTeam(primary); + when(userRepository.findBySupabaseId(uuid)).thenReturn(Optional.of(user)); + when(saasTeamExtensionService.isPersonal(primary)).thenReturn(false); + when(saasTeamExtensionService.getMaxSeats(primary)).thenReturn(10); + + ResponseEntity response = controller.getUserPrimaryTeamBySupabaseId(uuid.toString()); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(body(response)).containsEntry("teamId", 10L); + assertThat(body(response)).containsEntry("userId", 2L); + assertThat(body(response)).containsEntry("supabaseUserId", uuid.toString()); + assertThat(body(response)).containsEntry("isPersonal", false); + assertThat(body(response)).containsEntry("maxSeats", 10); + } + + @Test + @DisplayName("malformed UUID -> 400 generic message") + void malformedUuid_isBadRequest() { + ResponseEntity response = controller.getUserPrimaryTeamBySupabaseId("not-a-uuid"); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + assertThat(body(response)) + .containsEntry("error", "Invalid UUID format or user not found"); + } + + @Test + @DisplayName("unknown user -> 400 generic message") + void unknownUser_isBadRequest() { + UUID uuid = UUID.randomUUID(); + when(userRepository.findBySupabaseId(uuid)).thenReturn(Optional.empty()); + + ResponseEntity response = controller.getUserPrimaryTeamBySupabaseId(uuid.toString()); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + assertThat(body(response)) + .containsEntry("error", "Invalid UUID format or user not found"); + } + + @Test + @DisplayName("user with no primary team -> 404") + void noPrimaryTeam_isNotFound() { + UUID uuid = UUID.randomUUID(); + User user = user(2L, "bob", "bob@x.com"); + user.setSupabaseId(uuid); + user.setTeam(null); + when(userRepository.findBySupabaseId(uuid)).thenReturn(Optional.of(user)); + + ResponseEntity response = controller.getUserPrimaryTeamBySupabaseId(uuid.toString()); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.NOT_FOUND); + assertThat(body(response)).containsEntry("error", "User has no primary team"); + } + } + + @Nested + @DisplayName("getTeamInfo") + class GetTeamInfo { + + @Test + @DisplayName("happy path: leader sees full payload with members and seat math") + void happyPath_leader() { + stubCurrentUser(); + Team team = team(10L, "Acme"); + when(teamRepository.findById(10L)).thenReturn(Optional.of(team)); + User bob = user(2L, "bob", "bob@x.com"); + when(membershipRepository.findByTeamId(10L)) + .thenReturn(List.of(membership(team, bob, TeamRole.MEMBER))); + when(membershipRepository.findByTeamIdAndUserId(10L, currentUser.getId())) + .thenReturn(Optional.of(membership(team, currentUser, TeamRole.LEADER))); + when(saasTeamExtensionService.isPersonal(team)).thenReturn(false); + when(saasTeamExtensionService.getMaxSeats(team)).thenReturn(10); + when(saasTeamExtensionService.getSeatsUsed(team)).thenReturn(2); + + ResponseEntity response = controller.getTeamInfo(10L); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(body(response)).containsEntry("teamId", 10L); + assertThat(body(response)).containsEntry("name", "Acme"); + assertThat(body(response)).containsEntry("isPersonal", false); + assertThat(body(response)).containsEntry("maxSeats", 10); + assertThat(body(response)).containsEntry("seatsUsed", 2); + assertThat(body(response)).containsEntry("availableSeats", 8); + assertThat(body(response)).containsEntry("isLeader", true); + assertThat(body(response)).containsKey("members"); + } + + @Test + @DisplayName("non-leader member sees isLeader=false") + void nonLeader() { + stubCurrentUser(); + Team team = team(10L, "Acme"); + when(teamRepository.findById(10L)).thenReturn(Optional.of(team)); + when(membershipRepository.findByTeamId(10L)).thenReturn(List.of()); + when(membershipRepository.findByTeamIdAndUserId(10L, currentUser.getId())) + .thenReturn(Optional.of(membership(team, currentUser, TeamRole.MEMBER))); + when(saasTeamExtensionService.isPersonal(team)).thenReturn(false); + when(saasTeamExtensionService.getMaxSeats(team)).thenReturn(5); + when(saasTeamExtensionService.getSeatsUsed(team)).thenReturn(1); + + ResponseEntity response = controller.getTeamInfo(10L); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(body(response)).containsEntry("isLeader", false); + } + + @Test + @DisplayName("no membership row -> isLeader defaults to false") + void noMembershipRow_leaderFalse() { + stubCurrentUser(); + Team team = team(10L, "Acme"); + when(teamRepository.findById(10L)).thenReturn(Optional.of(team)); + when(membershipRepository.findByTeamId(10L)).thenReturn(List.of()); + when(membershipRepository.findByTeamIdAndUserId(10L, currentUser.getId())) + .thenReturn(Optional.empty()); + when(saasTeamExtensionService.isPersonal(team)).thenReturn(false); + when(saasTeamExtensionService.getMaxSeats(team)).thenReturn(5); + when(saasTeamExtensionService.getSeatsUsed(team)).thenReturn(1); + + ResponseEntity response = controller.getTeamInfo(10L); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(body(response)).containsEntry("isLeader", false); + } + + @Test + @DisplayName("team not found -> 404") + void teamNotFound_isNotFound() { + when(teamRepository.findById(10L)).thenReturn(Optional.empty()); + + ResponseEntity response = controller.getTeamInfo(10L); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.NOT_FOUND); + assertThat(body(response)).containsEntry("error", "Team not found"); + } + + @Test + @DisplayName("unexpected error -> 500") + void unexpectedError_isServerError() { + when(teamRepository.findById(10L)).thenThrow(new RuntimeException("db")); + + ResponseEntity response = controller.getTeamInfo(10L); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.INTERNAL_SERVER_ERROR); + assertThat(body(response)).containsEntry("error", "Failed to fetch team info"); + } + } + + @Nested + @DisplayName("DTO value holders") + class Dtos { + + @Test + @DisplayName("TeamDetailsDTO wires constructor fields verbatim") + void teamDetailsDto() { + SaasTeamController.TeamDetailsDTO dto = + new SaasTeamController.TeamDetailsDTO( + 1L, "Acme", "STANDARD", false, 3, 10, 4, 10, true); + assertThat(dto.getTeamId()).isEqualTo(1L); + assertThat(dto.getName()).isEqualTo("Acme"); + assertThat(dto.getTeamType()).isEqualTo("STANDARD"); + assertThat(dto.getIsPersonal()).isFalse(); + assertThat(dto.getMemberCount()).isEqualTo(3); + assertThat(dto.getSeatCount()).isEqualTo(10); + assertThat(dto.getSeatsUsed()).isEqualTo(4); + assertThat(dto.getMaxSeats()).isEqualTo(10); + assertThat(dto.getIsLeader()).isTrue(); + } + + @Test + @DisplayName("InviteUserRequest is a mutable POJO") + void inviteUserRequest() { + InviteUserRequest r = new InviteUserRequest(); + r.setTeamId(5L); + r.setEmail("x@y.com"); + assertThat(r.getTeamId()).isEqualTo(5L); + assertThat(r.getEmail()).isEqualTo("x@y.com"); + } + } + + // ===== shared test infrastructure ===== + + private static void addRole(User user, String roleId) { + // Authority's (String, User) ctor self-registers on the user's authority set. + new stirling.software.proprietary.security.model.Authority(roleId, user); + } + + /** + * Binds a Spring transaction context to the current thread so {@code @Transactional} handlers + * can call {@link TransactionAspectSupport#currentTransactionStatus()} on their error paths + * without a live Spring transaction, and verify {@code setRollbackOnly()} on the resulting + * status. + * + *

Spring's {@code TransactionInfo} type is {@code protected} and its {@code bindToThread()} + * plus the backing {@code transactionInfoHolder} ThreadLocal are {@code private}, so the whole + * binding is performed reflectively. {@link #unbind()} clears the ThreadLocal again so the + * binding never leaks into sibling tests. + */ + private static final class TransactionSupport { + + @SuppressWarnings("unchecked") + private static final ThreadLocal HOLDER = resolveHolder(); + + private final TransactionStatus status; + + private TransactionSupport(TransactionStatus status) { + this.status = status; + HOLDER.set(newTransactionInfo(status)); + } + + @SuppressWarnings("unchecked") + private static ThreadLocal resolveHolder() { + try { + java.lang.reflect.Field field = + TransactionAspectSupport.class.getDeclaredField("transactionInfoHolder"); + field.setAccessible(true); + return (ThreadLocal) field.get(null); + } catch (ReflectiveOperationException e) { + throw new IllegalStateException("Unable to access transactionInfoHolder", e); + } + } + + /** + * Reflectively build a TransactionInfo exposing the given status (protected nested type). + */ + private static Object newTransactionInfo(TransactionStatus status) { + try { + Class infoClass = + Class.forName( + "org.springframework.transaction.interceptor." + + "TransactionAspectSupport$TransactionInfo"); + java.lang.reflect.Constructor ctor = + infoClass.getDeclaredConstructor( + org.springframework.transaction.PlatformTransactionManager.class, + org.springframework.transaction.interceptor.TransactionAttribute + .class, + String.class); + ctor.setAccessible(true); + Object info = ctor.newInstance(null, null, "test"); + java.lang.reflect.Method newStatus = + infoClass.getDeclaredMethod( + "newTransactionStatus", TransactionStatus.class); + newStatus.setAccessible(true); + newStatus.invoke(info, status); + return info; + } catch (ReflectiveOperationException e) { + throw new IllegalStateException("Unable to build TransactionInfo", e); + } + } + + static TransactionSupport bind() { + return new TransactionSupport(mock(TransactionStatus.class)); + } + + TransactionStatus status() { + return status; + } + + void unbind() { + HOLDER.remove(); + } + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/controller/UserRoleWebhookControllerTest.java b/app/saas/src/test/java/stirling/software/saas/controller/UserRoleWebhookControllerTest.java new file mode 100644 index 0000000000..e0ff221459 --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/controller/UserRoleWebhookControllerTest.java @@ -0,0 +1,556 @@ +package stirling.software.saas.controller; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.verifyNoInteractions; +import static org.mockito.Mockito.when; + +import java.security.Principal; +import java.util.Map; +import java.util.Optional; +import java.util.UUID; + +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.ArgumentCaptor; +import org.mockito.InjectMocks; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.http.HttpStatus; +import org.springframework.http.ResponseEntity; + +import stirling.software.proprietary.security.model.AuthenticationType; +import stirling.software.proprietary.security.model.User; +import stirling.software.proprietary.security.service.UserService; +import stirling.software.saas.model.SupabaseUser; +import stirling.software.saas.service.SaasUserAccountService; +import stirling.software.saas.service.SupabaseUserService; + +/** + * Pure-Mockito unit tests for {@link UserRoleWebhookController}. + * + *

The controller is built via {@code @RequiredArgsConstructor}, so {@link InjectMocks} wires the + * three mocked collaborators ({@link UserService}, {@link SaasUserAccountService}, {@link + * SupabaseUserService}) by type. Each handler is invoked directly and the returned {@link + * ResponseEntity} (status + body) is asserted, alongside collaborator interaction verification. No + * Spring context, DB, Supabase or network is involved; {@code @PreAuthorize} is a no-op outside the + * security proxy so authorization is not exercised here. + */ +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class UserRoleWebhookControllerTest { + + @Mock private UserService userService; + @Mock private SaasUserAccountService saasUserAccountService; + @Mock private SupabaseUserService supabaseUserService; + + @InjectMocks private UserRoleWebhookController controller; + + private static final String SUPABASE_ID = "11111111-2222-3333-4444-555555555555"; + + @Nested + @DisplayName("POST /upgrade") + class HandleUpgrade { + + @Test + @DisplayName("returns 200 with 'upgraded' message when a promotion happened") + void upgraded() { + when(saasUserAccountService.handleUpgrade(SUPABASE_ID)).thenReturn(true); + + ResponseEntity response = controller.handleUpgrade(SUPABASE_ID); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(response.getBody()).isEqualTo("User upgraded to PRO successfully"); + verify(saasUserAccountService).handleUpgrade(SUPABASE_ID); + } + + @Test + @DisplayName("returns 200 with 'already PRO' message when nothing changed") + void alreadyPro() { + when(saasUserAccountService.handleUpgrade(SUPABASE_ID)).thenReturn(false); + + ResponseEntity response = controller.handleUpgrade(SUPABASE_ID); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(response.getBody()).isEqualTo("User is already PRO"); + } + + @Test + @DisplayName( + "maps IllegalArgumentException (bad/unknown supabaseId) to 400 'Invalid request'") + void illegalArgumentMapsTo400() { + when(saasUserAccountService.handleUpgrade(SUPABASE_ID)) + .thenThrow(new IllegalArgumentException("Invalid Supabase ID format")); + + ResponseEntity response = controller.handleUpgrade(SUPABASE_ID); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + assertThat(response.getBody()).isEqualTo("Invalid request"); + } + + @Test + @DisplayName("maps any other exception to 500 'Error processing webhook'") + void unexpectedExceptionMapsTo500() { + when(saasUserAccountService.handleUpgrade(SUPABASE_ID)) + .thenThrow(new RuntimeException("db down")); + + ResponseEntity response = controller.handleUpgrade(SUPABASE_ID); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.INTERNAL_SERVER_ERROR); + assertThat(response.getBody()).isEqualTo("Error processing webhook"); + } + } + + @Nested + @DisplayName("POST /downgrade") + class HandleDowngrade { + + @Test + @DisplayName("returns 200 with 'downgraded' message when a demotion happened") + void downgraded() { + when(saasUserAccountService.handleDowngrade(SUPABASE_ID)).thenReturn(true); + + ResponseEntity response = controller.handleDowngrade(SUPABASE_ID); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(response.getBody()).isEqualTo("User downgraded to FREE successfully"); + verify(saasUserAccountService).handleDowngrade(SUPABASE_ID); + } + + @Test + @DisplayName("returns 200 with 'already FREE' message when nothing changed") + void alreadyFree() { + when(saasUserAccountService.handleDowngrade(SUPABASE_ID)).thenReturn(false); + + ResponseEntity response = controller.handleDowngrade(SUPABASE_ID); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(response.getBody()).isEqualTo("User is already on FREE tier"); + } + + @Test + @DisplayName("maps IllegalArgumentException to 400 'Invalid request'") + void illegalArgumentMapsTo400() { + when(saasUserAccountService.handleDowngrade(SUPABASE_ID)) + .thenThrow(new IllegalArgumentException("User not found for Supabase ID")); + + ResponseEntity response = controller.handleDowngrade(SUPABASE_ID); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + assertThat(response.getBody()).isEqualTo("Invalid request"); + } + + @Test + @DisplayName("maps any other exception to 500 'Error processing webhook'") + void unexpectedExceptionMapsTo500() { + when(saasUserAccountService.handleDowngrade(SUPABASE_ID)) + .thenThrow(new RuntimeException("boom")); + + ResponseEntity response = controller.handleDowngrade(SUPABASE_ID); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.INTERNAL_SERVER_ERROR); + assertThat(response.getBody()).isEqualTo("Error processing webhook"); + } + } + + @Nested + @DisplayName("POST /enable-metered-billing") + class EnableMeteredBilling { + + @Test + @DisplayName("returns 200 'enabled' when metered billing is newly turned on") + void enabled() { + when(saasUserAccountService.enableMeteredBilling(SUPABASE_ID)).thenReturn(true); + + ResponseEntity response = controller.enableMeteredBilling(SUPABASE_ID); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(response.getBody()).isEqualTo("Metered billing enabled successfully"); + verify(saasUserAccountService).enableMeteredBilling(SUPABASE_ID); + } + + @Test + @DisplayName("returns 200 'already enabled' when no change was made") + void alreadyEnabled() { + when(saasUserAccountService.enableMeteredBilling(SUPABASE_ID)).thenReturn(false); + + ResponseEntity response = controller.enableMeteredBilling(SUPABASE_ID); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(response.getBody()).isEqualTo("User already has metered billing enabled"); + } + + @Test + @DisplayName("maps IllegalArgumentException to 400 'Invalid request'") + void illegalArgumentMapsTo400() { + when(saasUserAccountService.enableMeteredBilling(SUPABASE_ID)) + .thenThrow(new IllegalArgumentException("Invalid Supabase ID format")); + + ResponseEntity response = controller.enableMeteredBilling(SUPABASE_ID); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + assertThat(response.getBody()).isEqualTo("Invalid request"); + } + + @Test + @DisplayName("maps any other exception to 500 'Error processing webhook'") + void unexpectedExceptionMapsTo500() { + when(saasUserAccountService.enableMeteredBilling(SUPABASE_ID)) + .thenThrow(new RuntimeException("stripe down")); + + ResponseEntity response = controller.enableMeteredBilling(SUPABASE_ID); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.INTERNAL_SERVER_ERROR); + assertThat(response.getBody()).isEqualTo("Error processing webhook"); + } + } + + @Nested + @DisplayName("POST /disable-metered-billing") + class DisableMeteredBilling { + + @Test + @DisplayName("returns 200 'disabled' when metered billing is newly turned off") + void disabled() { + when(saasUserAccountService.disableMeteredBilling(SUPABASE_ID)).thenReturn(true); + + ResponseEntity response = controller.disableMeteredBilling(SUPABASE_ID); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(response.getBody()).isEqualTo("Metered billing disabled successfully"); + verify(saasUserAccountService).disableMeteredBilling(SUPABASE_ID); + } + + @Test + @DisplayName("returns 200 'does not have' when no change was made") + void notEnabled() { + when(saasUserAccountService.disableMeteredBilling(SUPABASE_ID)).thenReturn(false); + + ResponseEntity response = controller.disableMeteredBilling(SUPABASE_ID); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(response.getBody()).isEqualTo("User does not have metered billing enabled"); + } + + @Test + @DisplayName("maps IllegalArgumentException to 400 'Invalid request'") + void illegalArgumentMapsTo400() { + when(saasUserAccountService.disableMeteredBilling(SUPABASE_ID)) + .thenThrow(new IllegalArgumentException("User not found for Supabase ID")); + + ResponseEntity response = controller.disableMeteredBilling(SUPABASE_ID); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + assertThat(response.getBody()).isEqualTo("Invalid request"); + } + + @Test + @DisplayName("maps any other exception to 500 'Error processing webhook'") + void unexpectedExceptionMapsTo500() { + when(saasUserAccountService.disableMeteredBilling(SUPABASE_ID)) + .thenThrow(new RuntimeException("kaboom")); + + ResponseEntity response = controller.disableMeteredBilling(SUPABASE_ID); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.INTERNAL_SERVER_ERROR); + assertThat(response.getBody()).isEqualTo("Error processing webhook"); + } + } + + @Nested + @DisplayName("POST /promptToAuthUser") + class PromptToAuthUser { + + private static final String USERNAME = "anon-user"; + private static final UUID LINKED_SUPABASE_ID = + UUID.fromString("aaaaaaaa-bbbb-cccc-dddd-eeeeeeeeeeee"); + + private Principal principal(String name) { + Principal p = org.mockito.Mockito.mock(Principal.class); + when(p.getName()).thenReturn(name); + return p; + } + + private User anonymousUser(UUID supabaseId) { + User user = new User(); + user.setUsername(USERNAME); + user.setSupabaseId(supabaseId); + user.setAuthenticationType(AuthenticationType.ANONYMOUS); + return user; + } + + private SupabaseUser supabaseUserWithEmail(String email) { + SupabaseUser su = new SupabaseUser(); + su.setId(LINKED_SUPABASE_ID); + su.setEmail(email); + su.setAnonymous(true); + return su; + } + + @Test + @DisplayName("happy path: synchronizes upgrade and returns 200 with userId/email body") + void happyPath() { + User current = anonymousUser(LINKED_SUPABASE_ID); + when(userService.findByUsername(USERNAME)).thenReturn(Optional.of(current)); + + SupabaseUser supabaseUser = supabaseUserWithEmail("new@stirling.com"); + when(supabaseUserService.getUser(LINKED_SUPABASE_ID)).thenReturn(supabaseUser); + + User upgraded = new User(); + upgraded.setId(42L); + upgraded.setEmail("new@stirling.com"); + upgraded.setUsername("new@stirling.com"); + when(saasUserAccountService.synchronizeUserUpgrade( + supabaseUser, "new@stirling.com", "google")) + .thenReturn(upgraded); + + ResponseEntity> response = + controller.promptToAuthUser("google", principal(USERNAME)); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(response.getBody()) + .containsEntry("message", "User upgrade synchronized successfully") + .containsEntry("userId", "42") + .containsEntry("email", "new@stirling.com"); + } + + @Test + @DisplayName("normalizes auth method to lowercase/trimmed before delegating") + void normalizesAuthMethod() { + User current = anonymousUser(LINKED_SUPABASE_ID); + when(userService.findByUsername(USERNAME)).thenReturn(Optional.of(current)); + SupabaseUser supabaseUser = supabaseUserWithEmail("a@b.com"); + when(supabaseUserService.getUser(LINKED_SUPABASE_ID)).thenReturn(supabaseUser); + + User upgraded = new User(); + upgraded.setId(7L); + upgraded.setEmail("a@b.com"); + when(saasUserAccountService.synchronizeUserUpgrade(any(), anyString(), anyString())) + .thenReturn(upgraded); + + ResponseEntity> response = + controller.promptToAuthUser(" GitHub ", principal(USERNAME)); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + ArgumentCaptor methodCaptor = ArgumentCaptor.forClass(String.class); + verify(saasUserAccountService) + .synchronizeUserUpgrade( + eq(supabaseUser), eq("a@b.com"), methodCaptor.capture()); + assertThat(methodCaptor.getValue()).isEqualTo("github"); + } + + @Test + @DisplayName("null authMethod is accepted and passed through as null") + void nullAuthMethodAccepted() { + User current = anonymousUser(LINKED_SUPABASE_ID); + when(userService.findByUsername(USERNAME)).thenReturn(Optional.of(current)); + SupabaseUser supabaseUser = supabaseUserWithEmail("a@b.com"); + when(supabaseUserService.getUser(LINKED_SUPABASE_ID)).thenReturn(supabaseUser); + + User upgraded = new User(); + upgraded.setId(7L); + upgraded.setEmail("a@b.com"); + when(saasUserAccountService.synchronizeUserUpgrade( + eq(supabaseUser), eq("a@b.com"), eq(null))) + .thenReturn(upgraded); + + ResponseEntity> response = + controller.promptToAuthUser(null, principal(USERNAME)); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + verify(saasUserAccountService).synchronizeUserUpgrade(supabaseUser, "a@b.com", null); + } + + @Test + @DisplayName("falls back to username in body when upgraded user has no email") + void emailFallsBackToUsername() { + User current = anonymousUser(LINKED_SUPABASE_ID); + when(userService.findByUsername(USERNAME)).thenReturn(Optional.of(current)); + SupabaseUser supabaseUser = supabaseUserWithEmail("canon@b.com"); + when(supabaseUserService.getUser(LINKED_SUPABASE_ID)).thenReturn(supabaseUser); + + User upgraded = new User(); + upgraded.setId(9L); + upgraded.setEmail(null); + upgraded.setUsername("fallback-username"); + when(saasUserAccountService.synchronizeUserUpgrade(any(), anyString(), any())) + .thenReturn(upgraded); + + ResponseEntity> response = + controller.promptToAuthUser("email", principal(USERNAME)); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(response.getBody()).containsEntry("email", "fallback-username"); + } + + @Test + @DisplayName("invalid auth method returns 400 without touching userService") + void invalidAuthMethodRejected() { + ResponseEntity> response = + controller.promptToAuthUser("myspace", principal(USERNAME)); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + assertThat(response.getBody()).containsEntry("error", "Invalid authentication method"); + verifyNoInteractions(userService); + verifyNoInteractions(saasUserAccountService); + } + + @Test + @DisplayName("unknown current user (IllegalStateException) maps to 404 'User not found'") + void currentUserNotFound() { + when(userService.findByUsername(USERNAME)).thenReturn(Optional.empty()); + + ResponseEntity> response = + controller.promptToAuthUser("email", principal(USERNAME)); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.NOT_FOUND); + assertThat(response.getBody()).containsEntry("error", "User not found"); + verify(saasUserAccountService, never()).synchronizeUserUpgrade(any(), any(), any()); + } + + @Test + @DisplayName("current user without a linked Supabase ID returns 400") + void noLinkedSupabaseId() { + User current = anonymousUser(null); + when(userService.findByUsername(USERNAME)).thenReturn(Optional.of(current)); + + ResponseEntity> response = + controller.promptToAuthUser("email", principal(USERNAME)); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + assertThat(response.getBody()) + .containsEntry("error", "No Supabase account linked to current user"); + verifyNoInteractions(supabaseUserService); + } + + @Test + @DisplayName("non-anonymous user is rejected with 400") + void nonAnonymousRejected() { + User current = anonymousUser(LINKED_SUPABASE_ID); + current.setAuthenticationType(AuthenticationType.WEB); + when(userService.findByUsername(USERNAME)).thenReturn(Optional.of(current)); + when(supabaseUserService.getUser(LINKED_SUPABASE_ID)) + .thenReturn(supabaseUserWithEmail("x@y.com")); + + ResponseEntity> response = + controller.promptToAuthUser("email", principal(USERNAME)); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + assertThat(response.getBody()) + .containsEntry("error", "Only anonymous users can be upgraded"); + verify(saasUserAccountService, never()).synchronizeUserUpgrade(any(), any(), any()); + } + + @Test + @DisplayName("falls back to local user email when Supabase email is blank") + void canonicalEmailFallsBackToLocal() { + User current = anonymousUser(LINKED_SUPABASE_ID); + current.setEmail("local@stirling.com"); + when(userService.findByUsername(USERNAME)).thenReturn(Optional.of(current)); + + SupabaseUser supabaseUser = supabaseUserWithEmail(" "); // blank + when(supabaseUserService.getUser(LINKED_SUPABASE_ID)).thenReturn(supabaseUser); + + User upgraded = new User(); + upgraded.setId(5L); + upgraded.setEmail("local@stirling.com"); + when(saasUserAccountService.synchronizeUserUpgrade( + eq(supabaseUser), eq("local@stirling.com"), anyString())) + .thenReturn(upgraded); + + ResponseEntity> response = + controller.promptToAuthUser("email", principal(USERNAME)); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + verify(saasUserAccountService) + .synchronizeUserUpgrade(supabaseUser, "local@stirling.com", "email"); + } + + @Test + @DisplayName("no email anywhere (Supabase and local both blank) returns 400") + void noEmailAnywhere() { + User current = anonymousUser(LINKED_SUPABASE_ID); + current.setEmail(null); + when(userService.findByUsername(USERNAME)).thenReturn(Optional.of(current)); + + SupabaseUser supabaseUser = supabaseUserWithEmail(null); + when(supabaseUserService.getUser(LINKED_SUPABASE_ID)).thenReturn(supabaseUser); + + ResponseEntity> response = + controller.promptToAuthUser("email", principal(USERNAME)); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.BAD_REQUEST); + assertThat(response.getBody()) + .containsEntry("error", "No email associated with user account"); + verify(saasUserAccountService, never()).synchronizeUserUpgrade(any(), any(), any()); + } + + @Test + @DisplayName("unexpected RuntimeException from sync maps to 500") + void unexpectedExceptionMapsTo500() { + User current = anonymousUser(LINKED_SUPABASE_ID); + when(userService.findByUsername(USERNAME)).thenReturn(Optional.of(current)); + SupabaseUser supabaseUser = supabaseUserWithEmail("a@b.com"); + when(supabaseUserService.getUser(LINKED_SUPABASE_ID)).thenReturn(supabaseUser); + when(saasUserAccountService.synchronizeUserUpgrade(any(), anyString(), any())) + .thenThrow(new RuntimeException("db down")); + + ResponseEntity> response = + controller.promptToAuthUser("email", principal(USERNAME)); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.INTERNAL_SERVER_ERROR); + assertThat(response.getBody()) + .containsEntry("error", "Failed to synchronize user upgrade"); + } + + @Test + @DisplayName("getUser throwing (Supabase row missing) maps to 500") + void supabaseUserMissingMapsTo500() { + User current = anonymousUser(LINKED_SUPABASE_ID); + when(userService.findByUsername(USERNAME)).thenReturn(Optional.of(current)); + when(supabaseUserService.getUser(LINKED_SUPABASE_ID)) + .thenThrow(new RuntimeException("Supabase user not found")); + + ResponseEntity> response = + controller.promptToAuthUser("email", principal(USERNAME)); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.INTERNAL_SERVER_ERROR); + assertThat(response.getBody()) + .containsEntry("error", "Failed to synchronize user upgrade"); + } + + @Test + @DisplayName("all allowed auth methods are accepted (none rejected as invalid)") + void allowedAuthMethodsAccepted() { + for (String method : + new String[] { + "email", "oauth", "google", "github", "apple", "azure", "linkedin_oidc" + }) { + User current = anonymousUser(LINKED_SUPABASE_ID); + when(userService.findByUsername(USERNAME)).thenReturn(Optional.of(current)); + SupabaseUser supabaseUser = supabaseUserWithEmail("a@b.com"); + when(supabaseUserService.getUser(LINKED_SUPABASE_ID)).thenReturn(supabaseUser); + User upgraded = new User(); + upgraded.setId(1L); + upgraded.setEmail("a@b.com"); + when(saasUserAccountService.synchronizeUserUpgrade(any(), anyString(), anyString())) + .thenReturn(upgraded); + + ResponseEntity> response = + controller.promptToAuthUser(method, principal(USERNAME)); + + assertThat(response.getStatusCode()) + .as("method %s should be accepted", method) + .isEqualTo(HttpStatus.OK); + } + } + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/payg/api/PaygWalletControllerTest.java b/app/saas/src/test/java/stirling/software/saas/payg/api/PaygWalletControllerTest.java new file mode 100644 index 0000000000..8433f5e5a3 --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/payg/api/PaygWalletControllerTest.java @@ -0,0 +1,489 @@ +package stirling.software.saas.payg.api; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.times; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.verifyNoInteractions; +import static org.mockito.Mockito.when; + +import java.math.BigDecimal; +import java.time.LocalDate; +import java.time.LocalDateTime; +import java.util.List; +import java.util.Optional; +import java.util.UUID; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.ArgumentCaptor; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.springframework.http.HttpStatus; +import org.springframework.http.ResponseEntity; +import org.springframework.security.authentication.AnonymousAuthenticationToken; +import org.springframework.security.core.Authentication; +import org.springframework.security.core.authority.SimpleGrantedAuthority; +import org.springframework.security.oauth2.jwt.Jwt; + +import stirling.software.common.model.enumeration.TeamRole; +import stirling.software.proprietary.model.Team; +import stirling.software.proprietary.security.database.repository.UserRepository; +import stirling.software.proprietary.security.model.User; +import stirling.software.saas.model.TeamMembership; +import stirling.software.saas.payg.api.PaygWalletController.UpdateCapRequest; +import stirling.software.saas.payg.api.WalletSnapshotResponse.MemberRow; +import stirling.software.saas.payg.billing.TeamBillingContext; +import stirling.software.saas.payg.billing.TeamBillingService; +import stirling.software.saas.payg.entitlement.EntitlementService; +import stirling.software.saas.payg.entitlement.EntitlementSnapshot; +import stirling.software.saas.payg.model.BillingCategory; +import stirling.software.saas.payg.model.EntitlementState; +import stirling.software.saas.payg.model.FeatureGate; +import stirling.software.saas.payg.model.FeatureSet; +import stirling.software.saas.payg.model.LedgerEntryType; +import stirling.software.saas.payg.repository.PaygShadowChargeRepository; +import stirling.software.saas.payg.repository.PaygTeamExtensionsRepository; +import stirling.software.saas.payg.repository.WalletLedgerRepository; +import stirling.software.saas.payg.repository.WalletPolicyRepository; +import stirling.software.saas.payg.wallet.WalletPolicy; +import stirling.software.saas.repository.TeamMembershipRepository; +import stirling.software.saas.security.EnhancedJwtAuthenticationToken; + +/** + * Pure-Mockito unit tests for {@link PaygWalletController}. Covers the documented role / state + * matrix — free vs subscribed, leader vs member, anonymous — plus the cap update endpoint's + * leader-only enforcement and cache invalidation. + */ +@ExtendWith(MockitoExtension.class) +class PaygWalletControllerTest { + + @Mock private EntitlementService entitlementService; + @Mock private TeamBillingService billingService; + @Mock private TeamMembershipRepository memberRepo; + @Mock private PaygTeamExtensionsRepository extRepo; + @Mock private WalletPolicyRepository policyRepo; + @Mock private WalletLedgerRepository ledgerRepo; + @Mock private PaygShadowChargeRepository shadowRepo; + @Mock private UserRepository userRepository; + + private PaygWalletController controller; + + @BeforeEach + void setUp() { + controller = + new PaygWalletController( + entitlementService, + billingService, + memberRepo, + extRepo, + policyRepo, + ledgerRepo, + shadowRepo, + userRepository); + } + + /** + * Free-team billing context: the one-time grant is fully unused (remaining == grant), no + * subscription facts, no monthly cap. The displayed limit comes from the snapshot, not here. + */ + private static TeamBillingContext freeBilling(long teamFreeGrant) { + LocalDateTime start = LocalDate.now().withDayOfMonth(1).atStartOfDay(); + return new TeamBillingContext( + false, + null, + start, + start.plusMonths(1), + teamFreeGrant, + teamFreeGrant, + null, + null, + null, + null); + } + + /** + * Subscribed billing context with a money cap (minor units) and a known per-doc rate. The grant + * is treated as exhausted (remaining 0) — typical for a team that has subscribed; {@code + * monthlyCapDocUnits} is the paid-doc ceiling {@code floor(capMoney / rate)}. + */ + private static TeamBillingContext subscribedBilling( + String subscriptionId, Long capMoneyMinor, Long monthlyCapDocUnits) { + LocalDateTime start = LocalDate.now().withDayOfMonth(1).atStartOfDay(); + return new TeamBillingContext( + true, + subscriptionId, + start, + start.plusMonths(1), + 500L, + 0L, + BigDecimal.valueOf(2), + "usd", + capMoneyMinor, + monthlyCapDocUnits); + } + + private void stubEmptyLedgerReads(long teamId) { + when(ledgerRepo.sumPeriodAmountByCategory( + eq(teamId), eq(LedgerEntryType.DEBIT), any(), any())) + .thenReturn(List.of()); + when(ledgerRepo.findTop20ByTeamIdOrderByIdDesc(teamId)).thenReturn(List.of()); + } + + // ----------------------------------------------------------------------------------------- + // GET /wallet + // ----------------------------------------------------------------------------------------- + + @Test + void getWallet_freeTier_returnsFreeShape() { + User user = userWithId(7L, UUID.randomUUID()); + Team team = teamWithId(42L); + when(userRepository.findBySupabaseId(any())).thenReturn(Optional.of(user)); + when(memberRepo.findPrimaryMembership(7L)) + .thenReturn(List.of(membership(team, user, TeamRole.MEMBER))); + when(billingService.forTeam(42L)).thenReturn(freeBilling(500L)); + when(entitlementService.getSnapshot(42L)).thenReturn(snapshot(0L, 500L)); + stubEmptyLedgerReads(42L); + + ResponseEntity resp = + controller.getWallet(jwtAuth(user.getSupabaseId())); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.OK); + WalletSnapshotResponse body = resp.getBody(); + assertThat(body).isNotNull(); + assertThat(body.status()).isEqualTo("free"); + assertThat(body.role()).isEqualTo("member"); + assertThat(body.capUsd()).isNull(); + assertThat(body.stripeSubscriptionId()).isNull(); + assertThat(body.noCap()).isFalse(); + assertThat(body.billableUsed()).isZero(); + assertThat(body.billableLimit()).isEqualTo(500); + assertThat(body.freeAllowance()).isEqualTo(500); + assertThat(body.pricePerDocMinor()).isNull(); + assertThat(body.currency()).isNull(); + assertThat(body.estimatedBillMinor()).isNull(); + assertThat(body.members()).isEmpty(); + assertThat(body.recent()).isEmpty(); + assertThat(body.categoryBreakdown().api()).isZero(); + assertThat(body.categoryBreakdown().ai()).isZero(); + assertThat(body.categoryBreakdown().automation()).isZero(); + } + + @Test + void getWallet_subscribedMember_returnsCapAndBreakdownFromLedger() { + User user = userWithId(8L, UUID.randomUUID()); + Team team = teamWithId(99L); + when(userRepository.findBySupabaseId(any())).thenReturn(Optional.of(user)); + when(memberRepo.findPrimaryMembership(8L)) + .thenReturn(List.of(membership(team, user, TeamRole.MEMBER))); + // $25 cap (2500 minor) at $0.02/doc → 1250 paid docs/month (the one-time grant is a + // separate pool, not added to the cap). + when(billingService.forTeam(99L)) + .thenReturn(subscribedBilling("sub_test_99", 2500L, 1250L)); + // 312 paid (metered) docs this period → estimate 312 × $0.02 = $6.24 (624 minor). + when(shadowRepo.sumPaidUnits(eq(99L), any(), any())).thenReturn(312L); + when(billingService.estimateBillMinor(any(), eq(312L))).thenReturn(Optional.of(624L)); + when(entitlementService.getSnapshot(99L)).thenReturn(snapshot(312L, 1250L)); + when(ledgerRepo.sumPeriodAmountByCategory(eq(99L), eq(LedgerEntryType.DEBIT), any(), any())) + .thenReturn( + List.of( + new Object[] {BillingCategory.API, 110L}, + new Object[] {BillingCategory.AI, 200L}, + new Object[] {BillingCategory.AUTOMATION, 2L})); + when(ledgerRepo.findTop20ByTeamIdOrderByIdDesc(99L)).thenReturn(List.of()); + + ResponseEntity resp = + controller.getWallet(jwtAuth(user.getSupabaseId())); + + WalletSnapshotResponse body = resp.getBody(); + assertThat(body.status()).isEqualTo("subscribed"); + assertThat(body.role()).isEqualTo("member"); + assertThat(body.capUsd()).isEqualTo(25); + assertThat(body.noCap()).isFalse(); + assertThat(body.billableUsed()).isEqualTo(312); + assertThat(body.spendUnitsThisPeriod()).isEqualTo(312); + assertThat(body.billableLimit()).isEqualTo(1250); + assertThat(body.freeAllowance()).isEqualTo(500); + assertThat(body.pricePerDocMinor()).isEqualByComparingTo(BigDecimal.valueOf(2)); + assertThat(body.currency()).isEqualTo("usd"); + assertThat(body.estimatedBillMinor()).isEqualTo(624L); + assertThat(body.members()).isEmpty(); + assertThat(body.categoryBreakdown().api()).isEqualTo(110); + assertThat(body.categoryBreakdown().ai()).isEqualTo(200); + assertThat(body.categoryBreakdown().automation()).isEqualTo(2); + assertThat(body.stripeSubscriptionId()).isEqualTo("sub_test_99"); + // Member role → ledger never queried per-user. + verify(ledgerRepo, never()).sumPeriodAmountForMember(any(), any(), any(), any(), any()); + } + + @Test + void getWallet_subscribedNoCap_returnsNoCapTrueAndNullLimit() { + User user = userWithId(9L, UUID.randomUUID()); + Team team = teamWithId(11L); + when(userRepository.findBySupabaseId(any())).thenReturn(Optional.of(user)); + when(memberRepo.findPrimaryMembership(9L)) + .thenReturn(List.of(membership(team, user, TeamRole.LEADER))); + when(billingService.forTeam(11L)).thenReturn(subscribedBilling("sub_nocap", null, null)); + when(entitlementService.getSnapshot(11L)).thenReturn(snapshot(50L, null)); + stubEmptyLedgerReads(11L); + when(memberRepo.findByTeamId(11L)).thenReturn(List.of()); + + ResponseEntity resp = + controller.getWallet(jwtAuth(user.getSupabaseId())); + + WalletSnapshotResponse body = resp.getBody(); + assertThat(body.status()).isEqualTo("subscribed"); + assertThat(body.role()).isEqualTo("leader"); + assertThat(body.capUsd()).isNull(); + assertThat(body.noCap()).isTrue(); + // Uncapped → no document ceiling to draw a bar against. + assertThat(body.billableLimit()).isNull(); + } + + @Test + void getWallet_leader_populatesMembers() { + User leader = userWithId(10L, UUID.randomUUID()); + User member = userWithId(11L, UUID.randomUUID()); + member.setUsername("alice"); + member.setEmail("alice@example.com"); + Team team = teamWithId(77L); + TeamMembership leaderRow = membership(team, leader, TeamRole.LEADER); + TeamMembership memberRow = membership(team, member, TeamRole.MEMBER); + + when(userRepository.findBySupabaseId(any())).thenReturn(Optional.of(leader)); + when(memberRepo.findPrimaryMembership(10L)).thenReturn(List.of(leaderRow)); + when(billingService.forTeam(77L)).thenReturn(freeBilling(500L)); + when(entitlementService.getSnapshot(77L)).thenReturn(snapshot(0L, 500L)); + stubEmptyLedgerReads(77L); + when(memberRepo.findByTeamId(77L)).thenReturn(List.of(leaderRow, memberRow)); + // Ledger returns signed (negative) debits. + when(ledgerRepo.sumPeriodAmountForMember(eq(77L), any(), any(), any(), any())) + .thenReturn(-42L); + + ResponseEntity resp = + controller.getWallet(jwtAuth(leader.getSupabaseId())); + + WalletSnapshotResponse body = resp.getBody(); + assertThat(body.role()).isEqualTo("leader"); + assertThat(body.members()).hasSize(2); + MemberRow secondMember = + body.members().stream() + .filter(m -> "alice".equals(m.name())) + .findFirst() + .orElseThrow(); + assertThat(secondMember.email()).isEqualTo("alice@example.com"); + assertThat(secondMember.spendUnits()).isEqualTo(42); + } + + @Test + void getWallet_anonymousIsRejected() { + Authentication anon = + new AnonymousAuthenticationToken( + "k", + "anonymousUser", + List.of(new SimpleGrantedAuthority("ROLE_ANONYMOUS"))); + + ResponseEntity resp = controller.getWallet(anon); + + // AuthenticationUtils.getCurrentUser throws SecurityException for "anonymousUser" since it + // has no Supabase id and is not a User principal — the controller maps that to 401. + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.UNAUTHORIZED); + verifyNoInteractions(entitlementService, billingService, memberRepo, extRepo, policyRepo); + } + + @Test + void getWallet_authenticatedNoTeam_returnsEmptyFreeShape() { + User user = userWithId(12L, UUID.randomUUID()); + when(userRepository.findBySupabaseId(any())).thenReturn(Optional.of(user)); + when(memberRepo.findPrimaryMembership(12L)).thenReturn(List.of()); + + ResponseEntity resp = + controller.getWallet(jwtAuth(user.getSupabaseId())); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.OK); + assertThat(resp.getBody().status()).isEqualTo("free"); + assertThat(resp.getBody().members()).isEmpty(); + // Entitlement service must not be queried for a teamless user (avoids null-key NPE). + verifyNoInteractions(entitlementService); + } + + // ----------------------------------------------------------------------------------------- + // PATCH /cap + // ----------------------------------------------------------------------------------------- + + @Test + void updateCap_leaderUpdatesUnitsAndInvalidates() { + User leader = userWithId(20L, UUID.randomUUID()); + Team team = teamWithId(33L); + when(userRepository.findBySupabaseId(any())).thenReturn(Optional.of(leader)); + when(memberRepo.findPrimaryMembership(20L)) + .thenReturn(List.of(membership(team, leader, TeamRole.LEADER))); + when(policyRepo.findByTeamId(33L)).thenReturn(Optional.empty()); + + ResponseEntity resp = + controller.updateCap( + new UpdateCapRequest(40, false), jwtAuth(leader.getSupabaseId())); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.NO_CONTENT); + ArgumentCaptor saved = ArgumentCaptor.forClass(WalletPolicy.class); + verify(policyRepo).save(saved.capture()); + // Rate unknown in this test (docCapForMoney → empty) → legacy conversion fallback so + // the cap stays enforced rather than silently lifting. + assertThat(saved.getValue().getCapUnits()).isEqualTo(4000L); // 40 USD * 100 units/USD + assertThat(saved.getValue().getCapSourceMoney()).isEqualTo(4000L); // 40 USD == 4000 cents + verify(entitlementService, times(1)).invalidate(33L); + } + + @Test + void updateCap_withKnownRate_storesDerivedDocCap() { + User leader = userWithId(24L, UUID.randomUUID()); + Team team = teamWithId(36L); + when(userRepository.findBySupabaseId(any())).thenReturn(Optional.of(leader)); + when(memberRepo.findPrimaryMembership(24L)) + .thenReturn(List.of(membership(team, leader, TeamRole.LEADER))); + when(policyRepo.findByTeamId(36L)).thenReturn(Optional.empty()); + TeamBillingContext billing = subscribedBilling("sub_36", null, null); + when(billingService.forTeam(36L)).thenReturn(billing); + // $28 cap at $0.02/doc → 1400 paid documents/month (the grant is a separate pool). + when(billingService.docCapForMoney(billing, 2800L)).thenReturn(Optional.of(1400L)); + + ResponseEntity resp = + controller.updateCap( + new UpdateCapRequest(28, false), jwtAuth(leader.getSupabaseId())); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.NO_CONTENT); + ArgumentCaptor saved = ArgumentCaptor.forClass(WalletPolicy.class); + verify(policyRepo).save(saved.capture()); + assertThat(saved.getValue().getCapSourceMoney()).isEqualTo(2800L); + assertThat(saved.getValue().getCapUnits()).isEqualTo(1400L); + verify(entitlementService).invalidate(36L); + } + + @Test + void updateCap_noCapTrue_clearsCapUnits() { + User leader = userWithId(21L, UUID.randomUUID()); + Team team = teamWithId(34L); + when(userRepository.findBySupabaseId(any())).thenReturn(Optional.of(leader)); + when(memberRepo.findPrimaryMembership(21L)) + .thenReturn(List.of(membership(team, leader, TeamRole.LEADER))); + WalletPolicy existing = new WalletPolicy(); + existing.setTeamId(34L); + existing.setCapUnits(1000L); + existing.setCapSourceMoney(1000L); + when(policyRepo.findByTeamId(34L)).thenReturn(Optional.of(existing)); + + ResponseEntity resp = + controller.updateCap( + new UpdateCapRequest(0, true), jwtAuth(leader.getSupabaseId())); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.NO_CONTENT); + ArgumentCaptor saved = ArgumentCaptor.forClass(WalletPolicy.class); + verify(policyRepo).save(saved.capture()); + assertThat(saved.getValue().getCapUnits()).isNull(); + assertThat(saved.getValue().getCapSourceMoney()).isNull(); + verify(entitlementService).invalidate(34L); + } + + @Test + void updateCap_memberIsForbidden() { + User member = userWithId(22L, UUID.randomUUID()); + Team team = teamWithId(35L); + when(userRepository.findBySupabaseId(any())).thenReturn(Optional.of(member)); + when(memberRepo.findPrimaryMembership(22L)) + .thenReturn(List.of(membership(team, member, TeamRole.MEMBER))); + + ResponseEntity resp = + controller.updateCap( + new UpdateCapRequest(50, false), jwtAuth(member.getSupabaseId())); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.FORBIDDEN); + verify(policyRepo, never()).save(any()); + verify(entitlementService, never()).invalidate(any()); + } + + @Test + void updateCap_noTeam_isForbidden() { + User user = userWithId(23L, UUID.randomUUID()); + when(userRepository.findBySupabaseId(any())).thenReturn(Optional.of(user)); + when(memberRepo.findPrimaryMembership(23L)).thenReturn(List.of()); + + ResponseEntity resp = + controller.updateCap( + new UpdateCapRequest(10, false), jwtAuth(user.getSupabaseId())); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.FORBIDDEN); + verify(policyRepo, never()).save(any()); + } + + @Test + void updateCap_anonymousIs401() { + Authentication anon = + new AnonymousAuthenticationToken( + "k", + "anonymousUser", + List.of(new SimpleGrantedAuthority("ROLE_ANONYMOUS"))); + + ResponseEntity resp = controller.updateCap(new UpdateCapRequest(10, false), anon); + + assertThat(resp.getStatusCode()).isEqualTo(HttpStatus.UNAUTHORIZED); + verifyNoInteractions(policyRepo, entitlementService); + } + + // ----------------------------------------------------------------------------------------- + // Fixtures + // ----------------------------------------------------------------------------------------- + + private static User userWithId(Long id, UUID supabaseId) { + User u = new User(); + u.setId(id); + u.setSupabaseId(supabaseId); + return u; + } + + private static Team teamWithId(Long id) { + Team t = new Team(); + t.setId(id); + t.setName("t-" + id); + return t; + } + + private static TeamMembership membership(Team team, User user, TeamRole role) { + TeamMembership m = new TeamMembership(); + m.setTeam(team); + m.setUser(user); + m.setRole(role); + return m; + } + + private static EntitlementSnapshot snapshot(long spend, Long cap) { + LocalDateTime start = LocalDate.now().withDayOfMonth(1).atStartOfDay(); + LocalDateTime end = start.plusMonths(1); + return new EntitlementSnapshot( + EntitlementState.FULL, + FeatureSet.FULL, + List.of( + FeatureGate.OFFSITE_PROCESSING, + FeatureGate.AUTOMATION, + FeatureGate.AI_SUPPORT, + FeatureGate.CLIENT_SIDE), + spend, + cap, + start, + end, + false); + } + + private static Authentication jwtAuth(UUID supabaseId) { + Jwt jwt = + Jwt.withTokenValue("token") + .header("alg", "RS256") + .claim("sub", supabaseId.toString()) + .claim("email", "user@example.com") + .build(); + return new EnhancedJwtAuthenticationToken( + jwt, List.of(), "user@example.com", supabaseId.toString()); + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/payg/billing/TeamBillingServiceTest.java b/app/saas/src/test/java/stirling/software/saas/payg/billing/TeamBillingServiceTest.java new file mode 100644 index 0000000000..66e2eb074c --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/payg/billing/TeamBillingServiceTest.java @@ -0,0 +1,121 @@ +package stirling.software.saas.payg.billing; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.Mockito.when; + +import java.util.Optional; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.mockito.Mockito; + +import stirling.software.saas.payg.policy.PaygTeamExtensions; +import stirling.software.saas.payg.policy.PricingPolicy; +import stirling.software.saas.payg.policy.PricingPolicyService; +import stirling.software.saas.payg.repository.PaygTeamExtensionsRepository; +import stirling.software.saas.payg.repository.WalletPolicyRepository; +import stirling.software.saas.payg.stripe.StripeSubscriptionDao; + +/** + * Unit tests for {@link TeamBillingService#forTeam(Long)} — specifically the {@code subscribed} + * determination, which is the single switch both the wallet UI ({@code status}) and the entitlement + * gate read. + * + *

Regression focus: a team is subscribed iff {@code payg_subscription_id} is set. {@code + * payg_unlink_subscription} nulls that column on {@code customer.subscription.deleted} but + * deliberately keeps {@code stripe_customer_id} (for a future re-subscribe), so a cancelled team + * must read as free again. An earlier fallback treated customer-id presence as subscribed, + * which kept every team that ever subscribed pinned to subscribed forever — the cancelled-team bug + * these tests lock down. + */ +class TeamBillingServiceTest { + + private static final long TEAM_ID = 100L; + + private PaygTeamExtensionsRepository extensionsRepository; + private WalletPolicyRepository walletPolicyRepository; + private PricingPolicyService pricingPolicyService; + private StripeSubscriptionDao subscriptionDao; + private TeamBillingService service; + + @BeforeEach + void setUp() { + extensionsRepository = Mockito.mock(PaygTeamExtensionsRepository.class); + walletPolicyRepository = Mockito.mock(WalletPolicyRepository.class); + pricingPolicyService = Mockito.mock(PricingPolicyService.class); + subscriptionDao = Mockito.mock(StripeSubscriptionDao.class); + service = + new TeamBillingService( + extensionsRepository, + walletPolicyRepository, + pricingPolicyService, + subscriptionDao); + + // Default grant so the free-tier fields are populated; individual tests don't depend on it + // beyond the cancelled-team case below, which asserts the grant survives. + PricingPolicy policy = Mockito.mock(PricingPolicy.class); + when(policy.getFreeTierUnits()).thenReturn(500L); + when(pricingPolicyService.getEffectivePolicy(TEAM_ID)).thenReturn(policy); + } + + private PaygTeamExtensions ext(String subscriptionId, String customerId, long freeRemaining) { + PaygTeamExtensions ext = new PaygTeamExtensions(); + ext.setTeamId(TEAM_ID); + ext.setPaygSubscriptionId(subscriptionId); + ext.setStripeCustomerId(customerId); + ext.setFreeUnitsRemaining(freeRemaining); + return ext; + } + + @Test + void subscribed_whenSubscriptionIdPresent() { + when(extensionsRepository.findById(TEAM_ID)) + .thenReturn(Optional.of(ext("sub_123", "cus_123", 0L))); + + TeamBillingContext ctx = service.forTeam(TEAM_ID); + + assertThat(ctx.subscribed()).isTrue(); + assertThat(ctx.subscriptionId()).isEqualTo("sub_123"); + } + + /** + * The cancelled-subscription regression: after {@code payg_unlink_subscription} the + * subscription id is null but the Stripe customer id remains. The team must read as NOT + * subscribed (drops to the free-grant gate), and must still surface its remaining free grant. + */ + @Test + void notSubscribed_afterCancellation_whenOnlyCustomerIdRemains() { + when(extensionsRepository.findById(TEAM_ID)) + .thenReturn(Optional.of(ext(null, "cus_123", 120L))); + + TeamBillingContext ctx = service.forTeam(TEAM_ID); + + assertThat(ctx.subscribed()).isFalse(); + assertThat(ctx.subscriptionId()).isNull(); + // The free grant survives cancellation and is what now gates the team. + assertThat(ctx.freeGrantUnits()).isEqualTo(500L); + assertThat(ctx.freeRemainingUnits()).isEqualTo(120L); + // Not subscribed → no monthly paid-doc cap. + assertThat(ctx.monthlyCapDocUnits()).isNull(); + } + + @Test + void notSubscribed_whenNoSubscriptionAndNoCustomer() { + when(extensionsRepository.findById(TEAM_ID)).thenReturn(Optional.of(ext(null, null, 500L))); + + TeamBillingContext ctx = service.forTeam(TEAM_ID); + + assertThat(ctx.subscribed()).isFalse(); + assertThat(ctx.subscriptionId()).isNull(); + } + + @Test + void notSubscribed_whenNoExtensionRow() { + when(extensionsRepository.findById(TEAM_ID)).thenReturn(Optional.empty()); + + TeamBillingContext ctx = service.forTeam(TEAM_ID); + + assertThat(ctx.subscribed()).isFalse(); + assertThat(ctx.subscriptionId()).isNull(); + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/payg/cap/CapEvaluatorTest.java b/app/saas/src/test/java/stirling/software/saas/payg/cap/CapEvaluatorTest.java new file mode 100644 index 0000000000..2118f85794 --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/payg/cap/CapEvaluatorTest.java @@ -0,0 +1,150 @@ +package stirling.software.saas.payg.cap; + +import static org.assertj.core.api.Assertions.assertThat; + +import org.junit.jupiter.api.Test; + +import stirling.software.saas.payg.cap.CapEvaluator.Evaluation; +import stirling.software.saas.payg.model.EntitlementState; +import stirling.software.saas.payg.model.FeatureGate; +import stirling.software.saas.payg.model.FeatureSet; + +class CapEvaluatorTest { + + // --------------------------------------------------------------------------------------- + // evaluate(): single-axis cap evaluation + // --------------------------------------------------------------------------------------- + + @Test + void nullCap_returnsFullStateAndFullGates() { + Evaluation e = CapEvaluator.evaluate(1_000_000L, null, 80, 100, FeatureSet.MINIMAL); + assertThat(e.state()).isEqualTo(EntitlementState.FULL); + assertThat(e.featureSet()).isEqualTo(FeatureSet.FULL); + assertThat(e.enabledGates()) + .containsExactlyInAnyOrder( + FeatureGate.OFFSITE_PROCESSING, + FeatureGate.AUTOMATION, + FeatureGate.AI_SUPPORT, + FeatureGate.CLIENT_SIDE); + } + + @Test + void zeroCap_treatedAsUnlimitedForSafety() { + // Defensive: a zero cap would divide-by-zero. The guard treats it as null (FULL). + Evaluation e = CapEvaluator.evaluate(50L, 0L, 80, 100, FeatureSet.MINIMAL); + assertThat(e.state()).isEqualTo(EntitlementState.FULL); + } + + @Test + void wellBelowWarn_returnsFull() { + Evaluation e = CapEvaluator.evaluate(10L, 100L, 80, 100, FeatureSet.MINIMAL); + assertThat(e.state()).isEqualTo(EntitlementState.FULL); + assertThat(e.featureSet()).isEqualTo(FeatureSet.FULL); + } + + @Test + void exactlyAtWarnThreshold_returnsWarned() { + // 80% of 100 = 80; pct = 80 / 100 * 100 = 80 → ≥ warnAtPct=80 → WARNED + Evaluation e = CapEvaluator.evaluate(80L, 100L, 80, 100, FeatureSet.MINIMAL); + assertThat(e.state()).isEqualTo(EntitlementState.WARNED); + // Warn band keeps the FULL feature set — only flags the FE to show a banner. + assertThat(e.featureSet()).isEqualTo(FeatureSet.FULL); + assertThat(e.enabledGates()).hasSize(4); + } + + @Test + void betweenWarnAndDegrade_returnsWarned() { + Evaluation e = CapEvaluator.evaluate(95L, 100L, 80, 100, FeatureSet.MINIMAL); + assertThat(e.state()).isEqualTo(EntitlementState.WARNED); + assertThat(e.featureSet()).isEqualTo(FeatureSet.FULL); + } + + @Test + void exactlyAtDegradeThreshold_returnsDegradedWithConfiguredSet() { + Evaluation e = CapEvaluator.evaluate(100L, 100L, 80, 100, FeatureSet.MINIMAL); + assertThat(e.state()).isEqualTo(EntitlementState.DEGRADED); + assertThat(e.featureSet()).isEqualTo(FeatureSet.MINIMAL); + // MINIMAL keeps manual server tools (OFFSITE_PROCESSING) + client-side; only + // AUTOMATION + AI_SUPPORT are blocked. + assertThat(e.enabledGates()) + .containsExactlyInAnyOrder(FeatureGate.OFFSITE_PROCESSING, FeatureGate.CLIENT_SIDE); + } + + @Test + void overDegradeThreshold_returnsDegradedAndCannotEscalate() { + Evaluation e = CapEvaluator.evaluate(500L, 100L, 80, 100, FeatureSet.MINIMAL); + assertThat(e.state()).isEqualTo(EntitlementState.DEGRADED); + assertThat(e.featureSet()).isEqualTo(FeatureSet.MINIMAL); + } + + @Test + void degradedFeatureSetClientOnly_dropsEverythingButClientSide() { + Evaluation e = CapEvaluator.evaluate(100L, 100L, 80, 100, FeatureSet.CLIENT_ONLY); + assertThat(e.state()).isEqualTo(EntitlementState.DEGRADED); + assertThat(e.featureSet()).isEqualTo(FeatureSet.CLIENT_ONLY); + assertThat(e.enabledGates()).containsExactly(FeatureGate.CLIENT_SIDE); + } + + @Test + void nullDegradedFeatureSet_fallsBackToMinimal() { + Evaluation e = CapEvaluator.evaluate(100L, 100L, 80, 100, null); + assertThat(e.state()).isEqualTo(EntitlementState.DEGRADED); + assertThat(e.featureSet()).isEqualTo(FeatureSet.MINIMAL); + } + + @Test + void misconfiguredThresholds_treatedAsNoCap() { + // warn > degrade is malformed; the evaluator returns FULL rather than risk wrong + // degradation. The admin write path is supposed to validate this. + Evaluation e = CapEvaluator.evaluate(100L, 100L, 110, 100, FeatureSet.MINIMAL); + assertThat(e.state()).isEqualTo(EntitlementState.FULL); + + // negative warn → bad config → FULL + Evaluation neg = CapEvaluator.evaluate(100L, 100L, -10, 100, FeatureSet.MINIMAL); + assertThat(neg.state()).isEqualTo(EntitlementState.FULL); + + // zero degrade → bad config → FULL (would otherwise degrade on any spend) + Evaluation zero = CapEvaluator.evaluate(0L, 100L, 80, 0, FeatureSet.MINIMAL); + assertThat(zero.state()).isEqualTo(EntitlementState.FULL); + } + + @Test + void zeroSpend_returnsFullEvenIfCapTiny() { + Evaluation e = CapEvaluator.evaluate(0L, 1L, 80, 100, FeatureSet.MINIMAL); + assertThat(e.state()).isEqualTo(EntitlementState.FULL); + } + + @Test + void warnAtZeroPct_warnsImmediately() { + // Edge case: warn-at-0% means any spend → WARNED. Allowed even if quirky. + Evaluation e = CapEvaluator.evaluate(1L, 100L, 0, 100, FeatureSet.MINIMAL); + assertThat(e.state()).isEqualTo(EntitlementState.WARNED); + } + + // --------------------------------------------------------------------------------------- + // gatesFor(): mapping FeatureSet → declared gates + // --------------------------------------------------------------------------------------- + + @Test + void gatesFor_full_listsAllFour() { + assertThat(CapEvaluator.gatesFor(FeatureSet.FULL)) + .containsExactlyInAnyOrder( + FeatureGate.OFFSITE_PROCESSING, + FeatureGate.AUTOMATION, + FeatureGate.AI_SUPPORT, + FeatureGate.CLIENT_SIDE); + } + + @Test + void gatesFor_minimal_keepsOffsiteAndClientSide() { + // MINIMAL keeps manual server tools (OFFSITE_PROCESSING) + client-side; AUTOMATION + + // AI_SUPPORT are the only gates dropped on degrade. + assertThat(CapEvaluator.gatesFor(FeatureSet.MINIMAL)) + .containsExactlyInAnyOrder(FeatureGate.OFFSITE_PROCESSING, FeatureGate.CLIENT_SIDE); + } + + @Test + void gatesFor_null_returnsEmpty() { + assertThat(CapEvaluator.gatesFor(null)).isEmpty(); + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/payg/cap/RequiresFeatureAnnotationRolloutTest.java b/app/saas/src/test/java/stirling/software/saas/payg/cap/RequiresFeatureAnnotationRolloutTest.java new file mode 100644 index 0000000000..c2d5fcd8b7 --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/payg/cap/RequiresFeatureAnnotationRolloutTest.java @@ -0,0 +1,52 @@ +package stirling.software.saas.payg.cap; + +import static org.assertj.core.api.Assertions.assertThat; + +import org.junit.jupiter.api.Test; +import org.springframework.core.annotation.AnnotationUtils; + +import stirling.software.saas.ai.controller.AiCreateController; +import stirling.software.saas.ai.controller.AiCreateInternalController; +import stirling.software.saas.ai.controller.AiProxyController; +import stirling.software.saas.payg.model.FeatureGate; + +/** + * Annotation-rollout guard. The saas {@code PaygChargeInterceptor} reads class-level + * {@code @RequiresFeature} via {@link AnnotationUtils#findAnnotation(Class, Class)} to decide + * whether a request bills as {@code AI}, {@code AUTOMATION}, or falls through to the auth-derived + * default. These tests pin the gate on each AI surface so the classification can't silently regress + * to {@code BYPASSED} if someone strips the annotation while refactoring. + * + *

Out of scope: {@code PipelineController} (in core) and {@code PolicyController} (in + * proprietary) — neither module can import {@code @RequiresFeature} from saas without a forbidden + * upward dependency. Their automation classification is enforced via the {@code + * X-Stirling-Automation} header set unconditionally by {@code InternalApiClient.post}; see the + * dedicated test in that module. + */ +class RequiresFeatureAnnotationRolloutTest { + + @Test + void aiCreateController_isClassifiedAsAiSupport() { + RequiresFeature ann = + AnnotationUtils.findAnnotation(AiCreateController.class, RequiresFeature.class); + assertThat(ann).isNotNull(); + assertThat(ann.value()).containsExactly(FeatureGate.AI_SUPPORT); + } + + @Test + void aiCreateInternalController_isClassifiedAsAiSupport() { + RequiresFeature ann = + AnnotationUtils.findAnnotation( + AiCreateInternalController.class, RequiresFeature.class); + assertThat(ann).isNotNull(); + assertThat(ann.value()).containsExactly(FeatureGate.AI_SUPPORT); + } + + @Test + void aiProxyController_isClassifiedAsAiSupport() { + RequiresFeature ann = + AnnotationUtils.findAnnotation(AiProxyController.class, RequiresFeature.class); + assertThat(ann).isNotNull(); + assertThat(ann.value()).containsExactly(FeatureGate.AI_SUPPORT); + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/payg/charge/JobChargeServiceTest.java b/app/saas/src/test/java/stirling/software/saas/payg/charge/JobChargeServiceTest.java index 00e5843165..8b3ce50b6a 100644 --- a/app/saas/src/test/java/stirling/software/saas/payg/charge/JobChargeServiceTest.java +++ b/app/saas/src/test/java/stirling/software/saas/payg/charge/JobChargeServiceTest.java @@ -17,14 +17,18 @@ import java.time.LocalDateTime; import java.util.HashMap; import java.util.List; import java.util.Map; +import java.util.Optional; import java.util.UUID; +import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.io.TempDir; import org.mockito.ArgumentCaptor; import org.mockito.Mockito; import org.springframework.mock.web.MockMultipartFile; +import org.springframework.transaction.support.TransactionSynchronization; +import org.springframework.transaction.support.TransactionSynchronizationManager; import org.springframework.web.multipart.MultipartFile; import stirling.software.saas.payg.docs.DocumentClassifier; @@ -33,13 +37,24 @@ import stirling.software.saas.payg.job.JobContext; import stirling.software.saas.payg.job.JobService; import stirling.software.saas.payg.job.JoinOrOpenResult; import stirling.software.saas.payg.job.ProcessingJob; +import stirling.software.saas.payg.meter.PaygMeterReportingService; +import stirling.software.saas.payg.model.BillingCategory; import stirling.software.saas.payg.model.JobSource; import stirling.software.saas.payg.model.JobStatus; +import stirling.software.saas.payg.model.LedgerBucket; +import stirling.software.saas.payg.model.LedgerEntryType; import stirling.software.saas.payg.model.ProcessType; +import stirling.software.saas.payg.model.ReferenceType; +import stirling.software.saas.payg.model.ShadowChargeStatus; +import stirling.software.saas.payg.policy.PaygTeamExtensions; import stirling.software.saas.payg.policy.PricingPolicy; import stirling.software.saas.payg.policy.PricingPolicyService; import stirling.software.saas.payg.repository.PaygShadowChargeRepository; +import stirling.software.saas.payg.repository.PaygTeamExtensionsRepository; +import stirling.software.saas.payg.repository.ProcessingJobRepository; +import stirling.software.saas.payg.repository.WalletLedgerRepository; import stirling.software.saas.payg.shadow.PaygShadowCharge; +import stirling.software.saas.payg.wallet.WalletLedgerEntry; /** * Exercises {@link JobChargeService} as an orchestrator: policy lookup, step-limit resolution, @@ -51,6 +66,10 @@ class JobChargeServiceTest { private PricingPolicyService policyService; private DocumentClassifier classifier; private PaygShadowChargeRepository shadowRepo; + private ProcessingJobRepository jobRepo; + private PaygTeamExtensionsRepository teamExtRepo; + private PaygMeterReportingService meterReporter; + private WalletLedgerRepository ledgerRepo; private JobChargeService service; @BeforeEach @@ -59,7 +78,32 @@ class JobChargeServiceTest { policyService = Mockito.mock(PricingPolicyService.class); classifier = Mockito.mock(DocumentClassifier.class); shadowRepo = Mockito.mock(PaygShadowChargeRepository.class); - service = new JobChargeService(jobService, policyService, classifier, shadowRepo); + jobRepo = Mockito.mock(ProcessingJobRepository.class); + teamExtRepo = Mockito.mock(PaygTeamExtensionsRepository.class); + meterReporter = Mockito.mock(PaygMeterReportingService.class); + ledgerRepo = Mockito.mock(WalletLedgerRepository.class); + // findByIdForUpdate defaults to Optional.empty() (Mockito) → no free grant consumed unless + // a test stubs the sidecar row. The free split is decided at openProcess time now, not at + // close, so the meter tests just set free_units_consumed on the shadow row directly. + service = + new JobChargeService( + jobService, + policyService, + classifier, + shadowRepo, + jobRepo, + teamExtRepo, + meterReporter, + ledgerRepo); + } + + @AfterEach + void tearDown() { + // Defensive: a previous test could have left a fake synchronization registered. Clearing + // ensures isolation when tests run in any order. + if (TransactionSynchronizationManager.isSynchronizationActive()) { + TransactionSynchronizationManager.clear(); + } } @Test @@ -76,7 +120,12 @@ class JobChargeServiceTest { ChargeOutcome out = service.openProcess( - new ChargeContext(42L, 100L, JobSource.WEB, ProcessType.SINGLE_TOOL), + new ChargeContext( + 42L, + 100L, + JobSource.WEB, + ProcessType.SINGLE_TOOL, + BillingCategory.API), List.of(in)); assertThat(out.disposition()).isEqualTo(ChargeOutcome.Disposition.JOINED); @@ -85,6 +134,7 @@ class JobChargeServiceTest { verify(classifier, never()).classify(any(MultipartFile.class), any()); verify(classifier, never()).classify(anyList(), any()); verify(shadowRepo, never()).save(any()); + verify(ledgerRepo, never()).save(any()); } @Test @@ -97,19 +147,25 @@ class JobChargeServiceTest { .thenReturn(new JoinOrOpenResult(newJob, JoinOrOpenResult.Disposition.OPENED)); JobInput in = jobInput(tmp, "in.pdf", "application/pdf"); - when(classifier.classify(any(MultipartFile.class), eq(policy))) + when(classifier.classify(any(MultipartFile.class), any(Path.class), eq(policy))) .thenReturn(new DocumentMetrics(50, 1024L, "application/pdf", 4)); ChargeOutcome out = service.openProcess( - new ChargeContext(42L, 100L, JobSource.WEB, ProcessType.SINGLE_TOOL), + new ChargeContext( + 42L, + 100L, + JobSource.WEB, + ProcessType.SINGLE_TOOL, + BillingCategory.API), List.of(in)); assertThat(out.disposition()).isEqualTo(ChargeOutcome.Disposition.OPENED); assertThat(out.units()).isEqualTo(4); - // Single-file path called single-file classifier overload, not the list one. - verify(classifier, times(1)).classify(any(MultipartFile.class), eq(policy)); - verify(classifier, never()).classify(anyList(), eq(policy)); + // Single-file path called single-file classifier overload (with Path), not the list one. + verify(classifier, times(1)) + .classify(any(MultipartFile.class), any(Path.class), eq(policy)); + verify(classifier, never()).classify(anyList(), anyList(), eq(policy)); ArgumentCaptor captor = ArgumentCaptor.forClass(PaygShadowCharge.class); verify(shadowRepo).save(captor.capture()); @@ -118,12 +174,85 @@ class JobChargeServiceTest { assertThat(row.getJobId()).isEqualTo(newJob.getId()); assertThat(row.getPolicyId()).isEqualTo(policy.getId()); assertThat(row.getPaygUnits()).isEqualTo(4); - // Legacy comparison not wired yet — zeroed until CreditService is wired in the follow-up. + // Legacy comparison removed with the legacy credit engine — always zeroed. assertThat(row.getLegacyCreditsCharged()).isZero(); assertThat(row.getDiffPct()).isZero(); + // PAYG analytics axis: billing_category + job_source are copied from the context so the + // row stays self-describing after processing_job rows are pruned. + assertThat(row.getBillingCategory()).isEqualTo(BillingCategory.API); + assertThat(row.getJobSource()).isEqualTo(JobSource.WEB); // Job entity carries the classified docUnits so close-time receipts can render correctly. assertThat(newJob.getDocUnits()).isEqualTo(4); + + // Live ledger DEBIT mirrors the shadow row: same units, stored NEGATIVE per the + // wallet_ledger sign convention, tied back to the job via reference. + ArgumentCaptor ledgerCaptor = + ArgumentCaptor.forClass(WalletLedgerEntry.class); + verify(ledgerRepo).save(ledgerCaptor.capture()); + WalletLedgerEntry debit = ledgerCaptor.getValue(); + assertThat(debit.getTeamId()).isEqualTo(100L); + assertThat(debit.getActorUserId()).isEqualTo(42L); + assertThat(debit.getEntryType()).isEqualTo(LedgerEntryType.DEBIT); + assertThat(debit.getBucket()).isEqualTo(LedgerBucket.CYCLE); + assertThat(debit.getAmountUnits()).isEqualTo(-4); + assertThat(debit.getReferenceType()).isEqualTo(ReferenceType.JOB); + assertThat(debit.getReferenceId()).isEqualTo(newJob.getId().toString()); + assertThat(debit.getPolicyId()).isEqualTo(policy.getId()); + assertThat(debit.getBillingCategory()).isEqualTo(BillingCategory.API); + } + + @Test + void openProcess_bypassedCategory_writesShadowRowButNoLedgerDebit(@TempDir Path tmp) + throws IOException { + // Manual UI work is never billed: the shadow row still lands (comparison audit trail) + // but the live wallet_ledger must stay untouched. + PricingPolicy policy = stubPolicy(/*minCharge*/ 1, Map.of(JobSource.WEB, 10)); + when(policyService.getEffectivePolicy(100L)).thenReturn(policy); + ProcessingJob newJob = openJob(UUID.randomUUID()); + when(jobService.joinOrOpen(any(JobContext.class), anyList())) + .thenReturn(new JoinOrOpenResult(newJob, JoinOrOpenResult.Disposition.OPENED)); + when(classifier.classify(any(MultipartFile.class), any(Path.class), eq(policy))) + .thenReturn(new DocumentMetrics(1, 100L, "application/pdf", 1)); + + service.openProcess( + new ChargeContext( + 42L, + 100L, + JobSource.WEB, + ProcessType.SINGLE_TOOL, + BillingCategory.BYPASSED), + List.of(jobInput(tmp, "in.pdf", "application/pdf"))); + + verify(shadowRepo).save(any(PaygShadowCharge.class)); + verify(ledgerRepo, never()).save(any()); + } + + @Test + void openProcess_openedAutomationContext_writesShadowRowWithAutomationCategory( + @TempDir Path tmp) throws IOException { + PricingPolicy policy = + stubPolicy(/*minCharge*/ 1, Map.of(JobSource.WEB, 10, JobSource.PIPELINE, 20)); + when(policyService.getEffectivePolicy(100L)).thenReturn(policy); + ProcessingJob newJob = openJob(UUID.randomUUID()); + when(jobService.joinOrOpen(any(JobContext.class), anyList())) + .thenReturn(new JoinOrOpenResult(newJob, JoinOrOpenResult.Disposition.OPENED)); + when(classifier.classify(any(MultipartFile.class), any(Path.class), eq(policy))) + .thenReturn(new DocumentMetrics(1, 100L, "application/pdf", 1)); + + service.openProcess( + new ChargeContext( + 42L, + 100L, + JobSource.PIPELINE, + ProcessType.AUTOMATION, + BillingCategory.AUTOMATION), + List.of(jobInput(tmp, "in.pdf", "application/pdf"))); + + ArgumentCaptor captor = ArgumentCaptor.forClass(PaygShadowCharge.class); + verify(shadowRepo).save(captor.capture()); + assertThat(captor.getValue().getBillingCategory()).isEqualTo(BillingCategory.AUTOMATION); + assertThat(captor.getValue().getJobSource()).isEqualTo(JobSource.PIPELINE); } @Test @@ -136,17 +265,22 @@ class JobChargeServiceTest { JobInput a = jobInput(tmp, "a.pdf", "application/pdf"); JobInput b = jobInput(tmp, "b.pdf", "application/pdf"); - when(classifier.classify(anyList(), eq(policy))) + when(classifier.classify(anyList(), anyList(), eq(policy))) .thenReturn(new DocumentMetrics(100, 2048L, "application/pdf", 7)); ChargeOutcome out = service.openProcess( - new ChargeContext(42L, 100L, JobSource.WEB, ProcessType.AUTOMATION), + new ChargeContext( + 42L, + 100L, + JobSource.WEB, + ProcessType.AUTOMATION, + BillingCategory.AUTOMATION), List.of(a, b)); assertThat(out.units()).isEqualTo(7); - verify(classifier, never()).classify(any(MultipartFile.class), any()); - verify(classifier, times(1)).classify(anyList(), eq(policy)); + verify(classifier, never()).classify(any(MultipartFile.class), any(Path.class), any()); + verify(classifier, times(1)).classify(anyList(), anyList(), eq(policy)); } @Test @@ -157,17 +291,112 @@ class JobChargeServiceTest { ProcessingJob newJob = openJob(UUID.randomUUID()); when(jobService.joinOrOpen(any(JobContext.class), anyList())) .thenReturn(new JoinOrOpenResult(newJob, JoinOrOpenResult.Disposition.OPENED)); - when(classifier.classify(any(MultipartFile.class), eq(policy))) + when(classifier.classify(any(MultipartFile.class), any(Path.class), eq(policy))) .thenReturn(new DocumentMetrics(10, 1024L, "application/pdf", 2)); ChargeOutcome out = service.openProcess( - new ChargeContext(42L, 100L, JobSource.WEB, ProcessType.SINGLE_TOOL), + new ChargeContext( + 42L, + 100L, + JobSource.WEB, + ProcessType.SINGLE_TOOL, + BillingCategory.API), List.of(jobInput(tmp, "in.pdf", "application/pdf"))); assertThat(out.units()).isEqualTo(5); } + @Test + void openProcess_drawsFreeGrant_storesSplitAndDecrementsCounter(@TempDir Path tmp) + throws IOException { + // Team has 10 free units left; a 4-unit job draws all 4 from the grant. The shadow row + // records free_units_consumed = 4 (so nothing meters) and the counter drops to 6. + PricingPolicy policy = stubPolicy(1, Map.of(JobSource.WEB, 10)); + when(policyService.getEffectivePolicy(100L)).thenReturn(policy); + ProcessingJob newJob = openJob(UUID.randomUUID()); + when(jobService.joinOrOpen(any(JobContext.class), anyList())) + .thenReturn(new JoinOrOpenResult(newJob, JoinOrOpenResult.Disposition.OPENED)); + when(classifier.classify(any(MultipartFile.class), any(Path.class), eq(policy))) + .thenReturn(new DocumentMetrics(50, 1024L, "application/pdf", 4)); + + PaygTeamExtensions ext = new PaygTeamExtensions(); + ext.setTeamId(100L); + ext.setFreeUnitsRemaining(10L); + when(teamExtRepo.findByIdForUpdate(100L)).thenReturn(Optional.of(ext)); + + service.openProcess( + new ChargeContext( + 42L, 100L, JobSource.WEB, ProcessType.SINGLE_TOOL, BillingCategory.API), + List.of(jobInput(tmp, "in.pdf", "application/pdf"))); + + ArgumentCaptor captor = ArgumentCaptor.forClass(PaygShadowCharge.class); + verify(shadowRepo).save(captor.capture()); + assertThat(captor.getValue().getPaygUnits()).isEqualTo(4); + assertThat(captor.getValue().getFreeUnitsConsumed()).isEqualTo(4); + // Counter decremented in-place and persisted. + assertThat(ext.getFreeUnitsRemaining()).isEqualTo(6L); + verify(teamExtRepo).save(ext); + } + + @Test + void openProcess_grantStraddle_drawsRemainderFreeAndBillsTheRest(@TempDir Path tmp) + throws IOException { + // Only 3 free units left; a 10-unit job takes the 3 (counter → 0) and the other 7 bill. + PricingPolicy policy = stubPolicy(1, Map.of(JobSource.WEB, 10)); + when(policyService.getEffectivePolicy(100L)).thenReturn(policy); + ProcessingJob newJob = openJob(UUID.randomUUID()); + when(jobService.joinOrOpen(any(JobContext.class), anyList())) + .thenReturn(new JoinOrOpenResult(newJob, JoinOrOpenResult.Disposition.OPENED)); + when(classifier.classify(any(MultipartFile.class), any(Path.class), eq(policy))) + .thenReturn(new DocumentMetrics(50, 1024L, "application/pdf", 10)); + + PaygTeamExtensions ext = new PaygTeamExtensions(); + ext.setTeamId(100L); + ext.setFreeUnitsRemaining(3L); + when(teamExtRepo.findByIdForUpdate(100L)).thenReturn(Optional.of(ext)); + + service.openProcess( + new ChargeContext( + 42L, 100L, JobSource.WEB, ProcessType.SINGLE_TOOL, BillingCategory.API), + List.of(jobInput(tmp, "in.pdf", "application/pdf"))); + + ArgumentCaptor captor = ArgumentCaptor.forClass(PaygShadowCharge.class); + verify(shadowRepo).save(captor.capture()); + assertThat(captor.getValue().getPaygUnits()).isEqualTo(10); + assertThat(captor.getValue().getFreeUnitsConsumed()).isEqualTo(3); + assertThat(ext.getFreeUnitsRemaining()).isZero(); + verify(teamExtRepo).save(ext); + } + + @Test + void openProcess_exhaustedGrant_storesZeroFreeAndLeavesCounterUntouched(@TempDir Path tmp) + throws IOException { + // Grant already at 0 → nothing free, full units bill, counter not re-saved. + PricingPolicy policy = stubPolicy(1, Map.of(JobSource.WEB, 10)); + when(policyService.getEffectivePolicy(100L)).thenReturn(policy); + ProcessingJob newJob = openJob(UUID.randomUUID()); + when(jobService.joinOrOpen(any(JobContext.class), anyList())) + .thenReturn(new JoinOrOpenResult(newJob, JoinOrOpenResult.Disposition.OPENED)); + when(classifier.classify(any(MultipartFile.class), any(Path.class), eq(policy))) + .thenReturn(new DocumentMetrics(50, 1024L, "application/pdf", 5)); + + PaygTeamExtensions ext = new PaygTeamExtensions(); + ext.setTeamId(100L); + ext.setFreeUnitsRemaining(0L); + when(teamExtRepo.findByIdForUpdate(100L)).thenReturn(Optional.of(ext)); + + service.openProcess( + new ChargeContext( + 42L, 100L, JobSource.WEB, ProcessType.SINGLE_TOOL, BillingCategory.API), + List.of(jobInput(tmp, "in.pdf", "application/pdf"))); + + ArgumentCaptor captor = ArgumentCaptor.forClass(PaygShadowCharge.class); + verify(shadowRepo).save(captor.capture()); + assertThat(captor.getValue().getFreeUnitsConsumed()).isZero(); + verify(teamExtRepo, never()).save(any()); + } + @Test void openProcess_resolvesStepLimitFromPolicy_perJobSource(@TempDir Path tmp) throws IOException { @@ -178,11 +407,16 @@ class JobChargeServiceTest { ProcessingJob newJob = openJob(UUID.randomUUID()); when(jobService.joinOrOpen(any(JobContext.class), anyList())) .thenReturn(new JoinOrOpenResult(newJob, JoinOrOpenResult.Disposition.OPENED)); - when(classifier.classify(any(MultipartFile.class), eq(policy))) + when(classifier.classify(any(MultipartFile.class), any(Path.class), eq(policy))) .thenReturn(new DocumentMetrics(1, 100L, "application/pdf", 1)); service.openProcess( - new ChargeContext(42L, 100L, JobSource.PIPELINE, ProcessType.AUTOMATION), + new ChargeContext( + 42L, + 100L, + JobSource.PIPELINE, + ProcessType.AUTOMATION, + BillingCategory.AUTOMATION), List.of(jobInput(tmp, "in.pdf", "application/pdf"))); ArgumentCaptor ctxCaptor = ArgumentCaptor.forClass(JobContext.class); @@ -202,11 +436,16 @@ class JobChargeServiceTest { ProcessingJob newJob = openJob(UUID.randomUUID()); when(jobService.joinOrOpen(any(JobContext.class), anyList())) .thenReturn(new JoinOrOpenResult(newJob, JoinOrOpenResult.Disposition.OPENED)); - when(classifier.classify(any(MultipartFile.class), eq(policy))) + when(classifier.classify(any(MultipartFile.class), any(Path.class), eq(policy))) .thenReturn(new DocumentMetrics(1, 100L, "application/pdf", 1)); service.openProcess( - new ChargeContext(42L, 100L, JobSource.DESKTOP_APP, ProcessType.SINGLE_TOOL), + new ChargeContext( + 42L, + 100L, + JobSource.DESKTOP_APP, + ProcessType.SINGLE_TOOL, + BillingCategory.API), List.of(jobInput(tmp, "in.pdf", "application/pdf"))); ArgumentCaptor ctxCaptor = ArgumentCaptor.forClass(JobContext.class); @@ -220,12 +459,491 @@ class JobChargeServiceTest { () -> service.openProcess( new ChargeContext( - 42L, 100L, JobSource.WEB, ProcessType.SINGLE_TOOL), + 42L, + 100L, + JobSource.WEB, + ProcessType.SINGLE_TOOL, + BillingCategory.API), List.of())) .isInstanceOf(IllegalArgumentException.class) .hasMessageContaining("inputs must not be empty"); } + @Test + void markFirstStepFailed_flipsShadowRowAndClosesProcess() { + UUID jobId = UUID.randomUUID(); + PaygShadowCharge row = chargedShadowRow(jobId, 100L, 4, BillingCategory.API); + row.setPolicyId(7L); + when(shadowRepo.findFirstByJobIdOrderByIdAsc(jobId)).thenReturn(java.util.Optional.of(row)); + ProcessingJob job = openJob(jobId); + when(jobRepo.findById(jobId)).thenReturn(java.util.Optional.of(job)); + + service.markFirstStepFailed(jobId, "first-step-5xx:503"); + + assertThat(row.getStatus()).isEqualTo(ShadowChargeStatus.REFUNDED); + assertThat(row.getRefundedAt()).isNotNull(); + assertThat(row.getRefundReason()).isEqualTo("first-step-5xx:503"); + assertThat(job.getStatus()).isEqualTo(JobStatus.CLOSED); + assertThat(job.getClosedAt()).isNotNull(); + verify(shadowRepo).save(row); + verify(jobRepo).save(job); + + // Compensating REFUND entry: positive amount mirroring the openProcess debit, same JOB + // reference so the pair nets to zero for the period. + ArgumentCaptor ledgerCaptor = + ArgumentCaptor.forClass(WalletLedgerEntry.class); + verify(ledgerRepo).save(ledgerCaptor.capture()); + WalletLedgerEntry refund = ledgerCaptor.getValue(); + assertThat(refund.getTeamId()).isEqualTo(100L); + assertThat(refund.getEntryType()).isEqualTo(LedgerEntryType.REFUND); + assertThat(refund.getBucket()).isEqualTo(LedgerBucket.CYCLE); + assertThat(refund.getAmountUnits()).isEqualTo(4); + assertThat(refund.getReferenceType()).isEqualTo(ReferenceType.JOB); + assertThat(refund.getReferenceId()).isEqualTo(jobId.toString()); + assertThat(refund.getPolicyId()).isEqualTo(7L); + assertThat(refund.getBillingCategory()).isEqualTo(BillingCategory.API); + // This row consumed no free units, so the grant counter is left alone. + verify(teamExtRepo, never()).restoreFreeUnits(eq(100L), Mockito.anyLong()); + } + + @Test + void markFirstStepFailed_withFreeConsumed_restoresGrantToCounter() { + // A first-step failure is pre-meter: nothing billed to Stripe, but the grant moved at + // charge time. The refund must hand exactly those free units back to the team's counter. + UUID jobId = UUID.randomUUID(); + PaygShadowCharge row = chargedShadowRow(jobId, 100L, 10, 3, BillingCategory.API); + when(shadowRepo.findFirstByJobIdOrderByIdAsc(jobId)).thenReturn(Optional.of(row)); + when(jobRepo.findById(jobId)).thenReturn(Optional.of(openJob(jobId))); + + service.markFirstStepFailed(jobId, "first-step-5xx:503"); + + verify(teamExtRepo).restoreFreeUnits(100L, 3L); + } + + @Test + void markFirstStepFailed_alreadyRefunded_isNoOp() { + UUID jobId = UUID.randomUUID(); + PaygShadowCharge row = new PaygShadowCharge(); + row.setJobId(jobId); + row.setStatus(ShadowChargeStatus.REFUNDED); + when(shadowRepo.findFirstByJobIdOrderByIdAsc(jobId)).thenReturn(java.util.Optional.of(row)); + ProcessingJob job = openJob(jobId); + job.setStatus(JobStatus.CLOSED); + when(jobRepo.findById(jobId)).thenReturn(java.util.Optional.of(job)); + + service.markFirstStepFailed(jobId, "first-step-5xx:500"); + + verify(shadowRepo, never()).save(any()); + verify(jobRepo, never()).save(any()); + // No double-credit: the REFUND ledger entry only accompanies the CHARGED→REFUNDED flip. + verify(ledgerRepo, never()).save(any()); + } + + @Test + void markFirstStepFailed_noShadowRow_stillClosesProcess() { + UUID jobId = UUID.randomUUID(); + when(shadowRepo.findFirstByJobIdOrderByIdAsc(jobId)).thenReturn(java.util.Optional.empty()); + ProcessingJob job = openJob(jobId); + when(jobRepo.findById(jobId)).thenReturn(java.util.Optional.of(job)); + + service.markFirstStepFailed(jobId, "first-step-5xx:503"); + + assertThat(job.getStatus()).isEqualTo(JobStatus.CLOSED); + verify(jobRepo).save(job); + } + + @Test + void markFirstStepFailed_trimsLongRefundReason() { + UUID jobId = UUID.randomUUID(); + PaygShadowCharge row = new PaygShadowCharge(); + row.setStatus(ShadowChargeStatus.CHARGED); + when(shadowRepo.findFirstByJobIdOrderByIdAsc(jobId)).thenReturn(java.util.Optional.of(row)); + when(jobRepo.findById(jobId)).thenReturn(java.util.Optional.empty()); + + String oversized = "x".repeat(200); + service.markFirstStepFailed(jobId, oversized); + + assertThat(row.getRefundReason()).hasSize(128); + } + + @Test + void decrementStepCount_decrementsByOne() { + UUID jobId = UUID.randomUUID(); + ProcessingJob job = openJob(jobId); + job.setStepCount(3); + when(jobRepo.findById(jobId)).thenReturn(java.util.Optional.of(job)); + + service.decrementStepCount(jobId); + + assertThat(job.getStepCount()).isEqualTo(2); + verify(jobRepo).save(job); + } + + @Test + void decrementStepCount_floorAtOne_neverDrivesNegative() { + UUID jobId = UUID.randomUUID(); + ProcessingJob job = openJob(jobId); + job.setStepCount(1); + when(jobRepo.findById(jobId)).thenReturn(java.util.Optional.of(job)); + + service.decrementStepCount(jobId); + + assertThat(job.getStepCount()).isEqualTo(1); + verify(jobRepo, never()).save(any()); + } + + @Test + void decrementStepCount_missingJob_isNoOp() { + UUID jobId = UUID.randomUUID(); + when(jobRepo.findById(jobId)).thenReturn(java.util.Optional.empty()); + service.decrementStepCount(jobId); // must not throw + verify(jobRepo, never()).save(any()); + } + + // --- close() — meter reporting in afterCommit ----------------------------------------------- + + @Test + void close_subscribedTeam_postsMeterEventAfterCommit() { + UUID jobId = UUID.randomUUID(); + ProcessingJob job = openJob(jobId); + when(jobService.close(jobId)).thenReturn(job); + + PaygShadowCharge row = chargedShadowRow(jobId, 100L, 4, BillingCategory.API); + when(shadowRepo.findFirstByJobIdOrderByIdAsc(jobId)).thenReturn(Optional.of(row)); + + PaygTeamExtensions ext = new PaygTeamExtensions(); + ext.setTeamId(100L); + ext.setStripeCustomerId("cus_subscribed"); + ext.setPaygSubscriptionId("sub_test"); + when(teamExtRepo.findById(100L)).thenReturn(Optional.of(ext)); + // Row consumed no free units (free_units_consumed = 0) → all 4 are paid and meter. + + withTransactionSynchronization( + () -> { + service.close(jobId); + Mockito.verifyNoInteractions(meterReporter); + }); + + // afterCommit ran on tearDown of withTransactionSynchronization → meter posted now. + verify(meterReporter) + .recordUsage( + 100L, + "cus_subscribed", + 4, + BillingCategory.API, + "process:" + jobId + ":close", + jobId); + } + + @Test + void close_fullyFreeJob_doesNotPostMeterEvent() { + // The free-vs-paid split is fixed at charge time. A job whose 4 units all came from the + // one-time grant (free_units_consumed = 4) has nothing left to meter; the ledger DEBIT + // alone records the usage. + UUID jobId = UUID.randomUUID(); + ProcessingJob job = openJob(jobId); + when(jobService.close(jobId)).thenReturn(job); + + PaygShadowCharge row = chargedShadowRow(jobId, 100L, 4, 4, BillingCategory.API); + when(shadowRepo.findFirstByJobIdOrderByIdAsc(jobId)).thenReturn(Optional.of(row)); + + PaygTeamExtensions ext = new PaygTeamExtensions(); + ext.setTeamId(100L); + ext.setStripeCustomerId("cus_subscribed"); + ext.setPaygSubscriptionId("sub_test"); + when(teamExtRepo.findById(100L)).thenReturn(Optional.of(ext)); + + withTransactionSynchronization(() -> service.close(jobId)); + + Mockito.verifyNoInteractions(meterReporter); + } + + @Test + void close_partiallyFreeJob_metersOnlyThePaidPortion() { + // 20-unit job that drew 10 from the remaining grant at charge time (free_units_consumed = + // 10) → 10 paid units meter to Stripe. + UUID jobId = UUID.randomUUID(); + ProcessingJob job = openJob(jobId); + when(jobService.close(jobId)).thenReturn(job); + + PaygShadowCharge row = chargedShadowRow(jobId, 100L, 20, 10, BillingCategory.AUTOMATION); + when(shadowRepo.findFirstByJobIdOrderByIdAsc(jobId)).thenReturn(Optional.of(row)); + + PaygTeamExtensions ext = new PaygTeamExtensions(); + ext.setTeamId(100L); + ext.setStripeCustomerId("cus_subscribed"); + ext.setPaygSubscriptionId("sub_test"); + when(teamExtRepo.findById(100L)).thenReturn(Optional.of(ext)); + + withTransactionSynchronization(() -> service.close(jobId)); + + verify(meterReporter) + .recordUsage( + 100L, + "cus_subscribed", + 10, + BillingCategory.AUTOMATION, + "process:" + jobId + ":close", + jobId); + } + + @Test + void close_freeTierTeam_doesNotPostMeterEvent() { + UUID jobId = UUID.randomUUID(); + ProcessingJob job = openJob(jobId); + when(jobService.close(jobId)).thenReturn(job); + + PaygShadowCharge row = chargedShadowRow(jobId, 100L, 4, BillingCategory.API); + when(shadowRepo.findFirstByJobIdOrderByIdAsc(jobId)).thenReturn(Optional.of(row)); + + // No stripe_customer_id → treated as free-tier on this branch (pre-#6532). + PaygTeamExtensions ext = new PaygTeamExtensions(); + ext.setTeamId(100L); + ext.setStripeCustomerId(null); + when(teamExtRepo.findById(100L)).thenReturn(Optional.of(ext)); + + withTransactionSynchronization(() -> service.close(jobId)); + + Mockito.verifyNoInteractions(meterReporter); + } + + @Test + void close_noTeamExtensionsRow_doesNotPostMeterEvent() { + UUID jobId = UUID.randomUUID(); + ProcessingJob job = openJob(jobId); + when(jobService.close(jobId)).thenReturn(job); + + PaygShadowCharge row = chargedShadowRow(jobId, 100L, 4, BillingCategory.API); + when(shadowRepo.findFirstByJobIdOrderByIdAsc(jobId)).thenReturn(Optional.of(row)); + when(teamExtRepo.findById(100L)).thenReturn(Optional.empty()); + + withTransactionSynchronization(() -> service.close(jobId)); + + Mockito.verifyNoInteractions(meterReporter); + } + + @Test + void close_refundedShadowRow_doesNotPostMeterEvent() { + UUID jobId = UUID.randomUUID(); + ProcessingJob job = openJob(jobId); + when(jobService.close(jobId)).thenReturn(job); + + PaygShadowCharge row = chargedShadowRow(jobId, 100L, 4, BillingCategory.API); + row.setStatus(ShadowChargeStatus.REFUNDED); + when(shadowRepo.findFirstByJobIdOrderByIdAsc(jobId)).thenReturn(Optional.of(row)); + + withTransactionSynchronization(() -> service.close(jobId)); + + Mockito.verifyNoInteractions(meterReporter); + Mockito.verifyNoInteractions(teamExtRepo); + } + + @Test + void close_noShadowRow_doesNotPostMeterEvent() { + UUID jobId = UUID.randomUUID(); + ProcessingJob job = openJob(jobId); + when(jobService.close(jobId)).thenReturn(job); + + when(shadowRepo.findFirstByJobIdOrderByIdAsc(jobId)).thenReturn(Optional.empty()); + + withTransactionSynchronization(() -> service.close(jobId)); + + Mockito.verifyNoInteractions(meterReporter); + Mockito.verifyNoInteractions(teamExtRepo); + } + + @Test + void close_bypassedCategoryOnShadowRow_doesNotPostMeterEvent() { + // Defensive: BYPASSED rows shouldn't normally exist (the interceptor short-circuits + // before openProcess), but if one slips through we must not meter it. + UUID jobId = UUID.randomUUID(); + ProcessingJob job = openJob(jobId); + when(jobService.close(jobId)).thenReturn(job); + + PaygShadowCharge row = chargedShadowRow(jobId, 100L, 4, BillingCategory.BYPASSED); + when(shadowRepo.findFirstByJobIdOrderByIdAsc(jobId)).thenReturn(Optional.of(row)); + + withTransactionSynchronization(() -> service.close(jobId)); + + Mockito.verifyNoInteractions(meterReporter); + Mockito.verifyNoInteractions(teamExtRepo); + } + + @Test + void close_meterReporterThrowsRuntimeException_doesNotPropagate() { + // PaygMeterReportingService is documented to swallow; defence-in-depth in + // JobChargeService catches a misbehaving impl so the afterCommit hook can't poison the + // close flow. + UUID jobId = UUID.randomUUID(); + ProcessingJob job = openJob(jobId); + when(jobService.close(jobId)).thenReturn(job); + + PaygShadowCharge row = chargedShadowRow(jobId, 100L, 4, BillingCategory.AUTOMATION); + when(shadowRepo.findFirstByJobIdOrderByIdAsc(jobId)).thenReturn(Optional.of(row)); + + PaygTeamExtensions ext = new PaygTeamExtensions(); + ext.setTeamId(100L); + ext.setStripeCustomerId("cus_subscribed"); + ext.setPaygSubscriptionId("sub_test"); + when(teamExtRepo.findById(100L)).thenReturn(Optional.of(ext)); + + Mockito.doThrow(new RuntimeException("simulated meter failure")) + .when(meterReporter) + .recordUsage( + Mockito.anyLong(), + Mockito.anyString(), + Mockito.anyInt(), + Mockito.any(BillingCategory.class), + Mockito.anyString(), + Mockito.any(UUID.class)); + + // Should not throw — afterCommit's defence-in-depth wraps the call. + withTransactionSynchronization(() -> service.close(jobId)); + verify(meterReporter) + .recordUsage( + 100L, + "cus_subscribed", + 4, + BillingCategory.AUTOMATION, + "process:" + jobId + ":close", + jobId); + } + + @Test + void close_noActiveTransactionSync_skipsMeterPostButStillClosesJob() { + // Direct call without an outer @Transactional → no sync to register against. close() + // must still close the job; the meter post is implicitly deferred to whatever async path + // eventually wraps the call (or is never made, which is fine for ledger-only flows). + UUID jobId = UUID.randomUUID(); + ProcessingJob job = openJob(jobId); + when(jobService.close(jobId)).thenReturn(job); + + assertThat(TransactionSynchronizationManager.isSynchronizationActive()).isFalse(); + service.close(jobId); + + Mockito.verifyNoInteractions(meterReporter); + verify(jobService).close(jobId); + } + + // --- chargeStandalone() — non-file billable actions (e.g. AI Create) ----------------------- + + @Test + void chargeStandalone_subscribedTeam_chargesAndMetersPaidPortion() { + // AI Create-style charge: one standalone bookkeeping job, free split, ledger debit, meter. + long teamId = 100L; + PricingPolicy policy = stubPolicy(/*minCharge*/ 1, Map.of(JobSource.WEB, 10)); + when(policyService.getEffectivePolicy(teamId)).thenReturn(policy); + + UUID jobId = UUID.randomUUID(); + when(jobService.open(any(JobContext.class), eq(1))).thenReturn(openJob(jobId)); + when(jobService.close(jobId)).thenReturn(openJob(jobId)); + + // Subscribed, no free grant left → the whole unit is paid and meters. + PaygTeamExtensions ext = new PaygTeamExtensions(); + ext.setTeamId(teamId); + ext.setStripeCustomerId("cus_x"); + ext.setPaygSubscriptionId("sub_x"); + ext.setFreeUnitsRemaining(0L); + when(teamExtRepo.findByIdForUpdate(teamId)).thenReturn(Optional.of(ext)); + when(teamExtRepo.findById(teamId)).thenReturn(Optional.of(ext)); + when(shadowRepo.findFirstByJobIdOrderByIdAsc(jobId)) + .thenReturn(Optional.of(chargedShadowRow(jobId, teamId, 1, 0, BillingCategory.AI))); + + ChargeContext ctx = + new ChargeContext( + 7L, teamId, JobSource.WEB, ProcessType.SINGLE_TOOL, BillingCategory.AI); + ArgumentCaptor ledger = ArgumentCaptor.forClass(WalletLedgerEntry.class); + + withTransactionSynchronization(() -> service.chargeStandalone(ctx, 1)); + + verify(jobService).open(any(JobContext.class), eq(1)); + verify(jobService).close(jobId); + verify(ledgerRepo).save(ledger.capture()); + assertThat(ledger.getValue().getEntryType()).isEqualTo(LedgerEntryType.DEBIT); + assertThat(ledger.getValue().getAmountUnits()).isEqualTo(-1); + assertThat(ledger.getValue().getBillingCategory()).isEqualTo(BillingCategory.AI); + verify(shadowRepo).save(any(PaygShadowCharge.class)); + // Paid portion (1) metered to Stripe after commit, keyed by the standard process key. + verify(meterReporter) + .recordUsage( + eq(teamId), + eq("cus_x"), + eq(1), + eq(BillingCategory.AI), + eq("process:" + jobId + ":close"), + eq(jobId)); + } + + @Test + void chargeStandalone_freeTeamWithGrant_drawsGrantAndDoesNotMeter() { + long teamId = 100L; + PricingPolicy policy = stubPolicy(1, Map.of(JobSource.WEB, 10)); + when(policyService.getEffectivePolicy(teamId)).thenReturn(policy); + + UUID jobId = UUID.randomUUID(); + when(jobService.open(any(JobContext.class), eq(1))).thenReturn(openJob(jobId)); + when(jobService.close(jobId)).thenReturn(openJob(jobId)); + + // Free grant available, no subscription → the unit is drawn from the grant, nothing meters. + PaygTeamExtensions ext = new PaygTeamExtensions(); + ext.setTeamId(teamId); + ext.setFreeUnitsRemaining(50L); + when(teamExtRepo.findByIdForUpdate(teamId)).thenReturn(Optional.of(ext)); + when(teamExtRepo.findById(teamId)).thenReturn(Optional.of(ext)); + when(shadowRepo.findFirstByJobIdOrderByIdAsc(jobId)) + .thenReturn(Optional.of(chargedShadowRow(jobId, teamId, 1, 1, BillingCategory.AI))); + + ChargeContext ctx = + new ChargeContext( + 7L, teamId, JobSource.WEB, ProcessType.SINGLE_TOOL, BillingCategory.AI); + + withTransactionSynchronization(() -> service.chargeStandalone(ctx, 1)); + + assertThat(ext.getFreeUnitsRemaining()).isEqualTo(49L); + verify(meterReporter, never()) + .recordUsage(any(), any(), Mockito.anyInt(), any(), any(), any()); + } + + @Test + void chargeStandalone_bypassedCategory_throws() { + ChargeContext ctx = + new ChargeContext( + 7L, 100L, JobSource.WEB, ProcessType.SINGLE_TOOL, BillingCategory.BYPASSED); + assertThatThrownBy(() -> service.chargeStandalone(ctx, 1)) + .isInstanceOf(IllegalArgumentException.class); + } + + private static void withTransactionSynchronization(Runnable body) { + TransactionSynchronizationManager.initSynchronization(); + try { + body.run(); + // Drain registered synchronizations to simulate a successful commit. + for (TransactionSynchronization sync : + TransactionSynchronizationManager.getSynchronizations()) { + sync.afterCommit(); + } + } finally { + TransactionSynchronizationManager.clear(); + } + } + + private static PaygShadowCharge chargedShadowRow( + UUID jobId, Long teamId, int units, BillingCategory category) { + return chargedShadowRow(jobId, teamId, units, 0, category); + } + + private static PaygShadowCharge chargedShadowRow( + UUID jobId, Long teamId, int units, int freeUnitsConsumed, BillingCategory category) { + PaygShadowCharge row = new PaygShadowCharge(); + row.setJobId(jobId); + row.setTeamId(teamId); + row.setPaygUnits(units); + row.setFreeUnitsConsumed(freeUnitsConsumed); + row.setStatus(ShadowChargeStatus.CHARGED); + row.setBillingCategory(category); + return row; + } + // --- helpers -------------------------------------------------------------------------------- private static PricingPolicy stubPolicy(int minCharge, Map stepLimits) { diff --git a/app/saas/src/test/java/stirling/software/saas/payg/entitlement/EntitlementGuardTest.java b/app/saas/src/test/java/stirling/software/saas/payg/entitlement/EntitlementGuardTest.java new file mode 100644 index 0000000000..d66e4b17fe --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/payg/entitlement/EntitlementGuardTest.java @@ -0,0 +1,619 @@ +package stirling.software.saas.payg.entitlement; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.lang.reflect.Method; +import java.time.Instant; +import java.time.LocalDateTime; +import java.util.HashMap; +import java.util.List; +import java.util.Map; +import java.util.Optional; +import java.util.UUID; + +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.mockito.Mockito; +import org.springframework.http.MediaType; +import org.springframework.mock.web.MockHttpServletRequest; +import org.springframework.mock.web.MockHttpServletResponse; +import org.springframework.security.authentication.AnonymousAuthenticationToken; +import org.springframework.security.core.authority.SimpleGrantedAuthority; +import org.springframework.security.core.context.SecurityContextHolder; +import org.springframework.security.oauth2.jwt.Jwt; +import org.springframework.web.method.HandlerMethod; + +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; + +import io.micrometer.core.instrument.MeterRegistry; +import io.micrometer.core.instrument.simple.SimpleMeterRegistry; + +import stirling.software.common.annotations.AutoJobPostMapping; +import stirling.software.proprietary.model.Team; +import stirling.software.proprietary.security.database.repository.UserRepository; +import stirling.software.proprietary.security.model.ApiKeyAuthenticationToken; +import stirling.software.proprietary.security.model.User; +import stirling.software.saas.payg.cap.RequiresFeature; +import stirling.software.saas.payg.model.EntitlementState; +import stirling.software.saas.payg.model.FeatureGate; +import stirling.software.saas.payg.model.FeatureSet; +import stirling.software.saas.security.EnhancedJwtAuthenticationToken; + +/** + * Pure-Mockito tests for {@link EntitlementGuard}. Covers the four decision-matrix cells: anonymous + * billable → 401, anonymous manual → pass, authenticated FULL → pass, authenticated DEGRADED for a + * billable route → 402. + */ +class EntitlementGuardTest { + + private EntitlementService entitlementService; + private UserRepository userRepository; + private MeterRegistry meterRegistry; + private EntitlementGuard guard; + + private final ObjectMapper json = new ObjectMapper(); + + @BeforeEach + void setUp() { + entitlementService = Mockito.mock(EntitlementService.class); + userRepository = Mockito.mock(UserRepository.class); + meterRegistry = new SimpleMeterRegistry(); + guard = new EntitlementGuard(entitlementService, userRepository, meterRegistry); + SecurityContextHolder.clearContext(); + } + + @AfterEach + void tearDown() { + SecurityContextHolder.clearContext(); + } + + // --------------------------------------------------------------------------------------- + // Scope: non-AutoJobPostMapping routes skip the guard entirely + // --------------------------------------------------------------------------------------- + + @Test + void nonHandlerMethod_passesThrough() throws Exception { + MockHttpServletRequest req = new MockHttpServletRequest(); + MockHttpServletResponse res = new MockHttpServletResponse(); + + boolean proceed = guard.preHandle(req, res, "someRawHandler"); + + assertThat(proceed).isTrue(); + assertThat(res.getStatus()).isEqualTo(200); + } + + @Test + void routeWithNeitherAnnotation_isSkipped() throws Exception { + HandlerMethod hm = handlerFor("plainEndpoint"); + MockHttpServletRequest req = new MockHttpServletRequest(); + MockHttpServletResponse res = new MockHttpServletResponse(); + + boolean proceed = guard.preHandle(req, res, hm); + + assertThat(proceed).isTrue(); + Mockito.verifyNoInteractions(entitlementService, userRepository); + } + + @Test + void routeWithRequiresFeatureButNoAutoJobPostMapping_isInScope() throws Exception { + // Regression: @RequiresFeature alone (e.g. JSON-bodied AI controllers) must be in-scope. + // Previously the guard short-circuited unless @AutoJobPostMapping was present, which meant + // a team without AI entitlement could hit /api/v1/ai/* freely. + UUID supabaseId = UUID.randomUUID(); + SecurityContextHolder.getContext().setAuthentication(jwtAuth(supabaseId)); + when(userRepository.findBySupabaseId(supabaseId)) + .thenReturn(Optional.of(userWithTeam(7L, 42L))); + when(entitlementService.getSnapshot(42L)).thenReturn(degradedSnapshot()); + + HandlerMethod hm = handlerFor("aiOnlyNoAutoJobPostMapping"); + MockHttpServletRequest req = new MockHttpServletRequest(); + MockHttpServletResponse res = new MockHttpServletResponse(); + + boolean proceed = guard.preHandle(req, res, hm); + + assertThat(proceed).isFalse(); + assertThat(res.getStatus()).isEqualTo(402); + JsonNode body = json.readTree(res.getContentAsByteArray()); + assertThat(body.get("error").asText()).isEqualTo("FEATURE_DEGRADED"); + assertThat(body.get("missingGates").get(0).asText()).isEqualTo("AI_SUPPORT"); + verify(entitlementService).getSnapshot(42L); + } + + @Test + void routeWithClassLevelRequiresFeatureOnly_isInScope() throws Exception { + // Mirrors AiCreateController shape: @RequiresFeature lives on the @RestController class, + // not the method. Must still be picked up by the guard. + SecurityContextHolder.getContext() + .setAuthentication( + new AnonymousAuthenticationToken( + "key", + "anonymousUser", + List.of(new SimpleGrantedAuthority("ROLE_ANONYMOUS")))); + + HandlerMethod hm = handlerForClassLevel("classLevelAiMethod"); + MockHttpServletRequest req = new MockHttpServletRequest(); + MockHttpServletResponse res = new MockHttpServletResponse(); + + boolean proceed = guard.preHandle(req, res, hm); + + assertThat(proceed).isFalse(); + assertThat(res.getStatus()).isEqualTo(401); + JsonNode body = json.readTree(res.getContentAsByteArray()); + assertThat(body.get("error").asText()).isEqualTo("SIGNUP_REQUIRED"); + assertThat(body.get("category").asText()).isEqualTo("AI"); + } + + @Test + void aiToolRoute_noAnnotation_isInScopeAndGatedOnAiSupport() throws Exception { + // /api/v1/ai/tools/** controllers live in the proprietary module and carry no + // @RequiresFeature; the guard recognises them by path and gates on AI_SUPPORT. A degraded + // team is 402'd even though the handler has no annotation. + UUID supabaseId = UUID.randomUUID(); + SecurityContextHolder.getContext().setAuthentication(jwtAuth(supabaseId)); + when(userRepository.findBySupabaseId(supabaseId)) + .thenReturn(Optional.of(userWithTeam(7L, 42L))); + when(entitlementService.getSnapshot(42L)).thenReturn(degradedSnapshot()); + + HandlerMethod hm = handlerFor("plainEndpoint"); // no annotations + MockHttpServletRequest req = new MockHttpServletRequest(); + req.setRequestURI("/api/v1/ai/tools/pdf-comment-agent"); + MockHttpServletResponse res = new MockHttpServletResponse(); + + boolean proceed = guard.preHandle(req, res, hm); + + assertThat(proceed).isFalse(); + assertThat(res.getStatus()).isEqualTo(402); + JsonNode body = json.readTree(res.getContentAsByteArray()); + assertThat(body.get("error").asText()).isEqualTo("FEATURE_DEGRADED"); + assertThat(body.get("missingGates").get(0).asText()).isEqualTo("AI_SUPPORT"); + verify(entitlementService).getSnapshot(42L); + } + + @Test + void aiToolRoute_anonymous_returns401WithAiCategory() throws Exception { + SecurityContextHolder.getContext() + .setAuthentication( + new AnonymousAuthenticationToken( + "key", + "anonymousUser", + List.of(new SimpleGrantedAuthority("ROLE_ANONYMOUS")))); + + HandlerMethod hm = handlerFor("plainEndpoint"); + MockHttpServletRequest req = new MockHttpServletRequest(); + req.setRequestURI("/api/v1/ai/tools/math-auditor-agent"); + MockHttpServletResponse res = new MockHttpServletResponse(); + + boolean proceed = guard.preHandle(req, res, hm); + + assertThat(proceed).isFalse(); + assertThat(res.getStatus()).isEqualTo(401); + JsonNode body = json.readTree(res.getContentAsByteArray()); + assertThat(body.get("error").asText()).isEqualTo("SIGNUP_REQUIRED"); + assertThat(body.get("category").asText()).isEqualTo("AI"); + } + + // --------------------------------------------------------------------------------------- + // Anonymous user + // --------------------------------------------------------------------------------------- + + @Test + void anonymousUser_billableRoute_returns401SignupRequired() throws Exception { + SecurityContextHolder.getContext() + .setAuthentication( + new AnonymousAuthenticationToken( + "key", + "anonymousUser", + List.of(new SimpleGrantedAuthority("ROLE_ANONYMOUS")))); + + HandlerMethod hm = handlerFor("automationOnly"); + MockHttpServletRequest req = new MockHttpServletRequest(); + MockHttpServletResponse res = new MockHttpServletResponse(); + + boolean proceed = guard.preHandle(req, res, hm); + + assertThat(proceed).isFalse(); + assertThat(res.getStatus()).isEqualTo(401); + assertThat(res.getContentType()).startsWith(MediaType.APPLICATION_JSON_VALUE); + JsonNode body = json.readTree(res.getContentAsByteArray()); + assertThat(body.get("error").asText()).isEqualTo("SIGNUP_REQUIRED"); + assertThat(body.get("category").asText()).isEqualTo("AUTOMATION"); + Mockito.verifyNoInteractions(entitlementService); + } + + @Test + void anonymousUser_aiRoute_returns401WithAiCategory() throws Exception { + SecurityContextHolder.getContext() + .setAuthentication( + new AnonymousAuthenticationToken( + "key", + "anonymousUser", + List.of(new SimpleGrantedAuthority("ROLE_ANONYMOUS")))); + + HandlerMethod hm = handlerFor("aiOnly"); + MockHttpServletRequest req = new MockHttpServletRequest(); + MockHttpServletResponse res = new MockHttpServletResponse(); + + boolean proceed = guard.preHandle(req, res, hm); + + assertThat(proceed).isFalse(); + JsonNode body = json.readTree(res.getContentAsByteArray()); + assertThat(body.get("category").asText()).isEqualTo("AI"); + } + + @Test + void anonymousUser_manualTool_passesThroughUnbilled() throws Exception { + SecurityContextHolder.getContext() + .setAuthentication( + new AnonymousAuthenticationToken( + "key", + "anonymousUser", + List.of(new SimpleGrantedAuthority("ROLE_ANONYMOUS")))); + + HandlerMethod hm = handlerFor("manualTool"); + MockHttpServletRequest req = new MockHttpServletRequest(); + MockHttpServletResponse res = new MockHttpServletResponse(); + + boolean proceed = guard.preHandle(req, res, hm); + + assertThat(proceed).isTrue(); + assertThat(res.getStatus()).isEqualTo(200); + Mockito.verifyNoInteractions(entitlementService); + } + + // --------------------------------------------------------------------------------------- + // Authenticated user — FULL vs DEGRADED + // --------------------------------------------------------------------------------------- + + @Test + void authenticatedUser_fullState_passesThrough() throws Exception { + UUID supabaseId = UUID.randomUUID(); + SecurityContextHolder.getContext().setAuthentication(jwtAuth(supabaseId)); + when(userRepository.findBySupabaseId(supabaseId)) + .thenReturn(Optional.of(userWithTeam(7L, 42L))); + when(entitlementService.getSnapshot(42L)).thenReturn(fullSnapshot()); + + HandlerMethod hm = handlerFor("automationOnly"); + MockHttpServletRequest req = new MockHttpServletRequest(); + MockHttpServletResponse res = new MockHttpServletResponse(); + + boolean proceed = guard.preHandle(req, res, hm); + + assertThat(proceed).isTrue(); + assertThat(res.getStatus()).isEqualTo(200); + } + + @Test + void authenticatedUser_degradedAndBillable_returns402() throws Exception { + UUID supabaseId = UUID.randomUUID(); + SecurityContextHolder.getContext().setAuthentication(jwtAuth(supabaseId)); + when(userRepository.findBySupabaseId(supabaseId)) + .thenReturn(Optional.of(userWithTeam(7L, 42L))); + when(entitlementService.getSnapshot(42L)).thenReturn(degradedSnapshot()); + + HandlerMethod hm = handlerFor("automationOnly"); + MockHttpServletRequest req = new MockHttpServletRequest(); + MockHttpServletResponse res = new MockHttpServletResponse(); + + boolean proceed = guard.preHandle(req, res, hm); + + assertThat(proceed).isFalse(); + assertThat(res.getStatus()).isEqualTo(402); + assertThat(res.getContentType()).startsWith(MediaType.APPLICATION_JSON_VALUE); + JsonNode body = json.readTree(res.getContentAsByteArray()); + assertThat(body.get("error").asText()).isEqualTo("FEATURE_DEGRADED"); + assertThat(body.get("state").asText()).isEqualTo("DEGRADED"); + // subscribed drives the client's modal choice (free-limit vs spend-cap). + assertThat(body.get("subscribed").asBoolean()).isFalse(); + assertThat(body.get("capUnits").asLong()).isEqualTo(500L); + assertThat(body.get("spendUnits").asLong()).isEqualTo(500L); + assertThat(body.get("missingGates").isArray()).isTrue(); + assertThat(body.get("missingGates").get(0).asText()).isEqualTo("AUTOMATION"); + } + + @Test + void authenticatedUser_degradedButManualTool_passesThrough() throws Exception { + // KEY assertion: DEGRADED+MINIMAL must still allow manual server tools (OFFSITE_PROCESSING) + // — that's the whole point of moving OFFSITE into MINIMAL. + UUID supabaseId = UUID.randomUUID(); + SecurityContextHolder.getContext().setAuthentication(jwtAuth(supabaseId)); + when(userRepository.findBySupabaseId(supabaseId)) + .thenReturn(Optional.of(userWithTeam(7L, 42L))); + when(entitlementService.getSnapshot(42L)).thenReturn(degradedSnapshot()); + + HandlerMethod hm = handlerFor("manualTool"); + MockHttpServletRequest req = new MockHttpServletRequest(); + MockHttpServletResponse res = new MockHttpServletResponse(); + + boolean proceed = guard.preHandle(req, res, hm); + + assertThat(proceed).isTrue(); + assertThat(res.getStatus()).isEqualTo(200); + } + + @Test + void authenticatedUser_aiRouteDegraded_returns402() throws Exception { + UUID supabaseId = UUID.randomUUID(); + SecurityContextHolder.getContext().setAuthentication(jwtAuth(supabaseId)); + when(userRepository.findBySupabaseId(supabaseId)) + .thenReturn(Optional.of(userWithTeam(7L, 42L))); + when(entitlementService.getSnapshot(42L)).thenReturn(degradedSnapshot()); + + HandlerMethod hm = handlerFor("aiOnly"); + MockHttpServletRequest req = new MockHttpServletRequest(); + MockHttpServletResponse res = new MockHttpServletResponse(); + + boolean proceed = guard.preHandle(req, res, hm); + + assertThat(proceed).isFalse(); + assertThat(res.getStatus()).isEqualTo(402); + JsonNode body = json.readTree(res.getContentAsByteArray()); + assertThat(body.get("missingGates").get(0).asText()).isEqualTo("AI_SUPPORT"); + } + + // --------------------------------------------------------------------------------------- + // API-key calls: billable, hard-stop when degraded (allowance/cap reached) + // --------------------------------------------------------------------------------------- + + @Test + void apiKeyCall_degradedFreeTeam_returns402AndAdvisesSubscribe() throws Exception { + // A plain server tool (OFFSITE_PROCESSING) reached via API key: the gate survives DEGRADED, + // so the gate loop would wave it through — but API usage is billable and must hard-stop + // once the free allowance is spent. + SecurityContextHolder.getContext().setAuthentication(apiKeyAuth(42L)); + when(entitlementService.getSnapshot(42L)).thenReturn(degradedSnapshot(false)); + + HandlerMethod hm = handlerFor("manualTool"); // OFFSITE default gate + MockHttpServletRequest req = new MockHttpServletRequest(); + MockHttpServletResponse res = new MockHttpServletResponse(); + + boolean proceed = guard.preHandle(req, res, hm); + + assertThat(proceed).isFalse(); + assertThat(res.getStatus()).isEqualTo(402); + JsonNode body = json.readTree(res.getContentAsByteArray()); + assertThat(body.get("error").asText()).isEqualTo("PAYG_LIMIT_REACHED"); + assertThat(body.get("subscribed").asBoolean()).isFalse(); + assertThat(body.get("message").asText()).contains("Subscribe"); + } + + @Test + void apiKeyCall_degradedSubscribedTeam_returns402AndAdvisesCap() throws Exception { + SecurityContextHolder.getContext().setAuthentication(apiKeyAuth(42L)); + when(entitlementService.getSnapshot(42L)).thenReturn(degradedSnapshot(true)); + + HandlerMethod hm = handlerFor("manualTool"); + MockHttpServletRequest req = new MockHttpServletRequest(); + MockHttpServletResponse res = new MockHttpServletResponse(); + + boolean proceed = guard.preHandle(req, res, hm); + + assertThat(proceed).isFalse(); + assertThat(res.getStatus()).isEqualTo(402); + JsonNode body = json.readTree(res.getContentAsByteArray()); + assertThat(body.get("error").asText()).isEqualTo("PAYG_LIMIT_REACHED"); + assertThat(body.get("subscribed").asBoolean()).isTrue(); + assertThat(body.get("message").asText()).contains("cap"); + } + + @Test + void apiKeyCall_withinAllowance_passes() throws Exception { + // Not degraded → API usage under the free allowance proceeds normally. + SecurityContextHolder.getContext().setAuthentication(apiKeyAuth(42L)); + when(entitlementService.getSnapshot(42L)).thenReturn(fullSnapshot()); + + HandlerMethod hm = handlerFor("manualTool"); + MockHttpServletRequest req = new MockHttpServletRequest(); + MockHttpServletResponse res = new MockHttpServletResponse(); + + boolean proceed = guard.preHandle(req, res, hm); + + assertThat(proceed).isTrue(); + assertThat(res.getStatus()).isEqualTo(200); + } + + @Test + void jwtWebManualTool_degraded_stillPasses() throws Exception { + // The "allow JWT web tool usage" guarantee: a web (JWT) user running an everyday server + // tool (OFFSITE_PROCESSING) is NOT blocked when the team is over allowance — only billable + // API/AI/automation calls hard-stop. (Web manual calls are BillingCategory.BYPASSED.) + UUID supabaseId = UUID.randomUUID(); + SecurityContextHolder.getContext().setAuthentication(jwtAuth(supabaseId)); + when(userRepository.findBySupabaseId(supabaseId)) + .thenReturn(Optional.of(userWithTeam(7L, 42L))); + when(entitlementService.getSnapshot(42L)).thenReturn(degradedSnapshot(false)); + + HandlerMethod hm = handlerFor("manualTool"); // OFFSITE — survives DEGRADED + MockHttpServletRequest req = new MockHttpServletRequest(); + MockHttpServletResponse res = new MockHttpServletResponse(); + + boolean proceed = guard.preHandle(req, res, hm); + + assertThat(proceed).isTrue(); + assertThat(res.getStatus()).isEqualTo(200); + } + + // --------------------------------------------------------------------------------------- + // Fail-open + // --------------------------------------------------------------------------------------- + + @Test + void snapshotLookupThrows_failsOpenAndPasses() throws Exception { + UUID supabaseId = UUID.randomUUID(); + SecurityContextHolder.getContext().setAuthentication(jwtAuth(supabaseId)); + when(userRepository.findBySupabaseId(supabaseId)) + .thenReturn(Optional.of(userWithTeam(7L, 42L))); + when(entitlementService.getSnapshot(42L)) + .thenThrow(new RuntimeException("transient DB outage")); + + HandlerMethod hm = handlerFor("automationOnly"); + MockHttpServletRequest req = new MockHttpServletRequest(); + MockHttpServletResponse res = new MockHttpServletResponse(); + + boolean proceed = guard.preHandle(req, res, hm); + + assertThat(proceed).isTrue(); + assertThat(res.getStatus()).isEqualTo(200); + } + + @Test + void noTeam_passesThrough() throws Exception { + UUID supabaseId = UUID.randomUUID(); + SecurityContextHolder.getContext().setAuthentication(jwtAuth(supabaseId)); + User u = new User(); + u.setId(7L); + u.setTeam(null); + when(userRepository.findBySupabaseId(supabaseId)).thenReturn(Optional.of(u)); + + HandlerMethod hm = handlerFor("automationOnly"); + MockHttpServletRequest req = new MockHttpServletRequest(); + MockHttpServletResponse res = new MockHttpServletResponse(); + + boolean proceed = guard.preHandle(req, res, hm); + + assertThat(proceed).isTrue(); + verify(entitlementService, never()).getSnapshot(any()); + } + + // --------------------------------------------------------------------------------------- + // resolveRequiredGates fallback + // --------------------------------------------------------------------------------------- + + @Test + void resolveRequiredGates_noAnnotation_defaultsToOffsiteProcessing() throws Exception { + HandlerMethod hm = handlerFor("manualTool"); + FeatureGate[] gates = EntitlementGuard.resolveRequiredGates(hm); + assertThat(gates).containsExactly(FeatureGate.OFFSITE_PROCESSING); + } + + @Test + void resolveRequiredGates_withAnnotation_usesAnnotationValue() throws Exception { + HandlerMethod hm = handlerFor("automationOnly"); + FeatureGate[] gates = EntitlementGuard.resolveRequiredGates(hm); + assertThat(gates).containsExactly(FeatureGate.AUTOMATION); + } + + // --------------------------------------------------------------------------------------- + // Helpers / fixture controller + // --------------------------------------------------------------------------------------- + + private static EntitlementSnapshot fullSnapshot() { + return new EntitlementSnapshot( + EntitlementState.FULL, + FeatureSet.FULL, + List.of( + FeatureGate.OFFSITE_PROCESSING, + FeatureGate.AUTOMATION, + FeatureGate.AI_SUPPORT, + FeatureGate.CLIENT_SIDE), + 0L, + 500L, + LocalDateTime.of(2026, 6, 1, 0, 0), + LocalDateTime.of(2026, 7, 1, 0, 0), + false); + } + + private static EntitlementSnapshot degradedSnapshot() { + return degradedSnapshot(false); + } + + private static EntitlementSnapshot degradedSnapshot(boolean subscribed) { + return new EntitlementSnapshot( + EntitlementState.DEGRADED, + FeatureSet.MINIMAL, + List.of(FeatureGate.OFFSITE_PROCESSING, FeatureGate.CLIENT_SIDE), + 500L, + 500L, + LocalDateTime.of(2026, 6, 1, 0, 0), + LocalDateTime.of(2026, 7, 1, 0, 0), + subscribed); + } + + private static ApiKeyAuthenticationToken apiKeyAuth(long teamId) { + User u = userWithTeam(99L, teamId); + return new ApiKeyAuthenticationToken( + u, "sk-test", List.of(new SimpleGrantedAuthority("ROLE_USER"))); + } + + private static User userWithTeam(long userId, long teamId) { + User u = new User(); + u.setId(userId); + Team t = new Team(); + t.setId(teamId); + u.setTeam(t); + return u; + } + + private static EnhancedJwtAuthenticationToken jwtAuth(UUID supabaseId) { + Map headers = new HashMap<>(); + headers.put("alg", "RS256"); + Map claims = new HashMap<>(); + claims.put("sub", supabaseId.toString()); + claims.put("email", "user@example.com"); + Jwt jwt = new Jwt("token", Instant.now(), Instant.now().plusSeconds(3600), headers, claims); + return new EnhancedJwtAuthenticationToken( + jwt, + List.of(new SimpleGrantedAuthority("ROLE_USER")), + "user@example.com", + supabaseId.toString()); + } + + private static HandlerMethod handlerFor(String methodName) throws NoSuchMethodException { + Method m = TestController.class.getDeclaredMethod(methodName); + return new HandlerMethod(new TestController(), m); + } + + private static HandlerMethod handlerForClassLevel(String methodName) + throws NoSuchMethodException { + Method m = ClassLevelAiController.class.getDeclaredMethod(methodName); + return new HandlerMethod(new ClassLevelAiController(), m); + } + + /** Fixture mounting route shapes the guard's resolver needs to discriminate. */ + static class TestController { + + @AutoJobPostMapping("/manual") + public String manualTool() { + return "ok"; + } + + @AutoJobPostMapping("/automation") + @RequiresFeature(FeatureGate.AUTOMATION) + public String automationOnly() { + return "ok"; + } + + @AutoJobPostMapping("/ai") + @RequiresFeature(FeatureGate.AI_SUPPORT) + public String aiOnly() { + return "ok"; + } + + /** Endpoint without any annotation — guard must skip. */ + public String plainEndpoint() { + return "ok"; + } + + /** AI-controller shape: @RequiresFeature with NO @AutoJobPostMapping (JSON body). */ + @RequiresFeature(FeatureGate.AI_SUPPORT) + public String aiOnlyNoAutoJobPostMapping() { + return "ok"; + } + } + + /** + * Mirrors {@code AiCreateController} layout: @RequiresFeature on the class, plain methods. The + * guard must pick up the class-level annotation. + */ + @RequiresFeature(FeatureGate.AI_SUPPORT) + static class ClassLevelAiController { + public String classLevelAiMethod() { + return "ok"; + } + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/payg/entitlement/EntitlementServiceTest.java b/app/saas/src/test/java/stirling/software/saas/payg/entitlement/EntitlementServiceTest.java new file mode 100644 index 0000000000..430e1b7e0e --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/payg/entitlement/EntitlementServiceTest.java @@ -0,0 +1,271 @@ +package stirling.software.saas.payg.entitlement; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.times; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.time.LocalDateTime; +import java.util.Optional; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.mockito.Mockito; + +import stirling.software.saas.payg.billing.TeamBillingContext; +import stirling.software.saas.payg.billing.TeamBillingService; +import stirling.software.saas.payg.model.EntitlementState; +import stirling.software.saas.payg.model.FeatureGate; +import stirling.software.saas.payg.model.FeatureSet; +import stirling.software.saas.payg.repository.WalletLedgerRepository; +import stirling.software.saas.payg.repository.WalletPolicyRepository; +import stirling.software.saas.payg.wallet.WalletPolicy; + +/** + * Unit tests for {@link EntitlementService}. Two branches (design 2026-06-11 — the free allowance + * is a one-time lifetime grant): + * + *

    + *
  • Unsubscribed — gated by the grant. Cap = grant size, spend = {@code grant − + * remaining}, both read straight from the billing context (no ledger query). Exhausted grant + * (remaining ≤ 0) → DEGRADED. + *
  • Subscribed — gated by the monthly money-derived doc cap. Spend = this period's net + * billable units ({@link WalletLedgerRepository#sumPeriodNetBillable} negated, refunds + * netted). + *
+ * + * Also covers cache hit/miss + the invalidate cascade. + */ +class EntitlementServiceTest { + + private static final LocalDateTime PERIOD_START = LocalDateTime.of(2026, 6, 9, 0, 0); + private static final LocalDateTime PERIOD_END = LocalDateTime.of(2026, 7, 9, 0, 0); + + private TeamBillingService billingService; + private WalletPolicyRepository walletPolicyRepo; + private WalletLedgerRepository ledgerRepo; + private EntitlementService service; + + @BeforeEach + void setUp() { + billingService = Mockito.mock(TeamBillingService.class); + walletPolicyRepo = Mockito.mock(WalletPolicyRepository.class); + ledgerRepo = Mockito.mock(WalletLedgerRepository.class); + service = new EntitlementService(billingService, walletPolicyRepo, ledgerRepo); + } + + @Test + void getSnapshot_nullTeamId_throws() { + assertThatThrownBy(() -> service.getSnapshot(null)) + .isInstanceOf(NullPointerException.class); + } + + @Test + void freeTeam_capIsTheGrantAndSpendIsUsedFromCounter() { + // Unsubscribed: cap = grant size, spend = grant − remaining — no ledger read. + stubBilling(42L, freeContext(500L, 400L)); + when(walletPolicyRepo.findByTeamId(42L)).thenReturn(Optional.empty()); + + EntitlementSnapshot snap = service.getSnapshot(42L); + + assertThat(snap.periodCapUnits()).isEqualTo(500L); + assertThat(snap.periodSpendUnits()).isEqualTo(100L); + // 100/500 = 20% — well below warn → FULL + assertThat(snap.state()).isEqualTo(EntitlementState.FULL); + assertThat(snap.featureSet()).isEqualTo(FeatureSet.FULL); + // The grant gate doesn't touch the ledger at all. + Mockito.verifyNoInteractions(ledgerRepo); + } + + @Test + void subscribedTeam_capIsTheMoneyDerivedDocCap() { + stubBilling(42L, subscribedContext(2000L)); + when(walletPolicyRepo.findByTeamId(42L)) + .thenReturn(Optional.of(walletPolicyThresholds(FeatureSet.MINIMAL))); + when(ledgerRepo.sumPeriodNetBillable(eq(42L), any(), any())).thenReturn(-500L); + + EntitlementSnapshot snap = service.getSnapshot(42L); + + assertThat(snap.periodCapUnits()).isEqualTo(2000L); + assertThat(snap.periodSpendUnits()).isEqualTo(500L); + // 500/2000 = 25% — FULL + assertThat(snap.state()).isEqualTo(EntitlementState.FULL); + } + + @Test + void subscribedTeam_refundsNetAgainstSpend() { + // Net billable = debits − refunds. A −300 net (e.g. 500 debited, 200 refunded) → 300 spend. + stubBilling(42L, subscribedContext(2000L)); + when(walletPolicyRepo.findByTeamId(42L)).thenReturn(Optional.empty()); + when(ledgerRepo.sumPeriodNetBillable(eq(42L), any(), any())).thenReturn(-300L); + + EntitlementSnapshot snap = service.getSnapshot(42L); + + assertThat(snap.periodSpendUnits()).isEqualTo(300L); + } + + @Test + void spendWindow_comesFromBillingContextNotCalendarMonth() { + stubBilling(42L, subscribedContext(2000L)); + when(walletPolicyRepo.findByTeamId(42L)).thenReturn(Optional.empty()); + when(ledgerRepo.sumPeriodNetBillable(eq(42L), any(), any())).thenReturn(0L); + + EntitlementSnapshot snap = service.getSnapshot(42L); + + // The subscription-anchored window flows through to both the snapshot and the SUM query. + assertThat(snap.periodStart()).isEqualTo(PERIOD_START); + assertThat(snap.periodEnd()).isEqualTo(PERIOD_END); + verify(ledgerRepo).sumPeriodNetBillable(eq(42L), eq(PERIOD_START), eq(PERIOD_END)); + } + + @Test + void exhaustedGrant_returnsDegradedWithMinimalGates() { + // Grant fully consumed (remaining 0) → billable categories hard-stop for an unsubscribed + // team. The displayed cap stays the grant size; spend reads as the full grant. + stubBilling(42L, freeContext(100L, 0L)); + when(walletPolicyRepo.findByTeamId(42L)) + .thenReturn(Optional.of(walletPolicyThresholds(FeatureSet.MINIMAL))); + + EntitlementSnapshot snap = service.getSnapshot(42L); + + assertThat(snap.state()).isEqualTo(EntitlementState.DEGRADED); + assertThat(snap.featureSet()).isEqualTo(FeatureSet.MINIMAL); + assertThat(snap.periodCapUnits()).isEqualTo(100L); + assertThat(snap.periodSpendUnits()).isEqualTo(100L); + // MINIMAL now keeps OFFSITE_PROCESSING + CLIENT_SIDE (manual tools); AUTOMATION + AI gone. + assertThat(snap.enabledGates()) + .containsExactlyInAnyOrder(FeatureGate.OFFSITE_PROCESSING, FeatureGate.CLIENT_SIDE); + assertThat(snap.enabledGates()) + .doesNotContain(FeatureGate.AUTOMATION, FeatureGate.AI_SUPPORT); + } + + @Test + void grantInWarnBand_returnsWarnedButFullFeatureSet() { + // grant 100, remaining 15 → used 85 = 85% (between warn 80 and degrade 100). + stubBilling(42L, freeContext(100L, 15L)); + when(walletPolicyRepo.findByTeamId(42L)) + .thenReturn(Optional.of(walletPolicyThresholds(FeatureSet.MINIMAL))); + + EntitlementSnapshot snap = service.getSnapshot(42L); + + assertThat(snap.state()).isEqualTo(EntitlementState.WARNED); + assertThat(snap.featureSet()).isEqualTo(FeatureSet.FULL); + assertThat(snap.enabledGates()).hasSize(4); + } + + @Test + void uncappedSubscribedTeam_nullCapNeverDegrades() { + stubBilling(42L, subscribedContext(null)); + when(walletPolicyRepo.findByTeamId(42L)).thenReturn(Optional.empty()); + when(ledgerRepo.sumPeriodNetBillable(eq(42L), any(), any())).thenReturn(-1_000_000L); + + EntitlementSnapshot snap = service.getSnapshot(42L); + + assertThat(snap.periodCapUnits()).isNull(); + assertThat(snap.state()).isEqualTo(EntitlementState.FULL); + } + + @Test + void positiveNetBillable_treatedAsZeroSpend() { + // Subscribed defensive: if refunds exceed debits (positive net), spend clamps to zero. + stubBilling(42L, subscribedContext(100L)); + when(walletPolicyRepo.findByTeamId(42L)).thenReturn(Optional.empty()); + when(ledgerRepo.sumPeriodNetBillable(eq(42L), any(), any())).thenReturn(50L); + + EntitlementSnapshot snap = service.getSnapshot(42L); + + assertThat(snap.periodSpendUnits()).isZero(); + assertThat(snap.state()).isEqualTo(EntitlementState.FULL); + } + + @Test + void cacheHit_secondCallSkipsLedgerLookup() { + stubBilling(42L, subscribedContext(500L)); + when(walletPolicyRepo.findByTeamId(42L)).thenReturn(Optional.empty()); + when(ledgerRepo.sumPeriodNetBillable(eq(42L), any(), any())).thenReturn(0L); + + service.getSnapshot(42L); + service.getSnapshot(42L); + service.getSnapshot(42L); + + // Only one underlying ledger SUM despite 3 calls — second + third hit the cache. + verify(ledgerRepo, times(1)).sumPeriodNetBillable(eq(42L), any(), any()); + assertThat(service.cacheSize()).isEqualTo(1); + } + + @Test + void invalidate_dropsCacheAndCascadesToBillingService() { + stubBilling(42L, subscribedContext(500L)); + when(walletPolicyRepo.findByTeamId(42L)).thenReturn(Optional.empty()); + when(ledgerRepo.sumPeriodNetBillable(eq(42L), any(), any())).thenReturn(0L); + + service.getSnapshot(42L); + service.invalidate(42L); + service.getSnapshot(42L); + + verify(ledgerRepo, times(2)).sumPeriodNetBillable(eq(42L), any(), any()); + // Window/cap facts must recompute together with the spend. + verify(billingService).invalidate(42L); + } + + @Test + void invalidate_otherTeamLeavesEntryAlone() { + when(billingService.forTeam(any())).thenReturn(subscribedContext(500L)); + when(walletPolicyRepo.findByTeamId(any())).thenReturn(Optional.empty()); + when(ledgerRepo.sumPeriodNetBillable(any(), any(), any())).thenReturn(0L); + + service.getSnapshot(42L); + service.invalidate(99L); + service.getSnapshot(42L); + + // Only one fetch for team 42 — 99 invalidate didn't touch its entry. + verify(ledgerRepo, times(1)).sumPeriodNetBillable(eq(42L), any(), any()); + } + + @Test + void currentMonthWindow_isStartOfMonthInclusiveToStartOfNextMonthExclusive() { + LocalDateTime mid = LocalDateTime.of(2026, 6, 15, 14, 30); + LocalDateTime[] w = EntitlementService.currentMonthWindow(mid); + assertThat(w[0]).isEqualTo(LocalDateTime.of(2026, 6, 1, 0, 0)); + assertThat(w[1]).isEqualTo(LocalDateTime.of(2026, 7, 1, 0, 0)); + } + + private void stubBilling(Long teamId, TeamBillingContext ctx) { + when(billingService.forTeam(teamId)).thenReturn(ctx); + } + + /** Unsubscribed team: gated by the one-time grant (size + remaining); no monthly cap. */ + private static TeamBillingContext freeContext(long grant, long remaining) { + return new TeamBillingContext( + false, null, PERIOD_START, PERIOD_END, grant, remaining, null, null, null, null); + } + + /** + * Subscribed team: monthly money-derived paid-doc cap (null = uncapped); grant treated as + * exhausted (remaining 0 — doesn't gate a paying team). + */ + private static TeamBillingContext subscribedContext(Long monthlyCapDocUnits) { + return new TeamBillingContext( + true, + "sub_test", + PERIOD_START, + PERIOD_END, + 500L, + 0L, + java.math.BigDecimal.valueOf(2), + "usd", + monthlyCapDocUnits == null ? null : monthlyCapDocUnits * 2, + monthlyCapDocUnits); + } + + private static WalletPolicy walletPolicyThresholds(FeatureSet degradedSet) { + WalletPolicy p = new WalletPolicy(); + p.setDegradedFeatureSet(degradedSet); + p.setWarnAtPct(80); + p.setDegradeAtPct(100); + return p; + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/payg/filter/PaygChargeInterceptorTest.java b/app/saas/src/test/java/stirling/software/saas/payg/filter/PaygChargeInterceptorTest.java new file mode 100644 index 0000000000..128c336c32 --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/payg/filter/PaygChargeInterceptorTest.java @@ -0,0 +1,830 @@ +package stirling.software.saas.payg.filter; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyList; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.verifyNoInteractions; +import static org.mockito.Mockito.when; + +import java.io.IOException; +import java.lang.reflect.Method; +import java.nio.charset.StandardCharsets; +import java.util.List; +import java.util.Optional; +import java.util.UUID; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.mockito.ArgumentCaptor; +import org.springframework.mock.web.MockHttpServletRequest; +import org.springframework.mock.web.MockHttpServletResponse; +import org.springframework.mock.web.MockMultipartFile; +import org.springframework.mock.web.MockMultipartHttpServletRequest; +import org.springframework.security.authentication.UsernamePasswordAuthenticationToken; +import org.springframework.security.core.authority.SimpleGrantedAuthority; +import org.springframework.security.core.context.SecurityContextHolder; +import org.springframework.web.method.HandlerMethod; + +import io.micrometer.core.instrument.MeterRegistry; +import io.micrometer.core.instrument.simple.SimpleMeterRegistry; + +import stirling.software.common.annotations.AutoJobPostMapping; +import stirling.software.common.model.ApplicationProperties; +import stirling.software.common.util.TempFileManager; +import stirling.software.common.util.TempFileRegistry; +import stirling.software.proprietary.security.database.repository.UserRepository; +import stirling.software.proprietary.security.model.User; +import stirling.software.saas.payg.charge.ChargeOutcome; +import stirling.software.saas.payg.charge.JobChargeService; +import stirling.software.saas.payg.job.JobService; +import stirling.software.saas.payg.model.JobStepStatus; + +/** + * Pure-Mockito tests for {@link PaygChargeInterceptor}. Real {@link TempFileManager} but mocked + * downstream charge/job services so the test runs without a Spring context. + */ +class PaygChargeInterceptorTest { + + private JobChargeService chargeService; + private JobService jobService; + private UserRepository userRepository; + private PaygOutputExtractor outputExtractor; + private PaygFilterProperties properties; + private MeterRegistry meterRegistry; + private TempFileManager tempFileManager; + private PaygChargeInterceptor interceptor; + + @BeforeEach + void setUp() { + chargeService = org.mockito.Mockito.mock(JobChargeService.class); + jobService = org.mockito.Mockito.mock(JobService.class); + userRepository = org.mockito.Mockito.mock(UserRepository.class); + outputExtractor = org.mockito.Mockito.mock(PaygOutputExtractor.class); + properties = new PaygFilterProperties(); + meterRegistry = new SimpleMeterRegistry(); + tempFileManager = new TempFileManager(new TempFileRegistry(), new ApplicationProperties()); + interceptor = + new PaygChargeInterceptor( + chargeService, + jobService, + userRepository, + tempFileManager, + outputExtractor, + properties, + meterRegistry); + SecurityContextHolder.clearContext(); + } + + @Test + void preHandle_filterDisabled_isShortCircuitNoop() throws Exception { + properties.setEnabled(false); + MockHttpServletRequest req = new MockHttpServletRequest(); + MockHttpServletResponse res = new MockHttpServletResponse(); + + boolean cont = interceptor.preHandle(req, res, handlerMethodForFakeController()); + + assertThat(cont).isTrue(); + verifyNoInteractions(chargeService); + } + + @Test + void preHandle_handlerNotAnnotated_isShortCircuitNoop() throws Exception { + MockHttpServletRequest req = new MockHttpServletRequest(); + MockHttpServletResponse res = new MockHttpServletResponse(); + + boolean cont = interceptor.preHandle(req, res, handlerMethodForPlain()); + + assertThat(cont).isTrue(); + verifyNoInteractions(chargeService); + } + + @Test + void preHandle_noAuth_isShortCircuitNoop() throws Exception { + MockMultipartHttpServletRequest req = newMultipart(); + req.addFile(new MockMultipartFile("file", "x.pdf", "application/pdf", "hi".getBytes())); + + boolean cont = + interceptor.preHandle( + req, new MockHttpServletResponse(), handlerMethodForFakeController()); + + assertThat(cont).isTrue(); + verifyNoInteractions(chargeService); + } + + @Test + void preHandle_noMultipartParts_isShortCircuitNoop() throws Exception { + // Authenticated but no file parts. + authenticateWithUser(makeUser(1L, null)); + MockHttpServletRequest req = new MockHttpServletRequest(); + + boolean cont = + interceptor.preHandle( + req, new MockHttpServletResponse(), handlerMethodForFakeController()); + + assertThat(cont).isTrue(); + verifyNoInteractions(chargeService); + } + + @Test + void preHandle_openedDisposition_stashesJobIdAndDisposition() throws Exception { + // API-key auth → BillingCategory.API → billable path engaged. + authenticateWithApiKey(makeUser(7L, 42L)); + UUID jobId = UUID.randomUUID(); + when(chargeService.openProcess(any(), anyList())) + .thenReturn(new ChargeOutcome(jobId, 4, ChargeOutcome.Disposition.OPENED)); + + MockMultipartHttpServletRequest req = newMultipart(); + req.addFile(new MockMultipartFile("file", "x.pdf", "application/pdf", "abc".getBytes())); + + boolean cont = + interceptor.preHandle( + req, new MockHttpServletResponse(), handlerMethodForFakeController()); + + assertThat(cont).isTrue(); + assertThat(req.getAttribute(PaygChargeInterceptor.ATTR_JOB_ID)).isEqualTo(jobId); + assertThat(req.getAttribute(PaygChargeInterceptor.ATTR_DISPOSITION)) + .isEqualTo(ChargeOutcome.Disposition.OPENED); + assertThat(req.getAttribute(PaygChargeInterceptor.ATTR_INPUT_TEMP_FILES)).isNotNull(); + assertThat(meterRegistry.counter("payg.filter.calls", "disposition", "OPENED").count()) + .isEqualTo(1.0); + } + + @Test + void preHandle_chargeServiceThrows_failsOpenAndIncrementsErrorCounter() throws Exception { + // API-key auth → API category → reaches openProcess so the throw can be observed. + authenticateWithApiKey(makeUser(7L, 42L)); + when(chargeService.openProcess(any(), anyList())).thenThrow(new RuntimeException("boom")); + + MockMultipartHttpServletRequest req = newMultipart(); + req.addFile(new MockMultipartFile("file", "x.pdf", "application/pdf", "abc".getBytes())); + + boolean cont = + interceptor.preHandle( + req, new MockHttpServletResponse(), handlerMethodForFakeController()); + + assertThat(cont).isTrue(); // fail-open + assertThat(req.getAttribute(PaygChargeInterceptor.ATTR_FAILED)).isEqualTo(Boolean.TRUE); + assertThat(meterRegistry.counter("payg.filter.errors").count()).isEqualTo(1.0); + } + + @Test + void afterCompletion_2xx_appendsStepAndRecordsOutputs() throws Exception { + authenticateWithApiKey(makeUser(7L, 42L)); + UUID jobId = UUID.randomUUID(); + when(chargeService.openProcess(any(), anyList())) + .thenReturn(new ChargeOutcome(jobId, 1, ChargeOutcome.Disposition.OPENED)); + + MockMultipartHttpServletRequest req = newMultipart(); + req.addFile(new MockMultipartFile("file", "x.pdf", "application/pdf", "abc".getBytes())); + MockHttpServletResponse res = new MockHttpServletResponse(); + res.setStatus(200); + res.setContentType("application/pdf"); + + // Install a wrapper as the filter would. Pre-populate it with some bytes so + // materialisedPath returns non-null. + PaygResponseBodyWrapper wrapper = new PaygResponseBodyWrapper(res, tempFileManager, 1024); + wrapper.getOutputStream().write("body".getBytes(StandardCharsets.UTF_8)); + req.setAttribute(PaygResponseBodyWrapperFilter.REQUEST_ATTRIBUTE, wrapper); + + interceptor.preHandle(req, res, handlerMethodForFakeController()); + + when(outputExtractor.extract(eq("application/pdf"), any())) + .thenReturn( + List.of( + new PaygOutputExtractor.ExtractedPdf( + wrapper.materialisedPath(), null))); + + interceptor.afterCompletion(req, res, handlerMethodForFakeController(), null); + + ArgumentCaptor status = ArgumentCaptor.forClass(JobStepStatus.class); + verify(jobService).appendStep(eq(jobId), any(), status.capture(), any(), any(), any()); + assertThat(status.getValue()).isEqualTo(JobStepStatus.OK); + verify(jobService).recordOutput(eq(jobId), any()); + verify(chargeService, never()).markFirstStepFailed(any(), any()); + verify(chargeService, never()).decrementStepCount(any()); + // Success on an OPENED process is the primary meter trigger — fires now, not at close. + verify(chargeService).meterJobUsage(jobId); + } + + @Test + void afterCompletion_2xx_joined_doesNotMeter() throws Exception { + // A JOINED follow-up step (chained tool on the same document) added no units when it + // joined — it must not re-meter; the OPENED step already did. + authenticateWithApiKey(makeUser(7L, 42L)); + UUID jobId = UUID.randomUUID(); + when(chargeService.openProcess(any(), anyList())) + .thenReturn(new ChargeOutcome(jobId, 0, ChargeOutcome.Disposition.JOINED)); + + MockMultipartHttpServletRequest req = newMultipart(); + req.addFile(new MockMultipartFile("file", "x.pdf", "application/pdf", "abc".getBytes())); + MockHttpServletResponse res = new MockHttpServletResponse(); + res.setStatus(200); + + interceptor.preHandle(req, res, handlerMethodForFakeController()); + interceptor.afterCompletion(req, res, handlerMethodForFakeController(), null); + + verify(chargeService, never()).meterJobUsage(any()); + } + + @Test + void afterCompletion_5xx_opened_callsMarkFirstStepFailed() throws Exception { + authenticateWithApiKey(makeUser(7L, 42L)); + UUID jobId = UUID.randomUUID(); + when(chargeService.openProcess(any(), anyList())) + .thenReturn(new ChargeOutcome(jobId, 1, ChargeOutcome.Disposition.OPENED)); + + MockMultipartHttpServletRequest req = newMultipart(); + req.addFile(new MockMultipartFile("file", "x.pdf", "application/pdf", "abc".getBytes())); + MockHttpServletResponse res = new MockHttpServletResponse(); + res.setStatus(503); + + interceptor.preHandle(req, res, handlerMethodForFakeController()); + interceptor.afterCompletion(req, res, handlerMethodForFakeController(), null); + + verify(chargeService).markFirstStepFailed(eq(jobId), eq("first-step-5xx:503")); + verify(chargeService, never()).decrementStepCount(any()); + verify(jobService, never()).recordOutput(any(), any()); + verify(jobService) + .appendStep(eq(jobId), any(), eq(JobStepStatus.FAILED), any(), any(), eq("503")); + assertThat(meterRegistry.counter("payg.filter.refunds").count()).isEqualTo(1.0); + // First-step failure refunds — never meter it. + verify(chargeService, never()).meterJobUsage(any()); + } + + @Test + void afterCompletion_5xx_joined_callsDecrementStepCount() throws Exception { + authenticateWithApiKey(makeUser(7L, 42L)); + UUID jobId = UUID.randomUUID(); + when(chargeService.openProcess(any(), anyList())) + .thenReturn(new ChargeOutcome(jobId, 0, ChargeOutcome.Disposition.JOINED)); + + MockMultipartHttpServletRequest req = newMultipart(); + req.addFile(new MockMultipartFile("file", "x.pdf", "application/pdf", "abc".getBytes())); + MockHttpServletResponse res = new MockHttpServletResponse(); + res.setStatus(500); + + interceptor.preHandle(req, res, handlerMethodForFakeController()); + interceptor.afterCompletion(req, res, handlerMethodForFakeController(), null); + + verify(chargeService).decrementStepCount(jobId); + verify(chargeService, never()).markFirstStepFailed(any(), any()); + } + + @Test + void afterCompletion_4xx_appendsFailedStepNoRefundNoOutputs() throws Exception { + authenticateWithApiKey(makeUser(7L, 42L)); + UUID jobId = UUID.randomUUID(); + when(chargeService.openProcess(any(), anyList())) + .thenReturn(new ChargeOutcome(jobId, 1, ChargeOutcome.Disposition.OPENED)); + + MockMultipartHttpServletRequest req = newMultipart(); + req.addFile(new MockMultipartFile("file", "x.pdf", "application/pdf", "abc".getBytes())); + MockHttpServletResponse res = new MockHttpServletResponse(); + res.setStatus(422); + + interceptor.preHandle(req, res, handlerMethodForFakeController()); + interceptor.afterCompletion(req, res, handlerMethodForFakeController(), null); + + verify(chargeService, never()).markFirstStepFailed(any(), any()); + verify(chargeService, never()).decrementStepCount(any()); + verify(jobService, never()).recordOutput(any(), any()); + verify(jobService) + .appendStep(eq(jobId), any(), eq(JobStepStatus.FAILED), any(), any(), eq("422")); + // 4xx is a full charge (customer paid for the attempt), so it still meters. + verify(chargeService).meterJobUsage(jobId); + } + + @Test + void afterCompletion_preHandleFailedFlag_cleansUpWithoutCharging() throws Exception { + MockMultipartHttpServletRequest req = newMultipart(); + req.setAttribute(PaygChargeInterceptor.ATTR_FAILED, Boolean.TRUE); + MockHttpServletResponse res = new MockHttpServletResponse(); + + interceptor.afterCompletion(req, res, handlerMethodForFakeController(), null); + + verifyNoInteractions(chargeService); + verifyNoInteractions(jobService); + } + + @Test + void afterCompletion_noJobId_isNoop() throws Exception { + MockMultipartHttpServletRequest req = newMultipart(); + MockHttpServletResponse res = new MockHttpServletResponse(); + interceptor.afterCompletion(req, res, handlerMethodForFakeController(), null); + verifyNoInteractions(chargeService); + verifyNoInteractions(jobService); + } + + @Test + void afterCompletion_maxBytesExceeded_skipsOutputRecording() throws Exception { + properties.getResponse().setMaxBytes(2L); + authenticateWithApiKey(makeUser(7L, 42L)); + UUID jobId = UUID.randomUUID(); + when(chargeService.openProcess(any(), anyList())) + .thenReturn(new ChargeOutcome(jobId, 1, ChargeOutcome.Disposition.OPENED)); + + MockMultipartHttpServletRequest req = newMultipart(); + req.addFile(new MockMultipartFile("file", "x.pdf", "application/pdf", "abc".getBytes())); + MockHttpServletResponse res = new MockHttpServletResponse(); + res.setStatus(200); + res.setContentType("application/pdf"); + + PaygResponseBodyWrapper wrapper = new PaygResponseBodyWrapper(res, tempFileManager, 1024); + wrapper.getOutputStream().write("1234567890".getBytes(StandardCharsets.UTF_8)); + req.setAttribute(PaygResponseBodyWrapperFilter.REQUEST_ATTRIBUTE, wrapper); + + interceptor.preHandle(req, res, handlerMethodForFakeController()); + interceptor.afterCompletion(req, res, handlerMethodForFakeController(), null); + + verify(outputExtractor, never()).extract(any(), any()); + verify(jobService, never()).recordOutput(any(), any()); + } + + @Test + void preHandle_desktopClientHeader_setsJobSourceDesktopApp() throws Exception { + // API-key (billable) so the call reaches openProcess; the desktop header still wins the + // source mapping (checked before the API-key branch in determineSource). A manual JWT call + // is BYPASSED and never opens a process, so source wouldn't be recorded. + authenticateWithApiKey(makeUser(7L, 42L)); + UUID jobId = UUID.randomUUID(); + when(chargeService.openProcess(any(), anyList())) + .thenReturn(new ChargeOutcome(jobId, 1, ChargeOutcome.Disposition.OPENED)); + org.mockito.ArgumentCaptor ctxCaptor = + org.mockito.ArgumentCaptor.forClass( + stirling.software.saas.payg.charge.ChargeContext.class); + + MockMultipartHttpServletRequest req = newMultipart(); + req.addFile(new MockMultipartFile("file", "x.pdf", "application/pdf", "abc".getBytes())); + req.addHeader("X-Stirling-Client", "desktop"); + + interceptor.preHandle(req, new MockHttpServletResponse(), handlerMethodForFakeController()); + + verify(chargeService).openProcess(ctxCaptor.capture(), anyList()); + assertThat(ctxCaptor.getValue().source()) + .isEqualTo(stirling.software.saas.payg.model.JobSource.DESKTOP_APP); + } + + @Test + void preHandle_toolId_prefersBestMatchingPattern() throws Exception { + // API-key (billable) so doPreHandle runs and records tool_id; a manual JWT call is + // BYPASSED before that point, so no tool_id is stored. + authenticateWithApiKey(makeUser(7L, 42L)); + UUID jobId = UUID.randomUUID(); + when(chargeService.openProcess(any(), anyList())) + .thenReturn(new ChargeOutcome(jobId, 1, ChargeOutcome.Disposition.OPENED)); + + MockMultipartHttpServletRequest req = newMultipart(); + // Raw URI contains a path variable; the matched pattern is what audit rollups want. + req.setRequestURI("/api/v1/security/add-password/extra/segment"); + req.setAttribute( + org.springframework.web.servlet.HandlerMapping.BEST_MATCHING_PATTERN_ATTRIBUTE, + "/api/v1/security/add-password"); + req.addFile(new MockMultipartFile("file", "x.pdf", "application/pdf", "abc".getBytes())); + + interceptor.preHandle(req, new MockHttpServletResponse(), handlerMethodForFakeController()); + + assertThat(req.getAttribute(PaygChargeInterceptor.ATTR_TOOL_ID)) + .isEqualTo("/api/v1/security/add-password"); + } + + @Test + void preHandle_toolId_truncatesAndCountsWhenLongerThan128() throws Exception { + // API-key (billable) so doPreHandle runs and records tool_id (a manual JWT call is + // BYPASSED). + authenticateWithApiKey(makeUser(7L, 42L)); + UUID jobId = UUID.randomUUID(); + when(chargeService.openProcess(any(), anyList())) + .thenReturn(new ChargeOutcome(jobId, 1, ChargeOutcome.Disposition.OPENED)); + + MockMultipartHttpServletRequest req = newMultipart(); + String oversized = "/api/v1/" + "x".repeat(200); + req.setRequestURI(oversized); + // No matching pattern attribute — falls back to URI. + req.addFile(new MockMultipartFile("file", "x.pdf", "application/pdf", "abc".getBytes())); + + interceptor.preHandle(req, new MockHttpServletResponse(), handlerMethodForFakeController()); + + Object stored = req.getAttribute(PaygChargeInterceptor.ATTR_TOOL_ID); + assertThat(stored).isInstanceOf(String.class); + assertThat(((String) stored)).hasSize(128); + assertThat(meterRegistry.counter("payg.filter.errors").count()).isEqualTo(1.0); + } + + @Test + void preHandle_pipelineHeader_setsJobSourcePipeline() throws Exception { + authenticateWithApiKey(makeUser(7L, 42L)); + UUID jobId = UUID.randomUUID(); + when(chargeService.openProcess(any(), anyList())) + .thenReturn(new ChargeOutcome(jobId, 1, ChargeOutcome.Disposition.OPENED)); + org.mockito.ArgumentCaptor ctxCaptor = + org.mockito.ArgumentCaptor.forClass( + stirling.software.saas.payg.charge.ChargeContext.class); + + MockMultipartHttpServletRequest req = newMultipart(); + req.addFile(new MockMultipartFile("file", "x.pdf", "application/pdf", "abc".getBytes())); + req.addHeader("X-Stirling-Automation", "true"); + + interceptor.preHandle(req, new MockHttpServletResponse(), handlerMethodForFakeController()); + + verify(chargeService).openProcess(ctxCaptor.capture(), anyList()); + assertThat(ctxCaptor.getValue().source()) + .isEqualTo(stirling.software.saas.payg.model.JobSource.PIPELINE); + } + + // --- BillingCategory categorisation + bypass fast-path ------------------------------------- + + @Test + void preHandle_manualToolJwt_isBypassedAndSkipsOpenProcess() throws Exception { + // JWT-authenticated, plain @AutoJobPostMapping endpoint, no automation header → BYPASSED. + // The interceptor must skip openProcess entirely (no temp files, no DB writes) and bump + // the payg.filter.bypassed counter. + authenticateWithUser(makeUser(7L, 42L)); + + MockMultipartHttpServletRequest req = newMultipart(); + req.addFile(new MockMultipartFile("file", "x.pdf", "application/pdf", "abc".getBytes())); + + boolean cont = + interceptor.preHandle( + req, new MockHttpServletResponse(), handlerMethodForFakeController()); + + assertThat(cont).isTrue(); + verify(chargeService, never()).openProcess(any(), anyList()); + assertThat(req.getAttribute(PaygChargeInterceptor.ATTR_JOB_ID)).isNull(); + assertThat(req.getAttribute(PaygChargeInterceptor.ATTR_INPUT_TEMP_FILES)).isNull(); + assertThat(meterRegistry.counter("payg.filter.bypassed").count()).isEqualTo(1.0); + } + + @Test + void preHandle_apiKeyAuth_setsBillingCategoryApi() throws Exception { + authenticateWithApiKey(makeUser(7L, 42L)); + UUID jobId = UUID.randomUUID(); + when(chargeService.openProcess(any(), anyList())) + .thenReturn(new ChargeOutcome(jobId, 1, ChargeOutcome.Disposition.OPENED)); + org.mockito.ArgumentCaptor ctxCaptor = + org.mockito.ArgumentCaptor.forClass( + stirling.software.saas.payg.charge.ChargeContext.class); + + MockMultipartHttpServletRequest req = newMultipart(); + req.addFile(new MockMultipartFile("file", "x.pdf", "application/pdf", "abc".getBytes())); + + interceptor.preHandle(req, new MockHttpServletResponse(), handlerMethodForFakeController()); + + verify(chargeService).openProcess(ctxCaptor.capture(), anyList()); + assertThat(ctxCaptor.getValue().billingCategory()) + .isEqualTo(stirling.software.saas.payg.model.BillingCategory.API); + } + + @Test + void preHandle_requiresFeatureAutomation_setsBillingCategoryAutomation() throws Exception { + // JWT auth on a @RequiresFeature(AUTOMATION) endpoint → AUTOMATION category. + authenticateWithUser(makeUser(7L, 42L)); + UUID jobId = UUID.randomUUID(); + when(chargeService.openProcess(any(), anyList())) + .thenReturn(new ChargeOutcome(jobId, 1, ChargeOutcome.Disposition.OPENED)); + org.mockito.ArgumentCaptor ctxCaptor = + org.mockito.ArgumentCaptor.forClass( + stirling.software.saas.payg.charge.ChargeContext.class); + + MockMultipartHttpServletRequest req = newMultipart(); + req.addFile(new MockMultipartFile("file", "x.pdf", "application/pdf", "abc".getBytes())); + + interceptor.preHandle(req, new MockHttpServletResponse(), handlerMethodForAutomation()); + + verify(chargeService).openProcess(ctxCaptor.capture(), anyList()); + assertThat(ctxCaptor.getValue().billingCategory()) + .isEqualTo(stirling.software.saas.payg.model.BillingCategory.AUTOMATION); + } + + @Test + void preHandle_requiresFeatureAiSupport_setsBillingCategoryAi() throws Exception { + // JWT auth on a @RequiresFeature(AI_SUPPORT) endpoint → AI category. + authenticateWithUser(makeUser(7L, 42L)); + UUID jobId = UUID.randomUUID(); + when(chargeService.openProcess(any(), anyList())) + .thenReturn(new ChargeOutcome(jobId, 1, ChargeOutcome.Disposition.OPENED)); + org.mockito.ArgumentCaptor ctxCaptor = + org.mockito.ArgumentCaptor.forClass( + stirling.software.saas.payg.charge.ChargeContext.class); + + MockMultipartHttpServletRequest req = newMultipart(); + req.addFile(new MockMultipartFile("file", "x.pdf", "application/pdf", "abc".getBytes())); + + interceptor.preHandle(req, new MockHttpServletResponse(), handlerMethodForAi()); + + verify(chargeService).openProcess(ctxCaptor.capture(), anyList()); + assertThat(ctxCaptor.getValue().billingCategory()) + .isEqualTo(stirling.software.saas.payg.model.BillingCategory.AI); + } + + @Test + void preHandle_requiresFeatureWithoutAutoJobPostMapping_reachesCategoryGate() throws Exception { + // Regression: AI controllers carry @RequiresFeature but NO @AutoJobPostMapping. They must + // still flow past the short-circuit gate so determineCategory runs (and so future + // multipart-bearing @RequiresFeature routes bill correctly). API-key auth + + // @RequiresFeature + // — even without multipart inputs — should land in the BillingCategory.API branch via + // determineCategory's auth check, then short-circuit inside doPreHandle because there are + // no multipart parts. + authenticateWithApiKey(makeUser(7L, 42L)); + MockMultipartHttpServletRequest req = newMultipart(); + // No file parts — emulates a JSON-bodied AI controller request that happens to be wrapped + // as multipart. doPreHandle short-circuits with no openProcess call. + req.addFile(new MockMultipartFile("file", "x.pdf", "application/pdf", new byte[0])); + + boolean cont = + interceptor.preHandle( + req, new MockHttpServletResponse(), handlerMethodForAiNoAutoJob()); + + assertThat(cont).isTrue(); + verify(chargeService, never()).openProcess(any(), anyList()); + // Importantly: not counted as BYPASSED — the AI category was determined correctly. + assertThat(meterRegistry.counter("payg.filter.bypassed").count()).isEqualTo(0.0); + } + + @Test + void preHandle_aiEndpointWithoutAutoJobPostMapping_categoryIsAi() throws Exception { + // With multipart parts present + @RequiresFeature(AI_SUPPORT) but no @AutoJobPostMapping: + // the interceptor must run determineCategory and tag the ChargeContext as AI. + authenticateWithApiKey(makeUser(7L, 42L)); + UUID jobId = UUID.randomUUID(); + when(chargeService.openProcess(any(), anyList())) + .thenReturn(new ChargeOutcome(jobId, 1, ChargeOutcome.Disposition.OPENED)); + org.mockito.ArgumentCaptor ctxCaptor = + org.mockito.ArgumentCaptor.forClass( + stirling.software.saas.payg.charge.ChargeContext.class); + + MockMultipartHttpServletRequest req = newMultipart(); + req.addFile(new MockMultipartFile("file", "x.pdf", "application/pdf", "abc".getBytes())); + + interceptor.preHandle(req, new MockHttpServletResponse(), handlerMethodForAiNoAutoJob()); + + verify(chargeService).openProcess(ctxCaptor.capture(), anyList()); + assertThat(ctxCaptor.getValue().billingCategory()) + .isEqualTo(stirling.software.saas.payg.model.BillingCategory.AI); + } + + @Test + void preHandle_classLevelRequiresFeatureOnly_isInScope() throws Exception { + // Mirrors AiCreateController shape: @RequiresFeature on the @RestController class. + // The interceptor must resolve it via beanType lookup and not short-circuit as + // "no annotation". + authenticateWithApiKey(makeUser(7L, 42L)); + UUID jobId = UUID.randomUUID(); + when(chargeService.openProcess(any(), anyList())) + .thenReturn(new ChargeOutcome(jobId, 1, ChargeOutcome.Disposition.OPENED)); + org.mockito.ArgumentCaptor ctxCaptor = + org.mockito.ArgumentCaptor.forClass( + stirling.software.saas.payg.charge.ChargeContext.class); + + MockMultipartHttpServletRequest req = newMultipart(); + req.addFile(new MockMultipartFile("file", "x.pdf", "application/pdf", "abc".getBytes())); + + interceptor.preHandle(req, new MockHttpServletResponse(), handlerMethodForClassLevelAi()); + + verify(chargeService).openProcess(ctxCaptor.capture(), anyList()); + assertThat(ctxCaptor.getValue().billingCategory()) + .isEqualTo(stirling.software.saas.payg.model.BillingCategory.AI); + } + + @Test + void preHandle_aiEndpointWithAutomationHeader_automationWinsByPrecedence() throws Exception { + // X-Stirling-Automation: true on an @RequiresFeature(AI_SUPPORT) endpoint → AUTOMATION + // (header beats annotation by design — pipeline-driven AI counts as automation usage). + authenticateWithUser(makeUser(7L, 42L)); + UUID jobId = UUID.randomUUID(); + when(chargeService.openProcess(any(), anyList())) + .thenReturn(new ChargeOutcome(jobId, 1, ChargeOutcome.Disposition.OPENED)); + org.mockito.ArgumentCaptor ctxCaptor = + org.mockito.ArgumentCaptor.forClass( + stirling.software.saas.payg.charge.ChargeContext.class); + + MockMultipartHttpServletRequest req = newMultipart(); + req.addFile(new MockMultipartFile("file", "x.pdf", "application/pdf", "abc".getBytes())); + req.addHeader("X-Stirling-Automation", "true"); + + interceptor.preHandle(req, new MockHttpServletResponse(), handlerMethodForAi()); + + verify(chargeService).openProcess(ctxCaptor.capture(), anyList()); + assertThat(ctxCaptor.getValue().billingCategory()) + .isEqualTo(stirling.software.saas.payg.model.BillingCategory.AUTOMATION); + } + + @Test + void preHandle_aiToolRoute_inScopeAndCategoryAi() throws Exception { + // AI document tools (/api/v1/ai/tools/**) live in the proprietary module and carry no PAYG + // annotation. The interceptor recognises them by path → in scope + AI category, so a direct + // multipart call opens a charge. Handler has NO annotations (handlerMethodForPlain). + authenticateWithUser(makeUser(7L, 42L)); + UUID jobId = UUID.randomUUID(); + when(chargeService.openProcess(any(), anyList())) + .thenReturn(new ChargeOutcome(jobId, 1, ChargeOutcome.Disposition.OPENED)); + org.mockito.ArgumentCaptor ctxCaptor = + org.mockito.ArgumentCaptor.forClass( + stirling.software.saas.payg.charge.ChargeContext.class); + + MockMultipartHttpServletRequest req = newMultipart(); + req.setRequestURI("/api/v1/ai/tools/pdf-comment-agent"); + req.addFile( + new MockMultipartFile("fileInput", "x.pdf", "application/pdf", "abc".getBytes())); + + interceptor.preHandle(req, new MockHttpServletResponse(), handlerMethodForPlain()); + + verify(chargeService).openProcess(ctxCaptor.capture(), anyList()); + assertThat(ctxCaptor.getValue().billingCategory()) + .isEqualTo(stirling.software.saas.payg.model.BillingCategory.AI); + } + + @Test + void preHandle_aiToolRoute_withAutomationHeader_isAutomation() throws Exception { + // An AI tool dispatched inside a policy / AI workflow carries X-Stirling-Automation: true → + // AUTOMATION wins over the AI path rule (the header is checked first). + authenticateWithUser(makeUser(7L, 42L)); + UUID jobId = UUID.randomUUID(); + when(chargeService.openProcess(any(), anyList())) + .thenReturn(new ChargeOutcome(jobId, 1, ChargeOutcome.Disposition.OPENED)); + org.mockito.ArgumentCaptor ctxCaptor = + org.mockito.ArgumentCaptor.forClass( + stirling.software.saas.payg.charge.ChargeContext.class); + + MockMultipartHttpServletRequest req = newMultipart(); + req.setRequestURI("/api/v1/ai/tools/pdf-comment-agent"); + req.addFile( + new MockMultipartFile("fileInput", "x.pdf", "application/pdf", "abc".getBytes())); + req.addHeader("X-Stirling-Automation", "true"); + + interceptor.preHandle(req, new MockHttpServletResponse(), handlerMethodForPlain()); + + verify(chargeService).openProcess(ctxCaptor.capture(), anyList()); + assertThat(ctxCaptor.getValue().billingCategory()) + .isEqualTo(stirling.software.saas.payg.model.BillingCategory.AUTOMATION); + } + + @Test + void preHandle_plainRouteNoAnnotations_stillShortCircuits() throws Exception { + // Guard against the path rule being too broad: a non-AI-tools route with no annotations + // must + // still short-circuit (BYPASSED path), unaffected by the AI-tools recognition. + authenticateWithUser(makeUser(7L, 42L)); + + MockMultipartHttpServletRequest req = newMultipart(); // URI = /api/v1/security/test-tool + req.addFile(new MockMultipartFile("file", "x.pdf", "application/pdf", "abc".getBytes())); + + boolean cont = + interceptor.preHandle(req, new MockHttpServletResponse(), handlerMethodForPlain()); + + assertThat(cont).isTrue(); + verify(chargeService, never()).openProcess(any(), anyList()); + } + + // --- helpers -------------------------------------------------------------------------------- + + private MockMultipartHttpServletRequest newMultipart() { + MockMultipartHttpServletRequest r = new MockMultipartHttpServletRequest(); + r.setRequestURI("/api/v1/security/test-tool"); + return r; + } + + private void authenticateWithUser(User user) { + UsernamePasswordAuthenticationToken token = + new UsernamePasswordAuthenticationToken( + "supabase-id-here", null, List.of(new SimpleGrantedAuthority("ROLE_USER"))); + SecurityContextHolder.getContext().setAuthentication(token); + // resolveUser does UUID.fromString on the name, so use a real UUID string. + String supabaseId = UUID.randomUUID().toString(); + UsernamePasswordAuthenticationToken realToken = + new UsernamePasswordAuthenticationToken( + supabaseId, null, List.of(new SimpleGrantedAuthority("ROLE_USER"))); + SecurityContextHolder.getContext().setAuthentication(realToken); + when(userRepository.findBySupabaseId(UUID.fromString(supabaseId))) + .thenReturn(Optional.of(user)); + } + + private static User makeUser(Long id, Long teamId) { + User user = new User(); + try { + // User entity doesn't expose Lombok setters for all fields. Set via reflection. + java.lang.reflect.Field idField = User.class.getDeclaredField("id"); + idField.setAccessible(true); + idField.set(user, id); + } catch (ReflectiveOperationException e) { + throw new RuntimeException(e); + } + if (teamId != null) { + stirling.software.proprietary.model.Team team = + new stirling.software.proprietary.model.Team(); + team.setId(teamId); + user.setTeam(team); + } + return user; + } + + private static HandlerMethod handlerMethodForFakeController() { + try { + Method m = FakeController.class.getDeclaredMethod("handleAuto"); + return new HandlerMethod(new FakeController(), m); + } catch (NoSuchMethodException e) { + throw new RuntimeException(e); + } + } + + private static HandlerMethod handlerMethodForPlain() { + try { + Method m = FakeController.class.getDeclaredMethod("handlePlain"); + return new HandlerMethod(new FakeController(), m); + } catch (NoSuchMethodException e) { + throw new RuntimeException(e); + } + } + + private static HandlerMethod handlerMethodForAutomation() { + try { + Method m = FakeController.class.getDeclaredMethod("handleAutomation"); + return new HandlerMethod(new FakeController(), m); + } catch (NoSuchMethodException e) { + throw new RuntimeException(e); + } + } + + private static HandlerMethod handlerMethodForAi() { + try { + Method m = FakeController.class.getDeclaredMethod("handleAi"); + return new HandlerMethod(new FakeController(), m); + } catch (NoSuchMethodException e) { + throw new RuntimeException(e); + } + } + + private static HandlerMethod handlerMethodForAiNoAutoJob() { + try { + Method m = FakeController.class.getDeclaredMethod("handleAiNoAutoJob"); + return new HandlerMethod(new FakeController(), m); + } catch (NoSuchMethodException e) { + throw new RuntimeException(e); + } + } + + private static HandlerMethod handlerMethodForClassLevelAi() { + try { + Method m = ClassLevelAiController.class.getDeclaredMethod("classLevelAi"); + return new HandlerMethod(new ClassLevelAiController(), m); + } catch (NoSuchMethodException e) { + throw new RuntimeException(e); + } + } + + private void authenticateWithApiKey(User user) { + stirling.software.proprietary.security.model.ApiKeyAuthenticationToken token = + new stirling.software.proprietary.security.model.ApiKeyAuthenticationToken( + user, "test-api-key", List.of(new SimpleGrantedAuthority("ROLE_API"))); + SecurityContextHolder.getContext().setAuthentication(token); + } + + static class FakeController { + @AutoJobPostMapping(value = "/x", resourceWeight = 1) + public void handleAuto() {} + + public void handlePlain() {} + + @AutoJobPostMapping(value = "/auto", resourceWeight = 1) + @stirling.software.saas.payg.cap.RequiresFeature( + stirling.software.saas.payg.model.FeatureGate.AUTOMATION) + public void handleAutomation() {} + + @AutoJobPostMapping(value = "/ai", resourceWeight = 1) + @stirling.software.saas.payg.cap.RequiresFeature( + stirling.software.saas.payg.model.FeatureGate.AI_SUPPORT) + public void handleAi() {} + + /** + * AI-controller shape: @RequiresFeature without @AutoJobPostMapping (JSON body / proxy). + */ + @stirling.software.saas.payg.cap.RequiresFeature( + stirling.software.saas.payg.model.FeatureGate.AI_SUPPORT) + public void handleAiNoAutoJob() {} + } + + /** Mirrors AiCreateController layout: @RequiresFeature on the class, plain methods. */ + @stirling.software.saas.payg.cap.RequiresFeature( + stirling.software.saas.payg.model.FeatureGate.AI_SUPPORT) + static class ClassLevelAiController { + public void classLevelAi() {} + } + + /** Placeholder so AutoCloseable resources flow in some helper methods. */ + @SuppressWarnings("unused") + private static void closeQuietly(AutoCloseable c) { + try { + c.close(); + } catch (Exception ignored) { + // ignored + } + } + + @SuppressWarnings("unused") + private static IOException unused() { + return null; + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/payg/filter/PaygOutputExtractorTest.java b/app/saas/src/test/java/stirling/software/saas/payg/filter/PaygOutputExtractorTest.java new file mode 100644 index 0000000000..81684b389c --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/payg/filter/PaygOutputExtractorTest.java @@ -0,0 +1,184 @@ +package stirling.software.saas.payg.filter; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.io.ByteArrayOutputStream; +import java.io.IOException; +import java.nio.charset.StandardCharsets; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.List; +import java.util.zip.ZipEntry; +import java.util.zip.ZipOutputStream; + +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.io.TempDir; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.common.util.TempFileManager; +import stirling.software.common.util.TempFileRegistry; + +class PaygOutputExtractorTest { + + private final TempFileManager tempFileManager = + new TempFileManager(new TempFileRegistry(), new ApplicationProperties()); + private final PaygOutputExtractor extractor = new PaygOutputExtractor(tempFileManager); + + @Test + void pdfContentType_returnsBodyPathVerbatim(@TempDir Path tmp) throws IOException { + Path body = tmp.resolve("body.pdf"); + Files.write(body, pdfBytes("hello")); + List out = extractor.extract("application/pdf", body); + assertThat(out).hasSize(1); + assertThat(out.get(0).path()).isEqualTo(body); + assertThat(out.get(0).ownedTempFile()).isNull(); + } + + @Test + void pdfContentType_withParameters_stillMatches(@TempDir Path tmp) throws IOException { + Path body = tmp.resolve("body.pdf"); + Files.write(body, pdfBytes("x")); + List out = + extractor.extract("application/pdf; charset=binary", body); + assertThat(out).hasSize(1); + } + + @Test + void pdfContentType_butBodyMissingPdfMagic_returnsEmpty(@TempDir Path tmp) throws IOException { + // Tool sets application/pdf but writes a non-PDF payload — we should NOT record it as + // OUTPUT lineage. Mirrors the magic-byte gate already enforced on the ZIP branch. + Path body = tmp.resolve("misleading.pdf"); + Files.write(body, "this is not actually a pdf".getBytes(StandardCharsets.UTF_8)); + assertThat(extractor.extract("application/pdf", body)).isEmpty(); + } + + @Test + void zipContentType_extractsOnlyPdfEntriesWithValidMagicBytes(@TempDir Path tmp) + throws IOException { + Path zip = tmp.resolve("out.zip"); + try (ZipOutputStream zos = new ZipOutputStream(Files.newOutputStream(zip))) { + writeEntry(zos, "doc1.pdf", pdfBytes("one")); + writeEntry(zos, "doc2.pdf", pdfBytes("two")); + writeEntry(zos, "fake.pdf", "not a pdf at all".getBytes(StandardCharsets.UTF_8)); + writeEntry(zos, "notes.txt", "ignored".getBytes(StandardCharsets.UTF_8)); + } + + List out = extractor.extract("application/zip", zip); + try { + assertThat(out).hasSize(2); + for (PaygOutputExtractor.ExtractedPdf p : out) { + byte[] head = new byte[5]; + Files.newInputStream(p.path()).read(head); + assertThat(head).startsWith("%PDF-".getBytes(StandardCharsets.UTF_8)); + assertThat(p.ownedTempFile()).isNotNull(); + } + } finally { + for (PaygOutputExtractor.ExtractedPdf p : out) { + p.close(); + } + } + } + + @Test + void nonPdfNonZipContentType_returnsEmpty(@TempDir Path tmp) throws IOException { + Path body = tmp.resolve("body.json"); + Files.writeString(body, "{\"error\":\"bad request\"}"); + assertThat(extractor.extract("application/json", body)).isEmpty(); + assertThat(extractor.extract("text/plain", body)).isEmpty(); + assertThat(extractor.extract(null, body)).isEmpty(); + } + + @Test + void nullBodyPath_returnsEmpty() { + assertThat(extractor.extract("application/pdf", null)).isEmpty(); + } + + @Test + void corruptZip_failsClosedAndReturnsEmpty(@TempDir Path tmp) throws IOException { + Path corrupt = tmp.resolve("garbage.zip"); + Files.write(corrupt, "not a zip file at all".getBytes(StandardCharsets.UTF_8)); + // Should NOT throw — fail-open: empty list, response still serves normally. + List out = extractor.extract("application/zip", corrupt); + assertThat(out).isEmpty(); + } + + @Test + void zipContentType_emptyZip_returnsEmpty(@TempDir Path tmp) throws IOException { + Path zip = tmp.resolve("empty.zip"); + try (ZipOutputStream zos = new ZipOutputStream(Files.newOutputStream(zip))) { + // no entries + } + assertThat(extractor.extract("application/zip", zip)).isEmpty(); + } + + @Test + void octetStreamWithPdfMagic_treatedAsPdf(@TempDir Path tmp) throws IOException { + // Stirling tool endpoints sometimes set Content-Type to + // application/octet-stream for streamed responses even when the body + // is a real PDF. The extractor must sniff magic bytes when the + // declared Content-Type is generic. + Path body = tmp.resolve("body.bin"); + Files.write(body, pdfBytes("real-pdf-content")); + List out = + extractor.extract("application/octet-stream", body); + assertThat(out).hasSize(1); + assertThat(out.get(0).path()).isEqualTo(body); + } + + @Test + void octetStreamWithZipMagic_unpackedAsZip(@TempDir Path tmp) throws IOException { + Path zip = tmp.resolve("body.bin"); + try (ZipOutputStream zos = new ZipOutputStream(Files.newOutputStream(zip))) { + writeEntry(zos, "page1.pdf", pdfBytes("a")); + writeEntry(zos, "page2.pdf", pdfBytes("b")); + } + List out = + extractor.extract("application/octet-stream", zip); + try { + assertThat(out).hasSize(2); + } finally { + for (PaygOutputExtractor.ExtractedPdf p : out) { + p.close(); + } + } + } + + @Test + void nullContentTypeWithZipMagic_unpackedAsZip(@TempDir Path tmp) throws IOException { + Path zip = tmp.resolve("body.bin"); + try (ZipOutputStream zos = new ZipOutputStream(Files.newOutputStream(zip))) { + writeEntry(zos, "page1.pdf", pdfBytes("a")); + } + List out = extractor.extract(null, zip); + try { + assertThat(out).hasSize(1); + } finally { + for (PaygOutputExtractor.ExtractedPdf p : out) { + p.close(); + } + } + } + + @Test + void octetStreamWithoutMagic_returnsEmpty(@TempDir Path tmp) throws IOException { + Path body = tmp.resolve("body.bin"); + Files.write(body, "neither pdf nor zip".getBytes(StandardCharsets.UTF_8)); + assertThat(extractor.extract("application/octet-stream", body)).isEmpty(); + } + + private static void writeEntry(ZipOutputStream zos, String name, byte[] data) + throws IOException { + zos.putNextEntry(new ZipEntry(name)); + zos.write(data); + zos.closeEntry(); + } + + private static byte[] pdfBytes(String payload) throws IOException { + ByteArrayOutputStream out = new ByteArrayOutputStream(); + // Minimal "looks like a PDF" — just the magic-byte prefix + filler. The extractor only + // checks magic bytes, not full PDF validity. + out.write("%PDF-1.4\n".getBytes(StandardCharsets.UTF_8)); + out.write(payload.getBytes(StandardCharsets.UTF_8)); + return out.toByteArray(); + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/payg/filter/PaygResponseBodyWrapperTest.java b/app/saas/src/test/java/stirling/software/saas/payg/filter/PaygResponseBodyWrapperTest.java new file mode 100644 index 0000000000..64e6ccb573 --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/payg/filter/PaygResponseBodyWrapperTest.java @@ -0,0 +1,210 @@ +package stirling.software.saas.payg.filter; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +import java.io.IOException; +import java.io.PrintWriter; +import java.nio.charset.StandardCharsets; +import java.nio.file.Files; +import java.nio.file.Path; + +import org.junit.jupiter.api.Test; +import org.springframework.mock.web.MockHttpServletResponse; + +import stirling.software.common.model.ApplicationProperties; +import stirling.software.common.util.TempFileManager; +import stirling.software.common.util.TempFileRegistry; + +class PaygResponseBodyWrapperTest { + + private final TempFileManager tempFileManager = + new TempFileManager(new TempFileRegistry(), new ApplicationProperties()); + + @Test + void inMemory_smallWrites_clientReceivesAndPathContainsSameBytes() throws Exception { + MockHttpServletResponse downstream = new MockHttpServletResponse(); + try (PaygResponseBodyWrapper wrapper = + new PaygResponseBodyWrapper(downstream, tempFileManager, 1024)) { + wrapper.getOutputStream().write("hello world".getBytes(StandardCharsets.UTF_8)); + wrapper.getOutputStream().flush(); + + assertThat(downstream.getContentAsString()).isEqualTo("hello world"); + assertThat(wrapper.bytesWritten()).isEqualTo(11); + + Path materialised = wrapper.materialisedPath(); + assertThat(materialised).isNotNull(); + assertThat(Files.readString(materialised, StandardCharsets.UTF_8)) + .isEqualTo("hello world"); + } + } + + @Test + void spill_crossThresholdMidChunk_capturesEverything() throws Exception { + MockHttpServletResponse downstream = new MockHttpServletResponse(); + try (PaygResponseBodyWrapper wrapper = + new PaygResponseBodyWrapper(downstream, tempFileManager, 8)) { + // First write: 5 bytes — fits in memory. + wrapper.getOutputStream().write("HELLO".getBytes(StandardCharsets.UTF_8)); + // Second write: 6 bytes — crosses the 8-byte threshold mid-chunk; whole chunk + // ends up on disk per the design (we don't split the crossing chunk). + wrapper.getOutputStream().write(" WORLD".getBytes(StandardCharsets.UTF_8)); + wrapper.getOutputStream().flush(); + + assertThat(downstream.getContentAsString()).isEqualTo("HELLO WORLD"); + assertThat(wrapper.bytesWritten()).isEqualTo(11); + + Path path = wrapper.materialisedPath(); + assertThat(Files.readString(path, StandardCharsets.UTF_8)).isEqualTo("HELLO WORLD"); + } + } + + @Test + void spill_largeSingleWrite_capturedToDisk() throws Exception { + MockHttpServletResponse downstream = new MockHttpServletResponse(); + byte[] payload = new byte[64 * 1024]; // 64 KiB + for (int i = 0; i < payload.length; i++) { + payload[i] = (byte) (i % 251); + } + try (PaygResponseBodyWrapper wrapper = + new PaygResponseBodyWrapper(downstream, tempFileManager, 1024)) { + wrapper.getOutputStream().write(payload); + wrapper.getOutputStream().flush(); + + assertThat(downstream.getContentAsByteArray()).isEqualTo(payload); + assertThat(wrapper.bytesWritten()).isEqualTo(payload.length); + + byte[] disk = Files.readAllBytes(wrapper.materialisedPath()); + assertThat(disk).isEqualTo(payload); + } + } + + @Test + void writer_pathTeesThroughToBuffer() throws Exception { + MockHttpServletResponse downstream = new MockHttpServletResponse(); + downstream.setCharacterEncoding("UTF-8"); + try (PaygResponseBodyWrapper wrapper = + new PaygResponseBodyWrapper(downstream, tempFileManager, 1024)) { + PrintWriter writer = wrapper.getWriter(); + writer.write("a-string-from-writer"); + writer.flush(); + + assertThat(downstream.getContentAsString()).isEqualTo("a-string-from-writer"); + Path path = wrapper.materialisedPath(); + assertThat(Files.readString(path, StandardCharsets.UTF_8)) + .isEqualTo("a-string-from-writer"); + } + } + + @Test + void mixingOutputStreamAndWriter_throwsPerServletSpec() throws Exception { + MockHttpServletResponse downstream = new MockHttpServletResponse(); + try (PaygResponseBodyWrapper wrapper = + new PaygResponseBodyWrapper(downstream, tempFileManager, 1024)) { + wrapper.getOutputStream(); + assertThatThrownBy(wrapper::getWriter).isInstanceOf(IllegalStateException.class); + } + + MockHttpServletResponse downstream2 = new MockHttpServletResponse(); + try (PaygResponseBodyWrapper wrapper2 = + new PaygResponseBodyWrapper(downstream2, tempFileManager, 1024)) { + wrapper2.getWriter(); + assertThatThrownBy(wrapper2::getOutputStream).isInstanceOf(IllegalStateException.class); + } + } + + @Test + void noBytesWritten_materialisedPathIsNull() throws Exception { + MockHttpServletResponse downstream = new MockHttpServletResponse(); + try (PaygResponseBodyWrapper wrapper = + new PaygResponseBodyWrapper(downstream, tempFileManager, 1024)) { + // Don't touch the output stream. + assertThat(wrapper.materialisedPath()).isNull(); + assertThat(wrapper.bytesWritten()).isZero(); + } + } + + @Test + void resetBuffer_inMemory_clearsAccumulatedBytes() throws Exception { + MockHttpServletResponse downstream = new MockHttpServletResponse(); + try (PaygResponseBodyWrapper wrapper = + new PaygResponseBodyWrapper(downstream, tempFileManager, 1024)) { + wrapper.getOutputStream().write("draft".getBytes(StandardCharsets.UTF_8)); + wrapper.resetBuffer(); + assertThat(wrapper.bytesWritten()).isZero(); + assertThat(wrapper.materialisedPath()).isNull(); + + wrapper.getOutputStream().write("final".getBytes(StandardCharsets.UTF_8)); + wrapper.getOutputStream().flush(); + assertThat(Files.readString(wrapper.materialisedPath(), StandardCharsets.UTF_8)) + .isEqualTo("final"); + } + } + + @Test + void resetBuffer_afterSpill_dropsSpillFileAndStartsFresh() throws Exception { + MockHttpServletResponse downstream = new MockHttpServletResponse(); + try (PaygResponseBodyWrapper wrapper = + new PaygResponseBodyWrapper(downstream, tempFileManager, 4)) { + wrapper.getOutputStream().write("draftover".getBytes(StandardCharsets.UTF_8)); // spill + assertThat(wrapper.bytesWritten()).isEqualTo(9); + wrapper.resetBuffer(); + assertThat(wrapper.bytesWritten()).isZero(); + assertThat(wrapper.materialisedPath()).isNull(); + + // Write again — should land in fresh memory, not the dropped spill. + wrapper.getOutputStream().write("ok".getBytes(StandardCharsets.UTF_8)); + wrapper.getOutputStream().flush(); + assertThat(Files.readString(wrapper.materialisedPath(), StandardCharsets.UTF_8)) + .isEqualTo("ok"); + } + } + + @Test + void close_isIdempotent() throws IOException { + MockHttpServletResponse downstream = new MockHttpServletResponse(); + PaygResponseBodyWrapper wrapper = + new PaygResponseBodyWrapper(downstream, tempFileManager, 8); + wrapper.getOutputStream().write("more-than-eight-bytes".getBytes(StandardCharsets.UTF_8)); + wrapper.materialisedPath(); // forces flush + wrapper.close(); + wrapper.close(); // second call must not throw + } + + @Test + void materialisedPath_isStableAcrossCalls_whenInMemory() throws Exception { + MockHttpServletResponse downstream = new MockHttpServletResponse(); + try (PaygResponseBodyWrapper wrapper = + new PaygResponseBodyWrapper(downstream, tempFileManager, 1024)) { + wrapper.getOutputStream().write("abc".getBytes(StandardCharsets.UTF_8)); + Path first = wrapper.materialisedPath(); + Path second = wrapper.materialisedPath(); + assertThat(first).isEqualTo(second); + } + } + + @Test + void singleByteWrites_areCorrectlyAccountedAcrossThreshold() throws Exception { + MockHttpServletResponse downstream = new MockHttpServletResponse(); + try (PaygResponseBodyWrapper wrapper = + new PaygResponseBodyWrapper(downstream, tempFileManager, 3)) { + for (int b : "ABCDE".getBytes(StandardCharsets.UTF_8)) { + wrapper.getOutputStream().write(b); + } + wrapper.getOutputStream().flush(); + assertThat(downstream.getContentAsString()).isEqualTo("ABCDE"); + assertThat(wrapper.bytesWritten()).isEqualTo(5); + assertThat(Files.readString(wrapper.materialisedPath(), StandardCharsets.UTF_8)) + .isEqualTo("ABCDE"); + } + } + + @Test + void negativeThreshold_rejected() { + assertThatThrownBy( + () -> + new PaygResponseBodyWrapper( + new MockHttpServletResponse(), tempFileManager, -1)) + .isInstanceOf(IllegalArgumentException.class); + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/payg/filter/PaygWebMvcConfigTest.java b/app/saas/src/test/java/stirling/software/saas/payg/filter/PaygWebMvcConfigTest.java new file mode 100644 index 0000000000..b766fec4a5 --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/payg/filter/PaygWebMvcConfigTest.java @@ -0,0 +1,22 @@ +package stirling.software.saas.payg.filter; + +import static org.assertj.core.api.Assertions.assertThat; + +import org.junit.jupiter.api.Test; + +/** + * Locks the interceptor ordering that makes "a blocked request is never charged" hold structurally: + * the {@link stirling.software.saas.payg.entitlement.EntitlementGuard} must run BEFORE the {@link + * PaygChargeInterceptor}. Spring skips a later interceptor's {@code preHandle} once an earlier one + * returns {@code false}, so guard-first means a refused (402) request never reaches {@code + * openProcess}. If these constants are ever reordered the wrong way, refused requests would start + * billing again — this test fails first. + */ +class PaygWebMvcConfigTest { + + @Test + void entitlementGuardRunsBeforeChargeInterceptor() { + assertThat(PaygWebMvcConfig.ENTITLEMENT_GUARD_ORDER) + .isLessThan(PaygWebMvcConfig.INTERCEPTOR_ORDER); + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/payg/job/JobServiceTest.java b/app/saas/src/test/java/stirling/software/saas/payg/job/JobServiceTest.java index 864860e8bd..ca7033405a 100644 --- a/app/saas/src/test/java/stirling/software/saas/payg/job/JobServiceTest.java +++ b/app/saas/src/test/java/stirling/software/saas/payg/job/JobServiceTest.java @@ -322,10 +322,17 @@ class JobServiceTest { return p; } - /** Test double: programmable detector + records observations. */ + /** + * Test double: programmable detector + records observations. Synthesises one signature per + * path; every path seen by extractSignatures is reverse-mapped so the post-dedupe flow + * (extractSignatures → detect/record by signature) keeps the path-based test assertions working + * unchanged. + */ private static class FakeDetector implements HashLineageDetector { private final Map matches = new HashMap<>(); private final Set observations = new java.util.HashSet<>(); + private final Map + pathBySignature = new HashMap<>(); void willMatch(Path input, LineageMatch match) { matches.put(input, match); @@ -335,14 +342,53 @@ class JobServiceTest { return observations.contains(jobId + "|" + file + "|" + kind); } + private stirling.software.saas.payg.lineage.LineageSignature sigFor(Path file) { + stirling.software.saas.payg.lineage.LineageSignature sig = + new stirling.software.saas.payg.lineage.LineageSignature( + "test", Integer.toHexString(file.toString().hashCode())); + pathBySignature.put(sig, file); + return sig; + } + @Override public Optional detect(Long userId, Path inputFile) { return Optional.ofNullable(matches.get(inputFile)); } + @Override + public Optional detect( + Long userId, Set signatures) { + for (stirling.software.saas.payg.lineage.LineageSignature sig : signatures) { + Path p = pathBySignature.get(sig); + if (p != null && matches.containsKey(p)) { + return Optional.of(matches.get(p)); + } + } + return Optional.empty(); + } + @Override public void record(UUID jobId, Path file, ArtifactKind kind) { observations.add(jobId + "|" + file + "|" + kind); } + + @Override + public void record( + UUID jobId, + Set signatures, + ArtifactKind kind) { + for (stirling.software.saas.payg.lineage.LineageSignature sig : signatures) { + Path p = pathBySignature.get(sig); + if (p != null) { + observations.add(jobId + "|" + p + "|" + kind); + } + } + } + + @Override + public Set extractSignatures( + Path file) { + return Set.of(sigFor(file)); + } } } diff --git a/app/saas/src/test/java/stirling/software/saas/payg/job/StaleJobCloserTest.java b/app/saas/src/test/java/stirling/software/saas/payg/job/StaleJobCloserTest.java index e526bf7cae..6f921a6fdb 100644 --- a/app/saas/src/test/java/stirling/software/saas/payg/job/StaleJobCloserTest.java +++ b/app/saas/src/test/java/stirling/software/saas/payg/job/StaleJobCloserTest.java @@ -1,35 +1,72 @@ package stirling.software.saas.payg.job; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.never; import static org.mockito.Mockito.verify; import static org.mockito.Mockito.when; +import java.util.List; +import java.util.UUID; + import org.junit.jupiter.api.Test; import org.mockito.Mockito; +import stirling.software.saas.payg.charge.JobChargeService; + /** - * Smoke test for the scheduler wiring. The interesting close logic is exercised in {@link - * JobServiceTest#closeStale_closesAllStaleJobs}; here we just confirm the scheduler bean delegates - * to {@code JobService.closeStale} and tolerates an empty result without erroring. + * Smoke test for the scheduler wiring. Confirms the scheduler closes each stale job through {@link + * JobChargeService#close} (so the Stripe meter afterCommit hook fires per job), isolates per-job + * failures, and tolerates an empty stale set without erroring. The close/meter logic itself is + * exercised in {@code JobChargeServiceTest}. */ class StaleJobCloserTest { - @Test - void closeStale_invokesJobService() { - JobService jobService = Mockito.mock(JobService.class); - when(jobService.closeStale()).thenReturn(3); - - new StaleJobCloser(jobService).closeStale(); - - verify(jobService).closeStale(); + private static ProcessingJob job(UUID id) { + ProcessingJob j = new ProcessingJob(); + j.setId(id); + return j; } @Test - void closeStale_zeroClosedDoesNotThrow() { + void closeStale_closesEachStaleJobThroughChargeService() { JobService jobService = Mockito.mock(JobService.class); - when(jobService.closeStale()).thenReturn(0); + JobChargeService chargeService = Mockito.mock(JobChargeService.class); + UUID a = UUID.randomUUID(); + UUID b = UUID.randomUUID(); + when(jobService.findStale()).thenReturn(List.of(job(a), job(b))); - new StaleJobCloser(jobService).closeStale(); + new StaleJobCloser(jobService, chargeService).closeStale(); - verify(jobService).closeStale(); + // Routed through the charge service (meter hook), NOT the bulk closeStale flip. + verify(chargeService).close(a); + verify(chargeService).close(b); + verify(jobService, never()).closeStale(); + } + + @Test + void closeStale_oneJobFailing_doesNotStrandTheRest() { + JobService jobService = Mockito.mock(JobService.class); + JobChargeService chargeService = Mockito.mock(JobChargeService.class); + UUID bad = UUID.randomUUID(); + UUID good = UUID.randomUUID(); + when(jobService.findStale()).thenReturn(List.of(job(bad), job(good))); + when(chargeService.close(bad)).thenThrow(new RuntimeException("boom")); + + // Must not propagate — the sweep continues to the next job. + new StaleJobCloser(jobService, chargeService).closeStale(); + + verify(chargeService).close(bad); + verify(chargeService).close(good); + } + + @Test + void closeStale_emptyStaleSet_doesNotTouchChargeService() { + JobService jobService = Mockito.mock(JobService.class); + JobChargeService chargeService = Mockito.mock(JobChargeService.class); + when(jobService.findStale()).thenReturn(List.of()); + + new StaleJobCloser(jobService, chargeService).closeStale(); + + verify(chargeService, never()).close(any()); } } diff --git a/app/saas/src/test/java/stirling/software/saas/payg/meter/PaygMeterReconcileSchedulerTest.java b/app/saas/src/test/java/stirling/software/saas/payg/meter/PaygMeterReconcileSchedulerTest.java new file mode 100644 index 0000000000..f9ab02c3d1 --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/payg/meter/PaygMeterReconcileSchedulerTest.java @@ -0,0 +1,121 @@ +package stirling.software.saas.payg.meter; + +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyInt; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.verifyNoInteractions; +import static org.mockito.Mockito.when; + +import java.time.Duration; +import java.util.List; +import java.util.UUID; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.mockito.Mockito; +import org.springframework.data.domain.Pageable; + +import io.micrometer.core.instrument.MeterRegistry; +import io.micrometer.core.instrument.simple.SimpleMeterRegistry; + +import stirling.software.saas.payg.policy.PaygTeamExtensions; +import stirling.software.saas.payg.repository.PaygMeterEventLogRepository; +import stirling.software.saas.payg.repository.PaygTeamExtensionsRepository; + +/** + * Unit tests for {@link PaygMeterReconcileScheduler}: retries unposted events for still-subscribed + * teams under the same idempotency key, skips teams that have since unsubscribed, and no-ops when + * disabled. + */ +class PaygMeterReconcileSchedulerTest { + + private PaygMeterEventLogRepository logRepo; + private PaygTeamExtensionsRepository teamExtRepo; + private PaygMeterReportingService meterReportingService; + private MeterRegistry meterRegistry; + + @BeforeEach + void setUp() { + logRepo = Mockito.mock(PaygMeterEventLogRepository.class); + teamExtRepo = Mockito.mock(PaygTeamExtensionsRepository.class); + meterReportingService = Mockito.mock(PaygMeterReportingService.class); + meterRegistry = new SimpleMeterRegistry(); + when(logRepo.countStuck(any())).thenReturn(0L); + } + + private PaygMeterReconcileScheduler scheduler(boolean enabled) { + return new PaygMeterReconcileScheduler( + logRepo, + teamExtRepo, + meterReportingService, + enabled, + Duration.ofMinutes(5), + 100, + meterRegistry); + } + + private static PaygMeterEventLog row(Long teamId, UUID jobId, String key, int units) { + PaygMeterEventLog e = new PaygMeterEventLog(); + e.setTeamId(teamId); + e.setJobId(jobId); + e.setIdempotencyKey(key); + e.setUnits(units); + return e; + } + + private static PaygTeamExtensions ext(Long teamId, String customerId, String subscriptionId) { + PaygTeamExtensions ext = new PaygTeamExtensions(); + ext.setTeamId(teamId); + ext.setStripeCustomerId(customerId); + ext.setPaygSubscriptionId(subscriptionId); + return ext; + } + + @Test + void reconcile_subscribedTeam_reMetersUnderSameKey() { + UUID jobId = UUID.randomUUID(); + when(logRepo.findRetryable(any(), any(), any(Pageable.class))) + .thenReturn(List.of(row(100L, jobId, "process:" + jobId + ":close", 4))); + when(teamExtRepo.findAllById(any())).thenReturn(List.of(ext(100L, "cus_live", "sub_live"))); + + scheduler(true).reconcile(); + + verify(meterReportingService) + .recordUsage(100L, "cus_live", 4, null, "process:" + jobId + ":close", jobId); + } + + @Test + void reconcile_teamUnsubscribedSince_isSkipped() { + UUID jobId = UUID.randomUUID(); + when(logRepo.findRetryable(any(), any(), any(Pageable.class))) + .thenReturn(List.of(row(100L, jobId, "process:" + jobId + ":close", 4))); + // Customer still present but no live subscription → no longer billable. + when(teamExtRepo.findAllById(any())).thenReturn(List.of(ext(100L, "cus_live", null))); + + scheduler(true).reconcile(); + + verify(meterReportingService, never()) + .recordUsage(any(), any(), anyInt(), any(), any(), any()); + } + + @Test + void reconcile_missingTeamExtensions_isSkipped() { + UUID jobId = UUID.randomUUID(); + when(logRepo.findRetryable(any(), any(), any(Pageable.class))) + .thenReturn(List.of(row(100L, jobId, "process:" + jobId + ":close", 4))); + when(teamExtRepo.findAllById(any())).thenReturn(List.of()); + + scheduler(true).reconcile(); + + verify(meterReportingService, never()) + .recordUsage(any(), any(), anyInt(), any(), any(), any()); + } + + @Test + void reconcile_disabled_isNoOp() { + scheduler(false).reconcile(); + + verifyNoInteractions(logRepo, teamExtRepo, meterReportingService); + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/payg/meter/PaygMeterReportingServiceTest.java b/app/saas/src/test/java/stirling/software/saas/payg/meter/PaygMeterReportingServiceTest.java new file mode 100644 index 0000000000..7ee2408276 --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/payg/meter/PaygMeterReportingServiceTest.java @@ -0,0 +1,245 @@ +package stirling.software.saas.payg.meter; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatCode; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyInt; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.times; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.net.ConnectException; +import java.util.Map; +import java.util.UUID; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.mockito.ArgumentCaptor; +import org.mockito.Mockito; +import org.springframework.http.HttpEntity; +import org.springframework.http.HttpHeaders; +import org.springframework.http.HttpMethod; +import org.springframework.http.HttpStatus; +import org.springframework.http.ResponseEntity; +import org.springframework.web.client.ResourceAccessException; +import org.springframework.web.client.RestTemplate; + +import io.micrometer.core.instrument.Counter; +import io.micrometer.core.instrument.MeterRegistry; +import io.micrometer.core.instrument.simple.SimpleMeterRegistry; + +import stirling.software.saas.payg.model.BillingCategory; +import stirling.software.saas.payg.repository.PaygMeterEventLogRepository; + +/** + * Covers the contract documented on {@link PaygMeterReportingService#recordUsage}: never throws, + * skips when endpoint is blank, counts non-2xx and exceptions on {@code payg.meter.errors}, and + * wraps every POST in a durable {@code payg_meter_event_log} row (pending → posted / failed). + */ +class PaygMeterReportingServiceTest { + + private static final String ENDPOINT = + "https://example.supabase.co/functions/v1/meter-payg-units"; + private static final String TOKEN = "test-service-role-token"; + private static final UUID JOB = UUID.fromString("00000000-0000-0000-0000-0000000000aa"); + + private RestTemplate restTemplate; + private PaygMeterEventLogRepository eventLogRepository; + private MeterRegistry meterRegistry; + private Counter errorsCounter; + + @BeforeEach + void setUp() { + restTemplate = Mockito.mock(RestTemplate.class); + eventLogRepository = Mockito.mock(PaygMeterEventLogRepository.class); + meterRegistry = new SimpleMeterRegistry(); + errorsCounter = meterRegistry.counter("payg.meter.errors"); + } + + private PaygMeterReportingService newService(String endpoint, String token) { + return new PaygMeterReportingService( + endpoint, token, restTemplate, eventLogRepository, meterRegistry); + } + + @Test + void recordUsage_happyPath_postsBodyLogsPendingThenPosted() { + when(restTemplate.exchange(eq(ENDPOINT), eq(HttpMethod.POST), any(), eq(String.class))) + .thenReturn(new ResponseEntity<>("{\"ok\":true}", HttpStatus.OK)); + + PaygMeterReportingService service = newService(ENDPOINT, TOKEN); + service.recordUsage(100L, "cus_abc", 5, BillingCategory.API, "process:job1:close", JOB); + + @SuppressWarnings("unchecked") + ArgumentCaptor>> entityCaptor = + ArgumentCaptor.forClass(HttpEntity.class); + verify(restTemplate) + .exchange( + eq(ENDPOINT), + eq(HttpMethod.POST), + entityCaptor.capture(), + eq(String.class)); + + HttpEntity> sent = entityCaptor.getValue(); + Map body = sent.getBody(); + assertThat(body).isNotNull(); + // JSON number — the edge fn type-checks team_id and ignores strings. + assertThat(body.get("team_id")).isEqualTo(100L); + assertThat(body.get("stripe_customer_id")).isEqualTo("cus_abc"); + assertThat(body.get("units")).isEqualTo(5); + assertThat(body.get("idempotency_key")).isEqualTo("process:job1:close"); + assertThat(body.get("metadata")).isEqualTo(Map.of("category", "API")); + + HttpHeaders headers = sent.getHeaders(); + assertThat(headers.getFirst("Authorization")).isEqualTo("Bearer " + TOKEN); + assertThat(headers.getContentType()).isNotNull(); + assertThat(headers.getContentType().toString()).startsWith("application/json"); + + // Durable audit: pending row written before the POST, stamped posted after success. + verify(eventLogRepository).insertPending(100L, JOB, "process:job1:close", 5); + verify(eventLogRepository).markPosted("process:job1:close"); + verify(eventLogRepository, never()).markFailed(any(), any(), any()); + assertThat(errorsCounter.count()).isZero(); + } + + @Test + void recordUsage_5xxResponse_marksFailedIncrementsErrorCounterAndDoesNotThrow() { + when(restTemplate.exchange(eq(ENDPOINT), eq(HttpMethod.POST), any(), eq(String.class))) + .thenReturn(new ResponseEntity<>("oops", HttpStatus.INTERNAL_SERVER_ERROR)); + + PaygMeterReportingService service = newService(ENDPOINT, TOKEN); + assertThatCode( + () -> + service.recordUsage( + 100L, + "cus_abc", + 3, + BillingCategory.AUTOMATION, + "process:job2:close", + JOB)) + .doesNotThrowAnyException(); + + verify(eventLogRepository).insertPending(100L, JOB, "process:job2:close", 3); + verify(eventLogRepository).markFailed(eq("process:job2:close"), eq("500"), any()); + verify(eventLogRepository, never()).markPosted(any()); + assertThat(errorsCounter.count()).isEqualTo(1.0); + } + + @Test + void recordUsage_connectionRefused_marksFailedIncrementsErrorCounterAndDoesNotThrow() { + when(restTemplate.exchange(eq(ENDPOINT), eq(HttpMethod.POST), any(), eq(String.class))) + .thenThrow(new ResourceAccessException("connect refused", new ConnectException())); + + PaygMeterReportingService service = newService(ENDPOINT, TOKEN); + assertThatCode( + () -> + service.recordUsage( + 100L, + "cus_abc", + 7, + BillingCategory.AI, + "process:job3:close", + JOB)) + .doesNotThrowAnyException(); + + verify(eventLogRepository).insertPending(100L, JOB, "process:job3:close", 7); + verify(eventLogRepository).markFailed(eq("process:job3:close"), eq("exception"), any()); + assertThat(errorsCounter.count()).isEqualTo(1.0); + } + + @Test + void recordUsage_runtimeException_swallowed() { + when(restTemplate.exchange(eq(ENDPOINT), eq(HttpMethod.POST), any(), eq(String.class))) + .thenThrow(new RuntimeException("boom")); + + PaygMeterReportingService service = newService(ENDPOINT, TOKEN); + assertThatCode( + () -> + service.recordUsage( + 100L, + "cus_abc", + 7, + BillingCategory.AI, + "process:job4:close", + JOB)) + .doesNotThrowAnyException(); + + verify(eventLogRepository).markFailed(eq("process:job4:close"), eq("exception"), any()); + assertThat(errorsCounter.count()).isEqualTo(1.0); + } + + @Test + void recordUsage_logPendingFailure_stillPostsAndDoesNotThrow() { + // A DB hiccup writing the audit row must not stop us metering the customer. + Mockito.doThrow(new RuntimeException("db down")) + .when(eventLogRepository) + .insertPending(any(), any(), any(), anyInt()); + when(restTemplate.exchange(eq(ENDPOINT), eq(HttpMethod.POST), any(), eq(String.class))) + .thenReturn(new ResponseEntity<>("{}", HttpStatus.OK)); + + PaygMeterReportingService service = newService(ENDPOINT, TOKEN); + assertThatCode( + () -> + service.recordUsage( + 100L, + "cus_abc", + 2, + BillingCategory.API, + "process:job9:close", + JOB)) + .doesNotThrowAnyException(); + + verify(restTemplate).exchange(eq(ENDPOINT), eq(HttpMethod.POST), any(), eq(String.class)); + } + + @Test + void recordUsage_blankEndpoint_noopsAndDoesNotCallRestTemplateOrLog() { + PaygMeterReportingService service = newService("", TOKEN); + service.recordUsage(100L, "cus_abc", 5, BillingCategory.API, "process:job5:close", JOB); + + verify(restTemplate, never()).exchange(any(String.class), any(), any(), any(Class.class)); + verify(eventLogRepository, never()).insertPending(any(), any(), any(), anyInt()); + assertThat(errorsCounter.count()).isZero(); + } + + @Test + void recordUsage_nullEndpoint_noopsAndDoesNotCallRestTemplate() { + PaygMeterReportingService service = newService(null, TOKEN); + service.recordUsage(100L, "cus_abc", 5, BillingCategory.API, "process:job6:close", JOB); + + verify(restTemplate, never()).exchange(any(String.class), any(), any(), any(Class.class)); + verify(eventLogRepository, never()).insertPending(any(), any(), any(), anyInt()); + assertThat(errorsCounter.count()).isZero(); + } + + @Test + void recordUsage_zeroUnits_noopsAndDoesNotCallRestTemplateOrLog() { + PaygMeterReportingService service = newService(ENDPOINT, TOKEN); + service.recordUsage(100L, "cus_abc", 0, BillingCategory.API, "process:job7:close", JOB); + + verify(restTemplate, never()).exchange(any(String.class), any(), any(), any(Class.class)); + verify(eventLogRepository, never()).insertPending(any(), any(), any(), anyInt()); + assertThat(errorsCounter.count()).isZero(); + } + + @Test + void recordUsage_blankServiceRoleToken_postsWithoutAuthorizationHeader() { + when(restTemplate.exchange(eq(ENDPOINT), eq(HttpMethod.POST), any(), eq(String.class))) + .thenReturn(new ResponseEntity<>("{}", HttpStatus.OK)); + + PaygMeterReportingService service = newService(ENDPOINT, ""); + service.recordUsage(100L, "cus_abc", 5, BillingCategory.API, "process:job8:close", JOB); + + @SuppressWarnings("unchecked") + ArgumentCaptor>> entityCaptor = + ArgumentCaptor.forClass(HttpEntity.class); + verify(restTemplate, times(1)) + .exchange( + eq(ENDPOINT), + eq(HttpMethod.POST), + entityCaptor.capture(), + eq(String.class)); + assertThat(entityCaptor.getValue().getHeaders().getFirst("Authorization")).isNull(); + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/payg/model/PaygEntitiesSmokeTest.java b/app/saas/src/test/java/stirling/software/saas/payg/model/PaygEntitiesSmokeTest.java index bffc17bd93..f1a24d0357 100644 --- a/app/saas/src/test/java/stirling/software/saas/payg/model/PaygEntitiesSmokeTest.java +++ b/app/saas/src/test/java/stirling/software/saas/payg/model/PaygEntitiesSmokeTest.java @@ -107,6 +107,16 @@ class PaygEntitiesSmokeTest { assertThat(entry.getAmountUnits()).isEqualTo(-4); } + @Test + void walletLedgerEntry_billingCategoryRoundTrips() { + WalletLedgerEntry entry = new WalletLedgerEntry(); + // Default (unset) is null — captured by both the legacy debit path and pre-V16 rows. + assertThat(entry.getBillingCategory()).isNull(); + + entry.setBillingCategory(BillingCategory.AUTOMATION); + assertThat(entry.getBillingCategory()).isEqualTo(BillingCategory.AUTOMATION); + } + @Test void walletPolicy_carriesSensibleDefaults() { WalletPolicy policy = new WalletPolicy(); @@ -148,4 +158,29 @@ class PaygEntitiesSmokeTest { assertThat(row.getDiffPct()).isNegative(); } + + @Test + void paygShadowCharge_billingCategoryAndJobSourceRoundTrip() { + PaygShadowCharge row = new PaygShadowCharge(); + assertThat(row.getBillingCategory()).isNull(); + assertThat(row.getJobSource()).isNull(); + + row.setBillingCategory(BillingCategory.AI); + row.setJobSource(JobSource.API); + + assertThat(row.getBillingCategory()).isEqualTo(BillingCategory.AI); + assertThat(row.getJobSource()).isEqualTo(JobSource.API); + } + + @Test + void billingCategory_listingOrderIsStable() { + // No downstream relies on ordinal() today, but the comment in the enum claims BYPASSED is + // declared first as the default sentinel — guard against an accidental reorder. + assertThat(BillingCategory.values()) + .containsExactly( + BillingCategory.BYPASSED, + BillingCategory.API, + BillingCategory.AI, + BillingCategory.AUTOMATION); + } } diff --git a/app/saas/src/test/java/stirling/software/saas/payg/policy/PricingPolicyServiceTest.java b/app/saas/src/test/java/stirling/software/saas/payg/policy/PricingPolicyServiceTest.java index dace97d367..3c552d205a 100644 --- a/app/saas/src/test/java/stirling/software/saas/payg/policy/PricingPolicyServiceTest.java +++ b/app/saas/src/test/java/stirling/software/saas/payg/policy/PricingPolicyServiceTest.java @@ -98,6 +98,33 @@ class PricingPolicyServiceTest { .hasMessageContaining("No default pricing_policy row"); } + @Test + void nullTeamId_returnsDefaultPolicyDirectly_noCacheLookup() { + PricingPolicy result = service.getEffectivePolicy(null); + + assertThat(result).isEqualTo(defaultPolicy); + verify(extensionsRepo, never()).findById(any()); + // Team-less users go straight to the default — no cache pollution either. + assertThat(service.cacheSize()).isZero(); + } + + @Test + void nullTeamIdUncached_returnsDefaultPolicyDirectly() { + PricingPolicy result = service.getEffectivePolicyUncached(null); + + assertThat(result).isEqualTo(defaultPolicy); + verify(extensionsRepo, never()).findById(any()); + } + + @Test + void nullTeamId_noDefaultExists_throws() { + when(policyRepo.findFirstByIsDefaultTrue()).thenReturn(Optional.empty()); + + assertThatThrownBy(() -> service.getEffectivePolicy(null)) + .isInstanceOf(IllegalStateException.class) + .hasMessageContaining("No default pricing_policy row"); + } + @Test void secondCallHitsCache_noRepoLookup() { when(extensionsRepo.findById(42L)).thenReturn(Optional.empty()); diff --git a/app/saas/src/test/java/stirling/software/saas/payg/policy/admin/PricingPolicyAdminControllerTest.java b/app/saas/src/test/java/stirling/software/saas/payg/policy/admin/PricingPolicyAdminControllerTest.java index 595ce0ecd0..3216361919 100644 --- a/app/saas/src/test/java/stirling/software/saas/payg/policy/admin/PricingPolicyAdminControllerTest.java +++ b/app/saas/src/test/java/stirling/software/saas/payg/policy/admin/PricingPolicyAdminControllerTest.java @@ -26,9 +26,8 @@ import stirling.software.saas.payg.policy.admin.PolicyDtos.PolicyResponse; import stirling.software.saas.payg.policy.admin.PolicyDtos.TeamOverrideRequest; /** - * Tests {@link PricingPolicyAdminController} as a plain Java unit (matching {@code - * CreditControllerApiKeyTest}'s style — no MockMvc layer). Covers happy paths and the controller's - * error mapping (4xx for validation, 404 for missing rows). + * Tests {@link PricingPolicyAdminController} as a plain Java unit (no MockMvc layer). Covers happy + * paths and the controller's error mapping (4xx for validation, 404 for missing rows). */ @ExtendWith(MockitoExtension.class) class PricingPolicyAdminControllerTest { diff --git a/app/saas/src/test/java/stirling/software/saas/security/SupabaseAuthenticationFilterTest.java b/app/saas/src/test/java/stirling/software/saas/security/SupabaseAuthenticationFilterTest.java index 42238502fd..e906655edc 100644 --- a/app/saas/src/test/java/stirling/software/saas/security/SupabaseAuthenticationFilterTest.java +++ b/app/saas/src/test/java/stirling/software/saas/security/SupabaseAuthenticationFilterTest.java @@ -30,7 +30,6 @@ import org.springframework.security.oauth2.jwt.Jwt; import org.springframework.security.oauth2.jwt.JwtDecoder; import org.springframework.security.oauth2.jwt.JwtException; -import stirling.software.proprietary.model.Team; import stirling.software.proprietary.security.model.ApiKeyAuthenticationToken; import stirling.software.proprietary.security.model.AuthenticationType; import stirling.software.proprietary.security.model.User; @@ -46,7 +45,6 @@ class SupabaseAuthenticationFilterTest { @Mock private TeamService teamService; @Mock private UserService userService; @Mock private SupabaseUserService supabaseUserService; - @Mock private stirling.software.saas.service.CreditService creditService; @Mock private stirling.software.saas.service.SaasTeamService saasTeamService; @Mock private JwtDecoder jwtDecoder; @@ -60,12 +58,7 @@ class SupabaseAuthenticationFilterTest { SecurityContextHolder.clearContext(); filter = new SupabaseAuthenticationFilter( - teamService, - userService, - supabaseUserService, - creditService, - saasTeamService, - jwtDecoder); + teamService, userService, supabaseUserService, saasTeamService, jwtDecoder); request = new MockHttpServletRequest(); response = new MockHttpServletResponse(); chain = new MockFilterChain(); @@ -161,7 +154,6 @@ class SupabaseAuthenticationFilterTest { when(supabaseUserService.getUser(supabaseId)) .thenReturn(supabaseUserMatching(supabaseId, "bob@example.com", false)); when(userService.findBySupabaseId(supabaseId)).thenReturn(Optional.empty()); - when(teamService.getOrCreateDefaultTeam()).thenReturn(new Team()); when(userService.saveUser(any())).thenAnswer(inv -> inv.getArgument(0)); request.setRequestURI("/api/v1/something"); @@ -172,6 +164,9 @@ class SupabaseAuthenticationFilterTest { verify(userService, times(1)).saveUser(any(User.class)); verify(supabaseUserService).createSupabaseUser(supabaseId, "bob@example.com", false); + // New users get their own personal team, never the shared Default team. + verify(saasTeamService).ensurePersonalTeam(any(User.class)); + verify(teamService, never()).getOrCreateDefaultTeam(); assertThat(SecurityContextHolder.getContext().getAuthentication()) .isInstanceOf(EnhancedJwtAuthenticationToken.class); } @@ -185,7 +180,6 @@ class SupabaseAuthenticationFilterTest { when(supabaseUserService.getUser(supabaseId)) .thenReturn(supabaseUserMatching(supabaseId, "carol@example.com", false)); when(userService.findBySupabaseId(supabaseId)).thenReturn(Optional.empty()); - when(teamService.getOrCreateDefaultTeam()).thenReturn(new Team()); when(userService.saveUser(any(User.class))) .thenAnswer( inv -> { @@ -214,7 +208,6 @@ class SupabaseAuthenticationFilterTest { when(supabaseUserService.getUser(supabaseId)) .thenReturn(supabaseUserMatching(supabaseId, "dave@example.com", false)); when(userService.findBySupabaseId(supabaseId)).thenReturn(Optional.empty()); - when(teamService.getOrCreateDefaultTeam()).thenReturn(new Team()); when(userService.saveUser(any(User.class))) .thenAnswer( inv -> { @@ -243,7 +236,6 @@ class SupabaseAuthenticationFilterTest { when(supabaseUserService.getUser(supabaseId)) .thenReturn(supabaseUserMatching(supabaseId, "eve@example.com", false)); when(userService.findBySupabaseId(supabaseId)).thenReturn(Optional.empty()); - when(teamService.getOrCreateDefaultTeam()).thenReturn(new Team()); when(userService.saveUser(any(User.class))) .thenAnswer( inv -> { diff --git a/app/saas/src/test/java/stirling/software/saas/security/SupabaseSecurityConfigTest.java b/app/saas/src/test/java/stirling/software/saas/security/SupabaseSecurityConfigTest.java index 5c21b88a23..49779fce8e 100644 --- a/app/saas/src/test/java/stirling/software/saas/security/SupabaseSecurityConfigTest.java +++ b/app/saas/src/test/java/stirling/software/saas/security/SupabaseSecurityConfigTest.java @@ -9,17 +9,25 @@ import java.util.List; import java.util.Map; import java.util.UUID; +import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.Test; import org.springframework.security.authentication.AbstractAuthenticationToken; import org.springframework.security.core.GrantedAuthority; +import org.springframework.security.core.context.SecurityContextHolder; import org.springframework.security.oauth2.core.OAuth2TokenValidatorResult; import org.springframework.security.oauth2.jwt.Jwt; +import stirling.software.proprietary.security.model.User; import stirling.software.saas.security.SupabaseSecurityConfig.SupabaseTokenValidator; /** Unit tests for the JWT-claim → Spring authorities mapping. */ class SupabaseSecurityConfigTest { + @AfterEach + void clearSecurityContext() { + SecurityContextHolder.clearContext(); + } + @Test void anonymousJwtGetsLimitedApiUserRole() { Jwt jwt = jwtWith(true, null, null, null, List.of()); @@ -75,6 +83,47 @@ class SupabaseSecurityConfigTest { .doesNotContain("ROLE_", "PERM_"); } + @Test + void carriesUserPrincipalFromContextBuiltByAuthFilter() { + Jwt jwt = jwtWith(false, "alice@example.com", "authenticated", null, List.of()); + User user = new User(); + SecurityContextHolder.getContext() + .setAuthentication( + new EnhancedJwtAuthenticationToken( + jwt, List.of(), "alice@example.com", jwt.getSubject(), user)); + + AbstractAuthenticationToken auth = SupabaseSecurityConfig.toAuthentication(jwt); + + assertThat(auth.getPrincipal()).isSameAs(user); + } + + @Test + void principalStaysJwtWithoutContextUser() { + Jwt jwt = jwtWith(false, "alice@example.com", "authenticated", null, List.of()); + + AbstractAuthenticationToken auth = SupabaseSecurityConfig.toAuthentication(jwt); + + assertThat(auth.getPrincipal()).isSameAs(jwt); + } + + @Test + void ignoresContextUserForDifferentSubject() { + Jwt jwt = jwtWith(false, "alice@example.com", "authenticated", null, List.of()); + Jwt other = jwtWith(false, "bob@example.com", "authenticated", null, List.of()); + SecurityContextHolder.getContext() + .setAuthentication( + new EnhancedJwtAuthenticationToken( + other, + List.of(), + "bob@example.com", + other.getSubject(), + new User())); + + AbstractAuthenticationToken auth = SupabaseSecurityConfig.toAuthentication(jwt); + + assertThat(auth.getPrincipal()).isSameAs(jwt); + } + private static Jwt jwtWith( boolean anonymous, String email, diff --git a/app/saas/src/test/java/stirling/software/saas/security/TeamLeaderPolicyManagementAuthorityTest.java b/app/saas/src/test/java/stirling/software/saas/security/TeamLeaderPolicyManagementAuthorityTest.java new file mode 100644 index 0000000000..70cd360d7c --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/security/TeamLeaderPolicyManagementAuthorityTest.java @@ -0,0 +1,40 @@ +package stirling.software.saas.security; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.Mockito.when; + +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; + +/** SaaS policy context: team leader may edit; scoping uses the user's team. */ +@ExtendWith(MockitoExtension.class) +class TeamLeaderPolicyManagementAuthorityTest { + + @Mock private TeamSecurityExpressions teamSecurity; + + private TeamLeaderPolicyManagementAuthority authority() { + return new TeamLeaderPolicyManagementAuthority(teamSecurity); + } + + @Test + void teamLeaderMayEditPolicies() { + when(teamSecurity.isCurrentUserTeamLeader()).thenReturn(true); + assertTrue(authority().canEditPolicies()); + } + + @Test + void nonLeaderMayNot() { + when(teamSecurity.isCurrentUserTeamLeader()).thenReturn(false); + assertFalse(authority().canEditPolicies()); + } + + @Test + void currentUserTeamIdDelegatesToTeamSecurity() { + when(teamSecurity.currentUserTeamId()).thenReturn(9L); + assertEquals(9L, authority().currentUserTeamId()); + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/security/TeamSecurityExpressionsTest.java b/app/saas/src/test/java/stirling/software/saas/security/TeamSecurityExpressionsTest.java new file mode 100644 index 0000000000..23226e4467 --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/security/TeamSecurityExpressionsTest.java @@ -0,0 +1,118 @@ +package stirling.software.saas.security; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.Mockito.lenient; +import static org.mockito.Mockito.when; + +import java.util.List; +import java.util.Optional; + +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.springframework.security.authentication.UsernamePasswordAuthenticationToken; +import org.springframework.security.core.context.SecurityContextHolder; + +import stirling.software.common.model.enumeration.TeamRole; +import stirling.software.proprietary.model.Team; +import stirling.software.proprietary.security.model.User; +import stirling.software.proprietary.security.service.UserService; +import stirling.software.saas.model.TeamMembership; +import stirling.software.saas.repository.TeamMembershipRepository; + +/** + * {@link TeamSecurityExpressions#isCurrentUserTeamLeader()} — used to gate policy editing on SaaS. + */ +@ExtendWith(MockitoExtension.class) +class TeamSecurityExpressionsTest { + + @Mock private TeamMembershipRepository membershipRepository; + @Mock private UserService userService; + + private static final long TEAM_ID = 2L; + private static final long USER_ID = 1L; + + private TeamSecurityExpressions expressions() { + return new TeamSecurityExpressions(membershipRepository, userService); + } + + @AfterEach + void clearContext() { + SecurityContextHolder.clearContext(); + } + + private void authenticateAsUserWithTeam(boolean hasTeam) { + User user = new User(); + user.setId(USER_ID); + if (hasTeam) { + Team team = new Team(); + team.setId(TEAM_ID); + user.setTeam(team); + } + // API-key auth path: the principal is the User entity itself. + SecurityContextHolder.getContext() + .setAuthentication(new UsernamePasswordAuthenticationToken(user, null, List.of())); + } + + private TeamMembership membershipWithRole(TeamRole role) { + TeamMembership membership = new TeamMembership(); + membership.setRole(role); + return membership; + } + + @Test + void leaderOfOwnTeamIsLeader() { + authenticateAsUserWithTeam(true); + when(membershipRepository.findByTeamIdAndUserId(TEAM_ID, USER_ID)) + .thenReturn(Optional.of(membershipWithRole(TeamRole.LEADER))); + assertTrue(expressions().isCurrentUserTeamLeader()); + } + + @Test + void regularMemberIsNotLeader() { + authenticateAsUserWithTeam(true); + when(membershipRepository.findByTeamIdAndUserId(TEAM_ID, USER_ID)) + .thenReturn(Optional.of(membershipWithRole(TeamRole.MEMBER))); + assertFalse(expressions().isCurrentUserTeamLeader()); + } + + @Test + void noMembershipIsNotLeader() { + authenticateAsUserWithTeam(true); + when(membershipRepository.findByTeamIdAndUserId(TEAM_ID, USER_ID)) + .thenReturn(Optional.empty()); + assertFalse(expressions().isCurrentUserTeamLeader()); + } + + @Test + void userWithoutTeamIsNotLeader() { + authenticateAsUserWithTeam(false); + assertFalse(expressions().isCurrentUserTeamLeader()); + } + + @Test + void currentUserTeamIdReturnsTheUsersTeam() { + authenticateAsUserWithTeam(true); + assertEquals(TEAM_ID, expressions().currentUserTeamId()); + } + + @Test + void currentUserTeamIdIsNullWithoutTeam() { + authenticateAsUserWithTeam(false); + assertNull(expressions().currentUserTeamId()); + } + + @Test + void unauthenticatedIsNotLeader() { + // No authentication set on the context. + lenient() + .when(membershipRepository.findByTeamIdAndUserId(TEAM_ID, USER_ID)) + .thenReturn(Optional.of(membershipWithRole(TeamRole.LEADER))); + assertFalse(expressions().isCurrentUserTeamLeader()); + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/service/AnonymousUserCleanupServiceTest.java b/app/saas/src/test/java/stirling/software/saas/service/AnonymousUserCleanupServiceTest.java new file mode 100644 index 0000000000..40ddcc7ac6 --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/service/AnonymousUserCleanupServiceTest.java @@ -0,0 +1,351 @@ +package stirling.software.saas.service; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.times; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.verifyNoInteractions; +import static org.mockito.Mockito.when; + +import java.time.LocalDateTime; +import java.util.ArrayList; +import java.util.List; +import java.util.UUID; +import java.util.stream.Stream; + +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.ArgumentCaptor; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.test.util.ReflectionTestUtils; + +import stirling.software.proprietary.security.database.repository.UserRepository; +import stirling.software.saas.repository.SupabaseUserRepository; + +/** + * Unit tests for {@link AnonymousUserCleanupService}. + * + *

The service is a {@code @Scheduled} cleanup job. Its three {@code @Value} fields ({@code + * anonEnabled}, {@code retentionDays}, {@code batchSize}) are field-injected, so each scenario + * primes them with {@link ReflectionTestUtils}. The {@code cleanup()} method is invoked directly + * (no Spring scheduler). Both repositories return id {@link Stream}s that the service partitions + * into fixed-size batches via {@code Collectors.groupingBy}, then deletes each batch. + * + *

Important: the streams are consumed inside try-with-resources, so the SupabaseUser stream is a + * {@code Stream} and the User stream is a {@code Stream}; each stub must return a fresh + * stream (a {@code Stream} is single-use). + */ +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class AnonymousUserCleanupServiceTest { + + @Mock private UserRepository userRepository; + @Mock private SupabaseUserRepository supabaseUserRepository; + + private AnonymousUserCleanupService service; + + private AnonymousUserCleanupService newService( + boolean anonEnabled, int retentionDays, int batchSize) { + AnonymousUserCleanupService s = + new AnonymousUserCleanupService(userRepository, supabaseUserRepository); + ReflectionTestUtils.setField(s, "anonEnabled", anonEnabled); + ReflectionTestUtils.setField(s, "retentionDays", retentionDays); + ReflectionTestUtils.setField(s, "batchSize", batchSize); + return s; + } + + private static List uuids(int n) { + List out = new ArrayList<>(n); + for (int i = 0; i < n; i++) { + out.add(UUID.randomUUID()); + } + return out; + } + + private static List longs(int n) { + List out = new ArrayList<>(n); + for (long i = 0; i < n; i++) { + out.add(i); + } + return out; + } + + @Nested + @DisplayName("guard clauses - no work performed") + class Guards { + + @Test + @DisplayName("anonymous auth disabled: returns early, touches neither repository") + void anonDisabled_noop() { + service = newService(false, 30, 100); + + service.cleanup(); + + verifyNoInteractions(userRepository, supabaseUserRepository); + } + + @Test + @DisplayName("retentionDays == 0: returns early, touches neither repository") + void zeroRetention_noop() { + service = newService(true, 0, 100); + + service.cleanup(); + + verifyNoInteractions(userRepository, supabaseUserRepository); + } + + @Test + @DisplayName("negative retentionDays: returns early, touches neither repository") + void negativeRetention_noop() { + service = newService(true, -5, 100); + + service.cleanup(); + + verifyNoInteractions(userRepository, supabaseUserRepository); + } + } + + @Nested + @DisplayName("cutoff date derivation") + class CutoffDate { + + @Test + @DisplayName("queries both repositories with a cutoff ~retentionDays before now") + void cutoffIsRetentionDaysBeforeNow() { + service = newService(true, 30, 100); + when(supabaseUserRepository.findByCreatedAtBeforeAndIsAnonymousTrue( + any(LocalDateTime.class))) + .thenReturn(Stream.empty()); + when(userRepository.findByUsernameIsNullAndCreatedAtBefore(any(LocalDateTime.class))) + .thenReturn(Stream.empty()); + + LocalDateTime expectedLower = LocalDateTime.now().minusDays(30).minusMinutes(1); + LocalDateTime expectedUpper = LocalDateTime.now().minusDays(30).plusMinutes(1); + + service.cleanup(); + + ArgumentCaptor supaCutoff = ArgumentCaptor.forClass(LocalDateTime.class); + verify(supabaseUserRepository) + .findByCreatedAtBeforeAndIsAnonymousTrue(supaCutoff.capture()); + assertThat(supaCutoff.getValue()).isBetween(expectedLower, expectedUpper); + + ArgumentCaptor userCutoff = ArgumentCaptor.forClass(LocalDateTime.class); + verify(userRepository).findByUsernameIsNullAndCreatedAtBefore(userCutoff.capture()); + assertThat(userCutoff.getValue()).isBetween(expectedLower, expectedUpper); + } + + @Test + @DisplayName("both repositories receive the same cutoff instant") + void sameCutoffForBothRepositories() { + service = newService(true, 7, 100); + when(supabaseUserRepository.findByCreatedAtBeforeAndIsAnonymousTrue( + any(LocalDateTime.class))) + .thenReturn(Stream.empty()); + when(userRepository.findByUsernameIsNullAndCreatedAtBefore(any(LocalDateTime.class))) + .thenReturn(Stream.empty()); + + service.cleanup(); + + ArgumentCaptor supaCutoff = ArgumentCaptor.forClass(LocalDateTime.class); + verify(supabaseUserRepository) + .findByCreatedAtBeforeAndIsAnonymousTrue(supaCutoff.capture()); + ArgumentCaptor userCutoff = ArgumentCaptor.forClass(LocalDateTime.class); + verify(userRepository).findByUsernameIsNullAndCreatedAtBefore(userCutoff.capture()); + + assertThat(supaCutoff.getValue()).isEqualTo(userCutoff.getValue()); + } + } + + @Nested + @DisplayName("empty result sets") + class EmptyStreams { + + @Test + @DisplayName("no stale users: queries run but no batch delete is issued") + void noStaleUsers_noDeletes() { + service = newService(true, 30, 100); + when(supabaseUserRepository.findByCreatedAtBeforeAndIsAnonymousTrue( + any(LocalDateTime.class))) + .thenReturn(Stream.empty()); + when(userRepository.findByUsernameIsNullAndCreatedAtBefore(any(LocalDateTime.class))) + .thenReturn(Stream.empty()); + + service.cleanup(); + + verify(supabaseUserRepository) + .findByCreatedAtBeforeAndIsAnonymousTrue(any(LocalDateTime.class)); + verify(userRepository).findByUsernameIsNullAndCreatedAtBefore(any(LocalDateTime.class)); + verify(supabaseUserRepository, never()).deleteAllByIdInBatch(any()); + verify(userRepository, never()).deleteAllByIdInBatch(any()); + } + } + + @Nested + @DisplayName("single-batch deletion") + class SingleBatch { + + @Test + @DisplayName("count below batch size deletes everything in one batch per repository") + void belowBatchSize_singleDelete() { + service = newService(true, 30, 100); + List supaIds = uuids(3); + List userIds = longs(5); + when(supabaseUserRepository.findByCreatedAtBeforeAndIsAnonymousTrue( + any(LocalDateTime.class))) + .thenReturn(supaIds.stream()); + when(userRepository.findByUsernameIsNullAndCreatedAtBefore(any(LocalDateTime.class))) + .thenReturn(userIds.stream()); + + service.cleanup(); + + verify(supabaseUserRepository, times(1)).deleteAllByIdInBatch(supaIds); + verify(userRepository, times(1)).deleteAllByIdInBatch(userIds); + } + + @Test + @DisplayName("count exactly equal to batch size still produces exactly one batch") + void exactlyBatchSize_singleDelete() { + service = newService(true, 30, 4); + List supaIds = uuids(4); + when(supabaseUserRepository.findByCreatedAtBeforeAndIsAnonymousTrue( + any(LocalDateTime.class))) + .thenReturn(supaIds.stream()); + when(userRepository.findByUsernameIsNullAndCreatedAtBefore(any(LocalDateTime.class))) + .thenReturn(Stream.empty()); + + service.cleanup(); + + verify(supabaseUserRepository, times(1)).deleteAllByIdInBatch(supaIds); + } + } + + @Nested + @DisplayName("multi-batch partitioning") + class MultiBatch { + + @Test + @DisplayName("supabase ids split into fixed-size batches with a final remainder batch") + void supabaseIdsPartitioned() { + service = newService(true, 30, 2); + // 5 ids, batch 2 -> batches of [2,2,1] + List supaIds = uuids(5); + when(supabaseUserRepository.findByCreatedAtBeforeAndIsAnonymousTrue( + any(LocalDateTime.class))) + .thenReturn(supaIds.stream()); + when(userRepository.findByUsernameIsNullAndCreatedAtBefore(any(LocalDateTime.class))) + .thenReturn(Stream.empty()); + + service.cleanup(); + + @SuppressWarnings("unchecked") + ArgumentCaptor> captor = ArgumentCaptor.forClass(List.class); + verify(supabaseUserRepository, times(3)).deleteAllByIdInBatch(captor.capture()); + + List> batches = captor.getAllValues(); + // Two full batches and one remainder batch (order of values() not asserted on). + assertThat(batches).extracting(List::size).containsExactlyInAnyOrder(2, 2, 1); + // Every id is deleted exactly once across all batches. + assertThat(batches.stream().flatMap(List::stream)) + .containsExactlyInAnyOrderElementsOf(supaIds); + } + + @Test + @DisplayName("user ids split into fixed-size batches; even multiple yields no remainder") + void userIdsPartitionedEvenly() { + service = newService(true, 30, 3); + // 6 ids, batch 3 -> batches of [3,3] + List userIds = longs(6); + when(supabaseUserRepository.findByCreatedAtBeforeAndIsAnonymousTrue( + any(LocalDateTime.class))) + .thenReturn(Stream.empty()); + when(userRepository.findByUsernameIsNullAndCreatedAtBefore(any(LocalDateTime.class))) + .thenReturn(userIds.stream()); + + service.cleanup(); + + @SuppressWarnings("unchecked") + ArgumentCaptor> captor = ArgumentCaptor.forClass(List.class); + verify(userRepository, times(2)).deleteAllByIdInBatch(captor.capture()); + + List> batches = captor.getAllValues(); + assertThat(batches).extracting(List::size).containsExactlyInAnyOrder(3, 3); + // Within a batch, encounter order is preserved by groupingBy. + assertThat(batches).anySatisfy(b -> assertThat(b).containsExactly(0L, 1L, 2L)); + assertThat(batches).anySatisfy(b -> assertThat(b).containsExactly(3L, 4L, 5L)); + } + + @Test + @DisplayName("batch size of 1 produces one delete call per id") + void batchSizeOne_oneDeletePerId() { + service = newService(true, 30, 1); + List userIds = longs(4); + when(supabaseUserRepository.findByCreatedAtBeforeAndIsAnonymousTrue( + any(LocalDateTime.class))) + .thenReturn(Stream.empty()); + when(userRepository.findByUsernameIsNullAndCreatedAtBefore(any(LocalDateTime.class))) + .thenReturn(userIds.stream()); + + service.cleanup(); + + @SuppressWarnings("unchecked") + ArgumentCaptor> captor = ArgumentCaptor.forClass(List.class); + verify(userRepository, times(4)).deleteAllByIdInBatch(captor.capture()); + assertThat(captor.getAllValues()).extracting(List::size).containsExactly(1, 1, 1, 1); + } + } + + @Nested + @DisplayName("both repositories are cleaned in one run") + class BothRepositories { + + @Test + @DisplayName("supabase users are processed before legacy users, each in their own batches") + void supabaseThenUsers() { + service = newService(true, 30, 2); + List supaIds = uuids(3); // [2,1] + List userIds = longs(3); // [2,1] + when(supabaseUserRepository.findByCreatedAtBeforeAndIsAnonymousTrue( + any(LocalDateTime.class))) + .thenReturn(supaIds.stream()); + when(userRepository.findByUsernameIsNullAndCreatedAtBefore(any(LocalDateTime.class))) + .thenReturn(userIds.stream()); + + service.cleanup(); + + verify(supabaseUserRepository, times(2)).deleteAllByIdInBatch(any()); + verify(userRepository, times(2)).deleteAllByIdInBatch(any()); + } + } + + @Nested + @DisplayName("stream lifecycle") + class StreamLifecycle { + + @Test + @DisplayName("each streamed result is closed via try-with-resources") + void streamsAreClosed() { + service = newService(true, 30, 100); + + boolean[] supaClosed = {false}; + boolean[] userClosed = {false}; + Stream supaStream = uuids(2).stream().onClose(() -> supaClosed[0] = true); + Stream userStream = longs(2).stream().onClose(() -> userClosed[0] = true); + when(supabaseUserRepository.findByCreatedAtBeforeAndIsAnonymousTrue( + any(LocalDateTime.class))) + .thenReturn(supaStream); + when(userRepository.findByUsernameIsNullAndCreatedAtBefore(any(LocalDateTime.class))) + .thenReturn(userStream); + + service.cleanup(); + + assertThat(supaClosed[0]).as("supabase id stream closed").isTrue(); + assertThat(userClosed[0]).as("user id stream closed").isTrue(); + } + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/service/RateLimitServiceTest.java b/app/saas/src/test/java/stirling/software/saas/service/RateLimitServiceTest.java new file mode 100644 index 0000000000..7af2fc4a6c --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/service/RateLimitServiceTest.java @@ -0,0 +1,229 @@ +package stirling.software.saas.service; + +import static org.assertj.core.api.Assertions.assertThat; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; + +/** + * Unit tests for {@link RateLimitService}. + * + *

The service keeps per-team attempt counts in two in-memory {@link + * java.util.concurrent.ConcurrentHashMap} buckets (hourly cap 50, daily cap 150) keyed by {@code + * "team:" + teamId}. The reset clock is {@link System#currentTimeMillis()} with 1h / 1d windows, so + * nothing created inside a test ever expires during the run - every assertion below is + * deterministic pure arithmetic with no sleeps or fake clock. A fresh service instance is built per + * test so the in-memory buckets start clean. + */ +class RateLimitServiceTest { + + private static final int HOURLY_LIMIT = 50; + private static final int DAILY_LIMIT = 150; + + private RateLimitService service; + + @BeforeEach + void setUp() { + service = new RateLimitService(); + } + + @Nested + @DisplayName("allowInvitation - hourly limit") + class AllowInvitationHourly { + + @Test + @DisplayName("first invitation for a team is allowed") + void firstInvitation_allowed() { + assertThat(service.allowInvitation(1L)).isTrue(); + } + + @Test + @DisplayName("exactly the hourly limit (50) of invitations are all allowed") + void upToHourlyLimit_allAllowed() { + for (int i = 1; i <= HOURLY_LIMIT; i++) { + assertThat(service.allowInvitation(1L)) + .as("invitation #%d should be allowed", i) + .isTrue(); + } + } + + @Test + @DisplayName("the 51st invitation within the hour is rejected") + void overHourlyLimit_rejected() { + for (int i = 1; i <= HOURLY_LIMIT; i++) { + service.allowInvitation(1L); + } + + // count would become 51 > 50 -> rejected + assertThat(service.allowInvitation(1L)).isFalse(); + } + + @Test + @DisplayName("once over the hourly cap, subsequent attempts stay rejected") + void staysRejectedOnceOverHourly() { + for (int i = 1; i <= HOURLY_LIMIT + 1; i++) { + service.allowInvitation(1L); + } + + assertThat(service.allowInvitation(1L)).isFalse(); + assertThat(service.allowInvitation(1L)).isFalse(); + } + } + + @Nested + @DisplayName("allowInvitation - per-team isolation") + class PerTeamIsolation { + + @Test + @DisplayName("different teams have independent counters") + void differentTeams_independent() { + // Exhaust team 1's hourly quota. + for (int i = 1; i <= HOURLY_LIMIT; i++) { + service.allowInvitation(1L); + } + assertThat(service.allowInvitation(1L)).isFalse(); + + // Team 2 is untouched and fully allowed. + assertThat(service.allowInvitation(2L)).isTrue(); + assertThat(service.getRemainingInvitations(2L)).isEqualTo(HOURLY_LIMIT - 1); + } + + @Test + @DisplayName("null teamId is keyed as its own bucket and behaves like any team") + void nullTeamId_hasOwnBucket() { + assertThat(service.allowInvitation(null)).isTrue(); + // key becomes "team:null"; remaining drops by one for that key. + assertThat(service.getRemainingInvitations(null)).isEqualTo(HOURLY_LIMIT - 1); + // A real team is unaffected. + assertThat(service.getRemainingInvitations(1L)).isEqualTo(HOURLY_LIMIT); + } + } + + @Nested + @DisplayName("getRemainingInvitations") + class GetRemainingInvitations { + + @Test + @DisplayName("returns the full hourly allowance when no invitation has been recorded") + void noBucket_returnsFullAllowance() { + assertThat(service.getRemainingInvitations(7L)).isEqualTo(HOURLY_LIMIT); + } + + @Test + @DisplayName("decreases by one after a single allowed invitation") + void afterOneInvitation_decrementsByOne() { + service.allowInvitation(7L); + + assertThat(service.getRemainingInvitations(7L)).isEqualTo(HOURLY_LIMIT - 1); + } + + @Test + @DisplayName("tracks the running count across several invitations") + void tracksRunningCount() { + for (int i = 0; i < 10; i++) { + service.allowInvitation(7L); + } + + assertThat(service.getRemainingInvitations(7L)).isEqualTo(HOURLY_LIMIT - 10); + } + + @Test + @DisplayName("is zero exactly when the hourly limit has been fully consumed") + void atHourlyLimit_remainingIsZero() { + for (int i = 1; i <= HOURLY_LIMIT; i++) { + service.allowInvitation(7L); + } + + assertThat(service.getRemainingInvitations(7L)).isZero(); + } + + @Test + @DisplayName("never goes negative once invitations are rejected past the cap") + void overLimit_remainingClampedAtZero() { + for (int i = 1; i <= HOURLY_LIMIT + 5; i++) { + service.allowInvitation(7L); + } + + // Math.max(0, 50 - count) clamps at zero even though attempts exceeded the cap. + assertThat(service.getRemainingInvitations(7L)).isZero(); + } + } + + @Nested + @DisplayName("allowInvitation - daily limit and hourly rollback") + class DailyLimitAndRollback { + + @Test + @DisplayName("daily cap (150) blocks attempts even though it spans multiple hourly windows") + void belowDailyLimit_acrossTeams_doesNotInterfere() { + // A single team can never reach the daily cap within one hourly window because the + // hourly cap (50) trips first. Verify the first 50 are allowed and 51st is blocked, + // confirming the hourly gate is the binding constraint here. + for (int i = 1; i <= HOURLY_LIMIT; i++) { + assertThat(service.allowInvitation(9L)).isTrue(); + } + assertThat(service.allowInvitation(9L)).isFalse(); + } + + @Test + @DisplayName("rejection at the hourly gate does not consume the remaining count further") + void hourlyRejection_doesNotAdvanceRemaining() { + for (int i = 1; i <= HOURLY_LIMIT; i++) { + service.allowInvitation(9L); + } + assertThat(service.getRemainingInvitations(9L)).isZero(); + + // Rejected attempts push the internal count past the cap, but remaining stays clamped. + service.allowInvitation(9L); + assertThat(service.getRemainingInvitations(9L)).isZero(); + } + } + + @Nested + @DisplayName("cleanupExpiredBuckets") + class CleanupExpiredBuckets { + + @Test + @DisplayName("runs cleanly when there are no buckets at all") + void emptyState_noError() { + // Nothing recorded yet; cleanup must be a harmless no-op. + service.cleanupExpiredBuckets(); + } + + @Test + @DisplayName("does not evict fresh (unexpired) buckets, preserving their counts") + void doesNotEvictFreshBuckets() { + service.allowInvitation(11L); + service.allowInvitation(11L); + assertThat(service.getRemainingInvitations(11L)).isEqualTo(HOURLY_LIMIT - 2); + + // Buckets just created have reset times an hour/day out, so none are expired. + service.cleanupExpiredBuckets(); + + // Count survived the sweep. + assertThat(service.getRemainingInvitations(11L)).isEqualTo(HOURLY_LIMIT - 2); + } + + @Test + @DisplayName("is idempotent across repeated invocations") + void repeatedInvocations_stable() { + service.allowInvitation(12L); + + service.cleanupExpiredBuckets(); + service.cleanupExpiredBuckets(); + service.cleanupExpiredBuckets(); + + assertThat(service.getRemainingInvitations(12L)).isEqualTo(HOURLY_LIMIT - 1); + } + } + + @Test + @DisplayName("daily limit constant is wider than the hourly limit (sanity on configured caps)") + void dailyWiderThanHourly() { + // This documents the relationship the service relies on: the hourly cap always trips first + // for a single uninterrupted burst, so a lone team can't reach the daily cap in one window. + assertThat(DAILY_LIMIT).isGreaterThan(HOURLY_LIMIT); + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/service/SaasTeamExtensionServiceTest.java b/app/saas/src/test/java/stirling/software/saas/service/SaasTeamExtensionServiceTest.java new file mode 100644 index 0000000000..d65e41902b --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/service/SaasTeamExtensionServiceTest.java @@ -0,0 +1,585 @@ +package stirling.software.saas.service; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.verifyNoInteractions; +import static org.mockito.Mockito.when; + +import java.util.Optional; + +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.ArgumentCaptor; +import org.mockito.InjectMocks; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; + +import stirling.software.proprietary.model.Team; +import stirling.software.saas.model.SaasTeamExtensions; +import stirling.software.saas.repository.SaasTeamExtensionsRepository; + +/** + * Unit tests for {@link SaasTeamExtensionService}. + * + *

The service is a thin read/write facade over {@link SaasTeamExtensionsRepository}. Reads + * return safe defaults when no row exists (non-personal, STANDARD type, seatsUsed=0, maxSeats=1, + * createdBy=null, hasAvailableSeats=true, canInviteMembers=true); writes create the row lazily via + * {@code getOrCreate}. Pure delegation + Optional mapping, so everything is mocked at the + * repository boundary. + */ +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class SaasTeamExtensionServiceTest { + + private static final Long TEAM_ID = 7L; + + @Mock private SaasTeamExtensionsRepository repository; + + @InjectMocks private SaasTeamExtensionService service; + + /** A team with a non-null id (the common case). */ + private static Team team(Long id) { + Team t = new Team(); + t.setId(id); + t.setName("team-" + id); + return t; + } + + private static Team team() { + return team(TEAM_ID); + } + + /** A fresh extension row for the given team carrying entity defaults. */ + private static SaasTeamExtensions ext(Team team) { + return new SaasTeamExtensions(team); + } + + @Nested + @DisplayName("getOrCreate") + class GetOrCreate { + + @Test + @DisplayName("returns the existing row without creating a new one") + void existingRow_returnedWithoutSave() { + Team team = team(); + SaasTeamExtensions existing = ext(team); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.of(existing)); + + SaasTeamExtensions result = service.getOrCreate(team); + + assertThat(result).isSameAs(existing); + verify(repository, never()).save(any()); + } + + @Test + @DisplayName("lazily creates and saves a new row when none exists") + void missingRow_savesNew() { + Team team = team(); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.empty()); + // save echoes back its argument so the returned instance is the freshly built row. + when(repository.save(any(SaasTeamExtensions.class))) + .thenAnswer(inv -> inv.getArgument(0)); + + SaasTeamExtensions result = service.getOrCreate(team); + + ArgumentCaptor captor = + ArgumentCaptor.forClass(SaasTeamExtensions.class); + verify(repository).save(captor.capture()); + SaasTeamExtensions saved = captor.getValue(); + // The new row is bound to the team and carries entity defaults. + assertThat(saved.getTeam()).isSameAs(team); + assertThat(saved.getTeamId()).isEqualTo(TEAM_ID); + assertThat(saved.getTeamType()).isEqualTo(SaasTeamExtensions.TEAM_TYPE_STANDARD); + assertThat(saved.isPersonal()).isFalse(); + assertThat(saved.getSeatsUsed()).isZero(); + assertThat(saved.getMaxSeats()).isEqualTo(1); + assertThat(result).isSameAs(saved); + } + } + + @Nested + @DisplayName("isPersonal") + class IsPersonal { + + @Test + @DisplayName("null team short-circuits to false without touching the repository") + void nullTeam_false() { + assertThat(service.isPersonal(null)).isFalse(); + verifyNoInteractions(repository); + } + + @Test + @DisplayName("team with null id short-circuits to false without touching the repository") + void nullId_false() { + assertThat(service.isPersonal(team(null))).isFalse(); + verifyNoInteractions(repository); + } + + @Test + @DisplayName("no row defaults to false") + void noRow_defaultsFalse() { + Team team = team(); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.empty()); + + assertThat(service.isPersonal(team)).isFalse(); + } + + @Test + @DisplayName("reflects the persisted personal flag when a row exists") + void existingPersonalRow_true() { + Team team = team(); + SaasTeamExtensions row = ext(team); + row.setIsPersonal(true); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.of(row)); + + assertThat(service.isPersonal(team)).isTrue(); + } + } + + @Nested + @DisplayName("getTeamType") + class GetTeamType { + + @Test + @DisplayName("null team defaults to STANDARD without a lookup") + void nullTeam_standard() { + assertThat(service.getTeamType(null)).isEqualTo(SaasTeamExtensions.TEAM_TYPE_STANDARD); + verifyNoInteractions(repository); + } + + @Test + @DisplayName("team with null id defaults to STANDARD without a lookup") + void nullId_standard() { + assertThat(service.getTeamType(team(null))) + .isEqualTo(SaasTeamExtensions.TEAM_TYPE_STANDARD); + verifyNoInteractions(repository); + } + + @Test + @DisplayName("no row defaults to STANDARD") + void noRow_standard() { + Team team = team(); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.empty()); + + assertThat(service.getTeamType(team)).isEqualTo(SaasTeamExtensions.TEAM_TYPE_STANDARD); + } + + @Test + @DisplayName("returns the persisted PERSONAL type when a row exists") + void existingRow_personalType() { + Team team = team(); + SaasTeamExtensions row = ext(team); + row.setTeamType(SaasTeamExtensions.TEAM_TYPE_PERSONAL); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.of(row)); + + assertThat(service.getTeamType(team)).isEqualTo(SaasTeamExtensions.TEAM_TYPE_PERSONAL); + } + } + + @Nested + @DisplayName("getSeatsUsed") + class GetSeatsUsed { + + @Test + @DisplayName("no row defaults to 0") + void noRow_zero() { + Team team = team(); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.empty()); + + assertThat(service.getSeatsUsed(team)).isZero(); + } + + @Test + @DisplayName("returns the persisted seatsUsed value") + void existingRow_value() { + Team team = team(); + SaasTeamExtensions row = ext(team); + row.setSeatsUsed(5); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.of(row)); + + assertThat(service.getSeatsUsed(team)).isEqualTo(5); + } + } + + @Nested + @DisplayName("getMaxSeats") + class GetMaxSeats { + + @Test + @DisplayName("no row defaults to 1 (not 0)") + void noRow_one() { + Team team = team(); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.empty()); + + assertThat(service.getMaxSeats(team)).isEqualTo(1); + } + + @Test + @DisplayName("returns the persisted maxSeats value") + void existingRow_value() { + Team team = team(); + SaasTeamExtensions row = ext(team); + row.setMaxSeats(25); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.of(row)); + + assertThat(service.getMaxSeats(team)).isEqualTo(25); + } + } + + @Nested + @DisplayName("getCreatedByUserId") + class GetCreatedByUserId { + + @Test + @DisplayName("no row defaults to null") + void noRow_null() { + Team team = team(); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.empty()); + + assertThat(service.getCreatedByUserId(team)).isNull(); + } + + @Test + @DisplayName("row present but creator unset is null") + void existingRow_unsetCreator_null() { + Team team = team(); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.of(ext(team))); + + assertThat(service.getCreatedByUserId(team)).isNull(); + } + + @Test + @DisplayName("returns the persisted creator id") + void existingRow_value() { + Team team = team(); + SaasTeamExtensions row = ext(team); + row.setCreatedByUserId(99L); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.of(row)); + + assertThat(service.getCreatedByUserId(team)).isEqualTo(99L); + } + } + + @Nested + @DisplayName("hasAvailableSeats") + class HasAvailableSeats { + + @Test + @DisplayName("no row defaults to true (optimistic)") + void noRow_true() { + Team team = team(); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.empty()); + + assertThat(service.hasAvailableSeats(team)).isTrue(); + } + + @Test + @DisplayName("standard team always has seats regardless of usage") + void standardTeam_alwaysTrue() { + Team team = team(); + SaasTeamExtensions row = ext(team); + row.setIsPersonal(false); + row.setSeatsUsed(1000); + row.setMaxSeats(1); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.of(row)); + + assertThat(service.hasAvailableSeats(team)).isTrue(); + } + + @Test + @DisplayName("personal team with a free seat returns true") + void personalTeam_underCap_true() { + Team team = team(); + SaasTeamExtensions row = ext(team); + row.setIsPersonal(true); + row.setSeatsUsed(0); + row.setMaxSeats(1); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.of(row)); + + assertThat(service.hasAvailableSeats(team)).isTrue(); + } + + @Test + @DisplayName("personal team at its seat cap returns false (boundary)") + void personalTeam_atCap_false() { + Team team = team(); + SaasTeamExtensions row = ext(team); + row.setIsPersonal(true); + row.setSeatsUsed(1); + row.setMaxSeats(1); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.of(row)); + + assertThat(service.hasAvailableSeats(team)).isFalse(); + } + } + + @Nested + @DisplayName("canInviteMembers") + class CanInviteMembers { + + @Test + @DisplayName("no row defaults to true") + void noRow_true() { + Team team = team(); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.empty()); + + assertThat(service.canInviteMembers(team)).isTrue(); + } + + @Test + @DisplayName("standard team can invite") + void standardTeam_true() { + Team team = team(); + SaasTeamExtensions row = ext(team); + row.setIsPersonal(false); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.of(row)); + + assertThat(service.canInviteMembers(team)).isTrue(); + } + + @Test + @DisplayName("personal team can never invite") + void personalTeam_false() { + Team team = team(); + SaasTeamExtensions row = ext(team); + row.setIsPersonal(true); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.of(row)); + + assertThat(service.canInviteMembers(team)).isFalse(); + } + } + + @Nested + @DisplayName("incrementSeatsUsed") + class IncrementSeatsUsed { + + @Test + @DisplayName( + "ensures the row exists then delegates the atomic increment, returning its result") + void existingRow_delegatesIncrement() { + Team team = team(); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.of(ext(team))); + when(repository.incrementSeatsUsed(TEAM_ID)).thenReturn(1); + + int result = service.incrementSeatsUsed(team); + + assertThat(result).isEqualTo(1); + verify(repository).incrementSeatsUsed(TEAM_ID); + // Row already existed -> no lazy creation. + verify(repository, never()).save(any()); + } + + @Test + @DisplayName("creates the row first when missing, then increments") + void missingRow_createsThenIncrements() { + Team team = team(); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.empty()); + when(repository.save(any(SaasTeamExtensions.class))) + .thenAnswer(inv -> inv.getArgument(0)); + when(repository.incrementSeatsUsed(TEAM_ID)).thenReturn(1); + + int result = service.incrementSeatsUsed(team); + + assertThat(result).isEqualTo(1); + verify(repository).save(any(SaasTeamExtensions.class)); + verify(repository).incrementSeatsUsed(TEAM_ID); + } + + @Test + @DisplayName("returns 0 when the atomic update hits the personal-team cap") + void capHit_returnsZero() { + Team team = team(); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.of(ext(team))); + when(repository.incrementSeatsUsed(TEAM_ID)).thenReturn(0); + + assertThat(service.incrementSeatsUsed(team)).isZero(); + } + } + + @Nested + @DisplayName("decrementSeatsUsed") + class DecrementSeatsUsed { + + @Test + @DisplayName("delegates straight to the atomic decrement without creating a row") + void delegatesDecrement() { + Team team = team(); + when(repository.decrementSeatsUsed(TEAM_ID)).thenReturn(1); + + int result = service.decrementSeatsUsed(team); + + assertThat(result).isEqualTo(1); + verify(repository).decrementSeatsUsed(TEAM_ID); + verify(repository, never()).findByTeamId(any()); + verify(repository, never()).save(any()); + } + + @Test + @DisplayName("returns 0 when already floored at zero") + void alreadyZero_returnsZero() { + Team team = team(); + when(repository.decrementSeatsUsed(TEAM_ID)).thenReturn(0); + + assertThat(service.decrementSeatsUsed(team)).isZero(); + } + } + + @Nested + @DisplayName("setPersonal") + class SetPersonal { + + @Test + @DisplayName("marking personal sets the flag and the PERSONAL team type, then saves") + void markPersonal_setsFlagAndType() { + Team team = team(); + SaasTeamExtensions row = ext(team); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.of(row)); + + service.setPersonal(team, true); + + assertThat(row.isPersonal()).isTrue(); + assertThat(row.getTeamType()).isEqualTo(SaasTeamExtensions.TEAM_TYPE_PERSONAL); + verify(repository).save(row); + } + + @Test + @DisplayName( + "marking non-personal clears the flag and reverts to STANDARD type, then saves") + void markNonPersonal_revertsToStandard() { + Team team = team(); + SaasTeamExtensions row = ext(team); + row.setIsPersonal(true); + row.setTeamType(SaasTeamExtensions.TEAM_TYPE_PERSONAL); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.of(row)); + + service.setPersonal(team, false); + + assertThat(row.isPersonal()).isFalse(); + assertThat(row.getTeamType()).isEqualTo(SaasTeamExtensions.TEAM_TYPE_STANDARD); + verify(repository).save(row); + } + + @Test + @DisplayName("creates the row first when missing, then applies the personal flag") + void missingRow_createsThenSets() { + Team team = team(); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.empty()); + when(repository.save(any(SaasTeamExtensions.class))) + .thenAnswer(inv -> inv.getArgument(0)); + + service.setPersonal(team, true); + + // getOrCreate saves the new row, then setPersonal saves the mutated row. + ArgumentCaptor captor = + ArgumentCaptor.forClass(SaasTeamExtensions.class); + verify(repository, org.mockito.Mockito.times(2)).save(captor.capture()); + SaasTeamExtensions last = captor.getValue(); + assertThat(last.isPersonal()).isTrue(); + assertThat(last.getTeamType()).isEqualTo(SaasTeamExtensions.TEAM_TYPE_PERSONAL); + } + } + + @Nested + @DisplayName("setSeats") + class SetSeats { + + @Test + @DisplayName("writes both seatCount and maxSeats onto the existing row, then saves") + void existingRow_writesBoth() { + Team team = team(); + SaasTeamExtensions row = ext(team); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.of(row)); + + service.setSeats(team, 3, 10); + + assertThat(row.getSeatCount()).isEqualTo(3); + assertThat(row.getMaxSeats()).isEqualTo(10); + verify(repository).save(row); + } + + @Test + @DisplayName("does not touch seatsUsed (only seatCount and maxSeats)") + void doesNotChangeSeatsUsed() { + Team team = team(); + SaasTeamExtensions row = ext(team); + row.setSeatsUsed(4); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.of(row)); + + service.setSeats(team, 8, 8); + + assertThat(row.getSeatsUsed()).isEqualTo(4); + } + + @Test + @DisplayName("creates the row first when missing, then writes the seat fields") + void missingRow_createsThenWrites() { + Team team = team(); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.empty()); + when(repository.save(any(SaasTeamExtensions.class))) + .thenAnswer(inv -> inv.getArgument(0)); + + service.setSeats(team, 2, 5); + + ArgumentCaptor captor = + ArgumentCaptor.forClass(SaasTeamExtensions.class); + verify(repository, org.mockito.Mockito.times(2)).save(captor.capture()); + SaasTeamExtensions last = captor.getValue(); + assertThat(last.getSeatCount()).isEqualTo(2); + assertThat(last.getMaxSeats()).isEqualTo(5); + } + } + + @Nested + @DisplayName("setCreatedByUserId") + class SetCreatedByUserId { + + @Test + @DisplayName("writes the creator id onto the existing row, then saves") + void existingRow_writesCreator() { + Team team = team(); + SaasTeamExtensions row = ext(team); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.of(row)); + + service.setCreatedByUserId(team, 42L); + + assertThat(row.getCreatedByUserId()).isEqualTo(42L); + verify(repository).save(row); + } + + @Test + @DisplayName("accepts a null creator id (clearing)") + void nullCreator_cleared() { + Team team = team(); + SaasTeamExtensions row = ext(team); + row.setCreatedByUserId(5L); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.of(row)); + + service.setCreatedByUserId(team, null); + + assertThat(row.getCreatedByUserId()).isNull(); + verify(repository).save(row); + } + + @Test + @DisplayName("creates the row first when missing, then writes the creator id") + void missingRow_createsThenWrites() { + Team team = team(); + when(repository.findByTeamId(TEAM_ID)).thenReturn(Optional.empty()); + when(repository.save(any(SaasTeamExtensions.class))) + .thenAnswer(inv -> inv.getArgument(0)); + + service.setCreatedByUserId(team, 13L); + + ArgumentCaptor captor = + ArgumentCaptor.forClass(SaasTeamExtensions.class); + verify(repository, org.mockito.Mockito.times(2)).save(captor.capture()); + assertThat(captor.getValue().getCreatedByUserId()).isEqualTo(13L); + } + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/service/SaasUserAccountServiceTest.java b/app/saas/src/test/java/stirling/software/saas/service/SaasUserAccountServiceTest.java new file mode 100644 index 0000000000..5bf0a33e4c --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/service/SaasUserAccountServiceTest.java @@ -0,0 +1,479 @@ +package stirling.software.saas.service; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyBoolean; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.verifyNoInteractions; +import static org.mockito.Mockito.when; + +import java.util.Optional; +import java.util.UUID; + +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.InjectMocks; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; + +import stirling.software.common.model.enumeration.Role; +import stirling.software.proprietary.model.Team; +import stirling.software.proprietary.security.database.repository.UserRepository; +import stirling.software.proprietary.security.model.AuthenticationType; +import stirling.software.proprietary.security.model.Authority; +import stirling.software.proprietary.security.model.User; +import stirling.software.proprietary.security.service.UserService; +import stirling.software.saas.model.SupabaseUser; + +/** + * Unit tests for {@link SaasUserAccountService}. + * + *

The service is a thin orchestration layer over {@link UserService} and the saas extension + * services; every collaborator is mocked and the methods are invoked directly. Role state is driven + * through {@link User#getRolesAsString()} by attaching an {@link Authority} (its constructor + * registers itself on the user). Anonymous state is driven through {@link + * User#setAuthenticationType(AuthenticationType)}, which the entity stores lowercased so the + * service's case-insensitive comparison against {@code ANONYMOUS} matches. + */ +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class SaasUserAccountServiceTest { + + @Mock private UserService userService; + @Mock private UserRepository userRepository; + @Mock private UserRoleService userRoleService; + @Mock private SupabaseUserService supabaseUserService; + @Mock private SaasUserExtensionService saasUserExtensionService; + @Mock private SaasTeamExtensionService saasTeamExtensionService; + @Mock private SaasTeamService saasTeamService; + + @InjectMocks private SaasUserAccountService service; + + private static final String SUPABASE_ID = "11111111-2222-3333-4444-555555555555"; + private static final UUID SUPABASE_UUID = UUID.fromString(SUPABASE_ID); + + /** Build a user whose role string equals the given role id (e.g. "ROLE_USER"). */ + private static User userWithRole(String roleId) { + User u = new User(); + u.setUsername("alice@example.com"); + // Authority's constructor registers itself onto the user's authority set. + new Authority(roleId, u); + return u; + } + + private static Team team(Long id, String name) { + Team t = new Team(); + t.setId(id); + t.setName(name); + return t; + } + + @Nested + @DisplayName("getUserBySupabaseId") + class GetUserBySupabaseId { + + @Test + @DisplayName("returns the local user when the UUID parses and a row exists") + void returnsUser() { + User u = userWithRole(Role.USER.getRoleId()); + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.of(u)); + + assertThat(service.getUserBySupabaseId(SUPABASE_ID)).isSameAs(u); + } + + @Test + @DisplayName("throws with an 'invalid format' message when the id is not a UUID") + void invalidFormat_throws() { + assertThatThrownBy(() -> service.getUserBySupabaseId("not-a-uuid")) + .isInstanceOf(IllegalArgumentException.class) + .hasMessageContaining("Invalid Supabase ID format") + .hasMessageContaining("not-a-uuid"); + + verifyNoInteractions(userService); + } + + @Test + @DisplayName("throws 'user not found' when the UUID parses but no local row exists") + void notFound_throws() { + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.empty()); + + assertThatThrownBy(() -> service.getUserBySupabaseId(SUPABASE_ID)) + .isInstanceOf(IllegalArgumentException.class) + .hasMessageContaining("User not found for Supabase ID") + .hasMessageContaining(SUPABASE_ID); + } + } + + @Nested + @DisplayName("handleUpgrade") + class HandleUpgrade { + + @Test + @DisplayName("promotes a free (ROLE_USER) user to PRO and returns true") + void freeUser_isUpgraded() { + User u = userWithRole(Role.USER.getRoleId()); + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.of(u)); + + boolean upgraded = service.handleUpgrade(SUPABASE_ID); + + assertThat(upgraded).isTrue(); + verify(userRoleService).upgradeToPro(u); + } + + @Test + @DisplayName("returns false and does not re-upgrade a user already on PRO") + void proUser_isNoOp() { + User u = userWithRole(Role.PRO_USER.getRoleId()); + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.of(u)); + + boolean upgraded = service.handleUpgrade(SUPABASE_ID); + + assertThat(upgraded).isFalse(); + verify(userRoleService, never()).upgradeToPro(any()); + } + + @Test + @DisplayName("returns false for any non-free role (e.g. admin/other) without upgrading") + void otherRole_isNoOp() { + User u = userWithRole("ROLE_ADMIN"); + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.of(u)); + + assertThat(service.handleUpgrade(SUPABASE_ID)).isFalse(); + verify(userRoleService, never()).upgradeToPro(any()); + } + + @Test + @DisplayName("propagates the lookup failure when the supabase id has no local user") + void unknownUser_propagates() { + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.empty()); + + assertThatThrownBy(() -> service.handleUpgrade(SUPABASE_ID)) + .isInstanceOf(IllegalArgumentException.class); + verifyNoInteractions(userRoleService); + } + } + + @Nested + @DisplayName("handleDowngrade") + class HandleDowngrade { + + @Test + @DisplayName("downgrades a PRO user with no team to FREE and returns true") + void proWithoutTeam_isDowngraded() { + User u = userWithRole(Role.PRO_USER.getRoleId()); + // no team set -> getTeam() is null + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.of(u)); + + boolean downgraded = service.handleDowngrade(SUPABASE_ID); + + assertThat(downgraded).isTrue(); + verify(userRoleService).downgradeToFree(u); + } + + @Test + @DisplayName( + "downgrades a PRO user whose team is personal (personal team is not a shared PRO source)") + void proWithPersonalTeam_isDowngraded() { + User u = userWithRole(Role.PRO_USER.getRoleId()); + Team t = team(7L, "alice-personal"); + u.setTeam(t); + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.of(u)); + when(saasTeamExtensionService.isPersonal(t)).thenReturn(true); + + boolean downgraded = service.handleDowngrade(SUPABASE_ID); + + assertThat(downgraded).isTrue(); + verify(userRoleService).downgradeToFree(u); + } + + @Test + @DisplayName("keeps PRO (returns false) for a PRO user on a non-personal/shared team") + void proWithSharedTeam_keepsPro() { + User u = userWithRole(Role.PRO_USER.getRoleId()); + Team t = team(9L, "acme-team"); + u.setTeam(t); + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.of(u)); + when(saasTeamExtensionService.isPersonal(t)).thenReturn(false); + + boolean downgraded = service.handleDowngrade(SUPABASE_ID); + + assertThat(downgraded).isFalse(); + verify(userRoleService, never()).downgradeToFree(any()); + } + + @Test + @DisplayName("returns false without touching roles when the user is already FREE") + void freeUser_isNoOp() { + User u = userWithRole(Role.USER.getRoleId()); + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.of(u)); + + assertThat(service.handleDowngrade(SUPABASE_ID)).isFalse(); + verify(userRoleService, never()).downgradeToFree(any()); + verifyNoInteractions(saasTeamExtensionService); + } + } + + @Nested + @DisplayName("enableMeteredBilling") + class EnableMeteredBilling { + + @Test + @DisplayName("enables metered billing and returns true when it was off") + void wasOff_enables() { + User u = userWithRole(Role.PRO_USER.getRoleId()); + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.of(u)); + when(saasUserExtensionService.isMeteredBillingEnabled(u)).thenReturn(false); + + boolean result = service.enableMeteredBilling(SUPABASE_ID); + + assertThat(result).isTrue(); + verify(saasUserExtensionService).setMeteredBillingEnabled(u, true); + } + + @Test + @DisplayName("returns false and does not re-enable when metered billing is already on") + void alreadyOn_isNoOp() { + User u = userWithRole(Role.PRO_USER.getRoleId()); + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.of(u)); + when(saasUserExtensionService.isMeteredBillingEnabled(u)).thenReturn(true); + + boolean result = service.enableMeteredBilling(SUPABASE_ID); + + assertThat(result).isFalse(); + verify(saasUserExtensionService, never()).setMeteredBillingEnabled(any(), anyBoolean()); + } + } + + @Nested + @DisplayName("disableMeteredBilling") + class DisableMeteredBilling { + + @Test + @DisplayName("disables metered billing and returns true when it was on") + void wasOn_disables() { + User u = userWithRole(Role.USER.getRoleId()); + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.of(u)); + when(saasUserExtensionService.isMeteredBillingEnabled(u)).thenReturn(true); + + boolean result = service.disableMeteredBilling(SUPABASE_ID); + + assertThat(result).isTrue(); + verify(saasUserExtensionService).setMeteredBillingEnabled(u, false); + } + + @Test + @DisplayName("returns false and does nothing when metered billing is already off") + void alreadyOff_isNoOp() { + User u = userWithRole(Role.USER.getRoleId()); + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.of(u)); + when(saasUserExtensionService.isMeteredBillingEnabled(u)).thenReturn(false); + + boolean result = service.disableMeteredBilling(SUPABASE_ID); + + assertThat(result).isFalse(); + verify(saasUserExtensionService, never()).setMeteredBillingEnabled(any(), anyBoolean()); + } + } + + @Nested + @DisplayName("synchronizeUserUpgrade") + class SynchronizeUserUpgrade { + + private static SupabaseUser supabaseUser(boolean anonymous) { + SupabaseUser su = new SupabaseUser(); + su.setId(SUPABASE_UUID); + su.setAnonymous(anonymous); + return su; + } + + private static User anonymousLocalUser() { + User u = new User(); + u.setUsername("anon-handle"); + u.setAuthenticationType(AuthenticationType.ANONYMOUS); + return u; + } + + @Test + @DisplayName("throws IllegalStateException when no local user is linked to the supabase id") + void noLinkedUser_throws() { + SupabaseUser su = supabaseUser(true); + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.empty()); + + assertThatThrownBy( + () -> service.synchronizeUserUpgrade(su, "alice@example.com", "google")) + .isInstanceOf(IllegalStateException.class) + .hasMessageContaining("No local user linked to Supabase ID"); + + verifyNoInteractions(supabaseUserService); + } + + @Test + @DisplayName("flips the supabase anonymous mirror to false and saves it") + void anonymousMirror_isFlippedAndSaved() { + SupabaseUser su = supabaseUser(true); + User u = anonymousLocalUser(); + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.of(u)); + when(userService.saveUser(any(User.class))).thenAnswer(inv -> inv.getArgument(0)); + + service.synchronizeUserUpgrade(su, "alice@example.com", "web"); + + assertThat(su.isAnonymous()).isFalse(); + verify(supabaseUserService).save(su); + } + + @Test + @DisplayName("does not save the supabase mirror when it was already non-anonymous") + void nonAnonymousMirror_isNotSaved() { + SupabaseUser su = supabaseUser(false); + User u = anonymousLocalUser(); + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.of(u)); + when(userService.saveUser(any(User.class))).thenAnswer(inv -> inv.getArgument(0)); + + service.synchronizeUserUpgrade(su, "alice@example.com", "web"); + + verify(supabaseUserService, never()).save(any()); + } + + @Test + @DisplayName("promotes an anonymous local user to WEB and copies the email into username") + void anonymousLocalUser_promotedToWeb() { + SupabaseUser su = supabaseUser(false); + User u = anonymousLocalUser(); + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.of(u)); + when(userService.saveUser(any(User.class))).thenAnswer(inv -> inv.getArgument(0)); + + User result = service.synchronizeUserUpgrade(su, "alice@example.com", "web"); + + assertThat(result).isSameAs(u); + // setAuthenticationType stores the enum name lowercased. + assertThat(u.getAuthenticationType()).isEqualTo("web"); + assertThat(u.getEmail()).isEqualTo("alice@example.com"); + assertThat(u.getUsername()).isEqualTo("alice@example.com"); + verify(userService).saveUser(u); + // Upgrading from anon gives the user their own team. + verify(saasTeamService).ensurePersonalTeam(u); + } + + @Test + @DisplayName("maps a known OAuth provider (google) to OAUTH2 auth type") + void oauthProvider_mapsToOauth2() { + SupabaseUser su = supabaseUser(false); + User u = anonymousLocalUser(); + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.of(u)); + when(userService.saveUser(any(User.class))).thenAnswer(inv -> inv.getArgument(0)); + + service.synchronizeUserUpgrade(su, "alice@example.com", "google"); + + assertThat(u.getAuthenticationType()).isEqualTo("oauth2"); + } + + @Test + @DisplayName("maps a generic 'oauth' authMethod to OAUTH2") + void genericOauth_mapsToOauth2() { + SupabaseUser su = supabaseUser(false); + User u = anonymousLocalUser(); + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.of(u)); + when(userService.saveUser(any(User.class))).thenAnswer(inv -> inv.getArgument(0)); + + service.synchronizeUserUpgrade(su, "alice@example.com", "oauth"); + + assertThat(u.getAuthenticationType()).isEqualTo("oauth2"); + } + + @Test + @DisplayName("maps an unknown authMethod to WEB") + void unknownMethod_mapsToWeb() { + SupabaseUser su = supabaseUser(false); + User u = anonymousLocalUser(); + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.of(u)); + when(userService.saveUser(any(User.class))).thenAnswer(inv -> inv.getArgument(0)); + + service.synchronizeUserUpgrade(su, "alice@example.com", "carrier-pigeon"); + + assertThat(u.getAuthenticationType()).isEqualTo("web"); + } + + @Test + @DisplayName("maps a null authMethod to WEB") + void nullMethod_mapsToWeb() { + SupabaseUser su = supabaseUser(false); + User u = anonymousLocalUser(); + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.of(u)); + when(userService.saveUser(any(User.class))).thenAnswer(inv -> inv.getArgument(0)); + + service.synchronizeUserUpgrade(su, "alice@example.com", null); + + assertThat(u.getAuthenticationType()).isEqualTo("web"); + } + + @Test + @DisplayName( + "does not overwrite username/email when the email is blank, but still promotes the type") + void blankEmail_keepsUsername() { + SupabaseUser su = supabaseUser(false); + User u = anonymousLocalUser(); + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.of(u)); + when(userService.saveUser(any(User.class))).thenAnswer(inv -> inv.getArgument(0)); + + service.synchronizeUserUpgrade(su, " ", "google"); + + assertThat(u.getAuthenticationType()).isEqualTo("oauth2"); + assertThat(u.getUsername()).isEqualTo("anon-handle"); + assertThat(u.getEmail()).isNull(); + verify(userService).saveUser(u); + } + + @Test + @DisplayName("does not overwrite username/email when the email is null") + void nullEmail_keepsUsername() { + SupabaseUser su = supabaseUser(false); + User u = anonymousLocalUser(); + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.of(u)); + when(userService.saveUser(any(User.class))).thenAnswer(inv -> inv.getArgument(0)); + + service.synchronizeUserUpgrade(su, null, "web"); + + assertThat(u.getUsername()).isEqualTo("anon-handle"); + assertThat(u.getEmail()).isNull(); + } + + @Test + @DisplayName("leaves a non-anonymous local user untouched and never saves it") + void nonAnonymousLocalUser_untouched() { + SupabaseUser su = supabaseUser(false); + User u = new User(); + u.setUsername("existing@example.com"); + u.setAuthenticationType(AuthenticationType.WEB); + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.of(u)); + + User result = service.synchronizeUserUpgrade(su, "new@example.com", "google"); + + assertThat(result).isSameAs(u); + assertThat(u.getUsername()).isEqualTo("existing@example.com"); + assertThat(u.getAuthenticationType()).isEqualTo("web"); + verify(userService, never()).saveUser(any()); + verify(saasTeamService, never()).ensurePersonalTeam(any()); + } + + @Test + @DisplayName("anonymous comparison is case-insensitive (stored type is lowercased)") + void anonymousTypeMatchesCaseInsensitively() { + SupabaseUser su = supabaseUser(false); + User u = anonymousLocalUser(); + // sanity: the entity stored the lowercase form, exercising the equalsIgnoreCase branch + assertThat(u.getAuthenticationType()).isEqualTo("anonymous"); + when(userService.findBySupabaseId(SUPABASE_UUID)).thenReturn(Optional.of(u)); + when(userService.saveUser(any(User.class))).thenAnswer(inv -> inv.getArgument(0)); + + service.synchronizeUserUpgrade(su, "alice@example.com", "web"); + + verify(userService).saveUser(u); + } + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/service/StripeAfterCommitOrderingTest.java b/app/saas/src/test/java/stirling/software/saas/service/StripeAfterCommitOrderingTest.java index 6454dfe3b8..6b86f107de 100644 --- a/app/saas/src/test/java/stirling/software/saas/service/StripeAfterCommitOrderingTest.java +++ b/app/saas/src/test/java/stirling/software/saas/service/StripeAfterCommitOrderingTest.java @@ -11,7 +11,7 @@ import org.springframework.transaction.support.TransactionSynchronization; import org.springframework.transaction.support.TransactionSynchronizationManager; /** - * Pins the contract {@code CreditService.scheduleStripeReportAfterCommit} relies on: a {@link + * Pins the contract {@code JobChargeService.close} relies on for its Stripe meter post: a {@link * TransactionSynchronization#afterCommit()} hook fires after a successful commit and never on * rollback. */ diff --git a/app/saas/src/test/java/stirling/software/saas/service/SupabaseUserServiceTest.java b/app/saas/src/test/java/stirling/software/saas/service/SupabaseUserServiceTest.java new file mode 100644 index 0000000000..cb3d3635cb --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/service/SupabaseUserServiceTest.java @@ -0,0 +1,205 @@ +package stirling.software.saas.service; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.util.Optional; +import java.util.UUID; + +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.ArgumentCaptor; +import org.mockito.InjectMocks; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; + +import stirling.software.saas.model.SupabaseUser; +import stirling.software.saas.model.exception.UserNotFoundException; +import stirling.software.saas.repository.SupabaseUserRepository; + +/** + * Unit tests for {@link SupabaseUserService}. + * + *

Thin CRUD facade over {@link SupabaseUserRepository}. Every method either delegates to the + * repository or, in {@code getUser}, translates a {@link Optional#empty()} into a {@link + * UserNotFoundException}. The Supabase boundary is fully mocked - no DB, no network. + */ +@ExtendWith(MockitoExtension.class) +class SupabaseUserServiceTest { + + @Mock private SupabaseUserRepository supabaseUserRepository; + + @InjectMocks private SupabaseUserService service; + + private static final UUID SUPABASE_ID = UUID.fromString("11111111-2222-3333-4444-555555555555"); + private static final String EMAIL = "user@example.com"; + + private static SupabaseUser supabaseUser(UUID id, String email, boolean anonymous) { + SupabaseUser u = new SupabaseUser(); + u.setId(id); + u.setEmail(email); + u.setAnonymous(anonymous); + return u; + } + + @Nested + @DisplayName("getUser") + class GetUser { + + @Test + @DisplayName("returns the entity when the repository finds it by id") + void found_returnsEntity() { + SupabaseUser existing = supabaseUser(SUPABASE_ID, EMAIL, false); + when(supabaseUserRepository.findById(SUPABASE_ID)).thenReturn(Optional.of(existing)); + + SupabaseUser result = service.getUser(SUPABASE_ID); + + assertThat(result).isSameAs(existing); + verify(supabaseUserRepository).findById(SUPABASE_ID); + } + + @Test + @DisplayName("throws UserNotFoundException carrying the id when the repository is empty") + void missing_throwsUserNotFound() { + when(supabaseUserRepository.findById(SUPABASE_ID)).thenReturn(Optional.empty()); + + assertThatThrownBy(() -> service.getUser(SUPABASE_ID)) + .isInstanceOf(UserNotFoundException.class) + .hasMessageContaining(SUPABASE_ID.toString()) + .hasMessageContaining("not found"); + + verify(supabaseUserRepository, never()).save(any()); + } + + @Test + @DisplayName("looks the user up by the exact id it is given") + void passesIdThrough() { + UUID other = UUID.fromString("99999999-8888-7777-6666-555555555555"); + when(supabaseUserRepository.findById(other)) + .thenReturn(Optional.of(supabaseUser(other, "other@example.com", true))); + + SupabaseUser result = service.getUser(other); + + assertThat(result.getId()).isEqualTo(other); + verify(supabaseUserRepository).findById(other); + } + } + + @Nested + @DisplayName("createSupabaseUser") + class CreateSupabaseUser { + + @Test + @DisplayName("builds a SupabaseUser from the args and returns the saved entity") + void buildsAndSavesUser() { + // Repository echoes whatever it was handed. + when(supabaseUserRepository.save(any(SupabaseUser.class))) + .thenAnswer(invocation -> invocation.getArgument(0)); + + SupabaseUser result = service.createSupabaseUser(SUPABASE_ID, EMAIL, false); + + ArgumentCaptor captor = ArgumentCaptor.forClass(SupabaseUser.class); + verify(supabaseUserRepository).save(captor.capture()); + SupabaseUser saved = captor.getValue(); + assertThat(saved.getId()).isEqualTo(SUPABASE_ID); + assertThat(saved.getEmail()).isEqualTo(EMAIL); + assertThat(saved.isAnonymous()).isFalse(); + // The method returns exactly what the repository produced. + assertThat(result).isSameAs(saved); + } + + @Test + @DisplayName("propagates the anonymous flag when true") + void anonymousFlagTrue() { + when(supabaseUserRepository.save(any(SupabaseUser.class))) + .thenAnswer(invocation -> invocation.getArgument(0)); + + service.createSupabaseUser(SUPABASE_ID, EMAIL, true); + + ArgumentCaptor captor = ArgumentCaptor.forClass(SupabaseUser.class); + verify(supabaseUserRepository).save(captor.capture()); + assertThat(captor.getValue().isAnonymous()).isTrue(); + } + + @Test + @DisplayName("accepts a null email and persists it unchanged (no normalisation)") + void nullEmail_persistedAsNull() { + when(supabaseUserRepository.save(any(SupabaseUser.class))) + .thenAnswer(invocation -> invocation.getArgument(0)); + + service.createSupabaseUser(SUPABASE_ID, null, false); + + ArgumentCaptor captor = ArgumentCaptor.forClass(SupabaseUser.class); + verify(supabaseUserRepository).save(captor.capture()); + assertThat(captor.getValue().getEmail()).isNull(); + } + + @Test + @DisplayName("returns the repository's instance, not a freshly built one") + void returnsRepositoryInstance() { + SupabaseUser persisted = supabaseUser(SUPABASE_ID, EMAIL, false); + when(supabaseUserRepository.save(any(SupabaseUser.class))).thenReturn(persisted); + + SupabaseUser result = service.createSupabaseUser(SUPABASE_ID, EMAIL, false); + + assertThat(result).isSameAs(persisted); + } + + @Test + @DisplayName("does not read the user back before creating it") + void doesNotFindBeforeSave() { + when(supabaseUserRepository.save(any(SupabaseUser.class))) + .thenAnswer(invocation -> invocation.getArgument(0)); + + service.createSupabaseUser(SUPABASE_ID, EMAIL, false); + + verify(supabaseUserRepository, never()).findById(any()); + } + } + + @Nested + @DisplayName("save") + class Save { + + @Test + @DisplayName("delegates straight to the repository and returns its result") + void delegatesToRepository() { + SupabaseUser input = supabaseUser(SUPABASE_ID, EMAIL, false); + SupabaseUser persisted = supabaseUser(SUPABASE_ID, EMAIL, false); + when(supabaseUserRepository.save(input)).thenReturn(persisted); + + SupabaseUser result = service.save(input); + + assertThat(result).isSameAs(persisted); + verify(supabaseUserRepository).save(input); + } + + @Test + @DisplayName("passes a null entity through to the repository without guarding") + void nullEntity_passedThrough() { + when(supabaseUserRepository.save(null)).thenReturn(null); + + SupabaseUser result = service.save(null); + + assertThat(result).isNull(); + verify(supabaseUserRepository).save(null); + } + } + + @Test + @DisplayName("repository save failures bubble out of createSupabaseUser unchanged") + void createSupabaseUser_repositoryThrows_propagates() { + when(supabaseUserRepository.save(any(SupabaseUser.class))) + .thenThrow(new RuntimeException("constraint violation")); + + assertThatThrownBy(() -> service.createSupabaseUser(SUPABASE_ID, EMAIL, false)) + .isInstanceOf(RuntimeException.class) + .hasMessage("constraint violation"); + } +} diff --git a/app/saas/src/test/java/stirling/software/saas/service/UserRoleServiceTest.java b/app/saas/src/test/java/stirling/software/saas/service/UserRoleServiceTest.java new file mode 100644 index 0000000000..f73bf0daea --- /dev/null +++ b/app/saas/src/test/java/stirling/software/saas/service/UserRoleServiceTest.java @@ -0,0 +1,194 @@ +package stirling.software.saas.service; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.InOrder; +import org.mockito.Mock; +import org.mockito.Mockito; +import org.mockito.junit.jupiter.MockitoExtension; +import org.mockito.junit.jupiter.MockitoSettings; +import org.mockito.quality.Strictness; +import org.springframework.test.util.ReflectionTestUtils; + +import stirling.software.common.model.enumeration.Role; +import stirling.software.proprietary.security.database.repository.AuthorityRepository; +import stirling.software.proprietary.security.database.repository.UserRepository; +import stirling.software.proprietary.security.model.Authority; +import stirling.software.proprietary.security.model.User; + +/** + * Unit tests for {@link UserRoleService}. + * + *

The service is a thin orchestrator: it flips a user's {@link Authority} row and mirrors the + * role into the denormalized {@code roleName} column. All collaborators are mocked. + */ +@ExtendWith(MockitoExtension.class) +@MockitoSettings(strictness = Strictness.LENIENT) +class UserRoleServiceTest { + + @Mock private UserRepository userRepository; + @Mock private AuthorityRepository authorityRepository; + + private static final String ROLE_USER = Role.USER.getRoleId(); // "ROLE_USER" + private static final String ROLE_PRO_USER = Role.PRO_USER.getRoleId(); // "ROLE_PRO_USER" + + private UserRoleService service() { + return new UserRoleService(userRepository, authorityRepository); + } + + private static User user(long id, String username, String currentRole) { + User u = new User(); + u.setId(id); + u.setUsername(username); + u.setRoleName(currentRole); + return u; + } + + private static Authority authority(String currentRole) { + Authority a = new Authority(); + a.setId(7L); + a.setAuthority(currentRole); + return a; + } + + /** + * Reads the denormalized {@code roleName} column straight off the field. {@link + * User#getRoleName()} is overridden to derive the role from the authorities set (via {@link + * Role#fromString}), so it cannot observe the column that {@code changeRole} mirrors via {@code + * setRoleName}. + */ + private static String mirroredRoleName(User u) { + return (String) ReflectionTestUtils.getField(u, "roleName"); + } + + @Nested + @DisplayName("changeRole") + class ChangeRole { + + @Test + @DisplayName("flips the Authority row, mirrors roleName, and persists both") + void flipsAuthorityAndMirrorsRoleName() { + UserRoleService service = service(); + User u = user(42L, "alice@example.com", ROLE_USER); + Authority auth = authority(ROLE_USER); + when(authorityRepository.findByUserId(42L)).thenReturn(auth); + + service.changeRole(u, ROLE_PRO_USER); + + // Authority entity carries the new role and is saved. + assertThat(auth.getAuthority()).isEqualTo(ROLE_PRO_USER); + verify(authorityRepository).save(auth); + // Denormalized column mirrored on the User entity and saved. + assertThat(mirroredRoleName(u)).isEqualTo(ROLE_PRO_USER); + verify(userRepository).save(u); + } + + @Test + @DisplayName("persists the authority before the user (authority-first ordering)") + void persistsAuthorityBeforeUser() { + UserRoleService service = service(); + User u = user(42L, "alice@example.com", ROLE_USER); + Authority auth = authority(ROLE_USER); + when(authorityRepository.findByUserId(42L)).thenReturn(auth); + + service.changeRole(u, ROLE_PRO_USER); + + InOrder order = Mockito.inOrder(authorityRepository, userRepository); + order.verify(authorityRepository).save(auth); + order.verify(userRepository).save(u); + } + + @Test + @DisplayName("looks the authority up by the user's numeric id") + void looksUpAuthorityByUserId() { + UserRoleService service = service(); + User u = user(99L, "bob@example.com", ROLE_PRO_USER); + Authority auth = authority(ROLE_PRO_USER); + when(authorityRepository.findByUserId(99L)).thenReturn(auth); + + service.changeRole(u, ROLE_USER); + + verify(authorityRepository).findByUserId(99L); + assertThat(auth.getAuthority()).isEqualTo(ROLE_USER); + assertThat(mirroredRoleName(u)).isEqualTo(ROLE_USER); + } + + @Test + @DisplayName("setting the same role is a harmless no-op rewrite that still persists") + void sameRoleStillPersists() { + UserRoleService service = service(); + User u = user(42L, "carol@example.com", ROLE_USER); + Authority auth = authority(ROLE_USER); + when(authorityRepository.findByUserId(42L)).thenReturn(auth); + + service.changeRole(u, ROLE_USER); + + assertThat(auth.getAuthority()).isEqualTo(ROLE_USER); + verify(authorityRepository).save(auth); + verify(userRepository).save(u); + } + + @Test + @DisplayName("tolerates an empty authorities set on the User (logging reads roles as \"\")") + void emptyAuthoritiesSetIsTolerated() { + UserRoleService service = service(); + User u = user(42L, "dave@example.com", null); + // getRolesAsString() joins an empty set -> "" ; must not NPE in the debug log. + Authority auth = authority(null); + when(authorityRepository.findByUserId(42L)).thenReturn(auth); + + service.changeRole(u, ROLE_PRO_USER); + + assertThat(auth.getAuthority()).isEqualTo(ROLE_PRO_USER); + assertThat(mirroredRoleName(u)).isEqualTo(ROLE_PRO_USER); + } + } + + @Nested + @DisplayName("downgradeToFree") + class DowngradeToFree { + + @Test + @DisplayName("sets ROLE_USER, mirrors roleName, and persists both") + void setsUserRole() { + UserRoleService service = service(); + User u = user(42L, "eve@example.com", ROLE_PRO_USER); + Authority auth = authority(ROLE_PRO_USER); + when(authorityRepository.findByUserId(42L)).thenReturn(auth); + + service.downgradeToFree(u); + + assertThat(auth.getAuthority()).isEqualTo(ROLE_USER); + assertThat(mirroredRoleName(u)).isEqualTo(ROLE_USER); + verify(authorityRepository).save(auth); + verify(userRepository).save(u); + } + } + + @Nested + @DisplayName("upgradeToPro") + class UpgradeToPro { + + @Test + @DisplayName("sets ROLE_PRO_USER, mirrors roleName, and persists both") + void setsProRole() { + UserRoleService service = service(); + User u = user(42L, "ivan@example.com", ROLE_USER); + Authority auth = authority(ROLE_USER); + when(authorityRepository.findByUserId(42L)).thenReturn(auth); + + service.upgradeToPro(u); + + assertThat(auth.getAuthority()).isEqualTo(ROLE_PRO_USER); + assertThat(mirroredRoleName(u)).isEqualTo(ROLE_PRO_USER); + verify(authorityRepository).save(auth); + verify(userRepository).save(u); + } + } +} diff --git a/build.gradle b/build.gradle index 14df27ee5e..c7413605d1 100644 --- a/build.gradle +++ b/build.gradle @@ -78,7 +78,7 @@ springBoot { allprojects { group = 'stirling.software' - version = '2.12.0' + version = '2.13.0' configurations.configureEach { exclude group: "org.springframework.boot", module: "spring-boot-starter-tomcat" @@ -185,9 +185,10 @@ subprojects { allowInsecureProtocol = true } } - maven { url = "https://build.shibboleth.net/maven/releases" } - maven { url = "https://repository.jboss.org/" } + // Maven Central first; mirrors below are fallbacks for niche artifacts. mavenCentral() + maven { url = "https://repository.jboss.org/" } + maven { url = "https://build.shibboleth.net/maven/releases" } } configurations.configureEach { @@ -583,8 +584,9 @@ repositories { allowInsecureProtocol = true } } - maven { url = "https://build.shibboleth.net/maven/releases" } mavenCentral() + maven { url = "https://repository.jboss.org/" } + maven { url = "https://build.shibboleth.net/maven/releases" } } dependencies { diff --git a/devGuide/EXCEPTION_HANDLING_GUIDE.md b/devGuide/EXCEPTION_HANDLING_GUIDE.md index 666b70e2b1..f1ab95abea 100644 --- a/devGuide/EXCEPTION_HANDLING_GUIDE.md +++ b/devGuide/EXCEPTION_HANDLING_GUIDE.md @@ -16,7 +16,7 @@ Java forms the core of Stirling-PDF. When adding new features or handling errors 2. **Use `try-with-resources`** when working with streams or other closable resources to ensure clean-up even on failure. 3. **Return meaningful HTTP status codes** in controllers by throwing `ResponseStatusException` or using `@ExceptionHandler` methods. 4. **Log with context** using the project’s logging framework. Include identifiers or IDs that help trace the issue. -5. **Internationalise messages** by placing user-facing text in `messages_en_GB.properties` and referencing them with message keys. +5. **Internationalise messages** by placing user-facing text in `messages_en_US.properties` and referencing them with message keys. ## JavaScript @@ -55,11 +55,11 @@ except Exception as err: ## Internationalisation (i18n) -All user-visible error strings should be defined in the main translation file (`messages_en_GB.properties`). Other language files will use the same keys. Refer to messages in code rather than hard-coding text. +All user-visible error strings should be defined in the main translation file (`messages_en_US.properties`). Other language files will use the same keys. Refer to messages in code rather than hard-coding text. When creating new messages: -1. Add the English phrase to `messages_en_GB.properties`. +1. Add the English phrase to `messages_en_US.properties`. 2. Reference the message key in your Java, JavaScript, or Python code. 3. Update other localisation files as needed. diff --git a/devGuide/HowToAddNewLanguage.md b/devGuide/HowToAddNewLanguage.md index b2aef37420..09bc83d09d 100644 --- a/devGuide/HowToAddNewLanguage.md +++ b/devGuide/HowToAddNewLanguage.md @@ -16,7 +16,7 @@ Fork Stirling-PDF and create a new branch out of `main`. - Use hyphenated format: `pl-PL` (not underscore) 2. Copy the reference translation file: - - Source: `frontend/editor/public/locales/en-GB/translation.toml` + - Source: `frontend/editor/public/locales/en-US/translation.toml` - Destination: `frontend/editor/public/locales/pl-PL/translation.toml` 3. Translate all entries in the TOML file @@ -47,10 +47,10 @@ ignore = [ ## Add New Translation Tags > [!IMPORTANT] -> If you add any new translation tags, they must first be added to the `en-GB/translation.toml` file. This ensures consistency across all language files. +> If you add any new translation tags, they must first be added to the `en-US/translation.toml` file. This ensures consistency across all language files. -- New translation tags **must be added** to `frontend/editor/public/locales/en-GB/translation.toml` to maintain a reference for other languages. -- After adding the new tags to `en-GB/translation.toml`, add and translate them in the respective language file (e.g., `pl-PL/translation.toml`). +- New translation tags **must be added** to `frontend/editor/public/locales/en-US/translation.toml` to maintain a reference for other languages. +- After adding the new tags to `en-US/translation.toml`, add and translate them in the respective language file (e.g., `pl-PL/translation.toml`). - Use the scripts in `scripts/translations/` to validate and manage translations (see `scripts/translations/README.md`) Make sure to place the entry under the correct language section. This helps maintain the accuracy of translation progress statistics and ensures that the translation tool or scripts do not misinterpret the completion rate. diff --git a/devTools/package-lock.json b/devTools/package-lock.json index 2ea1628978..0555ac9543 100644 --- a/devTools/package-lock.json +++ b/devTools/package-lock.json @@ -894,10 +894,20 @@ "license": "MIT" }, "node_modules/js-yaml": { - "version": "4.1.1", - "resolved": "https://registry.npmjs.org/js-yaml/-/js-yaml-4.1.1.tgz", - "integrity": "sha512-qQKT4zQxXl8lLwBtHMWwaTcGfFOZviOJet3Oy/xmGk2gZH677CJM9EvtfdSkgWcATZhj/55JZ0rmy3myCT5lsA==", + "version": "4.2.0", + "resolved": "https://registry.npmjs.org/js-yaml/-/js-yaml-4.2.0.tgz", + "integrity": "sha512-ePWsvanv0DWuDRsW8dnt+R4jQ31SCRCQ7hhNcPXZPsoBZiemuZNYGf7adZdqX2D86j6rvKp3RpCxVTSb8WQlOw==", "dev": true, + "funding": [ + { + "type": "github", + "url": "https://github.com/sponsors/puzrin" + }, + { + "type": "github", + "url": "https://github.com/sponsors/nodeca" + } + ], "license": "MIT", "dependencies": { "argparse": "^2.0.1" diff --git a/docker/backend/Dockerfile b/docker/backend/Dockerfile new file mode 100644 index 0000000000..40abd00933 --- /dev/null +++ b/docker/backend/Dockerfile @@ -0,0 +1,127 @@ +# Stirling-PDF backend-only image — JAR built with -PbuildWithFrontend=false, UI ships separately. + +ARG BASE_VERSION=1.0.2 +ARG BASE_IMAGE=stirlingtools/stirling-pdf-base:${BASE_VERSION} + +# Stage 1: Build the Java application (backend only, no frontend) +FROM gradle:9.3.1-jdk25@sha256:85aec999629f4774a383cb792da4b598bdf5a7e69c4b9570bb70c0f919179183 AS app-build + +# JDK 25+: --add-exports is no longer accepted via JAVA_TOOL_OPTIONS; use JDK_JAVA_OPTIONS instead +ENV JDK_JAVA_OPTIONS="--add-exports=jdk.compiler/com.sun.tools.javac.api=ALL-UNNAMED \ + --add-exports=jdk.compiler/com.sun.tools.javac.file=ALL-UNNAMED \ + --add-exports=jdk.compiler/com.sun.tools.javac.parser=ALL-UNNAMED \ + --add-exports=jdk.compiler/com.sun.tools.javac.tree=ALL-UNNAMED \ + --add-exports=jdk.compiler/com.sun.tools.javac.util=ALL-UNNAMED" + +WORKDIR /app + +COPY build.gradle settings.gradle gradlew ./ +COPY gradle/ gradle/ +COPY app/core/build.gradle app/core/ +COPY app/common/build.gradle app/common/ +COPY app/proprietary/build.gradle app/proprietary/ + +# Use system gradle instead of gradlew to avoid SSL issues downloading gradle distribution on emulated arm64 +RUN gradle dependencies --no-daemon || true + +COPY . . + +ARG PROTOTYPES_BUILD=false +ARG STIRLING_FLAVOR=proprietary +ENV STIRLING_FLAVOR=${STIRLING_FLAVOR} + +# buildWithFrontend=false → backend-only JAR with API landing page. +# Bundle only the JPDFium native for this image's target arch. +ARG TARGETARCH +RUN JPDFIUM_PLATFORM="$([ "$TARGETARCH" = arm64 ] && echo linux-arm64 || echo linux-x64)" && \ + STIRLING_FLAVOR=${STIRLING_FLAVOR} \ + gradle clean build \ + -PbuildWithFrontend=false \ + -PjpdfiumPlatforms="$JPDFIUM_PLATFORM" \ + -PprototypesMode=${PROTOTYPES_BUILD} \ + -x spotlessApply -x spotlessCheck -x test -x sonarqube \ + --no-daemon + +# Stage 2: Extract Spring Boot Layers +FROM eclipse-temurin:25-jre-noble@sha256:b27ca47660a8fa837e47a8533b9b1a3a430295cf29ca28d91af4fd121572dc29 AS jar-extract +WORKDIR /tmp +COPY --from=app-build /app/app/core/build/libs/*.jar app.jar +RUN java -Djarmode=tools -jar app.jar extract --layers --destination /layers + + +# Stage 3: Final runtime image on top of pre-built base +FROM ${BASE_IMAGE} + +ARG VERSION_TAG + +WORKDIR /app + +# Application layers +COPY --link --from=jar-extract --chown=1000:1000 /layers/dependencies/ /app/ +COPY --link --from=jar-extract --chown=1000:1000 /layers/spring-boot-loader/ /app/ +COPY --link --from=jar-extract --chown=1000:1000 /layers/snapshot-dependencies/ /app/ +COPY --link --from=jar-extract --chown=1000:1000 /layers/application/ /app/ + +COPY --link --from=app-build --chown=1000:1000 \ + /app/build/libs/restart-helper.jar /restart-helper.jar +COPY --link --chown=1000:1000 scripts/ /scripts/ + +# Fonts go to system dir, root ownership is correct (world-readable) +COPY app/core/src/main/resources/static/fonts/*.ttf /usr/share/fonts/truetype/ + +# Permissions and configuration +RUN set -eux; \ + chmod +x /scripts/*; \ + ln -s /logs /app/logs; \ + ln -s /configs /app/configs; \ + ln -s /customFiles /app/customFiles; \ + ln -s /pipeline /app/pipeline; \ + ln -s /storage /app/storage; \ + chown -h stirlingpdfuser:stirlingpdfgroup /app/logs /app/configs /app/customFiles /app/pipeline /app/storage; \ + chown stirlingpdfuser:stirlingpdfgroup /app; \ + chmod 750 /tmp/stirling-pdf; \ + chmod 750 /tmp/stirling-pdf/heap_dumps; \ + fc-cache -f + +# Version file for scripts (init-without-ocr.sh reads /etc/stirling_version). +RUN echo "${VERSION_TAG:-dev}" > /etc/stirling_version + +# Environment variables +ENV VERSION_TAG=$VERSION_TAG \ + STIRLING_AOT_ENABLE="false" \ + STIRLING_JVM_PROFILE="balanced" \ + _JVM_OPTS_BALANCED="-XX:+ExitOnOutOfMemoryError -XX:+HeapDumpOnOutOfMemoryError -XX:HeapDumpPath=/configs/heap_dumps -XX:+UseG1GC -XX:MaxGCPauseMillis=200 -XX:G1HeapRegionSize=4m -XX:G1PeriodicGCInterval=60000 -XX:+UseStringDeduplication -XX:+UseCompactObjectHeaders -XX:+ExplicitGCInvokesConcurrent -Dspring.threads.virtual.enabled=true -Djava.awt.headless=true" \ + _JVM_OPTS_PERFORMANCE="-XX:+ExitOnOutOfMemoryError -XX:+HeapDumpOnOutOfMemoryError -XX:HeapDumpPath=/configs/heap_dumps -XX:+UseShenandoahGC -XX:ShenandoahGCMode=generational -XX:+UseCompactObjectHeaders -XX:+UseStringDeduplication -XX:+AlwaysPreTouch -XX:+ExplicitGCInvokesConcurrent -Dspring.threads.virtual.enabled=true -Djava.awt.headless=true" \ + JAVA_CUSTOM_OPTS="" \ + HOME=/home/stirlingpdfuser \ + PUID=1000 \ + PGID=1000 \ + UMASK=022 \ + STIRLING_TEMPFILES_DIRECTORY=/tmp/stirling-pdf \ + TMPDIR=/tmp/stirling-pdf \ + TEMP=/tmp/stirling-pdf \ + TMP=/tmp/stirling-pdf \ + DBUS_SESSION_BUS_ADDRESS=/dev/null \ + SAL_TMP=/tmp/stirling-pdf/libre + +# Metadata labels +LABEL org.opencontainers.image.title="Stirling-PDF Backend" \ + org.opencontainers.image.description="Backend-only version (no embedded UI) with Calibre, LibreOffice, Tesseract, OCRmyPDF" \ + org.opencontainers.image.source="https://github.com/Stirling-Tools/Stirling-PDF" \ + org.opencontainers.image.licenses="MIT" \ + org.opencontainers.image.vendor="Stirling-Tools" \ + org.opencontainers.image.url="https://www.stirlingpdf.com" \ + org.opencontainers.image.documentation="https://docs.stirlingpdf.com" \ + maintainer="Stirling-Tools" \ + org.opencontainers.image.authors="Stirling-Tools" \ + org.opencontainers.image.version="${VERSION_TAG}" \ + org.opencontainers.image.keywords="PDF, manipulation, backend, API, Spring Boot" + +EXPOSE 8080/tcp +STOPSIGNAL SIGTERM + +HEALTHCHECK --interval=30s --timeout=15s --start-period=120s --retries=5 \ + CMD curl -fs --max-time 10 http://localhost:8080${SYSTEM_ROOTURIPATH:-''}/api/v1/info/status || exit 1 + +ENTRYPOINT ["tini", "--", "/scripts/init.sh"] +CMD [] diff --git a/docker/embedded/Dockerfile b/docker/embedded/Dockerfile index 89913af09a..89ed06397a 100644 --- a/docker/embedded/Dockerfile +++ b/docker/embedded/Dockerfile @@ -5,7 +5,7 @@ ARG BASE_VERSION=1.0.2 ARG BASE_IMAGE=stirlingtools/stirling-pdf-base:${BASE_VERSION} # Stage 1: Build the Java application and frontend -FROM gradle:9.3.1-jdk25@sha256:85aec999629f4774a383cb792da4b598bdf5a7e69c4b9570bb70c0f919179183 AS app-build +FROM gradle:9.5.1-jdk25@sha256:8de3543f1772bb66be3b275893e5977b6d8bd2b0d25551faa5846a821d1f0600 AS app-build ARG TASK_VERSION=3.49.1 RUN apt-get update \ @@ -43,9 +43,13 @@ ARG PROTOTYPES_BUILD=false ARG STIRLING_FLAVOR=proprietary ENV STIRLING_FLAVOR=${STIRLING_FLAVOR} -RUN STIRLING_FLAVOR=${STIRLING_FLAVOR} \ +# Bundle only the JPDFium native for this image's target arch. +ARG TARGETARCH +RUN JPDFIUM_PLATFORM="$([ "$TARGETARCH" = arm64 ] && echo linux-arm64 || echo linux-x64)" && \ + STIRLING_FLAVOR=${STIRLING_FLAVOR} \ gradle clean build \ -PbuildWithFrontend=true \ + -PjpdfiumPlatforms="$JPDFIUM_PLATFORM" \ -PprototypesMode=${PROTOTYPES_BUILD} \ -x spotlessApply -x spotlessCheck -x test -x sonarqube \ --no-daemon diff --git a/docker/embedded/Dockerfile.fat b/docker/embedded/Dockerfile.fat index f9754b4721..0beb8406ec 100644 --- a/docker/embedded/Dockerfile.fat +++ b/docker/embedded/Dockerfile.fat @@ -6,7 +6,7 @@ ARG BASE_VERSION=1.0.2 ARG BASE_IMAGE=stirlingtools/stirling-pdf-base:${BASE_VERSION} # Stage 1: Build the Java application and frontend -FROM gradle:9.3.1-jdk25 AS app-build +FROM gradle:9.5.1-jdk25@sha256:8de3543f1772bb66be3b275893e5977b6d8bd2b0d25551faa5846a821d1f0600 AS app-build ARG TASK_VERSION=3.49.1 RUN apt-get update \ @@ -40,9 +40,13 @@ RUN gradle dependencies --no-daemon || true COPY . . -RUN DISABLE_ADDITIONAL_FEATURES=false \ +# Bundle only the JPDFium native for this image's target arch. +ARG TARGETARCH +RUN JPDFIUM_PLATFORM="$([ "$TARGETARCH" = arm64 ] && echo linux-arm64 || echo linux-x64)" && \ + DISABLE_ADDITIONAL_FEATURES=false \ gradle clean build \ -PbuildWithFrontend=true \ + -PjpdfiumPlatforms="$JPDFIUM_PLATFORM" \ -x spotlessApply -x spotlessCheck -x test -x sonarqube \ --no-daemon diff --git a/docker/embedded/Dockerfile.ultra-lite b/docker/embedded/Dockerfile.ultra-lite index d091d7b3fa..9f77716646 100644 --- a/docker/embedded/Dockerfile.ultra-lite +++ b/docker/embedded/Dockerfile.ultra-lite @@ -2,7 +2,7 @@ # Single JAR contains both frontend and backend with minimal dependencies # Stage 1: Build application with embedded frontend -FROM gradle:9.3.1-jdk25 AS build +FROM gradle:9.5.1-jdk25@sha256:8de3543f1772bb66be3b275893e5977b6d8bd2b0d25551faa5846a821d1f0600 AS build # Install Node.js and npm for frontend build ARG TASK_VERSION=3.49.1 @@ -39,17 +39,23 @@ RUN ./gradlew dependencies --no-daemon || true # Copy entire project COPY . . -# Build ultra-lite JAR with embedded frontend (minimal features) -RUN DISABLE_ADDITIONAL_FEATURES=true \ +# Build ultra-lite JAR with embedded frontend (minimal features). +# Bundle only the JPDFium native for this image's target arch. +ARG TARGETARCH +RUN JPDFIUM_PLATFORM="$([ "$TARGETARCH" = arm64 ] && echo linux-arm64 || echo linux-x64)" && \ + DISABLE_ADDITIONAL_FEATURES=true \ ./gradlew clean build \ -PbuildWithFrontend=true \ + -PjpdfiumPlatforms="$JPDFIUM_PLATFORM" \ -x spotlessApply -x spotlessCheck -x test -x sonarqube \ --no-daemon # Stage 2: Runtime image -FROM eclipse-temurin:25-jre-alpine +# glibc base (not Alpine/musl): JPDFium's PDFium natives are glibc-linked. +FROM eclipse-temurin:25-jre-noble@sha256:b27ca47660a8fa837e47a8533b9b1a3a430295cf29ca28d91af4fd121572dc29 -ENV LANG=C.UTF-8 \ +ENV DEBIAN_FRONTEND=noninteractive \ + LANG=C.UTF-8 \ LC_ALL=C.UTF-8 ARG VERSION_TAG @@ -87,22 +93,22 @@ ENV VERSION_TAG=$VERSION_TAG \ ENDPOINTS_GROUPS_TO_REMOVE=CLI # Install minimal dependencies -RUN echo "@main https://dl-cdn.alpinelinux.org/alpine/edge/main" | tee -a /etc/apk/repositories && \ - echo "@community https://dl-cdn.alpinelinux.org/alpine/edge/community" | tee -a /etc/apk/repositories && \ - echo "@testing https://dl-cdn.alpinelinux.org/alpine/edge/testing" | tee -a /etc/apk/repositories && \ - apk upgrade --no-cache -a && \ - apk add --no-cache \ +RUN mkdir -p $HOME /configs /logs /customFiles /pipeline/watchedFolders /pipeline/finishedFolders /storage /tmp/stirling-pdf /tmp/stirling-pdf/heap_dumps && \ + mkdir -p /usr/share/fonts/opentype/noto && \ + apt-get update && \ + apt-get install -y --no-install-recommends \ ca-certificates \ tzdata \ tini \ bash \ curl \ - shadow \ + procps \ util-linux && \ - mkdir -p $HOME /configs /logs /customFiles /pipeline/watchedFolders /pipeline/finishedFolders /storage /tmp/stirling-pdf /tmp/stirling-pdf/heap_dumps && \ - mkdir -p /usr/share/fonts/opentype/noto && \ + rm -rf /var/lib/apt/lists/* && \ # User permissions - addgroup -S stirlingpdfgroup && adduser -S stirlingpdfuser -G stirlingpdfgroup && \ + userdel -r ubuntu 2>/dev/null || true && \ + groupdel ubuntu 2>/dev/null || true && \ + groupadd -g 1000 stirlingpdfgroup && useradd -u 1000 -d $HOME -s /bin/bash -g stirlingpdfgroup stirlingpdfuser && \ chown -R stirlingpdfuser:stirlingpdfgroup $HOME /configs /customFiles /pipeline /storage /tmp/stirling-pdf # Copy scripts and built artifacts after OS package layer to maximize cache reuse. diff --git a/docker/frontend/Dockerfile b/docker/frontend/Dockerfile index e92f19ddc1..67e35f82c5 100644 --- a/docker/frontend/Dockerfile +++ b/docker/frontend/Dockerfile @@ -1,38 +1,44 @@ -# Frontend Dockerfile - React/Vite application +# check=skip=SecretsUsedInArgOrEnv +# Supabase publishable ARG is client-safe by design, not a real secret. + +# Stage 1: build FROM node:25-alpine@sha256:e80397b81fa93888b5f855e8bef37d9b18d3c5eb38b8731fc23d6d878647340f AS build WORKDIR /app -# Copy package files COPY frontend/package.json frontend/package-lock.json ./ - -# Install dependencies RUN npm ci -# Copy source code COPY frontend . -# Build the application (vite root is editor/, output lands in editor/dist/) -RUN npx vite build editor +# Generate material-symbols icon subset (normally done by task prepare:icons). +RUN node editor/scripts/generate-icons.js -# Production stage +# Defaults match prior behaviour. Supabase values are client-publishable build args. +ARG STIRLING_FLAVOR=proprietary +ARG VITE_BUILD_MODE=production +ARG VITE_SUPABASE_URL="" +# pragma: allowlist secret +ARG VITE_SUPABASE_PUBLISHABLE_DEFAULT_KEY="" + +# Build vite from editor/, output lands in editor/dist/. +RUN set -eu; \ + export STIRLING_FLAVOR="${STIRLING_FLAVOR}"; \ + if [ -n "${VITE_SUPABASE_URL}" ]; then export VITE_SUPABASE_URL="${VITE_SUPABASE_URL}"; fi; \ + if [ -n "${VITE_SUPABASE_PUBLISHABLE_DEFAULT_KEY}" ]; then export VITE_SUPABASE_PUBLISHABLE_DEFAULT_KEY="${VITE_SUPABASE_PUBLISHABLE_DEFAULT_KEY}"; fi; \ + npx vite build editor --mode "${VITE_BUILD_MODE}" + +# Stage 2: nginx FROM nginx:alpine@sha256:b0f7830b6bfaa1258f45d94c240ab668ced1b3651c8a222aefe6683447c7bf55 -# Copy built files from build stage COPY --from=build /app/editor/dist /usr/share/nginx/html - -# Copy nginx configuration and entrypoint COPY docker/frontend/nginx.conf /etc/nginx/nginx.conf COPY docker/frontend/entrypoint.sh /entrypoint.sh -# Make entrypoint executable RUN chmod +x /entrypoint.sh -# Expose port 80 (standard HTTP port) EXPOSE 80 -# Environment variables for flexibility ENV VITE_API_BASE_URL=http://backend:8080 -# Use custom entrypoint -ENTRYPOINT ["/entrypoint.sh"] \ No newline at end of file +ENTRYPOINT ["/entrypoint.sh"] diff --git a/docs/counter_translation.md b/docs/counter_translation.md index b2cdd74451..8468dfe1b8 100644 --- a/docs/counter_translation.md +++ b/docs/counter_translation.md @@ -3,7 +3,7 @@ ## Overview The script [`scripts/counter_translation.py`](../scripts/counter_translation.py) checks the translation progress of the property files in the directory `app/core/src/main/resources/`. -It compares each `messages_*.properties` file with the English reference file `messages_en_GB.properties` and calculates a percentage of completion for each language. +It compares each `messages_*.properties` file with the English reference file `messages_en_US.properties` and calculates a percentage of completion for each language. In addition to console output, the script automatically updates the progress badges in the project’s `README.md` and maintains the configuration file [`scripts/ignore_translation.toml`](../scripts/ignore_translation.toml), which lists translation keys to be ignored for each language. diff --git a/engine/.env b/engine/.env index e334a9bf90..a99abf0ccf 100644 --- a/engine/.env +++ b/engine/.env @@ -14,19 +14,30 @@ STIRLING_FAST_MODEL=anthropic:claude-haiku-4-5 STIRLING_SMART_MODEL_MAX_TOKENS=8192 STIRLING_FAST_MODEL_MAX_TOKENS=2048 -# RAG Configuration — retrieval-augmented generation is always on. -# Embedding provider credentials are handled natively (e.g. VOYAGE_API_KEY for VoyageAI). -STIRLING_RAG_EMBEDDING_MODEL=voyageai:voyage-4 +# Process-wide cap on concurrent model API calls, shared by both model tiers. +# Per-request fan-outs (chunked reasoner workers, contradiction detection) are +# bounded per request; this bounds their product across concurrent requests. +STIRLING_MODEL_MAX_CONCURRENCY=32 -# Vector store backend: "sqlite" (embedded) or "pgvector" (external Postgres). -STIRLING_RAG_BACKEND=sqlite +# Document store: the one database holding vector chunks, ordered page text, +# and ACL rows. Backend is "sqlite" (embedded sqlite-vec) or "pgvector" +# (external Postgres). +STIRLING_DOCUMENTS_BACKEND=sqlite # Path to the sqlite-vec database file (used when backend=sqlite). -STIRLING_RAG_STORE_PATH=data/rag.db +STIRLING_DOCUMENTS_SQLITE_PATH=data/rag.db # Postgres DSN for pgvector (used when backend=pgvector). Leave empty when backend=sqlite. # Example: postgresql://user:password@host:5432/dbname -STIRLING_RAG_PGVECTOR_DSN= +STIRLING_DOCUMENTS_PGVECTOR_DSN= + +# Connection pool bounds for the pgvector backend. +STIRLING_DOCUMENTS_PGVECTOR_POOL_MIN_SIZE=1 +STIRLING_DOCUMENTS_PGVECTOR_POOL_MAX_SIZE=10 + +# RAG Configuration - retrieval-augmented generation is always on. +# Embedding provider credentials are handled natively (e.g. VOYAGE_API_KEY for VoyageAI). +STIRLING_RAG_EMBEDDING_MODEL=voyageai:voyage-4 STIRLING_RAG_CHUNK_SIZE=512 STIRLING_RAG_CHUNK_OVERLAP=64 @@ -57,6 +68,11 @@ STIRLING_CHUNKED_REASONER_NOTES_CHAR_BUDGET=250000 STIRLING_MAX_PAGES=200 STIRLING_MAX_CHARACTERS=200000 +# Reject API requests that lack an X-User-Id header. Self-hosted deployments +# with security disabled have no user identity, so this is off by default. +# Multi-tenant (SaaS) deployments must set it to true. +STIRLING_REQUIRE_USER_ID=false + # PostHog analytics. Set STIRLING_POSTHOG_ENABLED=true and provide an API key to enable. STIRLING_POSTHOG_ENABLED=false STIRLING_POSTHOG_API_KEY=phc_VOdeYnlevc2T63m3myFGjeBlRcIusRgmhfx6XL5a1iz diff --git a/engine/Dockerfile b/engine/Dockerfile index c9c8d22391..375983142a 100644 --- a/engine/Dockerfile +++ b/engine/Dockerfile @@ -10,18 +10,27 @@ RUN apt-get update \ && rm /tmp/task.deb \ && rm -rf /var/lib/apt/lists/* -WORKDIR /app +# Source under /app/engine/ to match root Taskfile's `includes.engine.dir: engine`. +WORKDIR /app/engine -COPY pyproject.toml uv.lock Taskfile.yml .env ./ -COPY .taskfiles/ ./.taskfiles/ +COPY pyproject.toml uv.lock .env ./ COPY scripts/ ./scripts/ RUN --mount=type=cache,target=/root/.cache/uv \ uv sync --frozen --no-dev COPY src/ ./src/ -ENV PATH="/app/.venv/bin:$PATH" +WORKDIR /app +COPY Taskfile.yml ./ +COPY .taskfiles/ ./.taskfiles/ + +ENV PATH="/app/engine/.venv/bin:$PATH" ENV PYTHONUNBUFFERED=1 +ENV STIRLING_ENGINE_WORKERS=4 +# Container runs on a fixed port; skip the host-only free-port probe (its script +# is not shipped in the image). engine:run honours these. +ENV ENGINE_PORT_PROBE=false +ENV STIRLING_ENGINE_PORT=5001 EXPOSE 5001 diff --git a/engine/pyproject.toml b/engine/pyproject.toml index 4074c21ee2..fa972c8495 100644 --- a/engine/pyproject.toml +++ b/engine/pyproject.toml @@ -5,8 +5,9 @@ description = "AI Document Engine" requires-python = ">=3.13" dependencies = [ "fastapi>=0.116.0", + "jinja2>=3.1.0", "pgvector>=0.3.6", - "psycopg[binary]>=3.2", + "psycopg[binary,pool]>=3.2", "pydantic>=2.0.0", "pydantic-ai>=1.67.0", "pydantic-ai-slim[voyageai]>=1.67.0", diff --git a/engine/src/stirling/agents/__init__.py b/engine/src/stirling/agents/__init__.py index cddd0275c3..c22bd6c977 100644 --- a/engine/src/stirling/agents/__init__.py +++ b/engine/src/stirling/agents/__init__.py @@ -2,20 +2,20 @@ from .execution import ExecutionPlanningAgent from .orchestrator import OrchestratorAgent +from .pdf_create import PdfCreateAgent from .pdf_edit import PdfEditAgent, PdfEditParameterSelector, PdfEditPlanSelection from .pdf_questions import PdfQuestionAgent from .pdf_review import PdfReviewAgent -from .pdf_to_markdown import PdfToMarkdownAgent from .user_spec import UserSpecAgent __all__ = [ "ExecutionPlanningAgent", "OrchestratorAgent", + "PdfCreateAgent", "PdfEditAgent", "PdfEditParameterSelector", "PdfEditPlanSelection", "PdfQuestionAgent", "PdfReviewAgent", - "PdfToMarkdownAgent", "UserSpecAgent", ] diff --git a/engine/src/stirling/agents/orchestrator.py b/engine/src/stirling/agents/orchestrator.py index 4dbf0b65ab..c73d9dab32 100644 --- a/engine/src/stirling/agents/orchestrator.py +++ b/engine/src/stirling/agents/orchestrator.py @@ -8,17 +8,17 @@ from pydantic_ai import Agent from pydantic_ai.output import ToolOutput from pydantic_ai.tools import RunContext +from stirling.agents.pdf_create import PdfCreateAgent from stirling.agents.pdf_edit import PdfEditAgent from stirling.agents.pdf_questions import PdfQuestionAgent from stirling.agents.pdf_review import PdfReviewAgent -from stirling.agents.pdf_to_markdown import PdfToMarkdownAgent from stirling.agents.user_spec import UserSpecAgent from stirling.contracts import ( AgentDraftWorkflowResponse, + ConvertMarkdownResponse, ExtractedTextArtifact, OrchestratorRequest, OrchestratorResponse, - PageLayoutArtifact, PdfEditResponse, PdfQuestionOrchestrateResponse, PdfReviewOrchestrateResponse, @@ -27,7 +27,7 @@ from stirling.contracts import ( format_conversation_history, format_file_names, ) -from stirling.contracts.pdf_to_markdown import PdfToMarkdownOrchestrateResponse +from stirling.contracts.pdf_create import PdfCreateOrchestrateResponse from stirling.services import AppRuntime logger = logging.getLogger(__name__) @@ -72,9 +72,21 @@ class OrchestratorAgent: ), ), ToolOutput( - self.delegate_pdf_to_markdown, - name="delegate_pdf_to_markdown", - description=("Delegate requests to reconstruct a PDF as a Markdown document."), + self.delegate_pdf_ingest, + name="delegate_pdf_ingest", + description=( + "Delegate requests to convert a PDF to Markdown or extract its content as readable text." + ), + ), + ToolOutput( + self.delegate_pdf_create, + name="delegate_pdf_create", + description=( + "Delegate requests to create a new PDF document from scratch based on a" + " description. Use this when the user wants to generate a new document" + " (e.g. 'create an invoice', 'write a report', 'make a contract'," + " 'draft a letter'). No input file is required." + ), ), ToolOutput( self.unsupported_capability, @@ -92,8 +104,10 @@ class OrchestratorAgent: "Use delegate_pdf_review when the user wants the PDF returned with review" " comments attached — anything like 'review this', 'annotate with comments'," " 'leave feedback on the PDF'. " - "Use delegate_pdf_to_markdown for any request to convert a PDF to Markdown " - "or reconstruct its content as readable text. " + "Use delegate_pdf_create when the user wants to generate a new document from" + " scratch with no input file — invoices, reports, letters, contracts, etc. " + "Use delegate_pdf_ingest for any request to convert a PDF to Markdown " + "or extract its content as readable text. " "Use unsupported_capability when the user asks about the assistant itself " "or when none of the other outputs fit; supply a helpful message." ), @@ -133,8 +147,8 @@ class OrchestratorAgent: return await self._run_pdf_edit(request) case SupportedCapability.AGENT_DRAFT: return await self._run_agent_draft(request) - case SupportedCapability.PDF_TO_MARKDOWN: - return await self._run_pdf_to_markdown(request) + case SupportedCapability.PDF_CREATE: + return await self._run_pdf_create(request) case ( SupportedCapability.ORCHESTRATE | SupportedCapability.AGENT_REVISE @@ -163,11 +177,12 @@ class OrchestratorAgent: async def _run_agent_draft(self, request: OrchestratorRequest) -> AgentDraftWorkflowResponse: return await UserSpecAgent(self.runtime).orchestrate(request) - async def delegate_pdf_to_markdown(self, ctx: RunContext[OrchestratorDeps]) -> PdfToMarkdownOrchestrateResponse: - return await self._run_pdf_to_markdown(ctx.deps.request) - - async def _run_pdf_to_markdown(self, request: OrchestratorRequest) -> PdfToMarkdownOrchestrateResponse: - return await PdfToMarkdownAgent(self.runtime).orchestrate(request) + async def delegate_pdf_ingest(self, ctx: RunContext[OrchestratorDeps]) -> ConvertMarkdownResponse: + request = ctx.deps.request + return ConvertMarkdownResponse( + reason="PDF to Markdown requested — Java converts deterministically.", + files_to_ingest=request.files, + ) async def delegate_pdf_review(self, ctx: RunContext[OrchestratorDeps]) -> PdfReviewOrchestrateResponse: return await self._run_pdf_review(ctx.deps.request) @@ -175,6 +190,12 @@ class OrchestratorAgent: async def _run_pdf_review(self, request: OrchestratorRequest) -> PdfReviewOrchestrateResponse: return await PdfReviewAgent(self.runtime).orchestrate(request) + async def delegate_pdf_create(self, ctx: RunContext[OrchestratorDeps]) -> PdfCreateOrchestrateResponse: + return await self._run_pdf_create(ctx.deps.request) + + async def _run_pdf_create(self, request: OrchestratorRequest) -> PdfCreateOrchestrateResponse: + return await PdfCreateAgent(self.runtime).orchestrate(request) + async def unsupported_capability( self, ctx: RunContext[OrchestratorDeps], @@ -204,10 +225,5 @@ class OrchestratorAgent: file_names = [f.file_name for f in artifact.files] descriptions.append(f"- extracted_text: {total_pages} pages from {file_names}") continue - if isinstance(artifact, PageLayoutArtifact): - total_pages = sum(len(f.pages) for f in artifact.files) - file_names = [f.file_name for f in artifact.files] - descriptions.append(f"- page_layout: {total_pages} pages from {file_names}") - continue descriptions.append("- unknown artifact") return "\n".join(descriptions) diff --git a/engine/src/stirling/agents/pdf_create/__init__.py b/engine/src/stirling/agents/pdf_create/__init__.py new file mode 100644 index 0000000000..20e3b356a5 --- /dev/null +++ b/engine/src/stirling/agents/pdf_create/__init__.py @@ -0,0 +1,3 @@ +from .agent import PdfCreateAgent + +__all__ = ["PdfCreateAgent"] diff --git a/engine/src/stirling/agents/pdf_create/agent.py b/engine/src/stirling/agents/pdf_create/agent.py new file mode 100644 index 0000000000..bd159060df --- /dev/null +++ b/engine/src/stirling/agents/pdf_create/agent.py @@ -0,0 +1,443 @@ +"""PDF Create Agent — chunked multi-agent pipeline. + +Flow: + 1. MetaPlannerAgent (smart_model) analyses the request and produces DocumentMeta: + title, tone, shared terms, style, and cannot_do_reason. No sections yet. + 2. SectionPlannerAgent (smart_model) reads the meta and produces DocumentSections: + ordered list of PlannedSection with heading, type, depth, and key_points. + 3. Python assembles DocumentPlan from meta + sections, then groups sections into + chunks, each staying under the output-token ceiling. + 4. SectionWriterAgents (smart_model) run in parallel via asyncio.gather. + Each returns a WrittenSections with fully populated DocumentSection objects. + 5. The assembler collects sections in plan order → GeneratedDocument. + 6. Jinja renders the document to HTML. The LLM never writes HTML. + +The planner is split into two calls (meta then sections) so each LLM output schema +stays small enough for grammar compilation on all model tiers including Haiku. +""" + +from __future__ import annotations + +import asyncio +import logging +import re +from dataclasses import dataclass +from pathlib import Path + +from jinja2 import Environment, FileSystemLoader +from pydantic_ai import Agent +from pydantic_ai.output import NativeOutput + +from stirling.contracts import ( + EditCannotDoResponse, + EditPlanResponse, + OrchestratorRequest, + ToolOperationStep, + format_conversation_history, +) +from stirling.contracts.pdf_create import ( + DocumentMeta, + DocumentPlan, + DocumentSection, + DocumentSections, + GeneratedDocument, + PdfCreateOrchestrateResponse, + PlannedSection, + SectionDepth, + WrittenSections, +) +from stirling.models.agent_tool_models import AgentToolId, CreatePdfFromHtmlAgentParams +from stirling.services import AppRuntime + +logger = logging.getLogger(__name__) + +_TEMPLATES_DIR = Path(__file__).parent / "templates" + +# ── Token budget ────────────────────────────────────────────────────────────────────────────────── + +# Conservative per-section token estimates mapped from planner-assigned depth. +_DEPTH_TOKENS: dict[SectionDepth, int] = { + SectionDepth.BRIEF: 250, + SectionDepth.STANDARD: 550, + SectionDepth.DETAILED: 1200, +} + +# Maximum output tokens per writer call. Stays well below the quality cliff (~4k). +_CHUNK_CEILING = 3000 + +# Cap on simultaneous writer calls so a large document doesn't open a burst of LLM +# connections and trip provider rate limits. +_MAX_PARALLEL_WRITERS = 10 + +# ── Chunk dataclass ─────────────────────────────────────────────────────────────────────────────── + + +@dataclass(frozen=True) +class _Chunk: + index: int + sections: list[PlannedSection] + # Descriptions of neighbouring chunks from the plan — passed to writers as + # read-only context so they can open/close their sections naturally. + context_before: str | None + context_after: str | None + + +# ── Chunking logic ──────────────────────────────────────────────────────────────────────────────── + + +def _describe_sections(sections: list[PlannedSection]) -> str: + """One-line summary of a chunk used as neighbour context.""" + return "; ".join(f'"{s.heading}" ({s.type.value})' for s in sections) + + +def _make_chunks(sections: list[PlannedSection]) -> list[_Chunk]: + """Group planned sections into chunks, each under _CHUNK_CEILING output tokens. + + Section boundaries are atomic — a section is never split across chunks. + A single section whose estimated cost exceeds the ceiling gets its own chunk. + asyncio.gather preserves insertion order so chunk index is only used for logging. + """ + if not sections: + return [] + + groups: list[list[PlannedSection]] = [] + current: list[PlannedSection] = [] + current_tokens = 0 + + for section in sections: + cost = _DEPTH_TOKENS[section.depth] + if current and current_tokens + cost > _CHUNK_CEILING: + groups.append(current) + current = [section] + current_tokens = cost + else: + current.append(section) + current_tokens += cost + + if current: + groups.append(current) + + chunks: list[_Chunk] = [] + for i, group in enumerate(groups): + context_before = _describe_sections(groups[i - 1]) if i > 0 else None + context_after = _describe_sections(groups[i + 1]) if i < len(groups) - 1 else None + chunks.append( + _Chunk( + index=i, + sections=group, + context_before=context_before, + context_after=context_after, + ) + ) + + return chunks + + +# ── Prompts ─────────────────────────────────────────────────────────────────────────────────────── + +_META_PLANNER_SYSTEM_PROMPT = """\ +You are a document planner. Your job is Step 1 of 2: produce the document header — NOT the +section list (that comes in Step 2) and NOT any body text (section writers handle that). + +Analyse the user's request and produce a DocumentMeta with: + +- title, subtitle (if appropriate), reference_number (only if the user supplies one explicitly) + +- tone_brief: one sentence describing register and style + (e.g. "Formal legal language, third person, present tense." or + "Professional business tone, active voice.") + +- shared_terms: consistent names for key entities AND ground-truth facts used throughout + the document. Two rules: + 1. Capture EVERY value the user states explicitly that could be referenced in more than + one section. This includes — but is not limited to: + · Named parties, organisations, products, or systems + · Numeric values: amounts, quantities, percentages, durations, limits + · Identifiers: version numbers, reference codes, model names + · Dates and time periods + · Units of measure or currency + 2. For any fact that will appear in two or more sections and that the user did NOT specify + (e.g. a default time, a standard rate, a typical threshold), assign ONE specific value + here. Do NOT let multiple writers independently invent the same fact. + Examples: {"the Agreement": "this Non-Disclosure Agreement", "the Client": "Acme Corp", + "contract value": "£120,000", "notice period": "30 days"} + +- document_context: a single sentence anchoring the temporal or versioning context of the + document, if the user provides one. Leave empty if the user provides no such context. + +- style_primary_color: accent and heading colour. Set ONLY when the user explicitly names a + colour or colour scheme (e.g. "make it red", "use navy blue"). Use CSS named colours + (e.g. "magenta", "navy", "crimson") or hex values. Leave null if no colour is stated. +- style_background_color: page background colour. Set only if explicitly requested. +- style_body_text_color: body text colour. Set only if explicitly requested. + +- cannot_do_reason: set this ONLY when the request is not asking to create a document at all + (e.g. a question, a greeting, an edit request to an existing document). Never set it + because the document is large, complex, or technically detailed. Leave null otherwise. + +RULES: +1. Extract ALL information the user provides. Do not invent content. +2. Do not produce any sections — that is Step 2. +""" + +_SECTIONS_PLANNER_SYSTEM_PROMPT = """\ +You are a document planner. Your job is Step 2 of 2: produce the ordered section list for +a document whose header has already been decided. Do NOT write any body text. + +You will be given: + - The document meta (title, tone, shared terms, etc.) produced in Step 1 + - The original user request + +Produce a DocumentSections with an ordered list of PlannedSection objects. + +For each section choose: + type — the most appropriate section type: + text — prose paragraphs (narrative, obligations, terms, descriptions) + key_value — labelled fields (parties, dates, metadata, identifiers) + line_items — tables with column headers (expenses, schedules, item lists) + bullet_list — unordered items (requirements, responsibilities, definitions) + signature — sign-off blocks for named parties or roles + + depth — honest estimate of content volume: + brief (~250 tokens) — 1-2 items, a short paragraph, or a small table + standard (~550 tokens) — a few paragraphs, a medium table, or a moderate list + detailed (~1200 tokens) — long clauses, complex multi-row tables, or dense content + + key_points — specific points this section MUST cover, taken directly from the user's input. + These are instructions to the writer, not summaries. Be precise and complete. + Every fact, name, date, amount, and requirement the user provides must appear somewhere. + For large documents, include enough key_points that the writer can produce substantial + content. + +RULES: +1. Extract ALL information the user provides. Do not invent content. +2. Assign depth honestly — for a long detailed document most sections will be detailed. +3. For large documents, produce as many sections as needed — there is no section count limit. +4. Use the shared_terms from the meta exactly when writing key_points. +""" + +_WRITER_SYSTEM_PROMPT = """\ +You are a section writer for a structured document. +Write ONLY the sections assigned to you — no extras, no merging, no skipping. + +SECTION TYPES — produce sections of exactly the requested type: + text — prose paragraphs. Use \\n\\n between paragraphs. + key_value — list of (label, value) pairs. Labels ≤ 5 words. Values verbatim from the data. + line_items — table. Every row must have exactly as many cells as there are columns. + bullet_list — flat list of items. + signature — list of signatory names/roles. + +RULES: +1. Write ONLY the sections in your assignment list, in the order given. +2. Cover every key_point listed for each section. Do not omit any. +3. Use the shared_terms exactly — no paraphrasing or substituting alternatives. + Shared terms are ground truth. If your general knowledge or a common default would + produce a different value (e.g. a different duration, amount, date, or version number), + the shared term takes precedence. This applies everywhere in the document, including + boilerplate, FAQ, and summary sections. +4. Match the depth for each section: brief = concise, standard = moderate, \ +detailed = thorough. +5. Maintain the document's tone throughout. +6. Do not reference other sections by number (e.g. "as defined in Section 3"). +7. If a document_context is provided, use it to anchor any dates, versions, or time + references you generate. Do not invent a different temporal or versioning context. +""" + + +def _build_sections_prompt(meta: DocumentMeta, user_request: str, history: str) -> str: + lines: list[str] = [ + "Document meta from Step 1:", + f" Title: {meta.title}", + f" Tone: {meta.tone_brief}", + ] + if meta.subtitle: + lines.append(f" Subtitle: {meta.subtitle}") + if meta.document_context: + lines.append(f" Document context: {meta.document_context}") + if meta.shared_terms: + lines.append(" Shared terms:") + for term, referent in meta.shared_terms.items(): + lines.append(f" {term} → {referent}") + + lines.append(f"\nConversation history:\n{history}") + lines.append(f"\nUser request: {user_request}") + return "\n".join(lines) + + +def _build_writer_prompt(plan: DocumentPlan, chunk: _Chunk) -> str: + lines: list[str] = [ + f"Document: {plan.title}", + f"Tone: {plan.tone_brief}", + ] + + if plan.document_context: + lines.append(f"Document context: {plan.document_context}") + + if plan.shared_terms: + lines.append("Ground-truth facts and shared terms (use exactly — these override defaults):") + for term, referent in plan.shared_terms.items(): + lines.append(f" {term} → {referent}") + + if chunk.context_before: + lines.append(f"\nThe sections BEFORE yours cover: {chunk.context_before}") + if chunk.context_after: + lines.append(f"The sections AFTER yours cover: {chunk.context_after}") + + lines.append(f"\nWrite these {len(chunk.sections)} section(s) in order:") + for i, s in enumerate(chunk.sections, 1): + lines.append(f"\n--- Section {i} ---") + lines.append(f"Heading: {s.heading}") + lines.append(f"Type: {s.type.value}") + lines.append(f"Depth: {s.depth.value}") + lines.append("Key points to cover:") + for point in s.key_points: + lines.append(f" - {point}") + + return "\n".join(lines) + + +# ── Helpers ─────────────────────────────────────────────────────────────────────────────────────── + + +def _build_jinja_env() -> Environment: + return Environment( + loader=FileSystemLoader(str(_TEMPLATES_DIR)), + autoescape=True, + trim_blocks=True, + lstrip_blocks=True, + ) + + +def _safe_filename(title: str) -> str: + slug = re.sub(r"[^\w\s-]", "", title.lower()) + slug = re.sub(r"[\s_-]+", "-", slug).strip("-") + return (slug[:60] or "document") + ".pdf" + + +# ── Agent ───────────────────────────────────────────────────────────────────────────────────────── + + +class PdfCreateAgent: + def __init__(self, runtime: AppRuntime) -> None: + self.runtime = runtime + self._jinja_env = _build_jinja_env() + + self._meta_planner: Agent[None, DocumentMeta] = Agent( + model=runtime.smart_model, + output_type=NativeOutput(DocumentMeta), + system_prompt=_META_PLANNER_SYSTEM_PROMPT, + model_settings={**runtime.smart_model_settings, "temperature": 0.1}, + ) + + self._sections_planner: Agent[None, DocumentSections] = Agent( + model=runtime.smart_model, + output_type=NativeOutput(DocumentSections), + system_prompt=_SECTIONS_PLANNER_SYSTEM_PROMPT, + model_settings={**runtime.smart_model_settings, "temperature": 0.1}, + ) + + self._writer: Agent[None, WrittenSections] = Agent( + model=runtime.smart_model, + output_type=NativeOutput(WrittenSections), + system_prompt=_WRITER_SYSTEM_PROMPT, + model_settings={**runtime.smart_model_settings, "temperature": 0.3}, + ) + + async def orchestrate(self, request: OrchestratorRequest) -> PdfCreateOrchestrateResponse: + history = format_conversation_history(request.conversation_history) + + # ── Phase 1: plan meta ───────────────────────────────────────────────── + logger.info("[pdf-create] phase 1/6: planning document meta") + meta_prompt = f"Conversation history:\n{history}\n\nUser request: {request.user_message}" + meta_result = await self._meta_planner.run(meta_prompt) + meta = meta_result.output + + if meta.cannot_do_reason: + logger.info("[pdf-create] cannot_do: %s", meta.cannot_do_reason) + return EditCannotDoResponse(reason=meta.cannot_do_reason) + + logger.info("[pdf-create] meta: title=%r tone=%r", meta.title, meta.tone_brief) + + # ── Phase 2: plan sections ───────────────────────────────────────────── + logger.info("[pdf-create] phase 2/6: planning sections") + sections_prompt = _build_sections_prompt(meta, request.user_message, history) + sections_result = await self._sections_planner.run(sections_prompt) + planned_sections = sections_result.output + + if not planned_sections.sections: + logger.info("[pdf-create] sections planner returned empty sections") + return EditCannotDoResponse(reason="No document sections could be planned from the request.") + + plan = DocumentPlan.assemble(meta, planned_sections) + + # ── Phase 3: chunk ───────────────────────────────────────────────────── + chunks = _make_chunks(plan.sections) + logger.info( + "[pdf-create] phase 3/6: chunked — sections=%d chunks=%d", + len(plan.sections), + len(chunks), + ) + + # ── Phase 4: write in parallel, bounded ──────────────────────────────── + logger.info("[pdf-create] phase 4/6: writing %d chunk(s) in parallel", len(chunks)) + total_chunks = len(chunks) + semaphore = asyncio.Semaphore(_MAX_PARALLEL_WRITERS) + written_chunks: list[WrittenSections] = await asyncio.gather( + *[self._write_chunk(plan, chunk, total_chunks, semaphore) for chunk in chunks] + ) + + # ── Phase 5: assemble in plan order (gather preserves insertion order) ── + all_sections: list[DocumentSection] = [] + for written in written_chunks: + all_sections.extend(written.sections) + + logger.info("[pdf-create] phase 5/6: assembled %d sections", len(all_sections)) + + doc = GeneratedDocument( + title=plan.title, + subtitle=plan.subtitle, + reference_number=plan.reference_number, + style=plan.style, + sections=all_sections, + ) + + # ── Phase 6: render ──────────────────────────────────────────────────── + logger.info("[pdf-create] phase 6/6: rendering HTML") + html = self._render(doc) + filename = _safe_filename(plan.title) + logger.info( + "[pdf-create] done — filename=%r html_bytes=%d", + filename, + len(html), + ) + + return EditPlanResponse( + summary=f"Created {plan.title}", + steps=[ + ToolOperationStep( + tool=AgentToolId.CREATE_PDF_FROM_HTML_AGENT, + parameters=CreatePdfFromHtmlAgentParams( + html_content=html, + filename=filename, + ), + ) + ], + ) + + async def _write_chunk( + self, plan: DocumentPlan, chunk: _Chunk, total_chunks: int, semaphore: asyncio.Semaphore + ) -> WrittenSections: + async with semaphore: + prompt = _build_writer_prompt(plan, chunk) + result = await self._writer.run(prompt) + logger.info( + "[pdf-create] chunk %d/%d wrote %d sections", + chunk.index + 1, + total_chunks, + len(result.output.sections), + ) + return result.output + + def _render(self, doc: GeneratedDocument) -> str: + template = self._jinja_env.get_template("document.html.jinja2") + return template.render(doc=doc) diff --git a/engine/src/stirling/agents/pdf_create/templates/document.html.jinja2 b/engine/src/stirling/agents/pdf_create/templates/document.html.jinja2 new file mode 100644 index 0000000000..b969458f5e --- /dev/null +++ b/engine/src/stirling/agents/pdf_create/templates/document.html.jinja2 @@ -0,0 +1,301 @@ + + + + + +{%- if doc.style %} + +{%- endif %} + + + +

+
{{ doc.title }}
+ {%- if doc.subtitle %} +
{{ doc.subtitle }}
+ {%- endif %} + {%- if doc.reference_number %} +
{{ doc.reference_number }}
+ {%- endif %} +
+ +{%- for section in doc.sections %} + +{%- if section.type == "text" %} +
+ {%- if section.heading %} +

{{ section.heading }}

+ {%- endif %} +
+ {%- for para in section.body.split('\n\n') %} +

{{ para | replace('\n', ' ') }}

+ {%- endfor %} +
+
+ +{%- elif section.type == "key_value" %} +
+ {%- if section.heading %} +

{{ section.heading }}

+ {%- endif %} + + + {%- for label, value in section.pairs %} + + + + + {%- endfor %} + +
{{ label }}{{ value }}
+
+ +{%- elif section.type == "line_items" %} +
+ {%- if section.heading %} +

{{ section.heading }}

+ {%- endif %} + + + + {%- for col in section.columns %} + + {%- endfor %} + + + + {%- for row in section.rows %} + + {%- for cell in row %} + + {%- endfor %} + + {%- endfor %} + {%- if section.total_row %} + + {%- for cell in section.total_row %} + + {%- endfor %} + + {%- endif %} + +
{{ col }}
{{ cell }}
{{ cell }}
+
+ +{%- elif section.type == "bullet_list" %} +
+ {%- if section.heading %} +

{{ section.heading }}

+ {%- endif %} +
    + {%- for item in section.items %} +
  • {{ item }}
  • + {%- endfor %} +
+
+ +{%- elif section.type == "signature" %} +
+ {%- if section.heading %} +

{{ section.heading }}

+ {%- endif %} +
+ {%- for signatory in section.signatories %} +
+
+
{{ signatory }}
+
+ {%- endfor %} +
+
+ +{%- endif %} +{%- endfor %} + + + diff --git a/engine/src/stirling/agents/pdf_to_markdown/__init__.py b/engine/src/stirling/agents/pdf_to_markdown/__init__.py deleted file mode 100644 index d35ae05c7c..0000000000 --- a/engine/src/stirling/agents/pdf_to_markdown/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -from .agent import PdfToMarkdownAgent - -__all__ = ["PdfToMarkdownAgent"] diff --git a/engine/src/stirling/agents/pdf_to_markdown/agent.py b/engine/src/stirling/agents/pdf_to_markdown/agent.py deleted file mode 100644 index 8c0d7d8ee5..0000000000 --- a/engine/src/stirling/agents/pdf_to_markdown/agent.py +++ /dev/null @@ -1,435 +0,0 @@ -"""PDF to Markdown Agent. - -Converts a parsed PDF document into a single clean Markdown document, preserving -headings, paragraphs, and tables in reading order. -""" - -from __future__ import annotations - -import asyncio -import logging -import re -import time - -from pydantic import BaseModel, Field -from pydantic_ai import Agent -from pydantic_ai.output import NativeOutput - -from stirling.contracts import ( - EditCannotDoResponse, - GenerateFileResponse, - NeedContentFileRequest, - NeedContentResponse, - OrchestratorRequest, - PdfContentType, - SupportedCapability, - format_conversation_history, -) -from stirling.contracts.pdf_to_markdown import ( - PageLayout, - PageLayoutArtifact, - PdfToMarkdownCannotDoResponse, - PdfToMarkdownOrchestrateResponse, - PdfToMarkdownRequest, - PdfToMarkdownResponse, - PdfToMarkdownSuccessResponse, -) -from stirling.services import AppRuntime - -logger = logging.getLogger(__name__) - - -# Warn when output tokens are close to the typical model output limit (~8192 for most -# configurations). The actual limit is model-specific; this threshold catches likely truncation. -_OUTPUT_TOKEN_TRUNCATION_THRESHOLD = 7500 - -# Chunking limits — keep each LLM call to a manageable payload size. -# Fragment count is the primary driver of JSON payload size (each fragment carries x/y/width/ -# fontSize/bold metadata beyond its text). Page cap prevents low-text pages accumulating. -_MAX_CHUNK_FRAGMENTS = 1_000 -_MAX_CHUNK_PAGES = 10 - -# Max concurrent LLM calls — limits API rate pressure on large documents. -_MAX_PARALLEL_CHUNKS = 3 - -# ── LLM output model ──────────────────────────────────────────────────────────────────────────── - - -class _ReconstructionOutput(BaseModel): - markdown: str = Field(description="Full document reconstructed as clean Markdown.") - - -# ── Agent ──────────────────────────────────────────────────────────────────────────────────────── - - -class PdfToMarkdownAgent: - def __init__(self, runtime: AppRuntime) -> None: - self.runtime = runtime - self._sem = asyncio.Semaphore(_MAX_PARALLEL_CHUNKS) - self._reconstruct_agent = Agent( - model=runtime.smart_model, - output_type=NativeOutput(_ReconstructionOutput), - system_prompt=( - "You reconstruct PDF pages into clean Markdown from spatial fragment data.\n" - "Input: PAGE LAYOUT — per-fragment x/y/font data for structural analysis.\n\n" - "COLUMN DETECTION (for tables in page_layout):\n" - "- Look at the x-positions of fragments across 3+ consecutive lines.\n" - "- If fragments cluster at the same x-positions across multiple lines, those are table columns.\n" - "- Each distinct x-cluster is one column." - " Name them from the header row (the first line in the cluster).\n" - "- Do NOT merge values from different x-columns into one cell.\n\n" - "ROW DETECTION:\n" - "- Each unique y-coordinate (or group within 3pt) is one table row.\n" - "- Every line of layout data is its own row — do not merge rows.\n" - "- If a column has no fragment on a given y-row, that cell is empty.\n\n" - "TABLE RENDERING:\n" - "- Render as: | col1 | col2 | col3 |\n" - " | --- | --- | --- |\n" - " | val | val | val |\n" - "- One source row = one table row. Never collapse multiple rows into one.\n" - "- Preserve numeric values exactly (no rounding, no formatting changes).\n" - "- Bold cells: wrap with ** in the Markdown cell.\n" - "- CRITICAL: the separator row `| --- | --- |` appears EXACTLY ONCE per table, immediately\n" - " after the header row. NEVER put `| --- |` after a data row or between data rows.\n" - " NEVER put a blank line inside a table. All rows (header + data) must be consecutive.\n" - "- Do NOT produce a header-only table followed by a second table with the data rows.\n" - " One logical table = one markdown table block, with header, one separator, then all data.\n\n" - "GROUP HEADERS (label-only rows inside a table):\n" - "- A row is a group header when: the first column has text AND every numeric column is empty.\n" - "- Do NOT render group headers as table rows with empty cells.\n" - "- Break the table, emit the label as **bold text** on its own line," - " then start a new table for the rows that follow.\n" - "- Example labels: 'Policy functions', 'Non-current assets'.\n\n" - "TOTAL AND SUBTOTAL ROWS:\n" - "- Detect rows whose first cell contains (case-insensitive):" - " total, subtotal, surplus, balance, net, sum.\n" - "- These rows have numeric content — they are NOT group headers.\n" - "- Render the entire row in bold: | **Total income** | **1,234** | **5,678** |\n" - "- Keep total rows attached to the group they summarise.\n\n" - "MULTI-LEVEL TABLES (year or period as a row label):\n" - "- Detect when a row contains only a single label (a year like '2010' or period like 'Q1 2023')" - " with no numeric content, followed by repeated metric rows.\n" - "- Do NOT render the year as a table row.\n" - "- Normalise: add 'Year' as the first column, 'Metric' as the second," - " and repeat the year value on each metric row.\n\n" - "PROSE REGIONS:\n" - "- Lines where x-positions vary across lines (not repeating columns) are prose.\n" - "- Merge lines at the same x-level into paragraphs. Separate indented lines.\n\n" - "HEADINGS:\n" - "- A line is a heading when it is bold OR font_size ≥2pt above body.\n" - " CRITICAL EXCEPTION: a bold fragment is a TABLE HEADER CELL, not a document heading, when\n" - " the same y-row in page_layout contains other fragments at different x-positions.\n" - " Only classify a bold line as a document heading when it is the SOLE fragment on its y-row.\n" - " Example: 'Non-current assets' at y=120 with '2010'@x=350, '2009'@x=420, '2008'@x=490\n" - " → this is a table header row, NOT a heading. Render it as the first cell of the table.\n" - "- Use ## for section headings, ### for sub-headings. Use # only for the document title.\n\n" - "ORDERING:\n" - "- Process content top-to-bottom as it appears on the page.\n" - "- Interleave prose blocks and table blocks in page order.\n" - "- Do not move text that appears before a table to after it, or vice versa.\n\n" - "FIDELITY:\n" - "- Do NOT invent, summarise, or omit any content.\n" - "- Do NOT add commentary, metadata, or JSON — output Markdown only." - ), - model_settings={ - **runtime.smart_model_settings, - "temperature": 0.0, - "max_tokens": _OUTPUT_TOKEN_TRUNCATION_THRESHOLD, - }, - ) - - async def orchestrate(self, request: OrchestratorRequest) -> PdfToMarkdownOrchestrateResponse: - """Entry point for the orchestrator delegate. - - First turn: requests PAGE_LAYOUT extraction from Java via NeedContentResponse. - Resume turn: runs the LLM reconstruction and returns a write-file plan step. - """ - layout_artifact = next( - (a for a in request.artifacts if isinstance(a, PageLayoutArtifact)), - None, - ) - if layout_artifact is None: - return NeedContentResponse( - resume_with=SupportedCapability.PDF_TO_MARKDOWN, - reason="Page layout data is required to reconstruct the document.", - files=[ - NeedContentFileRequest(file=f, content_types=[PdfContentType.PAGE_LAYOUT]) for f in request.files - ], - max_pages=self.runtime.settings.max_pages, - max_characters=self.runtime.settings.max_characters, - ) - - page_layout = [page for entry in layout_artifact.files for page in entry.pages] - file_names = [f.name for f in request.files] - result = await self.handle( - PdfToMarkdownRequest( - user_message=request.user_message, - file_names=file_names, - conversation_history=request.conversation_history, - page_layout=page_layout, - ) - ) - if isinstance(result, PdfToMarkdownCannotDoResponse): - return EditCannotDoResponse(reason=result.reason) - - base = file_names[0].rsplit(".", 1)[0] if file_names else "document" - return GenerateFileResponse( - content=result.markdown, - filename=f"{base}-reconstruction.md", - summary="Reconstructed the document as a Markdown file.", - ) - - async def handle(self, request: PdfToMarkdownRequest) -> PdfToMarkdownResponse: - total_fragments = sum(len(line.fragments) for page in request.page_layout for line in page.lines) - logger.info( - "[pdf-to-markdown] received layout-pages=%d fragments=%d", - len(request.page_layout), - total_fragments, - ) - - if not request.page_layout: - logger.warning("[pdf-to-markdown] no content extracted from document; returning cannot_do") - return PdfToMarkdownCannotDoResponse( - reason=( - "No content was extracted from the document. " - "The file may be a scanned image PDF with no readable text. " - "Try running OCR on the document first." - ) - ) - - chunks = _build_page_chunks(request.page_layout) - logger.info("[pdf-to-markdown] chunks=%d (max %d in parallel)", len(chunks), _MAX_PARALLEL_CHUNKS) - - if len(chunks) == 1: - return await self._reconstruct_chunk(request, chunks[0], chunk_num=1, total_chunks=1) - - total = len(chunks) - results = await asyncio.gather( - *( - self._reconstruct_chunk(request, chunk, chunk_num=i + 1, total_chunks=total) - for i, chunk in enumerate(chunks) - ) - ) - - markdown_parts: list[str] = [] - for result in results: - if isinstance(result, PdfToMarkdownSuccessResponse) and result.markdown: - markdown_parts.append(result.markdown) - elif isinstance(result, PdfToMarkdownCannotDoResponse): - logger.warning("[pdf-to-markdown] chunk dropped: %s", result.reason) - - if not markdown_parts: - return PdfToMarkdownCannotDoResponse(reason="The document could not be reconstructed. All chunks failed.") - - logger.info("[pdf-to-markdown] assembly: %d/%d chunks produced output", len(markdown_parts), len(chunks)) - return PdfToMarkdownSuccessResponse(markdown="\n\n".join(markdown_parts)) - - async def _reconstruct_chunk( - self, - request: PdfToMarkdownRequest, - pages: list[PageLayout], - chunk_num: int, - total_chunks: int, - ) -> PdfToMarkdownResponse: - chunk_request = PdfToMarkdownRequest( - user_message=request.user_message, - file_names=request.file_names, - conversation_history=request.conversation_history, - page_layout=pages, - ) - try: - async with self._sem: - return await self._reconstruct_document(chunk_request, chunk_num, total_chunks) - except Exception as e: - logger.error("[pdf-to-markdown] chunk %d/%d failed: %s", chunk_num, total_chunks, e, exc_info=True) - return PdfToMarkdownCannotDoResponse( - reason="The document could not be reconstructed. The AI model failed to process it." - ) - - async def _reconstruct_document( - self, request: PdfToMarkdownRequest, chunk_num: int = 1, total_chunks: int = 1 - ) -> PdfToMarkdownSuccessResponse: - content = _build_reconstruction_prompt(request) - logger.info("[timing] chunk %d/%d llm-call prompt-chars=%d", chunk_num, total_chunks, len(content)) - t0 = time.monotonic() - result = await self._reconstruct_agent.run([content]) - llm_ms = int((time.monotonic() - t0) * 1000) - output: _ReconstructionOutput = result.output - usage = result.usage() - logger.info( - "[timing] chunk %d/%d llm-done ms=%d input-tokens=%s output-tokens=%s markdown-chars=%d", - chunk_num, - total_chunks, - llm_ms, - usage.input_tokens, - usage.output_tokens, - len(output.markdown), - ) - if usage.output_tokens and usage.output_tokens >= _OUTPUT_TOKEN_TRUNCATION_THRESHOLD: - logger.warning( - "[timing] chunk %d/%d output likely truncated (output-tokens=%d)", - chunk_num, - total_chunks, - usage.output_tokens, - ) - markdown = _remove_extra_separators(_fix_markdown_tables(_merge_orphaned_table_rows(output.markdown))) - return PdfToMarkdownSuccessResponse(markdown=markdown) - - -# ── Chunking ──────────────────────────────────────────────────────────────────────────────────── - - -def _build_page_chunks(pages: list[PageLayout]) -> list[list[PageLayout]]: - chunks: list[list[PageLayout]] = [] - current: list[PageLayout] = [] - current_fragments = 0 - for page in pages: - page_fragments = sum(len(line.fragments) for line in page.lines) - fragment_full = current and current_fragments + page_fragments > _MAX_CHUNK_FRAGMENTS - page_full = len(current) >= _MAX_CHUNK_PAGES - if fragment_full or page_full: - chunks.append(current) - current = [] - current_fragments = 0 - current.append(page) - current_fragments += page_fragments - if current: - chunks.append(current) - return chunks - - -# ── Prompt builders (module-level, no state) ──────────────────────────────────────────────────── - - -def _build_reconstruction_prompt(request: PdfToMarkdownRequest) -> str: - history = format_conversation_history(request.conversation_history) - file_names = ", ".join(request.file_names) if request.file_names else "Unknown files" - layout_section = _format_layout(request.page_layout) - - return ( - f"Files: {file_names}\n\n" - f"User request: {request.user_message}\n\n" - f"Conversation history:\n{history}\n\n" - "PAGE LAYOUT (structural source — x/y fragment positions):\n" - "Each line is: y=NNN | text@(x,y) fs=N text@(x,y) fs=N ...\n" - "- y=NNN is the vertical position (row). Lines close in y are the same visual row.\n" - "- x=NNN is the horizontal position (column). Consistent x across rows = a column.\n" - "- fs=N is font size. Larger = likely a heading.\n" - "- **bold** markers indicate bold text.\n\n" - f"{layout_section}" - ) - - -# ── LLM output post-processing ────────────────────────────────────────────────────────────────── - - -def _fix_markdown_tables(markdown: str) -> str: - """Remove blank lines between table rows produced by the LLM.""" - lines = markdown.split("\n") - result: list[str] = [] - i = 0 - while i < len(lines): - result.append(lines[i]) - if lines[i].strip().startswith("|"): - j = i + 1 - while j < len(lines) and lines[j].strip() == "": - j += 1 - if j < len(lines) and lines[j].strip().startswith("|"): - i = j - continue - i += 1 - return "\n".join(result) - - -_SEP_CELL = re.compile(r"^:?-+:?$") - - -def _is_sep_row(line: str) -> bool: - """Return True when a pipe row is a Markdown table separator (| --- | --- |).""" - stripped = line.strip() - if not stripped.startswith("|"): - return False - cells = [c.strip() for c in stripped.split("|") if c.strip()] - return bool(cells) and all(_SEP_CELL.match(c) for c in cells) - - -def _merge_orphaned_table_rows(markdown: str) -> str: - """Merge pipe-row blocks that lack a separator into the preceding table. - - When the LLM incorrectly breaks a table (e.g. on a false group-header), it emits - orphaned pipe rows with no header or separator. These are invalid markdown and get - merged back into the preceding table, discarding the intervening non-table content. - """ - lines = markdown.split("\n") - - segments: list[tuple[str, list[str]]] = [] - i = 0 - while i < len(lines): - if lines[i].strip().startswith("|"): - block: list[str] = [] - while i < len(lines) and lines[i].strip().startswith("|"): - block.append(lines[i]) - i += 1 - has_sep = any(_is_sep_row(row) for row in block) - segments.append(("table" if has_sep else "orphan", block)) - else: - block = [] - while i < len(lines) and not lines[i].strip().startswith("|"): - block.append(lines[i]) - i += 1 - segments.append(("prose", block)) - - result: list[tuple[str, list[str]]] = [] - last_table_idx: int | None = None - for seg_type, seg_lines in segments: - if seg_type == "orphan": - if last_table_idx is not None: - result = result[: last_table_idx + 1] - result[-1] = ("table", result[-1][1] + seg_lines) - else: - result.append((seg_type, seg_lines)) - else: - if seg_type == "table": - last_table_idx = len(result) - result.append((seg_type, seg_lines)) - - return "\n".join(line for _, seg_lines in result for line in seg_lines) - - -def _remove_extra_separators(markdown: str) -> str: - """Within each contiguous table block, keep only the first separator row.""" - lines = markdown.split("\n") - result: list[str] = [] - seen_sep = False - - for line in lines: - if not line.strip().startswith("|"): - seen_sep = False - result.append(line) - continue - if _is_sep_row(line): - if seen_sep: - continue - seen_sep = True - result.append(line) - - return "\n".join(result) - - -# ── Formatting helpers (module-level, no state) ────────────────────────────────────────────────── - - -def _format_layout(pages: list[PageLayout]) -> str: - if not pages: - return "None" - parts: list[str] = [] - for page in pages: - line_strs: list[str] = [] - for line in page.lines: - frags = " ".join( - f"{'**' if f.bold else ''}{f.text}{'**' if f.bold else ''}@({f.x:.0f},{f.y:.0f}) fs={f.font_size:.0f}" - for f in line.fragments - ) - line_strs.append(f"y={line.y:.0f} | {frags}") - parts.append(f"--- Page {page.page_number} ---\n" + "\n".join(line_strs)) - return "\n\n".join(parts) diff --git a/engine/src/stirling/api/agent_capabilities.py b/engine/src/stirling/api/agent_capabilities.py new file mode 100644 index 0000000000..2d552ea556 --- /dev/null +++ b/engine/src/stirling/api/agent_capabilities.py @@ -0,0 +1,160 @@ +""" +Curated registry of agent capabilities the MCP server (Java side) is allowed to publish. + +Internal sub-agents (currently only ``ExecutionPlanningAgent`` - it lives behind the orchestrator +and has no end-user-facing API surface) are intentionally absent. The handoff spec calls for +"user-facing" capabilities only; revisit this list when adding a new agent and ask whether MCP +clients should be able to invoke it directly. + +The Java side pulls ``/api/v1/agents/capabilities`` once at boot and again every few minutes; the +manifest is the authoritative source for the ``stirling_ai`` MCP tool's operation enum. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any + +from pydantic import BaseModel + +from stirling.contracts import ( + AgentDraftRequest, + AgentExecutionRequest, + AgentRevisionRequest, + Evidence, + FolioManifest, + PdfCommentRequest, + PdfEditRequest, + PdfQuestionRequest, +) + + +@dataclass(frozen=True) +class AgentCapability: + """One row in the curated manifest. + + Attributes: + id: stable capability identifier (used as the operation enum value in + ``stirling_ai``). Avoid renaming - clients persist these. + description: one-line human-friendly summary shown inside MCP tool descriptions. + input_model: Pydantic class whose JSON Schema becomes the capability's + ``input_schema``. Auto-derived; do not hand-write schemas. + mode: ``"sync"`` if the capability returns content inline, ``"async"`` if it returns a + plan that Java executes via the job pipeline. + required_scope: coarse OAuth scope. ``mcp.tools.read`` for pure-read capabilities + (Q&A, audits) and ``mcp.tools.write`` for anything that yields a plan / file. + route: HTTP path Java POSTs to when invoking this capability. When a capability does + not have a stable per-agent route yet, use the generic invoke fallback at + ``/api/v1/agents/invoke/{id}``. + """ + + id: str + description: str + input_model: type[BaseModel] + mode: str + required_scope: str + route: str + + +EXPOSED_CAPABILITIES: list[AgentCapability] = [ + AgentCapability( + id="pdf-question-answer", + description="Answer a natural-language question about a PDF document.", + input_model=PdfQuestionRequest, + mode="sync", + required_scope="mcp.tools.read", + route="/api/v1/pdf-question", + ), + AgentCapability( + id="pdf-edit-plan", + description=( + "Produce an edit plan (a structured sequence of PDF operations) from a" + " natural-language edit request. The plan is executed by Java through the job" + " pipeline; this capability does not modify files itself." + ), + input_model=PdfEditRequest, + mode="async", + required_scope="mcp.tools.write", + route="/api/v1/pdf-edit", + ), + AgentCapability( + id="agent-draft", + description=( + "Draft a structured agent specification from a free-text description of the task the user wants automated." + ), + input_model=AgentDraftRequest, + mode="sync", + required_scope="mcp.tools.read", + route="/api/v1/ai/agents/draft", + ), + AgentCapability( + id="agent-revise", + description=("Revise an existing draft agent specification based on user feedback or constraint changes."), + input_model=AgentRevisionRequest, + mode="sync", + required_scope="mcp.tools.read", + route="/api/v1/ai/agents/revise", + ), + AgentCapability( + id="math-audit-examine", + description=( + "Examine a folio manifest of financial / numeric documents and surface the" + " evidence that needs to be checked for arithmetic consistency." + ), + input_model=FolioManifest, + mode="sync", + required_scope="mcp.tools.read", + route="/api/v1/ai/math-auditor-agent/examine", + ), + AgentCapability( + id="math-audit-deliberate", + description=( + "Render a deliberated verdict on a single piece of evidence the examine step" + " surfaced (does the arithmetic check out, with what caveats)." + ), + input_model=Evidence, + mode="sync", + required_scope="mcp.tools.read", + route="/api/v1/ai/math-auditor-agent/deliberate", + ), + AgentCapability( + id="pdf-comment-generate", + description="Generate inline review comments for a PDF document.", + input_model=PdfCommentRequest, + mode="sync", + required_scope="mcp.tools.read", + route="/api/v1/pdf-comment/generate", + ), + AgentCapability( + id="agent-next-action", + description=( + "Decide the next execution step for an in-progress agent workflow. Returns a" + " ToolCall, Completed, or CannotContinue action." + ), + input_model=AgentExecutionRequest, + mode="sync", + required_scope="mcp.tools.read", + route="/api/v1/agents/next-action", + ), +] + + +def manifest_payload() -> dict[str, Any]: + """Serialize the curated registry to the wire shape consumed by Java. + + Schema is derived from ``input_model.model_json_schema()`` so we never hand-write JSON + Schema - the Pydantic model is the single source of truth. + """ + items: list[dict[str, Any]] = [] + for cap in EXPOSED_CAPABILITIES: + items.append( + { + "id": cap.id, + "description": cap.description, + "input_schema": cap.input_model.model_json_schema(), + "mode": cap.mode, + "required_scope": cap.required_scope, + "route": cap.route, + } + ) + return {"version": 1, "capabilities": items} diff --git a/engine/src/stirling/api/app.py b/engine/src/stirling/api/app.py index 1cf4b76c3e..7e6272c28b 100644 --- a/engine/src/stirling/api/app.py +++ b/engine/src/stirling/api/app.py @@ -18,8 +18,11 @@ from stirling.agents import ( ) from stirling.agents.ledger import MathAuditorAgent from stirling.agents.pdf_comment import PdfCommentAgent +from stirling.api.dependencies import enforce_required_user_id +from stirling.api.engine_auth import EngineSharedSecretMiddleware from stirling.api.middleware import UserIdMiddleware from stirling.api.routes import ( + agent_capabilities_router, agent_draft_router, document_router, execution_router, @@ -115,14 +118,19 @@ async def lifespan(fast_api: FastAPI): app = FastAPI(title="Stirling AI Engine", lifespan=lifespan, version="0.1.0") app.add_middleware(UserIdMiddleware) -app.include_router(orchestrator_router) -app.include_router(pdf_edit_router) -app.include_router(pdf_question_router) -app.include_router(agent_draft_router) -app.include_router(execution_router) -app.include_router(document_router) -app.include_router(ledger_router) -app.include_router(pdf_comments_router) +app.add_middleware(EngineSharedSecretMiddleware) +# Every router gets the same configurable identity gate; /health stays open +# for liveness probes. See enforce_required_user_id for the policy. +_user_gate = [Depends(enforce_required_user_id)] +app.include_router(orchestrator_router, dependencies=_user_gate) +app.include_router(pdf_edit_router, dependencies=_user_gate) +app.include_router(pdf_question_router, dependencies=_user_gate) +app.include_router(agent_draft_router, dependencies=_user_gate) +app.include_router(execution_router, dependencies=_user_gate) +app.include_router(document_router, dependencies=_user_gate) +app.include_router(ledger_router, dependencies=_user_gate) +app.include_router(pdf_comments_router, dependencies=_user_gate) +app.include_router(agent_capabilities_router, dependencies=_user_gate) @app.get("/health", response_model=HealthResponse) diff --git a/engine/src/stirling/api/dependencies.py b/engine/src/stirling/api/dependencies.py index cae20a8ae1..780e97c9d7 100644 --- a/engine/src/stirling/api/dependencies.py +++ b/engine/src/stirling/api/dependencies.py @@ -1,6 +1,8 @@ from __future__ import annotations -from fastapi import HTTPException, Request, status +from typing import Annotated + +from fastapi import Depends, HTTPException, Request, status from stirling.agents import ( ExecutionPlanningAgent, @@ -11,6 +13,7 @@ from stirling.agents import ( ) from stirling.agents.ledger import MathAuditorAgent from stirling.agents.pdf_comment import PdfCommentAgent +from stirling.config import AppSettings, load_settings from stirling.documents import DocumentService from stirling.models import UserId from stirling.services import AppRuntime, current_user_id @@ -67,3 +70,16 @@ def require_user_id() -> UserId: detail="X-User-Id header is required", ) return user_id + + +def enforce_required_user_id( + settings: Annotated[AppSettings, Depends(load_settings)], +) -> None: + """Router-level boundary gate, applied uniformly to every router.""" + if not settings.require_user_id: + return + if current_user_id.get() is None: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail="X-User-Id header is required", + ) diff --git a/engine/src/stirling/api/engine_auth.py b/engine/src/stirling/api/engine_auth.py new file mode 100644 index 0000000000..c30b9f7429 --- /dev/null +++ b/engine/src/stirling/api/engine_auth.py @@ -0,0 +1,99 @@ +""" +Shared-secret middleware that locks the engine to the trusted Java backend. + +Config (resolved via :class:`stirling.config.AppSettings`/pydantic-settings): +``STIRLING_ENGINE_SHARED_SECRET`` - non-public routes need ``X-Engine-Auth`` or 401. +``STIRLING_ENGINE_REQUIRE_AUTH`` - fail closed with 503 when truthy and no secret is set. +""" + +from __future__ import annotations + +import hmac +import logging +from collections.abc import Iterable + +from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint +from starlette.requests import Request +from starlette.responses import JSONResponse, Response +from starlette.types import ASGIApp + +from stirling.config import load_settings + +logger = logging.getLogger(__name__) + +_HEADER = "X-Engine-Auth" + +# Public paths (liveness + docs); everything else needs the secret when configured. +_PUBLIC_PREFIXES: tuple[str, ...] = ( + "/health", + "/docs", + "/redoc", + "/openapi.json", +) + + +class EngineSharedSecretMiddleware(BaseHTTPMiddleware): + """Reject non-public requests lacking the shared secret. + + Non-public path: secret set -> require matching X-Engine-Auth (else 401); else require flag + truthy -> 503 (fail-closed); else allow through. + + Secret/require values come from :class:`stirling.config.AppSettings` by default; tests can + pass them explicitly to avoid touching the lru-cached settings. + """ + + def __init__( + self, + app: ASGIApp, + public_prefixes: Iterable[str] = _PUBLIC_PREFIXES, + *, + secret: str | None = None, + require: bool | None = None, + ) -> None: + super().__init__(app) + self._public_prefixes = tuple(public_prefixes) + if secret is None or require is None: + settings = load_settings() + if secret is None: + secret = settings.engine_shared_secret + if require is None: + require = settings.engine_require_auth + self._secret = secret or "" + self._require = bool(require) + if self._secret: + logger.info( + "Engine shared-secret enforcement ENABLED: non-public routes require a valid %s" + " header (constant-time compared).", + _HEADER, + ) + elif self._require: + logger.error( + "STIRLING_ENGINE_REQUIRE_AUTH is enabled but STIRLING_ENGINE_SHARED_SECRET is not" + " set - the engine will REFUSE every non-public request (HTTP 503, fail-closed)" + " until a shared secret is configured.", + ) + else: + logger.warning( + "STIRLING_ENGINE_SHARED_SECRET not set - engine shared-secret enforcement is" + " DISABLED. The AI and document routes then trust the caller-supplied X-User-Id" + " header alone. Set this secret (and STIRLING_ENGINE_REQUIRE_AUTH=true) in any" + " deployment that exposes the engine beyond localhost.", + ) + + def _is_public(self, path: str) -> bool: + return any(path == p or path.startswith(p + "/") for p in self._public_prefixes) + + async def dispatch(self, request: Request, call_next: RequestResponseEndpoint) -> Response: + if not self._is_public(request.url.path): + if self._secret: + offered = request.headers.get(_HEADER) or "" + # Constant-time compare to avoid timing leaks. + if not hmac.compare_digest(offered, self._secret): + return JSONResponse({"detail": "Missing or invalid X-Engine-Auth header."}, status_code=401) + elif self._require: + # Fail closed: require flag set but no secret configured. + return JSONResponse( + {"detail": ("Engine authentication is required but no shared secret is configured.")}, + status_code=503, + ) + return await call_next(request) diff --git a/engine/src/stirling/api/routes/__init__.py b/engine/src/stirling/api/routes/__init__.py index 37572c5e76..1c1e4e1d23 100644 --- a/engine/src/stirling/api/routes/__init__.py +++ b/engine/src/stirling/api/routes/__init__.py @@ -1,3 +1,4 @@ +from .agent_capabilities import router as agent_capabilities_router from .agent_drafts import router as agent_draft_router from .documents import router as document_router from .execution import router as execution_router @@ -8,6 +9,7 @@ from .pdf_edit import router as pdf_edit_router from .pdf_questions import router as pdf_question_router __all__ = [ + "agent_capabilities_router", "agent_draft_router", "document_router", "execution_router", diff --git a/engine/src/stirling/api/routes/agent_capabilities.py b/engine/src/stirling/api/routes/agent_capabilities.py new file mode 100644 index 0000000000..75f1d0f823 --- /dev/null +++ b/engine/src/stirling/api/routes/agent_capabilities.py @@ -0,0 +1,22 @@ +"""GET ``/api/v1/agents/capabilities`` - the manifest the Java MCP server pulls at boot.""" + +from __future__ import annotations + +from typing import Any + +from fastapi import APIRouter + +from stirling.api.agent_capabilities import manifest_payload + +router = APIRouter(prefix="/api/v1/agents", tags=["agents"]) + + +@router.get("/capabilities") +def get_capabilities() -> dict[str, Any]: + """Return the curated agent capabilities manifest. + + Gated by ``EngineSharedSecretMiddleware`` when the ``STIRLING_ENGINE_SHARED_SECRET`` env var + is configured. In dev/local mode (no secret set), the endpoint is open - the engine binds to + localhost only by default, so this is acceptable while iterating. + """ + return manifest_payload() diff --git a/engine/src/stirling/config/__init__.py b/engine/src/stirling/config/__init__.py index c5044a2148..f038d02e55 100644 --- a/engine/src/stirling/config/__init__.py +++ b/engine/src/stirling/config/__init__.py @@ -1,10 +1,10 @@ """Configuration models and loaders for the Stirling AI service.""" -from .settings import ENGINE_ROOT, AppSettings, RagBackend, load_settings +from .settings import ENGINE_ROOT, AppSettings, DocumentsBackend, load_settings __all__ = [ "ENGINE_ROOT", "AppSettings", - "RagBackend", + "DocumentsBackend", "load_settings", ] diff --git a/engine/src/stirling/config/settings.py b/engine/src/stirling/config/settings.py index bd0499e6b3..1293cc371f 100644 --- a/engine/src/stirling/config/settings.py +++ b/engine/src/stirling/config/settings.py @@ -15,7 +15,7 @@ ENV_FILE = ENGINE_ROOT / ".env" ENV_LOCAL_FILE = ENGINE_ROOT / ".env.local" -class RagBackend(StrEnum): +class DocumentsBackend(StrEnum): SQLITE = "sqlite" PGVECTOR = "pgvector" @@ -27,12 +27,22 @@ class AppSettings(BaseSettings): fast_model_name: str = Field(validation_alias="STIRLING_FAST_MODEL") smart_model_max_tokens: int = Field(validation_alias="STIRLING_SMART_MODEL_MAX_TOKENS") fast_model_max_tokens: int = Field(validation_alias="STIRLING_FAST_MODEL_MAX_TOKENS") + # Process-wide ceiling on concurrent model API calls, shared by both model + # tiers. Per-request fan-outs (chunked reasoner, contradiction detection) + # carry their own per-request caps, but those multiply under concurrent + # traffic; this is the global backstop. + model_max_concurrency: int = Field(validation_alias="STIRLING_MODEL_MAX_CONCURRENCY") - # RAG settings — always on; the backend picks between embedded sqlite-vec and external pgvector. - rag_backend: RagBackend = Field(validation_alias="STIRLING_RAG_BACKEND") + # Document store: the one database holding vector chunks, ordered page + # text, and ACL rows - embedded sqlite-vec or external pgvector. + documents_backend: DocumentsBackend = Field(validation_alias="STIRLING_DOCUMENTS_BACKEND") + documents_sqlite_path: Path = Field(validation_alias="STIRLING_DOCUMENTS_SQLITE_PATH") + documents_pgvector_dsn: str = Field(validation_alias="STIRLING_DOCUMENTS_PGVECTOR_DSN") + documents_pgvector_pool_min_size: int = Field(validation_alias="STIRLING_DOCUMENTS_PGVECTOR_POOL_MIN_SIZE") + documents_pgvector_pool_max_size: int = Field(validation_alias="STIRLING_DOCUMENTS_PGVECTOR_POOL_MAX_SIZE") + + # RAG settings - always on. rag_embedding_model: str = Field(validation_alias="STIRLING_RAG_EMBEDDING_MODEL") - rag_store_path: Path = Field(validation_alias="STIRLING_RAG_STORE_PATH") - rag_pgvector_dsn: str = Field(validation_alias="STIRLING_RAG_PGVECTOR_DSN") rag_chunk_size: int = Field(validation_alias="STIRLING_RAG_CHUNK_SIZE") rag_chunk_overlap: int = Field(validation_alias="STIRLING_RAG_CHUNK_OVERLAP") rag_default_top_k: int = Field(validation_alias="STIRLING_RAG_TOP_K") @@ -88,6 +98,12 @@ class AppSettings(BaseSettings): max_pages: int = Field(validation_alias="STIRLING_MAX_PAGES") max_characters: int = Field(validation_alias="STIRLING_MAX_CHARACTERS") + # When true, API routes reject requests that lack an X-User-Id header at + # the boundary. Self-hosted deployments with security disabled have no + # user identity and leave this off; multi-tenant deployments turn it on so + # user-scoped work is never processed without a tenant attached. + require_user_id: bool = Field(validation_alias="STIRLING_REQUIRE_USER_ID") + log_level: str = Field(default="INFO", validation_alias="STIRLING_LOG_LEVEL") log_file: str = Field(default="", validation_alias="STIRLING_LOG_FILE") # When true, raises httpx + httpcore logger levels so every outgoing @@ -102,6 +118,11 @@ class AppSettings(BaseSettings): posthog_api_key: str = Field(validation_alias="STIRLING_POSTHOG_API_KEY") posthog_host: str = Field(validation_alias="STIRLING_POSTHOG_HOST") + # Shared secret enforced by EngineSharedSecretMiddleware. Empty disables enforcement + # unless engine_require_auth is set, in which case the engine fails closed (503). + engine_shared_secret: str = Field(default="", validation_alias="STIRLING_ENGINE_SHARED_SECRET") + engine_require_auth: bool = Field(default=False, validation_alias="STIRLING_ENGINE_REQUIRE_AUTH") + def _configure_logging(level_name: str, log_file: str, http_debug: bool) -> None: """Configure the ``stirling`` logger hierarchy.""" diff --git a/engine/src/stirling/contracts/__init__.py b/engine/src/stirling/contracts/__init__.py index 696749d7d7..6c8d99d120 100644 --- a/engine/src/stirling/contracts/__init__.py +++ b/engine/src/stirling/contracts/__init__.py @@ -13,6 +13,7 @@ from .common import ( AiFile, ArtifactKind, ConversationMessage, + ConvertMarkdownResponse, ExtractedFileText, GenerateFileResponse, MathAuditorToolReportArtifact, @@ -79,6 +80,15 @@ from .pdf_comments import ( PdfCommentResponse, TextChunk, ) +from .pdf_create import ( + DocumentMeta, + DocumentSections, + PdfCreateCannotDoResponse, + PdfCreateOrchestrateResponse, + PdfCreateRequest, + PdfCreateResponse, + PdfCreateSuccessResponse, +) from .pdf_edit import ( EditCannotDoResponse, EditClarificationRequest, @@ -96,17 +106,6 @@ from .pdf_questions import ( PdfQuestionTerminalResponse, ) from .pdf_review import PdfReviewOrchestrateResponse -from .pdf_to_markdown import ( - LayoutFragment, - LayoutLine, - PageLayout, - PageLayoutArtifact, - PageLayoutFileEntry, - PdfToMarkdownCannotDoResponse, - PdfToMarkdownOrchestrateResponse, - PdfToMarkdownRequest, - PdfToMarkdownResponse, -) from .progress import ( ProgressEvent, WholeDocCompressionRound, @@ -139,11 +138,9 @@ __all__ = [ "ConversationMessage", "DeleteDocumentResponse", "PurgeOwnerResponse", - "PdfToMarkdownCannotDoResponse", - "PdfToMarkdownOrchestrateResponse", - "PdfToMarkdownRequest", - "PdfToMarkdownResponse", "Discrepancy", + "DocumentMeta", + "DocumentSections", "DiscrepancyKind", "EditCannotDoResponse", "EditClarificationRequest", @@ -166,15 +163,11 @@ __all__ = [ "NeedContentFileRequest", "NeedContentResponse", "NeedIngestResponse", + "ConvertMarkdownResponse", "NextExecutionAction", "OrchestratorRequest", "OrchestratorResponse", - "LayoutFragment", - "LayoutLine", "Page", - "PageLayout", - "PageLayoutArtifact", - "PageLayoutFileEntry", "PageRange", "PageText", "PdfCommentInstruction", @@ -182,6 +175,11 @@ __all__ = [ "PdfCommentRequest", "PdfCommentResponse", "PdfContentType", + "PdfCreateCannotDoResponse", + "PdfCreateOrchestrateResponse", + "PdfCreateRequest", + "PdfCreateResponse", + "PdfCreateSuccessResponse", "PdfEditRequest", "PdfEditResponse", "PdfEditTerminalResponse", diff --git a/engine/src/stirling/contracts/common.py b/engine/src/stirling/contracts/common.py index 05103b1a4a..8f35c9ceb4 100644 --- a/engine/src/stirling/contracts/common.py +++ b/engine/src/stirling/contracts/common.py @@ -62,6 +62,7 @@ class WorkflowOutcome(StrEnum): CANNOT_CONTINUE = "cannot_continue" UNSUPPORTED_CAPABILITY = "unsupported_capability" GENERATE_FILE = "generate_file" + CONVERT_MARKDOWN = "convert_markdown" class ArtifactKind(StrEnum): @@ -87,11 +88,11 @@ class SupportedCapability(StrEnum): PDF_EDIT = "pdf_edit" PDF_QUESTION = "pdf_question" PDF_REVIEW = "pdf_review" + PDF_CREATE = "pdf_create" AGENT_DRAFT = "agent_draft" AGENT_REVISE = "agent_revise" AGENT_NEXT_ACTION = "agent_next_action" MATH_AUDITOR_AGENT = "math_auditor_agent" - PDF_TO_MARKDOWN = "pdf_to_markdown" class ConversationMessage(ApiModel): @@ -183,6 +184,19 @@ class NeedIngestResponse(ApiModel): content_types: list[PdfContentType] = Field(default_factory=list) +class ConvertMarkdownResponse(ApiModel): + """Terminal signal: convert the listed files to Markdown deterministically. + + This is a deterministic, non-AI conversion. Java runs the PDF→Markdown converter + (``PdfMarkdownConverter``) on each file and returns the resulting ``.md`` file(s) as a + completed result. There is no resume turn — the conversion output is the final answer. + """ + + outcome: Literal[WorkflowOutcome.CONVERT_MARKDOWN] = WorkflowOutcome.CONVERT_MARKDOWN + reason: str + files_to_ingest: list[AiFile] + + class ToolOperationStep(ApiModel): kind: Literal[StepKind.TOOL] = StepKind.TOOL tool: AnyToolId diff --git a/engine/src/stirling/contracts/orchestrator.py b/engine/src/stirling/contracts/orchestrator.py index 1bf0f6eb36..8b916ccaff 100644 --- a/engine/src/stirling/contracts/orchestrator.py +++ b/engine/src/stirling/contracts/orchestrator.py @@ -11,6 +11,7 @@ from .common import ( AiFile, ArtifactKind, ConversationMessage, + ConvertMarkdownResponse, ExtractedFileText, GenerateFileResponse, NeedContentResponse, @@ -23,7 +24,6 @@ from .common import ( from .execution import NextExecutionAction from .pdf_edit import PdfEditTerminalResponse from .pdf_questions import PdfQuestionTerminalResponse -from .pdf_to_markdown import PageLayoutArtifact class ExtractedTextArtifact(ApiModel): @@ -32,7 +32,7 @@ class ExtractedTextArtifact(ApiModel): WorkflowArtifact = Annotated[ - ExtractedTextArtifact | PageLayoutArtifact | ToolReportArtifact, + ExtractedTextArtifact | ToolReportArtifact, Field(discriminator="kind"), ] @@ -61,6 +61,7 @@ type OrchestratorResponse = Annotated[ | GenerateFileResponse | NeedContentResponse | NeedIngestResponse + | ConvertMarkdownResponse | AgentDraftResponse | NextExecutionAction | UnsupportedCapabilityResponse, diff --git a/engine/src/stirling/contracts/pdf_create.py b/engine/src/stirling/contracts/pdf_create.py new file mode 100644 index 0000000000..761a192dd1 --- /dev/null +++ b/engine/src/stirling/contracts/pdf_create.py @@ -0,0 +1,226 @@ +"""Contracts for the PDF Create Agent. + +The agent accepts a natural-language prompt and returns a single +CREATE_PDF_FROM_HTML_AGENT plan step carrying the rendered HTML. + +Pipeline: + 1. PlannerAgent (smart_model) → DocumentPlan: structured skeleton, no body text. + 2. Python chunks the plan by token budget. + 3. SectionWriterAgents (smart_model, parallel) → WrittenSections per chunk. + 4. Assembler collects sections in plan order → GeneratedDocument. + 5. Jinja renders GeneratedDocument → HTML. The LLM never writes HTML. +""" + +from __future__ import annotations + +import re +from enum import StrEnum +from typing import Annotated, Literal + +from pydantic import Field, field_validator + +from stirling.models import ApiModel + +from .common import ConversationMessage +from .pdf_edit import EditCannotDoResponse, EditPlanResponse + + +class SectionType(StrEnum): + TEXT = "text" + KEY_VALUE = "key_value" + LINE_ITEMS = "line_items" + BULLET_LIST = "bullet_list" + SIGNATURE = "signature" + + +class TextSection(ApiModel): + """One or more prose paragraphs. Use for introductions, summaries, and narrative content.""" + + type: Literal[SectionType.TEXT] = SectionType.TEXT + heading: str | None = None + body: str = Field(description="Paragraph text. Use \\n\\n to separate paragraphs.") + + +class KeyValueSection(ApiModel): + """Labelled fields. Use for contact info, dates, invoice details, and metadata.""" + + type: Literal[SectionType.KEY_VALUE] = SectionType.KEY_VALUE + heading: str | None = None + pairs: list[tuple[str, str]] = Field(description="List of (label, value) pairs.") + + +class LineItemsSection(ApiModel): + """A table with column headers and data rows. Use for invoices, expenses, schedules.""" + + type: Literal[SectionType.LINE_ITEMS] = SectionType.LINE_ITEMS + heading: str | None = None + columns: list[str] = Field(description="Column header names.") + rows: list[list[str]] = Field(description="Data rows; each row must match columns in length.") + total_row: list[str] | None = None + + +class BulletListSection(ApiModel): + """An unordered list. Use for requirements, responsibilities, or any enumerated items.""" + + type: Literal[SectionType.BULLET_LIST] = SectionType.BULLET_LIST + heading: str | None = None + items: list[str] + + +class SignatureSection(ApiModel): + """Signature blocks. Use when the document requires sign-off from named parties.""" + + type: Literal[SectionType.SIGNATURE] = SectionType.SIGNATURE + heading: str | None = None + signatories: list[str] = Field(description="Names or roles to sign, e.g. 'John Smith, CEO'.") + + +type DocumentSection = Annotated[ + TextSection | KeyValueSection | LineItemsSection | BulletListSection | SignatureSection, + Field(discriminator="type"), +] + + +# Named colour or hex only — anything else is dropped so a colour can't inject CSS into the +# +
+
+ +
+
${iconSvg}
+
${escapeHtml(name)}
+
${escapeHtml(description)}
+
+
`; +} + +const escapeHtml = (s) => + String(s).replace(/&/g, "&").replace(//g, ">"); + +// ---- render ---------------------------------------------------------------- +let _browser = null; +async function getBrowser() { + if (_browser) return _browser; + const puppeteer = require("puppeteer"); + _browser = await puppeteer.launch({ + headless: "new", + args: ["--no-sandbox"], + }); + return _browser; +} + +export async function renderOgCard({ + name, + description, + icon, + outFile, + theme = THEME, +}) { + const iconSvg = resolveIcon(icon); + const html = await buildHtml({ name, description, iconSvg, theme }); + const browser = await getBrowser(); + const page = await browser.newPage(); + await page.setViewport({ + width: theme.width, + height: theme.height, + deviceScaleFactor: 1, + }); + await page.setContent(html, { waitUntil: "networkidle0" }); + const fit = await page.evaluate( + async (family, weight, lineH, maxLines, minSize) => { + const el = document.querySelector(".title"); + // Force the actual web font (this exact weight) to load before measuring, + // else wrapping/height is computed against the fallback font. + try { + await document.fonts.load(`${weight} 80px "${family}"`); + await document.fonts.ready; + } catch { + /* best-effort font preload */ + } + // The title wraps within its max-width. Shrink only if it would exceed + // maxLines lines, so long names wrap to 2 lines instead of spilling off. + let size = parseFloat(getComputedStyle(el).fontSize); + while (size > minSize && el.offsetHeight > maxLines * lineH * size + 2) { + size -= 1; + el.style.fontSize = size + "px"; + } + return { size, lines: Math.round(el.offsetHeight / (lineH * size)) }; + }, + theme.fontFamily, + theme.titleWeight, + 1.06, + 2, + 40, + ); + if (process.env.OG_DEBUG) console.log(" fit:", JSON.stringify(fit)); + await fs.mkdir(path.dirname(outFile), { recursive: true }); + await page.screenshot({ path: outFile, type: "png" }); + await page.close(); + return outFile; +} + +export async function closeBrowser() { + if (_browser) await _browser.close(); + _browser = null; +} + +// ---- batch: generate cards for tools that currently have no bespoke art ---- +// material-symbols icon per tool (the white glyph shown in the card). +const MISSING_TOOL_ICONS = { + addText: "title", + annotate: "draw", + timestampPdf: "schedule", + bookletImposition: "menu-book-outline", + pdfTextEditor: "edit-document-outline", + formFill: "ballot-outline", + devApi: "api", + devFolderScanning: "folder-open-outline", + devSsoGuide: "key-outline", + devAirgapped: "cloud-off-outline", +}; + +const kebab = (id) => id.replace(/([A-Z])/g, "-$1").toLowerCase(); + +// English name/description live next to each tool as the `t(key, fallback)` default. +function readRegistryStrings() { + const src = require("node:fs").readFileSync( + path.join(ROOT, "src/core/data/useTranslatedToolRegistry.tsx"), + "utf8", + ); + const STR = '"((?:[^"\\\\]|\\\\.)*)"'; + const titles = {}, + descs = {}; + for (const m of src.matchAll( + new RegExp('t\\(\\s*"home\\.([A-Za-z0-9_]+)\\.title"\\s*,\\s*' + STR, "g"), + )) + titles[m[1]] = m[2]; + for (const m of src.matchAll( + new RegExp('t\\(\\s*"home\\.([A-Za-z0-9_]+)\\.desc"\\s*,\\s*' + STR, "g"), + )) + descs[m[1]] = m[2]; + return { titles, descs }; +} + +export async function generateMissing(theme = THEME) { + const { titles, descs } = readRegistryStrings(); + const results = []; + for (const [id, icon] of Object.entries(MISSING_TOOL_ICONS)) { + const out = path.join(ROOT, `public/og_images/${kebab(id)}.png`); + await renderOgCard({ + name: titles[id] || id, + description: descs[id] || "", + icon, + outFile: out, + theme, + }); + results.push(`${id} -> public/og_images/${kebab(id)}.png (${icon})`); + } + return results; +} + +// Each tool's app icon lives as `icon=""` just before its +// `name: t("home..title", …)`. Pair each title with the closest preceding icon. +function readRegistryIcons() { + const src = require("node:fs").readFileSync( + path.join(ROOT, "src/core/data/useTranslatedToolRegistry.tsx"), + "utf8", + ); + const icons = [...src.matchAll(/icon="([^"]+)"/g)].map((m) => ({ + pos: m.index, + name: m[1], + })); + const byId = {}; + for (const m of src.matchAll( + /name:\s*t\(\s*"home\.([A-Za-z0-9_]+)\.title"/g, + )) { + let best = null; + for (const ic of icons) + if (ic.pos < m.index && (!best || ic.pos > best.pos)) best = ic; + if (best) byId[m[1]] = best.name; + } + return byId; +} + +function iconExists(name) { + if (!name) return false; + try { + const { getIconData } = require("@iconify/utils"); + return !!getIconData( + require("@iconify-json/material-symbols/icons.json"), + name, + ); + } catch { + return false; + } +} + +// First candidate that resolves; also tries dropping a "-rounded" suffix. +function firstResolvableIcon(candidates) { + for (const c of candidates) { + if (iconExists(c)) return c; + const alt = c && c.replace(/-rounded$/, ""); + if (alt && alt !== c && iconExists(alt)) return alt; + } + return "description-outline"; +} + +const humanizeId = (id) => + id + .replace(/([A-Z])/g, " $1") + .replace(/^./, (c) => c.toUpperCase()) + .replace(/\s+/g, " ") + .trim(); + +// Regenerate a card for EVERY tool that has an image, so the whole set is one +// consistent style. Writes to each tool's existing filename (from ogImageMap). +export async function generateAll(theme = THEME) { + const { titles, descs } = readRegistryStrings(); + const regIcons = readRegistryIcons(); + const ogMap = JSON.parse( + require("node:fs").readFileSync( + path.join(ROOT, "src/core/data/ogImageMap.json"), + "utf8", + ), + ); + const results = []; + for (const [id, basename] of Object.entries(ogMap)) { + const icon = firstResolvableIcon([regIcons[id], MISSING_TOOL_ICONS[id]]); + await renderOgCard({ + name: titles[id] || humanizeId(id), + description: descs[id] || "", + icon, + outFile: path.join(ROOT, `public/og_images/${basename}.png`), + theme, + }); + results.push(`${id.padEnd(20)} -> ${basename}.png (${icon})`); + } + return results; +} + +// ---- CLI ------------------------------------------------------------------- +function parseArgs(argv) { + const a = {}; + for (let i = 0; i < argv.length; i++) { + if (argv[i].startsWith("--")) { + const k = argv[i].slice(2); + const v = argv[i + 1] && !argv[i + 1].startsWith("--") ? argv[++i] : true; + a[k] = v; + } + } + return a; +} + +if ( + import.meta.url === `file://${process.argv[1]}` || + process.argv[1]?.endsWith("generate-og-image.mjs") +) { + const a = parseArgs(process.argv.slice(2)); + const theme = { ...THEME }; + for (const [flag, key] of Object.entries({ + "bg-top": "bgTop", + "bg-mid": "bgMid", + "bg-bottom": "bgBottom", + rotate: "cardRotate", + "title-size": "titleSize", + font: "fontFamily", + })) { + if (a[flag] != null) theme[key] = isNaN(+a[flag]) ? a[flag] : +a[flag]; + } + const run = async () => { + if (a.all) { + const results = await generateAll(theme); + console.log( + `Regenerated ${results.length} tool cards:\n ` + results.join("\n "), + ); + } else if (a.missing) { + const results = await generateMissing(theme); + console.log( + "Generated cards for tools with no bespoke art:\n " + + results.join("\n "), + ); + } else if (a.name) { + const out = + a.out || + `public/og_images/${a.name.toLowerCase().replace(/\s+/g, "-")}.png`; + await renderOgCard({ + name: a.name, + description: a.desc || "", + icon: a.icon, + outFile: path.resolve(ROOT, out), + theme, + }); + console.log("wrote", out); + } else { + console.error( + "Provide --missing, or --name (and --desc, --icon, --out). See header for usage.", + ); + process.exitCode = 1; + } + await closeBrowser(); + }; + run().catch(async (e) => { + console.error(e); + await closeBrowser(); + process.exit(1); + }); +} diff --git a/frontend/editor/scripts/generate-og-metadata.mjs b/frontend/editor/scripts/generate-og-metadata.mjs new file mode 100644 index 0000000000..3fd4beaf3a --- /dev/null +++ b/frontend/editor/scripts/generate-og-metadata.mjs @@ -0,0 +1,234 @@ +// Generates Open Graph (OG) / SEO social-preview metadata from the single +// source of truth in the frontend (tool ids, URL aliases, translated English +// strings) plus the actual images in public/og_images. +// +// Outputs: +// src/core/data/ogImageMap.json - { toolId: imageBasename } (imported by the client) +// public/og-metadata.json - { default, byTool, byPath } (read by the backend at startup) +// +// Run: `node scripts/generate-og-metadata.mjs` (writes files) +// `node scripts/generate-og-metadata.mjs --check` (CI drift guard: fails if stale) +// +// Why a generator instead of hand-maintained JSON: tool ids, URL aliases and +// English copy already live in the codebase. Regenerating keeps OG metadata in +// lockstep with the tool registry and surfaces tools that have no art. + +import fs from "node:fs"; +import path from "node:path"; +import { fileURLToPath } from "node:url"; + +const HERE = path.dirname(fileURLToPath(import.meta.url)); +const ROOT = path.resolve(HERE, ".."); +const read = (p) => fs.readFileSync(path.join(ROOT, p), "utf8"); + +const SITE_NAME = "Stirling PDF"; +const SITE_TITLE = "Stirling PDF"; +const SITE_DESC = "The Free Adobe Acrobat alternative (10M+ Downloads)"; +const DEFAULT_IMAGE_BASENAME = "home"; + +// Tools whose art exists under a legacy v1 filename that does not match the +// tool id or any current URL slug. Verified against public/og_images contents. +const LEGACY_IMAGE_OVERRIDES = { + merge: "mergePdfs", + crop: "cropPdf", + getPdfInfo: "get-all-info-on-pdf", + validateSignature: "validate-pdf-signature", + replaceColor: "replace-and-invert-color", + scalePages: "adjust-page-size-scale", + adjustContrast: "adjust-colors-contrast", + autoRename: "auto-rename-pdf-file", + removeBlanks: "remove-blank-pages", + removePages: "remove", + scannerImageSplit: "detect-split-scanned-photos", +}; + +// --- parse tool ids --------------------------------------------------------- +const idSrc = read("src/core/types/toolId.ts"); +function idArray(name) { + const m = idSrc.match( + new RegExp("export const " + name + " = \\[([\\s\\S]*?)\\] as const"), + ); + return m ? [...m[1].matchAll(/"([^"]+)"/g)].map((x) => x[1]) : []; +} +const regularIds = idArray("CORE_REGULAR_TOOL_IDS"); +const superIds = idArray("CORE_SUPER_TOOL_IDS"); +const linkIds = idArray("CORE_LINK_TOOL_IDS"); +const allIds = [...regularIds, ...superIds, ...linkIds]; + +// --- parse URL aliases ------------------------------------------------------ +const mapSrc = read("src/core/utils/urlMapping.ts"); +const urlToTool = {}; +for (const m of mapSrc.matchAll(/"([^"]+)":\s*"([^"]+)"/g)) + urlToTool[m[1]] = m[2]; + +// --- parse English title/description fallbacks ------------------------------ +const regSrc = read("src/core/data/useTranslatedToolRegistry.tsx"); +const titleById = {}; +const descById = {}; +const STR = '"((?:[^"\\\\]|\\\\.)*)"'; +for (const m of regSrc.matchAll( + new RegExp('t\\(\\s*"home\\.([A-Za-z0-9_]+)\\.title"\\s*,\\s*' + STR, "g"), +)) + titleById[m[1]] = m[2]; +for (const m of regSrc.matchAll( + new RegExp('t\\(\\s*"home\\.([A-Za-z0-9_]+)\\.desc"\\s*,\\s*' + STR, "g"), +)) + descById[m[1]] = m[2]; + +// --- available images ------------------------------------------------------- +const imageDir = "public/og_images"; +const images = new Set( + fs + .readdirSync(path.join(ROOT, imageDir)) + .filter((f) => f.endsWith(".png")) + .map((f) => f.replace(/\.png$/, "")), +); + +const canonicalPath = (id) => "/" + id.replace(/([A-Z])/g, "-$1").toLowerCase(); +const aliasesByTool = {}; +for (const [p, id] of Object.entries(urlToTool)) + (aliasesByTool[id] ??= []).push(p); + +function resolveImage(id) { + const override = LEGACY_IMAGE_OVERRIDES[id]; + if (override) return images.has(override) ? override : null; + const candidates = [ + id, + canonicalPath(id).slice(1), + ...(aliasesByTool[id] || []).map((a) => a.slice(1)), + ]; + for (const c of candidates) if (images.has(c)) return c; + return null; +} + +const humanize = (id) => + id + .replace(/([A-Z]+)/g, " $1") + .replace(/^./, (c) => c.toUpperCase()) + .replace(/\s+/g, " ") + .trim(); +const titleFor = (id) => `${titleById[id] || humanize(id)} - ${SITE_NAME}`; +const descFor = (id) => descById[id] || SITE_DESC; + +// --- build outputs ---------------------------------------------------------- +const ogImageMap = {}; // toolId -> basename (only tools with art) +const byTool = {}; +const missing = []; +for (const id of allIds) { + const img = resolveImage(id); + if (img) ogImageMap[id] = img; + else missing.push(id); + byTool[id] = { + image: `/og_images/${img || DEFAULT_IMAGE_BASENAME}.png`, + title: titleFor(id), + description: descFor(id), + }; +} + +// path -> toolId for every canonical path and every alias +const byPath = {}; +for (const id of allIds) byPath[canonicalPath(id)] = id; +for (const [p, id] of Object.entries(urlToTool)) byPath[p] = id; + +// --- non-tool application routes -------------------------------------------- +// Every other URL the SPA serves also gets OG: auth, the file manager, the +// mobile scanner, and each settings section. These have no bespoke art (default +// image) but carry a page-specific title so shared links are labelled correctly. +// Keyed by path (tool ids never start with "/", so there is no collision). +const navKeys = ( + read("src/core/components/shared/config/types.ts") + .match(/export const VALID_NAV_KEYS = \[([\s\S]*?)\] as const/)?.[1] + .match(/"([^"]+)"/g) || [] +).map((s) => s.replace(/"/g, "")); + +const humanizeLabel = (s) => + s + .replace(/[-_]/g, " ") + .replace(/([a-z])([A-Z])/g, "$1 $2") + .replace(/\s+/g, " ") + .trim() + .replace(/\b\w/g, (c) => c.toUpperCase()); + +const pageTitles = { + "/login": "Sign In", + "/signup": "Sign Up", + "/mobile-scanner": "Mobile Scanner", + "/files": "Files", + "/settings": "Settings", +}; +for (const key of navKeys) + pageTitles[`/settings/${key}`] = `${humanizeLabel(key)} Settings`; + +for (const [routePath, label] of Object.entries(pageTitles)) { + byTool[routePath] = { + image: `/og_images/${DEFAULT_IMAGE_BASENAME}.png`, + title: `${label} - ${SITE_NAME}`, + description: SITE_DESC, + }; + byPath[routePath] = routePath; +} + +const manifest = { + default: { + image: `/og_images/${DEFAULT_IMAGE_BASENAME}.png`, + title: SITE_TITLE, + description: SITE_DESC, + }, + byTool, + byPath, +}; + +const mapJson = JSON.stringify(ogImageMap, null, 2) + "\n"; +const manifestJson = JSON.stringify(manifest, null, 2) + "\n"; +const mapPath = "src/core/data/ogImageMap.json"; +const manifestPath = "public/og-metadata.json"; + +const check = process.argv.includes("--check"); +if (check) { + const stale = []; + if (!fs.existsSync(path.join(ROOT, mapPath)) || read(mapPath) !== mapJson) + stale.push(mapPath); + if ( + !fs.existsSync(path.join(ROOT, manifestPath)) || + read(manifestPath) !== manifestJson + ) + stale.push(manifestPath); + if (stale.length) { + console.error( + "OG metadata is stale. Run `node scripts/generate-og-metadata.mjs`:\n " + + stale.join("\n "), + ); + process.exit(1); + } + console.log("OG metadata is up to date."); +} else { + fs.writeFileSync(path.join(ROOT, mapPath), mapJson); + fs.writeFileSync(path.join(ROOT, manifestPath), manifestJson); + console.log( + `Wrote ${mapPath} (${Object.keys(ogImageMap).length} tools with art)`, + ); + console.log(`Wrote ${manifestPath} (${Object.keys(byPath).length} paths)`); +} + +// --- report ----------------------------------------------------------------- +console.log( + `\nTools with OG image: ${allIds.length - missing.length}/${allIds.length}`, +); +console.log( + `Tools using the DEFAULT image (${DEFAULT_IMAGE_BASENAME}.png) - no bespoke art: ${missing.length}`, +); +for (const id of missing) { + const kind = superIds.includes(id) + ? "super" + : linkIds.includes(id) + ? "link" + : "regular"; + console.log(` ${kind.padEnd(8)} ${id.padEnd(20)} ${canonicalPath(id)}`); +} +const used = new Set(Object.values(ogImageMap)); +const orphans = [...images] + .filter((i) => !used.has(i) && i !== DEFAULT_IMAGE_BASENAME) + .sort(); +console.log( + `\nUnused images in ${imageDir} (${orphans.length}): ${orphans.join(", ")}`, +); diff --git a/frontend/editor/scripts/og-assets/stirling-lockup.png b/frontend/editor/scripts/og-assets/stirling-lockup.png new file mode 100644 index 0000000000..d34075798e Binary files /dev/null and b/frontend/editor/scripts/og-assets/stirling-lockup.png differ diff --git a/frontend/editor/scripts/og-prerender.d.mts b/frontend/editor/scripts/og-prerender.d.mts new file mode 100644 index 0000000000..9f2c58ec6d --- /dev/null +++ b/frontend/editor/scripts/og-prerender.d.mts @@ -0,0 +1,32 @@ +// Type declarations for og-prerender.mjs (plain ESM build helper). + +export interface OgEntry { + image: string; + title: string; + description: string; +} + +export interface OgInjectOptions { + ogBase?: string; + pageUrlPath?: string | null; +} + +export interface OgManifest { + default: OgEntry; + byTool: Record; + byPath: Record; +} + +export function escapeHtml(value: string): string; +export function buildOgTags(entry: OgEntry, opts?: OgInjectOptions): string; +export function injectOg( + html: string, + entry: OgEntry, + opts?: OgInjectOptions, +): string; +export function prerenderOg(args: { + distDir: string; + manifest: OgManifest; + ogBase?: string; + baseHref?: string; +}): Promise; diff --git a/frontend/editor/scripts/og-prerender.mjs b/frontend/editor/scripts/og-prerender.mjs new file mode 100644 index 0000000000..ee189fa7a2 --- /dev/null +++ b/frontend/editor/scripts/og-prerender.mjs @@ -0,0 +1,118 @@ +// Pure helpers for baking Open Graph / Twitter Card tags into prerendered HTML. +// Kept separate from vite.config so the logic is unit-testable without a full build. +// Used by the `prerender-og` Vite plugin (see vite.config.ts). + +import fs from "node:fs/promises"; +import path from "node:path"; + +export const escapeHtml = (value) => + String(value) + .replace(/&/g, "&") + .replace(//g, ">") + .replace(/"/g, """); + +const absolute = (urlPath, ogBase) => (ogBase ? ogBase + urlPath : urlPath); + +/** + * Build the OG/Twitter block for one route. + * @param {{image:string,title:string,description:string}} entry + * @param {{ogBase?:string, pageUrlPath?:string|null}} opts + */ +export function buildOgTags(entry, { ogBase = "", pageUrlPath = null } = {}) { + const title = escapeHtml(entry.title); + const description = escapeHtml(entry.description); + const imageUrl = absolute(entry.image, ogBase); + const image = escapeHtml(imageUrl); + const pageUrl = pageUrlPath + ? escapeHtml(absolute(pageUrlPath, ogBase)) + : null; + const lines = [ + "", + '', + '', + ``, + ``, + pageUrl ? `` : null, + ``, + imageUrl.startsWith("https") + ? `` + : null, + '', + '', + '', + '', + ``, + ``, + ``, + "", + ].filter(Boolean); + return lines.join("\n ") + "\n "; +} + +/** Inject route-specific , description and OG/Twitter tags into an HTML shell. */ +export function injectOg(html, entry, opts = {}) { + return html + .replace( + /<title>[\s\S]*?<\/title>/i, + () => `<title>${escapeHtml(entry.title)}`, + ) + .replace( + //i, + () => + ``, + ) + .replace("", ` ${buildOgTags(entry, opts)}`); +} + +const BASE_HREF_RE = //i; + +/** + * Write the root index.html (home preview) plus one file per route in the + * manifest: flat for single-segment routes (e.g. dist/compress.html) and nested + * for multi-segment ones (e.g. dist/settings/people.html). Returns the count of + * route pages written. + * + * `baseHref` is the absolute deploy base ("/" for a root deploy). Nested files + * need it because a relative `` would resolve their assets + * against the sub-path (e.g. /settings/) and 404; flat files and the root keep + * the build's relative base. + * @returns {Promise} + */ +export async function prerenderOg({ + distDir, + manifest, + ogBase = "", + baseHref = "/", +}) { + const template = await fs.readFile(path.join(distDir, "index.html"), "utf8"); + + await fs.writeFile( + path.join(distDir, "index.html"), + injectOg(template, manifest.default, { + ogBase, + pageUrlPath: ogBase ? "/" : null, + }), + ); + + let count = 0; + for (const [routePath, id] of Object.entries(manifest.byPath || {})) { + const segments = routePath.replace(/^\//, "").split("/"); + // Clean segments only - no traversal, no dots (e.g. /compress, /settings/people). + if (!segments.length || !segments.every((s) => /^[A-Za-z0-9_-]+$/.test(s))) + continue; + const entry = manifest.byTool[id] ?? manifest.default; + let html = injectOg(template, entry, { + ogBase, + pageUrlPath: ogBase ? routePath : null, + }); + const nested = segments.length > 1; + if (nested) + html = html.replace(BASE_HREF_RE, ``); + const outFile = path.join(distDir, ...segments) + ".html"; + if (nested) await fs.mkdir(path.dirname(outFile), { recursive: true }); + await fs.writeFile(outFile, html); + count++; + } + return count; +} diff --git a/frontend/editor/src-tauri/tauri.conf.json b/frontend/editor/src-tauri/tauri.conf.json index 6c6dcb376b..376f1027cd 100644 --- a/frontend/editor/src-tauri/tauri.conf.json +++ b/frontend/editor/src-tauri/tauri.conf.json @@ -1,7 +1,7 @@ { "$schema": "../node_modules/@tauri-apps/cli/config.schema.json", "productName": "Stirling-PDF", - "version": "2.12.0", + "version": "2.13.0", "identifier": "stirling.pdf.dev", "build": { "frontendDist": "../dist", @@ -17,14 +17,16 @@ "height": 800, "resizable": true, "fullscreen": false, + "dragDropEnabled": false, "additionalBrowserArgs": "--enable-features=CertVerifierBuiltinFeature" } ] }, "bundle": { "active": true, + "createUpdaterArtifacts": true, "publisher": "Stirling PDF Inc.", - "targets": ["deb", "rpm", "appimage", "dmg", "msi"], + "targets": ["deb", "rpm", "appimage", "dmg", "app", "msi"], "icon": [ "icons/icon.png", "icons/icon.icns", diff --git a/frontend/editor/src/cloud/LICENSE b/frontend/editor/src/cloud/LICENSE new file mode 100644 index 0000000000..d268556808 --- /dev/null +++ b/frontend/editor/src/cloud/LICENSE @@ -0,0 +1,51 @@ +Stirling PDF User License + +Copyright (c) 2025 Stirling PDF Inc. + +License Scope & Usage Rights + +Production use of the Stirling PDF Software is only permitted with a valid Stirling PDF User License. + +For purposes of this license, “the Software” refers to the Stirling PDF application and any associated documentation files +provided by Stirling PDF Inc. You or your organization may not use the Software in production, at scale, or for business-critical +processes unless you have agreed to, and remain in compliance with, the Stirling PDF Subscription Terms of Service +(https://www.stirlingpdf.com/terms) or another valid agreement with Stirling PDF, and hold an active User License subscription +covering the appropriate number of licensed users. + +Trial and Minimal Use + +You may use the Software without a paid subscription for the sole purposes of internal trial, evaluation, or minimal use, provided that: +* Use is limited to the capabilities and restrictions defined by the Software itself; +* You do not copy, distribute, sublicense, reverse-engineer, or use the Software in client-facing or commercial contexts. + +Continued use beyond this scope requires a valid Stirling PDF User License. + +Modifications and Derivative Works + +You may modify the Software only for development or internal testing purposes. Any such modifications or derivative works: + +* May not be deployed in production environments without a valid User License; +* May not be distributed or sublicensed; +* Remain the intellectual property of Stirling PDF and/or its licensors; +* May only be used, copied, or exploited in accordance with the terms of a valid Stirling PDF User License subscription. + +Prohibited Actions + +Unless explicitly permitted by a paid license or separate agreement, you may not: + +* Use the Software in production environments; +* Copy, merge, distribute, sublicense, or sell the Software; +* Remove or alter any licensing or copyright notices; +* Circumvent access restrictions or licensing requirements. + +Third-Party Components + +The Stirling PDF Software may include components subject to separate open source licenses. Such components remain governed by +their original license terms as provided by their respective owners. + +Disclaimer + +THE SOFTWARE IS PROVIDED “AS IS,” WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO WARRANTIES OF +MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE, OR NON-INFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE +LIABLE FOR ANY CLAIM, DAMAGES, OR OTHER LIABILITY, WHETHER IN CONTRACT, TORT, OR OTHERWISE, ARISING FROM, OUT OF, OR IN +CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. diff --git a/frontend/editor/src/cloud/_probe.ts b/frontend/editor/src/cloud/_probe.ts new file mode 100644 index 0000000000..b40e59d86e --- /dev/null +++ b/frontend/editor/src/cloud/_probe.ts @@ -0,0 +1 @@ +export const CLOUD_LAYER_PROBE = "cloud" as const; diff --git a/frontend/editor/src/cloud/auth/session.ts b/frontend/editor/src/cloud/auth/session.ts new file mode 100644 index 0000000000..2dc8da30fb --- /dev/null +++ b/frontend/editor/src/cloud/auth/session.ts @@ -0,0 +1,16 @@ +/** + * Auth/session seam (@app/auth/session). saas keeps a Supabase web session; + * desktop keeps a JWT in the Tauri secure store — cloud code reads the token + * through this seam instead. Default no-op; saas/ and desktop/ shadow it. + */ + +/** Minimal session shape shared across platforms. */ +export interface AppSession { + /** Bearer access token for authenticated API calls, or null when signed out. */ + accessToken: string | null; +} + +/** The current access token for authenticated API calls, or null when signed out. */ +export async function getAccessToken(): Promise { + return null; +} diff --git a/frontend/editor/src/cloud/auth/teamSession.ts b/frontend/editor/src/cloud/auth/teamSession.ts new file mode 100644 index 0000000000..1de09774bc --- /dev/null +++ b/frontend/editor/src/cloud/auth/teamSession.ts @@ -0,0 +1,37 @@ +/** + * Team-session auth seam (@app/auth/teamSession). {@link SaaSTeamContext} needs + * just two things from auth (can-use-teams + a post-membership refresh), so it + * consumes them through this narrow seam rather than reaching a platform auth + * surface directly. Default reports no access; saas/ and desktop/ shadow it. + */ + +/** + * The minimal auth surface {@link SaaSTeamContext} depends on. + * + * Each platform supplies its own implementation: + * - {@link canUseTeams} is {@code true} only for a signed-in, non-anonymous + * user (web: Supabase non-anonymous session; desktop: a valid authService + * token). When {@code false} the context stays empty and makes no API calls. + * - {@link refreshAfterMembershipChange} is invoked after a membership change + * (accepting an invite, leaving a team) so the platform can refresh any + * derived auth-tier state. On web this refreshes credits + the Supabase + * session; on desktop it is a no-op (no such derived state). + */ +export interface TeamAuth { + /** Whether the current session may load/manage teams (signed in, not anonymous). */ + canUseTeams: boolean; + /** Refresh derived auth state after a team membership change. */ + refreshAfterMembershipChange: () => Promise; +} + +/** + * Resolve the team-relevant auth surface for the current session. Cloud default + * reports no access + a no-op refresh; saas/desktop shadow it. Implemented as a + * hook so the saas/desktop impls can subscribe to live auth state. + */ +export function useTeamAuth(): TeamAuth { + return { + canUseTeams: false, + refreshAfterMembershipChange: async () => {}, + }; +} diff --git a/frontend/editor/src/cloud/components/UsageLimitModalHost.tsx b/frontend/editor/src/cloud/components/UsageLimitModalHost.tsx new file mode 100644 index 0000000000..1903ec4211 --- /dev/null +++ b/frontend/editor/src/cloud/components/UsageLimitModalHost.tsx @@ -0,0 +1,57 @@ +import { useEffect, useState } from "react"; +import { FreeLimitReachedModal } from "@app/components/shared/FreeLimitReachedModal"; +import { SpendCapReachedModal } from "@app/components/shared/SpendCapReachedModal"; +import { + FREE_LIMIT_MODAL_EVENT, + SPEND_CAP_MODAL_EVENT, +} from "@app/components/usageLimitModals"; +import { + PAYG_LIMIT_REACHED_EVENT, + type PaygLimitReachedDetail, +} from "@app/services/usageLimitBridge"; + +/** + * Always-mounted host for the usage-limit warning modals. Mount once (in + * App.tsx); it renders nothing until openFreeLimitModal()/openSpendCapModal() + * (see usageLimitModals.ts) fire their bridge events. Each modal is mounted + * only while open, so it reads the wallet (and animates in) on open rather + * than on app load. + * + *

Also bridges the server-side run paths (policy auto-run, AI agent): their tool + * calls run server-side, so a usage-limit 402 never reaches the apiClient interceptor + * that pops these modals for direct calls. Those proprietary paths broadcast {@link + * PAYG_LIMIT_REACHED_EVENT} (with the blocking 402's {@code subscribed} flag) instead; + * we open the matching modal here. + */ +export default function UsageLimitModalHost() { + const [freeOpen, setFreeOpen] = useState(false); + const [spendOpen, setSpendOpen] = useState(false); + + useEffect(() => { + const onFree = () => setFreeOpen(true); + const onSpend = () => setSpendOpen(true); + const onServerLimit = (e: Event) => { + const subscribed = (e as CustomEvent).detail + ?.subscribed; + if (subscribed) setSpendOpen(true); + else setFreeOpen(true); + }; + window.addEventListener(FREE_LIMIT_MODAL_EVENT, onFree); + window.addEventListener(SPEND_CAP_MODAL_EVENT, onSpend); + window.addEventListener(PAYG_LIMIT_REACHED_EVENT, onServerLimit); + return () => { + window.removeEventListener(FREE_LIMIT_MODAL_EVENT, onFree); + window.removeEventListener(SPEND_CAP_MODAL_EVENT, onSpend); + window.removeEventListener(PAYG_LIMIT_REACHED_EVENT, onServerLimit); + }; + }, []); + + return ( + <> + {freeOpen && setFreeOpen(false)} />} + {spendOpen && ( + setSpendOpen(false)} /> + )} + + ); +} diff --git a/frontend/editor/src/saas/components/onboarding/SaasOnboardingModal.tsx b/frontend/editor/src/cloud/components/onboarding/SaasOnboardingModal.tsx similarity index 82% rename from frontend/editor/src/saas/components/onboarding/SaasOnboardingModal.tsx rename to frontend/editor/src/cloud/components/onboarding/SaasOnboardingModal.tsx index 4af91fd020..5aa52034a4 100644 --- a/frontend/editor/src/saas/components/onboarding/SaasOnboardingModal.tsx +++ b/frontend/editor/src/cloud/components/onboarding/SaasOnboardingModal.tsx @@ -1,8 +1,8 @@ import React from "react"; import { Modal, Stack } from "@mantine/core"; -import DiamondOutlinedIcon from "@mui/icons-material/DiamondOutlined"; +import BoltRoundedIcon from "@mui/icons-material/BoltRounded"; +import GroupAddRoundedIcon from "@mui/icons-material/GroupAddRounded"; import { useTranslation } from "react-i18next"; -import LocalIcon from "@app/components/shared/LocalIcon"; import AnimatedSlideBackground from "@app/components/onboarding/slides/AnimatedSlideBackground"; import OnboardingStepper from "@app/components/onboarding/OnboardingStepper"; import { renderButtons } from "@app/components/onboarding/renderButtons"; @@ -14,6 +14,12 @@ import { Z_INDEX_OVER_FULLSCREEN_SURFACE } from "@app/styles/zIndex"; interface SaasOnboardingModalProps { opened: boolean; onClose: () => void; + /** + * Drop the closing "desktop-install" slide. Set by the desktop app, which + * reuses this flow but has no reason to pitch its own download. Defaults to + * false (slide shown) so the web (saas) flow is unchanged. + */ + hideDesktopInstall?: boolean; } export default function SaasOnboardingModal(props: SaasOnboardingModalProps) { @@ -48,18 +54,23 @@ export default function SaasOnboardingModal(props: SaasOnboardingModalProps) { ); } + if (slideDefinition.hero.type === "logo") { + return ( + Stirling logo + ); + } + return (

- {slideDefinition.hero.type === "rocket" && ( - + {slideDefinition.hero.type === "bolt" && ( + )} - {slideDefinition.hero.type === "diamond" && ( - + {slideDefinition.hero.type === "team" && ( + )}
); diff --git a/frontend/editor/src/saas/components/onboarding/renderButtons.tsx b/frontend/editor/src/cloud/components/onboarding/renderButtons.tsx similarity index 100% rename from frontend/editor/src/saas/components/onboarding/renderButtons.tsx rename to frontend/editor/src/cloud/components/onboarding/renderButtons.tsx diff --git a/frontend/editor/src/cloud/components/onboarding/saasFlowResolver.ts b/frontend/editor/src/cloud/components/onboarding/saasFlowResolver.ts new file mode 100644 index 0000000000..b6fb199730 --- /dev/null +++ b/frontend/editor/src/cloud/components/onboarding/saasFlowResolver.ts @@ -0,0 +1,34 @@ +import { SlideId } from "@app/components/onboarding/saasOnboardingFlowConfig"; + +export interface SaasFlowInputs { + /** Free-tier wallet with one-time allowance remaining — show the usage meter. */ + showUsageSlide: boolean; + /** Team leaders only — invited members and anonymous guests skip the team slide. */ + showTeamSlide: boolean; + /** + * Drop the closing "desktop-install" slide. The web (saas) flow pitches the + * desktop download, but the desktop app reuses this same flow and is already + * the desktop app, so it omits that slide. Defaults to false (slide shown). + */ + hideDesktopInstall?: boolean; +} + +/** + * Resolves the SaaS onboarding slide sequence. The free-editor pitch and + * desktop install bookend the flow; the usage meter and team slides slot in + * when their conditions hold. When {@link SaasFlowInputs.hideDesktopInstall} is + * set, the closing desktop-install slide is dropped (used by the desktop app, + * which has no reason to pitch its own download). + */ +export function resolveSaasFlow({ + showUsageSlide, + showTeamSlide, + hideDesktopInstall = false, +}: SaasFlowInputs): SlideId[] { + return [ + "free-editor", + ...(showUsageSlide ? (["usage"] as const) : []), + ...(showTeamSlide ? (["team"] as const) : []), + ...(hideDesktopInstall ? [] : (["desktop-install"] as const)), + ]; +} diff --git a/frontend/editor/src/saas/components/onboarding/saasOnboardingFlowConfig.ts b/frontend/editor/src/cloud/components/onboarding/saasOnboardingFlowConfig.ts similarity index 53% rename from frontend/editor/src/saas/components/onboarding/saasOnboardingFlowConfig.ts rename to frontend/editor/src/cloud/components/onboarding/saasOnboardingFlowConfig.ts index 342e1fe87f..91af54611b 100644 --- a/frontend/editor/src/saas/components/onboarding/saasOnboardingFlowConfig.ts +++ b/frontend/editor/src/cloud/components/onboarding/saasOnboardingFlowConfig.ts @@ -1,12 +1,12 @@ -import WelcomeSlide from "@app/components/onboarding/slides/WelcomeSlide"; +import FreeEditorSlide from "@app/components/onboarding/slides/FreeEditorSlide"; +import UsageSnapshotSlide from "@app/components/onboarding/slides/UsageSnapshotSlide"; +import TeamSlide from "@app/components/onboarding/slides/TeamSlide"; import DesktopInstallSlide from "@app/components/onboarding/slides/DesktopInstallSlide"; -import FreeTrialSlide from "@app/components/onboarding/slides/FreeTrialSlide"; import { SlideConfig } from "@app/types/types"; -import { TrialStatus } from "@app/auth/UseSession"; -export type SlideId = "welcome" | "free-trial" | "desktop-install"; +export type SlideId = "free-editor" | "usage" | "team" | "desktop-install"; -export type HeroType = "rocket" | "dual-icon" | "diamond"; +export type HeroType = "logo" | "bolt" | "team" | "dual-icon"; export type ButtonAction = "next" | "prev" | "close" | "download-selected"; @@ -23,7 +23,6 @@ export interface SlideFactoryParams { osUrl: string; osOptions?: OSOption[]; onDownloadUrlChange?: (url: string) => void; - trialStatus?: TrialStatus | null; } export interface HeroDefinition { @@ -48,47 +47,46 @@ export interface SlideDefinition { buttons: ButtonDefinition[]; } +const BACK_BUTTON: ButtonDefinition = { + key: "back", + type: "icon", + icon: "chevron-left", + group: "left", + action: "prev", +}; + +const NEXT_BUTTON: ButtonDefinition = { + key: "next", + type: "button", + label: "onboarding.buttons.next", + variant: "primary", + group: "right", + action: "next", +}; + export const SLIDE_DEFINITIONS: Record = { - welcome: { - id: "welcome", - createSlide: () => WelcomeSlide(), - hero: { type: "rocket" }, + "free-editor": { + id: "free-editor", + createSlide: () => FreeEditorSlide(), + hero: { type: "logo" }, + buttons: [{ ...NEXT_BUTTON, key: "free-editor-next" }], + }, + usage: { + id: "usage", + createSlide: () => UsageSnapshotSlide(), + hero: { type: "bolt" }, buttons: [ - { - key: "welcome-next", - type: "button", - label: "onboarding.buttons.next", - variant: "primary", - group: "right", - action: "next", - }, + { ...BACK_BUTTON, key: "usage-back" }, + { ...NEXT_BUTTON, key: "usage-next" }, ], }, - "free-trial": { - id: "free-trial", - createSlide: ({ trialStatus }) => { - if (!trialStatus) { - throw new Error("Trial status is required for free-trial slide"); - } - return FreeTrialSlide({ trialStatus }); - }, - hero: { type: "diamond" }, + team: { + id: "team", + createSlide: () => TeamSlide(), + hero: { type: "team" }, buttons: [ - { - key: "trial-back", - type: "icon", - icon: "chevron-left", - group: "left", - action: "prev", - }, - { - key: "trial-next", - type: "button", - label: "onboarding.buttons.next", - variant: "primary", - group: "right", - action: "next", - }, + { ...BACK_BUTTON, key: "team-back" }, + { ...NEXT_BUTTON, key: "team-next" }, ], }, "desktop-install": { @@ -97,13 +95,7 @@ export const SLIDE_DEFINITIONS: Record = { DesktopInstallSlide({ osLabel, osUrl, osOptions, onDownloadUrlChange }), hero: { type: "dual-icon" }, buttons: [ - { - key: "desktop-back", - type: "icon", - icon: "chevron-left", - group: "left", - action: "prev", - }, + { ...BACK_BUTTON, key: "desktop-back" }, { key: "desktop-skip", type: "button", @@ -123,8 +115,3 @@ export const SLIDE_DEFINITIONS: Record = { ], }, }; - -export const FLOW_SEQUENCES = { - saasTrialUser: ["welcome", "free-trial", "desktop-install"] as SlideId[], - saasPaidUser: ["welcome", "desktop-install"] as SlideId[], -}; diff --git a/frontend/editor/src/cloud/components/onboarding/slides/FreeEditorSlide.tsx b/frontend/editor/src/cloud/components/onboarding/slides/FreeEditorSlide.tsx new file mode 100644 index 0000000000..6391d0a0f7 --- /dev/null +++ b/frontend/editor/src/cloud/components/onboarding/slides/FreeEditorSlide.tsx @@ -0,0 +1,44 @@ +import React from "react"; +import { Trans } from "react-i18next"; +import { SlideConfig } from "@app/types/types"; +import { createLightSlideBackground } from "@app/components/onboarding/slides/unifiedBackgroundConfig"; +import styles from "@app/components/onboarding/slides/SaasOnboardingSlides.module.css"; + +// Stirling logo red (sampled from modern-logo/logo512.png) +const FREE_EDITOR_BACKGROUND = createLightSlideBackground( + [142, 49, 49], + "#F8E0E0", +); + +const FreeEditorBody = () => ( + + }} + defaults="We've added loads of new features, including Policies and Agent Chat." + /> + + }} + defaults="The editor is now completely free." + /> + + +); + +export default function FreeEditorSlide(): SlideConfig { + const title = ( + + ); + + return { + key: "free-editor", + title, + body: , + background: FREE_EDITOR_BACKGROUND, + }; +} diff --git a/frontend/editor/src/cloud/components/onboarding/slides/SaasOnboardingSlides.module.css b/frontend/editor/src/cloud/components/onboarding/slides/SaasOnboardingSlides.module.css new file mode 100644 index 0000000000..3fcdb93831 --- /dev/null +++ b/frontend/editor/src/cloud/components/onboarding/slides/SaasOnboardingSlides.module.css @@ -0,0 +1,91 @@ +/* Shared styles for the SaaS onboarding slides */ + +/* ── Slide 1: free editor pitch ─────────────────────────────────────────── */ + +.freeLine { + display: block; + margin-top: 10px; +} + +.freeHighlight { + background: linear-gradient(135deg, #6366f1, #ec4899); + -webkit-background-clip: text; + background-clip: text; + color: transparent; + font-weight: 700; +} + +/* ── Slide 2: usage meter snapshot ──────────────────────────────────────── */ + +.usageMeterWrap { + max-width: 420px; + margin: 16px auto 0; + text-align: left; +} + +/* ── Slide 3: team ──────────────────────────────────────────────────────── */ + +.teamCard { + max-width: 460px; + margin: 12px auto 0; + text-align: left; +} + +.memberList { + display: flex; + flex-direction: column; + max-height: 200px; + overflow-y: auto; + border: 1px solid var(--border-default, #d1d5db); + border-radius: 12px; + margin-bottom: 12px; +} + +.memberRow { + display: flex; + align-items: center; + gap: 12px; + padding: 10px 14px; +} + +.memberRow + .memberRow { + border-top: 1px solid var(--border-subtle, #e5e7eb); +} + +.memberIdentity { + flex: 1; + min-width: 0; + display: flex; + flex-direction: column; +} + +.memberName { + font-size: 14px; + font-weight: 600; + color: var(--onboarding-title, #111827); + overflow: hidden; + text-overflow: ellipsis; + white-space: nowrap; +} + +.memberEmail { + font-size: 12px; + color: var(--onboarding-body, #6b7280); + overflow: hidden; + text-overflow: ellipsis; + white-space: nowrap; +} + +.inviteRow { + display: flex; + gap: 8px; +} + +.inviteRow > :first-child { + flex: 1; +} + +.inviteFeedback { + font-size: 13px; + margin-top: 8px; +} diff --git a/frontend/editor/src/cloud/components/onboarding/slides/TeamSlide.tsx b/frontend/editor/src/cloud/components/onboarding/slides/TeamSlide.tsx new file mode 100644 index 0000000000..03eadf0f4c --- /dev/null +++ b/frontend/editor/src/cloud/components/onboarding/slides/TeamSlide.tsx @@ -0,0 +1,191 @@ +import React, { useEffect, useState } from "react"; +import { Badge, Button, TextInput } from "@mantine/core"; +import { useTranslation } from "react-i18next"; +import { SlideConfig } from "@app/types/types"; +import { createLightSlideBackground } from "@app/components/onboarding/slides/unifiedBackgroundConfig"; +import { useSaaSTeam } from "@app/contexts/SaaSTeamContext"; +import styles from "@app/components/onboarding/slides/SaasOnboardingSlides.module.css"; + +const TEAM_BACKGROUND = createLightSlideBackground([79, 70, 229], "#E0E7FF"); + +const EMAIL_PATTERN = /^[^\s@]+@[^\s@]+\.[^\s@]+$/; + +function useHasTeamMembers(): boolean { + const { teamMembers, teamInvitations } = useSaaSTeam(); + const pendingInvitations = teamInvitations.filter( + (invitation) => invitation.status === "PENDING", + ); + // The leader themselves is always in teamMembers, so "has a team" means + // anyone beyond them, or an invite already on its way. + return teamMembers.length > 1 || pendingInvitations.length > 0; +} + +function TeamSlideTitle() { + const { t } = useTranslation(); + const hasTeam = useHasTeamMembers(); + + return hasTeam + ? t("onboarding.saas.team.inviteTitle", "Invite members to your team") + : t("onboarding.saas.team.createTitle", "Create your team"); +} + +function InviteForm() { + const { t } = useTranslation(); + const { inviteUser } = useSaaSTeam(); + const [email, setEmail] = useState(""); + const [inviting, setInviting] = useState(false); + const [error, setError] = useState(null); + const [success, setSuccess] = useState(null); + + const emailValid = EMAIL_PATTERN.test(email); + + const handleInvite = async (e: React.FormEvent) => { + e.preventDefault(); + if (!emailValid) return; + + setInviting(true); + setError(null); + setSuccess(null); + + try { + await inviteUser(email); + setSuccess( + t("team.inviteSent", "Invitation sent to {{email}}", { email }), + ); + setEmail(""); + } catch (err) { + const inviteError = err as { response?: { data?: { error?: string } } }; + setError( + inviteError.response?.data?.error || + t("team.inviteError", "Failed to send invitation"), + ); + } finally { + setInviting(false); + } + }; + + return ( +
+ + setEmail(e.target.value)} + error={ + email && !emailValid + ? t("team.invite.invalidEmail", "Invalid email format") + : undefined + } + /> + + + {error && ( + + {error} + + )} + {success && ( + + {success} + + )} +
+ ); +} + +const TeamSlideBody = () => { + const { t } = useTranslation(); + const { teamMembers, teamInvitations, refreshTeams } = useSaaSTeam(); + const hasTeam = useHasTeamMembers(); + const pendingInvitations = teamInvitations.filter( + (invitation) => invitation.status === "PENDING", + ); + + // Onboarding shows right after first login, so make sure team data is fresh. + useEffect(() => { + refreshTeams(); + }, []); + + return ( + + {hasTeam + ? t( + "onboarding.saas.team.inviteBody", + "Everyone on your team shares files, automations and your plan. Add teammates by email and they'll get an invite.", + ) + : t( + "onboarding.saas.team.createBody", + "Work on documents together: teammates share files, automations and your plan. Add the first member by email to create your team.", + )} + + {hasTeam && ( + + {teamMembers.map((member) => ( + + + {member.username} + {member.email} + + + {member.role} + + + ))} + {pendingInvitations.map((invitation) => ( + + + + {invitation.inviteeEmail.split("@")[0]} + + + {invitation.inviteeEmail} + + + + {t("team.members.pending", "PENDING")} + + + ))} + + )} + + + + ); +}; + +export default function TeamSlide(): SlideConfig { + return { + key: "team", + title: , + body: , + background: TEAM_BACKGROUND, + }; +} diff --git a/frontend/editor/src/cloud/components/onboarding/slides/UsageSnapshotSlide.tsx b/frontend/editor/src/cloud/components/onboarding/slides/UsageSnapshotSlide.tsx new file mode 100644 index 0000000000..60a895d9f3 --- /dev/null +++ b/frontend/editor/src/cloud/components/onboarding/slides/UsageSnapshotSlide.tsx @@ -0,0 +1,45 @@ +import React from "react"; +import { useTranslation } from "react-i18next"; +import { SlideConfig } from "@app/types/types"; +import { createLightSlideBackground } from "@app/components/onboarding/slides/unifiedBackgroundConfig"; +import { + FreeMeterPanel, + useFreeSnapshot, +} from "@app/components/shared/config/configSections/usageMeters"; +import i18n from "@app/i18n"; +import styles from "@app/components/onboarding/slides/SaasOnboardingSlides.module.css"; + +const USAGE_BACKGROUND = createLightSlideBackground([249, 115, 22], "#FFEDD5"); + +const UsageSnapshotBody = () => { + const { t } = useTranslation(); + const snap = useFreeSnapshot(); + + return ( + + {t( + "onboarding.saas.usage.body", + "Automations, AI and API requests draw from your free allowance. Manual editing never counts against it.", + )} + {/* .payg provides the CSS variables the meter styles are scoped to */} + + + + + ); +}; + +export default function UsageSnapshotSlide(): SlideConfig { + return { + key: "usage-snapshot", + title: i18n.t( + "onboarding.saas.usage.title", + "Your free Processor allowance", + ), + body: , + background: USAGE_BACKGROUND, + }; +} diff --git a/frontend/editor/src/saas/components/onboarding/useSaasOnboardingState.ts b/frontend/editor/src/cloud/components/onboarding/useSaasOnboardingState.ts similarity index 77% rename from frontend/editor/src/saas/components/onboarding/useSaasOnboardingState.ts rename to frontend/editor/src/cloud/components/onboarding/useSaasOnboardingState.ts index e0728556f6..0b76c38a8f 100644 --- a/frontend/editor/src/saas/components/onboarding/useSaasOnboardingState.ts +++ b/frontend/editor/src/cloud/components/onboarding/useSaasOnboardingState.ts @@ -1,6 +1,8 @@ import { useCallback, useEffect, useMemo, useState, useRef } from "react"; import { useAuth } from "@app/auth/UseSession"; import { useOs } from "@app/hooks/useOs"; +import { useWallet } from "@app/hooks/useWallet"; +import { useSaaSTeam } from "@app/contexts/SaaSTeamContext"; import { SLIDE_DEFINITIONS, type ButtonAction, @@ -9,6 +11,7 @@ import { } from "@app/components/onboarding/saasOnboardingFlowConfig"; import { resolveSaasFlow } from "@app/components/onboarding/saasFlowResolver"; import { DOWNLOAD_URLS } from "@app/constants/downloads"; +import { openExternal } from "@app/platform/openExternal"; interface UseSaasOnboardingStateResult { currentStep: number; @@ -22,13 +25,22 @@ interface UseSaasOnboardingStateResult { interface UseSaasOnboardingStateProps { opened: boolean; onClose: () => void; + /** + * Drop the closing "desktop-install" slide. The desktop app reuses this + * flow but has no reason to pitch its own download. Defaults to false + * (slide shown) so the web (saas) flow is unchanged. + */ + hideDesktopInstall?: boolean; } export function useSaasOnboardingState({ opened, onClose, + hideDesktopInstall = false, }: UseSaasOnboardingStateProps): UseSaasOnboardingStateResult | null { - const { trialStatus, isPro, loading } = useAuth(); + const { loading } = useAuth(); + const { wallet } = useWallet(); + const { isTeamLeader } = useSaaSTeam(); const osType = useOs(); const selectedDownloadUrlRef = useRef(""); @@ -70,13 +82,16 @@ export function useSaasOnboardingState({ selectedDownloadUrlRef.current = url; }, []); - // Resolve flow based on trial status - const resolvedFlow = useMemo( - () => resolveSaasFlow(trialStatus, isPro), - [trialStatus, isPro], - ); + // Usage meter only makes sense for free-tier wallets with allowance left; + // the team slide is for leaders (anonymous guests are never leaders). + const showUsageSlide = wallet?.status === "free" && wallet.freeRemaining > 0; + const showTeamSlide = isTeamLeader; - const flowSlideIds = resolvedFlow.ids; + const flowSlideIds = useMemo( + () => + resolveSaasFlow({ showUsageSlide, showTeamSlide, hideDesktopInstall }), + [showUsageSlide, showTeamSlide, hideDesktopInstall], + ); const totalSteps = flowSlideIds.length; const maxIndex = Math.max(totalSteps - 1, 0); @@ -99,16 +114,8 @@ export function useSaasOnboardingState({ osUrl: os.url, osOptions, onDownloadUrlChange: handleDownloadUrlChange, - trialStatus: trialStatus ?? undefined, }); - }, [ - slideDefinition, - os.label, - os.url, - osOptions, - handleDownloadUrlChange, - trialStatus, - ]); + }, [slideDefinition, os.label, os.url, osOptions, handleDownloadUrlChange]); // Navigation functions const goNext = useCallback(() => { @@ -138,10 +145,11 @@ export function useSaasOnboardingState({ onClose(); return; case "download-selected": { - // Open download URL in new tab + // Open the download URL in the user's browser via the platform seam + // (saas opens a new tab, desktop hands off to the OS shell). const downloadUrl = selectedDownloadUrlRef.current || os.url; if (downloadUrl) { - window.open(downloadUrl, "_blank", "noopener,noreferrer"); + void openExternal(downloadUrl); } // Then advance to next slide or close if last if (currentStep === maxIndex) { diff --git a/frontend/editor/src/cloud/components/shared/FreeLimitReachedModal.tsx b/frontend/editor/src/cloud/components/shared/FreeLimitReachedModal.tsx new file mode 100644 index 0000000000..7db30887e7 --- /dev/null +++ b/frontend/editor/src/cloud/components/shared/FreeLimitReachedModal.tsx @@ -0,0 +1,213 @@ +import { useMemo } from "react"; +import { Modal, Stack, Button } from "@mantine/core"; +import { useTranslation } from "react-i18next"; +import CelebrationIcon from "@mui/icons-material/CelebrationOutlined"; +import AnimatedSlideBackground from "@app/components/onboarding/slides/AnimatedSlideBackground"; +import styles from "@app/components/onboarding/InitialOnboardingModal/InitialOnboardingModal.module.css"; +import { Z_INDEX_OVER_FULLSCREEN_SURFACE } from "@app/styles/zIndex"; +import { navigateToSettings } from "@app/utils/settingsNavigation"; +import { + FreeMeterPanel, + freeSnapshotFromWallet, +} from "@app/components/shared/config/configSections/usageMeters"; +import { useWallet } from "@app/hooks/useWallet"; + +interface FreeLimitReachedModalProps { + onClose: () => void; +} + +function readColor(varName: string, fallback: string): string { + return ( + getComputedStyle(document.documentElement) + .getPropertyValue(varName) + .trim() || fallback + ); +} + +export function FreeLimitReachedModal({ onClose }: FreeLimitReachedModalProps) { + const { t } = useTranslation(); + const { wallet, loading } = useWallet(); + + // Resolve theme colours once; reading the CSS vars on every render would + // force a style recalc. + const gradientStops = useMemo<[string, string]>( + () => [ + readColor("--color-primary-500", "#3b82f6"), + readColor("--color-primary-800", "#1e40af"), + ], + [], + ); + + // Hold the modal back until the wallet resolves so the meter never flashes + // placeholder numbers before the real ones land. + if (loading || !wallet) return null; + const snap = freeSnapshotFromWallet(wallet); + + const circles = [ + { + position: "bottom-left" as const, + size: 270, + color: "rgba(255, 255, 255, 0.25)", + opacity: 0.9, + amplitude: 24, + duration: 4.5, + offsetX: 18, + offsetY: 14, + }, + { + position: "top-right" as const, + size: 300, + color: "rgba(255, 255, 255, 0.2)", + opacity: 0.9, + amplitude: 28, + duration: 4.5, + delay: 0.5, + offsetX: 24, + offsetY: 18, + }, + ]; + + const handleUpgrade = () => { + onClose(); + navigateToSettings("plan"); + }; + + return ( + + +
+ +
+
+ +
+
+
+ +
+ +
+ {t("plan.freeLimit.title", "Woah, {{total}} PDFs Processed!", { + total: snap.billableUsed.toLocaleString(), + })} +
+ +
+
+ {t( + "plan.freeLimit.message", + "That's your whole free allowance for automation, AI and the API. Seriously impressive! Keep the momentum going for just pennies a day.", + )} +
+
+ + + +
+ +
+ + + +
+
+
+
+
+
+ ); +} diff --git a/frontend/editor/src/saas/components/shared/TrialExpiredModal.tsx b/frontend/editor/src/cloud/components/shared/SpendCapReachedModal.tsx similarity index 61% rename from frontend/editor/src/saas/components/shared/TrialExpiredModal.tsx rename to frontend/editor/src/cloud/components/shared/SpendCapReachedModal.tsx index 4f80a696b1..c124b9baec 100644 --- a/frontend/editor/src/saas/components/shared/TrialExpiredModal.tsx +++ b/frontend/editor/src/cloud/components/shared/SpendCapReachedModal.tsx @@ -1,65 +1,82 @@ +import { useMemo } from "react"; import { Modal, Stack, Button } from "@mantine/core"; import { useTranslation } from "react-i18next"; -import DiamondOutlinedIcon from "@mui/icons-material/DiamondOutlined"; +import TrendingUpIcon from "@mui/icons-material/TrendingUpOutlined"; import AnimatedSlideBackground from "@app/components/onboarding/slides/AnimatedSlideBackground"; import styles from "@app/components/onboarding/InitialOnboardingModal/InitialOnboardingModal.module.css"; import { Z_INDEX_OVER_FULLSCREEN_SURFACE } from "@app/styles/zIndex"; +import { navigateToSettings } from "@app/utils/settingsNavigation"; +import { + SpendCapMeterPanel, + spendCapSnapshotFromWallet, +} from "@app/components/shared/config/configSections/usageMeters"; +import { useWallet } from "@app/hooks/useWallet"; -interface TrialExpiredModalProps { - opened: boolean; +interface SpendCapReachedModalProps { onClose: () => void; - onSubscribe: () => void; } -export function TrialExpiredModal({ - opened, - onClose, - onSubscribe, -}: TrialExpiredModalProps) { - const { t } = useTranslation(); +function readColor(varName: string, fallback: string): string { + return ( + getComputedStyle(document.documentElement) + .getPropertyValue(varName) + .trim() || fallback + ); +} - // Use CSS variables for theme colors - const amberColor = - getComputedStyle(document.documentElement) - .getPropertyValue("--color-amber-500") - .trim() || "#f59e0b"; - const redColor = - getComputedStyle(document.documentElement) - .getPropertyValue("--color-red-500") - .trim() || "#ef4444"; - const gradientStops: [string, string] = [amberColor, redColor]; +export function SpendCapReachedModal({ onClose }: SpendCapReachedModalProps) { + const { t } = useTranslation(); + const { wallet, loading } = useWallet(); + + // Resolve theme colours once; reading the CSS vars on every render would + // force a style recalc. + const gradientStops = useMemo<[string, string]>( + () => [ + readColor("--color-green-500", "#22c55e"), + readColor("--color-green-700", "#15803d"), + ], + [], + ); + + // Hold the modal back until the wallet resolves so the meter never flashes + // placeholder numbers before the real ones land. + if (loading || !wallet) return null; + const snap = spendCapSnapshotFromWallet(wallet); const circles = [ { position: "bottom-left" as const, - size: 270, // 16.875rem + size: 270, color: "rgba(255, 255, 255, 0.25)", opacity: 0.9, - amplitude: 24, // 1.5rem + amplitude: 24, duration: 4.5, - offsetX: 18, // 1.125rem - offsetY: 14, // 0.875rem + offsetX: 18, + offsetY: 14, }, { position: "top-right" as const, - size: 300, // 18.75rem + size: 300, color: "rgba(255, 255, 255, 0.2)", opacity: 0.9, - amplitude: 28, // 1.75rem + amplitude: 28, duration: 4.5, delay: 0.5, - offsetX: 24, // 1.5rem - offsetY: 18, // 1.125rem + offsetX: 24, + offsetY: 18, }, ]; + const handleRaiseCap = () => { + onClose(); + navigateToSettings("plan"); + }; + return ( {}} // Prevent closing by clicking outside or ESC + opened + onClose={onClose} withCloseButton={false} - closeOnClickOutside={false} - closeOnEscape={false} centered size="lg" radius="lg" @@ -91,11 +108,11 @@ export function TrialExpiredModal({ gradientStops={gradientStops} circles={circles} isActive - slideKey="trial-expired" + slideKey="spend-cap-reached" />
- +
@@ -111,43 +128,36 @@ export function TrialExpiredModal({ >
- {t("plan.trial.expired", "Your Trial Has Ended")} + {t("plan.spendCap.title", "You're on a Roll!")}
{t( - "plan.trial.expiredMessage", - "Your 30-day Pro trial has expired. Subscribe to Pro to continue accessing premium features, or continue with our free tier.", + "plan.spendCap.message", + "You've made the most of this month's cap. That's a load of automation, AI and API work! Bump it up whenever you like to keep going.", )}
-
-
- {t( - "plan.trial.freeTierLimitations", - "Free tier includes basic PDF tools with usage limits.", - )} -
-
+
- {t("plan.trial.continueWithFree", "Continue with Free")} + {t("plan.spendCap.dismiss", "Not Now")}
diff --git a/frontend/editor/src/desktop/components/shared/TeamInvitationBanner.tsx b/frontend/editor/src/cloud/components/shared/TeamInvitationBanner.tsx similarity index 70% rename from frontend/editor/src/desktop/components/shared/TeamInvitationBanner.tsx rename to frontend/editor/src/cloud/components/shared/TeamInvitationBanner.tsx index 5389274986..aa07d823df 100644 --- a/frontend/editor/src/desktop/components/shared/TeamInvitationBanner.tsx +++ b/frontend/editor/src/cloud/components/shared/TeamInvitationBanner.tsx @@ -1,52 +1,34 @@ -import { useState, useEffect } from "react"; +import { useState } from "react"; import { Button, Group, Text } from "@mantine/core"; import { useTranslation } from "react-i18next"; import LocalIcon from "@app/components/shared/LocalIcon"; import { InfoBanner } from "@app/components/shared/InfoBanner"; import { useSaaSTeam } from "@app/contexts/SaaSTeamContext"; -import { useSaaSBilling } from "@app/contexts/SaasBillingContext"; -import { connectionModeService } from "@app/services/connectionModeService"; +/** + * SaaS-web team invitation banner. Shown at the top of the app when the + * signed-in user has a pending invitation to join a team. + * + * Ported from the desktop banner, with two differences: there is no + * {@code connectionMode} gate (web is always SaaS), and there is no explicit + * billing refresh — {@link useSaaSTeam.acceptInvitation} already refreshes + * credits and the session after the team membership changes. + */ export function TeamInvitationBanner() { const { t } = useTranslation(); const { receivedInvitations, acceptInvitation, rejectInvitation } = useSaaSTeam(); - const { refreshBilling } = useSaaSBilling(); const [processing, setProcessing] = useState(false); const [dismissed, setDismissed] = useState(false); - const [connectionMode, setConnectionMode] = useState(null); - // Load connection mode on mount - useEffect(() => { - connectionModeService - .getCurrentMode() - .then((mode) => setConnectionMode(mode)); - }, []); + const invitation = receivedInvitations[0]; // Show first invitation - // Accept invitation handler const handleAccept = async () => { - const invitation = receivedInvitations[0]; if (!invitation) return; - setProcessing(true); - try { await acceptInvitation(invitation.invitationToken); - console.log( - "[TeamInvitationBanner] Invitation accepted successfully:", - invitation.teamName, - ); - - // Wait briefly for backend to process team membership update - await new Promise((resolve) => setTimeout(resolve, 1000)); - - // Refresh billing after joining team (tier may have changed) - console.log( - "[TeamInvitationBanner] Refreshing billing after team join...", - ); - await refreshBilling(); - setDismissed(true); } catch (error) { console.error( @@ -58,16 +40,11 @@ export function TeamInvitationBanner() { } }; - // Reject invitation handler const handleReject = async () => { - const invitation = receivedInvitations[0]; if (!invitation) return; - setProcessing(true); - try { await rejectInvitation(invitation.invitationToken); - console.log("[TeamInvitationBanner] Invitation rejected"); setDismissed(true); } catch (error) { console.error( @@ -79,14 +56,9 @@ export function TeamInvitationBanner() { } }; - // Visibility logic - const shouldShow = - connectionMode === "saas" && !dismissed && receivedInvitations.length > 0; - + const shouldShow = !dismissed && receivedInvitations.length > 0; if (!shouldShow) return null; - const invitation = receivedInvitations[0]; // Show first invitation - const message = ( ; + +/** The Plan (billing) nav item — wallet-driven PAYG dashboard + spend cap. */ +export function createCloudPlanNavItem(t: Translate): ConfigNavItem { + return { + key: "plan", + label: t("config.plan", "Plan"), + icon: "credit-card", + component: , + }; +} + +/** The Team nav item — shared SaaS team management (invite/rename/members). */ +export function createCloudTeamNavItem(t: Translate): ConfigNavItem { + return { + key: "teams", + label: t("config.team", "Team"), + icon: "groups-rounded", + component: , + }; +} + +/** Billing nav section wrapping the Plan item, for leaves that group it (saas). */ +export function createCloudBillingSection(t: Translate): ConfigNavSection { + return { + title: "Billing", + items: [createCloudPlanNavItem(t)], + }; +} diff --git a/frontend/editor/src/cloud/components/shared/config/configSections/Payg.css b/frontend/editor/src/cloud/components/shared/config/configSections/Payg.css new file mode 100644 index 0000000000..8e1921d061 --- /dev/null +++ b/frontend/editor/src/cloud/components/shared/config/configSections/Payg.css @@ -0,0 +1,624 @@ +/* Pay-as-you-go settings screen — visual polish. + Uses the app's semantic theme tokens (theme.css) so it tracks light/dark. */ + +.payg { + --payg-accent: #0a8bff; + --payg-accent-soft: rgba(10, 139, 255, 0.1); + --payg-radius: 14px; + --payg-gap: 20px; + /* Cards sit ON the modal content bg; in light that's white-on-white so we + lean on border + shadow. Tokens are overridden per scheme below. */ + --payg-card-bg: var(--bg-surface); + --payg-card-border: var(--border-default); + --payg-inset-bg: var(--bg-raised); + --payg-divider: var(--border-subtle); +} + +/* Dark mode: the modal content bg is #2a2f36 and so is --bg-surface, so plain + cards vanish. Lift cards a shade above the modal and strengthen borders. */ +[data-mantine-color-scheme="dark"] .payg { + --payg-card-bg: #313842; + --payg-card-border: #3d444e; + --payg-inset-bg: #272c33; + --payg-accent-soft: rgba(10, 139, 255, 0.16); + --payg-divider: #3d444e; +} + +/* ── Header ─────────────────────────────────────────────────────────── */ +.payg-header__title { + margin: 0; + font-size: 1.35rem; + font-weight: 700; + letter-spacing: -0.01em; + color: var(--text-primary); +} +.payg-header__subtitle { + margin-top: 4px; + font-size: 0.875rem; + color: var(--text-muted); +} +.payg-role-pill { + display: inline-flex; + align-items: center; + gap: 6px; + padding: 4px 12px; + border-radius: 999px; + font-size: 0.75rem; + font-weight: 600; + white-space: nowrap; +} +.payg-role-pill[data-leader="true"] { + background: var(--payg-accent-soft); + color: var(--payg-accent); + border: 1px solid rgba(10, 139, 255, 0.25); +} +.payg-role-pill[data-leader="false"] { + background: var(--bg-muted); + color: var(--text-muted); + border: 1px solid var(--border-default); +} + +/* ── Plan header: free-vs-metered split (shared by free + subscribed) ──── */ +.payg-planhead { + padding: 14px 20px; + border-radius: var(--payg-radius); + background: var(--payg-card-bg); + border: 1px solid var(--payg-card-border); +} +.payg-planhead__top { + display: flex; + align-items: center; + justify-content: space-between; + gap: 12px; + margin-bottom: 12px; +} +.payg-planhead__eyebrow { + font-size: 0.78rem; + color: var(--text-muted); +} +.payg-planhead__split { + display: grid; + grid-template-columns: minmax(0, 1.5fr) minmax(0, 1fr); +} +.payg-planhead__col { + padding-right: 22px; +} +.payg-planhead__col--meter { + padding-right: 0; + padding-left: 22px; + border-left: 1px solid var(--payg-divider); +} +.payg-planhead__lbl { + display: inline-flex; + align-items: center; + gap: 6px; + font-size: 0.72rem; + font-weight: 700; + letter-spacing: 0.06em; + text-transform: uppercase; + margin-bottom: 8px; +} +.payg-planhead__lbl--free { + color: #10b981; +} +.payg-planhead__lbl--meter { + color: var(--payg-accent); +} +.payg-planhead__lbl-icon { + font-size: 1rem !important; +} +.payg-planhead__title { + margin: 0; + font-size: 1.05rem; + font-weight: 700; + color: var(--text-primary); + letter-spacing: -0.01em; + line-height: 1.25; +} +.payg-planhead__body { + margin: 5px 0 0; + font-size: 0.85rem; + color: var(--text-muted); + line-height: 1.5; +} +@media (max-width: 640px) { + .payg-planhead__split { + grid-template-columns: 1fr; + gap: 16px; + } + .payg-planhead__col { + padding-right: 0; + } + .payg-planhead__col--meter { + padding-left: 0; + padding-top: 16px; + border-left: none; + border-top: 1px solid var(--payg-divider); + } +} + +/* ── Generic card ───────────────────────────────────────────────────── */ +.payg-card { + background: var(--payg-card-bg); + border: 1px solid var(--payg-card-border); + border-radius: var(--payg-radius); + padding: 16px 20px; + box-shadow: var(--shadow-xs); +} +.payg-card__title { + font-size: 0.95rem; + font-weight: 650; + color: var(--text-primary); +} +.payg-card__subtitle { + font-size: 0.8125rem; + color: var(--text-muted); + margin-top: 2px; +} + +/* ── Hero usage panel ───────────────────────────────────────────────── */ +.payg-hero { + position: relative; + overflow: hidden; + border-radius: var(--payg-radius); + padding: 18px 22px; + border: 1px solid var(--payg-card-border); + background: var(--payg-card-bg); + box-shadow: var(--shadow-xs); +} +.payg-hero::before { + content: ""; + position: absolute; + inset: 0; + pointer-events: none; + opacity: 0.9; +} +.payg-hero[data-state="FULL"]::before { + background: linear-gradient( + 135deg, + rgba(10, 139, 255, 0.12) 0%, + rgba(10, 139, 255, 0) 55% + ); +} +.payg-hero[data-state="WARNED"]::before { + background: linear-gradient( + 135deg, + rgba(234, 179, 8, 0.16) 0%, + rgba(234, 179, 8, 0) 55% + ); +} +.payg-hero[data-state="DEGRADED"]::before { + background: linear-gradient( + 135deg, + rgba(239, 68, 68, 0.16) 0%, + rgba(239, 68, 68, 0) 55% + ); +} +.payg-hero__inner { + position: relative; + z-index: 1; +} +.payg-hero__eyebrow { + font-size: 0.6875rem; + font-weight: 700; + letter-spacing: 0.08em; + text-transform: uppercase; + color: var(--text-muted); +} +.payg-hero__figure { + display: flex; + align-items: baseline; + gap: 10px; + margin-top: 6px; +} +.payg-hero__spend { + font-size: 2.35rem; + font-weight: 750; + line-height: 1; + letter-spacing: -0.03em; + color: var(--text-primary); + font-variant-numeric: tabular-nums; +} +.payg-hero__cap { + font-size: 1rem; + color: var(--text-muted); + font-variant-numeric: tabular-nums; +} +.payg-hero__meta { + margin-top: 14px; + display: flex; + flex-wrap: wrap; + gap: 6px 18px; + font-size: 0.8125rem; + color: var(--text-secondary); +} +.payg-hero__meta-dot { + color: var(--border-strong); +} +.payg-hero__credit { + display: inline-flex; + align-items: center; + padding: 2px 9px; + border-radius: 999px; + font-size: 0.75rem; + font-weight: 600; + color: var(--color-green-700); + background: rgba(34, 197, 94, 0.14); +} +[data-mantine-color-scheme="dark"] .payg-hero__credit { + color: #4ade80; + background: rgba(34, 197, 94, 0.18); +} + +/* Make Mantine default-variant buttons read clearly on the lifted cards in + both themes (their stock border can wash out against --payg-card-bg). */ +.payg .mantine-Button-root[data-variant="default"] { + border-color: var(--payg-card-border); +} + +/* Dark mode: the stock default button is a flat slab that disappears into the + card. Lift it with a soft top-down gradient + brighter border, and warm the + hover so the buttons feel tactile against --payg-card-bg (#313842). Driven + through Mantine's --button-* vars. The same rule covers disabled buttons + (Cancel/Update cap before the form is dirty) so they read identically to the + always-enabled ones — Mantine only fades them via opacity. */ +[data-mantine-color-scheme="dark"] + .payg + .mantine-Button-root[data-variant="default"] { + --button-bg: linear-gradient(180deg, #3d4651 0%, #353d48 100%); + --button-hover: linear-gradient(180deg, #475160 0%, #3d4654 100%); + --button-bd: 1px solid #4b5563; + --button-color: var(--mantine-color-white); + box-shadow: + 0 1px 0 rgba(255, 255, 255, 0.04) inset, + 0 1px 2px rgba(0, 0, 0, 0.25); +} +/* Mantine recolours disabled buttons to a dim grey via a hard-coded color/bg + on the disabled selector, overriding --button-color. Re-assert white + the + gradient so disabled buttons match the rest. */ +[data-mantine-color-scheme="dark"] + .payg + .mantine-Button-root[data-variant="default"]:disabled, +[data-mantine-color-scheme="dark"] + .payg + .mantine-Button-root[data-variant="default"][data-disabled] { + background: var(--button-bg); + color: var(--mantine-color-white); +} + +/* Status chip */ +.payg-status { + display: inline-flex; + align-items: center; + gap: 7px; + padding: 6px 13px; + border-radius: 999px; + font-size: 0.8125rem; + font-weight: 600; + white-space: nowrap; +} +.payg-status__dot { + width: 8px; + height: 8px; + border-radius: 999px; +} +.payg-status[data-state="FULL"] { + background: var(--color-green-100); + color: var(--color-green-700); +} +.payg-status[data-state="FULL"] .payg-status__dot { + background: var(--color-green-500); + box-shadow: 0 0 0 3px rgba(34, 197, 94, 0.18); +} +.payg-status[data-state="WARNED"] { + background: var(--color-yellow-100); + color: var(--color-yellow-700); +} +.payg-status[data-state="WARNED"] .payg-status__dot { + background: var(--color-yellow-500); + box-shadow: 0 0 0 3px rgba(234, 179, 8, 0.2); +} +.payg-status[data-state="DEGRADED"] { + background: var(--color-red-100); + color: var(--color-red-700); +} +.payg-status[data-state="DEGRADED"] .payg-status__dot { + background: var(--color-red-500); + box-shadow: 0 0 0 3px rgba(239, 68, 68, 0.2); +} + +/* Segmented usage bar */ +.payg-bar { + margin-top: 18px; + height: 10px; + border-radius: 999px; + background: var(--bg-muted); + overflow: hidden; + position: relative; +} +.payg-bar__fill { + height: 100%; + border-radius: 999px; + transition: width 0.5s cubic-bezier(0.16, 1, 0.3, 1); +} +.payg-bar__fill[data-state="FULL"] { + background: linear-gradient(90deg, #0a8bff, #38bdf8); +} +.payg-bar__fill[data-state="WARNED"] { + background: linear-gradient(90deg, #f59e0b, #fbbf24); +} +.payg-bar__fill[data-state="DEGRADED"] { + background: linear-gradient(90deg, #dc2626, #f87171); +} + +/* ── "What counts as a document?" expandable help ───────────────────── */ +.payg-help { + margin-top: 12px; +} +.payg-help__toggle { + display: inline-flex; + align-items: center; + gap: 5px; + padding: 0; + border: none; + background: none; + cursor: pointer; + font: inherit; + font-size: 0.8125rem; + font-weight: 600; + color: var(--text-secondary); +} +.payg-help__toggle:hover { + color: var(--text-primary); +} +.payg-help__chevron { + transition: transform 0.15s ease; +} +.payg-help__toggle[aria-expanded="true"] .payg-help__chevron { + transform: rotate(180deg); +} +.payg-help__panel { + margin-top: 10px; + padding: 12px 14px; + border-radius: 10px; + border: 1px solid var(--payg-card-border); + background: var(--bg-muted); + font-size: 0.8125rem; + line-height: 1.45; + color: var(--text-secondary); +} +.payg-help__panel ul { + margin: 0; + padding-left: 18px; + display: grid; + gap: 6px; +} + +/* ── Cap preview strip ──────────────────────────────────────────────── */ +.payg-preview { + display: flex; + align-items: center; + gap: 12px; + padding: 13px 16px; + border-radius: 10px; + background: var(--payg-accent-soft); + border: 1px solid rgba(10, 139, 255, 0.2); +} +.payg-preview__icon { + color: var(--payg-accent); + display: flex; +} +.payg-preview__main { + font-size: 0.875rem; + color: var(--text-primary); + font-weight: 550; +} +.payg-preview__note { + font-size: 0.75rem; + color: var(--text-muted); + margin-top: 1px; +} + +/* ── Gates grid ─────────────────────────────────────────────────────── */ +.payg-gates { + display: grid; + grid-template-columns: repeat(2, 1fr); + gap: 10px; + margin-top: 4px; +} +@media (max-width: 720px) { + .payg-gates { + grid-template-columns: 1fr; + } +} +.payg-gate { + display: flex; + align-items: center; + gap: 12px; + padding: 13px 15px; + border-radius: 11px; + border: 1px solid var(--payg-card-border); + background: var(--payg-inset-bg); +} +.payg-gate[data-enabled="false"] { + border-style: dashed; + opacity: 0.85; +} +.payg-gate__chip { + display: flex; + align-items: center; + justify-content: center; + width: 32px; + height: 32px; + border-radius: 9px; + flex-shrink: 0; +} +.payg-gate[data-enabled="true"] .payg-gate__chip { + background: rgba(34, 197, 94, 0.16); + color: #22c55e; +} +.payg-gate[data-enabled="false"] .payg-gate__chip { + background: rgba(239, 68, 68, 0.16); + color: #f87171; +} +.payg-gate__label { + font-size: 0.8125rem; + line-height: 1.35; + color: var(--text-primary); + flex: 1; +} +.payg-gate__tag { + font-size: 0.625rem; + font-weight: 700; + letter-spacing: 0.05em; + text-transform: uppercase; + padding: 2px 8px; + border-radius: 999px; + flex-shrink: 0; + white-space: nowrap; +} +.payg-gate__tag[data-variant="on"] { + color: var(--text-muted); + background: var(--bg-muted); +} +.payg-gate__tag[data-variant="pause"] { + color: #dc2626; + background: rgba(239, 68, 68, 0.12); +} +[data-mantine-color-scheme="dark"] .payg-gate__tag[data-variant="pause"] { + color: #f87171; + background: rgba(239, 68, 68, 0.18); +} + +/* ── Member rows ────────────────────────────────────────────────────── */ +.payg-member { + display: flex; + align-items: center; + gap: 14px; + padding: 12px 4px; + border-bottom: 1px solid var(--payg-divider); +} +.payg-member:last-child { + border-bottom: none; +} +.payg-member__avatar { + width: 34px; + height: 34px; + border-radius: 999px; + display: flex; + align-items: center; + justify-content: center; + font-size: 0.8125rem; + font-weight: 650; + color: #fff; + flex-shrink: 0; +} +.payg-member__name { + font-size: 0.875rem; + font-weight: 550; + color: var(--text-primary); +} +.payg-member__email { + font-size: 0.75rem; + color: var(--text-muted); +} +.payg-member__usage { + text-align: right; + min-width: 120px; +} +.payg-member__usage-num { + font-size: 0.8125rem; + color: var(--text-primary); + font-variant-numeric: tabular-nums; +} +.payg-member__minibar { + height: 5px; + width: 90px; + border-radius: 999px; + background: var(--payg-inset-bg); + overflow: hidden; + margin-top: 5px; + margin-left: auto; +} +.payg-member__minibar-fill { + height: 100%; + border-radius: 999px; + background: var(--payg-accent); +} + +/* ── Activity feed ──────────────────────────────────────────────────── */ +.payg-activity-row { + display: flex; + align-items: center; + gap: 14px; + padding: 11px 4px; + border-bottom: 1px solid var(--payg-divider); +} +.payg-activity-row:last-child { + border-bottom: none; +} +.payg-activity__dot { + width: 9px; + height: 9px; + border-radius: 999px; + flex-shrink: 0; +} +.payg-activity__dot[data-kind="ai"] { + background: var(--category-color-formatting); /* purple */ +} +.payg-activity__dot[data-kind="automation"] { + background: var(--category-color-automation); /* pink */ +} +.payg-activity__label { + font-size: 0.8438rem; + color: var(--text-primary); + flex: 1; +} +.payg-activity__ts { + font-size: 0.75rem; + color: var(--text-muted); +} +.payg-activity__kind { + font-size: 0.625rem; + font-weight: 700; + letter-spacing: 0.05em; + text-transform: uppercase; + color: var(--text-muted); + background: var(--bg-muted); + padding: 2px 8px; + border-radius: 999px; +} +.payg-activity__units { + font-size: 0.8125rem; + font-weight: 650; + color: var(--text-primary); + font-variant-numeric: tabular-nums; + min-width: 58px; + text-align: right; +} + +/* ── Stripe CTA card ────────────────────────────────────────────────── */ +.payg-stripe { + display: flex; + align-items: center; + justify-content: space-between; + gap: 16px; + padding: 20px 24px; + border-radius: var(--payg-radius); + border: 1px solid rgba(99, 91, 255, 0.28); + background: linear-gradient( + 135deg, + rgba(99, 91, 255, 0.1) 0%, + rgba(99, 91, 255, 0.02) 100% + ); +} +.payg-stripe__title { + font-size: 0.9375rem; + font-weight: 650; + color: var(--text-primary); +} +.payg-stripe__subtitle { + font-size: 0.8125rem; + color: var(--text-muted); + margin-top: 2px; +} diff --git a/frontend/editor/src/cloud/components/shared/config/configSections/Payg.tsx b/frontend/editor/src/cloud/components/shared/config/configSections/Payg.tsx new file mode 100644 index 0000000000..ea0e7f3d32 --- /dev/null +++ b/frontend/editor/src/cloud/components/shared/config/configSections/Payg.tsx @@ -0,0 +1,804 @@ +/** + * Pay-as-you-go billing & usage section — the SUBSCRIBED leader/member views. + * + * All data comes from the real {@code Wallet} snapshot ({@code GET + * /api/v1/payg/wallet} via {@link useWallet}); there is no mock fallback. What + * the wallet doesn't provide, the UI doesn't show: + * + * - spend + cap render in the units/USD the backend actually returns + * ({@code spendUnitsThisPeriod}, {@code capUsd}, {@code noCap}) + * - the per-category breakdown comes from {@code categoryBreakdown} + * (wallet_category_summary view) + * - money-equivalent display and the units↔money cap preview need + * stripe.prices via Sync Engine and are deliberately absent until that + * ships + * - the activity feed renders {@code wallet.recent}, which is {@code []} in + * V1 — so it shows a real empty state, not fabricated rows + */ +import React, { useState } from "react"; +import { Button, Group, Stack, Text } from "@mantine/core"; +import { useRenderCount } from "@app/hooks/useRenderCount"; +import OpenInNewIcon from "@mui/icons-material/OpenInNew"; +import LockIcon from "@mui/icons-material/LockOutlined"; +import CheckIcon from "@mui/icons-material/CheckRounded"; +import HelpOutlineIcon from "@mui/icons-material/HelpOutlineRounded"; +import ExpandMoreIcon from "@mui/icons-material/ExpandMoreRounded"; +import BoltIcon from "@mui/icons-material/BoltRounded"; +import AllInclusiveIcon from "@mui/icons-material/AllInclusiveRounded"; +import { alert as showToast } from "@app/components/toast"; +// Relative (not @app/*) so the co-located CSS + sibling component resolve directly. +// eslint-disable-next-line no-restricted-imports +import "./Payg.css"; +// eslint-disable-next-line no-restricted-imports +import SpendCapControl from "./SpendCapControl"; +import { useTranslation } from "react-i18next"; +import type { Wallet } from "@app/hooks/useWallet"; + +// ─── Types ──────────────────────────────────────────────────────────────── + +type Gate = "OFFSITE_PROCESSING" | "AUTOMATION" | "AI_SUPPORT" | "CLIENT_SIDE"; + +interface PaygProps { + role: "LEADER" | "MEMBER"; + /** Real wallet snapshot from {@code useWallet}. Single source of truth. */ + wallet: Wallet; + /** + * Persist a cap change. Provided by {@code Plan} → {@code useWallet} for the + * leader view; absent on the member view (read-only). + */ + onSaveCap?: (capUsd: number | null) => Promise | void; + /** + * Open the Stripe Customer Portal. When omitted the Stripe card is hidden. + * On error the implementation shows a friendly toast and resolves — callers + * don't need to wrap in try/catch. + */ + onOpenPortal?: () => Promise; +} + +// ─── Helpers ────────────────────────────────────────────────────────────── + +// Stable-ish avatar colour from a string. +const AVATAR_COLORS = [ + "#0a8bff", + "#8b5cf6", + "#ec4899", + "#10b981", + "#f59e0b", + "#06b6d4", +]; +function avatarColor(seed: string): string { + let h = 0; + for (let i = 0; i < seed.length; i++) h = (h * 31 + seed.charCodeAt(i)) | 0; + return AVATAR_COLORS[Math.abs(h) % AVATAR_COLORS.length]; +} + +function gateLabel( + g: Gate, + t: (k: string, fallback: string) => string, +): string { + switch (g) { + case "OFFSITE_PROCESSING": + return t( + "payg.gates.offsite", + "Server tools (compress, OCR, convert, watermark…)", + ); + case "AUTOMATION": + return t("payg.gates.automation", "Automations & pipelines"); + case "AI_SUPPORT": + return t("payg.gates.ai", "AI tools (AI Create, suggestions, AI-OCR)"); + case "CLIENT_SIDE": + return t( + "payg.gates.client", + "Browser-only tools (viewer, page editor, file management)", + ); + } +} + +// ─── "What counts as a document?" help ────────────────────────────────────── + +/** + * Expandable explainer for the billing unit. Shared by the subscribed hero + * here and the free-tier hero in {@code PaygFree.tsx}. The bullets state the + * real charge mechanics (DefaultDocumentClassifier + JobChargeService) + * without hardcoding the policy-tunable thresholds: each non-empty file is + * at least one document; page count / file size can make it more; chained + * steps on the same file join the open process instead of re-charging; a + * first-step failure writes a compensating refund. + */ +export function DocHelp() { + const { t } = useTranslation(); + const [open, setOpen] = useState(false); + return ( +
+ + {open && ( +
+
    +
  • + {t( + "payg.docHelp.billable", + "Only automation runs, AI tools, and API calls count. Manual tools in the editor are always free.", + )} +
  • +
  • + {t( + "payg.docHelp.perFile", + "Each file you process counts as one PDF. Very long or very large files can count as more than one.", + )} +
  • +
  • + {t( + "payg.docHelp.chains", + "Running the same file through several steps of one automation counts it once, not once per step.", + )} +
  • +
  • + {t( + "payg.docHelp.refunds", + "If a job fails on its first step, the PDF is credited back automatically.", + )} +
  • +
+
+ )} +
+ ); +} + +// ─── Hero usage panel ─────────────────────────────────────────────────────── + +/** + * Format minor units of an ISO currency for display ("$2.24", "£0.40"). Only + * called when the backend resolved the rate — currency is always present + * alongside a non-null money amount. + */ +function formatMinor(minor: number, currency: string | null): string { + const code = (currency ?? "usd").toUpperCase(); + try { + return new Intl.NumberFormat(undefined, { + style: "currency", + currency: code, + }).format(minor / 100); + } catch { + return `${(minor / 100).toFixed(2)} ${code}`; + } +} + +/** Currency symbol for compact inline use; falls back to the ISO code. */ +function currencySymbol(currency: string | null): string { + switch ((currency ?? "").toLowerCase()) { + case "usd": + return "$"; + case "eur": + return "€"; + case "gbp": + return "£"; + default: + return currency ? currency.toUpperCase() + " " : "$"; + } +} + +/** + * Usage hero — gradient panel with the allowance bar. Every number is real: + * {@code billableUsed} is the ledger's period sum over the team's actual + * billing window (the Stripe subscription period); {@code billableLimit} is + * the backend-derived document ceiling (free allowance + what the money cap + * buys at the subscription Price's per-document rate, null when uncapped); + * {@code estimatedBillMinor} is spend beyond the allowance at that rate. + * Fields the backend couldn't resolve are null and simply not rendered. + */ +function UsageHero({ wallet }: { wallet: Wallet }) { + const { t } = useTranslation(); + + const periodEnd = new Date(wallet.billingPeriodEnd); + const daysLeft = Math.max( + 0, + Math.ceil((periodEnd.getTime() - Date.now()) / 86_400_000), + ); + const hasCap = !wallet.noCap && wallet.capUsd != null; + const breakdown = wallet.categoryBreakdown; + + const limit = wallet.billableLimit; + const hasLimit = limit != null && limit > 0; + const pct = hasLimit ? Math.min(100, (wallet.billableUsed / limit) * 100) : 0; + const state = hasLimit + ? pct >= 100 + ? "DEGRADED" + : pct >= 80 + ? "WARNED" + : "FULL" + : "FULL"; + const stateLabel = { + FULL: t("payg.state.full", "Healthy"), + WARNED: t("payg.state.warned", "Approaching cap"), + DEGRADED: t("payg.state.degraded", "Cap reached"), + }[state]; + + return ( +
+
+ +
+
+ {t("payg.usage.thisPeriod", "This billing period")} +
+
+ + {wallet.billableUsed.toLocaleString()} + + + {hasLimit + ? t( + "payg.usage.ofLimitProcessed", + "/ {{limit}} PDFs processed", + { limit: limit.toLocaleString() }, + ) + : t("payg.usage.processed", "PDFs processed")} + +
+
+
+ + {stateLabel} +
+
+ + {hasLimit && ( +
+
+
+ )} + +
+ + {t("payg.usage.firstFree", "First {{free}} free", { + free: wallet.freeAllowance.toLocaleString(), + })} + + {wallet.estimatedBillMinor != null && ( + <> + • + + {t("payg.usage.estBill", "≈ {{amount}} so far this period", { + amount: formatMinor( + wallet.estimatedBillMinor, + wallet.currency, + ), + })} + + + )} + • + + {hasCap + ? t("payg.usage.capLine", "{{cap}}/mo cap", { + cap: `${currencySymbol(wallet.currency)}${wallet.capUsd}`, + }) + : t("payg.usage.noCap", "No monthly cap")} + + • + + {daysLeft === 1 + ? t("payg.usage.resetsTomorrow", "Resets tomorrow") + : t("payg.usage.resetsIn", "Resets in {{days}} days", { + days: daysLeft, + })} + +
+ +
+ + {t( + "payg.usage.breakdown", + "AI {{ai}} • Automation {{automation}} • API {{api}}", + { + ai: breakdown.ai.toLocaleString(), + automation: breakdown.automation.toLocaleString(), + api: breakdown.api.toLocaleString(), + }, + )} + +
+ + +
+
+ ); +} + +// ─── Cap editor ───────────────────────────────────────────────────────────── + +interface CapEditorProps { + /** Current cap in major currency units; null = no cap set. */ + capUsd: number | null; + /** True when the leader explicitly disabled the cap. */ + noCap: boolean; + /** Per-document rate in minor units; null when unknown — preview hides. */ + pricePerDocMinor: number | null; + /** Currency of the rate; pairs with {@link CapEditorProps#pricePerDocMinor}. */ + currency: string | null; + /** + * Persist the cap change. Receives whole major units (matches the backend's + * {@code PATCH /api/v1/payg/cap} body) or null for no-cap. + */ + onSaveCap?: (capUsd: number | null) => Promise | void; +} + +/** + * Single-row cap editor: the shared {@link SpendCapControl} (preset chips + + * inline custom-entry pill + no-cap + inline Save) over a live "≈ N paid + * PDFs/month" estimate, wrapped in the plan-page card chrome + the cap-reached + * disclosure. Save-only — the working value is local, so abandoning the card + * abandons the edit. The very same control drives the upgrade checkout flow. + */ +function CapEditor({ + capUsd, + noCap, + pricePerDocMinor, + currency, + onSaveCap, +}: CapEditorProps) { + const { t } = useTranslation(); + const savedCap = noCap || capUsd == null ? null : capUsd; + const [working, setWorking] = useState(savedCap); + + return ( +
+ +
+
+ {t("payg.cap.title", "Monthly spending cap")} +
+
+ {t( + "payg.cap.subtitle", + "The most your team can spend per month. Billable processing pauses at the cap and resumes next period.", + )} +
+
+ + + + +
+
+ ); +} + +function CapReadOnly({ + capUsd, + noCap, +}: { + capUsd: number | null; + noCap: boolean; +}) { + const { t } = useTranslation(); + const hasCap = !noCap && capUsd != null; + return ( +
+ +
+ {t("payg.cap.title", "Monthly spending cap")} +
+ + + {hasCap ? `$${capUsd}` : t("payg.cap.noneShort", "No cap")} + + {hasCap && ( + {t("payg.cap.perMonth", "/ month")} + )} + +
+
+
+ {t( + "payg.member.askLeader", + "Only your team owner can change the cap.", + )} +
+
+
+ + +
+
+ ); +} + +// ─── Gates ──────────────────────────────────────────────────────────────── + +// What still works once the cap is hit. Everyday tools keep running; only AI +// and automation/pipelines pause until the cap resets or is raised. +const GATE_CAP_BEHAVIOR: { gate: Gate; staysAtCap: boolean }[] = [ + { gate: "CLIENT_SIDE", staysAtCap: true }, + { gate: "OFFSITE_PROCESSING", staysAtCap: true }, + { gate: "AUTOMATION", staysAtCap: false }, + { gate: "AI_SUPPORT", staysAtCap: false }, +]; + +/** + * Collapsed "what happens at the cap" disclosure rendered inside the cap + * card(s) to save vertical space. Mirrors the {@link DocHelp} toggle pattern; + * default-collapsed so it costs no height until the user opens it. + */ +function CapReachedHelp() { + const { t } = useTranslation(); + const [open, setOpen] = useState(false); + return ( +
+ + {open && ( +
+
+ {GATE_CAP_BEHAVIOR.map(({ gate, staysAtCap }) => ( +
+ + {staysAtCap ? ( + + ) : ( + + )} + + {gateLabel(gate, t)} + {!staysAtCap && ( + + {t("payg.gates.pauses", "pauses at cap")} + + )} +
+ ))} +
+
+ )} +
+ ); +} + +// ─── Per-member usage ──────────────────────────────────────────────────────── + +/** + * Leader-only roster of each teammate's billable usage this period. Display-only — per-member + * sub-cap enforcement isn't shipped, so there's no cap control here. + */ +function MemberUsage({ members }: { members: Wallet["members"] }) { + const { t } = useTranslation(); + return ( +
+ +
+
+ {t("payg.members.title", "Team member usage")} +
+
+ {t( + "payg.members.subtitle", + "Billable PDFs each teammate has processed this period.", + )} +
+
+
+ {members.map((m) => ( +
+ + {m.name.charAt(0).toUpperCase()} + +
+
{m.name}
+
{m.email}
+
+
+
+ {m.spendUnits.toLocaleString()}{" "} + + {t("payg.members.docs", "PDFs")} + +
+
+
+ ))} +
+
+
+ ); +} + +// ─── Activity feed ────────────────────────────────────────────────────────── + +/** + * Feature flag — the activity feed is hidden until the meter-event surface is + * built and polished (Wave 2). The backend returns {@code []} today, so an + * unpolished "No billable activity yet" card adds nothing. Flip to {@code true} + * once {@code wallet.recent} carries real rows. Kept as a flag (not deleted) so + * the renderer below stays wired and ready. + */ +const SHOW_ACTIVITY_FEED = false; + +/** + * Renders {@code wallet.recent} — the backend returns {@code []} in V1 (the + * meter-event surface ships in Wave 2), so today this shows a real empty + * state. The row renderer is ready for when the rows arrive; fields are read + * defensively because the activity-row shape isn't finalised yet (the Wallet + * type carries {@code Record} for the same reason). + */ +function ActivityFeed({ recent }: { recent: Wallet["recent"] }) { + const { t } = useTranslation(); + return ( +
+ +
+
+ {t("payg.activity.title", "Recent billable activity")} +
+
+ {t( + "payg.activity.subtitle", + "Only AI and automation draw from your budget. Everyday tools are free and aren't listed here.", + )} +
+
+ {recent.length === 0 ? ( + + {t("payg.activity.empty", "No billable activity yet this period.")} + + ) : ( +
+ {recent.map((r, i) => ( +
+ +
+
+ {String(r.label ?? "")} +
+
{String(r.ts ?? "")}
+
+ + {String(r.kind ?? "")} + + + {String(r.docUnits ?? 0)} {t("payg.activity.docs", "docs")} + +
+ ))} +
+ )} +
+
+ ); +} + +// ─── Stripe CTA ────────────────────────────────────────────────────────────── + +function StripePortalLink({ + onOpenPortal, +}: { + onOpenPortal: () => Promise; +}) { + const { t } = useTranslation(); + const [loading, setLoading] = useState(false); + + const handleClick = async () => { + setLoading(true); + try { + await onOpenPortal(); + } catch (e: unknown) { + // 503 = Supabase edge fn isn't configured (local dev without + // PORTAL_NOT_CONFIGURED env). 404 = no Stripe customer yet (e.g. the + // team was force-subscribed via dev hooks). Both are user-actionable + // in roughly the same way ("try again later or contact support") so + // we don't bother branching the copy. + console.warn("[Payg] portal session failed", e); + showToast({ + alertType: "warning", + title: t( + "payg.stripe.toast.unavailable.title", + "Billing portal unavailable", + ), + body: t( + "payg.stripe.toast.unavailable.body", + "Billing portal isn't available right now. Try again in a moment.", + ), + location: "bottom-right", + }); + } finally { + setLoading(false); + } + }; + + return ( +
+
+
+ {t("payg.stripe.title", "Manage billing in Stripe")} +
+
+ {t( + "payg.stripe.subtitle", + "Receipts, invoices, payment method, billing currency.", + )} +
+
+ +
+ ); +} + +// ─── Main component ─────────────────────────────────────────────────────── + +const Payg: React.FC = ({ + role, + wallet, + onSaveCap, + onOpenPortal, +}) => { + useRenderCount(role === "LEADER" ? "PaygLeader" : "PaygMember"); + const { t } = useTranslation(); + const isLeader = role === "LEADER"; + + const fmt = (iso: string) => + new Date(iso).toLocaleDateString(undefined, { + day: "numeric", + month: "short", + }); + + return ( +
+ + {/* The modal chrome already renders the section title ("Billing & + usage"), so we lead with the descriptive subtitle + role pill. */} +
+
+ + {t( + "payg.header.eyebrow", + "Processor plan · {{start}} – {{end}}", + { + start: fmt(wallet.billingPeriodStart), + end: fmt(wallet.billingPeriodEnd), + }, + )} + + + {isLeader + ? t("payg.role.leader", "Team owner") + : t("payg.role.member", "Member")} + +
+ +
+
+
+ + {t("payg.header.freeLabel", "Always free")} +
+

+ {t("payg.header.freeTitle", "Unlimited PDF editing")} +

+

+ {t( + "payg.header.freeBody", + "View, edit, merge, split, sign, watermark, compress, convert and manual OCR, as much as you want, no matter where you trigger it.", + )} +

+
+ +
+
+ + {t("payg.header.meterLabel", "Metered")} +
+

+ {t("payg.header.meterTitle", "Automation · AI · API")} +

+

+ {t( + "payg.header.meterBody", + "{{limit}} free PDFs to start, then billed per PDF up to your cap.", + { limit: wallet.freeAllowance.toLocaleString() }, + )} +

+
+
+
+ + + + {isLeader ? ( + + ) : ( + + )} + + {isLeader && wallet.members.length > 0 && ( + + )} + + {SHOW_ACTIVITY_FEED && } + + {isLeader && onOpenPortal && ( + + )} +
+
+ ); +}; + +export default Payg; + +// Convenience exports for the config nav to render either variant directly. +export interface PaygLeaderProps { + /** See {@link PaygProps#wallet}. */ + wallet: Wallet; + /** See {@link PaygProps#onSaveCap}. */ + onSaveCap?: (capUsd: number | null) => Promise | void; + /** See {@link PaygProps#onOpenPortal}. */ + onOpenPortal?: () => Promise; +} +export const PaygLeader: React.FC = ({ + wallet, + onSaveCap, + onOpenPortal, +}) => ( + +); +export const PaygMember: React.FC<{ wallet: Wallet }> = ({ wallet }) => ( + +); diff --git a/frontend/editor/src/cloud/components/shared/config/configSections/PaygFree.css b/frontend/editor/src/cloud/components/shared/config/configSections/PaygFree.css new file mode 100644 index 0000000000..17ef96e8b3 --- /dev/null +++ b/frontend/editor/src/cloud/components/shared/config/configSections/PaygFree.css @@ -0,0 +1,324 @@ +/* Styles for the free-tier Plan views (PaygFreeLeader + PaygFreeMember). + Extends Payg.css — the hero + bar + status pill reuse those classes so the + visual reads as a continuation of the Processor dashboard. New classes here + are prefixed `paygf-`. */ + +.payg-hero__head-row { + display: flex; + align-items: flex-start; + justify-content: space-between; + gap: 16px; + flex-wrap: nowrap; +} + +/* ── Editor plan card (always-free tools only, no billing window) ─────── */ +/* Reuses .payg-planhead chrome; the eyebrow sits in the flex top row so its + default 8px bottom margin would misalign it against the role pill. */ +.paygf-editorcard__eyebrow { + margin-bottom: 0; +} + +/* ── Processor plan card: two-column (pitch + benefits | meter + CTA) ──── */ +.paygf-proc { + gap: 14px; +} +.paygf-proc__eyebrow { + display: inline-flex; + align-items: center; + gap: 6px; + font-size: 0.72rem; + font-weight: 700; + letter-spacing: 0.06em; + text-transform: uppercase; + color: var(--payg-accent); +} +.paygf-proc__split { + display: grid; + grid-template-columns: minmax(0, 1.15fr) minmax(0, 1fr); + gap: 20px; +} +.paygf-proc__pitch { + min-width: 0; + display: flex; + flex-direction: column; + gap: 12px; +} +.paygf-proc__pitch .paygf-cta__subtitle { + margin-top: 0; +} +/* Force the benefits into a single column inside the narrower left column. */ +.paygf-proc__benefits { + grid-template-columns: 1fr; +} +.paygf-proc__aside { + display: flex; + flex-direction: column; + gap: 12px; + padding-left: 20px; + border-left: 1px solid var(--payg-divider); +} +.paygf-proc__cta { + width: 100%; + text-align: center; +} +.paygf-proc__reassure { + text-align: center; +} +.paygf-proc__membernote { + display: flex; + align-items: flex-start; + gap: 8px; + font-size: 0.8rem; + line-height: 1.45; + color: var(--text-muted); +} +.paygf-proc__membernote-icon { + flex-shrink: 0; + color: var(--text-muted); + font-size: 1.1rem !important; + margin-top: 1px; +} +@media (max-width: 640px) { + .paygf-proc__split { + grid-template-columns: 1fr; + } + .paygf-proc__aside { + padding-left: 0; + padding-top: 16px; + border-left: none; + border-top: 1px solid var(--payg-divider); + } +} + +/* ── Compact one-time free meter (inside the Processor aside) ─────────── */ +.paygf-meter { + padding: 14px 16px; + border-radius: 11px; + background: var(--payg-inset-bg); + border: 1px solid var(--payg-card-border); +} +.paygf-meter__top { + display: flex; + align-items: center; + justify-content: space-between; + gap: 10px; +} +.paygf-meter__figure { + display: flex; + align-items: baseline; + gap: 7px; +} +.paygf-meter__num { + font-size: 1.7rem; + font-weight: 750; + line-height: 1; + letter-spacing: -0.02em; + color: var(--text-primary); + font-variant-numeric: tabular-nums; +} +.paygf-meter__cap { + font-size: 0.85rem; + color: var(--text-muted); + font-variant-numeric: tabular-nums; +} +.paygf-meter .payg-bar { + margin-top: 11px; +} +.paygf-meter__meta { + margin-top: 10px; + display: flex; + flex-wrap: wrap; + align-items: center; + gap: 6px 10px; + font-size: 0.78rem; + color: var(--text-secondary); +} + +/* ── Leader CTA card ──────────────────────────────────────────────────── */ + +.paygf-cta { + padding: 18px 22px; + border-radius: var(--payg-radius); + background: var(--payg-card-bg); + /* Gradient border using padding-box / border-box trick so the inner bg + stays solid + the outline gets the brand gradient. */ + border: 1.5px solid transparent; + background: + linear-gradient(var(--payg-card-bg), var(--payg-card-bg)) padding-box, + linear-gradient(135deg, var(--payg-accent) 0%, #6c5ce7 100%) border-box; + box-shadow: 0 10px 28px -18px rgba(10, 139, 255, 0.4); + display: flex; + flex-direction: column; + gap: 16px; +} + +.paygf-cta__heading-row { + display: flex; + gap: 14px; + align-items: center; +} +.paygf-cta__icon { + flex-shrink: 0; + padding: 10px; + border-radius: 12px; + background: linear-gradient(135deg, var(--payg-accent) 0%, #6c5ce7 100%); + color: white !important; + font-size: 1.7rem !important; +} +.paygf-cta__heading-text { + flex: 1 1 auto; +} +.paygf-cta__title { + margin: 0; + font-size: 1.2rem; + font-weight: 700; + color: var(--text-primary); + letter-spacing: -0.01em; + line-height: 1.25; +} +.paygf-cta__subtitle { + margin: 4px 0 0; + font-size: 0.9rem; + color: var(--text-muted); + line-height: 1.5; +} + +.paygf-cta__benefits { + list-style: none; + margin: 0; + padding: 0; + display: grid; + grid-template-columns: repeat(auto-fit, minmax(220px, 1fr)); + gap: 10px 16px; +} +.paygf-cta__benefits li { + display: flex; + align-items: flex-start; + gap: 8px; + font-size: 0.85rem; + color: var(--text-primary); + line-height: 1.45; +} +.paygf-cta__check { + flex-shrink: 0; + color: var(--payg-accent); + margin-top: 2px; +} + +/* Footer row: button on the left, reassurance text on the right. + Right-side text wraps to second line below button on narrow viewports. */ +.paygf-cta__footer { + display: flex; + align-items: center; + justify-content: space-between; + gap: 14px; + flex-wrap: wrap; + padding-top: 4px; + border-top: 1px solid var(--payg-divider); + padding-top: 16px; +} +.paygf-cta__button { + padding: 13px 22px; + border: none; + border-radius: 11px; + background: linear-gradient(135deg, var(--payg-accent) 0%, #6c5ce7 100%); + color: white; + font-weight: 600; + font-size: 0.95rem; + font-family: inherit; + cursor: pointer; + transition: + transform 120ms ease, + box-shadow 120ms ease; + white-space: nowrap; +} +.paygf-cta__button:hover { + transform: translateY(-1px); + box-shadow: 0 8px 22px -6px rgba(10, 139, 255, 0.55); +} +.paygf-cta__reassurance { + margin: 0; + font-size: 0.78rem; + color: var(--text-muted); + text-align: right; +} + +/* ── Member ask-the-owner note ───────────────────────────────────────── */ + +.paygf-member-note { + padding: 18px 22px; + border-radius: var(--payg-radius); + background: var(--payg-inset-bg); + border: 1px solid var(--payg-card-border); + display: flex; + gap: 14px; + align-items: flex-start; +} +.paygf-member-note__icon { + flex-shrink: 0; + color: var(--text-muted); + padding: 8px; + background: var(--payg-card-bg); + border-radius: 10px; + font-size: 1.6rem !important; +} +.paygf-member-note__title { + margin: 0; + font-size: 1rem; + font-weight: 700; + color: var(--text-primary); + letter-spacing: -0.005em; + line-height: 1.3; +} +.paygf-member-note__body { + margin: 6px 0 0; + font-size: 0.875rem; + color: var(--text-muted); + line-height: 1.5; +} + +/* ── Explainer (free vs metered) — shared by leader + member ──────────── */ + +.paygf-explainer { + display: grid; + grid-template-columns: 1fr 1fr; + gap: 14px; + margin-top: var(--payg-gap); +} +@media (max-width: 640px) { + .paygf-explainer { + grid-template-columns: 1fr; + } +} +.paygf-explainer__col { + padding: 18px 20px; + border-radius: 12px; + background: var(--payg-card-bg); + border: 1px solid var(--payg-card-border); +} +.paygf-explainer__label { + display: inline-flex; + align-items: center; + gap: 6px; + font-size: 0.78rem; + font-weight: 700; + letter-spacing: 0.06em; + text-transform: uppercase; + color: var(--text-muted); + margin-bottom: 10px; +} +.paygf-explainer__icon { + font-size: 1rem !important; +} +.paygf-explainer__icon--free { + color: #10b981; +} +.paygf-explainer__icon--paid { + color: var(--payg-accent); +} +.paygf-explainer__text { + margin: 0; + font-size: 0.85rem; + color: var(--text-primary); + line-height: 1.55; +} diff --git a/frontend/editor/src/cloud/components/shared/config/configSections/PaygFree.tsx b/frontend/editor/src/cloud/components/shared/config/configSections/PaygFree.tsx new file mode 100644 index 0000000000..330bffa71e --- /dev/null +++ b/frontend/editor/src/cloud/components/shared/config/configSections/PaygFree.tsx @@ -0,0 +1,288 @@ +/** + * Free-tier Plan views (leader + member) shown before the team subscribes to + * Processor. Mirrors the visual language of {@code Payg.tsx} so the upgrade + * feels like enabling a switch, not visiting a new product. + * + *

Free tier model (set 2026-06): users only pay for + * automation, AI, and API operations. Manual tools + * — viewing, editing, signing, merging, splitting, conversion, manual OCR, + * watermarks, compression — are unmetered, no matter where they're triggered + * from. The distinction is the type of work (manual tool vs + * automation / AI / API), not where the click happens, because automation and + * AI also have UI surfaces. The one-time free grant (default 500) applies + * only to the three billable categories — it is a lifetime allowance, + * not a monthly one, and a team keeps any unused portion after subscribing. + * + *

Layout: a slim Editor plan card (always-free tools only — no dates, + * no metered split) on top, then a single Processor plan card that + * two-columns the upgrade pitch + benefits (left) against the one-time free + * meter stacked over the call-to-action (right). + * + *

Two variants: + * - {@link PaygFreeLeader} — the right column's CTA opens the upgrade modal. + * - {@link PaygFreeMember} — read-only; the CTA is replaced with an + * ask-the-owner note. + */ +import React, { useState } from "react"; +import { Stack } from "@mantine/core"; +import BoltIcon from "@mui/icons-material/BoltRounded"; +import AllInclusiveIcon from "@mui/icons-material/AllInclusiveRounded"; +import CheckIcon from "@mui/icons-material/CheckRounded"; +import LockIcon from "@mui/icons-material/LockOutlined"; +import { useTranslation } from "react-i18next"; +import { useRenderCount } from "@app/hooks/useRenderCount"; +import { useWallet } from "@app/hooks/useWallet"; +// eslint-disable-next-line no-restricted-imports +import "./Payg.css"; +// eslint-disable-next-line no-restricted-imports +import "./PaygFree.css"; +// eslint-disable-next-line no-restricted-imports +import UpgradeModal from "./UpgradeModal"; +// eslint-disable-next-line no-restricted-imports +import { DocHelp } from "./Payg"; +import { + FreeMeterPanel, + useFreeSnapshot, + type FreeSnapshot, +} from "@app/components/shared/config/configSections/usageMeters"; + +// ─── Editor plan card (always-free tools only) ──────────────────────────── + +interface EditorPlanCardProps { + /** Role pill text on the right. */ + pill: string; + /** LEADER pill colour treatment. */ + leader?: boolean; +} + +/** + * The top card: the free Editor plan. Manual tools only, no billing window — + * the one-time grant lives in the Processor card below, so there's no period + * to show here. + */ +function EditorPlanCard({ pill, leader }: EditorPlanCardProps) { + const { t } = useTranslation(); + return ( +

+
+ + + {t("payg.free.editor.eyebrow", "Editor plan · Always free")} + + + {pill} + +
+

+ {t("payg.free.header.freeTitle", "Unlimited PDF editing")} +

+

+ {t( + "payg.free.header.freeBody", + "View, edit, merge, split, sign, watermark, compress, convert and manual OCR, as much as you want, no matter where you trigger it.", + )} +

+
+ ); +} + +// ─── Processor plan card (two-column: pitch + benefits | meter + CTA) ────── + +interface ProcessorCardProps { + snap: FreeSnapshot; + /** Leaders get the live CTA; members get the ask-owner note. */ + isLeader: boolean; + /** Opens the upgrade modal — leader only. */ + onTurnOn?: () => void; +} + +function ProcessorCard({ snap, isLeader, onTurnOn }: ProcessorCardProps) { + const { t } = useTranslation(); + return ( +
+ + + {t("payg.free.proc.eyebrow", "Processor plan · metered")} + + +
+
+

+ {t("payg.free.cta.title", "Turn on the Processor plan")} +

+

+ {t( + "payg.free.cta.subtitle", + "Keep going past your {{limit}} free PDFs with automation, AI, and the API. Set a monthly ceiling, so you stay in control.", + { limit: snap.billableLimit.toLocaleString() }, + )} +

+ +
    +
  • + + + + {t("payg.free.cta.benefit1Title", "Automation pipelines")} + + {": "} + {t( + "payg.free.cta.benefit1Body", + "chain tools, schedule runs, batch process", + )} + +
  • +
  • + + + {t("payg.free.cta.benefit2Title", "AI tools")} + {": "} + {t( + "payg.free.cta.benefit2Body", + "summarise, classify, redact, AI-OCR", + )} + +
  • +
  • + + + + {t("payg.free.cta.benefit3Title", "API access")} + + {" — "} + {t( + "payg.free.cta.benefit3Body", + "call any Stirling endpoint programmatically", + )} + +
  • +
+ + +
+ +
+ + {isLeader ? ( + <> + + + {t( + "payg.free.cta.reassurance", + "No minimum · Set a $0 cap to test · Cancel anytime", + )} + + + ) : ( +
+ + + {t( + "payg.free.member.ownerOnly", + "Only your team owner can turn on Processor. Manual tools stay free for you to use as much as you like.", + )} + +
+ )} +
+
+
+ ); +} + +// ─── Free LEADER ────────────────────────────────────────────────────────── + +export interface PaygFreeLeaderProps { + /** + * Called when the user finishes the {@link UpgradeModal} checkout flow. + * Plumbed up to {@code PlanSection} so the page can flip to the subscribed + * view immediately. When undefined we fall back to a demo {@code alert} so + * the dev preview route still works in isolation. + */ + onUpgraded?: (result: { capUsd: number | null }) => void; +} + +function PaygFreeLeaderInner({ onUpgraded }: PaygFreeLeaderProps = {}) { + useRenderCount("PaygFreeLeader"); + const { t } = useTranslation(); + const snap = useFreeSnapshot(); + const { wallet } = useWallet(); + const [upgradeOpen, setUpgradeOpen] = useState(false); + + return ( +
+ + + setUpgradeOpen(true)} + /> + + + {wallet?.teamId != null && ( + setUpgradeOpen(false)} + onComplete={({ capUsd }) => { + setUpgradeOpen(false); + if (onUpgraded) { + onUpgraded({ capUsd }); + } else { + // Standalone fallback (dev preview route renders without a parent + // handler). Real flow always passes onUpgraded via PlanSection. + alert( + `Demo: subscription complete. Cap = ${ + capUsd === null ? "no cap" : `$${capUsd}/mo` + }.`, + ); + } + }} + /> + )} +
+ ); +} + +// ─── Free MEMBER ────────────────────────────────────────────────────────── + +function PaygFreeMemberInner() { + useRenderCount("PaygFreeMember"); + const { t } = useTranslation(); + const snap = useFreeSnapshot(); + + return ( +
+ + + + +
+ ); +} + +// React.memo so Plan re-rendering on loading/error toggles doesn't cascade +// down to these leaves. Plan passes a stable onUpgraded callback (hoisted in +// Plan.tsx) so the prop identity stays stable across wallet refetches. +export const PaygFreeLeader = React.memo(PaygFreeLeaderInner); +export const PaygFreeMember = React.memo(PaygFreeMemberInner); diff --git a/frontend/editor/src/cloud/components/shared/config/configSections/Plan.tsx b/frontend/editor/src/cloud/components/shared/config/configSections/Plan.tsx new file mode 100644 index 0000000000..ff34d9579e --- /dev/null +++ b/frontend/editor/src/cloud/components/shared/config/configSections/Plan.tsx @@ -0,0 +1,94 @@ +/** + * SaaS "Plan" page — the single entry point for billing, plan state, and + * usage. Branches on the team's wallet state and the viewer's role: + * + * - free + leader → {@link PaygFreeLeader} (upgrade CTA + manual-tools + * framing) + * - free + member → {@link PaygFreeMember} (ask-the-owner note) + * - subscribed + leader → {@link PaygLeader} (full dashboard, editable cap) + * - subscribed + member → {@link PaygMember} (member dashboard) + * + *

The hook handles loading + error states locally so the four view + * components stay focused on rendering the data they own. {@code Plan} is + * intentionally tiny (under 60 lines) so future "Plan-level" affordances — + * a top-level error toast, a subscription confirmation card, etc — have + * obvious places to land. + */ +import React, { useCallback } from "react"; +import { Alert, Center, Loader } from "@mantine/core"; +import { useTranslation } from "react-i18next"; +import { useWallet } from "@app/hooks/useWallet"; +import { useRenderCount } from "@app/hooks/useRenderCount"; +import { + PaygLeader, + PaygMember, +} from "@app/components/shared/config/configSections/Payg"; +import { + PaygFreeLeader, + PaygFreeMember, +} from "@app/components/shared/config/configSections/PaygFree"; + +const Plan: React.FC = () => { + useRenderCount("Plan"); + const { t } = useTranslation(); + const { wallet, loading, error, markSubscribed, updateCap, openPortal } = + useWallet(); + + // Stable callback so PaygFreeLeader's React.memo doesn't see a new prop + // identity on every Plan render (e.g. loading flips false→true→false on + // a refetch). Closing over the stable markSubscribed from useWallet + // means we don't need to add wallet state to deps. + const onUpgraded = useCallback( + ({ capUsd }: { capUsd: number | null }) => { + // Bridges the modal's local success → backend mock → refetch loop. + // Real Stripe flow: the customer.subscription.created webhook is + // what flips status; we still call markSubscribed locally so the + // optimistic refetch hits immediately. + void markSubscribed(capUsd); + }, + [markSubscribed], + ); + + if (loading && !wallet) { + return ( +

+ +
+ ); + } + + if (error || !wallet) { + return ( + + {error ?? + t( + "payg.error.body", + "We couldn't reach the billing service. Refresh the page to try again.", + )} + + ); + } + + if (wallet.status === "subscribed") { + return wallet.role === "leader" ? ( + + ) : ( + + ); + } + + // Free tier — only the leader sees the upgrade CTA. + if (wallet.role === "leader") { + return ; + } + return ; +}; + +export default Plan; diff --git a/frontend/editor/src/cloud/components/shared/config/configSections/SpendCapControl.css b/frontend/editor/src/cloud/components/shared/config/configSections/SpendCapControl.css new file mode 100644 index 0000000000..90083fb80e --- /dev/null +++ b/frontend/editor/src/cloud/components/shared/config/configSections/SpendCapControl.css @@ -0,0 +1,164 @@ +/* Reusable monthly spend-cap control — preset chips + an inline custom-entry + pill + a no-cap chip, with a live "≈ N PDFs/month" estimate beneath. Shared + by the subscribed plan-page cap editor (Payg.tsx) and the upgrade checkout + flow (UpgradeModal.tsx). + + Self-contained tokens so the control looks right whether it's dropped into + the Plan card (light/dark modal content) or the dark upgrade modal. We lean + on the app's semantic theme tokens and override per scheme, mirroring the + Payg.css / UpgradeModal.css treatments the control sits beside. */ + +.scc { + --scc-accent: #0a8bff; + --scc-accent-text: #0a8bff; + --scc-accent-soft: rgba(10, 139, 255, 0.12); + --scc-accent-border: rgba(10, 139, 255, 0.25); + --scc-chip-bg: var(--bg-muted); + --scc-chip-border: var(--border-default); + /* Consistent vertical rhythm between the control row, the estimate, and any + note — without it the blocks sit flush. */ + display: flex; + flex-direction: column; + gap: 14px; +} +[data-mantine-color-scheme="dark"] .scc { + --scc-accent-text: #66b8ff; + --scc-accent-soft: rgba(10, 139, 255, 0.16); + --scc-chip-bg: #272d35; + --scc-chip-border: #3d444e; +} + +/* ── Inline row: presets · custom · no-cap · (save) ──────────────────── */ +.scc-row { + display: flex; + flex-wrap: wrap; + align-items: center; + gap: 8px; +} + +/* Shared pill shape for presets, the custom-entry pill, and no-cap. Keeping a + single height/radius/border is what makes the custom field read as "one more + button" rather than a form input bolted on. */ +.scc-chip { + display: inline-flex; + align-items: center; + height: 34px; + padding: 0 16px; + border-radius: 999px; + border: 1px solid var(--scc-chip-border); + background: var(--scc-chip-bg); + color: var(--text-secondary); + font: inherit; + font-size: 0.85rem; + font-weight: 600; + cursor: pointer; + font-variant-numeric: tabular-nums; + transition: + color 0.12s ease, + border-color 0.12s ease, + background 0.12s ease; +} +.scc-chip:hover:not(:disabled) { + color: var(--text-primary); + border-color: var(--border-strong); +} +.scc-chip[data-selected="true"] { + background: var(--scc-accent-soft); + border-color: var(--scc-accent); + color: var(--scc-accent-text); +} +.scc-chip:disabled { + opacity: 0.45; + cursor: default; +} + +/* Custom-entry pill — Option A: a dashed pill that matches the presets and + reads as "or type your own", filling solid like a selected chip once it + carries a value. */ +.scc-custom { + display: inline-flex; + align-items: center; + gap: 1px; + height: 34px; + padding: 0 14px; + border-radius: 999px; + border: 1px dashed var(--scc-chip-border); + background: transparent; + cursor: text; + transition: + border-color 0.12s ease, + background 0.12s ease; +} +.scc-custom:hover { + border-color: var(--border-strong); +} +.scc-custom[data-active="true"] { + border-style: solid; + border-color: var(--scc-accent); + background: var(--scc-accent-soft); +} +.scc-custom__symbol { + color: var(--text-muted); + font-size: 0.85rem; + font-weight: 600; +} +.scc-custom[data-active="true"] .scc-custom__symbol, +.scc-custom[data-active="true"] .scc-custom__input { + color: var(--scc-accent-text); +} +.scc-custom__input { + width: 70px; + border: none; + outline: none; + background: transparent; + color: var(--text-primary); + font: inherit; + font-size: 0.85rem; + font-weight: 600; + font-variant-numeric: tabular-nums; +} +.scc-custom__input::placeholder { + color: var(--text-muted); + font-weight: 600; +} +.scc-custom:disabled, +.scc-custom[data-disabled="true"] { + opacity: 0.45; + cursor: default; +} + +.scc-row__spacer { + margin-left: auto; +} + +/* ── Live PDF estimate ───────────────────────────────────────────────── */ +.scc-estimate { + display: flex; + align-items: center; + gap: 11px; + padding: 12px 14px; + border-radius: 10px; + background: var(--scc-accent-soft); + border: 1px solid var(--scc-accent-border); +} +.scc-estimate__icon { + color: var(--scc-accent-text); + display: flex; +} +.scc-estimate__main { + font-size: 0.875rem; + font-weight: 550; + color: var(--text-primary); +} +.scc-estimate__sub { + font-size: 0.75rem; + color: var(--text-muted); + margin-top: 1px; +} + +/* Quiet helper line (e.g. "Shown in USD — change later in your currency"). */ +.scc-note { + font-size: 0.8125rem; + color: var(--text-muted); + line-height: 1.45; +} diff --git a/frontend/editor/src/cloud/components/shared/config/configSections/SpendCapControl.tsx b/frontend/editor/src/cloud/components/shared/config/configSections/SpendCapControl.tsx new file mode 100644 index 0000000000..bee9d0474d --- /dev/null +++ b/frontend/editor/src/cloud/components/shared/config/configSections/SpendCapControl.tsx @@ -0,0 +1,251 @@ +/** + * Reusable monthly spend-cap control. + * + * One inline row — preset chips, a custom-entry pill that matches the presets, + * a "No cap" chip, and (optionally) a Save button — over a live "≈ N PDFs / + * month" estimate. Extracted from the subscribed plan-page cap editor so the + * exact same control drives the upgrade checkout flow. + * + *

Currency-agnostic by design

+ * + * The control never decides a currency. It takes {@code pricePerDocMinor} + + * {@code currency} and renders whatever it's handed: the subscribed plan page + * passes the team's real Stripe-subscription rate/currency; the unsubscribed + * checkout flow passes a USD rate (Stripe hasn't assigned the team a currency + * yet) plus a {@code note} explaining the cap is editable later. When no rate + * is supplied the estimate simply hides. + * + *

Controlled

+ * + * Fully controlled via {@code capUsd} ({@code null} = no cap, {@code 0} = a + * real $0 cap that keeps everything free) + {@code onChange}. The parent owns + * the working value. When {@code onSave} is provided the control renders the + * inline Save button and computes "dirty" against {@code savedCapUsd}. + */ +import React, { useState } from "react"; +import { Button } from "@mantine/core"; +import DescriptionIcon from "@mui/icons-material/DescriptionOutlined"; +import LocalIcon from "@app/components/shared/LocalIcon"; +import { useTranslation } from "react-i18next"; +// eslint-disable-next-line no-restricted-imports +import "./SpendCapControl.css"; + +// Quick amounts offered everywhere — recognition over recall. +export const DEFAULT_CAP_PRESETS = [500, 1000, 2500, 5000] as const; + +export interface SpendCapControlProps { + /** Current cap in major currency units; {@code null} = no cap. Controlled. */ + capUsd: number | null; + /** Working-value setter. {@code null} signals no-cap. */ + onChange: (capUsd: number | null) => void; + /** Per-document rate in minor units; null/0 hides the estimate. May be fractional. */ + pricePerDocMinor?: number | null; + /** Lower-case ISO currency of the rate; pairs with {@link #pricePerDocMinor}. */ + currency?: string | null; + /** Quick-amount presets (major units). Defaults to {@link DEFAULT_CAP_PRESETS}. */ + presets?: readonly number[]; + /** + * When provided, the control renders an inline Save button. Receives whole + * major units, or {@code null} for no-cap. + */ + onSave?: (capUsd: number | null) => Promise | void; + /** Label for the Save button. */ + saveLabel?: string; + /** + * The persisted value to diff against for the dirty check. Same encoding as + * {@link #capUsd} ({@code null} = persisted no-cap). Only used with + * {@link #onSave}. + */ + savedCapUsd?: number | null; + /** Quiet helper line under the estimate (e.g. the USD / editable-later note). */ + note?: React.ReactNode; +} + +/** Format minor units of an ISO currency ("$2.24", "£0.40"). */ +function formatMinor( + minor: number, + currency: string | null | undefined, +): string { + const code = (currency ?? "usd").toUpperCase(); + try { + return new Intl.NumberFormat(undefined, { + style: "currency", + currency: code, + // Per-doc rates are often sub-cent (e.g. $0.02 → 2 minor, but a half-cent + // rate is 0.5). Allow up to 3 fraction digits so they don't round to $0. + maximumFractionDigits: 3, + }).format(minor / 100); + } catch { + return `${(minor / 100).toFixed(2)} ${code}`; + } +} + +/** Currency symbol for compact inline use; falls back to the ISO code. */ +function currencySymbol(currency: string | null | undefined): string { + switch ((currency ?? "").toLowerCase()) { + case "usd": + case "": + return "$"; + case "eur": + return "€"; + case "gbp": + return "£"; + default: + return currency!.toUpperCase() + " "; + } +} + +const SpendCapControl: React.FC = ({ + capUsd, + onChange, + pricePerDocMinor, + currency, + presets = DEFAULT_CAP_PRESETS, + onSave, + saveLabel, + savedCapUsd, + note, +}) => { + const { t } = useTranslation(); + const [saving, setSaving] = useState(false); + + const sym = currencySymbol(currency); + const isNoCap = capUsd === null; + const presetSelected = capUsd != null && presets.includes(capUsd); + // Custom is "active" when a cap is set that isn't one of the presets — i.e. + // the value came from the custom pill. + const customActive = capUsd != null && !presets.includes(capUsd); + + // Local mirror of the custom field's text so partial/empty entry doesn't get + // clobbered by the controlled value. Seeded from a non-preset incoming cap. + const [customText, setCustomText] = useState( + customActive ? String(capUsd) : "", + ); + + // Mirror of the backend's docCapForMoney: floor(capMinor / rate). The + // one-time free grant is a separate lifetime pool and is NOT added here — + // this is the paid PDFs the monthly cap buys. + const rate = + pricePerDocMinor != null && pricePerDocMinor > 0 ? pricePerDocMinor : null; + const previewDocs = + capUsd != null && rate != null ? Math.floor((capUsd * 100) / rate) : null; + + const dirty = onSave != null && capUsd !== (savedCapUsd ?? null); + + const selectPreset = (preset: number) => { + setCustomText(""); + onChange(preset); + }; + const selectNoCap = () => { + setCustomText(""); + onChange(null); + }; + const onCustomInput = (raw: string) => { + // Digits only; an empty field reads as "no custom value yet" → 0 so the + // estimate still renders sensibly without flipping to no-cap. + const cleaned = raw.replace(/[^0-9]/g, ""); + setCustomText(cleaned); + const v = cleaned === "" ? 0 : parseInt(cleaned, 10); + onChange(Number.isNaN(v) ? 0 : v); + }; + + const handleSave = async () => { + if (!onSave) return; + setSaving(true); + try { + await onSave(isNoCap ? null : Math.round(capUsd ?? 0)); + } finally { + setSaving(false); + } + }; + + return ( +
+
+ {presets.map((preset) => ( + + ))} + + {/* Custom-entry pill — dashed until it carries a value, then it fills + like a selected chip. */} + + + + + {onSave && ( + + )} +
+ + {previewDocs != null && ( +
+ +
+
+ {t("payg.cap.docsEstimate", "≈ {{docs}} processed PDFs / month", { + docs: previewDocs.toLocaleString(), + })} +
+
+ {t("payg.cap.docsRate", "at {{rate}} / PDF", { + rate: formatMinor(pricePerDocMinor ?? 0, currency), + })} +
+
+
+ )} + + {isNoCap && ( +
+ {t( + "payg.cap.noCapDesc", + "Usage is billed without an upper limit. You can re-enable a cap at any time.", + )} +
+ )} + + {note &&
{note}
} +
+ ); +}; + +export default SpendCapControl; diff --git a/frontend/editor/src/cloud/components/shared/config/configSections/StripeCheckoutPanel.tsx b/frontend/editor/src/cloud/components/shared/config/configSections/StripeCheckoutPanel.tsx new file mode 100644 index 0000000000..1ffc9f4ee2 --- /dev/null +++ b/frontend/editor/src/cloud/components/shared/config/configSections/StripeCheckoutPanel.tsx @@ -0,0 +1,300 @@ +/** + * Stripe Embedded Checkout panel — lives in its own module so it can be + * imported lazily. {@code @stripe/stripe-js} pulls a fairly chunky third-party + * SDK; we don't want it in the main bundle for users who never open the + * UpgradeModal, let alone reach step 2. + * + *

Lazy-load pattern

+ * + *
+ * // In UpgradeModal.tsx — only when the user advances to step 2:
+ * const StripeCheckoutPanel = React.lazy(
+ *   () => import("@app/components/shared/config/configSections/StripeCheckoutPanel"),
+ * );
+ * 
+ * + * The {@code loadStripe()} call inside this module deferred-imports the SDK + * itself, so the chunk graph is: + * + *
+ *   StripeCheckoutPanel.chunk.js
+ *     └─ @stripe/react-stripe-js  (pulled in by ESM static import here)
+ *     └─ @stripe/stripe-js        (pulled in by loadStripe inside fetchStripe)
+ * 
+ * + * Both chunks are eligible for Vite tree-shaking + lazy load; nothing in the + * main bundle references either package. + * + *

Architecture

+ * + * Stripe-touching code lives in Supabase edge functions, not the Java backend. + * This panel no longer talks to Supabase directly — it mints the PAYG checkout + * session through the {@code @app/services/billing} seam + * ({@link createCheckoutSession} with a {@code teamId}), which each platform + * implements (web supabase client vs Tauri fetch). The seam drives the + * {@code create-checkout-session} edge function. + * + *

The edge function is the canonical place Stripe Checkout + * Sessions get created - it uses the Stripe Sync Engine tables, has dedicated + * unit tests, and shares Stripe SDK / secret-key plumbing with the metering + + * webhook edge functions. Routing through Java would have meant a useless + * proxy hop + a second Stripe SDK to maintain. + * + *

Behaviour

+ * + *
    + *
  1. On mount: calls {@link createCheckoutSession} with the {@code teamId}, + * {@code currency} and billing email to obtain a {@code clientSecret}. The + * spending cap is NOT set here — it's an application-layer setting applied + * via {@code PATCH /payg/cap} after the subscription lands. + *
  2. If no Stripe publishable key is configured OR the edge function isn't + * deployed yet (errors out / returns a {@code cs_mock_} sentinel), render + * a clearly-labelled placeholder + "Continue with mock" button so the + * post-completion path stays testable. + *
  3. If only a hosted {@code url} comes back (the redirect fallback), hand it + * to the system browser via {@link openExternal}. + *
  4. Otherwise render the real {@code } + + * {@code } iframe. The Tauri webview has no CSP, so + * embedded checkout works on desktop too. + *
+ * + * The parent {@link UpgradeModal} passes {@code onComplete} which fires when + * either the real Stripe checkout emits its complete event OR the user + * presses "Continue with mock" in unconfigured environments. + */ +import React, { useEffect, useRef, useState } from "react"; +import { useTranslation } from "react-i18next"; +import { useAuth } from "@app/auth/UseSession"; +import { + createCheckoutSession, + getStripePublishableKey, +} from "@app/services/billing"; +import { openExternal } from "@app/platform/openExternal"; +import { getWalletDevPreview } from "@app/hooks/walletDevPreview"; + +// Eager static imports here are OK because this whole module is itself lazy- +// imported by the modal. They land in the same lazy chunk. +import { + EmbeddedCheckout, + EmbeddedCheckoutProvider, +} from "@stripe/react-stripe-js"; +import type { Stripe } from "@stripe/stripe-js"; + +export interface StripeCheckoutPanelProps { + /** + * The caller's team_id — required by the Supabase edge function, which can't + * derive it from the JWT alone (the function runs outside our app's Spring + * Security context and has no access to the {@code team_memberships} table + * other than via this hint). + */ + teamId: number; + /** Currency lower-case 3-letter ISO (e.g. {@code "gbp"}). Selects the Stripe Price. */ + currency?: string; + /** Cap in USD; null means no cap. Tracked locally; set on the wallet via PATCH after subscription. */ + capUsd: number | null; + /** Called when Stripe (or the mock continue button) signals completion. */ + onComplete: () => void; + /** Called when the call to /api/v1/payg/checkout fails. */ + onError?: (message: string) => void; +} + +// Singleton Stripe promise — created on first use and reused for the lifetime +// of the tab. {@code loadStripe} is dynamically imported so the actual SDK +// chunk is only pulled when this code path runs. +let stripePromise: Promise | null = null; +function getStripe(publishableKey: string): Promise { + if (stripePromise === null) { + stripePromise = import("@stripe/stripe-js").then((mod) => + mod.loadStripe(publishableKey), + ); + } + return stripePromise; +} + +const StripeCheckoutPanel: React.FC = ({ + teamId, + currency = "gbp", + // capUsd is part of the props contract but intentionally unused here — the cap is set + // application-side via PATCH /payg/cap after the subscription lands, not during checkout. + onComplete, + onError, +}) => { + const { t } = useTranslation(); + const { user } = useAuth(); + const [clientSecret, setClientSecret] = useState(null); + const [isMock, setIsMock] = useState(false); + const [loading, setLoading] = useState(true); + const [error, setError] = useState(null); + + // Billing email for the Stripe Checkout Session. Passing customer_email to + // Stripe prefills the email field AND locks it (Stripe's own behaviour for + // sessions created with customer_email or an attached customer) — the user + // can't bill a different address than the account they're signed in with. + const billingEmail = user?.email ?? null; + + // Stable ref for the error callback so we don't have to include it in the + // effect deps. + const onErrorRef = useRef(onError); + onErrorRef.current = onError; + + // Stash t() in a ref so the effect can read the current translator without + // forcing t into its deps. + const tRef = useRef(t); + tRef.current = t; + + const publishableKey = getStripePublishableKey(); + + // Dev preview route has no backend — skip the API call and go straight + // to the mock placeholder so the design + completion path stay testable. + // The dev-preview detection (import.meta.env.DEV + a /dev/ path) lives behind + // the walletDevPreview seam since cloud may not read either directly; it is + // a saas-only affordance and resolves to null on desktop / prod. + const devPreview = getWalletDevPreview() !== null; + + useEffect(() => { + if (devPreview) { + setClientSecret("cs_mock_devpreview"); + setIsMock(true); + setLoading(false); + return; + } + + // React 18 strict-mode dev mounts effects twice. We use `cancelled` to + // discard the first mount's response (its setState calls become no-ops), + // and the second mount's response wins. We deliberately do NOT short- + // circuit the second mount with a ref: that traps mount 1 as cancelled + // while the live mount never re-fetches, leaving `loading` stuck at true. + // Two network calls in dev is an acceptable cost; prod has no strict mode + // -> single fetch. + let cancelled = false; + async function createSession() { + try { + // Mint the PAYG checkout session through the billing seam. Passing a + // teamId routes it to the {@code create-checkout-session} edge function + // (not {@code create-payg-team-subscription} — that one subscribes + // directly without the embedded iframe and returns no client_secret). + // team_id is required because the edge fn runs outside our Spring + // Security context and can't resolve it from the JWT alone. The cap is + // *not* set during checkout — it's applied via PATCH /payg/cap after + // the subscription lands. The platform impl supplies the success/cancel + // return URL (browser origin on web, deep link on desktop). + const session = await createCheckoutSession({ + teamId, + currency, + // Maps to Stripe's customer_email when the team has no Stripe + // customer yet — prefills + locks the email field in Checkout. Teams + // with an existing customer get the email locked from the customer + // record instead; this field is ignored for them. + billingOwnerEmail: billingEmail, + }); + if (cancelled) return; + // Hosted-url fallback: no embedded iframe, hand the URL to the system + // browser. The deep-link / origin return URL brings the user back. + if (session.url && !session.clientSecret) { + await openExternal(session.url); + return; + } + if (!session.clientSecret) { + throw new Error("Edge function returned no client_secret"); + } + setClientSecret(session.clientSecret); + setIsMock( + Boolean(session.mock) || session.clientSecret.startsWith("cs_mock_"), + ); + } catch (e: unknown) { + if (cancelled) return; + const msg = + e instanceof Error + ? e.message + : tRef.current( + "payg.checkout.error.startFailed", + "Couldn't start checkout session", + ); + setError(msg); + onErrorRef.current?.(msg); + } finally { + if (!cancelled) setLoading(false); + } + } + void createSession(); + return () => { + cancelled = true; + }; + }, [teamId, currency, devPreview, billingEmail]); + + if (loading) { + return ( +
+
+ {t("payg.checkout.connecting", "Connecting to Stripe…")} +
+
+ ); + } + + if (error) { + return ( +
+
+ {t("payg.checkout.errorTitle", "Stripe error")} +
+
{error}
+
+ ); + } + + // Mock mode OR no publishable key configured → friendly placeholder. + const showMockPlaceholder = isMock || publishableKey.length === 0; + + if (showMockPlaceholder) { + return ( +
+
+ {t( + "payg.checkout.mock.title", + "Stripe Embedded Checkout (mock mode)", + )} +
+
+ {publishableKey.length === 0 + ? t( + "payg.checkout.mock.noKey", + "VITE_STRIPE_PUBLISHABLE_KEY is unset. Real iframe mounts here once configured.", + ) + : t( + "payg.checkout.mock.backend", + "Backend is in mock mode — no real Stripe session was created.", + )} +
+
+ +
+
+ ); + } + + if (!clientSecret) return null; + + return ( +
+ + + +
+ ); +}; + +export default StripeCheckoutPanel; diff --git a/frontend/editor/src/desktop/components/shared/config/configSections/SaaSTeamsSection.tsx b/frontend/editor/src/cloud/components/shared/config/configSections/TeamSection.tsx similarity index 75% rename from frontend/editor/src/desktop/components/shared/config/configSections/SaaSTeamsSection.tsx rename to frontend/editor/src/cloud/components/shared/config/configSections/TeamSection.tsx index b172f71284..6755a4d746 100644 --- a/frontend/editor/src/desktop/components/shared/config/configSections/SaaSTeamsSection.tsx +++ b/frontend/editor/src/cloud/components/shared/config/configSections/TeamSection.tsx @@ -10,25 +10,14 @@ import { Badge, ActionIcon, Menu, - List, - ThemeIcon, - Modal, - CloseButton, - Anchor, } from "@mantine/core"; import { useTranslation } from "react-i18next"; import { useSaaSTeam } from "@app/contexts/SaaSTeamContext"; -import { useSaaSBilling } from "@app/contexts/SaasBillingContext"; import LocalIcon from "@app/components/shared/LocalIcon"; import { Z_INDEX_OVER_CONFIG_MODAL } from "@app/styles/zIndex"; import apiClient from "@app/services/apiClient"; -/** - * Desktop SaaS Teams Section - * Allows team management for users connected to SaaS backend - * CRITICAL: Only shown when in SaaS mode (enforced by navigation) - */ -export function SaaSTeamsSection() { +const TeamSection: React.FC = () => { const { t } = useTranslation(); const { currentTeam, @@ -43,15 +32,10 @@ export function SaaSTeamsSection() { refreshTeams, } = useSaaSTeam(); - // Check Pro status via billing context - const { tier } = useSaaSBilling(); - const isPro = tier !== "free"; - const [inviteEmail, setInviteEmail] = useState(""); const [inviting, setInviting] = useState(false); const [error, setError] = useState(null); const [success, setSuccess] = useState(null); - const [featuresModalOpened, setFeaturesModalOpened] = useState(false); // Team rename state const [isEditingName, setIsEditingName] = useState(false); @@ -71,12 +55,6 @@ export function SaaSTeamsSection() { return () => clearInterval(interval); }, []); // Only run on mount/unmount - const navigateToPlan = () => { - window.dispatchEvent( - new CustomEvent("appConfig:navigate", { detail: { key: "planBilling" } }), - ); - }; - const handleInvite = async (e: React.FormEvent) => { e.preventDefault(); if (!inviteEmail.trim()) return; @@ -321,156 +299,6 @@ export function SaaSTeamsSection() {
- {/* Upgrade Banner for Free Users */} - {isPersonalTeam && !isPro && ( - } - > - -
- - {t( - "team.upgrade.title", - "Upgrade to Pro to unlock team features", - )} - - - {t( - "team.upgrade.description", - "Invite members, share credits, and more.", - )}{" "} - setFeaturesModalOpened(true)} - style={{ cursor: "pointer" }} - > - {t("common.learnMore", "Learn more")} - - -
- -
-
- )} - - {/* Team Features Modal */} - setFeaturesModalOpened(false)} - size="md" - centered - padding="xl" - withCloseButton={false} - zIndex={Z_INDEX_OVER_CONFIG_MODAL} - > -
- setFeaturesModalOpened(false)} - size="lg" - style={{ - position: "absolute", - top: -8, - right: -8, - zIndex: 1, - }} - /> - - {/* Header */} - - - {t("team.features.badge", "PRO FEATURE")} - - - {t("team.features.title", "Team Collaboration")} - - - {t( - "team.features.subtitle", - "Upgrade to Pro and unlock powerful team features", - )} - - - - {/* Features List */} - - - - } - > - - - {t("team.features.invite.title", "Invite team members")} - - - {t( - "team.features.invite.description", - "Add unlimited users with additional seat purchases", - )} - - - - - {t( - "team.features.credits.title", - "Share credits across your team", - )} - - - {t( - "team.features.credits.description", - "Pool resources for collaborative work", - )} - - - - - {t( - "team.features.dashboard.title", - "Team management dashboard", - )} - - - {t( - "team.features.dashboard.description", - "Control permissions, monitor usage, and manage members", - )} - - - - - {t("team.features.billing.title", "Centralized billing")} - - - {t( - "team.features.billing.description", - "One invoice for all team seats and usage", - )} - - - - - {/* CTA Button */} - - -
-
- {/* Error/Success Messages */} {error && ( setError(null)} withCloseButton> @@ -484,8 +312,8 @@ export function SaaSTeamsSection() { )} - {/* Invite Members (Pro Users) */} - {isTeamLeader && isPro && ( + {/* Invite Members */} + {isTeamLeader && (
{t("team.invite.title", "Invite Team Member")} @@ -706,4 +534,6 @@ export function SaaSTeamsSection() {
); -} +}; + +export default TeamSection; diff --git a/frontend/editor/src/cloud/components/shared/config/configSections/UpgradeModal.css b/frontend/editor/src/cloud/components/shared/config/configSections/UpgradeModal.css new file mode 100644 index 0000000000..2e7aafef5c --- /dev/null +++ b/frontend/editor/src/cloud/components/shared/config/configSections/UpgradeModal.css @@ -0,0 +1,447 @@ +/* Upgrade-to-Processor modal: two steps inside one frame. + Step 1: pick a monthly cap. + Step 2: Stripe Embedded Checkout (or placeholder when no test key set). + The frame doesn't change between steps — only the panel slides — so the + user feels like they're filling out one form, not navigating pages. */ + +.upm { + --upm-accent: #0a8bff; + --upm-accent-2: #6c5ce7; + --upm-accent-soft: rgba(10, 139, 255, 0.08); + --upm-success: #10b981; + --upm-radius: 18px; + --upm-card-bg: var(--bg-surface); + --upm-border: var(--border-default); + --upm-divider: var(--border-subtle); + --upm-text: var(--text-primary); + --upm-muted: var(--text-muted); +} +[data-mantine-color-scheme="dark"] .upm { + --upm-card-bg: #313842; + --upm-border: #3d444e; + --upm-divider: #3d444e; + --upm-accent-soft: rgba(10, 139, 255, 0.14); +} + +/* Backdrop locks the page and centres the modal. z-index matches + Z_INDEX_OVER_SETTINGS_MODAL (styles/zIndex.ts) — must sit above the + AppConfigModal's 1300 (Z_INDEX_OVER_FULLSCREEN_SURFACE). The config modal + also hides itself while we're open (appConfig:overlay event), so this is + belt-and-braces for the brief overlap during open/close transitions. */ +.upm-backdrop { + position: fixed; + inset: 0; + background: rgba(0, 0, 0, 0.55); + backdrop-filter: blur(4px); + z-index: 1400; + display: flex; + align-items: center; + justify-content: center; + padding: 24px; + animation: upm-fade-in 180ms ease; +} + +@keyframes upm-fade-in { + from { + opacity: 0; + } + to { + opacity: 1; + } +} + +.upm-frame { + background: var(--upm-card-bg); + border: 1px solid var(--upm-border); + border-radius: var(--upm-radius); + width: 100%; + /* Stripe Embedded Checkout flips from its single-column layout to the two-column + "horizontal" one (order summary beside the payment form) once the iframe is + ~1000px wide — measured empirically against a live test session; Stripe doesn't + document the breakpoint. The live iframe gets frame_width − 44px (the 22px body + padding each side; the live mount's own padding is removed — see + .upm-stripe-mount[data-state="live"]). So 1100px → ~1056px iframe, comfortably + past the ~1000px threshold. Because the frame stays width:100% under this cap, + narrow windows shrink it and Stripe falls back to single column — horizontal + when it fits, vertical when it doesn't. Cap + confirm steps read fine at this + width (their internal text is already capped — see upm-confirm__body, + upm-section-help). */ + max-width: 1100px; + max-height: calc(100vh - 48px); + display: flex; + flex-direction: column; + overflow: hidden; + box-shadow: 0 30px 80px -20px rgba(0, 0, 0, 0.4); + animation: upm-pop-in 220ms cubic-bezier(0.16, 1, 0.3, 1); +} + +@keyframes upm-pop-in { + from { + opacity: 0; + transform: translateY(12px) scale(0.97); + } + to { + opacity: 1; + transform: translateY(0) scale(1); + } +} + +/* ── Header: title + close + step dots ────────────────────────────────── */ +.upm-header { + display: flex; + align-items: center; + justify-content: space-between; + padding: 18px 22px 14px; + border-bottom: 1px solid var(--upm-divider); +} +.upm-header__title { + font-size: 1.05rem; + font-weight: 700; + letter-spacing: -0.005em; + color: var(--upm-text); + margin: 0; +} +.upm-header__close { + background: transparent; + border: none; + color: var(--upm-muted); + cursor: pointer; + width: 32px; + height: 32px; + border-radius: 8px; + display: inline-flex; + align-items: center; + justify-content: center; + transition: + background 120ms ease, + color 120ms ease; +} +.upm-header__close:hover { + background: var(--upm-divider); + color: var(--upm-text); +} + +/* Left cluster: optional back arrow + title. The back arrow only renders on the + checkout step (cap/confirm have no parent step to return to) — it replaces the + old footer "← Back" button. */ +.upm-header__left { + display: flex; + align-items: center; + gap: 6px; + min-width: 0; +} +.upm-header__back { + background: transparent; + border: none; + color: var(--upm-muted); + cursor: pointer; + width: 32px; + height: 32px; + border-radius: 8px; + display: inline-flex; + align-items: center; + justify-content: center; + flex-shrink: 0; + margin-left: -6px; + transition: + background 120ms ease, + color 120ms ease; +} +.upm-header__back:hover { + background: var(--upm-divider); + color: var(--upm-text); +} + +.upm-steps { + display: flex; + align-items: center; + gap: 8px; + padding: 12px 22px 14px; + background: var(--upm-accent-soft); + border-bottom: 1px solid var(--upm-divider); +} +.upm-step { + display: inline-flex; + align-items: center; + gap: 8px; + flex: 0 0 auto; + font-size: 0.78rem; + font-weight: 600; + color: var(--upm-muted); + transition: color 200ms ease; +} +.upm-step__dot { + width: 22px; + height: 22px; + border-radius: 50%; + background: var(--upm-card-bg); + border: 1.5px solid var(--upm-divider); + display: inline-flex; + align-items: center; + justify-content: center; + font-size: 0.7rem; + font-weight: 700; + transition: + background 200ms ease, + border-color 200ms ease, + color 200ms ease; +} +.upm-step[data-state="active"] { + color: var(--upm-text); +} +.upm-step[data-state="active"] .upm-step__dot { + background: linear-gradient(135deg, var(--upm-accent), var(--upm-accent-2)); + border-color: transparent; + color: white; +} +.upm-step[data-state="done"] { + color: var(--upm-success); +} +.upm-step[data-state="done"] .upm-step__dot { + background: var(--upm-success); + border-color: transparent; + color: white; +} +.upm-step__connector { + flex: 1 1 auto; + height: 1.5px; + background: var(--upm-divider); + margin: 0 4px; +} + +/* ── Body: slides between steps without remounting ───────────────────── */ +.upm-body { + padding: 22px; + overflow-y: auto; + flex: 1 1 auto; +} + +/* Hero promise reinforcement — only shown on step 1 */ +.upm-promise { + display: flex; + align-items: flex-start; + gap: 10px; + padding: 12px 14px; + background: var(--upm-accent-soft); + border-radius: 12px; + margin-bottom: 18px; + font-size: 0.85rem; + line-height: 1.45; + color: var(--upm-text); +} +.upm-promise__icon { + flex-shrink: 0; + color: var(--upm-accent); + margin-top: 1px; +} +.upm-promise__highlight { + font-weight: 700; +} + +.upm-section-title { + font-size: 0.95rem; + font-weight: 700; + color: var(--upm-text); + margin: 0 0 4px; + letter-spacing: -0.005em; +} +.upm-section-help { + font-size: 0.825rem; + color: var(--upm-muted); + margin: 0 0 14px; + line-height: 1.45; +} + +/* Cap selection (presets + custom + no-cap + estimate) now lives in the shared + SpendCapControl (SpendCapControl.css). The old upm-cap-* / upm-no-cap-toggle + rules were removed with it. */ + +/* What counts as "processed" — quiet helper text */ +.upm-help-card { + background: var(--upm-divider); + border-radius: 10px; + padding: 12px 14px; + margin-top: 14px; + font-size: 0.78rem; + line-height: 1.5; + color: var(--upm-muted); +} +.upm-help-card__title { + display: block; + font-weight: 700; + color: var(--upm-text); + margin-bottom: 4px; + font-size: 0.8rem; +} + +/* On the checkout step, step 1's "Set monthly ceiling" label is replaced (inside + the step bar) by the chosen ceiling + an inline Edit affordance — see + UpgradeModal.tsx. Kept here next to the step-bar markup it styles. */ +.upm-step__chosen { + display: inline-flex; + align-items: center; + gap: 8px; + color: var(--upm-text); + font-weight: 700; +} +.upm-step__edit { + background: transparent; + border: none; + color: var(--upm-accent); + font-weight: 600; + cursor: pointer; + font-family: inherit; + font-size: 0.75rem; + padding: 2px 8px; + border-radius: 6px; +} +.upm-step__edit:hover { + background: rgba(10, 139, 255, 0.12); +} + +/* Placeholder card for the loading / error / mock states — a dashed, centered + box with helper text. The live Stripe iframe is NOT meant to wear this chrome; + it gets a full-width reset below (data-state="live"). */ +.upm-stripe-mount { + min-height: 220px; + border: 1.5px dashed var(--upm-divider); + border-radius: 12px; + padding: 24px; + display: flex; + flex-direction: column; + align-items: center; + justify-content: center; + text-align: center; + gap: 8px; + color: var(--upm-muted); + font-size: 0.875rem; +} + +/* Live Stripe Embedded Checkout — strip the placeholder card chrome so the + iframe spans the full body width and Stripe renders its two-column + "horizontal" layout. Without this the dashed border + 24px padding + centering + above box the iframe ~48px narrower, dropping it back to single column. */ +.upm-stripe-mount[data-state="live"] { + display: block; + border: none; + padding: 0; + min-height: 0; + text-align: left; +} +.upm-stripe-mount__title { + font-weight: 700; + color: var(--upm-text); + font-size: 0.95rem; +} +.upm-stripe-mount__code { + font-family: ui-monospace, SFMono-Regular, monospace; + font-size: 0.75rem; + padding: 2px 6px; + background: var(--upm-card-bg); + border-radius: 4px; + border: 1px solid var(--upm-divider); + color: var(--upm-text); +} + +/* ── Footer / nav actions ─────────────────────────────────────────────── */ +.upm-footer { + display: flex; + align-items: center; + justify-content: flex-end; + gap: 12px; + padding: 16px 22px; + border-top: 1px solid var(--upm-divider); + background: var(--upm-card-bg); +} +.upm-footer__actions { + display: flex; + gap: 8px; +} +.upm-btn { + padding: 10px 18px; + border-radius: 10px; + border: 1.5px solid transparent; + font-weight: 600; + font-size: 0.9rem; + cursor: pointer; + transition: all 120ms ease; + font-family: inherit; + white-space: nowrap; +} +.upm-btn[data-variant="ghost"] { + background: transparent; + border-color: var(--upm-divider); + color: var(--upm-text); +} +.upm-btn[data-variant="ghost"]:hover { + border-color: var(--upm-accent); + color: var(--upm-accent); +} +.upm-btn[data-variant="primary"] { + background: linear-gradient(135deg, var(--upm-accent), var(--upm-accent-2)); + color: white; +} +.upm-btn[data-variant="primary"]:hover { + transform: translateY(-1px); + box-shadow: 0 6px 18px -6px rgba(10, 139, 255, 0.5); +} +.upm-btn[disabled] { + opacity: 0.5; + cursor: not-allowed; + transform: none !important; + box-shadow: none !important; +} + +/* ── Step 3: confirmation ─────────────────────────────────────────────── */ + +.upm-confirm { + display: flex; + flex-direction: column; + align-items: center; + text-align: center; + padding: 12px 8px 4px; + gap: 12px; +} +.upm-confirm__icon { + font-size: 64px !important; + color: #10b981; + filter: drop-shadow(0 6px 18px rgba(16, 185, 129, 0.35)); +} +.upm-confirm__title { + margin: 6px 0 0; + font-size: 1.5rem; + font-weight: 700; + letter-spacing: -0.015em; + color: var(--upm-text); +} +.upm-confirm__body { + margin: 0; + font-size: 0.95rem; + color: var(--upm-text-muted); + max-width: 380px; + line-height: 1.55; +} +.upm-confirm__summary { + display: flex; + justify-content: space-between; + align-items: center; + width: 100%; + max-width: 360px; + padding: 14px 18px; + background: var(--upm-inset-bg, rgba(0, 0, 0, 0.03)); + border: 1px solid var(--upm-divider); + border-radius: 12px; + margin-top: 6px; + font-size: 0.92rem; +} +.upm-confirm__summary > strong { + color: var(--upm-text); + font-size: 1.05rem; +} +.upm-confirm__note { + margin: 4px 0 0; + font-size: 0.78rem; + color: var(--upm-text-muted); + max-width: 380px; + line-height: 1.5; +} diff --git a/frontend/editor/src/cloud/components/shared/config/configSections/UpgradeModal.tsx b/frontend/editor/src/cloud/components/shared/config/configSections/UpgradeModal.tsx new file mode 100644 index 0000000000..0ca290611b --- /dev/null +++ b/frontend/editor/src/cloud/components/shared/config/configSections/UpgradeModal.tsx @@ -0,0 +1,516 @@ +/** + * Upgrade-to-Processor modal. Three sequential panels inside one frame: + * + * Step 1: Cap selection — local state only, no side effects + * Step 2: Stripe Checkout — POSTs to /api/v1/payg/checkout, mounts the + * Stripe Embedded Checkout iframe (lazy-loaded) + * Step 3: Confirmation — brief "Welcome to Processor" beat before + * the modal closes and the parent's + * {@code onComplete} triggers a wallet refetch + * + *

The Stripe SDK ({@code @stripe/stripe-js} + {@code @stripe/react-stripe-js}) + * is loaded via {@code React.lazy} on a dedicated module so the chunk only + * downloads when the user actually advances to step 2. The main bundle pays + * nothing for users who never open the modal. + * + *

Cap state is held locally — nothing reaches the backend until the user + * commits in step 2. A user who cancels mid-modal leaves no side effects. + */ +import React, { Suspense, useEffect, useState } from "react"; +import { createPortal } from "react-dom"; +import CloseIcon from "@mui/icons-material/CloseRounded"; +import ArrowBackIcon from "@mui/icons-material/ArrowBackRounded"; +import ShieldIcon from "@mui/icons-material/ShieldOutlined"; +import CheckCircleIcon from "@mui/icons-material/CheckCircleRounded"; +import { useTranslation } from "react-i18next"; +// eslint-disable-next-line no-restricted-imports +import "./UpgradeModal.css"; +// eslint-disable-next-line no-restricted-imports +import SpendCapControl from "./SpendCapControl"; + +/** + * Tell the AppConfigModal (or any other full-screen surface listening) that an + * upgrade overlay is opening/closing so it can hide itself rather than stack + * under us. Same window-event pattern the config modal already uses for + * appConfig:navigate / appConfig:notice. + */ +function dispatchOverlay(open: boolean) { + window.dispatchEvent( + new CustomEvent("appConfig:overlay", { detail: { open } }), + ); +} + +// Lazy-loaded so the @stripe/stripe-js bundle only downloads when the user +// reaches step 2. See StripeCheckoutPanel.tsx for the full pattern + the +// chunk-graph reasoning. +const StripeCheckoutPanel = React.lazy( + () => + import("@app/components/shared/config/configSections/StripeCheckoutPanel"), +); + +interface UpgradeModalProps { + open: boolean; + /** + * The caller's team_id. Threaded through to {@link StripeCheckoutPanel} so the + * {@code create-checkout-session} edge function can scope the subscription to + * the right team. + */ + teamId: number; + /** Called when the user closes the modal without completing checkout. */ + onClose: () => void; + /** + * Called after the user confirms cap + completes Stripe checkout. The parent + * is expected to refresh the wallet snapshot which will then show the + * subscribed state. + */ + onComplete: (result: { capUsd: number | null }) => void; + /** ISO 4217 currency code for the cap input. Default USD. */ + currency?: "USD" | "EUR" | "GBP"; + /** + * The team's one-time free grant in documents — the real {@code + * wallet.freeAllowance}, threaded from the free-leader view so the step copy + * quotes the backend's number instead of a hardcoded one. A lifetime grant, + * not a monthly one. + */ + freeLimit: number; + /** + * Per-document rate in minor units for the live "≈ N paid PDFs/month" + * estimate, threaded from {@code wallet.pricePerDocMinor}. For unsubscribed + * teams the backend resolves this from the default pricing policy's USD + * Price (Stripe hasn't assigned the team a currency yet). Null hides the + * estimate. + */ + pricePerDocMinor?: number | null; + /** Lower-case ISO currency of {@link #pricePerDocMinor} (e.g. {@code "usd"}). */ + rateCurrency?: string | null; +} + +type Step = "cap" | "checkout" | "confirm"; + +function currencySymbol(c: UpgradeModalProps["currency"]): string { + switch (c) { + case "EUR": + return "€"; + case "GBP": + return "£"; + default: + return "$"; + } +} + +export default function UpgradeModal({ + open, + teamId, + onClose, + onComplete, + currency = "USD", + freeLimit, + pricePerDocMinor, + rateCurrency, +}: UpgradeModalProps) { + const { t } = useTranslation(); + const [step, setStep] = useState("cap"); + const [capUsd, setCapUsd] = useState(500); + const [noCap, setNoCap] = useState(false); + + // The config modal hides itself while we're open (it listens for this event) + // so the upgrade flow visually REPLACES it instead of stacking inside it. + // Cleanup fires open=false on unmount too, so the config modal can't get + // stuck hidden if we unmount without a clean close. + useEffect(() => { + dispatchOverlay(open); + return () => dispatchOverlay(false); + }, [open]); + + if (!open) { + return null; + } + + const effectiveCap = noCap ? null : capUsd; + const sym = currencySymbol(currency); + + const goToCheckout = () => setStep("checkout"); + const goBackToCap = () => setStep("cap"); + const goToConfirm = () => setStep("confirm"); + + // Modal close → reset internal step so reopening starts at step 1. + const closeAndReset = () => { + setStep("cap"); + onClose(); + }; + + // Portal to document.body so the overlay escapes the config modal's portal / + // stacking context. Without this the fixed-position backdrop layers inside + // the Mantine modal (z-index 1300) instead of over the whole page, producing + // the modal-in-modal look. + return createPortal( +

+
+
e.stopPropagation()}> + {/* Header — title + close. Title stays constant; the step indicator + below tells the user where they are. */} +
+
+ {step === "checkout" && ( + + )} +

+ {step === "confirm" + ? t("payg.upgrade.title.confirm", "You're subscribed") + : t( + "payg.upgrade.title.default", + "Upgrade to Processor plan", + )} +

+
+ +
+ + {/* Step indicator. Hidden on the confirmation panel since the + modal is winding down at that point. */} + {step !== "confirm" && ( +
+
+ 1 + {step === "checkout" ? ( + + {effectiveCap === null + ? t("payg.upgrade.checkout.noCap", "No cap") + : t( + "payg.upgrade.checkout.capValue", + "{{symbol}}{{amount}} / month", + { symbol: sym, amount: effectiveCap }, + )} + + + ) : ( + + {t("payg.upgrade.steps.cap", "Set monthly ceiling")} + + )} +
+
+
+ 2 + + {t("payg.upgrade.steps.payment", "Add payment method")} + +
+
+ )} + +
+ {step === "cap" && ( + + )} + {step === "checkout" && ( + // Keyed on the cap value so editing the cap → returning to + // step 2 unmounts + remounts the panel, triggering a fresh + // POST /checkout for the new cap. Without the key, the + // StripeCheckoutPanel's fetchedRef short-circuits and the + // session keeps the stale cap. + + )} + {step === "confirm" && ( + + )} +
+ + {step !== "checkout" && ( +
+
+ {step === "cap" && ( + <> + + + + )} + {step === "confirm" && ( + + )} +
+
+ )} +
+
+
, + document.body, + ); +} + +// ─── Step 1: cap selection ────────────────────────────────────────────── + +interface CapStepProps { + capUsd: number; + setCapUsd: (v: number) => void; + noCap: boolean; + setNoCap: (v: boolean) => void; + pricePerDocMinor?: number | null; + rateCurrency?: string | null; +} + +function CapStep({ + capUsd, + setCapUsd, + noCap, + setNoCap, + pricePerDocMinor, + rateCurrency, +}: CapStepProps) { + const { t } = useTranslation(); + + return ( + <> +
+ +
+ + {t( + "payg.upgrade.promise.highlight", + "Manual tools stay free, always.", + )} + {" "} + {t( + "payg.upgrade.promise.body", + "You only pay for automation pipelines, AI tools, and API calls — the work that goes beyond a single click. Edit, merge, split, sign, compress as much as you want, no charge.", + )} +
+
+ +

+ {t("payg.upgrade.cap.title", "Set your monthly spend ceiling")} +

+

+ {t( + "payg.upgrade.cap.help", + "We'll never charge above this. Set $0 if you want to keep everything free while testing.", + )} +

+ + {/* Same control the subscribed plan page renders. null = no cap; the + shared control owns the presets, the inline custom-entry pill, the + no-cap chip, and the live processed-PDF estimate. */} + { + if (v === null) { + setNoCap(true); + } else { + setNoCap(false); + setCapUsd(v); + } + }} + pricePerDocMinor={pricePerDocMinor} + currency={rateCurrency} + note={t( + "payg.upgrade.cap.usdNote", + "Estimated in USD. You can adjust your cap any time after subscribing — in your own currency.", + )} + /> + +
+ + {t("payg.upgrade.help.title", "What we count toward billing")} + +
    +
  • + + {t("payg.upgrade.help.automationTitle", "Automation pipelines")} + + {" — "} + {t( + "payg.upgrade.help.automationBody", + "chained tools or scheduled runs that don't need clicks", + )} +
  • +
  • + {t("payg.upgrade.help.aiTitle", "AI tools")} + {" — "} + {t( + "payg.upgrade.help.aiBody", + "summarise, classify, redact, AI-OCR", + )} +
  • +
  • + {t("payg.upgrade.help.apiTitle", "API calls")} + {" — "} + {t( + "payg.upgrade.help.apiBody", + "programmatic access to any Stirling endpoint", + )} +
  • +
+
+ {t( + "payg.upgrade.help.footnote", + "Manual tools — viewing, editing, merging, splitting, signing, watermarking, compressing, manual OCR — are always free, even past 500. The distinction is the type of work, not where you click.", + )} +
+
+ + ); +} + +// ─── Step 2: Stripe Embedded Checkout (lazy-loaded) ──────────────────── + +interface CheckoutStepProps { + teamId: number; + effectiveCap: number | null; + currency: UpgradeModalProps["currency"]; + onComplete: () => void; +} + +function CheckoutStep({ + teamId, + effectiveCap, + currency, + onComplete, +}: CheckoutStepProps) { + const { t } = useTranslation(); + return ( + <> +

+ {t("payg.upgrade.checkout.title", "Add your payment method")} +

+

+ {t( + "payg.upgrade.checkout.help", + "Stripe handles your card details. Stirling never sees them.", + )} +

+ + +
+ {t("payg.upgrade.checkout.loading", "Loading checkout…")} +
+
+ } + > + + + + ); +} + +// ─── Step 3: confirmation ────────────────────────────────────────────── + +interface ConfirmationStepProps { + effectiveCap: number | null; + currency: UpgradeModalProps["currency"]; + freeLimit: number; +} + +function ConfirmationStep({ + effectiveCap, + currency, + freeLimit, +}: ConfirmationStepProps) { + const { t } = useTranslation(); + const sym = currencySymbol(currency); + return ( +
+ +

+ {t("payg.confirm.title", "Welcome to the Processor plan")} +

+

+ {t( + "payg.confirm.body", + "Your team can now process documents with automation, AI, and the API beyond your {{limit}} free PDFs.", + { limit: freeLimit.toLocaleString() }, + )} +

+
+ {t("payg.confirm.summaryLabel", "Monthly ceiling")} + + {effectiveCap === null + ? t("payg.confirm.noCap", "No cap") + : t("payg.confirm.capValue", "{{symbol}}{{amount}} / month", { + symbol: sym, + amount: effectiveCap, + })} + +
+

+ {t( + "payg.confirm.note", + "You can change your cap, cancel, or open the Stripe customer portal any time from this page.", + )} +

+
+ ); +} diff --git a/frontend/editor/src/cloud/components/shared/config/configSections/usageMeters.tsx b/frontend/editor/src/cloud/components/shared/config/configSections/usageMeters.tsx new file mode 100644 index 0000000000..f1012bd2dd --- /dev/null +++ b/frontend/editor/src/cloud/components/shared/config/configSections/usageMeters.tsx @@ -0,0 +1,202 @@ +/** + * Compact usage meters shared by the Plan section and the usage-limit warning + * modals. Kept in their own module (rather than inside Payg/PaygFree) so the + * modals can render a meter without pulling in the upgrade-checkout subtree + * (UpgradeModal, useWallet, etc.). Only depends on i18n + the co-located CSS. + */ +import { useMemo } from "react"; +import { useTranslation } from "react-i18next"; +import { useWallet, type Wallet } from "@app/hooks/useWallet"; +import "@app/components/shared/config/configSections/Payg.css"; +import "@app/components/shared/config/configSections/PaygFree.css"; + +export type MeterState = "FULL" | "WARNED" | "DEGRADED"; + +/** Warn/degrade band for a usage meter (mirrors the BE thresholds). */ +export function meterState( + used: number, + limit: number, +): { state: MeterState; pct: number } { + const pct = limit > 0 ? Math.min(100, (used / limit) * 100) : 100; + const state: MeterState = + pct >= 100 ? "DEGRADED" : pct >= 80 ? "WARNED" : "FULL"; + return { state, pct }; +} + +/** Currency symbol for compact inline use; falls back to the ISO code. */ +function currencySymbol(currency: string | null): string { + switch ((currency ?? "").toLowerCase()) { + case "usd": + return "$"; + case "eur": + return "€"; + case "gbp": + return "£"; + default: + return currency ? currency.toUpperCase() + " " : "$"; + } +} + +// ─── One-time free grant meter ────────────────────────────────────────────── + +export interface FreeSnapshot { + /** One-time free documents used so far (grant − remaining). */ + billableUsed: number; + /** The team's one-time free grant size in documents. */ + billableLimit: number; +} + +/** + * Derive the free-grant snapshot from a wallet. Null (not yet loaded) yields a + * zeroed view over the default 500 grant, the brief first-paint placeholder. + */ +export function freeSnapshotFromWallet(wallet: Wallet | null): FreeSnapshot { + if (!wallet) return { billableUsed: 0, billableLimit: 500 }; + return { + billableUsed: Math.max(0, wallet.freeAllowance - wallet.freeRemaining), + billableLimit: wallet.freeAllowance, + }; +} + +/** + * Read the free-grant snapshot from the live wallet. Falls back to a zeroed + * view over the default grant until the wallet loads. + */ +export function useFreeSnapshot(): FreeSnapshot { + const { wallet } = useWallet(); + return useMemo(() => freeSnapshotFromWallet(wallet), [wallet]); +} + +export function FreeMeterPanel({ snap }: { snap: FreeSnapshot }) { + const { t } = useTranslation(); + const { state, pct } = meterState(snap.billableUsed, snap.billableLimit); + const stateLabel = + state === "DEGRADED" + ? t("payg.free.state.limitReached", "Limit reached") + : state === "WARNED" + ? t("payg.free.state.approachingLimit", "Approaching limit") + : t("payg.free.state.plentyLeft", "Plenty left"); + + return ( +
+
+
+ + {snap.billableUsed.toLocaleString()} + + + {t("payg.free.hero.capSuffix", "/ {{limit}} free PDFs", { + limit: snap.billableLimit.toLocaleString(), + })} + +
+ + + {stateLabel} + +
+ +
+
+
+ +
+ + {t("payg.free.hero.metaCategories", "Automation · AI · API requests")} + +
+
+ ); +} + +// ─── Monthly spend-cap meter ──────────────────────────────────────────────── + +export interface SpendCapSnapshot { + /** Money spent so far this billing period, in major currency units. */ + spent: number; + /** The configured monthly spend cap, in major currency units. */ + cap: number; + /** ISO currency code of {@link spent}/{@link cap}; null falls back to "$". */ + currency: string | null; +} + +/** + * Derive the spend-vs-cap snapshot from a wallet. {@code estimatedBillMinor} is + * this period's charges in minor units; {@code capUsd} is the cap in major + * units. Null wallet (or no cap set) yields a zeroed view. + */ +export function spendCapSnapshotFromWallet( + wallet: Wallet | null, +): SpendCapSnapshot { + if (!wallet) return { spent: 0, cap: 0, currency: null }; + return { + spent: + wallet.estimatedBillMinor != null ? wallet.estimatedBillMinor / 100 : 0, + cap: wallet.capUsd ?? 0, + currency: wallet.currency, + }; +} + +/** + * Sibling of {@link FreeMeterPanel} for the money cap rather than the one-time + * free grant. Shares the same bar/status styling and the cap-state labels + * ({@code payg.state.*}) used by the Plan hero, so it reads as the same meter. + */ +export function SpendCapMeterPanel({ snap }: { snap: SpendCapSnapshot }) { + const { t } = useTranslation(); + const { state, pct } = meterState(snap.spent, snap.cap); + const stateLabel = + state === "DEGRADED" + ? t("payg.state.degraded", "Cap reached") + : state === "WARNED" + ? t("payg.state.warned", "Approaching cap") + : t("payg.state.full", "Healthy"); + const symbol = currencySymbol(snap.currency); + + return ( +
+
+
+ + {symbol} + {snap.spent.toLocaleString()} + + + {t("payg.spendCapMeter.capSuffix", "/ {{amount}} cap", { + amount: `${symbol}${snap.cap.toLocaleString()}`, + })} + +
+ + + {stateLabel} + +
+ +
+
+
+ +
+ + {t( + "payg.spendCapMeter.metaCategories", + "Automation · AI · API spend", + )} + + • + + {t("payg.spendCapMeter.resets", "Resets each billing period")} + +
+
+ ); +} diff --git a/frontend/editor/src/cloud/components/usageLimitModals.ts b/frontend/editor/src/cloud/components/usageLimitModals.ts new file mode 100644 index 0000000000..bb5fc59fef --- /dev/null +++ b/frontend/editor/src/cloud/components/usageLimitModals.ts @@ -0,0 +1,25 @@ +/** + * Imperative API for the usage-limit warning modals. + * + * Call these from anywhere (React or not) to pop a modal. They bridge to the + * always-mounted {@link ./UsageLimitModalHost} host via a window event. No + * arguments and no context plumbing: the modal reads the live wallet itself to + * fill in the usage figures. + * + * import { openFreeLimitModal } from "@app/components/usageLimitModals"; + * openFreeLimitModal(); + */ + +// Internal bridge events, not part of the public API. Use the helpers below. +export const FREE_LIMIT_MODAL_EVENT = "stirling:open-free-limit-modal"; +export const SPEND_CAP_MODAL_EVENT = "stirling:open-spend-cap-modal"; + +/** Open the "free limit reached" modal. Figures come from the live wallet. */ +export function openFreeLimitModal(): void { + window.dispatchEvent(new Event(FREE_LIMIT_MODAL_EVENT)); +} + +/** Open the "spend cap reached" modal. Figures come from the live wallet. */ +export function openSpendCapModal(): void { + window.dispatchEvent(new Event(SPEND_CAP_MODAL_EVENT)); +} diff --git a/frontend/editor/src/desktop/contexts/SaaSTeamContext.tsx b/frontend/editor/src/cloud/contexts/SaaSTeamContext.tsx similarity index 63% rename from frontend/editor/src/desktop/contexts/SaaSTeamContext.tsx rename to frontend/editor/src/cloud/contexts/SaaSTeamContext.tsx index 70f1bc9db8..3e762556bd 100644 --- a/frontend/editor/src/desktop/contexts/SaaSTeamContext.tsx +++ b/frontend/editor/src/cloud/contexts/SaaSTeamContext.tsx @@ -7,13 +7,16 @@ import { useCallback, } from "react"; import apiClient from "@app/services/apiClient"; -import { authService } from "@app/services/authService"; -import { connectionModeService } from "@app/services/connectionModeService"; +import { useTeamAuth } from "@app/auth/teamSession"; /** - * Desktop implementation of SaaS Team Context - * Provides team management for users connected to SaaS backend - * CRITICAL: Only active when in SaaS mode - all API calls check connection mode first + * Shared (cloud) SaaS Team Context. + * + * Provides team management for authenticated (non-anonymous) users. The + * platform-specific auth bits — whether teams may be used at all, and how to + * refresh derived auth state after a membership change — come from the + * {@code @app/auth/teamSession} seam (Supabase web session on saas, authService + * on desktop), keeping this context free of any platform auth coupling. */ interface Team { @@ -92,42 +95,11 @@ export function SaaSTeamProvider({ children }: { children: ReactNode }) { TeamInvitation[] >([]); const [loading, setLoading] = useState(true); - const [isSaasMode, setIsSaasMode] = useState(false); - const [isAuthenticated, setIsAuthenticated] = useState(false); - // Check if in SaaS mode and authenticated - useEffect(() => { - const checkAccess = async () => { - const mode = await connectionModeService.getCurrentMode(); - const auth = await authService.isAuthenticated(); - setIsSaasMode(mode === "saas"); - setIsAuthenticated(auth); - }; - - checkAccess(); - - // Subscribe to connection mode changes - const unsubscribe = - connectionModeService.subscribeToModeChanges(checkAccess); - return unsubscribe; - }, []); - - // Subscribe to auth changes - useEffect(() => { - const unsubscribe = authService.subscribeToAuth((status) => { - setIsAuthenticated(status === "authenticated"); - }); - return unsubscribe; - }, []); + const { canUseTeams, refreshAfterMembershipChange } = useTeamAuth(); const fetchMyTeams = useCallback(async () => { - // CRITICAL: Only fetch if in SaaS mode and authenticated - if (!isSaasMode || !isAuthenticated) { - console.log( - "[SaaSTeamContext] Skipping team fetch - not in SaaS mode or not authenticated", - ); - return null; - } + if (!canUseTeams) return null; try { const response = await apiClient.get("/api/v1/team/my", { @@ -136,56 +108,34 @@ export function SaaSTeamProvider({ children }: { children: ReactNode }) { setTeams(response.data); const activeTeam = response.data[0]; - console.log("[SaaSTeamContext] Current team set:", { - teamId: activeTeam?.teamId, - name: activeTeam?.name, - isPersonal: activeTeam?.isPersonal, - isLeader: activeTeam?.isLeader, - }); setCurrentTeam(activeTeam || null); return activeTeam || null; } catch (error) { console.error("[SaaSTeamContext] Failed to fetch teams:", error); return null; } - }, [isSaasMode, isAuthenticated]); + }, [canUseTeams]); - const fetchTeamMembers = useCallback( - async (teamId: number) => { - // CRITICAL: Only fetch if in SaaS mode and authenticated - if (!isSaasMode || !isAuthenticated) { - console.log( - "[SaaSTeamContext] Skipping members fetch - not in SaaS mode or not authenticated", - ); - return; - } - - try { - const response = await apiClient.get( - `/api/v1/team/${teamId}/members`, - { suppressErrorToast: true }, - ); - setTeamMembers(response.data); - } catch (error) { - console.error("[SaaSTeamContext] Failed to fetch team members:", error); - } - }, - [isSaasMode, isAuthenticated], - ); + const fetchTeamMembers = useCallback(async (teamId: number) => { + try { + const response = await apiClient.get( + `/api/v1/team/${teamId}/members`, + { suppressErrorToast: true }, + ); + setTeamMembers(response.data); + } catch (error) { + console.error("[SaaSTeamContext] Failed to fetch team members:", error); + } + }, []); const fetchTeamInvitations = useCallback( async (teamId?: number) => { - // CRITICAL: Only fetch if in SaaS mode and authenticated - if (!isSaasMode || !isAuthenticated || !teamId) { - return; - } + if (!canUseTeams || !teamId) return; try { const response = await apiClient.get( `/api/v1/team/${teamId}/invitations`, - { - suppressErrorToast: true, - }, + { suppressErrorToast: true }, ); setTeamInvitations(response.data); } catch (error) { @@ -195,26 +145,17 @@ export function SaaSTeamProvider({ children }: { children: ReactNode }) { ); } }, - [isSaasMode, isAuthenticated], + [canUseTeams], ); const fetchReceivedInvitations = useCallback(async () => { - // CRITICAL: Only fetch if in SaaS mode and authenticated - if (!isSaasMode || !isAuthenticated) { - return; - } - - console.log("[SaaSTeamContext] Fetching received team invitations"); + if (!canUseTeams) return; try { const response = await apiClient.get( "/api/v1/team/invitations/pending", { suppressErrorToast: true }, ); - console.log( - "[SaaSTeamContext] Received invitations response:", - response.data, - ); setReceivedInvitations(response.data); } catch (error) { console.error( @@ -222,14 +163,13 @@ export function SaaSTeamProvider({ children }: { children: ReactNode }) { error, ); } - }, [isSaasMode, isAuthenticated]); + }, [canUseTeams]); useEffect(() => { - if (isSaasMode && isAuthenticated) { + if (canUseTeams) { fetchMyTeams(); fetchReceivedInvitations(); } else { - // Clear state when not in SaaS mode or not authenticated setTeams([]); setCurrentTeam(null); setTeamMembers([]); @@ -237,15 +177,10 @@ export function SaaSTeamProvider({ children }: { children: ReactNode }) { setReceivedInvitations([]); setLoading(false); } - }, [isSaasMode, isAuthenticated, fetchMyTeams, fetchReceivedInvitations]); + }, [canUseTeams, fetchMyTeams, fetchReceivedInvitations]); useEffect(() => { - if ( - currentTeam && - !currentTeam.isPersonal && - isSaasMode && - isAuthenticated - ) { + if (currentTeam && !currentTeam.isPersonal) { fetchTeamMembers(currentTeam.teamId); // Only fetch invitations if user is team leader if (currentTeam.isLeader) { @@ -258,17 +193,10 @@ export function SaaSTeamProvider({ children }: { children: ReactNode }) { setTeamInvitations([]); } setLoading(false); - }, [ - currentTeam, - isSaasMode, - isAuthenticated, - fetchTeamMembers, - fetchTeamInvitations, - ]); + }, [currentTeam, fetchTeamMembers, fetchTeamInvitations]); const inviteUser = async (email: string) => { if (!currentTeam) throw new Error("No current team"); - if (!isSaasMode) throw new Error("Not in SaaS mode"); await apiClient.post("/api/v1/team/invite", { teamId: currentTeam.teamId, @@ -277,59 +205,7 @@ export function SaaSTeamProvider({ children }: { children: ReactNode }) { await fetchTeamInvitations(currentTeam.teamId); }; - const acceptInvitation = async (token: string) => { - if (!isSaasMode) throw new Error("Not in SaaS mode"); - - await apiClient.post(`/api/v1/team/invitations/${token}/accept`); - await fetchReceivedInvitations(); - await refreshTeams(); - // Note: Desktop doesn't have refreshCredits/refreshSession like SaaS - }; - - const rejectInvitation = async (token: string) => { - if (!isSaasMode) throw new Error("Not in SaaS mode"); - - await apiClient.post(`/api/v1/team/invitations/${token}/reject`); - await fetchReceivedInvitations(); - }; - - const cancelInvitation = async (invitationId: number) => { - if (!isSaasMode) throw new Error("Not in SaaS mode"); - - await apiClient.delete(`/api/v1/team/invitations/${invitationId}`); - if (currentTeam) { - await fetchTeamInvitations(currentTeam.teamId); - } - }; - - const removeMember = async (memberId: number) => { - if (!currentTeam) throw new Error("No current team"); - if (!isSaasMode) throw new Error("Not in SaaS mode"); - - await apiClient.delete( - `/api/v1/team/${currentTeam.teamId}/members/${memberId}`, - ); - await refreshTeams(); - await fetchTeamMembers(currentTeam.teamId); - }; - - const leaveTeam = async () => { - if (!currentTeam) throw new Error("No current team"); - if (!isSaasMode) throw new Error("Not in SaaS mode"); - - await apiClient.post(`/api/v1/team/${currentTeam.teamId}/leave`); - await refreshTeams(); - // Note: Desktop doesn't have refreshCredits/refreshSession like SaaS - }; - const refreshTeams = useCallback(async () => { - if (!isSaasMode || !isAuthenticated) { - console.log( - "[SaaSTeamContext] Skipping refresh - not in SaaS mode or not authenticated", - ); - return; - } - const newCurrentTeam = await fetchMyTeams(); await fetchReceivedInvitations(); if (newCurrentTeam && !newCurrentTeam.isPersonal) { @@ -340,14 +216,50 @@ export function SaaSTeamProvider({ children }: { children: ReactNode }) { } } }, [ - isSaasMode, - isAuthenticated, fetchMyTeams, fetchReceivedInvitations, fetchTeamMembers, fetchTeamInvitations, ]); + const acceptInvitation = async (token: string) => { + await apiClient.post(`/api/v1/team/invitations/${token}/accept`); + await fetchReceivedInvitations(); + await refreshTeams(); + await refreshAfterMembershipChange(); + }; + + const rejectInvitation = async (token: string) => { + await apiClient.post(`/api/v1/team/invitations/${token}/reject`); + await fetchReceivedInvitations(); + }; + + const cancelInvitation = async (invitationId: number) => { + await apiClient.delete(`/api/v1/team/invitations/${invitationId}`); + if (currentTeam) { + await fetchTeamInvitations(currentTeam.teamId); + } + }; + + const removeMember = async (memberId: number) => { + if (!currentTeam) throw new Error("No current team"); + + await apiClient.delete( + `/api/v1/team/${currentTeam.teamId}/members/${memberId}`, + ); + await refreshTeams(); + await fetchTeamMembers(currentTeam.teamId); + // No need to refresh session/credits: the team leader's status hasn't changed + }; + + const leaveTeam = async () => { + if (!currentTeam) throw new Error("No current team"); + + await apiClient.post(`/api/v1/team/${currentTeam.teamId}/leave`); + await refreshTeams(); + await refreshAfterMembershipChange(); + }; + const isTeamLeader = currentTeam?.isLeader ?? false; const isPersonalTeam = currentTeam?.isPersonal ?? true; diff --git a/frontend/editor/src/cloud/hooks/useWallet.ts b/frontend/editor/src/cloud/hooks/useWallet.ts new file mode 100644 index 0000000000..0222cc77c0 --- /dev/null +++ b/frontend/editor/src/cloud/hooks/useWallet.ts @@ -0,0 +1,432 @@ +/** + * Hook backing the PAYG Plan page. Wraps {@code GET /api/v1/payg/wallet} + * (served by {@code PaygWalletController} once Wave 1 BE lands; until then + * the dev preview route synthesises a wallet from localStorage) and exposes + * mutations for marking-subscribed and updating-the-cap. + * + *

Render efficiency

+ * + * The hook is designed so {@code Plan}, {@code PaygFreeLeader/Member}, and + * {@code PaygLeader/Member} re-render only on actual data change: + * + *
    + *
  • {@link Wallet} snapshot is stored as a plain object — but every + * successful fetch deep-compares with the previous snapshot and reuses + * the prior reference if the payload is unchanged. Consumers that hold + * a stable {@code wallet} reference get stable child memoisation. + *
  • The returned {@link UseWalletResult} keeps stable callback identities + * via {@code useCallback}. {@code Plan} can pass {@code markSubscribed} + * to {@code UpgradeModal} without forcing a remount. + *
  • {@code refetch / markSubscribed / updateCap} bump an internal counter + * that the {@code useEffect} watches — no global state plumbing. + *
  • A monotonic {@code requestId} ref drops stale responses so a slow + * refetch from tick=N can't overwrite a faster one from tick=N+1 + * (out-of-order resolution would otherwise show old data). + *
+ * + *

Mutation semantics

+ * + * Both {@code markSubscribed} and {@code updateCap} resolve only after the + * post-mutation wallet refetch completes. So callers like the cap-editor + * "Update cap" button that gate a {@code loading} state on the returned + * promise see the UI flip exactly once the new state is visible — no + * intermediate flash of the old value. + * + *

Dev preview fallback

+ * + * When the hook is rendered outside the saas app (e.g. on {@code + * /dev/payg-preview} during local design work) the {@code AppConfigContext} + * provider is not mounted and no backend is available. The hook detects that + * via the {@code @app/hooks/walletDevPreview} seam and, when it returns a live + * channel, falls back to a synthesised snapshot whose subscription state is + * read from {@code localStorage}. The detection + synthesis (which read + * {@code import.meta.env}, {@code window.location} and web storage — all banned + * in cloud/) live in the saas leaf's impl of that seam; this hook just consults + * it. Desktop's cascade falls through to the cloud default (no dev preview), so + * it always fetches the real wallet. + */ +import { useCallback, useEffect, useRef, useState } from "react"; +import apiClient from "@app/services/apiClient"; +import { createPortalSession } from "@app/services/billing"; +import { openExternal } from "@app/platform/openExternal"; +import { getWalletDevPreview } from "@app/hooks/walletDevPreview"; + +// ─── Public types ─────────────────────────────────────────────────────── + +export type WalletStatus = "free" | "subscribed"; +export type WalletRole = "leader" | "member"; + +/** + * A single team member's billing-relevant info — name + email for the avatar + * row, {@code spendUnits} for their per-member usage display. Mirrors a row of + * the backend's {@code members} array on {@code WalletSnapshot} (joined with + * {@code team_memberships}). + */ +export interface WalletMember { + /** Supabase user id of the member. */ + userId: string; + name: string; + email: string; + /** Member's current-period billable spend. */ + spendUnits: number; +} + +/** + * Per-category breakdown of current-period spend in billable units. The + * categories mirror the {@code FeatureGate} buckets the backend tracks: + * server-side tool calls ({@code api}), AI-backed tools ({@code ai}), and + * pipeline / automation runs ({@code automation}). Numbers sum to {@code + * billableUsed} (modulo rounding in mock data). + */ +export interface WalletCategoryBreakdown { + api: number; + ai: number; + automation: number; +} + +/** Mirror of the backend's {@code WalletSnapshot} record (the JSON returned from {@code GET /api/v1/payg/wallet}). */ +export interface Wallet { + /** + * The caller's primary team_id. Needed when invoking Supabase edge functions + * (create-checkout-session, etc.) that run outside Spring Security and have + * no other way to resolve the caller's team. May be null on the synthetic + * empty snapshot returned to anonymous / team-less callers. + */ + teamId: number | null; + status: WalletStatus; + role: WalletRole; + /** + * ISO yyyy-mm-dd. The Stripe subscription's current period when subscribed; + * the calendar month for free teams. + */ + billingPeriodStart: string; + billingPeriodEnd: string; + /** + * For a free team: the one-time free documents used so far ({@code + * freeAllowance − freeRemaining}). For a subscribed team: documents + * processed this month across automation + AI + API. + */ + billableUsed: number; + /** + * The team's document ceiling for the matching window: the one-time free + * grant ({@code freeAllowance}) for free teams; the monthly paid-doc cap + * {@code floor(cap / perDocRate)} for capped subscribed teams; null when + * subscribed with no cap (uncapped). + */ + billableLimit: number | null; + /** + * The team's one-time free document grant size — the "N" in "X of N free". + * A lifetime grant ({@code pricing_policy.free_tier_units}): it never resets + * and is not lost when the team subscribes. + */ + freeAllowance: number; + /** + * One-time free documents still available to the team + * ({@code payg_team_extensions.free_units_remaining}). 0 = grant exhausted. + * Survives subscribing — a subscribed team keeps any unused grant. + */ + freeRemaining: number; + /** + * Paid per-document rate in minor units of {@link Wallet#currency} (may be + * fractional); null when the rate can't be resolved — render "unknown", + * never substitute. + */ + pricePerDocMinor: number | null; + /** Lower-case ISO 4217 currency of the subscription's Stripe Price; null when unknown. */ + currency: string | null; + /** + * Estimated charges so far this period in minor units of currency: paid + * (Stripe-metered) documents this period × rate. The free portion was + * already netted out at charge time. Informational — the Stripe invoice + * is authoritative. Null when the rate is unknown. + */ + estimatedBillMinor: number | null; + /** Monthly cap in major currency units when subscribed; null when noCap or status=='free'. */ + capUsd: number | null; + /** Only meaningful when status=='subscribed'. */ + noCap: boolean; + /** Stripe subscription id when subscribed; null when free. */ + stripeSubscriptionId: string | null; + /** Current-period spend in billable units. */ + spendUnitsThisPeriod: number; + /** Per-category spend breakdown (api / ai / automation). */ + categoryBreakdown: WalletCategoryBreakdown; + /** + * Team members, populated for the leader view; empty for members or + * single-seat tenants. Leader-vs-member is still resolved via {@link + * Wallet#role} — this field just carries the per-member rows the leader's + * sub-cap table needs. + */ + members: WalletMember[]; + /** + * Recent billable-activity rows. V1 returns {@code []} from the backend; + * the field exists so the Plan page can render an empty state without + * branching on undefined. Each entry is a {@code Record} + * because the activity-row shape is not yet finalised — when the meter- + * event surface lands, this widens to a real interface. + */ + recent: Array>; +} + +export interface UseWalletResult { + wallet: Wallet | null; + loading: boolean; + error: string | null; + /** Force a refetch — e.g. after Stripe redirects back into the app. */ + refetch: () => Promise; + /** + * Dev-only side-channel that simulates the Stripe webhook flipping the + * team to subscribed. Used by {@code UpgradeModal} when the backend is + * running the mock checkout — the real flow waits for the webhook + * instead and the next {@code refetch} picks up the change. Resolves + * once the post-mutation refetch completes. + */ + markSubscribed: (capUsd: number | null) => Promise; + /** + * Update the team's monthly cap. {@code null} means "no cap". Resolves + * once the post-mutation refetch completes so a save-button + * {@code loading} state can be safely cleared on resolution. + */ + updateCap: (capUsd: number | null) => Promise; + /** + * Mint a Stripe Customer Portal session and send the user to it. Mints the + * session via the {@code @app/services/billing} seam (passing the caller's + * {@code teamId}, which the PAYG portal edge function needs to resolve the + * team outside Spring Security) and opens the returned URL via the + * {@code @app/platform/openExternal} seam — so web and desktop each route it + * the platform-appropriate way (new tab on web, system browser on desktop). + * Throws on error so the caller can show a friendly toast — notably 404 + * {@code team_not_subscribed}. + */ + openPortal: () => Promise; +} + +// ─── Implementation ───────────────────────────────────────────────────── + +/** + * Stable reference reuse — if the new payload deep-equals the previous one, + * return the previous object so React's reference check short-circuits child + * renders. Walks the top-level scalars first (cheapest), then the nested + * {@code categoryBreakdown} object, then the {@code members} array. The + * {@code recent} array is identity-compared only — Wave 1 always returns + * {@code []} so a reference-stability check is sufficient; we'll deepen + * this once the activity surface lands. + */ +function reuseIfEqual(prev: Wallet | null, next: Wallet): Wallet { + if (!prev) return next; + if ( + prev.status !== next.status || + prev.teamId !== next.teamId || + prev.role !== next.role || + prev.billingPeriodStart !== next.billingPeriodStart || + prev.billingPeriodEnd !== next.billingPeriodEnd || + prev.billableUsed !== next.billableUsed || + prev.billableLimit !== next.billableLimit || + prev.freeAllowance !== next.freeAllowance || + prev.freeRemaining !== next.freeRemaining || + prev.pricePerDocMinor !== next.pricePerDocMinor || + prev.currency !== next.currency || + prev.estimatedBillMinor !== next.estimatedBillMinor || + prev.capUsd !== next.capUsd || + prev.noCap !== next.noCap || + prev.stripeSubscriptionId !== next.stripeSubscriptionId || + prev.spendUnitsThisPeriod !== next.spendUnitsThisPeriod + ) { + return next; + } + if (prev.recent.length !== next.recent.length) { + return next; + } + if ( + prev.categoryBreakdown.api !== next.categoryBreakdown.api || + prev.categoryBreakdown.ai !== next.categoryBreakdown.ai || + prev.categoryBreakdown.automation !== next.categoryBreakdown.automation + ) { + return next; + } + if (prev.members.length !== next.members.length) { + return next; + } + for (let i = 0; i < prev.members.length; i++) { + const a = prev.members[i]; + const b = next.members[i]; + if ( + a.userId !== b.userId || + a.name !== b.name || + a.email !== b.email || + a.spendUnits !== b.spendUnits + ) { + return next; + } + } + // recent length-mismatch already returned `next` above; content (Wave 1 = []) is identical + // otherwise, so reuse the prior reference for stable child memoisation. + return prev; +} + +export function useWallet(): UseWalletResult { + // Resolved once: the dev-preview side-channel when rendered outside the real + // app (saas /dev/payg-preview route), else null (every real build + desktop). + // The detection + synthesis live behind the @app/hooks/walletDevPreview seam + // because they read import.meta.env / window.location / localStorage, which + // cloud/ may not touch directly. + const devPreview = useRef(getWalletDevPreview()).current; + + const [wallet, setWallet] = useState(null); + const [loading, setLoading] = useState(true); + const [error, setError] = useState(null); + const [refetchTick, setRefetchTick] = useState(0); + + // Monotonic request id — used to discard stale responses if a faster + // refetch lands first. Only the latest issued id is permitted to commit + // its result. + const latestReqId = useRef(0); + + // Promise tracking the most recent in-flight load. Mutations await this + // so their resolution semantics are "the new state is visible," not + // "the request fired." Cleared when no load is pending. + const inFlight = useRef | null>(null); + + useEffect(() => { + const reqId = ++latestReqId.current; + let cancelled = false; + + const promise = (async () => { + setLoading(true); + setError(null); + + if (devPreview) { + const synth = devPreview.buildWallet(devPreview.role()); + if (cancelled || reqId !== latestReqId.current) return; + setWallet((prev) => reuseIfEqual(prev, synth)); + setLoading(false); + return; + } + + try { + const res = await apiClient.get("/api/v1/payg/wallet"); + if (cancelled || reqId !== latestReqId.current) return; + setWallet((prev) => reuseIfEqual(prev, res.data)); + } catch (e: unknown) { + if (cancelled || reqId !== latestReqId.current) return; + console.warn("[useWallet] fetch failed", e); + setError(e instanceof Error ? e.message : "Failed to load wallet"); + } finally { + if (!cancelled && reqId === latestReqId.current) { + setLoading(false); + } + } + })(); + + inFlight.current = promise; + + return () => { + cancelled = true; + // Don't clear inFlight here — let it resolve so mutations awaiting it + // still see a definitive "load completed" point. The reqId guard + // upstream ensures stale results don't commit. + }; + }, [devPreview, refetchTick]); + + const refetch = useCallback(async () => { + setRefetchTick((t) => t + 1); + // Snapshot the next-tick promise so the caller awaits this refetch + // specifically — the in-flight ref will be updated to it on the next + // effect run, but we can't reference that synchronously, so settle for + // a microtask handoff: await the *current* effect to flush, then await + // the new in-flight promise. + await Promise.resolve(); + if (inFlight.current) { + await inFlight.current; + } + }, []); + + const markSubscribed = useCallback( + async (capUsd: number | null) => { + if (devPreview) { + devPreview.markSubscribed(); + await refetch(); + return; + } + const noCap = capUsd === null; + // The dev side-channel only exists when the BE mock service is + // running (FE-branch local dev). Once the real backend (PR #6574) + // is in play, /dev/mark-subscribed is removed and the webhook + // (customer.subscription.created) is what flips the team to + // subscribed. We swallow 404s so the modal's completion path — + // which awaits this promise before rendering the confirmation + // screen — doesn't error out on a perfectly normal "the real + // backend doesn't expose this dev hook" response. A subsequent + // refetch picks up the webhook-driven flip whenever it lands. + try { + await apiClient.post("/api/v1/payg/dev/mark-subscribed", { + capUsd: capUsd ?? 0, + noCap, + }); + } catch (e: unknown) { + const status = + typeof e === "object" && e !== null && "response" in e + ? (e as { response?: { status?: number } }).response?.status + : undefined; + if (status === 404) { + // Real BE in play — webhook will land the subscription + // state; log and continue. Loud-but-harmless so the dev + // notices their /dev/mark-subscribed isn't wired up. + console.info( + "[useWallet] /dev/mark-subscribed not available (404) — relying on Stripe webhook to flip subscription state", + ); + } else { + throw e; + } + } + await refetch(); + }, + [devPreview, refetch], + ); + + const updateCap = useCallback( + async (capUsd: number | null) => { + const noCap = capUsd === null; + if (devPreview) { + await refetch(); + return; + } + await apiClient.patch("/api/v1/payg/cap", { + capUsd: capUsd ?? 0, + noCap, + }); + await refetch(); + }, + [devPreview, refetch], + ); + + const openPortal = useCallback(async () => { + if (devPreview) { + // No real Stripe in dev preview — open a placeholder so the click still + // feels alive. Routed through the openExternal seam to stay portable. + await openExternal("https://billing.stripe.com/p/login/mock"); + return; + } + // Mint the portal session through the billing seam, passing teamId: the + // PAYG portal edge function needs it to resolve the caller's team outside + // Spring Security (its RPC enforces team membership). Then hand the URL to + // the openExternal seam so each platform routes it appropriately. The seam + // throws on error (e.g. 404 team_not_subscribed) so callers can toast. + const teamId = wallet?.teamId; + if (teamId == null) { + throw new Error("No team resolved yet"); + } + const { url } = await createPortalSession({ teamId }); + await openExternal(url); + }, [devPreview, wallet?.teamId]); + + return { + wallet, + loading, + error, + refetch, + markSubscribed, + updateCap, + openPortal, + }; +} diff --git a/frontend/editor/src/cloud/hooks/walletDevPreview.ts b/frontend/editor/src/cloud/hooks/walletDevPreview.ts new file mode 100644 index 0000000000..7a1477cb4f --- /dev/null +++ b/frontend/editor/src/cloud/hooks/walletDevPreview.ts @@ -0,0 +1,46 @@ +/** + * Wallet dev-preview seam (@app/hooks/walletDevPreview). + * + * The cloud/ layer is the SHARED hosted experience consumed by BOTH the saas + * (web) and desktop (Tauri) leaves, so it must stay platform-portable: it can't + * read {@code import.meta.env}, {@code window.location} or web storage directly + * (the cloud ESLint guardrail enforces this). The PAYG dev-preview route + * ({@code /dev/payg-preview}) is a saas-only local-design affordance that + * synthesises a wallet from {@code localStorage} when the real backend isn't + * mounted — all three of those banned reads. {@link useWallet} reaches that + * affordance through this seam instead. + * + * This module is the DEFAULT + the shared TypeScript contract: it reports + * "not a dev preview" so the real backend fetch always runs. The saas leaf + * shadows it with saas/hooks/walletDevPreview.ts, which supplies the synthesis; + * desktop has no dev-preview route, so the cascade falls through to this + * default. Returning {@code null} from {@link getWalletDevPreview} is the + * canonical "no dev preview active" signal. + */ +import type { Wallet, WalletRole } from "@app/hooks/useWallet"; + +/** + * The dev-preview side-channel {@link useWallet} consumes when rendered outside + * the real app (the saas {@code /dev/payg-preview} route). When active it stands + * in for the backend: {@link buildWallet} synthesises the snapshot and + * {@link markSubscribed} flips the simulated subscription state. {@code null} + * (the cloud default + every desktop build) means "no dev preview — fetch the + * real wallet". + */ +export interface WalletDevPreview { + /** Synthesise the dev-preview wallet snapshot (subscription state from storage). */ + buildWallet: (role: WalletRole) => Wallet; + /** Best-effort role read for the preview — flips per {@code ?role=member}. */ + role: () => WalletRole; + /** Flip the simulated subscription state to subscribed (persisted across reload). */ + markSubscribed: () => void; +} + +/** + * Resolve the active dev-preview side-channel, or {@code null} when we're in a + * real build / on a real route. The cloud default + desktop always return + * {@code null}; the saas leaf returns a live channel only on the dev route. + */ +export function getWalletDevPreview(): WalletDevPreview | null { + return null; +} diff --git a/frontend/editor/src/cloud/platform/openExternal.ts b/frontend/editor/src/cloud/platform/openExternal.ts new file mode 100644 index 0000000000..d4c06a37de --- /dev/null +++ b/frontend/editor/src/cloud/platform/openExternal.ts @@ -0,0 +1,28 @@ +/** + * Open-external-URL seam (@app/platform/openExternal). + * + * The cloud/ layer is the SHARED hosted experience consumed by BOTH the saas + * (web) and desktop (Tauri) leaves. Opening a URL in the user's real browser + * differs per platform — saas hands it to the browser via window.open / + * location.assign, desktop hands it to the OS via the Tauri shell plugin so it + * escapes the embedded webview. Cloud code must not reach either of those + * directly, so it opens external URLs through this seam. + * + * This module is the DEFAULT + the shared TypeScript contract. Real builds + * shadow it with saas/platform/openExternal.ts and desktop/platform/ + * openExternal.ts; this default body is only reached by the cloud-standalone + * typecheck, so it throws to make an accidental real-build resolution loud. + */ + +/** Opens an external URL in the user's system browser. */ +export type OpenExternal = (url: string) => Promise; + +/** + * Opens an external URL in the user's system browser. Each platform supplies + * its own implementation; this default is never reached in a real build. + */ +export const openExternal: OpenExternal = async ( + _url: string, +): Promise => { + throw new Error("openExternal: platform impl required"); +}; diff --git a/frontend/editor/src/cloud/services/billing.ts b/frontend/editor/src/cloud/services/billing.ts new file mode 100644 index 0000000000..7f049cd540 --- /dev/null +++ b/frontend/editor/src/cloud/services/billing.ts @@ -0,0 +1,80 @@ +/** + * Billing data seam (@app/services/billing). + * + * Creating Stripe Checkout / Customer Portal sessions touches platform-specific + * transport (saas: supabase-js web client; desktop: Tauri native HTTP with an + * explicit bearer + deep-link return URL). Cloud code can't reach those + * directly, so it mints sessions through this seam. This module is the DEFAULT + + * shared contract; saas/services/billing.ts and desktop/services/billing.ts + * shadow it. The default bodies throw so an accidental real-build resolution is + * loud (only the cloud-standalone typecheck reaches them). + */ + +/** + * Parameters for {@link createCheckoutSession}, which drives the PAYG + * {@code create-checkout-session} edge function (see StripeCheckoutPanel). The + * function runs outside Spring Security, so it needs the caller's {@link teamId} + * (it can't resolve the team from the JWT alone). The platform impl supplies the + * return URL itself (browser origin on web, deep-link scheme on desktop), so it + * is intentionally NOT part of this shape. + */ +export interface CheckoutParams { + /** The caller's team id. Required — scopes the PAYG subscription. */ + teamId: number; + /** Lower-case 3-letter ISO currency (e.g. {@code "gbp"}). Selects the Stripe Price. */ + currency?: string; + /** Billing email for the Checkout Session; maps to Stripe {@code customer_email} when the team has no customer yet. */ + billingOwnerEmail?: string | null; +} + +/** + * Result of {@link createCheckoutSession}. Embedded checkout yields a + * {@code clientSecret}; hosted checkout yields a {@code url}. Exactly one is set. + */ +export interface CheckoutSession { + /** Stripe Checkout Session client secret for embedded mode. */ + clientSecret?: string; + /** Hosted Stripe Checkout URL for redirect mode. */ + url?: string; + /** Non-prod sentinel: a stubbed secret (prefixed {@code cs_mock_}) renders a placeholder instead of a real iframe. */ + mock?: boolean; +} + +/** Result of {@link createPortalSession}. */ +export interface PortalSession { + /** Stripe Customer Portal URL to send the user to. */ + url: string; +} + +/** + * Parameters for {@link createPortalSession}. The PAYG portal edge function + * needs the caller's {@code teamId} (runs outside Spring Security). + */ +export interface PortalParams { + /** The caller's team id; required by the PAYG portal edge function. */ + teamId: number; +} + +/** Create a Stripe Checkout Session via the SaaS billing backend (platform impl required). */ +export async function createCheckoutSession( + _params: CheckoutParams, +): Promise { + throw new Error("billing: platform impl required"); +} + +/** Mint a Stripe Customer Portal session via the SaaS billing backend (platform impl required). */ +export async function createPortalSession( + _params: PortalParams, +): Promise { + throw new Error("billing: platform impl required"); +} + +/** + * The Stripe publishable key used to initialise {@code loadStripe()} for + * embedded checkout. Sourced through the seam because cloud code may not read + * {@code import.meta.env} directly. The cloud default returns "" so the checkout + * component falls back to its mock placeholder rather than throwing. + */ +export function getStripePublishableKey(): string { + return ""; +} diff --git a/frontend/editor/src/cloud/services/paygErrorInterceptor.ts b/frontend/editor/src/cloud/services/paygErrorInterceptor.ts new file mode 100644 index 0000000000..0f6a29ffb6 --- /dev/null +++ b/frontend/editor/src/cloud/services/paygErrorInterceptor.ts @@ -0,0 +1,149 @@ +/** + * Classifies and reacts to PAYG-specific error responses surfaced by the + * backend's {@code EntitlementGuard}. Three sentinels are recognised: + * + *
    + *
  • {@code 402 FEATURE_DEGRADED} — an authenticated (JWT/web) team hit a + * billable feature it no longer has: a free team that spent its one-time + * allowance, or a subscribed team over its monthly spending cap. Which + * one is told by the {@code subscribed} field on the body.
  • + *
  • {@code 402 PAYG_LIMIT_REACHED} — same situation reached via an API key + * (programmatic client). Also carries {@code subscribed}.
  • + *
  • {@code 401 SIGNUP_REQUIRED} — anonymous (guest) user hit a billable + * endpoint. Opens the signup modal (a different flow) via a + * {@code CustomEvent}.
  • + *
+ * + * For the two limit sentinels we pop the matching usage-limit modal (free → + * "free limit reached", subscribed → "spend cap reached") and show NO toast — + * the modal is the actionable surface. The modals read the live wallet for the + * usage figures, so we only need to decide which one to open. + * + * The classifier is exported separately from the handler so unit tests can + * exercise the parsing logic without touching the modal side effects. + */ +import { + openFreeLimitModal, + openSpendCapModal, +} from "@app/components/usageLimitModals"; + +/** + * Possible PAYG entitlement sentinels the EntitlementGuard returns. + * {@code null} when the error is not a PAYG entitlement response. + */ +export type PaygErrorKind = + | "FEATURE_DEGRADED" + | "PAYG_LIMIT_REACHED" + | "SIGNUP_REQUIRED"; + +/** + * Detail payload broadcast on {@code payg:signupRequired} when an anonymous + * user hits a billable endpoint. The listener (a Bootstrap component near + * the app root) opens a modal whose copy is parameterised by + * {@link #category}. + */ +export interface PaygSignupRequiredDetail { + /** Category that triggered the gate — {@code AI}, {@code AUTOMATION}, or {@code API}. */ + category: string | null; +} + +/** + * Inspect an axios-style error and decide whether it's one of the known + * PAYG sentinels. Returns the kind, or {@code null} if it isn't. + * + * The check is intentionally strict (status code AND body.error sentinel) + * so we don't hijack incidental 401/402 responses from other endpoints — + * notably the existing session-expired 401 flow. + */ +export function classifyPaygError(error: unknown): PaygErrorKind | null { + if (!error || typeof error !== "object") return null; + const response = (error as { response?: unknown }).response; + if (!response || typeof response !== "object") return null; + const status = (response as { status?: unknown }).status; + const data = (response as { data?: unknown }).data; + if (typeof status !== "number") return null; + if (!data || typeof data !== "object") return null; + const sentinel = (data as { error?: unknown }).error; + if (typeof sentinel !== "string") return null; + + if (status === 402 && sentinel === "FEATURE_DEGRADED") { + return "FEATURE_DEGRADED"; + } + if (status === 402 && sentinel === "PAYG_LIMIT_REACHED") { + return "PAYG_LIMIT_REACHED"; + } + if (status === 401 && sentinel === "SIGNUP_REQUIRED") { + return "SIGNUP_REQUIRED"; + } + return null; +} + +/** Extract {@code data.category} (a string) from an axios error, or {@code null}. */ +export function extractSignupCategory(error: unknown): string | null { + if (!error || typeof error !== "object") return null; + const response = (error as { response?: unknown }).response; + if (!response || typeof response !== "object") return null; + const data = (response as { data?: unknown }).data; + if (!data || typeof data !== "object") return null; + const category = (data as { category?: unknown }).category; + return typeof category === "string" ? category : null; +} + +/** + * Extract {@code data.subscribed} (a boolean) from an axios error. Returns + * {@code null} when absent so the caller can apply a default. A subscribed + * team that hits a limit is over its spending cap; an un-subscribed one has + * spent its free allowance. + */ +export function extractSubscribed(error: unknown): boolean | null { + if (!error || typeof error !== "object") return null; + const response = (error as { response?: unknown }).response; + if (!response || typeof response !== "object") return null; + const data = (response as { data?: unknown }).data; + if (!data || typeof data !== "object") return null; + const subscribed = (data as { subscribed?: unknown }).subscribed; + return typeof subscribed === "boolean" ? subscribed : null; +} + +/** + * Surface the appropriate UI for a classified PAYG error. + * + *
    + *
  • {@code FEATURE_DEGRADED} / {@code PAYG_LIMIT_REACHED} — pop the + * usage-limit modal (spend-cap when subscribed, free-limit otherwise) and + * show no toast. Defaults to the free-limit modal if {@code subscribed} + * is absent (most accounts at launch are free tier).
  • + *
  • {@code SIGNUP_REQUIRED} — dispatch {@code payg:signupRequired} so the + * signup-bootstrap listener opens its modal.
  • + *
+ * + * Safe to call multiple times — the modal hosts dedupe by their own open state. + * Suppress-respecting in spirit: these are user-facing gates, not transient + * error toasts, so we surface the modal even when the caller passed + * {@code suppressErrorToast} (that flag was for the generic error toast we are + * replacing with something more actionable). + */ +export function handlePaygError(kind: PaygErrorKind, error: unknown): void { + if (kind === "FEATURE_DEGRADED" || kind === "PAYG_LIMIT_REACHED") { + if (extractSubscribed(error) === true) { + openSpendCapModal(); + } else { + openFreeLimitModal(); + } + return; + } + + if (kind === "SIGNUP_REQUIRED") { + const category = extractSignupCategory(error); + try { + window.dispatchEvent( + new CustomEvent("payg:signupRequired", { + detail: { category }, + }), + ); + } catch { + // SSR / test environments without a real window — no-op. + } + return; + } +} diff --git a/frontend/editor/src/cloud/tsconfig.json b/frontend/editor/src/cloud/tsconfig.json new file mode 100644 index 0000000000..bddc25d1bf --- /dev/null +++ b/frontend/editor/src/cloud/tsconfig.json @@ -0,0 +1,21 @@ +{ + "extends": "../../tsconfig.json", + "compilerOptions": { + "baseUrl": "../../", + "paths": { + "@app/*": ["src/cloud/*", "src/proprietary/*", "src/core/*"], + "@cloud/*": ["src/cloud/*"], + "@proprietary/*": ["src/proprietary/*"], + "@core/*": ["src/core/*"], + "@shared/*": ["../shared/*"] + } + }, + "include": [ + "../global.d.ts", + "../*.js", + "../*.ts", + "../*.tsx", + "../core/setupTests.ts", + "." + ] +} diff --git a/frontend/editor/src/core/App.tsx b/frontend/editor/src/core/App.tsx index c417ef1f65..566e861520 100644 --- a/frontend/editor/src/core/App.tsx +++ b/frontend/editor/src/core/App.tsx @@ -3,7 +3,7 @@ import { Routes, Route } from "react-router-dom"; import { AppProviders } from "@app/components/AppProviders"; import { AppLayout } from "@app/components/AppLayout"; import { LoadingFallback } from "@app/components/shared/LoadingFallback"; -import { RainbowThemeProvider } from "@app/components/shared/RainbowThemeProvider"; +import { ThemeProvider } from "@app/components/shared/ThemeProvider"; import { PreferencesProvider } from "@app/contexts/PreferencesContext"; import HomePage from "@app/pages/HomePage"; import MobileScannerPage from "@app/pages/MobileScannerPage"; @@ -17,11 +17,12 @@ import "@app/styles/index.css"; // Import file ID debugging helpers (development only) import "@app/utils/fileIdSafety"; -// Minimal providers for mobile scanner - no API calls, no authentication -function MobileScannerProviders({ children }: { children: React.ReactNode }) { +// Minimal providers for the public, no-auth mobile-scanner page - no API +// calls, no authentication +function PublicRouteProviders({ children }: { children: React.ReactNode }) { return ( - {children} + {children} ); } @@ -34,9 +35,9 @@ export default function App() { + - + } /> diff --git a/frontend/editor/src/core/auth/UseSession.tsx b/frontend/editor/src/core/auth/UseSession.tsx index 8f9a44f52e..ee3314606c 100644 --- a/frontend/editor/src/core/auth/UseSession.tsx +++ b/frontend/editor/src/core/auth/UseSession.tsx @@ -14,6 +14,14 @@ export interface AuthContextType { * should treat the resulting string as opaque display text. */ displayName: string | null; + /** + * Whether the current session is an anonymous / guest one. Each layer + * derives this from its own native user shape (Supabase `is_anonymous` in + * SaaS, the Spring anonymous flag in proprietary). Always `false` in core + * OSS, which has no auth context. Consumers use it to gate account-only + * actions (cloud folders, MCP) without reaching into a layer-specific user. + */ + isAnonymous: boolean; loading: boolean; error: Error | null; signOut: () => Promise; @@ -29,6 +37,7 @@ export function useAuth(): AuthContextType { session: null, user: null, displayName: null, + isAnonymous: false, loading: false, error: null, signOut: async () => {}, diff --git a/frontend/editor/src/core/components/AppProviders.tsx b/frontend/editor/src/core/components/AppProviders.tsx index 6c00109454..f8855e2ae9 100644 --- a/frontend/editor/src/core/components/AppProviders.tsx +++ b/frontend/editor/src/core/components/AppProviders.tsx @@ -1,5 +1,5 @@ import { ReactNode, useEffect } from "react"; -import { RainbowThemeProvider } from "@app/components/shared/RainbowThemeProvider"; +import { ThemeProvider } from "@app/components/shared/ThemeProvider"; import { FileContextProvider } from "@app/contexts/FileContext"; import { NavigationProvider } from "@app/contexts/NavigationContext"; import { ToolRegistryProvider } from "@app/contexts/ToolRegistryProvider"; @@ -33,6 +33,7 @@ import AppConfigLoader from "@app/components/shared/AppConfigLoader"; import { UpdateStartupPopup } from "@app/components/shared/UpdateStartupPopup"; import { RedactionProvider } from "@app/contexts/RedactionContext"; import { FormFillProvider } from "@app/tools/formFill/FormFillContext"; +import { FolderFileContextProvider } from "@app/contexts/FolderFileContext"; import { FolderProvider } from "@app/contexts/FolderContext"; // Component to initialize scarf tracking (must be inside AppConfigProvider) @@ -113,7 +114,7 @@ export function AppProviders({ }: AppProvidersProps) { return ( - + - {children} + + {children} + @@ -169,7 +172,7 @@ export function AppProviders({ - + ); } diff --git a/frontend/editor/src/core/components/agents/AgentsPanel.tsx b/frontend/editor/src/core/components/agents/AgentsPanel.tsx deleted file mode 100644 index 282b4639b3..0000000000 --- a/frontend/editor/src/core/components/agents/AgentsPanel.tsx +++ /dev/null @@ -1,53 +0,0 @@ -/** - * Core stubs for the right-rail Agents UI. - * - * The real implementations live in {@code proprietary/components/agents/AgentsPanel.tsx} - * and shadow these stubs via the {@code @app/*} alias cascade when the proprietary - * build is active. Core builds render nothing, so the right rail collapses to the - * tool list unchanged. - */ - -/** Whether the right rail should reserve space for agents UI. False in core. */ -export function useAgentsEnabled(): boolean { - return false; -} - -/** - * Whether the agent chat overlay is currently open. Core builds have no chat, - * so this always returns false. Proprietary builds bridge to the ChatContext. - * Used by {@code RightSidebar} so the fullscreen tool picker can yield to the - * chat overlay just like it yields to a selected tool. - */ -export function useAgentChatOpen(): boolean { - return false; -} - -/** Inline "Agents" section rendered above the tool list in {@code ToolPicker}. */ -export function AgentsSection() { - return null; -} - -/** - * Icon-only agent button rendered in the collapsed (minimised) right rail. - * Returns null in core; proprietary renders the Stirling agent shortcut. - */ -export function AgentsCollapsedButton(_props: { onExpand: () => void }) { - return null; -} - -/** - * Full-rail chat overlay rendered inside {@code ToolPanel}. Covers the panel - * (including the search bar) when an agent conversation is active. - */ -export function AgentsChatOverlay() { - return null; -} - -/** - * Agents card rendered inside the fullscreen tool picker. Matches the visual - * language of the fullscreen category cards (gradient border, title, items). - * Returns null in core; proprietary renders the Stirling agent. - */ -export function AgentsFullscreenSection() { - return null; -} diff --git a/frontend/editor/src/core/components/annotation/shared/ColorControl.tsx b/frontend/editor/src/core/components/annotation/shared/ColorControl.tsx index ff2a3ff458..ca41f0f6a2 100644 --- a/frontend/editor/src/core/components/annotation/shared/ColorControl.tsx +++ b/frontend/editor/src/core/components/annotation/shared/ColorControl.tsx @@ -8,6 +8,7 @@ import { Group, } from "@mantine/core"; import { useState, useCallback, useEffect } from "react"; +import { useTranslation } from "react-i18next"; import ColorizeIcon from "@mui/icons-material/Colorize"; // safari and firefox do not support the eye dropper API, only edge, chrome and opera do. @@ -33,6 +34,7 @@ export function ColorControl({ label, disabled = false, }: ColorControlProps) { + const { t } = useTranslation(); const [opened, setOpened] = useState(false); // Buffer the colour locally so the picker stays responsive during drag. // Only propagate to the parent (which triggers expensive annotation updates) @@ -111,7 +113,9 @@ export function ColorControl({ /> {supportsEyeDropper && ( - + = ({ onImageChange, disabled = false, }) => { + const { t } = useTranslation(); const [, setImageData] = useState(null); const handleImageUpload = async (file: File | null) => { @@ -57,9 +59,12 @@ export const ImageTool: React.FC = ({ diff --git a/frontend/editor/src/core/components/chat/ChatContext.tsx b/frontend/editor/src/core/components/chat/ChatContext.tsx new file mode 100644 index 0000000000..cd8f080722 --- /dev/null +++ b/frontend/editor/src/core/components/chat/ChatContext.tsx @@ -0,0 +1,17 @@ +/** + * Core stub for the chat context. + * The real implementation lives in proprietary/components/chat/ChatContext.tsx + * and shadows this via the @app/* alias cascade in proprietary builds. + */ + +export function useChat() { + return { + messages: [] as never[], + isLoading: false, + progress: null, + progressLog: [] as never[], + sendMessage: async (_content: string) => {}, + cancelMessage: () => {}, + clearChat: () => {}, + }; +} diff --git a/frontend/editor/src/core/components/chat/ChatFAB.tsx b/frontend/editor/src/core/components/chat/ChatFAB.tsx new file mode 100644 index 0000000000..5488213693 --- /dev/null +++ b/frontend/editor/src/core/components/chat/ChatFAB.tsx @@ -0,0 +1,8 @@ +/** + * Core stub for the floating action button chat widget. + * The real implementation lives in proprietary/components/chat/ChatFAB.tsx + * and shadows this via the @app/* alias cascade in proprietary builds. + */ +export function ChatFAB() { + return null; +} diff --git a/frontend/editor/src/core/components/chat/ChatPanel.tsx b/frontend/editor/src/core/components/chat/ChatPanel.tsx new file mode 100644 index 0000000000..0d85c23fcc --- /dev/null +++ b/frontend/editor/src/core/components/chat/ChatPanel.tsx @@ -0,0 +1,14 @@ +/** + * Core stub for the chat panel. + * The real implementation lives in proprietary/components/chat/ChatPanel.tsx + * and shadows this via the @app/* alias cascade in proprietary builds. + */ + +export interface ChatPanelProps { + onBack: () => void; + backLabel: string; +} + +export function ChatPanel(_props: ChatPanelProps) { + return null; +} diff --git a/frontend/editor/src/core/components/fileEditor/AddFileCard.tsx b/frontend/editor/src/core/components/fileEditor/AddFileCard.tsx index 5fc62680b3..00c9e4b396 100644 --- a/frontend/editor/src/core/components/fileEditor/AddFileCard.tsx +++ b/frontend/editor/src/core/components/fileEditor/AddFileCard.tsx @@ -104,44 +104,44 @@ const AddFileCard = ({ }} onMouseLeave={() => setIsUploadHover(false)} > - + + )}
); }; diff --git a/frontend/editor/src/core/components/fileManager/EmptyFilesState.tsx b/frontend/editor/src/core/components/fileManager/EmptyFilesState.tsx index 777b80085c..422dd8306e 100644 --- a/frontend/editor/src/core/components/fileManager/EmptyFilesState.tsx +++ b/frontend/editor/src/core/components/fileManager/EmptyFilesState.tsx @@ -77,7 +77,7 @@ const EmptyFilesState: React.FC = () => { onMouseLeave={() => setIsUploadHover(false)} > + + + + + ); +} diff --git a/frontend/editor/src/core/components/filesPage/FileDetailsPanel.tsx b/frontend/editor/src/core/components/filesPage/FileDetailsPanel.tsx index 2608e4b30a..aaab5d3b08 100644 --- a/frontend/editor/src/core/components/filesPage/FileDetailsPanel.tsx +++ b/frontend/editor/src/core/components/filesPage/FileDetailsPanel.tsx @@ -1,21 +1,18 @@ import React, { useEffect, useMemo, useState } from "react"; import { useTranslation } from "react-i18next"; -import { ActionIcon, Badge, Button, Menu, Tooltip } from "@mantine/core"; +import { ActionIcon, Badge, Button, Tooltip } from "@mantine/core"; import CloseIcon from "@mui/icons-material/Close"; import OpenInNewIcon from "@mui/icons-material/OpenInNew"; -import VisibilityIcon from "@mui/icons-material/Visibility"; import DriveFileMoveIcon from "@mui/icons-material/DriveFileMove"; import DeleteIcon from "@mui/icons-material/Delete"; import DownloadIcon from "@mui/icons-material/Download"; import PictureAsPdfIcon from "@mui/icons-material/PictureAsPdf"; import HistoryIcon from "@mui/icons-material/History"; -import MoreVertIcon from "@mui/icons-material/MoreVert"; import KeyboardArrowDownIcon from "@mui/icons-material/KeyboardArrowDown"; import LinkIcon from "@mui/icons-material/Link"; import CloudUploadIcon from "@mui/icons-material/CloudUpload"; -import { FileId, ToolOperation } from "@app/types/file"; -import { ToolId } from "@app/types/toolId"; +import { FileId } from "@app/types/file"; import { FolderRecord } from "@app/types/folder"; import { StirlingFileStub } from "@app/types/fileContext"; import { formatFileSize, getFileDate } from "@app/utils/fileUtils"; @@ -27,6 +24,10 @@ import ToolChain from "@app/components/shared/ToolChain"; import ShareManagementModal from "@app/components/shared/ShareManagementModal"; import { useSharingEnabled } from "@app/hooks/useSharingEnabled"; import { fileStorage } from "@app/services/fileStorage"; +import { + VersionTimeline, + DetailField, +} from "@app/components/filesPage/VersionTimeline"; interface FileDetailsPanelProps { selectedFileIds: FileId[]; @@ -34,13 +35,16 @@ interface FileDetailsPanelProps { currentFolder: FolderRecord | null; onClose: () => void; onAddToWorkspace: (fileIds: FileId[]) => void; - onQuickView: (fileId: FileId) => void; onMove: (fileIds: FileId[]) => void; onRemove: (fileIds: FileId[]) => void; /** Save to server; only shown when at least one selected file is local-only. */ onSaveToServer?: (files: StirlingFileStub[]) => void; /** When set, Save to server renders disabled with this tooltip (storage off). */ saveToServerDisabledReason?: string | null; + /** On small screens, show a compact "Version journey" button instead of the + * full inline timeline (which opens onOpenVersionHistory). */ + compactVersions?: boolean; + onOpenVersionHistory?: () => void; } export function FileDetailsPanel({ @@ -49,11 +53,12 @@ export function FileDetailsPanel({ currentFolder, onClose, onAddToWorkspace, - onQuickView, onMove, onRemove, onSaveToServer, saveToServerDisabledReason, + compactVersions = false, + onOpenVersionHistory, }: FileDetailsPanelProps) { const { t } = useTranslation(); const { sharingEnabled } = useSharingEnabled(); @@ -68,6 +73,9 @@ export function FileDetailsPanel({ // Hooks must run before any early return. const [downloading, setDownloading] = useState(false); const [shareModalOpen, setShareModalOpen] = useState(false); + // Metadata (size/type/dates) is collapsed by default so the panel stays + // short and the action buttons keep their pinned footer in view. + const [fieldsOpen, setFieldsOpen] = useState(false); // Version chain for the selected file; empty for v1 or multi-select. const [versionChain, setVersionChain] = useState([]); const singleFileForChain = files.length === 1 ? files[0] : null; @@ -144,7 +152,11 @@ export function FileDetailsPanel({
{single ? ( <> -
+
{single.thumbnailUrl ? ( ) : ( @@ -174,36 +186,52 @@ export function FileDetailsPanel({ )}
-
- setFieldsOpen((o) => !o)} + aria-expanded={fieldsOpen} + > + {t("filesPage.fileInfo", "File info")} + - - - - -
+ + {fieldsOpen && ( +
+ + + + + +
+ )} {single.toolHistory && single.toolHistory.length > 0 && (
@@ -224,15 +252,27 @@ export function FileDetailsPanel({ shows WHICH tool was added at each step (the delta from the prior version) so the user can read the journey top-to-bottom. Long chains (> 6) collapse the middle. */} - {versionChain.length > 1 && ( - - )} + {versionChain.length > 1 && + (compactVersions && onOpenVersionHistory ? ( + + ) : ( + + ))} ) : (
@@ -246,120 +286,107 @@ export function FileDetailsPanel({ />
)} +
-
+ - {single && ( - - )} - - {/* Share is single-file only. When sharing is disabled in + {files.length === 1 + ? t("filesPage.addToWorkspace", "Add to workspace") + : t("filesPage.addToWorkspaceCount", "Add {{count}} to workspace", { + count: files.length, + })} + + + {/* Share is single-file only. When sharing is disabled in server config (storage.sharing.enabled=false) we still render the button - disabled with an explanatory tooltip - so users discover the feature exists and know how to enable it, rather than wondering why "share" is missing from the action stack on their build. */} - {single && ( - - - - )} - - {/* Save to server; shown when any selected file is local-only. When + + + )} + + {/* Save to server; shown when any selected file is local-only. When storage is off it stays visible but disabled with a tooltip (same treatment as Manage sharing above). */} - {onSaveToServer && localOnlyFiles.length > 0 && ( - - - - )} - -
+ + + )} +
{/* Single panel-level mount; gated on sharingEnabled. */} {single && sharingEnabled && ( @@ -372,309 +399,3 @@ export function FileDetailsPanel({ ); } - -function DetailField({ label, value }: { label: string; value: string }) { - return ( -
- {label} - {value} -
- ); -} - -/** Tool that produced `version` from `prior`; null for v1. */ -function deltaToolFor( - version: StirlingFileStub, - prior: StirlingFileStub | null, -): ToolOperation | null { - if (!prior) return null; - const priorLen = prior.toolHistory?.length ?? 0; - const curr = version.toolHistory ?? []; - return curr[priorLen] ?? null; -} - -interface VersionTimelineProps { - /** Chain sorted oldest-first. */ - chain: StirlingFileStub[]; - /** Currently selected version. */ - currentId: FileId; - onQuickView: (fileId: FileId) => void; - onAddToWorkspace: (fileIds: FileId[]) => void; - onRemove: (fileIds: FileId[]) => void; -} - -/** Version timeline with per-row tool deltas and collapse-when-long. */ -function VersionTimeline({ - chain, - currentId, - onQuickView, - onAddToWorkspace, - onRemove, -}: VersionTimelineProps) { - const { t } = useTranslation(); - const [expandedIds, setExpandedIds] = useState>(new Set()); - const [showAllCollapsed, setShowAllCollapsed] = useState(false); - - // Newest-first ordering. - const ordered = useMemo( - () => - [...chain].sort( - (a, b) => (b.versionNumber ?? 1) - (a.versionNumber ?? 1), - ), - [chain], - ); - - // Index by versionNumber for prior-version lookup. - const byVersionNumber = useMemo(() => { - const map = new Map(); - for (const v of chain) { - map.set(v.versionNumber ?? 1, v); - } - return map; - }, [chain]); - - // Collapse middle when long: 3 newest + ellipsis + 2 oldest. - const COLLAPSE_THRESHOLD = 6; - const collapsible = ordered.length > COLLAPSE_THRESHOLD; - type Row = - | { kind: "version"; version: StirlingFileStub } - | { - kind: "ellipsis"; - hidden: number; - }; - const rows: Row[] = useMemo(() => { - if (!collapsible || showAllCollapsed) { - return ordered.map((v) => ({ kind: "version", version: v }) as Row); - } - const head = ordered - .slice(0, 3) - .map((v) => ({ kind: "version", version: v }) as Row); - const tail = ordered - .slice(-2) - .map((v) => ({ kind: "version", version: v }) as Row); - const hidden = ordered.length - 5; - return [...head, { kind: "ellipsis", hidden }, ...tail]; - }, [collapsible, showAllCollapsed, ordered]); - - const toggleExpand = (id: FileId) => { - setExpandedIds((prev) => { - const next = new Set(prev); - if (next.has(id)) next.delete(id); - else next.add(id); - return next; - }); - }; - - return ( -
-
- - {t("filesPage.field.versionHistory", "Version journey")} - - {t("filesPage.versionsCount", "{{count}} versions", { - count: ordered.length, - })} - -
-
    - {rows.map((row, idx) => { - const isLast = idx === rows.length - 1; - if (row.kind === "ellipsis") { - return ( -
  1. -
    - - {!isLast && ( - - )} -
    - -
  2. - ); - } - const v = row.version; - const isActive = v.id === currentId; - const isExpanded = expandedIds.has(v.id); - const prior = byVersionNumber.get((v.versionNumber ?? 1) - 1) ?? null; - const delta = deltaToolFor(v, prior); - return ( -
  3. -
    - - {!isLast && ( - - )} -
    -
    - -
    - {formatFileSize(v.size)} - {v.lastModified ? ( - <> - · - - {getFileDate({ lastModified: v.lastModified })} - - - ) : null} - {!isActive && ( - <> - - - - e.stopPropagation()} - > - - - - - } - onClick={() => onQuickView(v.id)} - > - {t("filesPage.viewVersion", "View this version")} - - } - onClick={() => onAddToWorkspace([v.id])} - > - {t( - "filesPage.openVersionInWorkspace", - "Open in workspace", - )} - - } - onClick={() => { - void downloadFileFromStorage(v); - }} - > - {t( - "filesPage.downloadVersion", - "Download this version", - )} - - - } - onClick={() => onRemove([v.id])} - > - {t( - "filesPage.removeVersion", - "Remove this version", - )} - - - - - )} -
    - {isExpanded && ( - // Filename + full cumulative tool chain. -
    - - {v.toolHistory && v.toolHistory.length > 0 && ( -
    - - {t( - "filesPage.field.toolHistoryAtVersion", - "Cumulative tool chain", - )} - - -
    - )} -
    - )} -
    -
  4. - ); - })} -
- {collapsible && showAllCollapsed && ( - - )} -
- ); -} - -/** Translated tool name via `home.{toolId}.title`. */ -function ToolLabel({ toolId }: { toolId: ToolId }) { - const { t } = useTranslation(); - return {t(`home.${toolId}.title`, toolId)}; -} diff --git a/frontend/editor/src/core/components/filesPage/FileGrid.tsx b/frontend/editor/src/core/components/filesPage/FileGrid.tsx index 0e3c44ddcd..8600ed443d 100644 --- a/frontend/editor/src/core/components/filesPage/FileGrid.tsx +++ b/frontend/editor/src/core/components/filesPage/FileGrid.tsx @@ -2,13 +2,14 @@ import React, { useCallback, useMemo, useRef } from "react"; import { useTranslation } from "react-i18next"; import { ActionIcon, Button, Checkbox, Menu, Tooltip } from "@mantine/core"; import MoreVertIcon from "@mui/icons-material/MoreVert"; +import ShieldOutlinedIcon from "@mui/icons-material/ShieldOutlined"; import FolderIcon from "@mui/icons-material/Folder"; import PictureAsPdfIcon from "@mui/icons-material/PictureAsPdf"; import InsertDriveFileIcon from "@mui/icons-material/InsertDriveFile"; import DriveFileMoveIcon from "@mui/icons-material/DriveFileMove"; import DeleteIcon from "@mui/icons-material/Delete"; +import HistoryIcon from "@mui/icons-material/History"; import OpenInNewIcon from "@mui/icons-material/OpenInNew"; -import VisibilityIcon from "@mui/icons-material/Visibility"; import DriveFileRenameOutlineIcon from "@mui/icons-material/DriveFileRenameOutline"; import CloudUploadIcon from "@mui/icons-material/CloudUpload"; import UploadFileIcon from "@mui/icons-material/UploadFile"; @@ -18,6 +19,7 @@ import SearchIcon from "@mui/icons-material/Search"; import { FileId } from "@app/types/file"; import { FolderId, FolderRecord, ROOT_FOLDER_ID } from "@app/types/folder"; import { useFolders } from "@app/contexts/FolderContext"; +import { usePolicyFileBadges } from "@app/hooks/usePolicyFileBadges"; import { StirlingFileStub } from "@app/types/fileContext"; import { formatFileSize, getFileDate } from "@app/utils/fileUtils"; import { @@ -59,8 +61,6 @@ interface FileGridProps { onOpenFolder: (id: FolderId) => void; /** "Add to workspace". */ onOpenFile: (file: StirlingFileStub) => void; - /** "Quick view". */ - onQuickView: (file: StirlingFileStub) => void; onMoveFiles: ( fileIds: FileId[], targetFolderId: FolderId | null, @@ -79,6 +79,8 @@ interface FileGridProps { onPromptMoveFiles: (fileIds: FileId[]) => void; /** Per-file Save to server; hidden when file already has remoteStorageId. */ onSaveToServer?: (file: StirlingFileStub) => void; + /** Open the version-history modal for a file (only when it has >1 version). */ + onVersionHistory?: (file: StirlingFileStub) => void; /** When set, the Save to server item renders disabled with this tooltip. */ saveToServerDisabledReason?: string | null; /** When supplied the list-view column headers become sortable. */ @@ -357,7 +359,6 @@ function GridView({ onSelectFile, onOpenFolder, onOpenFile, - onQuickView, onMoveFiles, onMoveFolder, onRenameFolder, @@ -366,6 +367,7 @@ function GridView({ onRemoveFiles, onPromptMoveFiles, onSaveToServer, + onVersionHistory, saveToServerDisabledReason, }: FileGridProps) { return ( @@ -408,7 +410,6 @@ function GridView({ onSelectFile(entry.file!.id, e.shiftKey, e.metaKey || e.ctrlKey) } onDoubleClick={() => onOpenFile(entry.file!)} - onQuickView={() => onQuickView(entry.file!)} onRemove={() => onRemoveFiles([entry.file!.id])} onMove={() => { const target = selectedFileIds.has(entry.file!.id) @@ -419,6 +420,11 @@ function GridView({ onSaveToServer={ onSaveToServer ? () => onSaveToServer(entry.file!) : undefined } + onVersionHistory={ + onVersionHistory + ? () => onVersionHistory(entry.file!) + : undefined + } saveToServerDisabledReason={saveToServerDisabledReason} /> ); @@ -602,6 +608,31 @@ function FolderCard({ ); } +/** Shield badges for the policies that have run on a file. */ +function PolicyBadges({ fileId }: { fileId: string }) { + const badges = usePolicyFileBadges().get(fileId) ?? []; + if (badges.length === 0) return null; + return ( + + {badges.slice(0, 3).map((policy) => ( + + + + + + ))} + + ); +} + interface FileCardProps { file: StirlingFileStub; isSelected: boolean; @@ -613,11 +644,12 @@ interface FileCardProps { multiSelectActive: boolean; onClick: (e: React.MouseEvent) => void; onDoubleClick: () => void; - onQuickView: () => void; onRemove: () => void; onMove: () => void; /** Kebab Save to server; only fires when file is local-only. */ onSaveToServer?: () => void; + /** Open the version-history modal; shown only when file has >1 version. */ + onVersionHistory?: () => void; /** When set, the kebab Save to server is disabled with this tooltip. */ saveToServerDisabledReason?: string | null; } @@ -631,10 +663,10 @@ function FileCard({ multiSelectActive, onClick, onDoubleClick, - onQuickView, onRemove, onMove, onSaveToServer, + onVersionHistory, saveToServerDisabledReason, }: FileCardProps) { const { t } = useTranslation(); @@ -761,6 +793,7 @@ function FileCard({ {fileSize} · {fileDate} +
@@ -773,6 +806,7 @@ function FileCard({ size="sm" onClick={(e) => e.stopPropagation()} aria-label={t("filesPage.fileMenu", "File actions")} + data-testid="file-card-actions" > @@ -787,15 +821,6 @@ function FileCard({ > {t("filesPage.addToWorkspace", "Add to workspace")} - } - onClick={(e) => { - e.stopPropagation(); - onQuickView(); - }} - > - {t("filesPage.quickView", "Quick view")} - } @@ -803,6 +828,7 @@ function FileCard({ e.stopPropagation(); onMove(); }} + data-testid="file-menu-move-to" > {t("filesPage.moveTo", "Move to…")} @@ -834,6 +860,17 @@ function FileCard({ )} + {onVersionHistory && (file.versionNumber ?? 1) > 1 && ( + } + onClick={(e) => { + e.stopPropagation(); + onVersionHistory(); + }} + > + {t("filesPage.versionHistory", "Version history")} + + )} onOpenFile(entry.file!)} - onQuickView={() => onQuickView(entry.file!)} onRemove={() => onRemoveFiles([entry.file!.id])} onMove={() => { const target = selectedFileIds.has(entry.file!.id) @@ -1000,6 +1036,11 @@ function ListView({ onSaveToServer={ onSaveToServer ? () => onSaveToServer(entry.file!) : undefined } + onVersionHistory={ + onVersionHistory + ? () => onVersionHistory(entry.file!) + : undefined + } saveToServerDisabledReason={saveToServerDisabledReason} /> ); @@ -1198,11 +1239,12 @@ interface FileRowProps { multiSelectActive: boolean; onClick: (e: React.MouseEvent) => void; onOpen: () => void; - onQuickView: () => void; onRemove: () => void; onMove: () => void; /** Kebab Save to server; only fires when file is local-only. */ onSaveToServer?: () => void; + /** Open the version-history modal; shown only when file has >1 version. */ + onVersionHistory?: () => void; /** When set, the kebab Save to server is disabled with this tooltip. */ saveToServerDisabledReason?: string | null; } @@ -1216,10 +1258,10 @@ function FileRow({ multiSelectActive, onClick, onOpen, - onQuickView, onRemove, onMove, onSaveToServer, + onVersionHistory, saveToServerDisabledReason, }: FileRowProps) { const { t } = useTranslation(); @@ -1343,6 +1385,7 @@ function FileRow({ )} + {isInWorkspace && ( @@ -1361,6 +1404,7 @@ function FileRow({ size="sm" onClick={(e) => e.stopPropagation()} aria-label={t("filesPage.fileMenu", "File actions")} + data-testid="file-card-actions" > @@ -1375,15 +1419,6 @@ function FileRow({ > {t("filesPage.addToWorkspace", "Add to workspace")} - } - onClick={(e) => { - e.stopPropagation(); - onQuickView(); - }} - > - {t("filesPage.quickView", "Quick view")} - } @@ -1422,6 +1457,17 @@ function FileRow({ )} + {onVersionHistory && (file.versionNumber ?? 1) > 1 && ( + } + onClick={(e) => { + e.stopPropagation(); + onVersionHistory(); + }} + > + {t("filesPage.versionHistory", "Version history")} + + )} (null); + // Version-history modal target (opened from the card kebab). + const [versionHistoryFile, setVersionHistoryFile] = + useState(null); const folders = useFolders(); const { actions: fileActions } = useFileActions(); const { fileIds: activeWorkspaceFileIds } = useAllFiles(); @@ -108,17 +111,27 @@ export default function FileManagerView() { const isMobile = useIsMobile(); const isMobileUploadAvailable = Boolean(appConfig?.enableMobileScanner) && !isMobile; + // Guests (anonymous sessions) have no server-side storage, so every cloud + // action is account-only. Rather than let the click fire a guaranteed 401 + // (which surfaced as an error toast), we disable the control and explain why + // on hover - the same affordance the storage-disabled / wrong-tab gates use. + const { isAnonymous } = useAuth(); + const signInRequiredReason = isAnonymous + ? t("filesPage.signInRequired", "Sign in to use cloud storage.") + : null; // Server storage gate; mirrors ConfigController's storageEnabled // (enableLogin && storage.isEnabled). When off, Save-to-server stays // visible but disabled with an explanatory tooltip (discoverability beats // hiding - mirrors the New folder / Manage sharing gates in this view). const uploadEnabled = appConfig?.storageEnabled === true; - const saveToServerDisabledReason: string | null = uploadEnabled - ? null - : t( - "filesPage.saveToServerDisabledHint", - "Saving to the server isn't enabled on this server. Ask your admin to enable it.", - ); + const saveToServerDisabledReason: string | null = + signInRequiredReason ?? + (uploadEnabled + ? null + : t( + "filesPage.saveToServerDisabledHint", + "Saving to the server isn't enabled on this server. Ask your admin to enable it.", + )); const [mobileUploadModalOpen, setMobileUploadModalOpen] = useState(false); const { actions: navActions } = useNavigationActions(); const { requestNavigation } = useNavigationGuard(); @@ -156,6 +169,10 @@ export default function FileManagerView() { moveFilesTo, moveFolderTo, removeFiles, + deleteDialogFileIds, + deleteDialogOpen, + closeDeleteDialog, + confirmRemoveFiles, promptDeleteFolder, deleteFolder, deleteFolderDialog, @@ -163,6 +180,15 @@ export default function FileManagerView() { setFolderAppearance, } = filesPage; + // Resolve queued delete ids into stubs for the DeleteFilesDialog. + const deleteDialogFiles = useMemo( + () => + deleteDialogFileIds + .map((id) => fileMap.get(id)) + .filter((s): s is StirlingFileStub => Boolean(s)), + [deleteDialogFileIds, fileMap], + ); + const setCurrentFolderId = folders.setCurrentFolderId; const foldersById = folders.foldersById; const currentFolderId = folders.currentFolderId; @@ -192,10 +218,11 @@ export default function FileManagerView() { // Push folder selection into the URL while still on /files. useEffect(() => { - if (!window.location.pathname.startsWith("/files")) return; + const stripped = stripBasePath(window.location.pathname); + if (!stripped.startsWith("/files")) return; const target = currentFolderId === null ? "/files" : `/files/${currentFolderId}`; - if (window.location.pathname !== target) { + if (stripped !== target) { navigate(target, { replace: true }); } }, [currentFolderId, navigate]); @@ -527,30 +554,16 @@ export default function FileManagerView() { [handleNativeUpload], ); - // ─── add to workspace vs quick view ───────────────────────────────────── - // addToWorkspace: commit; no back affordance. - // quickView: peek; "Back to My Files" pill in WorkbenchBar. + // ─── add to workspace ─────────────────────────────────────────────────── const openFilesInWorkbench = useCallback( - async (fileIds: FileId[], options: { trackReturn: boolean }) => { + async (fileIds: FileId[]) => { const stubs = fileIds .map((id) => fileMap.get(id)) .filter((s): s is StirlingFileStub => Boolean(s)); if (stubs.length === 0) return; const proceed = async () => { - if (options.trackReturn) { - const returnRoute = - currentFolderId === null ? "/files" : `/files/${currentFolderId}`; - const folderRecord = currentFolderId - ? (foldersById.get(currentFolderId) ?? null) - : null; - const returnLabel = folderRecord - ? folderRecord.name - : t("filesPage.myFiles", "My Files"); - setFilesPageReturnRoute(returnRoute, returnLabel); - } else { - clearFilesPageReturnRoute(); - } + clearFilesPageReturnRoute(); // Server-only stubs have no bytes in IDB; download + ingest first. const materialized = await materializeServerStubs(stubs, { @@ -588,20 +601,12 @@ export default function FileManagerView() { navActions, navigate, requestNavigation, - currentFolderId, - foldersById, - t, + clearFilesPageReturnRoute, ], ); const handleAddToWorkspace = useCallback( - (fileIds: FileId[]) => - openFilesInWorkbench(fileIds, { trackReturn: false }), - [openFilesInWorkbench], - ); - - const handleQuickView = useCallback( - (fileId: FileId) => openFilesInWorkbench([fileId], { trackReturn: true }), + (fileIds: FileId[]) => openFilesInWorkbench(fileIds), [openFilesInWorkbench], ); @@ -817,6 +822,11 @@ export default function FileManagerView() { // null = New folder actionable; string = disabled tooltip reason. const newFolderDisabledReason: string | null = useMemo(() => { + // Guests can't use cloud folders at all - say so before any tab/storage + // hint, since switching tabs wouldn't help them. + if (signInRequiredReason) { + return signInRequiredReason; + } if (currentTab === "local") { return t( "filesPage.localFoldersUnavailable", @@ -840,7 +850,7 @@ export default function FileManagerView() { ); } return null; - }, [currentTab, folders.serverReachable, t]); + }, [signInRequiredReason, currentTab, folders.serverReachable, t]); return (
@@ -903,14 +913,17 @@ export default function FileManagerView() {
@@ -1193,19 +1205,6 @@ export default function FileManagerView() { {addLabel} - {selectedFiles.length === 1 && ( - - - - )} {/* Save to server; shown whenever local-only files are selected. When storage is off it stays visible but disabled, tooltip pointing at the admin. */} @@ -1489,7 +1488,6 @@ export default function FileManagerView() { onSetSelection={setSelectedFileIds} onOpenFolder={handleOpenFolder} onOpenFile={handleOpenFile} - onQuickView={(file) => handleQuickView(file.id)} onMoveFiles={moveFilesTo} onMoveFolder={moveFolderTo} onRenameFolder={openRenameFolderDialog} @@ -1512,6 +1510,7 @@ export default function FileManagerView() { onRemoveFiles={handleRemoveFiles} onPromptMoveFiles={promptMoveFiles} onSaveToServer={(file) => setSaveToServerTarget([file])} + onVersionHistory={(file) => setVersionHistoryFile(file)} saveToServerDisabledReason={saveToServerDisabledReason} // Center-of-grid CTAs when the empty state shows - same // handlers the corner header buttons use so behaviour @@ -1554,7 +1553,6 @@ export default function FileManagerView() { currentFolder={currentFolderRecord} onClose={() => clearSelection()} onAddToWorkspace={handleAddToWorkspace} - onQuickView={handleQuickView} onMove={promptMoveFiles} onRemove={handleRemoveFiles} onSaveToServer={(files) => setSaveToServerTarget(files)} @@ -1581,11 +1579,18 @@ export default function FileManagerView() { currentFolder={currentFolderRecord} onClose={() => setMobileDetailsOpen(false)} onAddToWorkspace={handleAddToWorkspace} - onQuickView={handleQuickView} onMove={promptMoveFiles} onRemove={handleRemoveFiles} onSaveToServer={(files) => setSaveToServerTarget(files)} saveToServerDisabledReason={saveToServerDisabledReason} + compactVersions + onOpenVersionHistory={() => { + const f = fileMap.get(selectedFiles[0]); + if (f) { + setMobileDetailsOpen(false); + setVersionHistoryFile(f); + } + }} /> )} @@ -1657,6 +1662,22 @@ export default function FileManagerView() { }} /> + {/* Cloud-aware delete; offers local/cloud/both when a file lives in both. */} + + + {/* Version journey in a modal (opened from the card kebab). */} + setVersionHistoryFile(null)} + file={versionHistoryFile} + onChanged={refresh} + /> + {/* Save-to-server modal; keyed on target so updates don't retarget. */} s.id).join(",")}`} diff --git a/frontend/editor/src/core/components/filesPage/FilesPage.css b/frontend/editor/src/core/components/filesPage/FilesPage.css index 9b3b7b206f..dcc11bd0f3 100644 --- a/frontend/editor/src/core/components/filesPage/FilesPage.css +++ b/frontend/editor/src/core/components/filesPage/FilesPage.css @@ -484,6 +484,24 @@ gap: 0.4rem; } +/* Policy activity badges (a shield per policy that has run on the file). */ +.files-page-policy-badges { + display: inline-flex; + align-items: center; + gap: 3px; + flex-shrink: 0; +} +.files-page-policy-badge { + display: inline-flex; + align-items: center; + justify-content: center; + width: 15px; + height: 15px; + border-radius: 4px; + /* `color` set inline to the policy accent; tint follows it. */ + background: color-mix(in srgb, currentColor 16%, transparent); +} + /* Parent-folder breadcrumb shown on cards/rows during recursive search so the user can tell which folder each hit lives in without navigating. */ .files-page-card-path { @@ -829,14 +847,55 @@ and bottom padding stays at 1.1rem so the content still breathes against the panel edges. */ padding: 0.3rem 1.1rem 1.1rem; + /* Grow to fill the panel and scroll internally so the action footer below + stays pinned and visible regardless of how tall the content gets. */ + flex: 1; + min-height: 0; overflow-y: auto; display: flex; flex-direction: column; gap: 0.65rem; } +/* Pinned action footer - always visible so the buttons never scroll off. */ +.files-page-details-actions { + flex-shrink: 0; + display: flex; + flex-direction: column; + gap: 0.4rem; + padding: 0.6rem 1.1rem; + border-top: 1px solid var(--border-subtle); + background: var(--bg-toolbar); +} + +/* Collapsible "File info" header (size/type/dates). */ +.files-page-details-collapse-toggle { + display: flex; + align-items: center; + justify-content: space-between; + width: 100%; + padding: 0.4rem 0.6rem; + background: var(--bg-surface); + border: 1px solid var(--border-subtle); + border-radius: 0.5rem; + cursor: pointer; + color: var(--text-secondary, var(--text-primary)); + font-size: 0.78rem; + text-transform: uppercase; + letter-spacing: 0.04em; +} + +.files-page-details-collapse-chevron { + transition: transform 0.15s ease; +} +.files-page-details-collapse-chevron.is-open { + transform: rotate(180deg); +} + .files-page-details-thumb { - aspect-ratio: 4 / 3; + /* Compact fixed height (was a tall 4:3 box that dominated the panel and + pushed the action buttons off the bottom). */ + height: 6.5rem; background: var(--bg-surface); border: 1px solid var(--border-subtle); border-radius: 0.6rem; @@ -859,6 +918,11 @@ object-fit: contain; } +/* On small screens (drawer) the preview isn't worth the vertical space. */ +.files-page-details-thumb.is-compact { + display: none; +} + .files-page-details-fieldlist { display: flex; flex-direction: column; diff --git a/frontend/editor/src/core/components/filesPage/MoveToFolderDialog.tsx b/frontend/editor/src/core/components/filesPage/MoveToFolderDialog.tsx index 18d8240792..8950f3f494 100644 --- a/frontend/editor/src/core/components/filesPage/MoveToFolderDialog.tsx +++ b/frontend/editor/src/core/components/filesPage/MoveToFolderDialog.tsx @@ -191,6 +191,7 @@ export function MoveToFolderDialog({ return creatingFolder ? ( {t( "filesPage.moveDialog.newFolderToggle", diff --git a/frontend/editor/src/core/components/filesPage/VersionHistoryModal.tsx b/frontend/editor/src/core/components/filesPage/VersionHistoryModal.tsx new file mode 100644 index 0000000000..077924155d --- /dev/null +++ b/frontend/editor/src/core/components/filesPage/VersionHistoryModal.tsx @@ -0,0 +1,125 @@ +import { useCallback, useEffect, useState } from "react"; +import { useTranslation } from "react-i18next"; +import { Center, Loader, Modal, Text } from "@mantine/core"; + +import type { FileId } from "@app/types/file"; +import type { StirlingFileStub } from "@app/types/fileContext"; +import { fileStorage } from "@app/services/fileStorage"; +import { useFileActions } from "@app/contexts/file/fileHooks"; +import { useNavigationActions } from "@app/contexts/NavigationContext"; +import { useIndexedDBRevision } from "@app/contexts/IndexedDBContext"; +import { VersionTimeline } from "@app/components/filesPage/VersionTimeline"; + +interface VersionHistoryModalProps { + opened: boolean; + onClose: () => void; + /** The file whose version journey to show (the current/leaf version). */ + file: StirlingFileStub | null; + /** Called after a destructive change so the launcher can refresh its list. */ + onChanged?: () => void; +} + +/** + * Self-contained modal that renders the same Version Journey timeline used in + * the details panel. Loads the chain itself and wires view/open/remove to the + * file context, so it can be opened from anywhere (sidebar or /files card). + */ +export function VersionHistoryModal({ + opened, + onClose, + file, + onChanged, +}: VersionHistoryModalProps) { + const { t } = useTranslation(); + const { actions: fileActions } = useFileActions(); + const { actions: navActions } = useNavigationActions(); + const dbRevision = useIndexedDBRevision(); + + const [chain, setChain] = useState([]); + const [loading, setLoading] = useState(false); + + useEffect(() => { + if (!opened || !file) { + setChain([]); + return; + } + let cancelled = false; + setLoading(true); + const rootId = (file.originalFileId ?? file.id) as FileId; + fileStorage + .getHistoryChainStubs(rootId) + .then((c) => { + if (!cancelled) setChain(c); + }) + .catch((err) => { + console.error("Failed to load version history", err); + if (!cancelled) setChain([]); + }) + .finally(() => { + if (!cancelled) setLoading(false); + }); + return () => { + cancelled = true; + }; + // dbRevision so the chain refreshes after a version is removed. + }, [opened, file, dbRevision]); + + const stubById = useCallback( + (id: FileId) => chain.find((c) => c.id === id), + [chain], + ); + + const handleAddToWorkspace = useCallback( + (ids: FileId[]) => { + const stubs = ids + .map((id) => stubById(id)) + .filter((s): s is StirlingFileStub => Boolean(s)); + if (stubs.length === 0) return; + void fileActions.addStirlingFileStubs(stubs); + navActions.setWorkbench("fileEditor"); + onClose(); + }, + [stubById, fileActions, navActions, onClose], + ); + + const handleRemove = useCallback( + async (ids: FileId[]) => { + await fileActions.removeFiles(ids, true); + onChanged?.(); + // If only one version remains there's no journey to show. + const remaining = chain.filter((c) => !ids.includes(c.id)); + if (remaining.length <= 1) onClose(); + }, + [fileActions, chain, onChanged, onClose], + ); + + return ( + + {loading ? ( +
+ +
+ ) : chain.length > 1 && file ? ( + + ) : ( + + {t( + "filesPage.versionHistoryEmpty", + "This file has no earlier versions.", + )} + + )} +
+ ); +} diff --git a/frontend/editor/src/core/components/filesPage/VersionTimeline.tsx b/frontend/editor/src/core/components/filesPage/VersionTimeline.tsx new file mode 100644 index 0000000000..937db8c384 --- /dev/null +++ b/frontend/editor/src/core/components/filesPage/VersionTimeline.tsx @@ -0,0 +1,316 @@ +import { useMemo, useState } from "react"; +import { useTranslation } from "react-i18next"; +import { ActionIcon, Badge, Menu } from "@mantine/core"; +import OpenInNewIcon from "@mui/icons-material/OpenInNew"; +import DeleteIcon from "@mui/icons-material/Delete"; +import DownloadIcon from "@mui/icons-material/Download"; +import HistoryIcon from "@mui/icons-material/History"; +import MoreVertIcon from "@mui/icons-material/MoreVert"; +import KeyboardArrowDownIcon from "@mui/icons-material/KeyboardArrowDown"; + +import { FileId, ToolOperation } from "@app/types/file"; +import { ToolId } from "@app/types/toolId"; +import { StirlingFileStub } from "@app/types/fileContext"; +import { formatFileSize, getFileDate } from "@app/utils/fileUtils"; +import { downloadFileFromStorage } from "@app/utils/downloadUtils"; +import ToolChain from "@app/components/shared/ToolChain"; + +/** Small label/value row; shared with FileDetailsPanel. */ +export function DetailField({ + label, + value, +}: { + label: string; + value: string; +}) { + return ( +
+ {label} + {value} +
+ ); +} + +/** Tool that produced `version` from `prior`; null for v1. */ +function deltaToolFor( + version: StirlingFileStub, + prior: StirlingFileStub | null, +): ToolOperation | null { + if (!prior) return null; + const priorLen = prior.toolHistory?.length ?? 0; + const curr = version.toolHistory ?? []; + return curr[priorLen] ?? null; +} + +/** Translated tool name via `home.{toolId}.title`. */ +function ToolLabel({ toolId }: { toolId: ToolId }) { + const { t } = useTranslation(); + return {t(`home.${toolId}.title`, toolId)}; +} + +export interface VersionTimelineProps { + /** Chain sorted oldest-first. */ + chain: StirlingFileStub[]; + /** Currently selected version. */ + currentId: FileId; + onAddToWorkspace: (fileIds: FileId[]) => void; + onRemove: (fileIds: FileId[]) => void; +} + +/** Version timeline with per-row tool deltas and collapse-when-long. */ +export function VersionTimeline({ + chain, + currentId, + onAddToWorkspace, + onRemove, +}: VersionTimelineProps) { + const { t } = useTranslation(); + const [expandedIds, setExpandedIds] = useState>(new Set()); + const [showAllCollapsed, setShowAllCollapsed] = useState(false); + + // Newest-first ordering. + const ordered = useMemo( + () => + [...chain].sort( + (a, b) => (b.versionNumber ?? 1) - (a.versionNumber ?? 1), + ), + [chain], + ); + + // Index by versionNumber for prior-version lookup. + const byVersionNumber = useMemo(() => { + const map = new Map(); + for (const v of chain) { + map.set(v.versionNumber ?? 1, v); + } + return map; + }, [chain]); + + // Collapse middle when long: 3 newest + ellipsis + 2 oldest. + const COLLAPSE_THRESHOLD = 6; + const collapsible = ordered.length > COLLAPSE_THRESHOLD; + type Row = + | { kind: "version"; version: StirlingFileStub } + | { + kind: "ellipsis"; + hidden: number; + }; + const rows: Row[] = useMemo(() => { + if (!collapsible || showAllCollapsed) { + return ordered.map((v) => ({ kind: "version", version: v }) as Row); + } + const head = ordered + .slice(0, 3) + .map((v) => ({ kind: "version", version: v }) as Row); + const tail = ordered + .slice(-2) + .map((v) => ({ kind: "version", version: v }) as Row); + const hidden = ordered.length - 5; + return [...head, { kind: "ellipsis", hidden }, ...tail]; + }, [collapsible, showAllCollapsed, ordered]); + + const toggleExpand = (id: FileId) => { + setExpandedIds((prev) => { + const next = new Set(prev); + if (next.has(id)) next.delete(id); + else next.add(id); + return next; + }); + }; + + return ( +
+
+ + {t("filesPage.field.versionHistory", "Version journey")} + + {t("filesPage.versionsCount", "{{count}} versions", { + count: ordered.length, + })} + +
+
    + {rows.map((row, idx) => { + const isLast = idx === rows.length - 1; + if (row.kind === "ellipsis") { + return ( +
  1. +
    + + {!isLast && ( + + )} +
    + +
  2. + ); + } + const v = row.version; + const isActive = v.id === currentId; + const isExpanded = expandedIds.has(v.id); + const prior = byVersionNumber.get((v.versionNumber ?? 1) - 1) ?? null; + const delta = deltaToolFor(v, prior); + return ( +
  3. +
    + + {!isLast && ( + + )} +
    +
    + +
    + {formatFileSize(v.size)} + {v.lastModified ? ( + <> + · + + {getFileDate({ lastModified: v.lastModified })} + + + ) : null} + {/* Kebab on every row - the original/active version also + needs download + open-in-workspace. */} + + + + e.stopPropagation()} + > + + + + + } + onClick={() => onAddToWorkspace([v.id])} + > + {t( + "filesPage.openVersionInWorkspace", + "Open in workspace", + )} + + } + onClick={() => { + void downloadFileFromStorage(v); + }} + > + {t( + "filesPage.downloadVersion", + "Download this version", + )} + + + } + onClick={() => onRemove([v.id])} + > + {t("filesPage.removeVersion", "Remove this version")} + + + +
    + {isExpanded && ( + // Filename + full cumulative tool chain. +
    + + {v.toolHistory && v.toolHistory.length > 0 && ( +
    + + {t( + "filesPage.field.toolHistoryAtVersion", + "Cumulative tool chain", + )} + + +
    + )} +
    + )} +
    +
  4. + ); + })} +
+ {collapsible && showAllCollapsed && ( + + )} +
+ ); +} diff --git a/frontend/editor/src/core/components/layout/Workbench.tsx b/frontend/editor/src/core/components/layout/Workbench.tsx index 9d330ba7ba..25b15fd263 100644 --- a/frontend/editor/src/core/components/layout/Workbench.tsx +++ b/frontend/editor/src/core/components/layout/Workbench.tsx @@ -2,7 +2,6 @@ import { useEffect, useState, Suspense, lazy } from "react"; import { useTranslation } from "react-i18next"; import KeyboardArrowDownIcon from "@mui/icons-material/KeyboardArrowDown"; import { Box, Loader, Center } from "@mantine/core"; -import { useRainbowThemeContext } from "@app/components/shared/RainbowThemeProvider"; import { useToolWorkflow } from "@app/contexts/ToolWorkflowContext"; import { useFileHandler } from "@app/hooks/useFileHandler"; import { useFileState } from "@app/contexts/FileContext"; @@ -13,12 +12,13 @@ import { import { isBaseWorkbench } from "@app/types/workbench"; import { VIEWER_SUPPORTED_EXTENSIONS } from "@app/utils/fileUtils"; import { useAppConfig } from "@app/contexts/AppConfigContext"; +import { useCookieConsent } from "@app/hooks/useCookieConsent"; import styles from "@app/components/layout/Workbench.module.css"; import WorkbenchBar from "@app/components/shared/WorkbenchBar"; import LandingPage from "@app/components/shared/LandingPage"; -import Footer from "@app/components/shared/Footer"; import DismissAllErrorsButton from "@app/components/shared/DismissAllErrorsButton"; +import { ChatFAB } from "@app/components/chat/ChatFAB"; // Workbench panels are loaded on demand. Viewer pulls in pdfjs-dist and the // full @embedpdf plugin set; FileEditor/PageEditor are only needed once a file @@ -35,9 +35,12 @@ const FileManagerView = lazy( // No props needed - component uses contexts directly export default function Workbench() { - const { isRainbowMode } = useRainbowThemeContext(); const { config } = useAppConfig(); + // The consent banner used to be initialised by the footer; the legal links + // now live in Settings → Legal, so the workbench owns the banner lifecycle. + useCookieConsent({ analyticsEnabled: config?.enableAnalytics === true }); + // Use context-based hooks to eliminate all prop drilling const { selectors } = useFileState(); const { workbench: currentView } = useNavigationState(); @@ -199,14 +202,7 @@ export default function Workbench() { {/* Workbench Bar - animates in/out based on file presence */} {currentView !== "myFiles" && @@ -248,6 +244,9 @@ export default function Workbench() { {/* Dismiss All Errors Button */} + {/* Floating AI chat button + panel */} + + {/* Main content area */} - -
); } diff --git a/frontend/editor/src/core/components/onboarding/InitialOnboardingModal/InitialOnboardingModal.module.css b/frontend/editor/src/core/components/onboarding/InitialOnboardingModal/InitialOnboardingModal.module.css index 5b3aec2d91..c1d88f46cb 100644 --- a/frontend/editor/src/core/components/onboarding/InitialOnboardingModal/InitialOnboardingModal.module.css +++ b/frontend/editor/src/core/components/onboarding/InitialOnboardingModal/InitialOnboardingModal.module.css @@ -34,8 +34,8 @@ } .standaloneIcon { - width: 96px; - height: 96px; + width: 80px; + height: 80px; object-fit: contain; animation: heroLogoScale 0.25s ease forwards; } @@ -124,8 +124,6 @@ align-items: flex-start; justify-content: center; animation: heroLogoEnter 0.25s ease forwards; - position: relative; - top: 1rem; } .iconWrapper { @@ -175,8 +173,8 @@ } .downloadIcon { - width: 96px; - height: 96px; + width: 80px; + height: 80px; object-fit: contain; animation: heroLogoScale 0.25s ease forwards; } @@ -205,6 +203,7 @@ } .bodyCopy { + text-align: center; opacity: 0; transform: translateX(24px); animation: bodySlideIn 0.25s ease forwards; diff --git a/frontend/editor/src/core/components/onboarding/slides/DesktopInstallSlide.tsx b/frontend/editor/src/core/components/onboarding/slides/DesktopInstallSlide.tsx index eb3e743527..d8549d0ea9 100644 --- a/frontend/editor/src/core/components/onboarding/slides/DesktopInstallSlide.tsx +++ b/frontend/editor/src/core/components/onboarding/slides/DesktopInstallSlide.tsx @@ -1,7 +1,7 @@ import React from "react"; import { useTranslation } from "react-i18next"; import { SlideConfig } from "@app/types/types"; -import { UNIFIED_CIRCLE_CONFIG } from "@app/components/onboarding/slides/unifiedBackgroundConfig"; +import { UNIFIED_LIGHT_BACKGROUND } from "@app/components/onboarding/slides/unifiedBackgroundConfig"; import { DesktopInstallTitle, type OSOption, @@ -47,9 +47,6 @@ export default function DesktopInstallSlide({ ), body: , downloadUrl: osUrl, - background: { - gradientStops: ["#2563EB", "#0EA5E9"], - circles: UNIFIED_CIRCLE_CONFIG, - }, + background: UNIFIED_LIGHT_BACKGROUND, }; } diff --git a/frontend/editor/src/core/components/onboarding/slides/MFASetupSlide.tsx b/frontend/editor/src/core/components/onboarding/slides/MFASetupSlide.tsx index 94bba3010d..de76eadd6b 100644 --- a/frontend/editor/src/core/components/onboarding/slides/MFASetupSlide.tsx +++ b/frontend/editor/src/core/components/onboarding/slides/MFASetupSlide.tsx @@ -16,6 +16,7 @@ import { TextInput, } from "@mantine/core"; import { QRCodeSVG } from "qrcode.react"; +import { useTranslation } from "react-i18next"; import { SlideConfig } from "@app/types/types"; import { UNIFIED_CIRCLE_CONFIG } from "@app/components/onboarding/slides/unifiedBackgroundConfig"; import { accountService } from "@app/services/accountService"; @@ -31,6 +32,7 @@ interface MFASetupSlideProps { } function MFASetupContent({ onMfaSetupComplete }: MFASetupSlideProps) { + const { t } = useTranslation(); const [mfaSetupData, setMfaSetupData] = useState( null, ); @@ -179,7 +181,9 @@ function MFASetupContent({ onMfaSetupComplete }: MFASetupSlideProps) { {mfaLoading && !isReady && ( - Generating your QR code… + + {t("onboarding.mfa.qrCodeLoading", "Generating your QR code…")} + )} @@ -189,7 +193,10 @@ function MFASetupContent({ onMfaSetupComplete }: MFASetupSlideProps) { diff --git a/frontend/editor/src/core/components/onboarding/slides/SecurityCheckSlide.tsx b/frontend/editor/src/core/components/onboarding/slides/SecurityCheckSlide.tsx index 78ed4d6db4..02b3e8502c 100644 --- a/frontend/editor/src/core/components/onboarding/slides/SecurityCheckSlide.tsx +++ b/frontend/editor/src/core/components/onboarding/slides/SecurityCheckSlide.tsx @@ -37,11 +37,20 @@ export default function SecurityCheckSlide({