diff --git a/.github/workflows/pr-code-security.yml b/.github/workflows/pr-code-security.yml index fde10666c..3448776aa 100644 --- a/.github/workflows/pr-code-security.yml +++ b/.github/workflows/pr-code-security.yml @@ -4,13 +4,34 @@ on: pull_request: branches: [main] +permissions: + contents: read + jobs: secret-detection: name: Secret Detection - if: github.event_name == 'pull_request' - uses: prisma/.github/.github/workflows/secret_detection.yml@main - secrets: inherit - code-scanning: - name: Code Scanning - if: github.event_name == 'pull_request' - uses: prisma/.github/.github/workflows/code_scanning.yml@main + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2 + with: + # Full history so gitleaks can scan every commit in the PR range. + fetch-depth: 0 + persist-credentials: false + # We run the gitleaks CLI directly rather than gitleaks-action: the + # action wrapper requires a paid license for repos under a GitHub + # organization, while the gitleaks binary itself is MIT-licensed and + # free. Pinned by version and verified by SHA-256 before use. + - name: Install gitleaks + env: + GITLEAKS_VERSION: 8.30.1 + GITLEAKS_SHA256: 551f6fc83ea457d62a0d98237cbad105af8d557003051f41f3e7ca7b3f2470eb + run: | + set -euo pipefail + url="https://github.com/gitleaks/gitleaks/releases/download/v${GITLEAKS_VERSION}/gitleaks_${GITLEAKS_VERSION}_linux_x64.tar.gz" + curl -sSfL "$url" -o gitleaks.tar.gz + echo "${GITLEAKS_SHA256} gitleaks.tar.gz" | sha256sum -c - + tar -xzf gitleaks.tar.gz gitleaks + sudo install gitleaks /usr/local/bin/gitleaks + gitleaks version + - name: Scan git history for secrets + run: gitleaks git --no-banner --redact --exit-code 1 --config .gitleaks.toml . diff --git a/.github/workflows/security.yml b/.github/workflows/security.yml new file mode 100644 index 000000000..71985011d --- /dev/null +++ b/.github/workflows/security.yml @@ -0,0 +1,29 @@ +name: Security audit + +on: + push: + branches: [main] + pull_request: + schedule: + # Re-run weekly so newly-published advisories are caught even without a push. + - cron: "0 6 * * 1" + +permissions: + contents: read + +jobs: + cargo-deny: + name: cargo-deny (advisories, bans, sources, licenses) + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2 + with: + persist-credentials: false + - uses: dtolnay/rust-toolchain@e97e2d8cc328f1b50210efc529dca0028893a2d9 # v1 + with: + toolchain: stable + # Run cargo-deny on the runner directly (the musl container action conflicts with the repo's rust-toolchain file). + - name: Install cargo-deny + uses: taiki-e/install-action@cargo-deny + - name: Check advisories, bans, sources, licenses + run: cargo deny check advisories bans sources licenses diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index fec4c17a0..4de9c1fdc 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -1,39 +1,132 @@ name: Cargo tests on: + # qa + manual dispatch run the integration tier push: branches: - main + - qa pull_request: + workflow_dispatch: jobs: clippy: runs-on: ubuntu-latest steps: - - uses: actions/checkout@v1 - - uses: actions-rs/toolchain@v1 + - uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2 with: - components: clippy - override: true - - name: Install dependencies - run: sudo apt install -y openssl libkrb5-dev - - uses: actions-rs/clippy-check@v1 + persist-credentials: false + - uses: dtolnay/rust-toolchain@e97e2d8cc328f1b50210efc529dca0028893a2d9 # v1 with: - token: ${{ secrets.GITHUB_TOKEN }} - args: --features=all + toolchain: stable + components: clippy + - uses: Swatinem/rust-cache@c19371144df3bb44fab255c43d04cbc2ab54d1c4 # v2.9.1 + - name: Install dependencies + run: sudo apt-get update && sudo apt-get install -y openssl libkrb5-dev + - name: Clippy + # Advisory here: modernizes the retired actions-rs/clippy-check and reports + # lints without gating. The strict `-D warnings` gate lands together with its + # baseline-lint fixes in the feature stack, so main is never red in between. + run: cargo clippy --features=all format: runs-on: ubuntu-latest steps: - - uses: actions/checkout@v2 - - uses: actions-rs/toolchain@v1 + - uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2 with: - components: rustfmt - override: true - - uses: mbrobbel/rustfmt-check@master + persist-credentials: false + - uses: dtolnay/rust-toolchain@e97e2d8cc328f1b50210efc529dca0028893a2d9 # v1 with: - token: ${{ secrets.GITHUB_TOKEN }} + toolchain: stable + components: rustfmt + - name: Rustfmt + run: cargo fmt --check - cargo-test-linux: + msrv: + name: MSRV (1.88) + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2 + with: + persist-credentials: false + # Keep this in sync with `rust-version` in Cargo.toml. + - uses: dtolnay/rust-toolchain@e97e2d8cc328f1b50210efc529dca0028893a2d9 # v1 + with: + toolchain: "1.88" + - uses: Swatinem/rust-cache@c19371144df3bb44fab255c43d04cbc2ab54d1c4 # v2.9.1 + - name: Install dependencies + run: sudo apt-get update && sudo apt-get install -y openssl libkrb5-dev + - name: Check on MSRV + run: cargo check --features all + + # The `rustls-webpki-roots` feature (and the rustls trust-store code paths in + # general) are excluded from `--features=all` because the three TLS backends + # are mutually exclusive, so neither the clippy nor the MSRV job above builds + # them. This server-less lane keeps them clippy-clean, buildable, and unit- + # tested so a regression (e.g. a `webpki-roots` major bump changing + # `TLS_SERVER_ROOTS`) is caught in CI. + rustls-features: + name: rustls trust-store (unit) + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2 + with: + persist-credentials: false + - uses: dtolnay/rust-toolchain@e97e2d8cc328f1b50210efc529dca0028893a2d9 # v1 + with: + toolchain: stable + - uses: Swatinem/rust-cache@c19371144df3bb44fab255c43d04cbc2ab54d1c4 # v2.9.1 + - name: Install dependencies + run: sudo apt-get update && sudo apt-get install -y openssl libkrb5-dev + # Build + unit-test both rustls feature sets. `--lib` needs no live server. + # (Strict `-D warnings` clippy is deferred repo-wide until the baseline + # lints land, matching the advisory `clippy` job above.) + - name: Unit tests (rustls) + run: cargo test --no-default-features --features rustls,chrono,time,tds73 --lib + - name: Unit tests (rustls-webpki-roots) + run: cargo test --no-default-features --features rustls-webpki-roots,chrono,time,tds73 --lib + + semver: + name: semver-checks (advisory) + if: github.event_name == 'pull_request' + runs-on: ubuntu-latest + # Advisory during 0.x: reports API-breaking changes vs the target branch + # without blocking (breaking changes are allowed pre-1.0, but should be + # visible in review). + continue-on-error: true + steps: + - uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2 + with: + persist-credentials: false + fetch-depth: 0 + - uses: dtolnay/rust-toolchain@e97e2d8cc328f1b50210efc529dca0028893a2d9 # v1 + with: + toolchain: stable + - uses: Swatinem/rust-cache@c19371144df3bb44fab255c43d04cbc2ab54d1c4 # v2.9.1 + - name: Install dependencies + run: sudo apt-get update && sudo apt-get install -y openssl libkrb5-dev + - name: Install cargo-semver-checks + uses: taiki-e/install-action@cargo-semver-checks + - name: Fetch baseline branch + # `actions/checkout` leaves the PR base branch without a local ref, so + # `--baseline-rev origin/` can't be resolved; fetch it explicitly + # and compare against FETCH_HEAD. + run: git fetch --no-tags --depth=100 origin "${{ github.base_ref }}" + - name: Check for semver-breaking changes + # Scope to one coherent, non-conflicting feature set. Checking all + # features at once enables the mutually-exclusive TLS backends + # (native-tls + rustls + vendored-openssl) together, whose duplicate + # `TlsStream` definitions make rustdoc fail to build (E0428). These are + # also the crate's original features, so they exist in every baseline + # slice this PR is compared against. + run: >- + cargo semver-checks --baseline-rev FETCH_HEAD + --only-explicit-features + --features rustls,chrono,time,tds73,rust_decimal,bigdecimal + + smoke: + name: integration smoke (SQL 2022) + # Fast real-server signal on the dev lane; the full matrix runs on qa. + if: github.event_name == 'pull_request' || (github.event_name == 'push' && github.ref == 'refs/heads/dev') runs-on: ubuntu-latest strategy: @@ -51,34 +144,143 @@ jobs: - "--no-default-features --features=time" - "--no-default-features --features=rustls" - "--no-default-features --features=vendored-openssl" + - "--no-default-features --features=rustls,chrono,time,tds73,rust_decimal,bigdecimal" env: TIBERIUS_TEST_CONNECTION_STRING: "server=tcp:localhost,1433;user=SA;password=;TrustServerCertificate=true" RUSTFLAGS: "-Dwarnings" steps: - - uses: actions/checkout@v2 + - uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2 + with: + persist-credentials: false - - uses: actions-rs/toolchain@v1 + - uses: dtolnay/rust-toolchain@e97e2d8cc328f1b50210efc529dca0028893a2d9 # v1 + with: + toolchain: stable + + - name: Compute cache key + shell: bash + run: | + key="${{ matrix.features }}" + key="${key//,/+}" + echo "RUST_CACHE_KEY=$key" >> "$GITHUB_ENV" - - uses: actions/cache@v2 + - uses: Swatinem/rust-cache@c19371144df3bb44fab255c43d04cbc2ab54d1c4 # v2.9.1 with: - path: | - ~/.cargo/registry - ~/.cargo/git - target - key: ${{ runner.os }}-cargo-${{ matrix.features }} + shared-key: ${{ env.RUST_CACHE_KEY }} - name: Start SQL Server ${{matrix.database}} - run: DOCKER_BUILDKIT=1 docker-compose -f docker-compose.yml up -d mssql-${{matrix.database}} + run: DOCKER_BUILDKIT=1 docker compose -f docker-compose.yml up -d mssql-${{matrix.database}} - name: Install dependencies - run: sudo apt install -y openssl libkrb5-dev + run: sudo apt-get update && sudo apt-get install -y openssl libkrb5-dev + + - name: Wait for SQL Server + # A listening port is not readiness: SQL Server binds 1433 before the SA + # login and databases finish initializing, so tests started too early race + # it and hit sporadic connection/login failures. Gate on an authenticated + # `SELECT 1` from a throwaway mssql-tools container (works uniformly across + # the full server images and azure-sql-edge, which ships no in-box sqlcmd). + run: | + pw='' + for _ in $(seq 1 60); do + if docker run --rm --network host mcr.microsoft.com/mssql-tools \ + /opt/mssql-tools/bin/sqlcmd -S localhost,1433 -U SA -P "$pw" -Q "SELECT 1" >/dev/null 2>&1; then + echo "SQL Server ready (authenticated login succeeded)"; exit 0 + fi + sleep 3 + done + echo "SQL Server did not accept an authenticated login in time" >&2 + docker compose -f docker-compose.yml logs mssql-${{matrix.database}} || true + exit 1 - name: Run tests run: cargo test ${{matrix.features}} - cargo-test-windows: + integration-macos: + name: macos (SQL 2022 via colima) + if: github.event_name == 'workflow_dispatch' || (github.event_name == 'push' && github.ref == 'refs/heads/qa') + # Intel runner: hosted macOS has no Linux Docker daemon, so we host the + # linux/amd64 SQL Server container inside a colima (Lima) Linux VM. An x86_64 + # runner runs that image natively — no qemu emulation — which keeps the VM + # fast enough for the full integration suite instead of build + unit only. + runs-on: macos-15-intel + continue-on-error: ${{ matrix.soft_fail }} + strategy: + fail-fast: false + matrix: + # Exercise every TLS backend: + include: + # rustls (TLS 1.3) + - features: "--no-default-features --features=rustls,chrono,time,tds73,rust_decimal,bigdecimal" + soft_fail: false + # vendored-openssl (opentls, statically-linked OpenSSL) — cross-platform. + - features: "--no-default-features --features=vendored-openssl,chrono,time,tds73,rust_decimal,bigdecimal" + soft_fail: false + # NOTE: a `--features=all` lane is intentionally omitted here. On macOS + # `native-tls` resolves to Apple Secure Transport, which cannot complete + # the SQL Server TLS handshake (it works on Linux=OpenSSL and + # Windows=SChannel, both covered elsewhere). Running it here could only + # ever fail, so the two TLS backends that do work on macOS (rustls and + # vendored-openssl, above) provide the macOS integration coverage. + env: + TIBERIUS_TEST_CONNECTION_STRING: "server=tcp:localhost,1433;user=SA;password=;TrustServerCertificate=true" + RUSTFLAGS: "-Dwarnings" + steps: + - uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2 + with: + persist-credentials: false + - uses: dtolnay/rust-toolchain@e97e2d8cc328f1b50210efc529dca0028893a2d9 # v1 + with: + toolchain: stable + - name: Compute cache key + run: | + key="${{ matrix.features }}" + key="${key//,/+}" + echo "RUST_CACHE_KEY=$key" >> "$GITHUB_ENV" + - uses: Swatinem/rust-cache@c19371144df3bb44fab255c43d04cbc2ab54d1c4 # v2.9.1 + with: + shared-key: ${{ env.RUST_CACHE_KEY }} + - name: Install build dependencies (openssl, krb5) + run: | + brew install openssl krb5 + # krb5 is keg-only; expose its pkg-config so integrated-auth-gssapi + # and any openssl-linking features build against Homebrew's copy. + echo "PKG_CONFIG_PATH=$(brew --prefix krb5)/lib/pkgconfig:$(brew --prefix openssl)/lib/pkgconfig" >> "$GITHUB_ENV" + - name: Start Docker (colima) + run: | + brew install colima docker docker-compose + # SQL Server wants ~2 GB RAM; give the VM headroom for it + the build. + colima start --cpu 3 --memory 6 --disk 20 + docker version + - name: Start SQL Server 2022 + run: docker compose -f docker-compose.yml up -d mssql-2022 + - name: Wait for SQL Server + # A listening port is not readiness: SQL Server binds 1433 before the SA + # login and databases finish initializing, so the connection-heavy bulk + # tests can race it and hit "Login failed for user 'SA'" (18456). Gate on + # an actual authenticated `SELECT 1` succeeding instead. + run: | + pw='' + for _ in $(seq 1 60); do + for bin in /opt/mssql-tools18/bin/sqlcmd /opt/mssql-tools/bin/sqlcmd; do + if docker compose -f docker-compose.yml exec -T mssql-2022 \ + "$bin" -S localhost -U SA -P "$pw" -C -Q "SELECT 1" >/dev/null 2>&1; then + echo "SQL Server ready (authenticated login succeeded)"; exit 0 + fi + done + sleep 3 + done + echo "SQL Server did not accept an authenticated login in time" >&2 + docker compose -f docker-compose.yml logs mssql-2022 || true + exit 1 + - name: Run tests + run: cargo test ${{ matrix.features }} + + integration-windows: + name: windows (SQL 2019, integrated auth) + if: github.event_name == 'workflow_dispatch' || (github.event_name == 'push' && github.ref == 'refs/heads/qa') runs-on: windows-latest strategy: @@ -96,41 +298,39 @@ jobs: TIBERIUS_TEST_CONNECTION_STRING: "server=tcp:127.0.0.1,1433;IntegratedSecurity=true;TrustServerCertificate=true" steps: - - uses: actions/checkout@v2 + - uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2 + with: + persist-credentials: false + + - uses: dtolnay/rust-toolchain@e97e2d8cc328f1b50210efc529dca0028893a2d9 # v1 + with: + toolchain: stable + + - name: Compute cache key + shell: bash + run: | + key="${{ matrix.features }}" + key="${key//,/+}" + echo "RUST_CACHE_KEY=$key" >> "$GITHUB_ENV" - - uses: actions-rs/toolchain@v1 + - uses: Swatinem/rust-cache@c19371144df3bb44fab255c43d04cbc2ab54d1c4 # v2.9.1 + with: + shared-key: ${{ env.RUST_CACHE_KEY }} - name: Set required PowerShell modules id: psmodulecache - uses: potatoqualitee/psmodulecache@v1 + uses: potatoqualitee/psmodulecache@ee5e9494714abf56f6efbfa51527b2aec5c761b8 # v6.2.1 with: modules-to-cache: SqlServer - - name: Setup PowerShell module cache - id: cacher - uses: actions/cache@v2 - with: - path: ${{ steps.psmodulecache.outputs.modulepath }} - key: ${{ steps.psmodulecache.outputs.keygen }} - - name: Setup Chocolatey download cache id: chococache - uses: actions/cache@v2 + uses: actions/cache@27d5ce7f107fe9357f9df03efb73ab90386fccae # v5.0.5 with: path: C:\Users\runneradmin\AppData\Local\Temp\chocolatey\ key: chocolatey-install - - name: Setup Cargo build cache - uses: actions/cache@v2 - with: - path: | - C:\Users\runneradmin\.cargo\registry - C:\Users\runneradmin\.cargo\git - target - key: ${{ runner.os }}-cargo - - name: Install required PowerShell modules - if: steps.cacher.outputs.cache-hit != 'true' shell: powershell run: | Set-PSRepository PSGallery -InstallationPolicy Trusted @@ -187,31 +387,3 @@ jobs: - name: Run normal tests shell: powershell run: cargo test ${{matrix.features}} - - cargo-test-macos: - runs-on: macos-12 - - strategy: - fail-fast: false - matrix: - database: - - 2019 - features: - - "--no-default-features --features=rustls,chrono,time,tds73,sql-browser-async-std,sql-browser-tokio,sql-browser-smol,integrated-auth-gssapi,rust_decimal,bigdecimal" - - "--no-default-features --features=vendored-openssl" - - env: - TIBERIUS_TEST_CONNECTION_STRING: "server=tcp:localhost,1433;user=SA;password=;TrustServerCertificate=true" - - steps: - - uses: actions/checkout@v2 - - - uses: actions-rs/toolchain@v1 - - - uses: docker-practice/actions-setup-docker@master - - - name: Start SQL Server ${{matrix.database}} - run: DOCKER_BUILDKIT=1 docker-compose -f docker-compose.yml up -d mssql-${{matrix.database}} - - - name: Run tests - run: cargo test ${{matrix.features}} diff --git a/.gitleaks.toml b/.gitleaks.toml new file mode 100644 index 000000000..3439483e1 --- /dev/null +++ b/.gitleaks.toml @@ -0,0 +1,10 @@ +# gitleaks config for the PR secret scan: full default rule set, with the +# self-signed TLS test fixtures under docker/certs/ allowlisted. +[extend] +useDefault = true + +[allowlist] +description = "Self-signed TLS test fixtures for the local integration-test SQL Server container" +paths = [ + '''docker/certs/.*''', +] diff --git a/CHANGELOG.md b/CHANGELOG.md index fed7001e5..42611d30c 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,83 @@ # Changes +## Version 0.13.0 + +- feat: TLS trust configuration is revamped around two orthogonal axes plus a + bypass, unifying community PRs #330 and #290: + - `Config::trust_cert_ca_bundle(bytes)` (and the `ConfigBuilder` mirror) trusts + additional CA certificates supplied as in-memory bytes, without writing them + to a temporary file. The bytes are auto-detected: a `-----BEGIN` marker is + parsed as a multi-certificate PEM bundle (e.g. the AWS RDS root bundle), + otherwise they are treated as a single DER certificate. + - `Config::trust_webpki_roots()` (and the `ConfigBuilder` mirror) bases trust + on a compiled-in snapshot of Mozilla's root CA store instead of the OS trust + store. rustls-only, behind the new `rustls-webpki-roots` feature. Note: the + bundled roots are a pinned snapshot that goes stale (missing newly added or + newly distrusted CAs) unless the dependency is updated and the app rebuilt. + - `Config::trust_cert_ca(path)` now accepts **multi-certificate** files (the + previous "exactly one certificate" restriction is lifted); every certificate + in the file is trusted, on all three TLS backends. +- BREAKING: repeated `trust_cert_ca` calls now **accumulate** rather than + replace-last-wins. `trust_cert_ca(a); trust_cert_ca(b)` (and any mix with + `trust_cert_ca_bundle`) trusts every supplied CA, layered on top of the base + trust anchors. Code that relied on a later call overriding an earlier one must + now set the CA only once. +- fix: a CA source (file or in-memory bundle) that yields zero usable + certificates is now a hard error naming the source, on every backend, instead + of silently degrading to base-roots-only trust. The `native-tls` and + `vendored-openssl` backends also now load *all* certificates from a + multi-certificate CA file/bundle (previously only the first was used) and + preserve the path plus underlying I/O error in load failures, matching rustls. +- BREAKING: the connection-string `encrypt` default is now `Required` (was + `Off`) when a TLS backend is enabled, matching modern ADO.NET; without a TLS + backend it remains `NotSupported`. +- BREAKING: removed the `sql-browser-async-std` feature and the async-std SQL + Browser integration. +- feat: `Command`/RPC API for parameterized stored-procedure calls, plus a + `#[derive(TableValueRow)]` macro (in `tiberius-macros`) for table-valued + parameters. +- feat: `sspi-rs` feature for Windows-style SSPI/NTLM authentication on Unix via + the pure-Rust `sspi` crate (no Kerberos required). +- feat: `serde` feature adding `Serialize`/`Deserialize` impls for query result + types (`Row`, `Column`, `ColumnData`, `Numeric`, and the time/xml types). +- feat: client-certificate authentication, including PEM/DER key files + (`Config::client_certificate`) and PKCS#12 bundles + (`Config::client_certificate_pkcs12`). +- BREAKING: credentials (SQL Server / Windows passwords, the AAD bearer token + and the PKCS#12 password) are now stored as `secrecy::SecretString` instead of + `zeroize::Zeroizing`. They are still zeroized on drop, and their + `Debug` now renders as `SecretBox([REDACTED])` (was ``). The + `AuthMethod::AADToken` tuple variant consequently holds a `SecretString`: code + that pattern-matched it and read the token via `Deref`/`Display` must now call + `secrecy::ExposeSecret::expose_secret`. Constructing auth via + `AuthMethod::aad_token`/`sql_server`/`windows` is unchanged. +- feat: connection & command timeouts (closes #375 and #360), matching + ADO.NET's two-knob model and backed by a runtime-agnostic timer so they apply + under any async runtime: + - `Config::handshake_timeout` bounds the whole post-TCP handshake (prelogin, + TLS negotiation and login), surfacing a `TimedOut` error instead of hanging + forever when a server accepts the TCP connection and then stalls + mid-handshake — the reported failure against `azure-sql-edge` on macOS + (#375) and the stalled-peer case in #360. Defaults to 15s (ADO.NET + `Connect Timeout` parity). The handshake also emits per-stage `tracing` + DEBUG events so a stall can be pinpointed. + - `Config::command_timeout` bounds each server round-trip while reading + command results (`query`/`execute`/`simple_query`, the `bulk_insert` + acknowledgement and `column_metadata`). It measures per-round-trip stall, + not total enumeration: the deadline resets on every delivered token, so a + slow consumer never trips it — only a stalled server does. Defaults to 30s + (ADO.NET `Command Timeout` parity). + - BREAKING: both knobs now default to a bounded value (15s handshake / 30s + command) where pre-0.13 they were effectively unbounded. A command that + legitimately runs longer than 30s between server round-trips (e.g. a long + `WAITFOR`, a big sort/aggregate or a slow stored procedure) will now fail + with a `TimedOut` error unless you raise or disable `command_timeout`. Pass + `None` to either knob to restore the pre-0.13 wait-indefinitely behaviour. +- chore: upgraded the rustls stack to 0.23 (tokio-rustls 0.26) and resolved the + associated advisories. +- fix: numerous decode-path hardening fixes (protocol errors instead of panics + or stream desyncs on hostile server input across the codec/token modules). + ## Version 0.12.3 - feat: improve column type accuracy (#347) - fix: encoding of zero-length values for large varlen columns (#315) diff --git a/Cargo.toml b/Cargo.toml index fb45b46e5..16bd8b76c 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -8,24 +8,24 @@ authors = [ description = "A TDS (MSSQL) driver" documentation = "https://docs.rs/tiberius/" edition = "2021" +rust-version = "1.88" keywords = ["tds", "mssql", "sql"] license = "MIT/Apache-2.0" name = "tiberius" readme = "README.md" repository = "https://github.com/prisma/tiberius" -version = "0.12.3" +version = "0.13.0" [workspace] -members = ["runtimes-macro"] +members = ["runtimes-macro", "tiberius-macros"] [[test]] path = "tests/query.rs" name = "query" [[test]] -path = "tests/named-instance-async.rs" -name = "named-instance-async" -required-features = ["sql-browser-async-std"] +path = "tests/command.rs" +name = "command" [[test]] path = "tests/named-instance-tokio.rs" @@ -37,29 +37,46 @@ path = "tests/named-instance-smol.rs" name = "named-instance-smol" required-features = ["sql-browser-smol"] +[[test]] +path = "tests/serde.rs" +name = "serde" +required-features = ["serde"] + +[[example]] +name = "named-pipes" +path = "examples/named-pipes.rs" + [dependencies] enumflags2 = "0.7" byteorder = "1.0" encoding_rs = "0.8" once_cell = "1.3" -thiserror = "1.0" +thiserror = "2" bytes = "1.0" -pretty-hex = "0.3" +pretty-hex = "0.4" pin-project-lite = "0.2" -asynchronous-codec = "0.6" +asynchronous-codec = "0.7" async-trait = "0.1" connection-string = "0.2" num-traits = "0.2" uuid = "1.0" - -winauth = { version = "0.0.4", optional = true } +zeroize = "1.8.2" +secrecy = "0.10" + +# Cross-platform: the NTLMv2 client (`winauth::NtlmV2Client`) is pure Rust and +# backs `AuthMethod::Windows` on every target. Only its `windows` SSPI module +# (used by `AuthMethod::Integrated`) is Windows-specific. +[dependencies.winauth] +version = "0.0.5" +optional = true [target.'cfg(unix)'.dependencies] -libgssapi = { version = "0.8.1", optional = true, default-features = false } +libgssapi = { version = "0.11", optional = true, default-features = false } +sspi = { version = "0.18", optional = true } +libc = "0.2" [dependencies.async-native-tls] -version = "0.4" -features = ["runtime-async-std"] +version = "0.6" optional = true [dependencies.tokio] @@ -72,11 +89,6 @@ version = "0.7" features = ["compat"] optional = true -[dependencies.async-std] -version = "1" -optional = true -features = ["attributes"] - [dependencies.chrono] version = "0.4" optional = true @@ -91,6 +103,13 @@ version = "0.3" default-features = false features = ["io", "sink"] +# Runtime-agnostic timer used to bound the connection handshake (see +# `Config::handshake_timeout`). Its `Delay` future works under any executor, so +# the generic, runtime-agnostic `Connection::connect` path can enforce a timeout +# without depending on tokio/smol timers. +[dependencies.futures-timer] +version = "3" + [dependencies.tracing] features = ["log"] version = "0.1" @@ -100,33 +119,55 @@ version = "1.6" optional = true [dependencies.bigdecimal_] -version = "0.3" +version = "0.4" optional = true package = "bigdecimal" +[dependencies.serde] +version = "1.0" +optional = true +features = ["derive", "rc"] + [dependencies.async-io] -version = "1.8" +version = "2" optional = true [dependencies.async-net] -version = "1.7" +version = "2" optional = true [dependencies.futures-lite] -version = "1.12.0" +version = "2" optional = true [dependencies.tokio-rustls] -version = "0.24.0" +version = "0.26.4" +optional = true + +[dependencies.rustls-native-certs] +version = "0.8" +optional = true + +# Version-floor pin for RUSTSEC-2026-0104 (rustls-webpki < 0.103.13); pulled in via tokio-rustls. +[dependencies.rustls-webpki] +version = "0.103.13" optional = true -features = ["dangerous_configuration"] +default-features = false -[dependencies.rustls-pemfile] +# Shared certificate DER/PEM types used by the backend-agnostic CA loader +# (`src/client/tls_stream/certs.rs`). Enabled by every TLS backend so all three +# parse trust anchors through one code path. This is the same `rustls-pki-types` +# crate `tokio-rustls` re-exports, so the `CertificateDer` handed to rustls is +# the exact same type. +[dependencies.rustls-pki-types] version = "1" optional = true -[dependencies.rustls-native-certs] -version = "0.6" +# Compiled-in Mozilla root CA bundle for the `rustls-webpki-roots` feature. +# Pinned to its current major (1.x). Note: this is a *pinned snapshot* that goes +# stale unless the dependency is updated. +[dependencies.webpki-roots] +version = "1" optional = true [dependencies.opentls] @@ -161,50 +202,74 @@ features = [ ] version = "1.0" -[dev-dependencies.async-std] -features = ["attributes"] -version = "1" +[dev-dependencies.smol] +version = "2" [dev-dependencies.runtimes-macro] path = "./runtimes-macro" +[dependencies.tiberius-macros] +path = "./tiberius-macros" +version = "0.1.0" + [dev-dependencies] names = "0.14" anyhow = "1" -env_logger = "0.9" -azure_identity = "0.5.0" -oauth2 = "4.2.3" -url = "2.2.2" -reqwest = "0.11.10" +env_logger = "0.11" +azure_identity = "1.0" +azure_core = "1" +url = "2.5" +reqwest = "0.13" paste = "1.0" -indicatif = "0.17" +indicatif = "0.18" chrono = "0.4.38" -indoc = "1.0.7" +indoc = "2" +serde_json = "1.0" [package.metadata.docs.rs] -features = ["all", "docs"] +features = ["all"] +# docs.rs builds on nightly with this cfg set, enabling #[doc(cfg(...))] +# annotations (feature(doc_cfg)) without requiring nightly for normal builds. +rustdoc-args = ["--cfg", "docsrs"] + +[lints.rust] +unexpected_cfgs = { level = "warn", check-cfg = ['cfg(docsrs)'] } [features] all = [ "chrono", "time", "tds73", - "sql-browser-async-std", + "tds80", "sql-browser-tokio", "sql-browser-smol", "integrated-auth-gssapi", "rust_decimal", "bigdecimal", "native-tls", + "serde", + "sspi-rs", ] -default = ["tds73", "winauth", "native-tls"] +default = ["tds80", "winauth", "native-tls"] tds73 = [] -docs = [] -sql-browser-async-std = ["async-std"] +# Enables TDS 8.0 support, including the `Strict` encryption level (TLS before +# the TDS prelogin, TDS 8.0 "strict" mode). Requires a TLS backend. +tds80 = ["tds73"] sql-browser-tokio = ["tokio", "tokio-util"] sql-browser-smol = ["async-io", "async-net", "futures-lite"] integrated-auth-gssapi = ["libgssapi"] bigdecimal = ["bigdecimal_"] -rustls = ["tokio-rustls", "tokio-util", "rustls-pemfile", "rustls-native-certs"] -native-tls = ["async-native-tls"] -vendored-openssl = ["opentls", "schannel"] +rustls = ["tokio-rustls", "tokio-util", "rustls-native-certs", "rustls-webpki", "rustls-pki-types"] +# Bundle a compiled-in snapshot of Mozilla's root CA store, selectable via +# `Config::trust_webpki_roots()` (rustls only). Implies `rustls`; the bundled +# roots are a pinned snapshot that goes stale without dependency updates. +rustls-webpki-roots = ["rustls", "dep:webpki-roots"] +native-tls = ["async-native-tls", "rustls-pki-types"] +vendored-openssl = ["opentls", "rustls-pki-types", "schannel"] +# Enables Windows-style SSPI/NTLM authentication (`AuthMethod::Windows`) on Unix +# platforms without requiring Kerberos, via the pure-Rust `sspi` crate. On +# Windows the same authentication is provided by the `winauth` feature. +sspi-rs = ["sspi"] +# Optional serde Serialize/Deserialize impls for query result types +# (Row, Column, ColumnData, Numeric, ColumnType and time/xml types). +serde = ["dep:serde", "uuid/serde"] diff --git a/README.md b/README.md index 44398dc55..84a102ea7 100644 --- a/README.md +++ b/README.md @@ -44,10 +44,12 @@ A native Microsoft SQL Server (TDS) client for Rust. | `time` | Read and write date and time values using `time` crate types. | `disabled` | | `rust_decimal` | Read and write `numeric`/`decimal` values using `rust_decimal`'s `Decimal`. | `disabled` | | `bigdecimal` | Read and write `numeric`/`decimal` values using `bigdecimal`'s `BigDecimal`. | `disabled` | -| `sql-browser-async-std` | SQL Browser implementation for the `TcpStream` of async-std. | `disabled` | | `sql-browser-tokio` | SQL Browser implementation for the `TcpStream` of Tokio. | `disabled` | | `sql-browser-smol` | SQL Browser implementation for the `TcpStream` of smol. | `disabled` | | `integrated-auth-gssapi` | Support for using Integrated Auth via GSSAPI | `disabled` | +| `winauth` | Windows-only SSPI/NTLM integrated authentication (`AuthMethod::Windows`). | `enabled` | +| `sspi-rs` | Windows-style SSPI/NTLM authentication on Unix via the pure-Rust `sspi` crate (no Kerberos required). | `disabled` | +| `serde` | `serde` `Serialize`/`Deserialize` impls for query result types (`Row`, `Column`, `ColumnData`, `Numeric`, etc.). | `disabled` | ### Supported protocols diff --git a/deny.toml b/deny.toml new file mode 100644 index 000000000..f515700c0 --- /dev/null +++ b/deny.toml @@ -0,0 +1,43 @@ +# Run locally with: cargo deny check advisories bans sources +# CI runs the same via .github/workflows/security.yml. +# +# Policy: a vulnerability or yanked crate in the default-built graph fails the build; ignored advisories are only reachable via dev-deps or opt-in features. + +[advisories] +yanked = "deny" +ignore = [ + # The following are ALL dev-dependency-only (test harness + the aad-auth + # example) and are never compiled into the published library. + { id = "RUSTSEC-2024-0375", reason = "atty: dev-dependency only (via `names` -> clap 3); not shipped" }, + { id = "RUSTSEC-2024-0370", reason = "proc-macro-error: dev-dependency only (via `names` -> clap 3); not shipped" }, + { id = "RUSTSEC-2024-0436", reason = "paste: dev/test only (tests/bulk.rs + azure_identity example); not shipped" }, +] + +[licenses] +# Allow-list for the current dependency graph (verified locally with +# `cargo deny check licenses`). Every crate resolves to at least one of these +# via its SPDX expression; add new entries here rather than loosening the policy. +allow = [ + "MIT", + "MIT-0", + "Apache-2.0", + "Apache-2.0 WITH LLVM-exception", + "BSD-1-Clause", + "BSD-2-Clause", + "BSD-3-Clause", + "ISC", + "CC0-1.0", + "Unicode-3.0", + "Unlicense", +] +# Confidence threshold for detecting a license from its text (0.0 - 1.0). +confidence-threshold = 0.8 + +[bans] +multiple-versions = "warn" +wildcards = "allow" + +[sources] +unknown-registry = "deny" +unknown-git = "deny" +allow-registry = ["https://github.com/rust-lang/crates.io-index"] diff --git a/docker-compose.yml b/docker-compose.yml index db5f3a39a..2aef9c6e4 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -1,6 +1,7 @@ version: "3" services: mssql-2022: + platform: linux/amd64 build: context: docker/ dockerfile: docker-mssql-2022.dockerfile @@ -12,6 +13,7 @@ services: - "1433:1433" mssql-2019: + platform: linux/amd64 build: context: docker/ dockerfile: docker-mssql-2019.dockerfile @@ -23,6 +25,7 @@ services: - "1433:1433" mssql-2017: + platform: linux/amd64 build: context: docker/ dockerfile: docker-mssql-2017.dockerfile @@ -34,6 +37,7 @@ services: - "1433:1433" mssql-azure-sql-edge: + platform: linux/amd64 build: context: docker/ dockerfile: docker-azure-sql-edge.dockerfile diff --git a/docker/certs/customCA.srl b/docker/certs/customCA.srl index 618df7789..a02a6570a 100644 --- a/docker/certs/customCA.srl +++ b/docker/certs/customCA.srl @@ -1 +1 @@ -0DAEECC45C07F5E06E0DD1B05115C3CFD1A46D9C +0DAEECC45C07F5E06E0DD1B05115C3CFD1A46D9D diff --git a/docker/certs/generate-signed-cert.sh b/docker/certs/generate-signed-cert.sh index dc3086f29..db3858cca 100755 --- a/docker/certs/generate-signed-cert.sh +++ b/docker/certs/generate-signed-cert.sh @@ -5,8 +5,10 @@ set -o pipefail # Skript creates a custom-signed certificate # Parameter1 = name of the cert +# Parameter2 = validity in days (default 1825) CERT_KEY_NAME=$1 +CERT_DAYS=${2:-1825} CERT_FILE=$CERT_KEY_NAME.crt export CERT_CN=$CERT_KEY_NAME @@ -32,7 +34,7 @@ openssl x509 -req \ -CAserial customCA.srl \ -out $CERT_FILE \ -passin file:passphrase.txt \ - -days 200 + -days $CERT_DAYS echo Generating PEM format openssl rsa -in ${CERT_KEY_NAME}.key -out ${CERT_KEY_NAME}-nopassword.key diff --git a/docker/certs/server-full.crt b/docker/certs/server-full.crt index 31ceafd70..1128cc190 100644 --- a/docker/certs/server-full.crt +++ b/docker/certs/server-full.crt @@ -1,33 +1,33 @@ -----BEGIN CERTIFICATE----- -MIIFVDCCAzygAwIBAgIUDa7sxFwH9eBuDdGwURXDz9GkbZwwDQYJKoZIhvcNAQEL -BQAwDzENMAsGA1UEAwwEQWNtZTAeFw0yNDA2MDMxMTQwMzNaFw0yNDEyMjAxMTQw -MzNaMEAxCzAJBgNVBAYTAkRFMQ0wCwYDVQQKDARBY21lMREwDwYDVQQLDAhUaWJl +MIIFVDCCAzygAwIBAgIUDa7sxFwH9eBuDdGwURXDz9GkbZ0wDQYJKoZIhvcNAQEL +BQAwDzENMAsGA1UEAwwEQWNtZTAeFw0yNjA1MTExOTA2MzlaFw0yNzExMDIxOTA2 +MzlaMEAxCzAJBgNVBAYTAkRFMQ0wCwYDVQQKDARBY21lMREwDwYDVQQLDAhUaWJl cml1czEPMA0GA1UEAwwGc2VydmVyMIICIjANBgkqhkiG9w0BAQEFAAOCAg8AMIIC -CgKCAgEAztKC7UloJuxGMaOslWm7vEDcd8YkcC9P4PMqDTS0qgr/IXeK1LB1Pt2w -iEY4Bz/Bd3boj2IMgRzT9gjtJoD6Y3Aa32UWp1TgrDtLQ6Bns30d6sNdk7xJ5m9v -qM3ZpJSdLNKolvldcdbUWQkthKUCArNQzHUoHI70PNZGKE6iikWoqvOv4xUq3L8J -e5Ows8fw8NY8TyaJAiHE8zOH0kUyRGaVp2+ku6qNHLFPaLk/iJjlMs1CfsdUNjNN -/N5YhwYxF7ikIhsnNXV7/AHKQeM0z5jlD74VwnquuyXc0Mgq4I99xg7nJXQNLKdU -X7thDJ8BJdKM7i8KKn/UgDoU2USIiF1x8GsqZzFR//LS9lt+n/utduEdBX7Ut0rr -nv2lQZhL4313hyzdv0f5gaEjCAndQXu/oq9SutJDAa3uszHejiyBEWgpfY7xiaTT -xf5XMTue+hbwruXLlX+H0tdH9W/BWuT7+RR3H35nKZ4FLyNG0g3joL5la3WIhRHb -9PP5hZSB6Mf1mnWuBWiJ63MJzAVsfuwyBMir8feRbj+YvI6azPXfkz874OdWnN9F -Zi6GUWy3z4UAwnC0OXO5WwH56gHfZi9u2S70Zho4jPPnF3OP2KrVJSQNrc9qwC1M -0HJNcYw9O4ERnI5OYkclEafrK98VVRPhnuKLDak31jenUh4nwNECAwEAAaN3MHUw +CgKCAgEAqnTgLxZ/eCpB46PPJqOE2IJnopLlpkK2wfp/3b7Wqskiitnr3Llw6iuk +Z0UJQJ38kIkW/UqPiyIsjEfSsRFoGhb3KofTIRd+U7Xug3wLNU1HoxUJUvXKndPk +TOaTkxaHm7wBj4oHIrGuEZGoeOpzI1BeKhhxT3xoqnuA3DjR0umMcPLwsrN1Q4O8 ++RD0xZm1sKO/nSx2rN1UfD62MFf+YW2mkjBj7UQnsgANcm5aHHj9l9osPBtOTQ+I +da1ycsJIbOJ7LhfSCTzXN6a/cLuBtAWOdgmARQf1n/TX75AcPOVJCg3fyPVTYvq2 +eSfWrYbK6cnRCzI0Sdi2oP0gPHKU3pgGKPSg6sg/WFHvGQRkj+H5AhgkGSfAlghY +sECsduqLJaDZJ+2qxC6c4fGyCYRc29BzdrE51x6VzVL7nwMTVVUeSfRzE5QyrgNX +0TXJUv1qyjo4MG3cqLNRYo73Am8+jFaxCn+a5MKavKOAW+958bdS1NmfeZFXeydG +MCjufiRlF0GBESFnv7JIE+kgn1PYPIrcBbOp7UKAl8VS1bth97eeRIeR/tCyD50r +05b4xj98+KXGLWncPoQ8ojL+9wjagPlVodRJrR+E5HVvG6kN470jLbPaClerEjx8 +SqN8VWRb+J84TjL++DaL0kf7Mjyq5cMwhacYPQtPHULLqzoGL+0CAwEAAaN3MHUw FAYDVR0RBA0wC4IJbG9jYWxob3N0MB0GA1UdJQQWMBQGCCsGAQUFBwMBBggrBgEF -BQcDAjAdBgNVHQ4EFgQUn6la/z79UFTu+LlDc6aDXG+6Tv0wHwYDVR0jBBgwFoAU -RHcTzm1u6x8WiXeAWDblHzwBt9kwDQYJKoZIhvcNAQELBQADggIBAA6sCw60Cr1V -aeFXxpzYKc3dtfKjuD6d5K6kwRkrt2AlsSfEk9fVu4SXbYeISXkL42g9nI02ce4j -o2iCeabgBT7HQVMsSx3KzlCXzXW2ACtma1D87RRQjBJinbCLSHaksZxSsMK6J+3u -MxLIgYIbxP9xGt8PLURkJq5tvJua8WZhdvaUXD1YdLANIzenCL6gHuW6WkzmHJ7E -c5rX/p8njJe7hse0ng04B9eQpuTPGUXYxOs7yMvSb5fNqZZr1EAVhBphDVjR6TuD -KTrh8vCDqHDj1xm00sbnYjzah/znmq+8XAvYGlf7DpuT68ipR914UDGvG4vKcdLz -x+3mcT3tOLfCT0VqlieWiJEdotk6EvFyubP034VxIqwr53ew2+e4m3dw39/HZ+Y1 -tggXWwlFpkZS/knLje9kz7F/EOReA4WknFSfm07B0Yv7qZNgTc/Kptw7FgPFTDLL -Cah96vwSny66C1iaRV4ALdAa1/ZNSkD/D6y1oTFGQVgy4KezjwlTA0EvmIS+wves -7jXoTSqO1iBRRl2DfHnzBtWHP1XtSTo7rqDHj6WOb/rEkTsgXqdnA5RQokj8zjLq -zaNaREfrAw55tuOASw0TbWLlv3qDofUlZyqOE6oCgCCjN/0KyqWm5m8lTUJKo6qg -HTMZ5IJXU9f1XKtMHLdGRpx0YiEGTw0e +BQcDAjAdBgNVHQ4EFgQUKTH2Ri4hNDGnL4ifUg7HEbEwhbQwHwYDVR0jBBgwFoAU +RHcTzm1u6x8WiXeAWDblHzwBt9kwDQYJKoZIhvcNAQELBQADggIBAByBbh6Mj+jp +z0Rb2vdiEV4sK0o+ad96p74ZJdiyeTLki8fLSxtKlnlrlhAzY/YFr49KQJKOzbHM +X1aoieL4Si72eprWREyNcXuD2N7tuVnw8/p3WpqW7IKBXSDDdqkdppc1B+LvTBwI ++FXSdou7dPuHgim8fHmoz/ogj+Zf1gvog3ohcnAtj9kN0zfQoBjeyjQ7v81uQ0sx +K8AO+yg/P/IWSNfzEMEGRxT91as9IrV+nmvIfe7k14ljDdJDsf+FRkea9UBOhtJW +G7cqFeWTCBV7W8bjFB0kBF9HE09E2B7hUtYZwOpVruhxXdy7WdzQzx8RWXG/bJnS +qML1bw+sdY+RtfbOr1jy8ctcAg+OmBbR0qLDQeuWlXqjTtxoHViMZpa6lNtIuD8+ +1e3+iFJ53djOSgSZ6XW163HI9353nrr1dXtlx7kdPZsb5Z3FXvL940rLiIx69ftR +dP4hP7iWstrUsrwnk6E3OmVwzc+pD8f72ztFhcqI81rmvgJ/MufGvaKoB254OibT +ng4pgs4NF2kKSFmqhXG1dTen2XRlg4ZecLrcCcotdcFX4qPEGcPjjQ4UEEaYhgFW +yWmTUWJEMO9BtqSxUFTZiQ8Ul0cJs16CyAC+oxGhaM92r7w/2xZ7fH4MHGyzJcm6 +WY7hfVHCK4+xjXMLn+k5qZYVEPUPe+0s -----END CERTIFICATE----- -----BEGIN CERTIFICATE----- MIIE/zCCAuegAwIBAgIUATFLyERaRfsQiPasMC5l0vrBMUMwDQYJKoZIhvcNAQEL diff --git a/docker/certs/server.crt b/docker/certs/server.crt index 95e4d43e4..2804eb8af 100644 --- a/docker/certs/server.crt +++ b/docker/certs/server.crt @@ -1,31 +1,31 @@ -----BEGIN CERTIFICATE----- -MIIFVDCCAzygAwIBAgIUDa7sxFwH9eBuDdGwURXDz9GkbZwwDQYJKoZIhvcNAQEL -BQAwDzENMAsGA1UEAwwEQWNtZTAeFw0yNDA2MDMxMTQwMzNaFw0yNDEyMjAxMTQw -MzNaMEAxCzAJBgNVBAYTAkRFMQ0wCwYDVQQKDARBY21lMREwDwYDVQQLDAhUaWJl +MIIFVDCCAzygAwIBAgIUDa7sxFwH9eBuDdGwURXDz9GkbZ0wDQYJKoZIhvcNAQEL +BQAwDzENMAsGA1UEAwwEQWNtZTAeFw0yNjA1MTExOTA2MzlaFw0yNzExMDIxOTA2 +MzlaMEAxCzAJBgNVBAYTAkRFMQ0wCwYDVQQKDARBY21lMREwDwYDVQQLDAhUaWJl cml1czEPMA0GA1UEAwwGc2VydmVyMIICIjANBgkqhkiG9w0BAQEFAAOCAg8AMIIC -CgKCAgEAztKC7UloJuxGMaOslWm7vEDcd8YkcC9P4PMqDTS0qgr/IXeK1LB1Pt2w -iEY4Bz/Bd3boj2IMgRzT9gjtJoD6Y3Aa32UWp1TgrDtLQ6Bns30d6sNdk7xJ5m9v -qM3ZpJSdLNKolvldcdbUWQkthKUCArNQzHUoHI70PNZGKE6iikWoqvOv4xUq3L8J -e5Ows8fw8NY8TyaJAiHE8zOH0kUyRGaVp2+ku6qNHLFPaLk/iJjlMs1CfsdUNjNN -/N5YhwYxF7ikIhsnNXV7/AHKQeM0z5jlD74VwnquuyXc0Mgq4I99xg7nJXQNLKdU -X7thDJ8BJdKM7i8KKn/UgDoU2USIiF1x8GsqZzFR//LS9lt+n/utduEdBX7Ut0rr -nv2lQZhL4313hyzdv0f5gaEjCAndQXu/oq9SutJDAa3uszHejiyBEWgpfY7xiaTT -xf5XMTue+hbwruXLlX+H0tdH9W/BWuT7+RR3H35nKZ4FLyNG0g3joL5la3WIhRHb -9PP5hZSB6Mf1mnWuBWiJ63MJzAVsfuwyBMir8feRbj+YvI6azPXfkz874OdWnN9F -Zi6GUWy3z4UAwnC0OXO5WwH56gHfZi9u2S70Zho4jPPnF3OP2KrVJSQNrc9qwC1M -0HJNcYw9O4ERnI5OYkclEafrK98VVRPhnuKLDak31jenUh4nwNECAwEAAaN3MHUw +CgKCAgEAqnTgLxZ/eCpB46PPJqOE2IJnopLlpkK2wfp/3b7Wqskiitnr3Llw6iuk +Z0UJQJ38kIkW/UqPiyIsjEfSsRFoGhb3KofTIRd+U7Xug3wLNU1HoxUJUvXKndPk +TOaTkxaHm7wBj4oHIrGuEZGoeOpzI1BeKhhxT3xoqnuA3DjR0umMcPLwsrN1Q4O8 ++RD0xZm1sKO/nSx2rN1UfD62MFf+YW2mkjBj7UQnsgANcm5aHHj9l9osPBtOTQ+I +da1ycsJIbOJ7LhfSCTzXN6a/cLuBtAWOdgmARQf1n/TX75AcPOVJCg3fyPVTYvq2 +eSfWrYbK6cnRCzI0Sdi2oP0gPHKU3pgGKPSg6sg/WFHvGQRkj+H5AhgkGSfAlghY +sECsduqLJaDZJ+2qxC6c4fGyCYRc29BzdrE51x6VzVL7nwMTVVUeSfRzE5QyrgNX +0TXJUv1qyjo4MG3cqLNRYo73Am8+jFaxCn+a5MKavKOAW+958bdS1NmfeZFXeydG +MCjufiRlF0GBESFnv7JIE+kgn1PYPIrcBbOp7UKAl8VS1bth97eeRIeR/tCyD50r +05b4xj98+KXGLWncPoQ8ojL+9wjagPlVodRJrR+E5HVvG6kN470jLbPaClerEjx8 +SqN8VWRb+J84TjL++DaL0kf7Mjyq5cMwhacYPQtPHULLqzoGL+0CAwEAAaN3MHUw FAYDVR0RBA0wC4IJbG9jYWxob3N0MB0GA1UdJQQWMBQGCCsGAQUFBwMBBggrBgEF -BQcDAjAdBgNVHQ4EFgQUn6la/z79UFTu+LlDc6aDXG+6Tv0wHwYDVR0jBBgwFoAU -RHcTzm1u6x8WiXeAWDblHzwBt9kwDQYJKoZIhvcNAQELBQADggIBAA6sCw60Cr1V -aeFXxpzYKc3dtfKjuD6d5K6kwRkrt2AlsSfEk9fVu4SXbYeISXkL42g9nI02ce4j -o2iCeabgBT7HQVMsSx3KzlCXzXW2ACtma1D87RRQjBJinbCLSHaksZxSsMK6J+3u -MxLIgYIbxP9xGt8PLURkJq5tvJua8WZhdvaUXD1YdLANIzenCL6gHuW6WkzmHJ7E -c5rX/p8njJe7hse0ng04B9eQpuTPGUXYxOs7yMvSb5fNqZZr1EAVhBphDVjR6TuD -KTrh8vCDqHDj1xm00sbnYjzah/znmq+8XAvYGlf7DpuT68ipR914UDGvG4vKcdLz -x+3mcT3tOLfCT0VqlieWiJEdotk6EvFyubP034VxIqwr53ew2+e4m3dw39/HZ+Y1 -tggXWwlFpkZS/knLje9kz7F/EOReA4WknFSfm07B0Yv7qZNgTc/Kptw7FgPFTDLL -Cah96vwSny66C1iaRV4ALdAa1/ZNSkD/D6y1oTFGQVgy4KezjwlTA0EvmIS+wves -7jXoTSqO1iBRRl2DfHnzBtWHP1XtSTo7rqDHj6WOb/rEkTsgXqdnA5RQokj8zjLq -zaNaREfrAw55tuOASw0TbWLlv3qDofUlZyqOE6oCgCCjN/0KyqWm5m8lTUJKo6qg -HTMZ5IJXU9f1XKtMHLdGRpx0YiEGTw0e +BQcDAjAdBgNVHQ4EFgQUKTH2Ri4hNDGnL4ifUg7HEbEwhbQwHwYDVR0jBBgwFoAU +RHcTzm1u6x8WiXeAWDblHzwBt9kwDQYJKoZIhvcNAQELBQADggIBAByBbh6Mj+jp +z0Rb2vdiEV4sK0o+ad96p74ZJdiyeTLki8fLSxtKlnlrlhAzY/YFr49KQJKOzbHM +X1aoieL4Si72eprWREyNcXuD2N7tuVnw8/p3WpqW7IKBXSDDdqkdppc1B+LvTBwI ++FXSdou7dPuHgim8fHmoz/ogj+Zf1gvog3ohcnAtj9kN0zfQoBjeyjQ7v81uQ0sx +K8AO+yg/P/IWSNfzEMEGRxT91as9IrV+nmvIfe7k14ljDdJDsf+FRkea9UBOhtJW +G7cqFeWTCBV7W8bjFB0kBF9HE09E2B7hUtYZwOpVruhxXdy7WdzQzx8RWXG/bJnS +qML1bw+sdY+RtfbOr1jy8ctcAg+OmBbR0qLDQeuWlXqjTtxoHViMZpa6lNtIuD8+ +1e3+iFJ53djOSgSZ6XW163HI9353nrr1dXtlx7kdPZsb5Z3FXvL940rLiIx69ftR +dP4hP7iWstrUsrwnk6E3OmVwzc+pD8f72ztFhcqI81rmvgJ/MufGvaKoB254OibT +ng4pgs4NF2kKSFmqhXG1dTen2XRlg4ZecLrcCcotdcFX4qPEGcPjjQ4UEEaYhgFW +yWmTUWJEMO9BtqSxUFTZiQ8Ul0cJs16CyAC+oxGhaM92r7w/2xZ7fH4MHGyzJcm6 +WY7hfVHCK4+xjXMLn+k5qZYVEPUPe+0s -----END CERTIFICATE----- diff --git a/docker/certs/server.key b/docker/certs/server.key index 7e60bb02e..71c4e52fd 100644 --- a/docker/certs/server.key +++ b/docker/certs/server.key @@ -1,52 +1,52 @@ -----BEGIN PRIVATE KEY----- -MIIJQgIBADANBgkqhkiG9w0BAQEFAASCCSwwggkoAgEAAoICAQDO0oLtSWgm7EYx -o6yVabu8QNx3xiRwL0/g8yoNNLSqCv8hd4rUsHU+3bCIRjgHP8F3duiPYgyBHNP2 -CO0mgPpjcBrfZRanVOCsO0tDoGezfR3qw12TvEnmb2+ozdmklJ0s0qiW+V1x1tRZ -CS2EpQICs1DMdSgcjvQ81kYoTqKKRaiq86/jFSrcvwl7k7Czx/Dw1jxPJokCIcTz -M4fSRTJEZpWnb6S7qo0csU9ouT+ImOUyzUJ+x1Q2M0383liHBjEXuKQiGyc1dXv8 -AcpB4zTPmOUPvhXCeq67JdzQyCrgj33GDucldA0sp1Rfu2EMnwEl0ozuLwoqf9SA -OhTZRIiIXXHwaypnMVH/8tL2W36f+6124R0FftS3Suue/aVBmEvjfXeHLN2/R/mB -oSMICd1Be7+ir1K60kMBre6zMd6OLIERaCl9jvGJpNPF/lcxO576FvCu5cuVf4fS -10f1b8Fa5Pv5FHcffmcpngUvI0bSDeOgvmVrdYiFEdv08/mFlIHox/Wada4FaInr -cwnMBWx+7DIEyKvx95FuP5i8jprM9d+TPzvg51ac30VmLoZRbLfPhQDCcLQ5c7lb -AfnqAd9mL27ZLvRmGjiM8+cXc4/YqtUlJA2tz2rALUzQck1xjD07gRGcjk5iRyUR -p+sr3xVVE+Ge4osNqTfWN6dSHifA0QIDAQABAoICAADFLMzFjAZPlVIWYQRYLcVd -ZDjLt4tlqLVusGSW0niq5HD3ZxBkVRZyKMf0I32m65F2Y1az27YwIVuyZDAzVSNh -Sa9U6vr97F2F1cGbZ4F2DQJInpjID+okVnkNZbLoxQZThUJVLMd5kGZBvA45N1cD -XBDb25WyJFeU6HNaWh171Y1H7arxw2xpp3dS6Sq9OxDpilVU4FgeQDOT6LzEKlQS -AfsK9dUHVUHS6Pfbz0BS6fEYzbdnRoFyatcfDJs5nx2Oj+lq2pg2zxq01sAMsJ/Y -ittWdtIn5u5OXXp3UV4PWL1/5RVZD5q/x4cY/Xs4nR5rAKB7Mz1t5xCgbr8Ro9TE -9PVzrbGy8hCWW0Yz+zhwIsDrtkQ7RGIg95W7IjaxnrjCUszK0xG1hXpce1qg1EN0 -rF4u7pU0qEWw4piLfIXepVZxVo27dOYj9qEpDkGiVYXCJ3+HifHBt5tE/rVkStF3 -dzihxyk5E7F4wJd9tz2xAMxFSgG3IeEZ3IOCxFWJib6micXZJ2n6N9uuUnHGW3D2 -o7FC02G1gXsxxgY871b8G6mFyGhmfEJxqrIvek8fBvvgOPWKnroLqJprxYow6miE -QU6yC4C/1RZgn/l6kj9jz2r6BY2nVjhHjbLGTh9bsqf5dCPdJV01FsVMiJqUzg5+ -HR5XJSf1hXRx/egBYdaBAoIBAQD3Hb12rwXRVaf38wth4VMaZr1Dxgkt0/X58LTf -SXPzGMChqnhBKdNHPv4pfWpBbvKBPWUcd+uBylgABl4xD8QH6VcspRWdgAJjul4K -RCRdWJtt0nxOqU4KitaBWOM7d6Ec3oCCaOZI5ZT+6Hj+X/RmAwd9acNM8NQ5166y -AyVQfO+2QvWRgLWxyYnBIRYkPU0L+ItkBxWpe0W8bRCj2ilAP+UCH0VSGMsnkzKw -y2HQtLGu8EBODmoW36qeYFYf6iKTMQpdtwyRYjjVq5smYSfJPy5WvdIOvcbcpI4I -Edpd1GvdjcwdfTKPiCvhDgpjQUCEOeLaKvszSFAxsSyyMFRRAoIBAQDWQfBWEwLT -jFZ9N07xkMxG4qA28KUXIHZ53DkEQmrDYQWSpJ6OfrhQgwtX9CtTMoyrG4gw1IDJ -lAcx91o6GVkC4CP8+ssvhPZi+KD9iVAI61hg3gVyxvndXgYg2xBeJ8IBm7Jkg5HK -A9tZW8jEfH+nO6HhszY0r9VNov2naRwGGZ9JgGpcMvFN5taXOhierfk3L63zaJPJ -Mx8Aaspxlk7u9ommZ1jkdpmczUzPfEpyRfSD9qoKxA4GOYPxDCUSkAyy6XzlF4rg -AKetXg5yDNa2Y4MXfbIK40Oh1wz7e9yZDjovSxonjC141RD8ybyOXhfsK67oMMME -J0gxhBR3vASBAoIBAG0jJVoVUmxxeA15ub0w1pMCbPRRshwbULdiJ3+14Q+sDudX -cmTVJAqDN5z7VsIvTcrmYpGAJPLdeqAIL/FbFSipVWbSQgmdT3DcDkxaa/UN/Rcz -rtLO0zi0uKfHqhPJcc5eNkNiMNJhErzBzy4JEtc630P0QdzpP9GMAAt+eCxkATpt -uCbawWQTrlMtWaoHqM9wpZ83wcloOBRP1tmGsFE/5tRZGzR23sJLsEeEi16xbwfj -84KFuzT+80ufIGpX7Y00S2+4OES9LHyxnYQFxJyM2tpUW0FHb1xjEJdfyyFFf54J -0ev0LzBU44wxt0S+vM+pARd5hBfSCBjqNuM7lQECggEALhpmMr9IfmjWO39pN0Wn -DyG4w9moTH+pvrMKecYo3v3Dizhs/dB6rKhmCnj50Z8w8ais94TiaX22xqOpAJNv -udStKcR1cDY2JjnFuoiPdjvd+ooLthTmsyGGRA+fSANaFaqBCmvdNRD7ZBEB9HWt -qjiEruI3KcMkLN6DokBVzWI6CkDdohU8Iz0ms8fGgG6DD8LstVGtaz/azeYsxaBI -P9dA61OVpyN2Dm2Gt6bRBiHTaYnsMQDa27AImhe46nOgp+bh/xG/yk+ZxQ5WIWht -0zU6ghWD+B/K78osevi+ERkkoASTDit1pWiDjUGDl0bb8u+7ZS8I553kRPNczB7j -AQKCAQEA9wJW7rWBuIVMUymSqynSvy4SqClOX2IKFbsJqqe3PO5dby/8YnxPXOZK -lq7gSXWfSgTN29JY5beVBLJI66spSTiz6AP4/iWQqCpzw9VM0Gv7GxIasZmfP+tp -l4JV8+yAElOFd1IhjV3RKGU1fGPGJfstIBt5eXQCSVQyQaFYQeGYE0KU5AUD6lvY -6R9irgVicVa9x1eq5HVcTVYb0gFs4zSZ1YlpqTc/i1ttZEWGyzmOK5cMX2iOeou7 -H/IZyIjtTm6edWgUANXhZdDss3gBUitLUpne579efdPCTJ4vqRjEA8tjZeGgmJpf -Oeu1HE+LelnM2vOc9TtbJC9FrC8nYw== +MIIJQgIBADANBgkqhkiG9w0BAQEFAASCCSwwggkoAgEAAoICAQCqdOAvFn94KkHj +o88mo4TYgmeikuWmQrbB+n/dvtaqySKK2evcuXDqK6RnRQlAnfyQiRb9So+LIiyM +R9KxEWgaFvcqh9MhF35Tte6DfAs1TUejFQlS9cqd0+RM5pOTFoebvAGPigcisa4R +kah46nMjUF4qGHFPfGiqe4DcONHS6Yxw8vCys3VDg7z5EPTFmbWwo7+dLHas3VR8 +PrYwV/5hbaaSMGPtRCeyAA1ybloceP2X2iw8G05ND4h1rXJywkhs4nsuF9IJPNc3 +pr9wu4G0BY52CYBFB/Wf9NfvkBw85UkKDd/I9VNi+rZ5J9athsrpydELMjRJ2Lag +/SA8cpTemAYo9KDqyD9YUe8ZBGSP4fkCGCQZJ8CWCFiwQKx26osloNkn7arELpzh +8bIJhFzb0HN2sTnXHpXNUvufAxNVVR5J9HMTlDKuA1fRNclS/WrKOjgwbdyos1Fi +jvcCbz6MVrEKf5rkwpq8o4Bb73nxt1LU2Z95kVd7J0YwKO5+JGUXQYERIWe/skgT +6SCfU9g8itwFs6ntQoCXxVLVu2H3t55Eh5H+0LIPnSvTlvjGP3z4pcYtadw+hDyi +Mv73CNqA+VWh1EmtH4TkdW8bqQ3jvSMts9oKV6sSPHxKo3xVZFv4nzhOMv74NovS +R/syPKrlwzCFpxg9C08dQsurOgYv7QIDAQABAoICADtlW4L893DzZJ9Cgtnna9CX +7C3Zux0qrQ090RV/PMUpLhitJAN3ONHYYEK96yHxi0MAChs7wnYMc/JzyoZ51skU +jI7s4lRrH9FimViGvk8V/SrmFygpzq8dWTW0uOKtnJZXNkICqkbcHBgyJb7wjytU +g2NuvfkhFEWnoHjccbzpNc9b0CSs5OUgQBaX4nsCey2weYH2runAfAKJRanl15Wy +hDL3mrJgJ+beHtFrg4ndXRxvYS+Woju2+GltBW7YpS0P5DVlBoLCiQny2E2bgPAu +aXxXBjPHuL7CrgXjtPtBOCjBOeQIHETmsPPZvnQb/pPlh6q7lT3QPp8tZPC7SoUN +t20ISZWgo74qt1gzisCNJ/GjmI0QiS96hKGw9iMYnJjZNFuUZx7oZDP3JnTFffwA +Ks9MAskUu1rxc+4H+PpGj+z4+0fq9iYEZ4EPWSreB+mt117xzlbUJlfGlmfa4O5A +srtLae/JIQJsefU+yBzXix+tPngSbiexxcYKiOsRqpoWdz+Prjq+v07pQO6IXt47 +9DDW9RtfkFDiqchoqQHfYb1p05vPqJltwTJVwiLaGBCdqkZYWZjYFIJeMNpgYhlM +2h9YdQVBdbRf554UHW1zXlmkRyD0jrI5MKmPqwl5hIFvPdhWb16V3At6Fj7Ky0VO +vzHI3QbDDj1N2pzpXSQfAoIBAQDmI8Sdumi9zaGffwHPrsTosa+btDri3fyEyeOm +3ZsbLYGos8sJdGLMdf+XvTF+4i+yw+pAtN7teUbbgLaJv00NjKA5QSp+RYKNmpWH +cMkSEQloGBZnFDqk1NRuTQ6LvQmr4Isxdg5wugFXBtvwmC7BjMzvE6pH1usubDAY +8zv2By0W63IX01WKWRkFRoSF74XoJjc0fjngW8csqhr103DihapNkSkTU6HIVvau +CvjKclkdZ2YMeGh7fNwkthA8oZcdHneeQwzJzFPFEg/juk/ggFsmB6U5U3wA1awh +ac0KWit0qN0nmZ8GAEh4KWSwu3am9yw50MvRC9AvCTdfzjU/AoIBAQC9nEDF360l +ldoGtEhiz/HuEYs/g3B3BnXvsH6YGEXpOyFt/XUNTQboBw39Csy7v1FOgJnavXHw +b3HQcIFaZNEUmZO0UgAnQxzHmGQ2gGCKYUAb4cDb85N/n2+y0Fen1jlOhvs7RlBh +atnaIfaXJ/xqcXy/5UTA5396KSPkYsE8WA34gUwG+cnldPYvy/W9gn4jeHNIrzsi +1R9kfDxcg2IqU056oR8P6PZ3tSSToxMb1Q8QtBC2FFwWpU0xDO0372DvSOcW61am +otYSoDp7FO3XmtuJl7UW5wuWZHoD86iJBcPGZnwIbEJbjo5BUF9HHZoOWIqIPGYf +f6tO9g+cm3PTAoIBADDbXQ1DGqNYuTwcAW1uo9zmg+phO7MX/1jNZ2fwWdJOOd1v +teXe8G6JimZTQuO17vxbfSqZe04c1f8ZdycNFrWOqiEdhYDjDtEzBRWIyxbryPxx +SKg/cie2CxcTgsgFrLzxYXtxnaUux8QK77xHAn4SfxsuKJMxvCHR0/AoCw2y/k6E +U2dddSZ2vcoR62ZnsBzVqBibx3uq4EDKKAkSBz//smTfMUIqGglm9N2D9Mc9uU91 +uQNiuIOmwTGF+TJ195e19R0DDP72Qr5ulDL7RaPae/851kiyQXwH4JADXwUYmWsd +wj167nierMPdvcOLOKg/hwMLIYnSoTKrGTdclo8CggEBAKyzIRwZeu986bSpiDTY +Chc4y4fyBAGlVM4YB3YoxaSFQxGXhYGz4tJ7enY72/Y1b6z83SWq35iLKTMdBfR7 +VyRYLXxUI+ee7Ruu5bfufgAMTAQZPzwXQwU/BtHriatJJ7EqqLF4fcX9OKfBv4Q1 +22ZoL6PpAxJgyG9QAW0HtdFssmzh94lzAj2IpqMqNo2Bybos/3P4hvhW/dzce24Y +DNVYQ2bWUiB/o92sk8AVDFaRXMNt/rqZGLdXoFNI3tfPpI7N7A2oFKh6MFmOrzVj +/q4eUk+kakCN+LPmmGv5Bkynf4W52schM9+InHFI7z8q6yKd6q/js3CFLFcjL10J +ChkCggEAFXZjrjb0iAF34Oel0tCsH8Vm0td7wgIOov0YZSoafRQRbBN1uFyLIjOl +5kuK5vGGHIFSr+4fsD+GsDKXf9D0NCp7E+kPKKfsS7HobDcZ/FshhxxxcmSp6KbZ +Cs2AaMwq1wW2lyQtFDxLsR7ACWfp1MvT6ZpaPE/4bVW325Bsav1qf/HmqGP82KCO +d0FesLXKZ41hJHyYENkIjXzglAL25TOum19A+8digoI6tuCeOodEuMvME6AgW0EC +NyVO+NrVA4YqkOklwLeoTrvpzSQ+TymgMKM36rcnR/zSUIIAVfa7maEmRUgN+YlG +6FJg2C5WHHHhAOWiS+gM1HBaeLSLwA== -----END PRIVATE KEY----- diff --git a/docker/certs/server.pem b/docker/certs/server.pem index 7acbb192f..4fc2c9526 100644 --- a/docker/certs/server.pem +++ b/docker/certs/server.pem @@ -1,83 +1,83 @@ -----BEGIN PRIVATE KEY----- -MIIJQgIBADANBgkqhkiG9w0BAQEFAASCCSwwggkoAgEAAoICAQDO0oLtSWgm7EYx -o6yVabu8QNx3xiRwL0/g8yoNNLSqCv8hd4rUsHU+3bCIRjgHP8F3duiPYgyBHNP2 -CO0mgPpjcBrfZRanVOCsO0tDoGezfR3qw12TvEnmb2+ozdmklJ0s0qiW+V1x1tRZ -CS2EpQICs1DMdSgcjvQ81kYoTqKKRaiq86/jFSrcvwl7k7Czx/Dw1jxPJokCIcTz -M4fSRTJEZpWnb6S7qo0csU9ouT+ImOUyzUJ+x1Q2M0383liHBjEXuKQiGyc1dXv8 -AcpB4zTPmOUPvhXCeq67JdzQyCrgj33GDucldA0sp1Rfu2EMnwEl0ozuLwoqf9SA -OhTZRIiIXXHwaypnMVH/8tL2W36f+6124R0FftS3Suue/aVBmEvjfXeHLN2/R/mB -oSMICd1Be7+ir1K60kMBre6zMd6OLIERaCl9jvGJpNPF/lcxO576FvCu5cuVf4fS -10f1b8Fa5Pv5FHcffmcpngUvI0bSDeOgvmVrdYiFEdv08/mFlIHox/Wada4FaInr -cwnMBWx+7DIEyKvx95FuP5i8jprM9d+TPzvg51ac30VmLoZRbLfPhQDCcLQ5c7lb -AfnqAd9mL27ZLvRmGjiM8+cXc4/YqtUlJA2tz2rALUzQck1xjD07gRGcjk5iRyUR -p+sr3xVVE+Ge4osNqTfWN6dSHifA0QIDAQABAoICAADFLMzFjAZPlVIWYQRYLcVd -ZDjLt4tlqLVusGSW0niq5HD3ZxBkVRZyKMf0I32m65F2Y1az27YwIVuyZDAzVSNh -Sa9U6vr97F2F1cGbZ4F2DQJInpjID+okVnkNZbLoxQZThUJVLMd5kGZBvA45N1cD -XBDb25WyJFeU6HNaWh171Y1H7arxw2xpp3dS6Sq9OxDpilVU4FgeQDOT6LzEKlQS -AfsK9dUHVUHS6Pfbz0BS6fEYzbdnRoFyatcfDJs5nx2Oj+lq2pg2zxq01sAMsJ/Y -ittWdtIn5u5OXXp3UV4PWL1/5RVZD5q/x4cY/Xs4nR5rAKB7Mz1t5xCgbr8Ro9TE -9PVzrbGy8hCWW0Yz+zhwIsDrtkQ7RGIg95W7IjaxnrjCUszK0xG1hXpce1qg1EN0 -rF4u7pU0qEWw4piLfIXepVZxVo27dOYj9qEpDkGiVYXCJ3+HifHBt5tE/rVkStF3 -dzihxyk5E7F4wJd9tz2xAMxFSgG3IeEZ3IOCxFWJib6micXZJ2n6N9uuUnHGW3D2 -o7FC02G1gXsxxgY871b8G6mFyGhmfEJxqrIvek8fBvvgOPWKnroLqJprxYow6miE -QU6yC4C/1RZgn/l6kj9jz2r6BY2nVjhHjbLGTh9bsqf5dCPdJV01FsVMiJqUzg5+ -HR5XJSf1hXRx/egBYdaBAoIBAQD3Hb12rwXRVaf38wth4VMaZr1Dxgkt0/X58LTf -SXPzGMChqnhBKdNHPv4pfWpBbvKBPWUcd+uBylgABl4xD8QH6VcspRWdgAJjul4K -RCRdWJtt0nxOqU4KitaBWOM7d6Ec3oCCaOZI5ZT+6Hj+X/RmAwd9acNM8NQ5166y -AyVQfO+2QvWRgLWxyYnBIRYkPU0L+ItkBxWpe0W8bRCj2ilAP+UCH0VSGMsnkzKw -y2HQtLGu8EBODmoW36qeYFYf6iKTMQpdtwyRYjjVq5smYSfJPy5WvdIOvcbcpI4I -Edpd1GvdjcwdfTKPiCvhDgpjQUCEOeLaKvszSFAxsSyyMFRRAoIBAQDWQfBWEwLT -jFZ9N07xkMxG4qA28KUXIHZ53DkEQmrDYQWSpJ6OfrhQgwtX9CtTMoyrG4gw1IDJ -lAcx91o6GVkC4CP8+ssvhPZi+KD9iVAI61hg3gVyxvndXgYg2xBeJ8IBm7Jkg5HK -A9tZW8jEfH+nO6HhszY0r9VNov2naRwGGZ9JgGpcMvFN5taXOhierfk3L63zaJPJ -Mx8Aaspxlk7u9ommZ1jkdpmczUzPfEpyRfSD9qoKxA4GOYPxDCUSkAyy6XzlF4rg -AKetXg5yDNa2Y4MXfbIK40Oh1wz7e9yZDjovSxonjC141RD8ybyOXhfsK67oMMME -J0gxhBR3vASBAoIBAG0jJVoVUmxxeA15ub0w1pMCbPRRshwbULdiJ3+14Q+sDudX -cmTVJAqDN5z7VsIvTcrmYpGAJPLdeqAIL/FbFSipVWbSQgmdT3DcDkxaa/UN/Rcz -rtLO0zi0uKfHqhPJcc5eNkNiMNJhErzBzy4JEtc630P0QdzpP9GMAAt+eCxkATpt -uCbawWQTrlMtWaoHqM9wpZ83wcloOBRP1tmGsFE/5tRZGzR23sJLsEeEi16xbwfj -84KFuzT+80ufIGpX7Y00S2+4OES9LHyxnYQFxJyM2tpUW0FHb1xjEJdfyyFFf54J -0ev0LzBU44wxt0S+vM+pARd5hBfSCBjqNuM7lQECggEALhpmMr9IfmjWO39pN0Wn -DyG4w9moTH+pvrMKecYo3v3Dizhs/dB6rKhmCnj50Z8w8ais94TiaX22xqOpAJNv -udStKcR1cDY2JjnFuoiPdjvd+ooLthTmsyGGRA+fSANaFaqBCmvdNRD7ZBEB9HWt -qjiEruI3KcMkLN6DokBVzWI6CkDdohU8Iz0ms8fGgG6DD8LstVGtaz/azeYsxaBI -P9dA61OVpyN2Dm2Gt6bRBiHTaYnsMQDa27AImhe46nOgp+bh/xG/yk+ZxQ5WIWht -0zU6ghWD+B/K78osevi+ERkkoASTDit1pWiDjUGDl0bb8u+7ZS8I553kRPNczB7j -AQKCAQEA9wJW7rWBuIVMUymSqynSvy4SqClOX2IKFbsJqqe3PO5dby/8YnxPXOZK -lq7gSXWfSgTN29JY5beVBLJI66spSTiz6AP4/iWQqCpzw9VM0Gv7GxIasZmfP+tp -l4JV8+yAElOFd1IhjV3RKGU1fGPGJfstIBt5eXQCSVQyQaFYQeGYE0KU5AUD6lvY -6R9irgVicVa9x1eq5HVcTVYb0gFs4zSZ1YlpqTc/i1ttZEWGyzmOK5cMX2iOeou7 -H/IZyIjtTm6edWgUANXhZdDss3gBUitLUpne579efdPCTJ4vqRjEA8tjZeGgmJpf -Oeu1HE+LelnM2vOc9TtbJC9FrC8nYw== +MIIJQgIBADANBgkqhkiG9w0BAQEFAASCCSwwggkoAgEAAoICAQCqdOAvFn94KkHj +o88mo4TYgmeikuWmQrbB+n/dvtaqySKK2evcuXDqK6RnRQlAnfyQiRb9So+LIiyM +R9KxEWgaFvcqh9MhF35Tte6DfAs1TUejFQlS9cqd0+RM5pOTFoebvAGPigcisa4R +kah46nMjUF4qGHFPfGiqe4DcONHS6Yxw8vCys3VDg7z5EPTFmbWwo7+dLHas3VR8 +PrYwV/5hbaaSMGPtRCeyAA1ybloceP2X2iw8G05ND4h1rXJywkhs4nsuF9IJPNc3 +pr9wu4G0BY52CYBFB/Wf9NfvkBw85UkKDd/I9VNi+rZ5J9athsrpydELMjRJ2Lag +/SA8cpTemAYo9KDqyD9YUe8ZBGSP4fkCGCQZJ8CWCFiwQKx26osloNkn7arELpzh +8bIJhFzb0HN2sTnXHpXNUvufAxNVVR5J9HMTlDKuA1fRNclS/WrKOjgwbdyos1Fi +jvcCbz6MVrEKf5rkwpq8o4Bb73nxt1LU2Z95kVd7J0YwKO5+JGUXQYERIWe/skgT +6SCfU9g8itwFs6ntQoCXxVLVu2H3t55Eh5H+0LIPnSvTlvjGP3z4pcYtadw+hDyi +Mv73CNqA+VWh1EmtH4TkdW8bqQ3jvSMts9oKV6sSPHxKo3xVZFv4nzhOMv74NovS +R/syPKrlwzCFpxg9C08dQsurOgYv7QIDAQABAoICADtlW4L893DzZJ9Cgtnna9CX +7C3Zux0qrQ090RV/PMUpLhitJAN3ONHYYEK96yHxi0MAChs7wnYMc/JzyoZ51skU +jI7s4lRrH9FimViGvk8V/SrmFygpzq8dWTW0uOKtnJZXNkICqkbcHBgyJb7wjytU +g2NuvfkhFEWnoHjccbzpNc9b0CSs5OUgQBaX4nsCey2weYH2runAfAKJRanl15Wy +hDL3mrJgJ+beHtFrg4ndXRxvYS+Woju2+GltBW7YpS0P5DVlBoLCiQny2E2bgPAu +aXxXBjPHuL7CrgXjtPtBOCjBOeQIHETmsPPZvnQb/pPlh6q7lT3QPp8tZPC7SoUN +t20ISZWgo74qt1gzisCNJ/GjmI0QiS96hKGw9iMYnJjZNFuUZx7oZDP3JnTFffwA +Ks9MAskUu1rxc+4H+PpGj+z4+0fq9iYEZ4EPWSreB+mt117xzlbUJlfGlmfa4O5A +srtLae/JIQJsefU+yBzXix+tPngSbiexxcYKiOsRqpoWdz+Prjq+v07pQO6IXt47 +9DDW9RtfkFDiqchoqQHfYb1p05vPqJltwTJVwiLaGBCdqkZYWZjYFIJeMNpgYhlM +2h9YdQVBdbRf554UHW1zXlmkRyD0jrI5MKmPqwl5hIFvPdhWb16V3At6Fj7Ky0VO +vzHI3QbDDj1N2pzpXSQfAoIBAQDmI8Sdumi9zaGffwHPrsTosa+btDri3fyEyeOm +3ZsbLYGos8sJdGLMdf+XvTF+4i+yw+pAtN7teUbbgLaJv00NjKA5QSp+RYKNmpWH +cMkSEQloGBZnFDqk1NRuTQ6LvQmr4Isxdg5wugFXBtvwmC7BjMzvE6pH1usubDAY +8zv2By0W63IX01WKWRkFRoSF74XoJjc0fjngW8csqhr103DihapNkSkTU6HIVvau +CvjKclkdZ2YMeGh7fNwkthA8oZcdHneeQwzJzFPFEg/juk/ggFsmB6U5U3wA1awh +ac0KWit0qN0nmZ8GAEh4KWSwu3am9yw50MvRC9AvCTdfzjU/AoIBAQC9nEDF360l +ldoGtEhiz/HuEYs/g3B3BnXvsH6YGEXpOyFt/XUNTQboBw39Csy7v1FOgJnavXHw +b3HQcIFaZNEUmZO0UgAnQxzHmGQ2gGCKYUAb4cDb85N/n2+y0Fen1jlOhvs7RlBh +atnaIfaXJ/xqcXy/5UTA5396KSPkYsE8WA34gUwG+cnldPYvy/W9gn4jeHNIrzsi +1R9kfDxcg2IqU056oR8P6PZ3tSSToxMb1Q8QtBC2FFwWpU0xDO0372DvSOcW61am +otYSoDp7FO3XmtuJl7UW5wuWZHoD86iJBcPGZnwIbEJbjo5BUF9HHZoOWIqIPGYf +f6tO9g+cm3PTAoIBADDbXQ1DGqNYuTwcAW1uo9zmg+phO7MX/1jNZ2fwWdJOOd1v +teXe8G6JimZTQuO17vxbfSqZe04c1f8ZdycNFrWOqiEdhYDjDtEzBRWIyxbryPxx +SKg/cie2CxcTgsgFrLzxYXtxnaUux8QK77xHAn4SfxsuKJMxvCHR0/AoCw2y/k6E +U2dddSZ2vcoR62ZnsBzVqBibx3uq4EDKKAkSBz//smTfMUIqGglm9N2D9Mc9uU91 +uQNiuIOmwTGF+TJ195e19R0DDP72Qr5ulDL7RaPae/851kiyQXwH4JADXwUYmWsd +wj167nierMPdvcOLOKg/hwMLIYnSoTKrGTdclo8CggEBAKyzIRwZeu986bSpiDTY +Chc4y4fyBAGlVM4YB3YoxaSFQxGXhYGz4tJ7enY72/Y1b6z83SWq35iLKTMdBfR7 +VyRYLXxUI+ee7Ruu5bfufgAMTAQZPzwXQwU/BtHriatJJ7EqqLF4fcX9OKfBv4Q1 +22ZoL6PpAxJgyG9QAW0HtdFssmzh94lzAj2IpqMqNo2Bybos/3P4hvhW/dzce24Y +DNVYQ2bWUiB/o92sk8AVDFaRXMNt/rqZGLdXoFNI3tfPpI7N7A2oFKh6MFmOrzVj +/q4eUk+kakCN+LPmmGv5Bkynf4W52schM9+InHFI7z8q6yKd6q/js3CFLFcjL10J +ChkCggEAFXZjrjb0iAF34Oel0tCsH8Vm0td7wgIOov0YZSoafRQRbBN1uFyLIjOl +5kuK5vGGHIFSr+4fsD+GsDKXf9D0NCp7E+kPKKfsS7HobDcZ/FshhxxxcmSp6KbZ +Cs2AaMwq1wW2lyQtFDxLsR7ACWfp1MvT6ZpaPE/4bVW325Bsav1qf/HmqGP82KCO +d0FesLXKZ41hJHyYENkIjXzglAL25TOum19A+8digoI6tuCeOodEuMvME6AgW0EC +NyVO+NrVA4YqkOklwLeoTrvpzSQ+TymgMKM36rcnR/zSUIIAVfa7maEmRUgN+YlG +6FJg2C5WHHHhAOWiS+gM1HBaeLSLwA== -----END PRIVATE KEY----- -----BEGIN CERTIFICATE----- -MIIFVDCCAzygAwIBAgIUDa7sxFwH9eBuDdGwURXDz9GkbZwwDQYJKoZIhvcNAQEL -BQAwDzENMAsGA1UEAwwEQWNtZTAeFw0yNDA2MDMxMTQwMzNaFw0yNDEyMjAxMTQw -MzNaMEAxCzAJBgNVBAYTAkRFMQ0wCwYDVQQKDARBY21lMREwDwYDVQQLDAhUaWJl +MIIFVDCCAzygAwIBAgIUDa7sxFwH9eBuDdGwURXDz9GkbZ0wDQYJKoZIhvcNAQEL +BQAwDzENMAsGA1UEAwwEQWNtZTAeFw0yNjA1MTExOTA2MzlaFw0yNzExMDIxOTA2 +MzlaMEAxCzAJBgNVBAYTAkRFMQ0wCwYDVQQKDARBY21lMREwDwYDVQQLDAhUaWJl cml1czEPMA0GA1UEAwwGc2VydmVyMIICIjANBgkqhkiG9w0BAQEFAAOCAg8AMIIC -CgKCAgEAztKC7UloJuxGMaOslWm7vEDcd8YkcC9P4PMqDTS0qgr/IXeK1LB1Pt2w -iEY4Bz/Bd3boj2IMgRzT9gjtJoD6Y3Aa32UWp1TgrDtLQ6Bns30d6sNdk7xJ5m9v -qM3ZpJSdLNKolvldcdbUWQkthKUCArNQzHUoHI70PNZGKE6iikWoqvOv4xUq3L8J -e5Ows8fw8NY8TyaJAiHE8zOH0kUyRGaVp2+ku6qNHLFPaLk/iJjlMs1CfsdUNjNN -/N5YhwYxF7ikIhsnNXV7/AHKQeM0z5jlD74VwnquuyXc0Mgq4I99xg7nJXQNLKdU -X7thDJ8BJdKM7i8KKn/UgDoU2USIiF1x8GsqZzFR//LS9lt+n/utduEdBX7Ut0rr -nv2lQZhL4313hyzdv0f5gaEjCAndQXu/oq9SutJDAa3uszHejiyBEWgpfY7xiaTT -xf5XMTue+hbwruXLlX+H0tdH9W/BWuT7+RR3H35nKZ4FLyNG0g3joL5la3WIhRHb -9PP5hZSB6Mf1mnWuBWiJ63MJzAVsfuwyBMir8feRbj+YvI6azPXfkz874OdWnN9F -Zi6GUWy3z4UAwnC0OXO5WwH56gHfZi9u2S70Zho4jPPnF3OP2KrVJSQNrc9qwC1M -0HJNcYw9O4ERnI5OYkclEafrK98VVRPhnuKLDak31jenUh4nwNECAwEAAaN3MHUw +CgKCAgEAqnTgLxZ/eCpB46PPJqOE2IJnopLlpkK2wfp/3b7Wqskiitnr3Llw6iuk +Z0UJQJ38kIkW/UqPiyIsjEfSsRFoGhb3KofTIRd+U7Xug3wLNU1HoxUJUvXKndPk +TOaTkxaHm7wBj4oHIrGuEZGoeOpzI1BeKhhxT3xoqnuA3DjR0umMcPLwsrN1Q4O8 ++RD0xZm1sKO/nSx2rN1UfD62MFf+YW2mkjBj7UQnsgANcm5aHHj9l9osPBtOTQ+I +da1ycsJIbOJ7LhfSCTzXN6a/cLuBtAWOdgmARQf1n/TX75AcPOVJCg3fyPVTYvq2 +eSfWrYbK6cnRCzI0Sdi2oP0gPHKU3pgGKPSg6sg/WFHvGQRkj+H5AhgkGSfAlghY +sECsduqLJaDZJ+2qxC6c4fGyCYRc29BzdrE51x6VzVL7nwMTVVUeSfRzE5QyrgNX +0TXJUv1qyjo4MG3cqLNRYo73Am8+jFaxCn+a5MKavKOAW+958bdS1NmfeZFXeydG +MCjufiRlF0GBESFnv7JIE+kgn1PYPIrcBbOp7UKAl8VS1bth97eeRIeR/tCyD50r +05b4xj98+KXGLWncPoQ8ojL+9wjagPlVodRJrR+E5HVvG6kN470jLbPaClerEjx8 +SqN8VWRb+J84TjL++DaL0kf7Mjyq5cMwhacYPQtPHULLqzoGL+0CAwEAAaN3MHUw FAYDVR0RBA0wC4IJbG9jYWxob3N0MB0GA1UdJQQWMBQGCCsGAQUFBwMBBggrBgEF -BQcDAjAdBgNVHQ4EFgQUn6la/z79UFTu+LlDc6aDXG+6Tv0wHwYDVR0jBBgwFoAU -RHcTzm1u6x8WiXeAWDblHzwBt9kwDQYJKoZIhvcNAQELBQADggIBAA6sCw60Cr1V -aeFXxpzYKc3dtfKjuD6d5K6kwRkrt2AlsSfEk9fVu4SXbYeISXkL42g9nI02ce4j -o2iCeabgBT7HQVMsSx3KzlCXzXW2ACtma1D87RRQjBJinbCLSHaksZxSsMK6J+3u -MxLIgYIbxP9xGt8PLURkJq5tvJua8WZhdvaUXD1YdLANIzenCL6gHuW6WkzmHJ7E -c5rX/p8njJe7hse0ng04B9eQpuTPGUXYxOs7yMvSb5fNqZZr1EAVhBphDVjR6TuD -KTrh8vCDqHDj1xm00sbnYjzah/znmq+8XAvYGlf7DpuT68ipR914UDGvG4vKcdLz -x+3mcT3tOLfCT0VqlieWiJEdotk6EvFyubP034VxIqwr53ew2+e4m3dw39/HZ+Y1 -tggXWwlFpkZS/knLje9kz7F/EOReA4WknFSfm07B0Yv7qZNgTc/Kptw7FgPFTDLL -Cah96vwSny66C1iaRV4ALdAa1/ZNSkD/D6y1oTFGQVgy4KezjwlTA0EvmIS+wves -7jXoTSqO1iBRRl2DfHnzBtWHP1XtSTo7rqDHj6WOb/rEkTsgXqdnA5RQokj8zjLq -zaNaREfrAw55tuOASw0TbWLlv3qDofUlZyqOE6oCgCCjN/0KyqWm5m8lTUJKo6qg -HTMZ5IJXU9f1XKtMHLdGRpx0YiEGTw0e +BQcDAjAdBgNVHQ4EFgQUKTH2Ri4hNDGnL4ifUg7HEbEwhbQwHwYDVR0jBBgwFoAU +RHcTzm1u6x8WiXeAWDblHzwBt9kwDQYJKoZIhvcNAQELBQADggIBAByBbh6Mj+jp +z0Rb2vdiEV4sK0o+ad96p74ZJdiyeTLki8fLSxtKlnlrlhAzY/YFr49KQJKOzbHM +X1aoieL4Si72eprWREyNcXuD2N7tuVnw8/p3WpqW7IKBXSDDdqkdppc1B+LvTBwI ++FXSdou7dPuHgim8fHmoz/ogj+Zf1gvog3ohcnAtj9kN0zfQoBjeyjQ7v81uQ0sx +K8AO+yg/P/IWSNfzEMEGRxT91as9IrV+nmvIfe7k14ljDdJDsf+FRkea9UBOhtJW +G7cqFeWTCBV7W8bjFB0kBF9HE09E2B7hUtYZwOpVruhxXdy7WdzQzx8RWXG/bJnS +qML1bw+sdY+RtfbOr1jy8ctcAg+OmBbR0qLDQeuWlXqjTtxoHViMZpa6lNtIuD8+ +1e3+iFJ53djOSgSZ6XW163HI9353nrr1dXtlx7kdPZsb5Z3FXvL940rLiIx69ftR +dP4hP7iWstrUsrwnk6E3OmVwzc+pD8f72ztFhcqI81rmvgJ/MufGvaKoB254OibT +ng4pgs4NF2kKSFmqhXG1dTen2XRlg4ZecLrcCcotdcFX4qPEGcPjjQ4UEEaYhgFW +yWmTUWJEMO9BtqSxUFTZiQ8Ul0cJs16CyAC+oxGhaM92r7w/2xZ7fH4MHGyzJcm6 +WY7hfVHCK4+xjXMLn+k5qZYVEPUPe+0s -----END CERTIFICATE----- diff --git a/docker/docker-azure-sql-edge.dockerfile b/docker/docker-azure-sql-edge.dockerfile index 14279c405..d4a009035 100644 --- a/docker/docker-azure-sql-edge.dockerfile +++ b/docker/docker-azure-sql-edge.dockerfile @@ -1,5 +1,9 @@ FROM mcr.microsoft.com/azure-sql-edge:latest -COPY --chmod=440 certs/server.* /certs/ -COPY --chmod=440 certs/customCA.* /certs/ +USER root +COPY certs/server.* /certs/ +RUN chmod 440 /certs/server.* +COPY certs/customCA.* /certs/ +RUN chmod 440 /certs/customCA.* COPY --chown=mssql docker-mssql.conf /var/opt/mssql/mssql.conf +USER mssql diff --git a/docker/docker-mssql-2017.dockerfile b/docker/docker-mssql-2017.dockerfile index 28a3dd4f4..ec4ccf451 100644 --- a/docker/docker-mssql-2017.dockerfile +++ b/docker/docker-mssql-2017.dockerfile @@ -1,5 +1,8 @@ FROM mcr.microsoft.com/mssql/server:2017-latest -COPY --chmod=440 certs/server.* /certs/ -COPY --chmod=440 certs/customCA.* /certs/ +USER root +COPY certs/server.* /certs/ +RUN chmod 440 /certs/server.* +COPY certs/customCA.* /certs/ +RUN chmod 440 /certs/customCA.* COPY docker-mssql.conf /var/opt/mssql/mssql.conf diff --git a/docker/docker-mssql-2019.dockerfile b/docker/docker-mssql-2019.dockerfile index 02ffdec0d..097a1d24d 100644 --- a/docker/docker-mssql-2019.dockerfile +++ b/docker/docker-mssql-2019.dockerfile @@ -1,5 +1,9 @@ FROM mcr.microsoft.com/mssql/server:2019-latest -COPY --chmod=440 certs/server.* /certs/ -COPY --chmod=440 certs/customCA.* /certs/ +USER root +COPY certs/server.* /certs/ +RUN chmod 440 /certs/server.* +COPY certs/customCA.* /certs/ +RUN chmod 440 /certs/customCA.* COPY --chown=mssql docker-mssql.conf /var/opt/mssql/mssql.conf +USER mssql diff --git a/docker/docker-mssql-2022.dockerfile b/docker/docker-mssql-2022.dockerfile index 930d3026c..aefdd64e6 100644 --- a/docker/docker-mssql-2022.dockerfile +++ b/docker/docker-mssql-2022.dockerfile @@ -1,5 +1,9 @@ FROM mcr.microsoft.com/mssql/server:2022-latest -COPY --chmod=444 certs/server.* /certs/ -COPY --chmod=444 certs/customCA.* /certs/ +USER root +COPY certs/server.* /certs/ +RUN chmod 444 /certs/server.* +COPY certs/customCA.* /certs/ +RUN chmod 444 /certs/customCA.* COPY --chown=mssql docker-mssql.conf /var/opt/mssql/mssql.conf +USER mssql diff --git a/docker/test-server.sh b/docker/test-server.sh new file mode 100755 index 000000000..a0239280c --- /dev/null +++ b/docker/test-server.sh @@ -0,0 +1,86 @@ +#!/usr/bin/env bash +# +# Start a SQL Server for the test suite, with podman or docker. +# +# ./docker/test-server.sh up # build, start, wait until it accepts connections +# ./docker/test-server.sh down # stop and remove +# ./docker/test-server.sh logs # follow the server log +# +# Then: +# +# export TIBERIUS_TEST_CONNECTION_STRING='server=tcp:localhost,1433;user=SA;password=;IntegratedSecurity=true;TrustServerCertificate=true' +# cargo test +# +# IMAGE selects the flavour; the default works on both x86_64 and arm64. +# The full SQL Server images are x86_64 only, so on an arm64 machine +# (Apple silicon) they either refuse to run or run under emulation. + +set -euo pipefail + +ENGINE="${ENGINE:-$(command -v podman >/dev/null 2>&1 && echo podman || echo docker)}" +NAME="${NAME:-tiberius-test-mssql}" +PORT="${PORT:-1433}" +IMAGE="${IMAGE:-azure-sql-edge}" +PASSWORD='' +HERE="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" + +case "${1:-up}" in + up) + echo "engine: $ENGINE image: $IMAGE port: $PORT" + "$ENGINE" build -q -f "$HERE/docker-$IMAGE.dockerfile" -t "$NAME:local" "$HERE" + "$ENGINE" rm -f "$NAME" >/dev/null 2>&1 || true + "$ENGINE" run -d --name "$NAME" \ + -e ACCEPT_EULA=Y \ + -e "MSSQL_SA_PASSWORD=$PASSWORD" \ + -e "SA_PASSWORD=$PASSWORD" \ + -p "$PORT:1433" \ + "$NAME:local" >/dev/null + + # The port opens well before the server will answer, so poll the log + # rather than the socket. + # + # The log is captured into a variable and matched there, rather than + # piped into `grep -q`. Under `set -o pipefail`, `grep -q` exits on the + # first match, the writer upstream dies of SIGPIPE, and the pipeline + # reports failure even though the match succeeded — so the wait never + # ends. + echo -n "waiting for SQL Server" + for _ in $(seq 1 120); do + logs="$("$ENGINE" logs "$NAME" 2>&1 || true)" + + case "$logs" in + *"SQL Server is now ready for client connections"*) + echo " — ready" + exit 0 + ;; + esac + + running="$("$ENGINE" ps --format '{{.Names}}' || true)" + case "$running" in + *"$NAME"*) ;; + *) + echo " — container exited:" + "$ENGINE" logs --tail 30 "$NAME" || true + exit 1 + ;; + esac + + echo -n . + sleep 2 + done + echo " — gave up; last lines:" + "$ENGINE" logs --tail 30 "$NAME" + exit 1 + ;; + down) + "$ENGINE" rm -f "$NAME" >/dev/null 2>&1 || true + echo "removed $NAME" + ;; + logs) + "$ENGINE" logs -f "$NAME" + ;; + *) + echo "usage: $0 {up|down|logs}" >&2 + exit 2 + ;; +esac diff --git a/docs/GUIDE.md b/docs/GUIDE.md new file mode 100644 index 000000000..11e5683d0 --- /dev/null +++ b/docs/GUIDE.md @@ -0,0 +1,325 @@ +# Tiberius Guide + +A practical tour of the driver. For the full API see [docs.rs](https://docs.rs/tiberius). +Every example imports the crate as `tiberius` (the package is `tiberius`; see +the [README](../README.md#installation)). + +- [Connecting](#connecting) +- [Configuration](#configuration) +- [Encryption & TLS](#encryption--tls) +- [Authentication](#authentication) +- [Querying](#querying) +- [Reading rows](#reading-rows) +- [Bulk insert](#bulk-insert) +- [Stored procedures, OUT params & TVPs](#stored-procedures-out-params--tvps) +- [Transactions](#transactions) +- [`IN (…)` lists](#in--lists) +- [Named instances (SQL Browser)](#named-instances-sql-browser) +- [Query cancellation](#query-cancellation) +- [Connection pooling](#connection-pooling) +- [Error handling](#error-handling) + +## Connecting + +Tiberius is runtime-independent: you create the `TcpStream` and hand it to the +[`Client`]. + +**Tokio** (wrap the stream with `tokio_util::compat`): + +```rust +use tiberius::{Client, Config, AuthMethod}; +use tokio::net::TcpStream; +use tokio_util::compat::TokioAsyncWriteCompatExt; + +# async fn f() -> anyhow::Result<()> { +let mut config = Config::new(); +config.host("localhost"); +config.port(1433); +config.authentication(AuthMethod::sql_server("SA", "")); +config.trust_cert(); // dev only + +let tcp = TcpStream::connect(config.get_addr()).await?; +tcp.set_nodelay(true)?; +let mut client = Client::connect(config, tcp.compat_write()).await?; +# Ok(()) } +``` + +**smol** (pass the stream directly — no compat layer): + +```rust,ignore +let tcp = smol::net::TcpStream::connect(config.get_addr()).await?; +tcp.set_nodelay(true)?; +let mut client = tiberius::Client::connect(config, tcp).await?; +``` + +## Configuration + +Build a [`Config`] fluently, or parse a connection string: + +```rust,ignore +// ADO.NET +let config = Config::from_ado_string( + "Server=tcp:localhost,1433;User Id=SA;Password=pw;Encrypt=strict;", +)?; + +// JDBC +let config = Config::from_jdbc_string( + "jdbc:sqlserver://localhost:1433;user=SA;password=pw;encrypt=true", +)?; + +// Builder +let config = Config::builder() + .host("localhost") + .port(1433) + .authentication(AuthMethod::sql_server("SA", "pw")) + .build(); +``` + +## Encryption & TLS + +TLS is on by default. Pick a backend via a feature flag (mutually exclusive): +`native-tls` (default), `rustls`, or `vendored-openssl`. + +Encryption levels (`Config::encryption`): `NotSupported`, `Off`, `Required` +(default), and **`Strict`** — TDS 8.0, TLS *before* the pre-login, with the +`tds/8.0` ALPN (requires the `tds80` feature; native-tls or rustls). + +```rust,ignore +use tiberius::EncryptionLevel; +config.encryption(EncryptionLevel::Strict); +config.hostname_in_certificate("my-sql-host"); // validate against a specific CN/SAN +config.client_certificate("client.pem", "client.key"); // mutual TLS +// or: config.client_certificate_pkcs12("client.pfx", "password"); +``` + +## Authentication + +```rust,ignore +// SQL Server login (password buffers are zeroized) +config.authentication(AuthMethod::sql_server("user", "pw")); + +// Windows integrated auth: SSPI on Windows (winauth), NTLM on Unix (sspi-rs) +config.authentication(AuthMethod::windows("user", "pw")); + +// Kerberos on Unix (integrated-auth-gssapi feature) +config.authentication(AuthMethod::Integrated); + +// Azure AD access token +config.authentication(AuthMethod::aad_token(token)); +``` + +## Querying + +From the [`Client`] when parameters are known at the call site: + +```rust,ignore +// Rows back +let stream = client.query("SELECT @P1, @P2", &[&1i32, &"foo"]).await?; + +// Rows affected +let result = client.execute("UPDATE t SET x = @P1 WHERE id = @P2", &[&1i32, &2i32]).await?; +println!("{} rows", result.total()); +``` + +For dynamic or owned parameters, use the [`Query`] object: + +```rust,ignore +use tiberius::Query; +let mut select = Query::new("SELECT @P1, @P2"); +for p in ["a", "b"] { select.bind(p); } +let stream = select.query(&mut client).await?; +``` + +## Reading rows + +A query returns a stream. Collect it, or take the first row: + +```rust,ignore +// All rows of the first result set +let rows = client.query("SELECT id, name FROM users", &[]).await? + .into_first_result().await?; + +for row in rows { + let id: i32 = row.get("id").unwrap(); + let name: &str = row.get("name").unwrap(); +} + +// Just the first row +let row = client.query("SELECT 1 AS n", &[]).await?.into_row().await?.unwrap(); +let n: i32 = row.get("n").unwrap(); +``` + +Type mappings are available via `FromSql`/`ToSql`, with optional `chrono`, +`time`, `rust_decimal`, `bigdecimal`, and `serde` support behind their features. +`Client::column_metadata()` exposes column type, size, precision/scale, +nullability and identity flags. + +## Bulk insert + +Efficiently stream many rows into a table: + +```rust,ignore +use tiberius::IntoRow; + +let mut req = client.bulk_insert("dbo.target").await?; // all columns +// or a specific column list: +// let mut req = client.bulk_insert_columns("dbo.target", &["foo", "bar"]).await?; + +for i in 0..1000i32 { + req.send(i.into_row()).await?; +} +let res = req.finalize().await?; +println!("{} rows", res.total()); +``` + +## Stored procedures, OUT params & TVPs + +Named RPC with input, output and table-valued parameters is supported via the +[`Command`] API. A TVP row type derives `TableValueRow`: + +```rust,ignore +use tiberius::{Command, TableValueRow}; + +// Fields may be `Copy` scalars (`i32`, `Numeric`, `Uuid`, …), owned columns +// (`String`, `Vec`), or borrowed `&str`/`&[u8]` (with a struct lifetime). +#[derive(TableValueRow)] +struct Item { + #[colname = "Id"] id: i32, + #[colname = "Name"] name: String, +} + +let rows = vec![ + Item { id: 1, name: "one".to_owned() }, + Item { id: 2, name: "two".to_owned() }, +]; + +let mut cmd = Command::new("dbo.InsertItems"); +cmd.bind_table_with_dbtype("items", "dbo.ItemList", rows); +let mut stream = cmd.exec(&mut client).await?; +``` + +## Transactions + +Real Transaction Manager requests (not T-SQL batches), with isolation levels: + +```rust,ignore +client.begin_transaction().await?; +// ... work ... +client.commit_transaction().await?; +// or client.rollback_transaction().await?; +``` + +## `IN (…)` lists + +Helpers make variable-length `IN` lists and the 2,100-parameter limit ergonomic — +see the `Query` docs on docs.rs for `in_clause`/parameter-expansion helpers. + +## Named instances (SQL Browser) + +On Windows, a named instance's port is resolved through SQL Browser. Enable +`sql-browser-tokio` or `sql-browser-smol` and use the `SqlBrowser` extension: + +```rust,ignore +use tiberius::SqlBrowser; +use tokio::net::TcpStream; +use tokio_util::compat::TokioAsyncWriteCompatExt; + +config.port(1434); +config.instance_name("INSTANCE"); +let tcp = TcpStream::connect_named(&config).await?; +let mut client = Client::connect(config, tcp.compat_write()).await?; +``` + +## Query cancellation + +```rust,ignore +client.cancel_query().await?; // sends a TDS Attention signal +``` + +## Connection pooling + +Pooling is delegated to the async pool crates rather than built in. This keeps +Tiberius runtime-agnostic (a pool has to pick a runtime's timers and tasks) and +lets the connection lifecycle — sizing, health checks, idle reaping — evolve +independently of the driver. Use [`bb8`](https://crates.io/crates/bb8), +[`deadpool`](https://crates.io/crates/deadpool), or +[`mobc`](https://crates.io/crates/mobc) with a small connection manager. + +Because Tiberius has no MARS (one in-flight request per connection), a pool is +also the natural way to get concurrency: check out a connection per task. + +Here is a minimal [`bb8`](https://crates.io/crates/bb8) manager over Tokio: + +```rust,ignore +use bb8::{ManageConnection, Pool}; +use tiberius::{Client, Config}; +use tokio::net::TcpStream; +use tokio_util::compat::{Compat, TokioAsyncWriteCompatExt}; + +struct TiberiusManager { + config: Config, +} + +#[async_trait::async_trait] +impl ManageConnection for TiberiusManager { + type Connection = Client>; + type Error = tiberius::error::Error; + + async fn connect(&self) -> Result { + let tcp = TcpStream::connect(self.config.get_addr()).await?; + tcp.set_nodelay(true)?; + Client::connect(self.config.clone(), tcp.compat_write()).await + } + + async fn is_valid(&self, conn: &mut Self::Connection) -> Result<(), Self::Error> { + // Cheap round-trip to confirm the connection is still alive. + conn.simple_query("SELECT 1").await?.into_row().await?; + Ok(()) + } + + fn has_broken(&self, _conn: &mut Self::Connection) -> bool { + false + } +} + +#[tokio::main] +async fn main() -> anyhow::Result<()> { + let mut config = Config::new(); + config.host("localhost"); + config.port(1433); + config.authentication(tiberius::AuthMethod::sql_server("SA", "")); + config.trust_cert(); // don't do this in production + + let pool = Pool::builder() + .max_size(16) + .build(TiberiusManager { config }) + .await?; + + // Check out a connection; it returns to the pool on drop. + let mut conn = pool.get().await?; + let row = conn + .query("SELECT @P1 AS n", &[&1i32]) + .await? + .into_row() + .await? + .unwrap(); + assert_eq!(Some(1i32), row.get("n")); + Ok(()) +} +``` + +The [`deadpool`](https://crates.io/crates/deadpool) and +[`mobc`](https://crates.io/crates/mobc) integrations follow the same shape — +implement their manager trait with the `connect`/`is_valid` logic above. + +## Error handling + +All fallible calls return `tiberius::Result` (`Err` is [`tiberius::error::Error`]). +Server-side errors surface as `Error::Server(TokenError { .. })` carrying the SQL +Server error code, state, class, message, procedure and line. + +[`Client`]: https://docs.rs/tiberius/latest/tiberius/struct.Client.html +[`Config`]: https://docs.rs/tiberius/latest/tiberius/struct.Config.html +[`Query`]: https://docs.rs/tiberius/latest/tiberius/struct.Query.html +[`Command`]: https://docs.rs/tiberius/latest/tiberius/struct.Command.html +[`tiberius::error::Error`]: https://docs.rs/tiberius/latest/tiberius/error/enum.Error.html diff --git a/docs/TDS_COMPATIBILITY.md b/docs/TDS_COMPATIBILITY.md new file mode 100644 index 000000000..2349b86aa --- /dev/null +++ b/docs/TDS_COMPATIBILITY.md @@ -0,0 +1,98 @@ +# TDS Protocol Compatibility + +This document tracks `tiberius`'s coverage of the Microsoft Tabular Data +Stream (MS-TDS) protocol. It is a maintained snapshot — code references may +drift; treat the ratings as the source of truth and file an issue for +discrepancies. + +Legend: ✅ full · 🟡 partial · ❌ missing · N/A not applicable. + +## Protocol-version support + +| TDS version | SQL Server | Rating | Notes | +|---|---|---|---| +| **8.0** | 2022 | ✅ Full | TLS-before-prelogin "strict" mode (`EncryptionLevel::Strict`), `tds/8.0` ALPN, and client-certificate (mutual-TLS) login on native-tls + rustls. 8.0 reuses 7.4 tokens over mandatory TLS; the only backend caveat is that opentls cannot advertise ALPN (see below). | +| **7.4** | 2012–2019 | ✅ Full | Login (`FeatureLevel::SqlServerN`), routing ENVCHANGE, `fReadOnlyIntent`, FedAuth prelogin option + FeatureExt, FEATUREEXTACK, FEDAUTHINFO (0xEE), SESSIONSTATE (0xE4). The only deferral is *transparent reconnect* — SESSIONSTATE is decoded and stored, but not yet replayed to silently re-establish a dropped session. | +| **7.3 A/B** | 2008 / R2 | ✅ Full (`tds73`) | date / time / datetime2 / datetimeoffset types, NBCROW. | +| **7.2** | 2005 | ✅ Full | PLP / varchar(max), XML, MARS / transaction-descriptor headers, SQL_VARIANT (read + write), UDT (0xF0) raw-value decode. | +| **7.1** | 2000 | ✅ Full | Collation, UCS-2 strings, `n`-prefixed var-len types, LOGIN7 layout. | +| **7.0** | 7.0 | ❌ None | Legacy fixed non-nullable types unsupported; the client only negotiates 7.4. Not a target. | + +There is **no Microsoft "TDS 6.0"** — the Microsoft protocol line is +7.0 → 7.1 → 7.2 → 7.3A → 7.3B → 7.4 → 8.0. Sybase-era TDS 4.2/5.0 predate +MS-TDS and are a separate protocol lineage, out of scope for this driver. +TDS 8.0 has no distinct LOGIN7 version — it reuses 7.4 tokens over mandatory +TLS, so "8.0" is a transport/ALPN distinction. + +**Summary:** TDS 7.1 through 8.0 are fully supported. The only remaining +elements are optional and rarely used: transparent session recovery (the +SESSIONSTATE token is decoded and stored, just not replayed on reconnect), +CLR UDT *object* deserialization (raw bytes are surfaced), and `tds/8.0` ALPN +on the opentls backend (an upstream-crate limitation) — none of which any +standard query, bulk-load, RPC, or transaction path depends on. + +## Feature matrix + +### Client → server messages +| Message | Status | +|---|---| +| PRELOGIN (0x12) | ✅ (VERSION/ENCRYPTION/INSTOPT/THREADID/MARS emitted; server options decoded and INSTOPT validated) | +| LOGIN7 (0x10) | ✅ | +| SQL Batch (0x01) | ✅ | +| RPC request (0x03) | ✅ by-ID procs and named procs, incl. OUT params and table-valued parameters (#328) | +| Bulk load (0x07) | ✅ whole-table and column-list | +| SSPI (0x11) | ✅ | +| FedAuth token | ✅ via LOGIN7 FeatureExt | +| Attention / cancel (0x06) | ✅ (`Client::cancel_query`) | +| Transaction Manager request (0x14) | ✅ (`begin`/`commit`/`rollback` with isolation levels) | + +### Server → client tokens +| Token | Status | +|---|---| +| COLMETADATA, ROW, NBCROW, DONE/PROC/INPROC | ✅ | +| ENVCHANGE, ERROR/INFO, LOGINACK | ✅ | +| RETURNVALUE, RETURNSTATUS, ORDER, SSPI | ✅ | +| FEATUREEXTACK | ✅ | +| ALTMETADATA / ALTROW (compute-by) | ✅ | +| COLINFO (0xA5) | ✅ | +| TABNAME (0xA4) | ✅ | +| FEDAUTHINFO (0xEE) | ✅ (STSURL + SPN surfaced for AAD flows) | +| SESSIONSTATE (0xE4) | ✅ decoded + stored (transparent-reconnect replay not yet implemented) | + +### Data types +| Type | Status | +|---|---| +| Fixed-len ints/bit/float/money/datetime(4) | ✅ | +| Nullable var-len (Intn/Bitn/Floatn/Guid/Money/Datetimen/Decimaln/Numericn) | ✅ | +| Char/binary + collation, Text/NText/Image | ✅ | +| PLP (max types), XML | ✅ | +| date / time / datetime2 / datetimeoffset (7.3) | ✅ (`tds73`) | +| Numeric / Decimal (incl. automatic scale rescaling of params) | ✅ | +| SQL_VARIANT (0x62) | ✅ read + write | +| UDT (0xF0) | ✅ raw-value decode (surfaced as bytes; CLR object deserialization out of scope) | + +### Encryption & auth +| Feature | Status | +|---|---| +| Encryption NotSupported / Off / On / Required | ✅ | +| Strict (TDS 8.0, TLS-first) | ✅ (`tds80`) | +| `tds/8.0` ALPN | 🟡 native-tls + rustls ✅, opentls ❌ (backend cannot advertise ALPN) | +| ENCRYPT_CLIENT_CERT (mutual TLS) | ✅ (`Config::client_certificate` / `client_certificate_pkcs12`) | +| SQL auth (zeroized) | ✅ | +| Windows NTLM/SSPI (Windows `winauth`; Unix `sspi-rs`) | ✅ | +| Kerberos/GSSAPI (Unix `integrated-auth-gssapi`) | ✅ | +| AAD / federated token | ✅ (`AuthMethod::aad_token`) | + +## Remaining optional items + +None of the following block standard operation; they are tracked for +completeness and would be additive, backward-compatible features: + +- **Transparent session recovery** — the SESSIONSTATE token is already decoded + and stored; replaying it to silently re-establish a dropped connection is the + remaining step. +- **CLR UDT object deserialization** — UDT values are decoded to raw bytes; + interpreting them into CLR objects (e.g. `geometry`) is left to the caller. +- **opentls `tds/8.0` ALPN** — verified infeasible with `opentls` 0.2.1's public + API (no ALPN setter; the wrapped `SslConnector` is private). Use native-tls or + rustls for TDS 8.0 strict mode. Tracked upstream against the `opentls` crate. diff --git a/examples/aad-auth.rs b/examples/aad-auth.rs index 8ef41c472..ebc23165b 100644 --- a/examples/aad-auth.rs +++ b/examples/aad-auth.rs @@ -8,41 +8,32 @@ //! - CLIENT_SECRET: service principal secret; //! - TENANT_ID: tenant id of service principal and sql instance; //! - SERVER: SQL server URI -use azure_identity::client_credentials_flow; -use oauth2::{ClientId, ClientSecret}; -use std::{env, sync::Arc}; +use azure_core::credentials::{Secret, TokenCredential}; +use azure_identity::ClientSecretCredential; +use std::env; use tiberius::{AuthMethod, Client, Config, Query}; use tokio::net::TcpStream; use tokio_util::compat::TokioAsyncWriteCompatExt; #[tokio::main] async fn main() -> anyhow::Result<()> { - // following code will retrive token with AAD Service Principal Auth - let client_id = - ClientId::new(env::var("CLIENT_ID").expect("Missing CLIENT_ID environment variable.")); - let client_secret = ClientSecret::new( - env::var("CLIENT_SECRET").expect("Missing CLIENT_SECRET environment variable."), - ); + let client_id = env::var("CLIENT_ID").expect("Missing CLIENT_ID environment variable."); + let client_secret = + env::var("CLIENT_SECRET").expect("Missing CLIENT_SECRET environment variable."); let tenant_id = env::var("TENANT_ID").expect("Missing TENANT_ID environment variable."); - let client = Arc::new(reqwest::Client::new()); - // This will give you the final token to use in authorization. - let token = client_credentials_flow::perform( - client, - &client_id, - &client_secret, - &["https://management.azure.com/"], - &tenant_id, - ) - .await?; + let credential = + ClientSecretCredential::new(&tenant_id, client_id, Secret::new(client_secret), None)?; + + let token = credential + .get_token(&["https://database.windows.net/.default"], None) + .await?; let mut config = Config::new(); let server = env::var("SERVER").expect("Missing SERVER environment variable."); config.host(server); config.port(1433); - config.authentication(AuthMethod::AADToken( - token.access_token().secret().to_string(), - )); + config.authentication(AuthMethod::aad_token(token.token.secret())); config.trust_cert(); let tcp = TcpStream::connect(config.get_addr()).await?; diff --git a/examples/async-std.rs b/examples/async-std.rs deleted file mode 100644 index 88fcf1c8d..000000000 --- a/examples/async-std.rs +++ /dev/null @@ -1,50 +0,0 @@ -use async_std::net::TcpStream; -use once_cell::sync::Lazy; -use std::env; -use tiberius::{Client, Config}; - -static CONN_STR: Lazy = Lazy::new(|| { - env::var("TIBERIUS_TEST_CONNECTION_STRING").unwrap_or_else(|_| { - "server=tcp:localhost,1433;IntegratedSecurity=true;TrustServerCertificate=true".to_owned() - }) -}); - -#[cfg(not(all(windows, feature = "sql-browser-async-std")))] -#[async_std::main] -async fn main() -> anyhow::Result<()> { - let config = Config::from_ado_string(&CONN_STR)?; - - let tcp = TcpStream::connect(config.get_addr()).await?; - tcp.set_nodelay(true)?; - - let mut client = Client::connect(config, tcp).await?; - - let stream = client.query("SELECT @P1", &[&1i32]).await?; - let row = stream.into_row().await?.unwrap(); - - println!("{:?}", row); - assert_eq!(Some(1), row.get(0)); - - Ok(()) -} - -#[cfg(all(windows, feature = "sql-browser-async-std"))] -#[async_std::main] -async fn main() -> anyhow::Result<()> { - use tiberius::SqlBrowser; - - let config = Config::from_ado_string(&CONN_STR)?; - - let tcp = TcpStream::connect_named(&config).await?; - tcp.set_nodelay(true)?; - - let mut client = Client::connect(config, tcp).await?; - - let stream = client.query("SELECT @P1", &[&1i32]).await?; - let row = stream.into_row().await?.unwrap(); - - println!("{:?}", row); - assert_eq!(Some(1), row.get(0)); - - Ok(()) -} diff --git a/examples/named-pipes.rs b/examples/named-pipes.rs new file mode 100644 index 000000000..fb135dbb2 --- /dev/null +++ b/examples/named-pipes.rs @@ -0,0 +1,43 @@ +//! Connecting to SQL Server over a Windows named pipe. +//! +//! SQL Server exposes a named pipe endpoint (by default +//! `\\.\pipe\sql\query` for the default instance). Because a named pipe +//! implements `AsyncRead`/`AsyncWrite`, it can be handed to +//! [`Client::connect`] exactly like a TCP stream once it is wrapped with the +//! `tokio-util` compatibility layer. +//! +//! Named pipes are a Windows-only transport, so the real example is compiled +//! only on Windows; on other platforms `main` panics with an unsupported +//! message. See tiberius issues #131 and #53 for background. + +#[cfg(windows)] +#[tokio::main] +async fn main() -> anyhow::Result<()> { + use tiberius::{AuthMethod, Client, Config}; + use tokio::net::windows::named_pipe::ClientOptions; + use tokio_util::compat::TokioAsyncWriteCompatExt; + + // The default named pipe for a default SQL Server instance. A named + // instance uses `\\.\pipe\MSSQL$\sql\query`. + const PIPE_NAME: &str = r"\\.\pipe\sql\query"; + + let mut config = Config::new(); + config.authentication(AuthMethod::Integrated); + config.trust_cert(); + + let pipe = ClientOptions::new().open(PIPE_NAME)?; + let mut client = Client::connect(config, pipe.compat_write()).await?; + + let stream = client.query("SELECT @P1", &[&1i32]).await?; + let row = stream.into_row().await?.unwrap(); + + println!("{row:?}"); + assert_eq!(Some(1), row.get(0)); + + Ok(()) +} + +#[cfg(not(windows))] +fn main() { + panic!("Named pipe connections are only supported on Windows."); +} diff --git a/runtimes-macro/Cargo.toml b/runtimes-macro/Cargo.toml index 6bf114b2a..4b0f68311 100644 --- a/runtimes-macro/Cargo.toml +++ b/runtimes-macro/Cargo.toml @@ -1,14 +1,19 @@ +# Test-only helper crate (dev-dependency, path-only): the `#[test_on_runtimes]` +# attribute that runs each integration test on tokio and smol. Not shipped, not +# published. `license`/`publish` set so the cargo-deny license gate is satisfied +# and it can never be accidentally published. [package] name = "runtimes-macro" version = "0.1.0" authors = ["Eric Sheppard "] -edition = "2018" +edition = "2021" +license = "MIT OR Apache-2.0" +publish = false [lib] proc-macro = true [dependencies] quote = "1" -syn = "1" -darling = "0.14" +syn = { version = "2", features = ["full"] } proc-macro2 = "1" diff --git a/runtimes-macro/src/lib.rs b/runtimes-macro/src/lib.rs index cc1d2cabc..94a6b0464 100644 --- a/runtimes-macro/src/lib.rs +++ b/runtimes-macro/src/lib.rs @@ -1,50 +1,70 @@ +//! Internal test-only proc-macro for tiberius. +//! +//! `#[test_on_runtimes]` takes an `async fn(client) -> Result<()>` and generates +//! one integration test per supported async runtime, so every test proves the +//! (runtime-independent) driver works on each of them. Currently: **tokio** and +//! **smol**. extern crate proc_macro; -use darling::FromMeta; -#[derive(Debug, FromMeta)] -struct MacroArgs { - #[darling(default)] - connection_string: Option, +use proc_macro::TokenStream; +use quote::{format_ident, quote}; +use syn::{parse_macro_input, ItemFn, LitStr}; + +/// Optional `connection_string = "IDENT"` attribute argument naming the `&str` +/// constant to connect with. Defaults to `CONN_STR`. +struct Args { + conn_str: String, } -#[proc_macro_attribute] -pub fn test_on_runtimes( - args: proc_macro::TokenStream, - input: proc_macro::TokenStream, -) -> proc_macro::TokenStream { - let attr_args = syn::parse_macro_input!(args as syn::AttributeArgs); - - let args = match MacroArgs::from_list(&attr_args) { - Ok(v) => v, - Err(e) => { - return proc_macro::TokenStream::from(e.write_errors()); - } - }; +impl syn::parse::Parse for Args { + fn parse(input: syn::parse::ParseStream<'_>) -> syn::Result { + let mut conn_str = String::from("CONN_STR"); - let func = syn::parse_macro_input!(input as syn::ItemFn); + if !input.is_empty() { + let ident: syn::Ident = input.parse()?; + if ident != "connection_string" { + return Err(syn::Error::new( + ident.span(), + "expected `connection_string = \"...\"`", + )); + } + input.parse::()?; + let lit: LitStr = input.parse()?; + conn_str = lit.value(); + } - let conn_str_ident_str = args.connection_string.unwrap_or_else(|| "CONN_STR".into()); + Ok(Args { conn_str }) + } +} - let conn_str_ident = - proc_macro2::Ident::new(&conn_str_ident_str, proc_macro2::Span::call_site()); +#[proc_macro_attribute] +pub fn test_on_runtimes(args: TokenStream, input: TokenStream) -> TokenStream { + let args = parse_macro_input!(args as Args); + let func = parse_macro_input!(input as ItemFn); + let conn_str_ident = format_ident!("{}", args.conn_str); let func_name = func.sig.ident.clone(); - let async_std_test = quote::format_ident!("{}_{}", func_name, "async_std"); - let tokio_test = quote::format_ident!("{}_{}", func_name, "tokio"); + let tokio_test = format_ident!("{}_tokio", func_name); + let smol_test = format_ident!("{}_smol", func_name); - let tokens = quote::quote! { + let tokens = quote! { #func #[test] - fn #async_std_test()-> Result<()> { + fn #tokio_test() -> Result<()> { LOGGER_SETUP.call_once(|| { - env_logger::init(); + let _ = env_logger::builder().is_test(true).try_init(); }); - async_std::task::block_on(async { + + use tokio_util::compat::TokioAsyncWriteCompatExt; + + let rt = tokio::runtime::Runtime::new()?; + + rt.block_on(async { let config = tiberius::Config::from_ado_string(&#conn_str_ident)?; - let tcp = async_std::net::TcpStream::connect(config.get_addr()).await?; + let tcp = tokio::net::TcpStream::connect(config.get_addr()).await?; tcp.set_nodelay(true)?; - let mut client = tiberius::Client::connect(config, tcp).await?; + let client = tiberius::Client::connect(config, tcp.compat_write()).await?; #func_name(client).await?; Ok(()) @@ -52,19 +72,16 @@ pub fn test_on_runtimes( } #[test] - fn #tokio_test()-> Result<()> { + fn #smol_test() -> Result<()> { LOGGER_SETUP.call_once(|| { - env_logger::init(); + let _ = env_logger::builder().is_test(true).try_init(); }); - use tokio_util::compat::TokioAsyncWriteCompatExt; - - let mut rt = tokio::runtime::Runtime::new()?; - rt.block_on(async { + smol::block_on(async { let config = tiberius::Config::from_ado_string(&#conn_str_ident)?; - let tcp = tokio::net::TcpStream::connect(config.get_addr()).await?; + let tcp = smol::net::TcpStream::connect(config.get_addr()).await?; tcp.set_nodelay(true)?; - let mut client = tiberius::Client::connect(config, tcp.compat_write()).await?; + let client = tiberius::Client::connect(config, tcp).await?; #func_name(client).await?; Ok(()) @@ -72,5 +89,5 @@ pub fn test_on_runtimes( } }; - proc_macro::TokenStream::from(tokens) + TokenStream::from(tokens) } diff --git a/src/bulk_options.rs b/src/bulk_options.rs new file mode 100644 index 000000000..45da4b7a5 --- /dev/null +++ b/src/bulk_options.rs @@ -0,0 +1,64 @@ +//! Options controlling the behaviour of a bulk-insert (`INSERT BULK`) +//! operation, mirroring the knobs exposed by `SqlBulkCopyOptions` in +//! ADO.NET's `SqlBulkCopy`. + +use enumflags2::{bitflags, BitFlags}; + +/// A single bulk-copy option. Combine several with the `|` operator to build a +/// [`SqlBulkCopyOptions`] set, e.g. +/// `SqlBulkCopyOption::KeepIdentity | SqlBulkCopyOption::TableLock`. +/// +/// See the MS docs for `SqlBulkCopyOptions`: +/// +#[bitflags] +#[repr(u32)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub enum SqlBulkCopyOption { + /// Preserve source identity values. When not specified, identity values are + /// assigned by the destination. Implemented by keeping the identity column + /// in the bulk column list (the `INSERT BULK` `WITH (...)` grammar has no + /// `KEEP_IDENTITY` keyword), matching ADO.NET's `SqlBulkCopy`. + KeepIdentity = 1 << 0, + /// Check constraints while data is being inserted. By default, constraints + /// are not checked. Emits `CHECK_CONSTRAINTS`. + CheckConstraints = 1 << 1, + /// Obtain a bulk update lock for the duration of the bulk-copy operation. + /// When not specified, row locks are used. Emits `TABLOCK`. + TableLock = 1 << 2, + /// Preserve null values in the destination table regardless of the settings + /// for default values. When not specified, null values are replaced by + /// default values where applicable. Emits `KEEP_NULLS`. + KeepNulls = 1 << 3, + /// Cause the server to fire the insert triggers for the rows being inserted + /// into the database. Emits `FIRE_TRIGGERS`. + FireTriggers = 1 << 4, +} + +/// A set of [`SqlBulkCopyOption`] flags controlling an `INSERT BULK`. +/// +/// Build one by combining flags with `|`, or use [`SqlBulkCopyOptions::empty`] +/// (also the [`Default`]) for no options: +/// +/// ``` +/// # use tiberius::{SqlBulkCopyOption, SqlBulkCopyOptions}; +/// let opts: SqlBulkCopyOptions = +/// SqlBulkCopyOption::KeepIdentity | SqlBulkCopyOption::TableLock; +/// assert!(opts.contains(SqlBulkCopyOption::KeepIdentity)); +/// assert!(SqlBulkCopyOptions::empty().is_empty()); +/// ``` +pub type SqlBulkCopyOptions = BitFlags; + +/// The sort order of a column, used as a bulk-insert `ORDER (...)` hint. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub enum SortOrder { + /// Ascending order (`ASC`). + Ascending, + /// Descending order (`DESC`). + Descending, +} + +/// An order hint for a bulk insert: the column name and the [`SortOrder`] the +/// incoming rows are already sorted by. Passed as a slice, e.g. +/// `&[("id", SortOrder::Ascending)]`, these become the `ORDER (...)` clause of +/// the `INSERT BULK` statement. +pub type ColumnOrderHint<'a> = (&'a str, SortOrder); diff --git a/src/client.rs b/src/client.rs index 688721d10..3aaf4b0d0 100644 --- a/src/client.rs +++ b/src/client.rs @@ -1,4 +1,4 @@ -mod auth; +pub(crate) mod auth; mod config; mod connection; @@ -14,6 +14,8 @@ pub use auth::*; pub use config::*; pub(crate) use connection::*; +use crate::bulk_options::{ColumnOrderHint, SortOrder, SqlBulkCopyOption, SqlBulkCopyOptions}; +use crate::tds::codec::RpcValue; use crate::tds::stream::ReceivedToken; use crate::{ result::ExecuteResult, @@ -21,9 +23,12 @@ use crate::{ codec::{self, IteratorJoin}, stream::{QueryStream, TokenStream}, }, - BulkLoadRequest, ColumnFlag, SqlReadBytes, ToSql, + BulkLoadRequest, MetaDataColumn, SqlReadBytes, ToSql, +}; +use codec::{ + BatchRequest, ColumnData, IsolationLevel, PacketHeader, RpcParam, RpcProcId, TokenRpcRequest, + TransactionManagerRequest, }; -use codec::{BatchRequest, ColumnData, PacketHeader, RpcParam, RpcProcId, TokenRpcRequest}; use enumflags2::BitFlags; use futures_util::io::{AsyncRead, AsyncWrite}; use futures_util::stream::TryStreamExt; @@ -57,6 +62,19 @@ use std::{borrow::Cow, fmt::Debug}; /// # } /// ``` /// +/// # Cancellation safety +/// +/// A single [`Client`] drives one connection and one request at a time. If a +/// `query`/`execute`/`simple_query` future — or the result stream it returns — +/// is dropped before the request has been sent in full and the response fully +/// consumed (for example under a `tokio::time::timeout` or a `select!` branch +/// that loses the race), the connection may be left mid-message and out of sync +/// with the server. A cancelled *write* is detected and any further use of that +/// connection fails cleanly; a result stream dropped mid-response cannot be +/// recovered. In both cases the safe course is to drop the `Client` and open a +/// new connection (a connection pool should discard the connection on error) +/// rather than reuse it. +/// /// [`Config`]: struct.Config.html #[derive(Debug)] pub struct Client { @@ -68,6 +86,11 @@ impl Client { /// options required to connect to the database using an established /// tcp connection /// + /// Note: `tcp_stream` is a connected stream, so some parts of the `Config` + /// (such as multi-subnet failover, which selects between resolved + /// addresses) must be handled while establishing that stream, outside of + /// this constructor. + /// /// [`Config`]: struct.Config.html pub async fn connect(config: Config, tcp_stream: S) -> crate::Result> { Ok(Client { @@ -89,6 +112,12 @@ impl Client { /// This API is not quite suitable for dynamic query parameters. In these /// cases using a [`Query`] object might be easier. /// + /// # Errors + /// + /// Returns an error if the statement cannot be sent, if the server reports + /// an error while executing it, or if the connection fails during the + /// request. + /// /// # Example /// /// ```no_run @@ -148,9 +177,16 @@ impl Client { /// if fighting too much with the compiler, using a [`Query`] object might be /// easier. /// + /// # Errors + /// + /// Returns an error if the statement cannot be sent, if the server reports + /// an error while executing it, or if the connection fails during the + /// request. Note that per-statement server errors may instead surface while + /// consuming the returned [`QueryStream`]. + /// /// # Example /// - /// ``` + /// ```no_run /// # use tiberius::Config; /// # use tokio_util::compat::TokioAsyncWriteCompatExt; /// # use std::env; @@ -202,9 +238,16 @@ impl Client { /// Execute multiple queries, delimited with `;` and return multiple result /// sets; one for each query. /// + /// # Errors + /// + /// Returns an error if the batch cannot be sent, if the server reports an + /// error while executing it, or if the connection fails during the request. + /// Per-statement server errors may instead surface while consuming the + /// returned [`QueryStream`]. + /// /// # Example /// - /// ``` + /// ```no_run /// # use tiberius::Config; /// # use tokio_util::compat::TokioAsyncWriteCompatExt; /// # use std::env; @@ -251,13 +294,38 @@ impl Client { Ok(result) } - /// Execute a `BULK INSERT` statement, efficiantly storing a large number of + /// Execute a `BULK INSERT` statement, efficiently storing a large number of /// rows to a specified table. Note: make sure the input row follows the same /// schema as the table, otherwise calling `send()` will return an error. /// + /// This is equivalent to `bulk_insert_columns(table, &["*"])`, inserting into + /// all of a table's columns. + /// + /// # Security + /// + /// `table` is interpolated **directly** into the SQL batch sent to the + /// server. SQL Server does not allow table (or column) identifiers to be + /// supplied as bound parameters, so this value cannot be parameterized — it + /// becomes part of the SQL text verbatim. The caller MUST therefore pass a + /// **trusted, hard-coded or otherwise validated** identifier and MUST NOT + /// pass untrusted or user-supplied input, which would open a SQL injection + /// vector. As cheap defense-in-depth this method rejects obviously-malformed + /// identifiers (NUL/ASCII control characters or an unbalanced `]` bracket), + /// but that guard is not a substitute for passing trusted input. + /// + /// # Errors + /// + /// Returns an error if `table` is a malformed identifier, if the query for + /// the column metadata and collations fails, or if the server rejects the + /// `INSERT BULK` statement. Row-level failures surface later from [`send`] and + /// [`finalize`] on the returned request. + /// + /// [`send`]: BulkLoadRequest::send + /// [`finalize`]: BulkLoadRequest::finalize + /// /// # Example /// - /// ``` + /// ```no_run /// # use tiberius::{Config, IntoRow}; /// # use tokio_util::compat::TokioAsyncWriteCompatExt; /// # use std::env; @@ -300,41 +368,225 @@ impl Client { &'a mut self, table: &'a str, ) -> crate::Result> { - // Start the bulk request - self.connection.flush_stream().await?; - - // retrieve column metadata from server - let query = format!("SELECT TOP 0 * FROM {}", table); - - let req = BatchRequest::new(query, self.connection.context().transaction_descriptor()); + self.bulk_insert_columns(table, &["*"]).await + } - let id = self.connection.context_mut().next_packet_id(); - self.connection.send(PacketHeader::batch(id), req).await?; + /// Execute a `BULK INSERT` statement, efficiently storing a large number of + /// rows to a specified table. Note: make sure the input row follows the same + /// schema as the column list, otherwise calling `send()` will return an error. + /// + /// # Security + /// + /// Both `table` and the entries of `columns` are interpolated **directly** + /// into the SQL batches sent to the server (the `SELECT` used to fetch + /// column metadata and the `INSERT BULK` statement). SQL Server does not + /// allow identifiers to be supplied as bound parameters, so these values + /// cannot be parameterized — they become part of the SQL text verbatim. The + /// caller MUST therefore pass **trusted, hard-coded or otherwise validated** + /// identifiers and MUST NOT pass untrusted or user-supplied input, which + /// would open a SQL injection vector. As cheap defense-in-depth this method + /// rejects an obviously-malformed `table` (NUL/ASCII control characters or an + /// unbalanced `]` bracket), but that guard is not a substitute for passing + /// trusted input. + /// + /// # Errors + /// + /// Returns an error if `table` is a malformed identifier, if the query for + /// the column metadata and collations fails, or if the server rejects the + /// `INSERT BULK` statement. Row-level failures surface later from [`send`] and + /// [`finalize`] on the returned request. + /// + /// [`send`]: BulkLoadRequest::send + /// [`finalize`]: BulkLoadRequest::finalize + /// + /// # Example + /// + /// ```no_run + /// # use tiberius::{Config, IntoRow}; + /// # use tokio_util::compat::TokioAsyncWriteCompatExt; + /// # use std::env; + /// # #[tokio::main] + /// # async fn main() -> Result<(), Box> { + /// # let c_str = env::var("TIBERIUS_TEST_CONNECTION_STRING").unwrap_or( + /// # "server=tcp:localhost,1433;integratedSecurity=true;TrustServerCertificate=true".to_owned(), + /// # ); + /// # let config = Config::from_ado_string(&c_str)?; + /// # let tcp = tokio::net::TcpStream::connect(config.get_addr()).await?; + /// # tcp.set_nodelay(true)?; + /// # let mut client = tiberius::Client::connect(config, tcp.compat_write()).await?; + /// let create_table = r#" + /// CREATE TABLE ##bulk_test_columns ( + /// id INT IDENTITY PRIMARY KEY, + /// foo INT NOT NULL, + /// bar FLOAT NOT NULL + /// ) + /// "#; + /// + /// client.simple_query(create_table).await?; + /// + /// // Start the bulk insert with the client. + /// let mut req = client.bulk_insert_columns("##bulk_test_columns", &["foo", "bar"]).await?; + /// + /// for (i, j) in [(0i32, 0f64), (1i32, 1f64), (2i32, 2f64)] { + /// let row = (i, j).into_row(); + /// + /// // The request will handle flushing to the wire in an optimal way, + /// // balancing between memory usage and IO performance. + /// req.send(row).await?; + /// } + /// + /// // The request must be finalized. + /// let res = req.finalize().await?; + /// assert_eq!(3, res.total()); + /// # Ok(()) + /// # } + /// ``` + pub async fn bulk_insert_columns<'a>( + &'a mut self, + table: &'a str, + columns: &'a [&'a str], + ) -> crate::Result> { + self.bulk_insert_with_options(table, columns, SqlBulkCopyOptions::empty(), &[]) + .await + } - let token_stream = TokenStream::new(&mut self.connection).try_unfold(); + /// Execute a `BULK INSERT` statement like [`bulk_insert_columns`], with + /// additional control over the emitted `WITH (...)` clause. + /// + /// `options` is a set of [`SqlBulkCopyOption`] flags (combine them with `|`, + /// or pass [`SqlBulkCopyOptions::empty`] for none), and `order_hints` + /// declares the sort order the incoming rows are already in via an + /// `ORDER (...)` clause. Pass `&["*"]` as `columns` to target every column, + /// exactly like [`bulk_insert`] does internally. + /// + /// When `options` is empty and `order_hints` is empty this emits the exact + /// same statement as [`bulk_insert_columns`] (no `WITH` clause). + /// + /// # Security + /// + /// `table`, every entry of `columns`, and every order-hint column name are + /// interpolated **directly** into the SQL batches sent to the server (T-SQL + /// does not allow identifiers to be parameterized). The caller MUST pass + /// **trusted, hard-coded or otherwise validated** identifiers and MUST NOT + /// pass untrusted or user-supplied input. As cheap defense-in-depth this + /// method rejects obviously-malformed identifiers — NUL/ASCII control + /// characters, an unbalanced `]` bracket, a top-level space, statement- + /// breaking punctuation (`;`, quotes, `-`, etc.) and tokens spliced onto a + /// closing `]`/`)` — for the table, columns *and* order-hint columns, but + /// that guard is not a substitute for passing trusted input. + /// + /// # Errors + /// + /// Returns an error if `table`, a column, or an order-hint column is a + /// malformed identifier, if the query for the column metadata and + /// collations fails, if the server reports a collation name that is not a + /// plain identifier, or if the server rejects the `INSERT BULK` statement. + /// Row-level failures surface later from [`send`] and [`finalize`] on the + /// returned request. + /// + /// # Collations + /// + /// The `INSERT BULK` statement declares each char, varchar, text, nchar, + /// nvarchar and ntext column with its own collation (`COLLATE `), as + /// SqlClient's `SqlBulkCopy` does, so the server reads the bulk data in the + /// column's code page rather than the database default's. The collation + /// names come from `sp_tablecollations_100`, run in the same batch as the + /// column metadata query (in `tempdb` for a `#` temp table). + /// + /// [`bulk_insert`]: #method.bulk_insert + /// [`bulk_insert_columns`]: #method.bulk_insert_columns + /// [`send`]: BulkLoadRequest::send + /// [`finalize`]: BulkLoadRequest::finalize + /// + /// # Example + /// + /// ```no_run + /// # use tiberius::{Config, IntoRow, SqlBulkCopyOption, SortOrder}; + /// # use tokio_util::compat::TokioAsyncWriteCompatExt; + /// # use std::env; + /// # #[tokio::main] + /// # async fn main() -> Result<(), Box> { + /// # let c_str = env::var("TIBERIUS_TEST_CONNECTION_STRING").unwrap_or( + /// # "server=tcp:localhost,1433;integratedSecurity=true;TrustServerCertificate=true".to_owned(), + /// # ); + /// # let config = Config::from_ado_string(&c_str)?; + /// # let tcp = tokio::net::TcpStream::connect(config.get_addr()).await?; + /// # tcp.set_nodelay(true)?; + /// # let mut client = tiberius::Client::connect(config, tcp.compat_write()).await?; + /// let mut req = client + /// .bulk_insert_with_options( + /// "##bulk_test", + /// &["id", "val"], + /// SqlBulkCopyOption::KeepIdentity | SqlBulkCopyOption::TableLock, + /// &[("id", SortOrder::Ascending)], + /// ) + /// .await?; + /// + /// for i in [0i32, 1i32, 2i32] { + /// req.send((i, i).into_row()).await?; + /// } + /// + /// let res = req.finalize().await?; + /// # Ok(()) + /// # } + /// ``` + pub async fn bulk_insert_with_options<'a>( + &'a mut self, + table: &'a str, + columns: &'a [&'a str], + options: SqlBulkCopyOptions, + order_hints: &'a [ColumnOrderHint<'a>], + ) -> crate::Result> { + // `table` is interpolated directly into the SQL batch (identifiers cannot + // be parameterized in T-SQL). Reject obviously-malformed/dangerous input + // as cheap defense-in-depth; see the `# Security` note above. + validate_bulk_table_identifier(table)?; - let columns = token_stream - .try_fold(None, |mut columns, token| async move { - if let ReceivedToken::NewResultset(metadata) = token { - columns = Some(metadata.columns.clone()); - }; + // Each `columns` entry is likewise interpolated directly into the SQL + // (both the metadata `SELECT` and the `INSERT BULK` column list), so it + // gets the same cheap defense-in-depth guard as `table`. + for column in columns { + validate_bulk_column_identifier(column)?; + } - Ok(columns) - }) - .await?; + // Order-hint column names are interpolated into the `ORDER (...)` clause, + // so they get the exact same guard (and, below, the same bracket + // escaping) as the regular column list. + for &(column, _) in order_hints { + validate_bulk_column_identifier(column)?; + } - // now start bulk upload - let columns: Vec<_> = columns - .ok_or_else(|| { - crate::Error::Protocol("expecting column metadata from query but not found".into()) - })? + // Retrieve column metadata from the server, keeping only the updateable + // columns as bulk targets (read-only/computed columns are skipped). + // + // Identity columns are normally read-only (so filtered out, letting the + // server assign values), but `KEEP_IDENTITY` is implemented — exactly as + // ADO.NET's `SqlBulkCopy` does — by *including* the identity column in + // the bulk column list so the caller supplies explicit values; there is + // no `KEEP_IDENTITY` keyword in the `INSERT BULK` `WITH (...)` grammar. + // So when the flag is set, identity columns are additionally retained. + // + // The same batch lists the collation of every column, which the + // `INSERT BULK` column list declares for character columns. + let keep_identity = options.contains(SqlBulkCopyOption::KeepIdentity); + let (columns, collations) = self.fetch_column_metadata(table, columns, true).await?; + let mut columns: Vec<_> = columns .into_iter() - .filter(|column| column.base.flags.contains(ColumnFlag::Updateable)) + .filter(|column| bulk_column_is_target(&column.base, keep_identity)) .collect(); + // `text`/`ntext`/`image` columns must carry the destination TableName in + // the COLMETADATA we emit for the bulk load (MS-TDS §2.2.7.4). Record the + // target table on every column; the encoder only emits it for those + // types, so this is a no-op on the wire for all other columns. + for column in columns.iter_mut() { + column.base.table_name = Some(table.to_string()); + } + + // now start bulk upload self.connection.flush_stream().await?; - let col_data = columns.iter().map(|c| format!("{}", c)).join(", "); - let query = format!("INSERT BULK {} ({})", table, col_data); + let col_data = bulk_column_list(&columns, &collations)?; + let query = build_insert_bulk_sql(table, &col_data, options, order_hints); let req = BatchRequest::new(query, self.connection.context().transaction_descriptor()); let id = self.connection.context_mut().next_packet_id(); @@ -347,22 +599,293 @@ impl Client { BulkLoadRequest::new(&mut self.connection, columns) } + /// Retrieve the column metadata for a set of columns of a table, including + /// the column names, types (with their size, precision and scale) and flags + /// such as nullability and whether a column is an identity column. + /// + /// Pass `&["*"]` as `columns` to return the metadata for every column of the + /// table. + /// + /// # Security + /// + /// Both `table` and the entries of `columns` are interpolated **directly** + /// into the SQL batch sent to the server (`SELECT TOP 0 {columns} FROM + /// {table}`). SQL Server does not allow identifiers to be supplied as bound + /// parameters, so these values cannot be parameterized — they become part of + /// the SQL text verbatim. The caller MUST therefore pass **trusted, + /// hard-coded or otherwise validated** identifiers and MUST NOT pass + /// untrusted or user-supplied input, which would open a SQL injection + /// vector. As cheap defense-in-depth this method rejects obviously-malformed + /// identifiers (NUL/ASCII control characters or an unbalanced `]` bracket), + /// but that guard is not a substitute for passing trusted input. + /// + /// # Errors + /// + /// Returns an error if `table` is a malformed identifier, if the metadata + /// `SELECT` fails, or if the server reports an error while executing it. + /// + /// ```no_run + /// # use tiberius::Config; + /// # use tokio_util::compat::TokioAsyncWriteCompatExt; + /// # use std::env; + /// # #[tokio::main] + /// # async fn main() -> Result<(), Box> { + /// # let c_str = env::var("TIBERIUS_TEST_CONNECTION_STRING").unwrap_or( + /// # "server=tcp:localhost,1433;integratedSecurity=true;TrustServerCertificate=true".to_owned(), + /// # ); + /// # let config = Config::from_ado_string(&c_str)?; + /// # let tcp = tokio::net::TcpStream::connect(config.get_addr()).await?; + /// # tcp.set_nodelay(true)?; + /// # let mut client = tiberius::Client::connect(config, tcp.compat_write()).await?; + /// let meta = client.column_metadata("some_table", &["*"]).await?; + /// assert!(meta[0].base().is_identity()); + /// # Ok(()) + /// # } + /// ``` + pub async fn column_metadata( + &mut self, + table: &str, + columns: &[&str], + ) -> crate::Result>> { + // `table` and each `column` are interpolated directly into the SQL + // batch below (identifiers cannot be parameterized in T-SQL). Reject + // obviously-malformed/dangerous input as cheap defense-in-depth; see the + // `# Security` note above. + validate_bulk_table_identifier(table)?; + for column in columns { + validate_bulk_column_identifier(column)?; + } + + Ok(self.fetch_column_metadata(table, columns, false).await?.0) + } + + /// Fetch the metadata of `columns` of `table` like [`column_metadata`], + /// and with `collations` also the `(column name, collation name)` of + /// every column of the table from `sp_tablecollations_100` (see + /// [`table_collations_sql`]), in the same batch. Without `collations` the + /// list is empty. + /// + /// `table` and `columns` must already be validated. + /// + /// [`column_metadata`]: #method.column_metadata + async fn fetch_column_metadata( + &mut self, + table: &str, + columns: &[&str], + collations: bool, + ) -> crate::Result<(Vec>, Vec<(String, Option)>)> { + self.connection.flush_stream().await?; + + // Ask the server for the column layout without returning any rows. + let columns = columns.join(", "); + let mut query = format!("SELECT TOP 0 {columns} FROM {table}"); + if collations { + let version = self.connection.context().version(); + query.push_str("; "); + query.push_str(&table_collations_sql(table, version)); + } + + let req = BatchRequest::new(query, self.connection.context().transaction_descriptor()); + let id = self.connection.context_mut().next_packet_id(); + self.connection.send(PacketHeader::batch(id), req).await?; + + let token_stream = TokenStream::new(&mut self.connection).try_unfold(); + + // The `SELECT TOP 0` result set comes first and has no rows; every + // row is one of the collation list. + let (columns, collations) = token_stream + .try_fold( + (None, Vec::new()), + |(mut columns, mut collations), token| async move { + match token { + ReceivedToken::NewResultset(metadata) if columns.is_none() => { + columns = Some(metadata.columns.clone()); + } + ReceivedToken::Row(row) => { + let string = |i| match row.get(i) { + Some(ColumnData::String(s)) => s.as_ref().map(|s| s.to_string()), + _ => None, + }; + if let Some(name) = string(1) { + collations.push((name, string(3))); + } + } + _ => {} + } + + Ok((columns, collations)) + }, + ) + .await?; + + let columns = columns.ok_or_else(|| { + crate::Error::Protocol("expecting column metadata from query but not found".into()) + })?; + + // Own the column names so the returned metadata is not tied to the + // lifetime of the token stream. + let columns = columns + .into_iter() + .map(|c| MetaDataColumn { + base: c.base, + col_name: std::borrow::Cow::Owned(c.col_name.into_owned()), + }) + .collect(); + + Ok((columns, collations)) + } + + /// Sends a TDS Attention signal to the server (packet type `0x06`, + /// MS-TDS section 2.2.1.6) to cancel the request that is currently in + /// flight on this connection, and drains the acknowledging token stream so + /// the connection can be reused for further queries. + /// + /// The server responds to the Attention signal by aborting the running + /// batch or RPC and returning a `DONE` token with the `DONE_ATTN` status + /// bit set. This method waits for that acknowledgement before returning, + /// discarding any remaining rows or tokens from the cancelled request. + /// + /// # Query cancellation and futures + /// + /// Dropping a [`query`], [`execute`] or [`simple_query`] future (for + /// example when a `tokio::time::timeout` elapses or a `select!` branch is + /// cancelled) stops the client from polling the stream, but it does *not* + /// tell the server to stop working on the request. To actually cancel the + /// in-flight work on the server, keep the [`Client`] and call + /// `cancel_query` on it. Because `cancel_query` borrows the client + /// mutably, it can only be issued once the borrowing result stream has + /// been dropped — typically from a separate task holding the client, or + /// after a cancelled/timed-out future has released its borrow. + /// + /// [`query`]: #method.query + /// [`execute`]: #method.execute + /// [`simple_query`]: #method.simple_query + pub async fn cancel_query(&mut self) -> crate::Result<()> { + self.connection.cancel_request().await?; + Ok(()) + } + /// Closes this database connection explicitly. pub async fn close(self) -> crate::Result<()> { self.connection.close().await } + /// Begins a new transaction using a Transaction Manager request + /// (`TM_BEGIN_XACT`, MS-TDS 2.2.6.8) instead of a `BEGIN TRAN` T-SQL + /// batch. + /// + /// On success the server replies with a `BeginTransaction` environment + /// change token whose descriptor is stored in the connection context and + /// automatically attached to subsequent requests, scoping them to the + /// transaction. Commit the work with [`commit_transaction`] or discard it + /// with [`rollback_transaction`]. + /// + /// The transaction uses the server's default isolation level. Use + /// [`begin_transaction_with_isolation`] to request a specific one. + /// + /// [`commit_transaction`]: #method.commit_transaction + /// [`rollback_transaction`]: #method.rollback_transaction + /// [`begin_transaction_with_isolation`]: #method.begin_transaction_with_isolation + pub async fn begin_transaction(&mut self) -> crate::Result<()> { + self.begin_transaction_with_isolation(IsolationLevel::Unspecified) + .await + } + + /// Begins a new transaction with an explicit isolation level using a + /// Transaction Manager request (`TM_BEGIN_XACT`, MS-TDS 2.2.6.8). + /// + /// See [`begin_transaction`] for details on transaction scoping. + /// + /// [`begin_transaction`]: #method.begin_transaction + pub async fn begin_transaction_with_isolation( + &mut self, + isolation_level: IsolationLevel, + ) -> crate::Result<()> { + let req = TransactionManagerRequest::begin( + self.connection.context().transaction_descriptor(), + isolation_level, + "", + ); + + self.send_transaction_manager_request(req).await + } + + /// Commits the active transaction using a Transaction Manager request + /// (`TM_COMMIT_XACT`, MS-TDS 2.2.6.8). + /// + /// After a successful commit the connection is no longer scoped to a + /// transaction. + pub async fn commit_transaction(&mut self) -> crate::Result<()> { + let req = TransactionManagerRequest::commit( + self.connection.context().transaction_descriptor(), + "", + ); + + self.send_transaction_manager_request(req).await + } + + /// Rolls back the active transaction using a Transaction Manager request + /// (`TM_ROLLBACK_XACT`, MS-TDS 2.2.6.8). + /// + /// After a successful rollback the connection is no longer scoped to a + /// transaction. + pub async fn rollback_transaction(&mut self) -> crate::Result<()> { + let req = TransactionManagerRequest::rollback( + self.connection.context().transaction_descriptor(), + "", + ); + + self.send_transaction_manager_request(req).await + } + + /// Creates a named savepoint in the active transaction using a Transaction + /// Manager request (`TM_SAVE_XACT`, MS-TDS 2.2.6.8). + /// + /// The savepoint can later be targeted by a T-SQL `ROLLBACK TRANSACTION + /// ` to undo work performed after it while keeping the surrounding + /// transaction open. + pub async fn save_transaction<'a>( + &mut self, + name: impl Into>, + ) -> crate::Result<()> { + let req = TransactionManagerRequest::save( + self.connection.context().transaction_descriptor(), + name, + ); + + self.send_transaction_manager_request(req).await + } + + async fn send_transaction_manager_request( + &mut self, + req: TransactionManagerRequest<'_>, + ) -> crate::Result<()> { + self.connection.flush_stream().await?; + + let id = self.connection.context_mut().next_packet_id(); + self.connection + .send(PacketHeader::transaction_manager(id), req) + .await?; + + // The server responds with a DONE token (plus an ENVCHANGE token that + // the token stream applies to the connection context, updating the + // active transaction descriptor). + TokenStream::new(&mut self.connection).flush_done().await?; + + Ok(()) + } + pub(crate) fn rpc_params<'a>(query: impl Into>) -> Vec> { vec![ RpcParam { name: Cow::Borrowed("stmt"), flags: BitFlags::empty(), - value: ColumnData::String(Some(query.into())), + value: RpcValue::Scalar(ColumnData::String(Some(query.into()))), }, RpcParam { name: Cow::Borrowed("params"), flags: BitFlags::empty(), - value: ColumnData::I32(Some(0)), + value: RpcValue::Scalar(ColumnData::I32(Some(0))), }, ] } @@ -388,12 +911,12 @@ impl Client { rpc_params.push(RpcParam { name: Cow::Owned(format!("@P{}", i + 1)), flags: BitFlags::empty(), - value: param, + value: RpcValue::Scalar(param), }); } if let Some(params) = rpc_params.iter_mut().find(|x| x.name == "params") { - params.value = ColumnData::String(Some(param_str.into())); + params.value = RpcValue::Scalar(ColumnData::String(Some(param_str.into()))); } let req = TokenRpcRequest::new( @@ -407,4 +930,1089 @@ impl Client { Ok(()) } + + /// Sends a named-procedure RPC request with the given parameters. The caller + /// is responsible for flushing the connection beforehand and for consuming + /// the resulting token stream. + pub(crate) async fn rpc_run_command<'a, 'b>( + &'a mut self, + command_name: Cow<'b, str>, + rpc_params: Vec>, + ) -> crate::Result<()> + where + 'a: 'b, + { + let req = TokenRpcRequest::new( + command_name, + rpc_params, + self.connection.context().transaction_descriptor(), + ); + + let id = self.connection.context_mut().next_packet_id(); + self.connection.send(PacketHeader::rpc(id), req).await?; + + Ok(()) + } + + /// Runs a batch query solely to retrieve its column metadata. Used to + /// resolve the column layout of a table-valued parameter type. + pub(crate) async fn query_run_for_metadata<'b>( + &mut self, + query: String, + ) -> crate::Result>>> { + self.connection.flush_stream().await?; + + let req = BatchRequest::new(query, self.connection.context().transaction_descriptor()); + + let id = self.connection.context_mut().next_packet_id(); + self.connection.send(PacketHeader::batch(id), req).await?; + + let token_stream = TokenStream::new(&mut self.connection).try_unfold(); + + let columns = token_stream + .try_fold(None, |mut columns, token| async move { + if let ReceivedToken::NewResultset(metadata) = token { + columns = Some(metadata.columns.clone()); + }; + + Ok(columns) + }) + .await?; + + Ok(columns) + } +} + +/// Build the `INSERT BULK () [WITH (...)]` statement text. +/// +/// `col_data` is the already-formatted column list (each column bracket-quoted +/// with its type, joined by `, `). Active [`SqlBulkCopyOptions`] flags and any +/// `order_hints` are collected into a `WITH (...)` clause; when both are empty +/// no `WITH` clause is emitted, so this reproduces the plain +/// `INSERT BULK
()` used by `bulk_insert_columns`. +/// +/// Order-hint column names are interpolated verbatim as already-valid SQL +/// identifiers — exactly like `table` and the `columns` entries elsewhere in +/// this module — so a caller may pass a plain (`id`), bracket-quoted (`[id]`, +/// `[my]]col]`) or multi-part name and get that same text back. Callers are +/// expected to have run every identifier through [`validate_bulk_column_identifier`] +/// first; do NOT re-bracket here, or an already-quoted name like `[id]` would be +/// double-escaped into the wrong identifier `[[id]]]`. +/// +/// Note there is deliberately no `KEEP_IDENTITY` keyword: the `INSERT BULK` +/// `WITH (...)` grammar has no such option (unlike the textual `BULK INSERT` +/// statement). [`SqlBulkCopyOption::KeepIdentity`] is honoured by the caller +/// keeping the identity column in `col_data` instead. +fn build_insert_bulk_sql( + table: &str, + col_data: &str, + options: SqlBulkCopyOptions, + order_hints: &[ColumnOrderHint<'_>], +) -> String { + let mut clauses: Vec = Vec::new(); + + // NB: `KeepIdentity` intentionally emits no keyword here — it is a + // column-inclusion concern handled in `bulk_insert_with_options`. + if options.contains(SqlBulkCopyOption::CheckConstraints) { + clauses.push("CHECK_CONSTRAINTS".to_owned()); + } + if options.contains(SqlBulkCopyOption::TableLock) { + clauses.push("TABLOCK".to_owned()); + } + if options.contains(SqlBulkCopyOption::KeepNulls) { + clauses.push("KEEP_NULLS".to_owned()); + } + if options.contains(SqlBulkCopyOption::FireTriggers) { + clauses.push("FIRE_TRIGGERS".to_owned()); + } + + if !order_hints.is_empty() { + let hints = order_hints + .iter() + .map(|(col, order)| { + // Interpolate the (already-validated) identifier verbatim, like + // `table`/`columns`. Re-bracketing here would double-escape an + // already-quoted name (`[id]` -> `[[id]]]`). + format!( + "{col} {}", + match order { + SortOrder::Ascending => "ASC", + SortOrder::Descending => "DESC", + } + ) + }) + .join(", "); + + clauses.push(format!("ORDER({hints})")); + } + + let mut query = format!("INSERT BULK {table} ({col_data})"); + + if !clauses.is_empty() { + query.push_str(" WITH ("); + query.push_str(&clauses.join(", ")); + query.push(')'); + } + + query +} + +/// Build the `INSERT BULK` column list: each column bracket-quoted with its +/// type, and every character column followed by ` COLLATE `. +/// +/// SQL Server reads the bulk data of a column as the type declared here. A +/// char/varchar/text declaration without `COLLATE` takes the database default +/// collation, so the server would read the bytes in that code page and convert +/// them to the column's, corrupting any text whose column collation differs. +/// Like SqlClient's `SqlBulkCopy`, the column's own collation is declared for +/// char, varchar, text, nchar, nvarchar and ntext columns. +/// +/// `collations` holds the `(column name, collation name)` rows of +/// `sp_tablecollations_100`. A column is matched by exact name, else by a +/// unique case-insensitive name (a select list may spell a name in another +/// case than the table does). A character column without a row or with a +/// `NULL` collation is declared without `COLLATE`. +/// +/// # Errors +/// +/// Returns [`Error::Protocol`](crate::Error::Protocol) if a collation name to +/// be declared is not a plain identifier (see [`validate_collation_name`]). +fn bulk_column_list( + columns: &[MetaDataColumn<'_>], + collations: &[(String, Option)], +) -> crate::Result { + let mut list = Vec::with_capacity(columns.len()); + + for column in columns { + let mut item = format!("{}", column); + + if declares_collation(&column.base.ty) { + if let Some(collation) = column_collation(&column.col_name, collations) { + validate_collation_name(collation)?; + item.push_str(" COLLATE "); + item.push_str(collation); + } + } + + list.push(item); + } + + Ok(list.join(", ")) +} + +/// Whether the `INSERT BULK` declaration of a column of type `ty` carries a +/// `COLLATE` clause: the character types SqlClient declares one for. +fn declares_collation(ty: &crate::tds::codec::TypeInfo) -> bool { + use crate::tds::codec::{TypeInfo, VarLenType}; + + matches!( + ty, + TypeInfo::VarLenSized(ctx) if matches!( + ctx.r#type(), + VarLenType::BigChar + | VarLenType::BigVarChar + | VarLenType::Text + | VarLenType::NChar + | VarLenType::NVarchar + | VarLenType::NText + ) + ) +} + +/// The collation name listed for the column `name`: the row with exactly this +/// name, else the only row whose name matches ignoring case. +fn column_collation<'a>(name: &str, collations: &'a [(String, Option)]) -> Option<&'a str> { + let row = match collations.iter().find(|(n, _)| n == name) { + Some(row) => row, + None => { + let name = name.to_lowercase(); + let mut rows = collations.iter().filter(|(n, _)| n.to_lowercase() == name); + match (rows.next(), rows.next()) { + (Some(row), None) => row, + _ => return None, + } + } + }; + + row.1.as_deref() +} + +/// Build the batch statement that lists the collation of every column of +/// `table`, as SqlClient's `SqlBulkCopy` does: +/// +/// ```text +/// EXEC ..sp_tablecollations_100 N'.
' +/// ``` +/// +/// Its result set has one row per column, ordered by column id, of `colid`, +/// `name`, `tds_collation_100` and `collation_100` (the collation name). +/// +/// The procedure reads the catalog it is called in, so it is called in the +/// table's database: the catalog part of `table`, else `tempdb` for a temp +/// table (a name starting with `#`), else the current database (`..`). Schema +/// and table name are bracket-quoted and put in a string literal. Before SQL +/// Server 2008 (TDS 7.3) the procedure is `sp_tablecollations_90`. +/// +/// `table` must have passed [`validate_bulk_table_identifier`]. +fn table_collations_sql(table: &str, version: crate::FeatureLevel) -> String { + let parts = split_multipart_identifier(table); + let part = |from_end: usize| { + parts + .len() + .checked_sub(from_end + 1) + .map_or("", |i| parts[i].as_str()) + }; + + let table_name = part(0); + let schema = part(1); + let catalog = part(2); + + let catalog = if catalog.is_empty() { + if table_name.starts_with('#') { + "tempdb".to_owned() + } else { + String::new() + } + } else { + quote_identifier(catalog) + }; + + let literal = |part: &str| { + if part.is_empty() { + String::new() + } else { + quote_identifier(&part.replace('\'', "''")) + } + }; + + let procedure = if version >= crate::FeatureLevel::SqlServer2008 { + "sp_tablecollations_100" + } else { + "sp_tablecollations_90" + }; + + format!( + "EXEC {catalog}..{procedure} N'{}.{}'", + literal(schema), + literal(table_name) + ) +} + +/// Bracket-quote `name`, doubling any `]` in it. +fn quote_identifier(name: &str) -> String { + format!("[{}]", name.replace(']', "]]")) +} + +/// Split a multi-part identifier (`server.catalog.schema.table`, any leading +/// parts omitted or empty) into its unquoted parts: `[a]]b].c` gives `a]b` +/// and `c`. `ident` must have passed [`validate_sql_identifier`], so every +/// bracket-quoted segment is terminated. +fn split_multipart_identifier(ident: &str) -> Vec { + let mut parts = vec![String::new()]; + let mut chars = ident.chars().peekable(); + let mut in_bracket = false; + + while let Some(c) = chars.next() { + let part = parts.last_mut().expect("parts is never empty"); + + if in_bracket { + if c == ']' { + if chars.peek() == Some(&']') { + chars.next(); + part.push(']'); + } else { + in_bracket = false; + } + } else { + part.push(c); + } + } else { + match c { + '[' => in_bracket = true, + '.' => parts.push(String::new()), + c => part.push(c), + } + } + } + + parts +} + +/// Reject a server-reported collation name that is not a plain identifier of +/// ASCII letters, digits and `_`, as every SQL Server collation name is. The +/// name is put into the `INSERT BULK` statement unquoted. +fn validate_collation_name(name: &str) -> crate::Result<()> { + if !name.is_empty() && name.bytes().all(|b| b.is_ascii_alphanumeric() || b == b'_') { + Ok(()) + } else { + Err(crate::Error::Protocol( + format!("invalid collation name {name:?} for a bulk insert column").into(), + )) + } +} + +/// Decide whether a server-reported column is a bulk-insert target. +/// +/// Read-only/computed columns are skipped so the server assigns their values. +/// Identity columns are read-only by default (filtered out, letting the server +/// assign them), but when [`SqlBulkCopyOption::KeepIdentity`] is set they are +/// additionally retained so the caller can supply explicit values — matching +/// ADO.NET's `SqlBulkCopy` (there is no `KEEP_IDENTITY` keyword in the +/// `INSERT BULK` `WITH (...)` grammar). Factored out of +/// [`Client::bulk_insert_with_options`] so this selection logic is unit-testable +/// without a live server. +fn bulk_column_is_target( + base: &crate::tds::codec::BaseMetaDataColumn, + keep_identity: bool, +) -> bool { + base.is_updateable() || (keep_identity && base.is_identity()) +} + +/// Reject an obviously-malformed or dangerous bulk-insert table identifier. +/// +/// The `table` argument of [`Client::bulk_insert`] / [`Client::bulk_insert_columns`] +/// is interpolated directly into the SQL batch because T-SQL does not allow +/// identifiers to be parameterized. This guard is cheap defense-in-depth — it +/// does NOT make untrusted input safe. It only rejects input that cannot be a +/// legitimate identifier (see [`validate_sql_identifier`] for the exact rules). +/// +/// It deliberately does NOT try to quote or rewrite the identifier, so +/// multi-part names (`schema.table`), already-bracketed names (`[my table]`) and +/// temp tables (`##bulk_test`) keep working unchanged. +fn validate_bulk_table_identifier(table: &str) -> crate::Result<()> { + validate_sql_identifier("bulk insert table", table) +} + +/// Reject an obviously-malformed or dangerous bulk-insert `column` identifier. +/// +/// Column names are interpolated into the metadata `SELECT` and `INSERT BULK` +/// column list exactly like `table`, so they get the same cheap +/// defense-in-depth check. See [`validate_bulk_table_identifier`]. +fn validate_bulk_column_identifier(column: &str) -> crate::Result<()> { + validate_sql_identifier("bulk insert column", column) +} + +/// Reject an obviously-malformed or dangerous SQL identifier that will be +/// interpolated directly into a SQL batch because T-SQL does not allow +/// identifiers or type names to be parameterized. `what` names the kind of +/// identifier for the error message. +/// +/// This is cheap defense-in-depth — it does NOT make untrusted input safe. It +/// rejects input that cannot be a legitimate identifier: +/// +/// - a NUL byte or any ASCII control character; +/// - **outside** a `[...]` bracket-quoted segment, any character that is not +/// identifier-safe. Only alphanumerics and the punctuation needed for real +/// identifiers/type names are allowed — `_`, `.` (multi-part names like +/// `dbo.MyType`), `@`/`#` (variable/temp-style names), `*` (a bulk column +/// list may be `*`), and `(`, `)`, `,` (parameterized types such as +/// `decimal(10,2)` / `varchar(max)`). Parens must be balanced — an unbalanced +/// `)` *or* an unclosed `(` is rejected — and both a space and a comma are allowed only *inside* +/// those parens (e.g. `decimal(18, 4)`); a top-level space or comma is +/// rejected because it would let one identifier split into several SQL tokens +/// (a top-level comma would splice a verbatim order hint `a,b` into two +/// columns). This rejects statement-breaking characters such as `;`, quotes +/// and `-`, so a value like `1; DROP TABLE Users--` cannot slip through; and +/// - **inside** a `[...]` bracket-quoted segment anything is allowed except an +/// unescaped `]` (per the T-SQL bracket-escaping rule a literal `]` must be +/// doubled as `]]`). The segment must be terminated — an unclosed `[` running +/// to the end of the string is rejected. A `]` seen outside any bracket is +/// unbalanced and rejected. A closing `]`, and a top-level closing `)`, are themselves SQL +/// token boundaries, so a segment may only be followed by `.` (the next part +/// of a multi-part name) or the end of the string — this stops delimiter- +/// adjacent splicing such as `[t]UNION(...)` or `foo(1)UNION(...)` that needs +/// no space. +/// +/// It deliberately does NOT try to quote or rewrite the identifier, so +/// multi-part names (`schema.table`), already-bracketed names (`[my table]`, +/// `[dbo].[my table]`, `[weird]]name]`) and temp tables (`##bulk_test`) keep +/// working unchanged. This is shared by the bulk-insert guards here and by +/// `validate_db_type_identifier` in `src/command.rs`. +pub(crate) fn validate_sql_identifier(what: &str, ident: &str) -> crate::Result<()> { + if ident.chars().any(|c| c.is_ascii_control()) { + return Err(crate::Error::BulkInput( + format!("{what} identifier must not contain NUL or control characters").into(), + )); + } + + // Apply the T-SQL bracket rule: inside a `[...]` quoted identifier a literal + // `]` must be doubled (`]]`); a single `]` closes the bracket. A `]` seen + // outside of any bracket is unbalanced and rejected. Tracking bracket state + // keeps legitimate names like `[dbo].[my table]` and `[weird]]name]` + // working while catching stray closing brackets such as `Foo]`. + let mut chars = ident.chars().peekable(); + let mut in_bracket = false; + let mut paren_depth: u32 = 0; + while let Some(c) = chars.next() { + if in_bracket { + if c == ']' { + if chars.peek() == Some(&']') { + chars.next(); // consume the doubled `]]` escape + } else { + in_bracket = false; + // A quoted-identifier segment may only be followed by `.` + // (introducing the next part of a multi-part name) or the + // end of the string. Anything else — e.g. `[t]UNION ...` — + // splices a fresh SQL token directly onto the closing `]` + // (which is itself a token boundary, needing no space), so + // reject it. + match chars.peek() { + None | Some('.') => {} + Some(_) => { + return Err(crate::Error::BulkInput( + format!("{what} identifier has trailing characters after a bracket-quoted segment").into(), + )); + } + } + } + } + continue; + } + + match c { + '[' => in_bracket = true, + ']' => { + return Err(crate::Error::BulkInput( + format!("{what} identifier contains an unbalanced `]` bracket").into(), + )); + } + '(' => paren_depth += 1, + ')' => { + // A `)` with no matching `(` is unbalanced. Left unchecked + // (`saturating_sub`) it would sail through and, interpolated + // verbatim, could close a caller-supplied enclosing paren early + // — e.g. an order hint `col)` desyncing `ORDER(col) ...)`. Reject + // it exactly like an unbalanced `]`. + if paren_depth == 0 { + return Err(crate::Error::BulkInput( + format!("{what} identifier contains an unbalanced `)`").into(), + )); + } + paren_depth -= 1; + // A parameterized type ends at its closing paren; nothing is a + // legitimate continuation after a top-level `)`. Rejecting it + // stops `)`-delimited token splicing such as `foo(1)UNION(...)` + // (the `)` is a token boundary, so no space is required). + if paren_depth == 0 && chars.peek().is_some() { + return Err(crate::Error::BulkInput( + format!("{what} identifier has trailing characters after a closing `)`") + .into(), + )); + } + } + // A space is only legitimate inside a parameterized type's parens + // (e.g. `decimal(18, 4)`) or inside a bracket-quoted name (handled + // above). A *top-level* space would let a single identifier split + // into several SQL tokens (`t UNION SELECT ...`, `t WHERE ...`) + // without using any otherwise-blocked character, so reject it. + ' ' if paren_depth == 0 => { + return Err(crate::Error::BulkInput( + format!("{what} identifier contains a disallowed character").into(), + )); + } + // A comma is only legitimate inside a parameterized type's parens + // (`decimal(10,2)`). A *top-level* comma would splice a single + // identifier into two — e.g. a verbatim order hint `a,b` becoming + // two order columns `a` and `b` — so reject it like a top-level + // space. + ',' if paren_depth == 0 => { + return Err(crate::Error::BulkInput( + format!("{what} identifier contains a disallowed character").into(), + )); + } + // Outside a bracket-quoted segment only identifier-safe characters + // are permitted (see the doc comment). Everything else — `;`, + // quotes, `-`, etc. — is rejected. + c if c.is_alphanumeric() => {} + '_' | '.' | '@' | '#' | ',' | ' ' | '*' => {} + _ => { + return Err(crate::Error::BulkInput( + format!("{what} identifier contains a disallowed character").into(), + )); + } + } + } + + // Reject unbalanced *openers* left dangling at end of input, mirroring the + // rejection of unbalanced closers above. An unterminated `[` would quote and + // swallow whatever follows when interpolated (e.g. an order hint `[abc` + // becoming `ORDER([abc ASC))`); an unclosed `(` leaves an enclosing paren + // open. A legitimate identifier/type name never ends mid-bracket or with an + // open paren. + if in_bracket { + return Err(crate::Error::BulkInput( + format!("{what} identifier has an unterminated `[` bracket-quoted segment").into(), + )); + } + if paren_depth != 0 { + return Err(crate::Error::BulkInput( + format!("{what} identifier contains an unbalanced `(`").into(), + )); + } + + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::{ + build_insert_bulk_sql, bulk_column_is_target, validate_bulk_column_identifier, + validate_bulk_table_identifier, + }; + use crate::tds::codec::{BaseMetaDataColumn, ColumnFlag, TypeInfo, VarLenContext, VarLenType}; + use crate::{SortOrder, SqlBulkCopyOption, SqlBulkCopyOptions}; + use enumflags2::BitFlags; + + // Build a `BaseMetaDataColumn` carrying exactly `flags`, so the bulk-target + // selection predicate can be exercised without a live server. The column + // type is irrelevant to the predicate; a small fixed-len type stands in. + fn base_with_flags(flags: BitFlags) -> BaseMetaDataColumn { + BaseMetaDataColumn { + flags, + ty: TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Intn, 4, None)), + table_name: None, + } + } + + #[test] + fn bulk_column_is_target_selects_writeable_and_keep_identity() { + let updateable: BitFlags = ColumnFlag::Updateable.into(); + let unknown: BitFlags = ColumnFlag::UpdateableUnknown.into(); + let read_only = BitFlags::::empty(); + let identity_ro: BitFlags = ColumnFlag::Identity.into(); + // An identity column SQL Server also reports as updateable-unknown. + let identity_unknown = ColumnFlag::Identity | ColumnFlag::UpdateableUnknown; + + // Read/write and updateable-unknown columns are always targets, + // regardless of the KeepIdentity flag. + for keep in [false, true] { + assert!(bulk_column_is_target(&base_with_flags(updateable), keep)); + assert!(bulk_column_is_target(&base_with_flags(unknown), keep)); + // A plain read-only (non-identity) column is never a target. + assert!(!bulk_column_is_target(&base_with_flags(read_only), keep)); + } + + // A read-only identity column is skipped by default (server assigns the + // value) but retained when KeepIdentity is set. + assert!(!bulk_column_is_target(&base_with_flags(identity_ro), false)); + assert!(bulk_column_is_target(&base_with_flags(identity_ro), true)); + + // KeepIdentity only *adds* identity columns; it never drops a column that + // was already a target on its own merits. + assert!(bulk_column_is_target( + &base_with_flags(identity_unknown), + false + )); + assert!(bulk_column_is_target( + &base_with_flags(identity_unknown), + true + )); + + // A non-identity read-only column is not resurrected by KeepIdentity. + assert!(!bulk_column_is_target(&base_with_flags(read_only), true)); + } + + // A fixed, hand-written column list stands in for the server-provided + // `MetaDataColumn` Display output so the SQL-string builder can be unit + // tested without a live connection. + const COLS: &str = "[id] int, [val] int"; + + #[test] + fn build_sql_no_options_no_hints_emits_no_with_clause() { + // Empty options + empty hints must reproduce the plain statement that + // `bulk_insert_columns` has always emitted. + assert_eq!( + build_insert_bulk_sql("t", COLS, SqlBulkCopyOptions::empty(), &[]), + "INSERT BULK t ([id] int, [val] int)" + ); + } + + #[test] + fn build_sql_each_flag_emits_expected_keyword() { + // `KeepIdentity` is deliberately absent: the `INSERT BULK` `WITH (...)` + // grammar has no `KEEP_IDENTITY` keyword (it is a column-inclusion + // concern), so only these four flags map to keywords. + for (flag, keyword) in [ + (SqlBulkCopyOption::CheckConstraints, "CHECK_CONSTRAINTS"), + (SqlBulkCopyOption::TableLock, "TABLOCK"), + (SqlBulkCopyOption::KeepNulls, "KEEP_NULLS"), + (SqlBulkCopyOption::FireTriggers, "FIRE_TRIGGERS"), + ] { + assert_eq!( + build_insert_bulk_sql("t", COLS, flag.into(), &[]), + format!("INSERT BULK t ([id] int, [val] int) WITH ({keyword})"), + "flag {flag:?} did not emit {keyword}", + ); + } + } + + #[test] + fn build_sql_keep_identity_emits_no_with_keyword() { + // Regression guard: emitting `KEEP_IDENTITY` into the `WITH (...)` clause + // is a syntax error against the TDS `INSERT BULK` statement (verified + // against ADO.NET's `SqlBulkCopy` and go-mssqldb, neither of which emits + // it). `KeepIdentity` alone must therefore produce no `WITH` clause at + // all; the flag's effect is entirely in column selection. + let sql = build_insert_bulk_sql("t", COLS, SqlBulkCopyOption::KeepIdentity.into(), &[]); + assert_eq!(sql, "INSERT BULK t ([id] int, [val] int)", "got: {sql}"); + assert!(!sql.contains("KEEP_IDENTITY"), "got: {sql}"); + } + + #[test] + fn build_sql_combined_flags_join_in_fixed_order() { + // `KeepIdentity` contributes no keyword, so only TABLOCK/FIRE_TRIGGERS + // appear even though it is set. + let opts = SqlBulkCopyOption::KeepIdentity + | SqlBulkCopyOption::TableLock + | SqlBulkCopyOption::FireTriggers; + + assert_eq!( + build_insert_bulk_sql("t", COLS, opts, &[]), + "INSERT BULK t ([id] int, [val] int) WITH (TABLOCK, FIRE_TRIGGERS)" + ); + } + + #[test] + fn build_sql_all_flags() { + // `all()` includes `KeepIdentity`, which emits nothing, so the keyword + // list is the four real `WITH` options. + let opts = SqlBulkCopyOptions::all(); + assert_eq!( + build_insert_bulk_sql("t", COLS, opts, &[]), + "INSERT BULK t ([id] int, [val] int) WITH (CHECK_CONSTRAINTS, TABLOCK, KEEP_NULLS, FIRE_TRIGGERS)" + ); + } + + #[test] + fn build_sql_order_hints_asc_and_desc() { + // Plain column names are interpolated verbatim (like `columns`), not + // re-bracketed, and the reference `ORDER(` has no space. + let hints = [("id", SortOrder::Ascending), ("val", SortOrder::Descending)]; + assert_eq!( + build_insert_bulk_sql("t", COLS, SqlBulkCopyOptions::empty(), &hints), + "INSERT BULK t ([id] int, [val] int) WITH (ORDER(id ASC, val DESC))" + ); + } + + #[test] + fn build_sql_options_and_order_hints_combine() { + let hints = [("id", SortOrder::Ascending)]; + assert_eq!( + build_insert_bulk_sql("t", COLS, SqlBulkCopyOption::TableLock.into(), &hints), + "INSERT BULK t ([id] int, [val] int) WITH (TABLOCK, ORDER(id ASC))" + ); + } + + #[test] + fn build_sql_order_hint_bracketed_name_not_double_escaped() { + // Regression guard: order-hint column names are interpolated verbatim as + // already-valid identifiers, exactly like `table`/`columns`. An + // already-bracket-quoted name (which the validator accepts) must be + // emitted as-is — NOT re-bracketed into `[[id]]]` or `[[my]]]]col]]]`. + let hints = [ + ("[id]", SortOrder::Ascending), + ("[my]]col]", SortOrder::Descending), + ]; + assert_eq!( + build_insert_bulk_sql("t", COLS, SqlBulkCopyOptions::empty(), &hints), + "INSERT BULK t ([id] int, [val] int) WITH (ORDER([id] ASC, [my]]col] DESC))" + ); + } + + #[test] + fn build_sql_empty_hints_emits_no_order_clause() { + let sql = build_insert_bulk_sql("t", COLS, SqlBulkCopyOption::TableLock.into(), &[]); + assert!(!sql.contains("ORDER"), "got: {sql}"); + assert_eq!(sql, "INSERT BULK t ([id] int, [val] int) WITH (TABLOCK)"); + } + + #[test] + fn build_sql_order_hint_column_names_are_validated_like_columns() { + // The public API runs order-hint column names through the same guard as + // the column list; confirm that guard rejects a stray `]`. + assert!(validate_bulk_column_identifier("my]col").is_err()); + assert!(validate_bulk_column_identifier("[my]]col]").is_ok()); + } + + #[test] + fn accepts_normal_column_identifiers() { + for column in ["foo", "bar", "*", "[my col]", "[weird]]col]"] { + assert!( + validate_bulk_column_identifier(column).is_ok(), + "expected column {column:?} to be accepted", + ); + } + } + + #[test] + fn rejects_bad_column_identifiers() { + // control character + assert!(validate_bulk_column_identifier("foo\0bar").is_err()); + assert!(validate_bulk_column_identifier("foo\nbar").is_err()); + // lone / unbalanced closing bracket + assert!(validate_bulk_column_identifier("foo]").is_err()); + assert!(validate_bulk_column_identifier("a]b").is_err()); + } + + #[test] + fn accepts_normal_identifiers() { + for table in [ + "Foo", + "dbo.Foo", + "##bulk_test", + "#temp", + "[my table]", + "[dbo].[my table]", + "[weird]]name]", // doubled `]]` escape inside brackets + ] { + assert!( + validate_bulk_table_identifier(table).is_ok(), + "expected {table:?} to be accepted", + ); + } + } + + #[test] + fn rejects_control_characters() { + assert!(validate_bulk_table_identifier("Foo\0bar").is_err()); + assert!(validate_bulk_table_identifier("Foo\nbar").is_err()); + assert!(validate_bulk_table_identifier("Foo\tbar").is_err()); + } + + #[test] + fn rejects_unbalanced_closing_bracket() { + assert!(validate_bulk_table_identifier("Foo]").is_err()); + assert!(validate_bulk_table_identifier("[my] table]").is_err()); + assert!(validate_bulk_table_identifier("a]b").is_err()); + } + + #[test] + fn column_metadata_guards_reject_bad_identifiers() { + // `column_metadata` interpolates `table`/`columns` into + // `SELECT TOP 0 {columns} FROM {table}` and now validates both with the + // same guards as `bulk_insert*`. A malformed table or column identifier + // must be rejected before any SQL is built. + assert!(validate_bulk_table_identifier("Foo]").is_err()); + assert!(validate_bulk_table_identifier("Foo\nbar").is_err()); + assert!(validate_bulk_column_identifier("id]").is_err()); + assert!(validate_bulk_column_identifier("col\0").is_err()); + } + + #[test] + fn rejects_statement_breaking_identifiers() { + // The guard must reject statement-breakers (`;`, quotes, `--`, `=`) + // that would let an injected column/table escape the interpolated SQL. + for bad in [ + "1; DROP TABLE Users--", + "id = 1", + "foo'bar", + "foo\"bar", + "foo--bar", + "foo;bar", + ] { + assert!( + validate_bulk_column_identifier(bad).is_err(), + "expected column {bad:?} to be rejected", + ); + assert!( + validate_bulk_table_identifier(bad).is_err(), + "expected table {bad:?} to be rejected", + ); + } + } + + #[test] + fn accepts_star_and_bracketed_statement_breakers() { + // `*` must stay valid as a whole column list (bulk_insert uses `&["*"]`), + // and characters that would be statement-breakers outside brackets are + // fine when properly bracket-quoted. + assert!(validate_bulk_column_identifier("*").is_ok()); + assert!(validate_bulk_column_identifier("[weird;name]").is_ok()); + assert!(validate_bulk_table_identifier("[dbo].[my;table]").is_ok()); + } + + #[test] + fn rejects_top_level_space_but_allows_it_inside_parens_and_brackets() { + // A top-level space lets a single identifier split into multiple SQL + // tokens without any otherwise-blocked character, so it is rejected. + assert!(validate_bulk_table_identifier("t UNION SELECT name FROM sysobjects").is_err()); + assert!(validate_bulk_table_identifier("t WHERE 1=1").is_err()); + assert!(validate_bulk_column_identifier("a b").is_err()); + + // A space is still fine inside a parameterized type's parens... + assert!(super::validate_sql_identifier("db_type", "decimal(18, 4)").is_ok()); + // ...and inside a bracket-quoted name. + assert!(validate_bulk_table_identifier("[my table]").is_ok()); + // The canonical no-space forms keep working. + assert!(super::validate_sql_identifier("db_type", "decimal(10,2)").is_ok()); + assert!(super::validate_sql_identifier("db_type", "varchar(max)").is_ok()); + } + + #[test] + fn rejects_delimiter_adjacent_token_splicing() { + // A `]` or `)` is itself a SQL token boundary, so an attacker can splice + // a new token onto one without any (now-blocked) top-level space. Both + // forms must be rejected. + assert!( + validate_bulk_table_identifier("[RealTable]UNION(SELECT secret FROM users)").is_err() + ); + assert!(validate_bulk_table_identifier("[t]UNION").is_err()); + assert!(validate_bulk_table_identifier("[t](1)").is_err()); + assert!(super::validate_sql_identifier("t", "foo(1)UNION(SELECT x)").is_err()); + assert!(super::validate_sql_identifier("t", "decimal(10,2)x").is_err()); + + // ...while every legitimate multi-part / parameterized form still passes. + assert!(validate_bulk_table_identifier("[dbo].[my table]").is_ok()); + assert!(validate_bulk_table_identifier("[db].[schema].[tbl]").is_ok()); + assert!(validate_bulk_table_identifier("[weird]]name]").is_ok()); + assert!(super::validate_sql_identifier("db_type", "decimal(18,4)").is_ok()); + assert!(super::validate_sql_identifier("db_type", "numeric(38, 38)").is_ok()); + assert!(super::validate_sql_identifier("db_type", "varchar(max)").is_ok()); + } + + #[test] + fn rejects_unbalanced_closing_paren_and_top_level_comma() { + // Order-hint columns are interpolated verbatim into `ORDER( ... )`, so a + // value the guard accepts must not be able to close that paren early or + // splice a second column: a lone trailing `)` (no matching `(`) and a + // top-level comma are both rejected. + assert!(validate_bulk_column_identifier("col)").is_err()); + assert!(validate_bulk_column_identifier(")").is_err()); + assert!(validate_bulk_column_identifier("a,b").is_err()); + assert!(validate_bulk_table_identifier("t)").is_err()); + assert!(super::validate_sql_identifier("db_type", "decimal(10,2))").is_err()); + + // The symmetric openers must also be rejected: an unbalanced `(` never + // closed, and an unterminated `[` that would swallow trailing SQL when + // interpolated verbatim into `ORDER( ... )`. + assert!(validate_bulk_column_identifier("a(b").is_err()); + assert!(validate_bulk_column_identifier("(").is_err()); + assert!(validate_bulk_column_identifier("[abc").is_err()); + assert!(validate_bulk_table_identifier("dbo.[my table").is_err()); + assert!(super::validate_sql_identifier("db_type", "decimal(10,2").is_err()); + + // Balanced parens with an interior comma/space are still fine. + assert!(super::validate_sql_identifier("db_type", "decimal(10,2)").is_ok()); + assert!(super::validate_sql_identifier("db_type", "numeric(18, 4)").is_ok()); + // A bracket-quoted name may still contain `)`/`,` literally. + assert!(validate_bulk_column_identifier("[weird,name)]").is_ok()); + } +} + +// Server-free tests for the column collations declared in `INSERT BULK`. +#[cfg(test)] +mod bulk_collation_tests { + use super::{bulk_column_list, table_collations_sql, validate_collation_name}; + use crate::tds::codec::{ + BaseMetaDataColumn, ColumnFlag, FixedLenType, TypeInfo, VarLenContext, VarLenType, + }; + use crate::tds::Collation; + use crate::{FeatureLevel, MetaDataColumn}; + use std::borrow::Cow; + + fn column(name: &'static str, ty: TypeInfo) -> MetaDataColumn<'static> { + MetaDataColumn { + base: BaseMetaDataColumn { + flags: ColumnFlag::Updateable.into(), + ty, + table_name: None, + }, + col_name: Cow::Borrowed(name), + } + } + + fn var(ty: VarLenType, len: usize) -> TypeInfo { + // Latin1_General_CI_AS; the declared name comes from the server's + // collation list, not from these bytes. + let collation = match ty { + VarLenType::BigChar + | VarLenType::BigVarChar + | VarLenType::Text + | VarLenType::NChar + | VarLenType::NVarchar + | VarLenType::NText => Some(Collation::new(0x00d0_0409, 52)), + _ => None, + }; + TypeInfo::VarLenSized(VarLenContext::new(ty, len, collation)) + } + + fn collations(rows: &[(&str, Option<&str>)]) -> Vec<(String, Option)> { + rows.iter() + .map(|(name, collation)| (name.to_string(), collation.map(str::to_string))) + .collect() + } + + #[test] + fn character_columns_declare_their_collation() { + let columns = [ + column("id", TypeInfo::FixedLen(FixedLenType::Int4)), + column("c", var(VarLenType::BigChar, 6)), + column("v", var(VarLenType::BigVarChar, 20)), + column("vmax", var(VarLenType::BigVarChar, 0xffff)), + column("t", var(VarLenType::Text, 0x7fff_ffff)), + column("nc", var(VarLenType::NChar, 12)), + column("nv", var(VarLenType::NVarchar, 40)), + column("nt", var(VarLenType::NText, 0x7fff_ffff)), + column("b", var(VarLenType::BigVarBin, 16)), + column("img", var(VarLenType::Image, 0x7fff_ffff)), + ]; + let rows = collations(&[ + ("id", None), + ("c", Some("Cyrillic_General_CI_AS")), + ("v", Some("Cyrillic_General_CI_AS")), + ("vmax", Some("Chinese_PRC_CI_AS")), + ("t", Some("Cyrillic_General_CI_AS")), + ("nc", Some("Latin1_General_100_CI_AS_SC")), + ("nv", Some("Latin1_General_100_CI_AS_SC")), + ("nt", Some("Latin1_General_100_CI_AS_SC")), + ("b", None), + ("img", None), + ]); + + assert_eq!( + bulk_column_list(&columns, &rows).unwrap(), + "[id] int, \ + [c] char(6) COLLATE Cyrillic_General_CI_AS, \ + [v] varchar(20) COLLATE Cyrillic_General_CI_AS, \ + [vmax] varchar(max) COLLATE Chinese_PRC_CI_AS, \ + [t] text COLLATE Cyrillic_General_CI_AS, \ + [nc] nchar(12) COLLATE Latin1_General_100_CI_AS_SC, \ + [nv] nvarchar(40) COLLATE Latin1_General_100_CI_AS_SC, \ + [nt] ntext COLLATE Latin1_General_100_CI_AS_SC, \ + [b] varbinary(16), \ + [img] image" + ); + } + + #[test] + fn non_character_columns_never_declare_a_collation() { + // Even if the server listed a collation for them. + let columns = [ + column("id", TypeInfo::FixedLen(FixedLenType::Int4)), + column("b", var(VarLenType::BigVarBin, 16)), + ]; + let rows = collations(&[ + ("id", Some("Latin1_General_CI_AS")), + ("b", Some("Latin1_General_CI_AS")), + ]); + + assert_eq!( + bulk_column_list(&columns, &rows).unwrap(), + "[id] int, [b] varbinary(16)" + ); + } + + #[test] + fn collations_are_matched_by_column_name() { + let columns = [ + column("B", var(VarLenType::BigVarChar, 10)), + column("a", var(VarLenType::BigVarChar, 10)), + column("missing", var(VarLenType::BigVarChar, 10)), + column("no_collation", var(VarLenType::BigVarChar, 10)), + ]; + // Out of order, a case-insensitive match for `B`, an exact match for + // `a` that wins over `A`, no row for `missing` and a NULL collation. + let rows = collations(&[ + ("A", Some("Greek_CI_AS")), + ("a", Some("Cyrillic_General_CI_AS")), + ("b", Some("Hebrew_CI_AS")), + ("no_collation", None), + ]); + + assert_eq!( + bulk_column_list(&columns, &rows).unwrap(), + "[B] varchar(10) COLLATE Hebrew_CI_AS, \ + [a] varchar(10) COLLATE Cyrillic_General_CI_AS, \ + [missing] varchar(10), \ + [no_collation] varchar(10)" + ); + } + + #[test] + fn bad_collation_names_are_rejected() { + for name in [ + "", + "Latin1_General_CI_AS; DROP TABLE t", + "Latin1 General", + "Latin1_General_CI_AS'", + "[Latin1_General_CI_AS]", + "Latin1-General", + "Latin1_General_CI_AS\0", + "Кириллица", + ] { + assert!(validate_collation_name(name).is_err(), "{name:?}"); + } + for name in [ + "Latin1_General_CI_AS", + "SQL_Latin1_General_CP1_CI_AS", + "Latin1_General_100_CI_AS_SC_UTF8", + ] { + assert!(validate_collation_name(name).is_ok(), "{name:?}"); + } + + let columns = [column("v", var(VarLenType::BigVarChar, 10))]; + let rows = collations(&[("v", Some("Latin1_General_CI_AS) --"))]); + assert!(matches!( + bulk_column_list(&columns, &rows), + Err(crate::Error::Protocol(_)) + )); + } + + #[test] + fn table_collations_query_matches_sqlclient() { + let v = FeatureLevel::SqlServerN; + assert_eq!( + table_collations_sql("t", v), + "EXEC ..sp_tablecollations_100 N'.[t]'" + ); + assert_eq!( + table_collations_sql("dbo.t", v), + "EXEC ..sp_tablecollations_100 N'[dbo].[t]'" + ); + assert_eq!( + table_collations_sql("[my db].[dbo].[it's]]x]", v), + "EXEC [my db]..sp_tablecollations_100 N'[dbo].[it''s]]x]'" + ); + assert_eq!( + table_collations_sql("srv.db.s.t", v), + "EXEC [db]..sp_tablecollations_100 N'[s].[t]'" + ); + assert_eq!( + table_collations_sql("db..t", v), + "EXEC [db]..sp_tablecollations_100 N'.[t]'" + ); + } + + #[test] + fn table_collations_query_reads_temp_tables_from_tempdb() { + let v = FeatureLevel::SqlServerN; + assert_eq!( + table_collations_sql("#t", v), + "EXEC tempdb..sp_tablecollations_100 N'.[#t]'" + ); + assert_eq!( + table_collations_sql("##t", v), + "EXEC tempdb..sp_tablecollations_100 N'.[##t]'" + ); + assert_eq!( + table_collations_sql("[#t]", v), + "EXEC tempdb..sp_tablecollations_100 N'.[#t]'" + ); + assert_eq!( + table_collations_sql("tempdb..#t", v), + "EXEC [tempdb]..sp_tablecollations_100 N'.[#t]'" + ); + } + + #[test] + fn table_collations_query_uses_the_2005_procedure_before_2008() { + assert_eq!( + table_collations_sql("t", FeatureLevel::SqlServer2005), + "EXEC ..sp_tablecollations_90 N'.[t]'" + ); + assert_eq!( + table_collations_sql("t", FeatureLevel::SqlServer2008), + "EXEC ..sp_tablecollations_100 N'.[t]'" + ); + } } diff --git a/src/client/auth.rs b/src/client/auth.rs index 2003d99ba..a898eacb2 100644 --- a/src/client/auth.rs +++ b/src/client/auth.rs @@ -1,59 +1,81 @@ -use std::fmt::Debug; +use secrecy::{ExposeSecret, SecretString}; +use zeroize::Zeroize; -#[derive(Clone, PartialEq, Eq)] +/// Build a `SecretString` from owned bytes without leaving an un-zeroized copy +/// in freed heap. `SecretString::from(String)` routes through +/// `String::into_boxed_str()`, which reallocates and frees the source buffer +/// *without zeroizing* when the string has spare capacity. Copy the bytes into +/// an exact-sized `Box` and wipe the original. +pub(crate) fn secret_from_string(mut s: String) -> SecretString { + let boxed: Box = s.as_str().into(); + s.zeroize(); + SecretString::new(boxed) +} + +// Credentials are stored as `secrecy::SecretString`, which zeroizes the +// plaintext on drop and redacts it from `Debug`, so these types can derive +// `Debug` and still never print a secret. `SecretString` does not implement +// `PartialEq`/`Eq` (comparing secrets is deliberately opt-in), so the equality +// impls below are hand-written; they compare the exposed plaintext to preserve +// the previous derived behaviour (and the public `AuthMethod: Eq` bound). +#[derive(Clone, Debug)] pub struct SqlServerAuth { user: String, - password: String, + password: SecretString, } impl SqlServerAuth { - pub(crate) fn user(&self) -> &str { - &self.user - } - - pub(crate) fn password(&self) -> &str { - &self.password + pub(crate) fn into_credentials(self) -> (String, SecretString) { + (self.user, self.password) } } -impl Debug for SqlServerAuth { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("SqlServerAuth") - .field("user", &self.user) - .field("password", &"") - .finish() +impl PartialEq for SqlServerAuth { + fn eq(&self, other: &Self) -> bool { + self.user == other.user && self.password.expose_secret() == other.password.expose_secret() } } -#[derive(Clone, PartialEq, Eq)] -#[cfg(any(feature = "winauth", doc))] -#[cfg_attr(feature = "docs", doc(feature = "winauth"))] +impl Eq for SqlServerAuth {} + +#[derive(Clone, Debug)] +#[cfg(any(feature = "winauth", all(unix, feature = "sspi-rs"), doc))] +#[cfg_attr( + docsrs, + doc(cfg(any(feature = "winauth", all(unix, feature = "sspi-rs")))) +)] pub struct WindowsAuth { pub(crate) user: String, - pub(crate) password: String, + pub(crate) password: SecretString, pub(crate) domain: Option, } -#[cfg(any(feature = "winauth", doc))] -#[cfg_attr(feature = "docs", doc(feature = "winauth"))] -impl Debug for WindowsAuth { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("WindowsAuth") - .field("user", &self.user) - .field("password", &"") - .field("domain", &self.domain) - .finish() +#[cfg(any(feature = "winauth", all(unix, feature = "sspi-rs"), doc))] +impl PartialEq for WindowsAuth { + fn eq(&self, other: &Self) -> bool { + self.user == other.user + && self.domain == other.domain + && self.password.expose_secret() == other.password.expose_secret() } } +#[cfg(any(feature = "winauth", all(unix, feature = "sspi-rs"), doc))] +impl Eq for WindowsAuth {} + /// Defines the method of authentication to the server. -#[derive(Clone, Debug, PartialEq, Eq)] +#[derive(Clone, Debug)] pub enum AuthMethod { /// Authenticate directly with SQL Server. SqlServer(SqlServerAuth), - /// Authenticate with Windows credentials. - #[cfg(any(feature = "winauth", doc))] - #[cfg_attr(feature = "docs", doc(cfg(feature = "winauth")))] + /// Authenticate with Windows credentials (NTLMv2). The `winauth` feature + /// provides it on every platform via its pure-Rust NTLMv2 client; on Unix + /// the `sspi-rs` feature provides it via sspi-rs and takes precedence when + /// both are enabled. + #[cfg(any(feature = "winauth", all(unix, feature = "sspi-rs"), doc))] + #[cfg_attr( + docsrs, + doc(cfg(any(feature = "winauth", all(unix, feature = "sspi-rs")))) + )] Windows(WindowsAuth), /// Authenticate as the currently logged in user. On Windows uses SSPI and /// Kerberos on Unix platforms. @@ -63,29 +85,56 @@ pub enum AuthMethod { doc ))] #[cfg_attr( - feature = "docs", + docsrs, doc(cfg(any(windows, all(unix, feature = "integrated-auth-gssapi")))) )] Integrated, /// Authenticate with an AAD token. The token should encode an AAD user/service principal /// which has access to SQL Server. - AADToken(String), + AADToken(SecretString), #[doc(hidden)] None, } +// `SecretString` has no `PartialEq`, so `AuthMethod`'s public equality is +// hand-written. It mirrors the old derived behaviour: same variant + equal +// fields (secrets compared via their exposed plaintext). +impl PartialEq for AuthMethod { + fn eq(&self, other: &Self) -> bool { + match (self, other) { + (Self::SqlServer(a), Self::SqlServer(b)) => a == b, + #[cfg(any(feature = "winauth", all(unix, feature = "sspi-rs"), doc))] + (Self::Windows(a), Self::Windows(b)) => a == b, + #[cfg(any( + all(windows, feature = "winauth"), + all(unix, feature = "integrated-auth-gssapi"), + doc + ))] + (Self::Integrated, Self::Integrated) => true, + (Self::AADToken(a), Self::AADToken(b)) => a.expose_secret() == b.expose_secret(), + (Self::None, Self::None) => true, + _ => false, + } + } +} + +impl Eq for AuthMethod {} + impl AuthMethod { /// Construct a new SQL Server authentication configuration. pub fn sql_server(user: impl ToString, password: impl ToString) -> Self { Self::SqlServer(SqlServerAuth { user: user.to_string(), - password: password.to_string(), + password: secret_from_string(password.to_string()), }) } /// Construct a new Windows authentication configuration. - #[cfg(any(feature = "winauth", doc))] - #[cfg_attr(feature = "docs", doc(cfg(feature = "winauth")))] + #[cfg(any(feature = "winauth", all(unix, feature = "sspi-rs"), doc))] + #[cfg_attr( + docsrs, + doc(cfg(any(feature = "winauth", all(unix, feature = "sspi-rs")))) + )] pub fn windows(user: impl AsRef, password: impl ToString) -> Self { let (domain, user) = match user.as_ref().find('\\') { Some(idx) => (Some(&user.as_ref()[..idx]), &user.as_ref()[idx + 1..]), @@ -94,13 +143,108 @@ impl AuthMethod { Self::Windows(WindowsAuth { user: user.to_string(), - password: password.to_string(), + password: secret_from_string(password.to_string()), domain: domain.map(|s| s.to_string()), }) } /// Construct a new configuration with AAD auth token. pub fn aad_token(token: impl ToString) -> Self { - Self::AADToken(token.to_string()) + Self::AADToken(secret_from_string(token.to_string())) + } +} + +#[cfg(test)] +mod tests { + use super::AuthMethod; + use secrecy::ExposeSecret; + + // Compile-time proof that the stored credential zeroizes its plaintext on + // drop: `secrecy::SecretString` implements `ZeroizeOnDrop`. + #[test] + fn stored_credentials_are_zeroize_on_drop() { + fn assert_zeroize_on_drop() {} + assert_zeroize_on_drop::(); + } + + #[test] + fn sql_server_password_can_be_consumed_and_exposed() { + let AuthMethod::SqlServer(auth) = AuthMethod::sql_server("sa", "secret") else { + unreachable!(); + }; + + let (user, password) = auth.into_credentials(); + + assert_eq!("sa", user); + // `expose_secret()` yields exactly the plaintext that was provided. + assert_eq!("secret", password.expose_secret()); + // The credential is dropped (and zeroized) at the end of this scope. + } + + #[test] + fn aad_token_exposes_the_right_value() { + let AuthMethod::AADToken(token) = AuthMethod::aad_token("aad-secret-token") else { + unreachable!(); + }; + assert_eq!("aad-secret-token", token.expose_secret()); + } + + #[test] + fn debug_redacts_credentials() { + let sql = format!("{:?}", AuthMethod::sql_server("sa", "sql-secret")); + assert!(!sql.contains("sql-secret"), "SQL password leaked: {sql}"); + // The non-secret user is still visible for diagnostics. + assert!(sql.contains("sa"), "user should be shown: {sql}"); + assert!(sql.contains("REDACTED"), "password not redacted: {sql}"); + + let aad = format!("{:?}", AuthMethod::aad_token("aad-secret-token")); + assert!(!aad.contains("aad-secret-token"), "AAD token leaked: {aad}"); + assert!(aad.contains("REDACTED"), "AAD token not redacted: {aad}"); + } + + #[test] + fn secret_from_string_preserves_value_with_spare_capacity() { + // A `String` with spare capacity is exactly the input shape that would + // trigger the leaky `SecretString::from(String)` reallocation path. This + // test verifies the helper preserves the value for that shape. The + // zeroization of the freed source is guaranteed by construction + // (`s.zeroize()` before drop) and is not directly unit-observable in + // safe Rust, so it is not asserted here. + let mut s = String::with_capacity(64); + s.push_str("pw"); + assert_eq!(super::secret_from_string(s).expose_secret(), "pw"); + } + + #[test] + fn debug_none_variant() { + assert_eq!(format!("{:?}", AuthMethod::None), "None"); + } + + #[cfg(any(feature = "winauth", all(unix, feature = "sspi-rs")))] + #[test] + fn windows_auth_parses_domain_and_debug_redacts() { + // `DOMAIN\user` form exercises the domain-splitting branch of `windows()`. + let auth = AuthMethod::windows("DOMAIN\\user", "win-secret"); + let dbg = format!("{:?}", auth); + assert!(dbg.contains("Windows"), "variant name missing: {dbg}"); + assert!(dbg.contains("DOMAIN"), "domain not preserved: {dbg}"); + assert!(dbg.contains("user"), "user not preserved: {dbg}"); + assert!(!dbg.contains("win-secret"), "password leaked: {dbg}"); + assert!(dbg.contains("REDACTED"), "password not redacted: {dbg}"); + + // No backslash exercises the domain-less branch. + let plain = AuthMethod::windows("plainuser", "pw"); + let dbg = format!("{:?}", plain); + assert!(dbg.contains("plainuser"), "user not preserved: {dbg}"); + assert!(dbg.contains("None"), "domain should be None: {dbg}"); + } + + #[cfg(any( + all(windows, feature = "winauth"), + all(unix, feature = "integrated-auth-gssapi") + ))] + #[test] + fn integrated_debug() { + assert_eq!(format!("{:?}", AuthMethod::Integrated), "Integrated"); } } diff --git a/src/client/config.rs b/src/client/config.rs index fff68bc15..db2ffd96f 100644 --- a/src/client/config.rs +++ b/src/client/config.rs @@ -1,13 +1,30 @@ mod ado_net; +mod ado_parser; mod jdbc; use std::collections::HashMap; use std::path::PathBuf; +use std::time::Duration; + +/// Default upper bound on how long the connection handshake (TDS prelogin, TLS +/// negotiation and login) may take before [`Connection::connect`] gives up with +/// a timeout error. Matches the 15-second `Connect Timeout` default used by +/// ADO.NET / the Microsoft SQL Server drivers. +/// +/// [`Connection::connect`]: crate::Client::connect +const DEFAULT_HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(15); + +/// Default per-response deadline applied while reading command results (see +/// [`Config::command_timeout`]). Matches the 30-second `Command Timeout` / +/// `CommandTimeout` default used by ADO.NET / the Microsoft SQL Server drivers. +const DEFAULT_COMMAND_TIMEOUT: Duration = Duration::from_secs(30); use super::AuthMethod; use crate::EncryptionLevel; use ado_net::*; use jdbc::*; +#[cfg(any(feature = "native-tls", feature = "vendored-openssl"))] +use secrecy::SecretString; #[derive(Clone, Debug)] /// The `Config` struct contains all configuration information @@ -18,10 +35,17 @@ use jdbc::*; /// When using an [ADO.NET connection string], it can be /// constructed using the [`from_ado_string`] function. /// +/// Alternatively, a [`ConfigBuilder`] can be used for an ergonomic, +/// chainable construction. Create one via [`builder`], call its +/// setter methods and finalize it with [`build`]. +/// /// [`Client`]: struct.Client.html /// [ADO.NET connection string]: https://docs.microsoft.com/en-us/dotnet/framework/data/adonet/connection-strings /// [`from_ado_string`]: struct.Config.html#method.from_ado_string /// [`get_addr`]: struct.Config.html#method.get_addr +/// [`ConfigBuilder`]: struct.ConfigBuilder.html +/// [`builder`]: struct.Config.html#method.builder +/// [`build`]: struct.ConfigBuilder.html#method.build pub struct Config { pub(crate) host: Option, pub(crate) port: Option, @@ -32,14 +56,175 @@ pub struct Config { pub(crate) trust: TrustConfig, pub(crate) auth: AuthMethod, pub(crate) readonly: bool, + pub(crate) packet_size: Option, + pub(crate) hostname_in_certificate: Option, + pub(crate) client_name: Option, + pub(crate) multi_subnet_failover: bool, + pub(crate) handshake_timeout: Option, + pub(crate) command_timeout: Option, + pub(crate) lossy_utf16_decoding: bool, + pub(crate) lossy_codepage_decoding: bool, + #[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + ))] + pub(crate) client_cert: Option, +} + +/// How the server certificate is trusted. +/// +/// Two orthogonal axes plus a bypass: +/// +/// - [`source`](TrustConfig::source): the base set of trust anchors (mutually +/// exclusive — the OS store or a bundled Mozilla snapshot). +/// - [`extra_cas`](TrustConfig::extra_cas): additional CA certificates layered +/// *on top of* the source. This **accumulates**: every `trust_cert_ca` / +/// `trust_cert_ca_bundle` call appends one entry (composable) rather than +/// replacing the previous one. +/// - [`bypass`](TrustConfig::bypass): skip certificate validation entirely +/// (`trust_cert`). Mutually exclusive with configuring a `source` or +/// `extra_cas` — mixing them panics (or, from a connection string, errors) +/// rather than letting the bypass silently win. +/// +/// `bypass` is deliberately a distinct axis (not folded into `source`) so a +/// future, narrower "relax hostname verification only" mode can be +/// added later without a public-API rewrite. +#[derive(Clone)] +pub(crate) struct TrustConfig { + /// Base trust anchors. Mutually exclusive; defaults to [`RootSource::Native`]. + pub(crate) source: RootSource, + /// Extra CA certificates added on top of `source`. Accumulates in call order. + pub(crate) extra_cas: Vec, + /// If set, certificate validation (and hostname verification) is skipped + /// entirely. See [`Config::trust_cert`]. + pub(crate) bypass: bool, +} + +impl Default for TrustConfig { + fn default() -> Self { + Self { + source: RootSource::Native, + extra_cas: Vec::new(), + bypass: false, + } + } +} + +impl TrustConfig { + /// True when nothing has been customised: the OS trust store, no extra CAs, + /// no bypass. Used by tests to assert the default trust posture. + #[cfg(test)] + pub(crate) fn is_default(&self) -> bool { + matches!(self.source, RootSource::Native) && self.extra_cas.is_empty() && !self.bypass + } +} + +// Manual `Debug` following the crate's curated-Debug convention: never dump raw +// certificate bytes. `ExtraCa::Bundle` is summarised as `Bundle { len: N }`. +impl std::fmt::Debug for TrustConfig { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("TrustConfig") + .field("source", &self.source) + .field("extra_cas", &self.extra_cas) + .field("bypass", &self.bypass) + .finish() + } +} + +/// The base set of trust anchors used to validate the server certificate. +/// Mutually exclusive. +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub(crate) enum RootSource { + /// The operating system's trust store (the default). + #[default] + Native, + /// A compiled-in snapshot of Mozilla's root CA store (`webpki-roots`). + /// + /// rustls-only, and gated on the `rustls-webpki-roots` feature so the + /// variant cannot exist without the crate that consumes it. Selected via + /// [`Config::trust_webpki_roots`]. + #[cfg(feature = "rustls-webpki-roots")] + WebpkiRoots, +} + +/// An additional CA certificate source layered on top of [`RootSource`]. +#[derive(Clone)] +#[cfg_attr( + not(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + )), + allow(dead_code) +)] +pub(crate) enum ExtraCa { + /// A certificate file on disk (PEM `.pem`/`.crt`, or DER `.der`). May hold + /// multiple certificates when PEM. + File(PathBuf), + /// In-memory certificate bytes. Multi-certificate: sniffed as PEM (all + /// blocks parsed) when the bytes contain a `-----BEGIN` marker at the start + /// of a line, otherwise treated as a single DER certificate. + Bundle(Vec), +} + +// Manual `Debug`: summarise `Bundle` as `Bundle { len: N }` so raw certificate +// bytes are never dumped. +impl std::fmt::Debug for ExtraCa { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + ExtraCa::File(path) => f.debug_tuple("File").field(path).finish(), + ExtraCa::Bundle(bytes) => f.debug_struct("Bundle").field("len", &bytes.len()).finish(), + } + } +} + +/// A client certificate and its private key, presented to the server during the +/// TLS handshake to authenticate the *client* (mutual TLS / TDS 8.0 +/// `ENCRYPT_CLIENT_CERT`). +/// +/// Construct one indirectly via [`Config::client_certificate`] (PEM/DER +/// certificate + private-key files) or [`Config::client_certificate_pkcs12`] +/// (a PKCS#12 / PFX bundle, `native-tls` and `vendored-openssl` only). +#[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" +))] +#[derive(Clone, Debug)] +pub(crate) struct ClientCertificate { + pub(crate) source: ClientCertSource, } +#[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" +))] +// `SecretString` redacts the PKCS#12 password from `Debug` (and zeroizes it on +// drop), so this can derive `Debug` without leaking the password. #[derive(Clone, Debug)] -pub(crate) enum TrustConfig { - #[allow(dead_code)] - CaCertificateLocation(PathBuf), - TrustAll, - Default, +pub(crate) enum ClientCertSource { + /// A certificate file and a separate private-key file. Both may be PEM + /// (`.pem`/`.crt` for the certificate, `.pem`/`.key` for the key) or DER + /// (`.der`); the concrete format is detected from the file extension by the + /// active TLS backend. + /// + /// The `cert`/`key` fields are read only by the `rustls` and `native-tls` + /// backends; the `vendored-openssl` (opentls) backend rejects this variant + /// (it supports PKCS#12 only). In a `vendored-openssl`-only build the fields + /// are therefore never read, so silence dead-code exactly there rather than + /// unconditionally — this keeps the derived `Debug` (which no longer counts + /// as a read) while still compiling clean under `-Dwarnings`. + #[cfg_attr(not(any(feature = "rustls", feature = "native-tls")), allow(dead_code))] + CertAndKey { cert: PathBuf, key: PathBuf }, + /// A PKCS#12 / PFX bundle path together with its decryption password. Only + /// supported by the `native-tls` and `vendored-openssl` backends. + #[cfg(any(feature = "native-tls", feature = "vendored-openssl"))] + Pkcs12 { + path: PathBuf, + password: SecretString, + }, } impl Default for Config { @@ -62,9 +247,23 @@ impl Default for Config { feature = "vendored-openssl" )))] encryption: EncryptionLevel::NotSupported, - trust: TrustConfig::Default, + trust: TrustConfig::default(), auth: AuthMethod::None, readonly: false, + packet_size: None, + hostname_in_certificate: None, + client_name: None, + multi_subnet_failover: false, + handshake_timeout: Some(DEFAULT_HANDSHAKE_TIMEOUT), + command_timeout: Some(DEFAULT_COMMAND_TIMEOUT), + lossy_utf16_decoding: false, + lossy_codepage_decoding: false, + #[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + ))] + client_cert: None, } } } @@ -75,6 +274,33 @@ impl Config { Self::default() } + /// Create a new [`ConfigBuilder`] initialized with the default settings. + /// + /// This provides an ergonomic, chainable alternative to constructing a + /// [`Config`] via its individual setter methods. + /// + /// # Example + /// + /// ``` + /// # use tiberius::{Config, AuthMethod}; + /// let config = Config::builder() + /// .host("localhost") + /// .port(1433) + /// .database("master") + /// .authentication(AuthMethod::sql_server("SA", "")) + /// .build(); + /// + /// assert_eq!("localhost:1433", config.get_addr()); + /// ``` + /// + /// [`ConfigBuilder`]: struct.ConfigBuilder.html + /// [`Config`]: struct.Config.html + pub fn builder() -> ConfigBuilder { + ConfigBuilder { + inner: Self::default(), + } + } + /// A host or ip address to connect to. /// /// - Defaults to `localhost`. @@ -115,10 +341,27 @@ impl Config { self.application_name = Some(name.to_string()); } + /// Sets the TDS packet size for the connection. + /// + /// Larger packet sizes can improve bulk insert performance by reducing + /// the number of network round-trips. Valid values are 512 to 32767. + /// The server may negotiate a different size. + /// + /// - Defaults to 4096 bytes. + pub fn packet_size(&mut self, size: u32) { + self.packet_size = Some(size); + } + + /// Gets the configured packet size, if set. + pub fn get_packet_size(&self) -> Option { + self.packet_size + } + /// Set the preferred encryption level. /// - /// - With `tls` feature, defaults to `Required`. - /// - Without `tls` feature, defaults to `NotSupported`. + /// - With a TLS backend enabled (any of the `rustls`, `native-tls`, or + /// `vendored-openssl` features), defaults to `Required`. + /// - Without a TLS backend, defaults to `NotSupported`. pub fn encryption(&mut self, encryption: EncryptionLevel) { self.encryption = encryption; } @@ -129,32 +372,150 @@ impl Config { /// On production setting, the certificate should be added to the local key /// storage (or use `trust_cert_ca` instead), using this setting is potentially dangerous. /// + /// This also bypasses the rustls backend's certificate *version* checks, so + /// it is the escape hatch for the `invalid peer certificate: + /// UnsupportedCertVersion` failure some older or self-signed SQL Server + /// certificates trigger. Prefer [`trust_cert_ca`] when you + /// can point at the server's CA; reach for `trust_cert` only when you cannot. + /// Note that SQL Server performs a TLS handshake during login even when + /// `Encrypt=false`, so certificate errors can surface regardless of the + /// encryption level. + /// + /// # Security + /// + /// This accepts **any** certificate: it disables both certificate-chain + /// validation *and* hostname verification, so it provides no protection + /// against a man-in-the-middle. Only use it on a network you already trust. + /// Prefer [`trust_cert_ca`] / [`trust_cert_ca_bundle`] to pin the server's + /// CA instead. + /// + /// [`trust_cert_ca`]: Self::trust_cert_ca + /// [`trust_cert_ca_bundle`]: Self::trust_cert_ca_bundle + /// /// # Panics - /// Will panic in case `trust_cert_ca` was called before. + /// Will panic if any of [`trust_cert_ca`], [`trust_cert_ca_bundle`] or + /// [`trust_webpki_roots`] was called before — the trust bypass is mutually + /// exclusive with configuring trust anchors. + /// + /// [`trust_webpki_roots`]: Self::trust_webpki_roots /// /// - Defaults to `default`, meaning server certificate is validated against system-truststore. pub fn trust_cert(&mut self) { - if let TrustConfig::CaCertificateLocation(_) = &self.trust { - panic!("'trust_cert' and 'trust_cert_ca' are mutual exclusive! Only use one.") + if !self.trust.extra_cas.is_empty() || !matches!(self.trust.source, RootSource::Native) { + panic!( + "'trust_cert' and 'trust_cert_ca'/'trust_cert_ca_bundle'/'trust_webpki_roots' \ + are mutual exclusive! Only use one." + ) } - self.trust = TrustConfig::TrustAll; + self.trust.bypass = true; } - /// If set, the server certificate will be validated against the given CA certificate in - /// in addition to the system-truststore. + /// Trust an additional CA certificate (from a file) *in addition to* the + /// base trust anchors (the system trust store by default). /// Useful when using self-signed certificates on the server without having to disable the /// trust-chain. /// + /// The file may be PEM (`.pem`/`.crt`) or DER (`.der`); a PEM file may + /// contain **multiple** certificates and all of them are trusted. + /// + /// This **accumulates**: calling it more than once (or alongside + /// [`trust_cert_ca_bundle`]) trusts every supplied CA — repeated calls are + /// additive, not replace-last-wins. + /// /// # Panics - /// Will panic in case `trust_cert` was called before. + /// Will panic in case [`trust_cert`] was called before. + /// + /// [`trust_cert`]: Self::trust_cert + /// [`trust_cert_ca_bundle`]: Self::trust_cert_ca_bundle /// /// - Defaults to validating the server certificate is validated against system's certificate storage. pub fn trust_cert_ca(&mut self, path: impl ToString) { - if let TrustConfig::TrustAll = &self.trust { + if self.trust.bypass { panic!("'trust_cert' and 'trust_cert_ca' are mutual exclusive! Only use one.") - } else { - self.trust = TrustConfig::CaCertificateLocation(PathBuf::from(path.to_string())) } + self.trust + .extra_cas + .push(ExtraCa::File(PathBuf::from(path.to_string()))); + } + + /// Trust additional CA certificates supplied as in-memory bytes, *in + /// addition to* the base trust anchors. This avoids having to write an + /// in-memory certificate out to a temporary file, and accepts a whole + /// **bundle** of CA certificates (for example the AWS RDS root bundle). + /// + /// The byte format is auto-detected: if the bytes contain a `-----BEGIN` + /// marker at the start of a line they are parsed as PEM (every certificate + /// block is trusted), otherwise they are treated as a single DER + /// certificate. A leading UTF-8 byte-order mark is tolerated. + /// + /// Like [`trust_cert_ca`], this **accumulates**: each call appends to the + /// set of trusted CAs. + /// + /// # Panics + /// Will panic in case [`trust_cert`] was called before. + /// + /// [`trust_cert`]: Self::trust_cert + /// [`trust_cert_ca`]: Self::trust_cert_ca + pub fn trust_cert_ca_bundle(&mut self, bundle: impl Into>) { + if self.trust.bypass { + panic!("'trust_cert' and 'trust_cert_ca_bundle' are mutual exclusive! Only use one.") + } + self.trust.extra_cas.push(ExtraCa::Bundle(bundle.into())); + } + + /// Use a compiled-in snapshot of Mozilla's root CA store (via the + /// `webpki-roots` crate) as the base trust anchors instead of the operating + /// system's trust store. Any CAs added with [`trust_cert_ca`] / + /// [`trust_cert_ca_bundle`] are still layered on top. + /// + /// This is useful on platforms with no usable OS trust store. It is only + /// available with the `rustls` backend (via the `rustls-webpki-roots` + /// feature); with any other backend the method does not exist, so misuse is + /// a compile error. + /// + /// # Security + /// + /// The bundled roots are a **pinned snapshot** taken when this crate's + /// `webpki-roots` dependency was last updated. Unlike the OS trust store + /// they do not receive security updates on their own — they go stale (newly + /// distrusted or newly added CAs are missed) unless the dependency is + /// updated and the application rebuilt. + /// + /// # Panics + /// Will panic in case [`trust_cert`] was called before. + /// + /// [`trust_cert`]: Self::trust_cert + /// [`trust_cert_ca`]: Self::trust_cert_ca + /// [`trust_cert_ca_bundle`]: Self::trust_cert_ca_bundle + #[cfg(feature = "rustls-webpki-roots")] + #[cfg_attr(docsrs, doc(cfg(feature = "rustls-webpki-roots")))] + pub fn trust_webpki_roots(&mut self) { + if self.trust.bypass { + panic!("'trust_cert' and 'trust_webpki_roots' are mutual exclusive! Only use one.") + } + self.trust.source = RootSource::WebpkiRoots; + } + + /// Sets the hostname that the server certificate is validated against, + /// instead of the value given to [`host`]. + /// + /// This is useful when connecting through an IP address, a tunnel, or a + /// load balancer whose certificate carries a different subject/SAN than the + /// address used to reach it (see issue #340). + /// + /// - Defaults to the value of [`host`]. + /// + /// [`host`]: Config::host + pub fn hostname_in_certificate(&mut self, hostname: impl ToString) { + self.hostname_in_certificate = Some(hostname.to_string()); + } + + /// Sets the client / workstation name reported to the server in the login + /// record (queryable with `HOST_NAME()`). + /// + /// - Defaults to the local workstation id (the machine hostname). + pub fn client_name(&mut self, name: impl ToString) { + self.client_name = Some(name.to_string()); } /// Sets the authentication method. @@ -167,8 +528,302 @@ impl Config { /// Sets ApplicationIntent readonly. /// /// - Defaults to `false`. - pub fn readonly(&mut self, readnoly: bool) { - self.readonly = readnoly; + pub fn readonly(&mut self, readonly: bool) { + self.readonly = readonly; + } + + /// Enable multi-subnet failover. + /// + /// When enabled and the server host name resolves to more than one IP + /// address (for example, an Always On availability group listener spread + /// across subnets), connections are attempted to all resolved addresses in + /// parallel and the first one to succeed is used. This mirrors the ADO.NET + /// `MultiSubnetFailover` connection-string keyword. + /// + /// - Defaults to `false`. + pub fn multi_subnet_failover(&mut self, multi_subnet_failover: bool) { + self.multi_subnet_failover = multi_subnet_failover; + } + + /// Returns whether multi-subnet failover is enabled. + pub fn get_multi_subnet_failover(&self) -> bool { + self.multi_subnet_failover + } + + /// Sets an upper bound on how long the connection handshake may take. + /// + /// The handshake covers everything [`Client::connect`] does after it is + /// handed a connected TCP stream: the TDS prelogin exchange, the TLS + /// negotiation and the login. If the server accepts the TCP connection but + /// then stops responding mid-handshake (for example a TLS handshake that + /// stalls indefinitely), the connect future would otherwise hang forever + /// with no error. This bound makes such a stall surface as a + /// [`Error::Io`] with [`std::io::ErrorKind::TimedOut`] instead. + /// + /// The timer is runtime-agnostic (it does not depend on tokio or smol), so + /// it applies regardless of the async runtime driving the connection. Note + /// that it does **not** cover establishing the TCP connection itself, which + /// happens before the stream is passed to [`Client::connect`]; bound that + /// with your runtime's own connect timeout (e.g. `tokio::time::timeout`). + /// + /// Enable `tracing` at the `DEBUG` level to see which handshake stage was + /// reached, which pinpoints where a stall occurred. + /// + /// - Defaults to 15 seconds, matching ADO.NET's `Connect Timeout`. Pass + /// `None` to wait indefinitely. A zero duration is a degenerate bound + /// that fails as soon as the handshake would block. + /// + /// # Example + /// + /// ``` + /// # use tiberius::Config; + /// # use std::time::Duration; + /// let mut config = Config::new(); + /// // Fail fast if the handshake stalls. + /// config.handshake_timeout(Some(Duration::from_secs(5))); + /// ``` + /// + /// [`Client::connect`]: crate::Client::connect + /// [`Error::Io`]: crate::error::Error::Io + pub fn handshake_timeout(&mut self, timeout: Option) { + self.handshake_timeout = timeout; + } + + /// Returns the configured connection-handshake timeout, if any. + /// + /// See [`handshake_timeout`](Config::handshake_timeout). + pub fn get_handshake_timeout(&self) -> Option { + self.handshake_timeout + } + + /// Sets a per-response deadline applied while reading command results + /// (`query`, `execute`, `simple_query`, the `bulk_insert` server + /// acknowledgement and `column_metadata`). + /// + /// # Semantics + /// + /// This bounds **each server round-trip**, not the total time to enumerate a + /// result stream. Query results are lazy, caller-driven streams: the timer + /// only runs while the client is actively waiting on the server for the next + /// chunk of the response, and it is **reset every time a token is + /// delivered**. A slow *consumer* (code that pauses between pulling rows) + /// therefore never trips it — only a stalled *server* does. Concretely, if + /// the gap between the client asking for more data and the server delivering + /// the next token exceeds the deadline, the stream yields an [`Error::Io`] + /// with [`std::io::ErrorKind::TimedOut`]. This covers both the + /// time-to-first-response and any mid-stream stall. + /// + /// The timer is runtime-agnostic (it does not depend on tokio or smol). + /// + /// Note: once a command times out, its connection is left mid-response and + /// out of sync with the server; drop the [`Client`] and open a new one (a + /// pool should discard the connection) rather than reuse it. + /// + /// - Defaults to 30 seconds, matching ADO.NET's `Command Timeout`. Pass + /// `None` to wait indefinitely. A zero duration is a degenerate bound + /// that trips on the first server round-trip that is not answered + /// immediately. + /// + /// # Example + /// + /// ``` + /// # use tiberius::Config; + /// # use std::time::Duration; + /// let mut config = Config::new(); + /// // Fail a query if the server stops responding for 10s mid-result. + /// config.command_timeout(Some(Duration::from_secs(10))); + /// // Or wait forever (e.g. for a deliberately long-running batch): + /// config.command_timeout(None); + /// ``` + /// + /// [`Client`]: crate::Client + /// [`Error::Io`]: crate::error::Error::Io + pub fn command_timeout(&mut self, timeout: Option) { + self.command_timeout = timeout; + } + + /// Returns the configured per-response command timeout, if any. + /// + /// See [`command_timeout`](Config::command_timeout). + pub fn get_command_timeout(&self) -> Option { + self.command_timeout + } + + /// Controls how malformed UTF-16 in NVARCHAR/NTEXT row values is handled. + /// + /// SQL Server stores `NVARCHAR`/`NCHAR`/`NTEXT` as unchecked UCS-2/UTF-16, + /// so a column can legitimately hold lone (unpaired) surrogates or other + /// sequences that are not valid Unicode. By default tiberius decodes these + /// values *strictly*: a malformed sequence aborts the row stream with an + /// error — [`Error::Protocol`] for NVARCHAR/NCHAR and [`Error::Utf16`] for + /// NTEXT (an odd, desynced byte length is always [`Error::Protocol`], in + /// both modes; see below). Strict decoding + /// is the safer default because it also surfaces framing desyncs (a decode + /// that silently "succeeds" on garbage can mask a misaligned read). + /// + /// Enabling this option makes decoding *lossy* for the NVARCHAR/NCHAR and + /// NTEXT text arms only: each invalid UTF-16 sequence is replaced with the + /// Unicode replacement character (`U+FFFD`, `�`) instead of erroring, so a + /// row containing bad Unicode stays readable. + /// + /// Scope and guarantees: + /// + /// - Only NVARCHAR/NCHAR (`string`) and NTEXT (`text`) decoding is affected. + /// `XML` columns are always decoded strictly. Code-page + /// (`VARCHAR`/`CHAR`/`TEXT`) columns are controlled independently by + /// [`Config::lossy_codepage_decoding`]. + /// - Framing/length validation is always enforced: an odd byte length for a + /// UTF-16 value is still a protocol error in both modes, because it + /// indicates a desynced stream rather than merely bad Unicode. + /// + /// - Defaults to `false` (strict decoding). + /// + /// # Example + /// + /// ``` + /// # use tiberius::Config; + /// let mut config = Config::new(); + /// // Tolerate legacy rows that hold unchecked UCS-2 with lone surrogates. + /// config.lossy_utf16_decoding(true); + /// ``` + /// + /// [`Error::Protocol`]: crate::error::Error::Protocol + /// [`Error::Utf16`]: crate::error::Error::Utf16 + pub fn lossy_utf16_decoding(&mut self, lossy: bool) { + self.lossy_utf16_decoding = lossy; + } + + /// Returns whether lossy UTF-16 decoding is enabled for NVARCHAR/NTEXT + /// values. + /// + /// See [`lossy_utf16_decoding`](Config::lossy_utf16_decoding). + pub fn get_lossy_utf16_decoding(&self) -> bool { + self.lossy_utf16_decoding + } + + /// Enables lossy decoding for code-page character row values. + /// + /// Defaults to `false`: bytes that are invalid in the column's collation + /// cause an [`Error::Encoding`](crate::Error::Encoding) error. Enabling + /// this option replaces those sequences with `U+FFFD` so legacy rows + /// containing invalid bytes can still be read. + /// + /// Replacement follows `encoding_rs` / WHATWG decoding, not SQL Server's + /// `CONVERT(NVARCHAR)` behavior. For example, under `Chinese_PRC_CI_AS`, + /// bytes `61 81 20 62` decode to `"a\u{fffd} b"` and a lone `FF` decodes + /// to `"\u{fffd}"`; the server may instead use `?` or a private-use character. + /// CP437 and CP850 define every byte, so this option does not change their + /// decoded values. + /// + /// Applies to `CHAR`, `VARCHAR` (including `VARCHAR(MAX)`), `TEXT`, and + /// `CHAR`/`VARCHAR` values inside `SQL_VARIANT`. The column's declared + /// collation always determines the encoding; BOM-like bytes are not + /// stripped or used to select another encoding. Unknown collations and + /// malformed protocol lengths still produce errors. + /// + /// Unicode values are controlled independently by + /// [`Config::lossy_utf16_decoding`]. XML, metadata, protocol strings, and + /// outgoing character encoding are unaffected. + /// + /// ``` + /// let mut config = tiberius::Config::new(); + /// config.lossy_codepage_decoding(true); + /// ``` + pub fn lossy_codepage_decoding(&mut self, lossy: bool) { + self.lossy_codepage_decoding = lossy; + } + + /// Returns whether lossy code-page row decoding is enabled. + /// + /// See [`Config::lossy_codepage_decoding`] for the scope of this option. + pub fn get_lossy_codepage_decoding(&self) -> bool { + self.lossy_codepage_decoding + } + + /// Supplies a client certificate and private key used to authenticate the + /// client to the server during the TLS handshake (mutual TLS). This is + /// required for TDS 8.0 "strict" connections that use client-certificate + /// authentication (`ENCRYPT_CLIENT_CERT`), and may also be used with the + /// classic (pre-8.0) TLS handshake when the server requests a client + /// certificate. + /// + /// Both arguments are paths to files: + /// + /// - `cert`: the client certificate, PEM (`.pem`/`.crt`) or DER (`.der`). + /// - `key`: the matching private key, PEM (`.pem`/`.key`) or DER (`.der`, + /// PKCS#8). + /// + /// Backend support: + /// + /// - `rustls`: PEM and DER certificate/key files. + /// - `native-tls`: PEM certificate + PEM PKCS#8 key only (DER files are + /// rejected at connect time; use [`client_certificate_pkcs12`] for a + /// bundled DER identity). + /// - `vendored-openssl` (opentls): does not support separate certificate/key + /// files; use [`client_certificate_pkcs12`] instead. + /// + /// - Defaults to no client certificate. + /// + /// [`client_certificate_pkcs12`]: Config::client_certificate_pkcs12 + #[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + ))] + #[cfg_attr( + docsrs, + doc(cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + ))) + )] + pub fn client_certificate(&mut self, cert: impl Into, key: impl Into) { + self.client_cert = Some(ClientCertificate { + source: ClientCertSource::CertAndKey { + cert: cert.into(), + key: key.into(), + }, + }); + } + + /// Supplies a client identity from a PKCS#12 / PFX bundle (certificate, + /// private key and any chain, encrypted with `password`) used to + /// authenticate the client to the server during the TLS handshake (mutual + /// TLS). + /// + /// Only supported by the `native-tls` and `vendored-openssl` backends; the + /// `rustls` backend rejects PKCS#12 identities at connect time (supply + /// separate PEM/DER files via [`client_certificate`] instead). + /// + /// - Defaults to no client certificate. + /// + /// [`client_certificate`]: Config::client_certificate + #[cfg(any(feature = "native-tls", feature = "vendored-openssl"))] + #[cfg_attr( + docsrs, + doc(cfg(any(feature = "native-tls", feature = "vendored-openssl"))) + )] + pub fn client_certificate_pkcs12( + &mut self, + path: impl Into, + password: impl Into, + ) { + self.client_cert = Some(ClientCertificate { + source: ClientCertSource::Pkcs12 { + path: path.into(), + password: crate::client::auth::secret_from_string(password.into()), + }, + }); + } + + #[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + ))] + pub(crate) fn get_client_certificate(&self) -> Option<&ClientCertificate> { + self.client_cert.as_ref() } pub(crate) fn get_host(&self) -> &str { @@ -178,6 +833,17 @@ impl Config { .unwrap_or("localhost") } + #[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + ))] + pub(crate) fn get_hostname_in_certificate(&self) -> &str { + self.hostname_in_certificate + .as_deref() + .unwrap_or_else(|| self.get_host()) + } + pub(crate) fn get_port(&self) -> u16 { match (self.port, self.instance_name.as_ref()) { // A user-defined port, we must use that. @@ -210,8 +876,60 @@ impl Config { /// |`database`|``|The name of the database.| /// |`TrustServerCertificate`|`true`,`false`,`yes`,`no`|Specifies whether the driver trusts the server certificate when connecting using TLS. Cannot be used toghether with `TrustServerCertificateCA`| /// |`TrustServerCertificateCA`|``|Path to a `pem`, `crt` or `der` certificate file. Cannot be used together with `TrustServerCertificate`| - /// |`encrypt`|`true`,`false`,`yes`,`no`,`DANGER_PLAINTEXT`|Specifies whether the driver uses TLS to encrypt communication.| + /// |`encrypt`|`strict`,`true`,`false`,`yes`,`no`,`DANGER_PLAINTEXT`|Specifies whether the driver uses TLS to encrypt communication. `strict` (TDS 8.0) requires the `tds80` feature.| /// |`Application Name`, `ApplicationName`|``|Sets the application name for the connection.| + /// |`HostNameInCertificate`, `HostName In Certificate`|``|The hostname the server certificate is validated against. Defaults to the value of the `Server` keyword (host).| + /// |`WorkstationID`, `Workstation ID`|``|The client / workstation name reported to the server.| + /// |`MultiSubnetFailover`|`true`,`false`,`yes`,`no`|When enabled, connections are attempted in parallel to all IP addresses the server resolves to, and the first to succeed is used.| + /// + /// # Parsing model and special characters + /// + /// Pairs are separated by `;`, and each pair is split on its **first** `=`, + /// so a value may itself contain `=`: a base64 or otherwise generated + /// password with `=` padding (for example `password=Zm9vYmFy==`) parses + /// without any quoting. + /// + /// Characters that are still structural inside a *value* — `;`, a leading + /// `'`, `"` or `{`, and any leading or trailing spaces you want to keep — + /// must be quoted. Wrap the value in single quotes, double quotes, or braces + /// (`{...}`), choosing a style the value does not itself contain, and embed + /// the enclosing quote by doubling it (`"a""b"` → `a"b`): + /// + /// - `password='p@ss;word=42'` + /// - `password="p@ss;word=42"` + /// - `password={p@ss;word=42}` + /// + /// Do **not** URL-encode the value: `%24` is sent to the server literally, + /// not decoded back to `$`, which is a common cause of login failures. + /// A bare `}` needs no quoting, but brace quoting cannot hold a literal `}` + /// (there is no `}}` doubling), so a value that must be quoted *and* + /// contains a `}` has to use single or double quotes. Only ASCII values are + /// accepted by the parser. + /// + /// tiberius parses a **superset** of ADO.NET's documented rules: the ADO.NET + /// semantics above, plus `{...}` brace quoting as an extension (braces are + /// not part of ADO.NET itself). One consequence of that extension: because a + /// value that *begins* with `{` is read as brace-quoted, a value whose + /// literal first character is `{` must instead be single- or double-quoted + /// (e.g. `password='{literal-braces}'`). + /// + /// For non-ASCII passwords, or to avoid escaping entirely, build the + /// [`Config`] programmatically — the values are passed verbatim: + /// + /// ``` + /// # use tiberius::{Config, AuthMethod}; + /// // Quoted in an ADO.NET string: the `;`, `=` and `{}` are preserved. + /// let _config = Config::from_ado_string( + /// r#"server=tcp:localhost,1433;user=sa;password='p@ss;w{}rd=42'"#, + /// )?; + /// + /// // The programmatic API needs no escaping and accepts any value. + /// let mut config = Config::new(); + /// config.host("localhost"); + /// config.port(1433); + /// config.authentication(AuthMethod::sql_server("sa", "p@ss;w{}rd=42")); + /// # Ok::<(), tiberius::error::Error>(()) + /// ``` /// /// [ADO.NET connection string]: https://docs.microsoft.com/en-us/dotnet/framework/data/adonet/connection-strings pub fn from_ado_string(s: &str) -> crate::Result { @@ -223,6 +941,20 @@ impl Config { /// /// See [`from_ado_string`] method for supported parameters. /// + /// # Special characters in values + /// + /// The JDBC parser understands only brace quoting: wrap a value that + /// contains `;`, `=`, `:`, `{`, `\`, `/`, `[` or `]` in braces, e.g. + /// `password={p@ss;word}`. + /// Unlike ADO.NET, single and double quotes are **not** escape characters + /// here — they are ordinary value characters, and spaces are kept verbatim + /// without quoting (JDBC does not trim). Brace quoting cannot hold a literal + /// `}` (there is no `}}` doubling), though a bare `}` needs no quoting. Do + /// not URL-encode values, and note that only ASCII is accepted. + /// For non-ASCII values, or to avoid escaping entirely, build the + /// [`Config`] programmatically as shown in [`from_ado_string`]; the values + /// are passed verbatim. + /// /// [JDBC connection string]: https://docs.microsoft.com/en-us/sql/connect/jdbc/building-the-connection-url?view=sql-server-ver15 /// [`from_ado_string`]: #method.from_ado_string pub fn from_jdbc_string(s: &str) -> crate::Result { @@ -257,65 +989,407 @@ impl Config { builder.application_name(name); } - if s.trust_cert()? { + let trust_cert = s.trust_cert()?; + let trust_cert_ca = s.trust_cert_ca(); + + // `TrustServerCertificate` and `TrustServerCertificateCA` are mutually + // exclusive. Detect the conflict here and return an error instead of + // letting `trust_cert`/`trust_cert_ca` panic on external input. + if trust_cert && trust_cert_ca.is_some() { + return Err(crate::Error::Conversion( + "'TrustServerCertificate' and 'TrustServerCertificateCA' are \ + mutually exclusive; specify only one" + .into(), + )); + } + + if trust_cert { builder.trust_cert(); } - if let Some(ca) = s.trust_cert_ca() { + if let Some(ca) = trust_cert_ca { builder.trust_cert_ca(ca); } + if let Some(hostname_in_cert) = s.hostname_in_certificate() { + builder.hostname_in_certificate(hostname_in_cert); + } + builder.encryption(s.encrypt()?); builder.readonly(s.readonly()); + if let Some(client_name) = s.client_name() { + builder.client_name(client_name); + } + builder.multi_subnet_failover(s.multi_subnet_failover()?); + Ok(builder) } } -pub(crate) struct ServerDefinition { - host: Option, - port: Option, - instance: Option, +/// A builder for [`Config`], providing an ergonomic, chainable way to +/// construct a connection configuration. +/// +/// Create a builder with [`Config::builder`], set the desired options by +/// calling its methods (each returns the builder to allow chaining) and +/// finalize it with [`build`]. +/// +/// # Example +/// +/// ``` +/// # use tiberius::{Config, AuthMethod, EncryptionLevel}; +/// let config = Config::builder() +/// .host("localhost") +/// .port(1433) +/// .database("master") +/// .encryption(EncryptionLevel::NotSupported) +/// .authentication(AuthMethod::sql_server("SA", "")) +/// .build(); +/// ``` +/// +/// [`Config`]: struct.Config.html +/// [`Config::builder`]: struct.Config.html#method.builder +/// [`build`]: struct.ConfigBuilder.html#method.build +#[derive(Clone, Debug)] +pub struct ConfigBuilder { + inner: Config, } -pub(crate) trait ConfigString { - fn dict(&self) -> &HashMap; - - fn server(&self) -> crate::Result; - - fn authentication(&self) -> crate::Result { - let user = self - .dict() - .get("uid") - .or_else(|| self.dict().get("username")) - .or_else(|| self.dict().get("user")) - .or_else(|| self.dict().get("user id")) - .map(|s| s.as_str()); +impl ConfigBuilder { + /// A host or ip address to connect to. + /// + /// - Defaults to `localhost`. + pub fn host(mut self, host: impl ToString) -> Self { + self.inner.host = Some(host.to_string()); + self + } - let pw = self - .dict() - .get("password") - .or_else(|| self.dict().get("pwd")) - .map(|s| s.as_str()); + /// The server port. + /// + /// - Defaults to `1433`. + pub fn port(mut self, port: u16) -> Self { + self.inner.port = Some(port); + self + } - match self - .dict() - .get("integratedsecurity") - .or_else(|| self.dict().get("integrated security")) - { - #[cfg(all(windows, feature = "winauth"))] - Some(val) if val.to_lowercase() == "sspi" || Self::parse_bool(val)? => match (user, pw) - { - (None, None) => Ok(AuthMethod::Integrated), - _ => Ok(AuthMethod::windows(user.unwrap_or(""), pw.unwrap_or(""))), - }, - #[cfg(feature = "integrated-auth-gssapi")] - Some(val) if val.to_lowercase() == "sspi" || Self::parse_bool(val)? => { - Ok(AuthMethod::Integrated) - } - _ => Ok(AuthMethod::sql_server(user.unwrap_or(""), pw.unwrap_or(""))), - } + /// The database to connect to. + /// + /// - Defaults to `master`. + pub fn database(mut self, database: impl ToString) -> Self { + self.inner.database = Some(database.to_string()); + self + } + + /// The instance name as defined in the SQL Browser. Only available on + /// Windows platforms. + /// + /// If specified, the port is replaced with the value returned from the + /// browser. + /// + /// - Defaults to no name specified. + pub fn instance_name(mut self, name: impl ToString) -> Self { + self.inner.instance_name = Some(name.to_string()); + self + } + + /// Sets the application name to the connection, queryable with the + /// `APP_NAME()` command. + /// + /// - Defaults to no name specified. + pub fn application_name(mut self, name: impl ToString) -> Self { + self.inner.application_name = Some(name.to_string()); + self + } + + /// Set the preferred encryption level. + /// + /// - With a TLS backend enabled (any of the `rustls`, `native-tls`, or + /// `vendored-openssl` features), defaults to `Required`. + /// - Without a TLS backend, defaults to `NotSupported`. + pub fn encryption(mut self, encryption: EncryptionLevel) -> Self { + self.inner.encryption = encryption; + self + } + + /// If set, the server certificate will not be validated and it is accepted + /// as-is. + /// + /// On production setting, the certificate should be added to the local key + /// storage (or use `trust_cert_ca` instead), using this setting is potentially dangerous. + /// + /// # Panics + /// Will panic in case `trust_cert_ca` was called before. + /// + /// - Defaults to `default`, meaning server certificate is validated against system-truststore. + pub fn trust_cert(mut self) -> Self { + self.inner.trust_cert(); + self + } + + /// Trust an additional CA certificate (from a file), in addition to the base + /// trust anchors. Accumulates across calls. + /// + /// See [`Config::trust_cert_ca`] for details. + /// + /// # Panics + /// Will panic in case `trust_cert` was called before. + pub fn trust_cert_ca(mut self, path: impl ToString) -> Self { + self.inner.trust_cert_ca(path); + self + } + + /// Trust additional CA certificates supplied as in-memory bytes (a + /// PEM/DER bundle), in addition to the base trust anchors. Accumulates + /// across calls. + /// + /// See [`Config::trust_cert_ca_bundle`] for details. + /// + /// # Panics + /// Will panic in case `trust_cert` was called before. + pub fn trust_cert_ca_bundle(mut self, bundle: impl Into>) -> Self { + self.inner.trust_cert_ca_bundle(bundle); + self + } + + /// Use the compiled-in Mozilla root CA snapshot as the base trust anchors + /// (rustls only). + /// + /// See [`Config::trust_webpki_roots`] for details and the staleness caveat. + /// + /// # Panics + /// Will panic in case `trust_cert` was called before. + #[cfg(feature = "rustls-webpki-roots")] + #[cfg_attr(docsrs, doc(cfg(feature = "rustls-webpki-roots")))] + pub fn trust_webpki_roots(mut self) -> Self { + self.inner.trust_webpki_roots(); + self + } + + /// Sets the authentication method. + /// + /// - Defaults to `None`. + pub fn authentication(mut self, auth: AuthMethod) -> Self { + self.inner.auth = auth; + self + } + + /// Sets ApplicationIntent readonly. + /// + /// - Defaults to `false`. + pub fn readonly(mut self, readonly: bool) -> Self { + self.inner.readonly = readonly; + self + } + + /// Sets an upper bound on how long the connection handshake may take. + /// + /// See [`Config::handshake_timeout`] for details. Pass `None` to wait + /// indefinitely. + /// + /// - Defaults to 15 seconds. + pub fn handshake_timeout(mut self, timeout: Option) -> Self { + self.inner.handshake_timeout = timeout; + self + } + + /// Sets a per-response deadline applied while reading command results. + /// + /// See [`Config::command_timeout`] for the exact (per-round-trip) + /// semantics. Pass `None` to wait indefinitely. + /// + /// - Defaults to 30 seconds. + pub fn command_timeout(mut self, timeout: Option) -> Self { + self.inner.command_timeout = timeout; + self + } + + /// Enables lossy UTF-16 decoding for NVARCHAR/NTEXT row values. + /// + /// See [`Config::lossy_utf16_decoding`] for the exact semantics (strict is + /// the default; enabling replaces invalid UTF-16 with `U+FFFD` for the + /// NVARCHAR/NCHAR and NTEXT arms only). + /// + /// - Defaults to `false`. + pub fn lossy_utf16_decoding(mut self, lossy: bool) -> Self { + self.inner.lossy_utf16_decoding = lossy; + self + } + + /// Enables lossy code-page row decoding. Defaults to `false`. + /// + /// See [`Config::lossy_codepage_decoding`] for the scope of this option. + pub fn lossy_codepage_decoding(mut self, lossy: bool) -> Self { + self.inner.lossy_codepage_decoding = lossy; + self + } + + /// Supplies a client certificate and private key for mutual TLS. + /// + /// See [`Config::client_certificate`] for details and backend support. + #[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + ))] + #[cfg_attr( + docsrs, + doc(cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + ))) + )] + pub fn client_certificate(mut self, cert: impl Into, key: impl Into) -> Self { + self.inner.client_certificate(cert, key); + self + } + + /// Supplies a client identity from a PKCS#12 / PFX bundle for mutual TLS. + /// + /// See [`Config::client_certificate_pkcs12`] for details and backend + /// support. + #[cfg(any(feature = "native-tls", feature = "vendored-openssl"))] + #[cfg_attr( + docsrs, + doc(cfg(any(feature = "native-tls", feature = "vendored-openssl"))) + )] + pub fn client_certificate_pkcs12( + mut self, + path: impl Into, + password: impl Into, + ) -> Self { + self.inner.client_certificate_pkcs12(path, password); + self + } + + /// Produces the finalized [`Config`] from this builder. + /// + /// [`Config`]: struct.Config.html + pub fn build(self) -> Config { + self.inner + } +} + +impl From for ConfigBuilder { + fn from(config: Config) -> Self { + ConfigBuilder { inner: config } + } +} + +impl From for Config { + fn from(builder: ConfigBuilder) -> Self { + builder.inner + } +} + +pub(crate) struct ServerDefinition { + host: Option, + port: Option, + instance: Option, +} + +/// Wrap a [`connection_string`] parse failure with actionable guidance. +/// +/// The underlying ADO.NET / JDBC parser fails with a terse, low-level message +/// (for example "key-value pairs must be joined by a =") when a value contains +/// an unescaped special character such as `;`, `=` or `{` — a very common +/// cause of "cannot connect" reports for passwords with special characters. +/// Append a hint describing how to quote such values so the message is +/// self-explanatory. `quoting` spells out the escaping styles the specific +/// parser accepts (they differ between ADO.NET and JDBC) and `docs` names the +/// constructor whose rustdoc carries the full rules and the programmatic-build +/// workaround. +/// +/// The hint is phrased conditionally because `connection_string::Error` is an +/// opaque string with no error kind, so this wrapper cannot tell a quoting +/// failure apart from an unrelated one (a bad port, a mistyped sub-protocol, +/// …); "if a value … contains a special character" keeps the guidance honest +/// for every failure while still solving the common quoting case. +fn connection_string_error( + err: connection_string::Error, + quoting: &str, + docs: &str, +) -> crate::Error { + // `connection_string::Error`'s `Display` already prefixes `Conversion + // error: `; strip it so wrapping the text in `Error::Conversion` does not + // duplicate the prefix. + let err = err.to_string(); + let detail = err.strip_prefix("Conversion error: ").unwrap_or(&err); + hinted_conversion_error(detail, quoting, docs) +} + +/// Builds a `Conversion` error from a terse parser reason plus the shared +/// quoting hint, so every connection-string failure points the caller at the +/// escaping options and the programmatic `Config` API. +fn hinted_conversion_error(detail: &str, quoting: &str, docs: &str) -> crate::Error { + crate::Error::Conversion( + format!( + "{detail}. hint: if a value such as a password contains a special \ + character (for example `;`, `=` or `{{`), it must be quoted — \ + {quoting}. Do not URL-encode the value (`%24` is sent literally, \ + not decoded to `$`). Non-ASCII values are not accepted by the \ + connection-string parser; build the `Config` programmatically \ + (see `{docs}`) to pass any value without escaping." + ) + .into(), + ) +} + +pub(crate) trait ConfigString { + fn dict(&self) -> &HashMap; + + fn server(&self) -> crate::Result; + + fn authentication(&self) -> crate::Result { + let user = self + .dict() + .get("uid") + .or_else(|| self.dict().get("username")) + .or_else(|| self.dict().get("user")) + .or_else(|| self.dict().get("user id")) + .map(|s| s.as_str()); + + let pw = self + .dict() + .get("password") + .or_else(|| self.dict().get("pwd")) + .map(|s| s.as_str()); + + match self + .dict() + .get("integratedsecurity") + .or_else(|| self.dict().get("integrated security")) + { + #[cfg(all(windows, feature = "winauth"))] + Some(val) if Self::is_sspi_or_truthy(val)? => match (user, pw) { + (None, None) => Ok(AuthMethod::Integrated), + _ => Ok(AuthMethod::windows(user.unwrap_or(""), pw.unwrap_or(""))), + }, + // On Unix with `sspi-rs`, `IntegratedSecurity=SSPI` (or a truthy + // value) uses NTLM when a username/password is supplied, and falls + // back to Kerberos (`Integrated`) only if `integrated-auth-gssapi` + // is also enabled and no credentials are given. + #[cfg(all(unix, feature = "sspi-rs"))] + Some(val) if Self::is_sspi_or_truthy(val)? => match (user, pw) { + (Some(user), Some(pw)) => Ok(AuthMethod::windows(user, pw)), + #[cfg(feature = "integrated-auth-gssapi")] + (None, None) => Ok(AuthMethod::Integrated), + _ => Ok(AuthMethod::windows(user.unwrap_or(""), pw.unwrap_or(""))), + }, + #[cfg(all( + feature = "integrated-auth-gssapi", + not(all(unix, feature = "sspi-rs")) + ))] + Some(val) if Self::is_sspi_or_truthy(val)? => Ok(AuthMethod::Integrated), + // Default (no integrated security): SQL Server authentication. A + // missing user or password is intentionally passed through as an + // empty string rather than rejected here — validation of the + // credentials is deferred to the server LOGIN response, which + // returns a precise authentication error. Failing early would also + // break the (unusual but valid) case of an empty SQL login. + _ => Ok(AuthMethod::sql_server(user.unwrap_or(""), pw.unwrap_or(""))), + } } fn database(&self) -> Option { @@ -346,6 +1420,20 @@ pub(crate) trait ConfigString { .map(|ca| ca.to_string()) } + fn hostname_in_certificate(&self) -> Option { + self.dict() + .get("hostnameincertificate") + .or_else(|| self.dict().get("hostname in certificate")) + .map(|host| host.to_string()) + } + + fn client_name(&self) -> Option { + self.dict() + .get("workstationid") + .or_else(|| self.dict().get("workstation id")) + .map(|name| name.to_string()) + } + #[cfg(any( feature = "rustls", feature = "native-tls", @@ -358,9 +1446,20 @@ pub(crate) trait ConfigString { Ok(true) => Ok(EncryptionLevel::Required), Ok(false) => Ok(EncryptionLevel::Off), Err(_) if val == "DANGER_PLAINTEXT" => Ok(EncryptionLevel::NotSupported), + Err(_) if val.eq_ignore_ascii_case("strict") && cfg!(feature = "tds80") => { + Ok(EncryptionLevel::Strict) + } + Err(_) if val.eq_ignore_ascii_case("strict") => Err(crate::Error::Conversion( + "encrypt=strict requires the crate's `tds80` feature to be enabled".into(), + )), Err(e) => Err(e), }) - .unwrap_or(Ok(EncryptionLevel::Off)) + // When the `encrypt` keyword is omitted, default to requiring + // encryption — matching `Config::default()` and modern ADO.NET + // (`Encrypt=Mandatory`). Callers who want an unencrypted connection + // must opt out explicitly with `encrypt=false` (or + // `encrypt=DANGER_PLAINTEXT`). + .unwrap_or(Ok(EncryptionLevel::Required)) } #[cfg(not(any( @@ -369,7 +1468,38 @@ pub(crate) trait ConfigString { feature = "vendored-openssl" )))] fn encrypt(&self) -> crate::Result { - Ok(EncryptionLevel::NotSupported) + // Security (#305): a no-TLS build cannot encrypt, so an explicit request + // to do so must fail loudly rather than silently downgrade to plaintext. + // Token classification follows the with-TLS `encrypt()`: the values that + // mean "encryption on" (`true`/`yes`/`strict`) become the TLS-missing + // error, while opting out (`false`/`no`/`DANGER_PLAINTEXT`) and an + // omitted keyword still resolve to `NotSupported`. (`strict` reports the + // missing backend rather than the with-TLS `tds80`-required hint — a + // TLS backend is the more fundamental thing it needs here.) + let Some(val) = self.dict().get("encrypt") else { + return Ok(EncryptionLevel::NotSupported); + }; + + match Self::parse_bool(val) { + Ok(false) => Ok(EncryptionLevel::NotSupported), + Err(_) if val == "DANGER_PLAINTEXT" => Ok(EncryptionLevel::NotSupported), + Ok(true) => Err(Self::tls_backend_missing()), + Err(_) if val.eq_ignore_ascii_case("strict") => Err(Self::tls_backend_missing()), + Err(e) => Err(e), + } + } + + #[cfg(not(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + )))] + fn tls_backend_missing() -> crate::Error { + crate::Error::Tls( + "encryption was requested (`encrypt=...`) but the crate was compiled without a TLS \ + backend; enable one of the `native-tls`, `rustls` or `vendored-openssl` features." + .to_string(), + ) } fn parse_bool>(v: T) -> crate::Result { @@ -382,10 +1512,683 @@ pub(crate) trait ConfigString { } } + /// An `IntegratedSecurity` connection-string value selects Windows/SSPI auth + /// when it is the literal `SSPI` (case-insensitive) or a truthy boolean. + /// Uses `eq_ignore_ascii_case` so the `SSPI` comparison does not allocate. + /// + /// Only referenced by the integrated-auth match arms, which are themselves + /// feature-gated; gate the helper identically so builds without any + /// integrated-auth backend do not warn about it being unused. + #[cfg(any( + all(windows, feature = "winauth"), + all(unix, feature = "sspi-rs"), + feature = "integrated-auth-gssapi" + ))] + fn is_sspi_or_truthy>(v: T) -> crate::Result { + let v = v.as_ref(); + Ok(v.eq_ignore_ascii_case("sspi") || Self::parse_bool(v)?) + } + fn readonly(&self) -> bool { self.dict() .get("applicationintent") - .filter(|val| *val == "ReadOnly") + .filter(|val| val.trim().eq_ignore_ascii_case("ReadOnly")) .is_some() } + + fn multi_subnet_failover(&self) -> crate::Result { + self.dict() + .get("multisubnetfailover") + .map(Self::parse_bool) + .unwrap_or(Ok(false)) + } +} + +#[cfg(test)] +mod tests { + use super::*; + #[cfg(any(feature = "native-tls", feature = "vendored-openssl"))] + use secrecy::ExposeSecret; + + #[test] + fn config_debug_redacts_connection_password() { + // The connection password lives inside `auth: AuthMethod`; the whole + // `Config` Debug must never print it. + let config = Config::builder() + .authentication(AuthMethod::sql_server("SA", "conn-str-secret")) + .build(); + + let dbg = format!("{config:?}"); + assert!( + !dbg.contains("conn-str-secret"), + "connection password leaked in Config Debug: {dbg}" + ); + assert!(dbg.contains("REDACTED"), "password not redacted: {dbg}"); + } + + #[test] + fn config_builder_constructs_config() { + let config = Config::builder() + .host("db.example.com") + .port(4433) + .database("northwind") + .application_name("my-app") + .authentication(AuthMethod::sql_server("SA", "secret")) + .readonly(true) + .build(); + + assert_eq!("db.example.com", config.get_host()); + assert_eq!(4433, config.get_port()); + assert_eq!("db.example.com:4433", config.get_addr()); + assert_eq!(Some("northwind"), config.database.as_deref()); + assert_eq!(Some("my-app"), config.application_name.as_deref()); + assert!(config.readonly); + assert!(matches!(config.auth, AuthMethod::SqlServer(_))); + assert!(config.trust.is_default()); + } + + #[test] + fn config_builder_roundtrips_via_from() { + let config = Config::builder().host("localhost").port(1433).build(); + let builder: ConfigBuilder = config.into(); + let config = builder.database("master").build(); + + assert_eq!("localhost:1433", config.get_addr()); + assert_eq!(Some("master"), config.database.as_deref()); + } + + #[test] + fn config_from_builder_carries_builder_settings() { + // `From` must return the built inner config, not a default. + let config: Config = Config::builder().host("db.internal").port(2020).into(); + assert_eq!("db.internal", config.get_host()); + assert_eq!(2020, config.get_port()); + } + + #[test] + fn handshake_timeout_defaults_to_fifteen_seconds() { + // A sensible, ADO.NET-matching default so a stalled handshake surfaces + // an error instead of hanging forever. + let config = Config::new(); + assert_eq!( + config.get_handshake_timeout(), + Some(Duration::from_secs(15)) + ); + } + + #[test] + fn handshake_timeout_setter_roundtrips() { + // Table-driven over the values a caller can set, including opting out. + let cases = [ + Some(Duration::from_millis(1)), + Some(Duration::from_secs(30)), + None, + ]; + for want in cases { + let mut config = Config::new(); + config.handshake_timeout(want); + assert_eq!(config.get_handshake_timeout(), want, "for {want:?}"); + } + } + + #[test] + fn config_builder_sets_handshake_timeout() { + let config = Config::builder() + .handshake_timeout(Some(Duration::from_secs(3))) + .build(); + assert_eq!(config.get_handshake_timeout(), Some(Duration::from_secs(3))); + + // And can disable it. + let config = Config::builder().handshake_timeout(None).build(); + assert_eq!(config.get_handshake_timeout(), None); + } + + #[test] + fn config_builder_defaults_handshake_timeout() { + // The builder starts from `Config::default()`, so it inherits the bound. + let config = Config::builder().build(); + assert_eq!( + config.get_handshake_timeout(), + Some(Duration::from_secs(15)) + ); + } + + #[test] + fn command_timeout_defaults_to_thirty_seconds() { + // ADO.NET-matching default so a mid-result server stall surfaces an + // error instead of hanging forever. + let config = Config::new(); + assert_eq!(config.get_command_timeout(), Some(Duration::from_secs(30))); + } + + #[test] + fn command_timeout_setter_roundtrips() { + let cases = [ + Some(Duration::from_millis(1)), + Some(Duration::from_secs(60)), + None, + ]; + for want in cases { + let mut config = Config::new(); + config.command_timeout(want); + assert_eq!(config.get_command_timeout(), want, "for {want:?}"); + } + } + + #[test] + fn config_builder_sets_command_timeout() { + let config = Config::builder() + .command_timeout(Some(Duration::from_secs(7))) + .build(); + assert_eq!(config.get_command_timeout(), Some(Duration::from_secs(7))); + + let config = Config::builder().command_timeout(None).build(); + assert_eq!(config.get_command_timeout(), None); + } + + #[test] + fn config_builder_defaults_command_timeout() { + let config = Config::builder().build(); + assert_eq!(config.get_command_timeout(), Some(Duration::from_secs(30))); + } + + #[test] + fn lossy_utf16_decoding_defaults_to_false() { + let config = Config::new(); + assert!(!config.get_lossy_utf16_decoding()); + } + + #[test] + fn lossy_codepage_decoding_defaults_and_setter() { + let mut config = Config::new(); + assert!(!config.get_lossy_codepage_decoding()); + config.lossy_codepage_decoding(true); + assert!(config.get_lossy_codepage_decoding()); + assert!(!config.get_lossy_utf16_decoding()); + config.lossy_utf16_decoding(true); + config.lossy_codepage_decoding(false); + assert!(!config.get_lossy_codepage_decoding()); + assert!(config.get_lossy_utf16_decoding()); + } + + #[test] + fn config_builder_sets_lossy_codepage_decoding() { + assert!(!Config::builder().build().get_lossy_codepage_decoding()); + let config = Config::builder().lossy_codepage_decoding(true).build(); + assert!(config.get_lossy_codepage_decoding()); + assert!(!config.get_lossy_utf16_decoding()); + } + + #[test] + fn lossy_utf16_decoding_setter_roundtrips() { + let mut config = Config::new(); + config.lossy_utf16_decoding(true); + assert!(config.get_lossy_utf16_decoding()); + config.lossy_utf16_decoding(false); + assert!(!config.get_lossy_utf16_decoding()); + } + + #[test] + fn config_builder_sets_lossy_utf16_decoding() { + let config = Config::builder().lossy_utf16_decoding(true).build(); + assert!(config.get_lossy_utf16_decoding()); + } + + #[test] + fn config_builder_defaults_lossy_utf16_decoding() { + let config = Config::builder().build(); + assert!(!config.get_lossy_utf16_decoding()); + } + + #[test] + fn timeouts_are_independent() { + // The two knobs must not alias one another. + let mut config = Config::new(); + config.handshake_timeout(Some(Duration::from_secs(1))); + config.command_timeout(Some(Duration::from_secs(2))); + assert_eq!(config.get_handshake_timeout(), Some(Duration::from_secs(1))); + assert_eq!(config.get_command_timeout(), Some(Duration::from_secs(2))); + } + + #[test] + fn get_packet_size_reflects_the_set_value() { + let mut config = Config::new(); + assert_eq!(config.get_packet_size(), None); + config.packet_size(8192); + assert_eq!(config.get_packet_size(), Some(8192)); + } + + #[test] + fn from_jdbc_string_parses_host_and_port() { + let config = + Config::from_jdbc_string("jdbc:sqlserver://db.example.com:2345").expect("valid jdbc"); + assert_eq!("db.example.com", config.get_host()); + assert_eq!(2345, config.get_port()); + } + + #[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + ))] + #[test] + fn get_hostname_in_certificate_falls_back_to_host() { + let mut config = Config::new(); + config.host("real.host"); + // Unset: falls back to the connection host. + assert_eq!(config.get_hostname_in_certificate(), "real.host"); + // Set: returns the explicit certificate hostname. + config.hostname_in_certificate("cert.host"); + assert_eq!(config.get_hostname_in_certificate(), "cert.host"); + } + + #[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + ))] + #[test] + fn client_certificate_sets_cert_and_key_source() { + let mut config = Config::new(); + assert!(config.get_client_certificate().is_none()); + + config.client_certificate("/tmp/client.pem", "/tmp/client.key"); + + let cert = config + .get_client_certificate() + .expect("client certificate should be set"); + match &cert.source { + ClientCertSource::CertAndKey { cert, key } => { + assert_eq!(cert, &PathBuf::from("/tmp/client.pem")); + assert_eq!(key, &PathBuf::from("/tmp/client.key")); + } + #[allow(unreachable_patterns)] + other => panic!("expected CertAndKey source, got {other:?}"), + } + } + + #[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + ))] + #[test] + fn config_builder_sets_client_certificate() { + let config = Config::builder() + .host("localhost") + .client_certificate("cert.der", "key.der") + .build(); + + match &config + .get_client_certificate() + .expect("client certificate should be set") + .source + { + ClientCertSource::CertAndKey { cert, key } => { + assert_eq!(cert, &PathBuf::from("cert.der")); + assert_eq!(key, &PathBuf::from("key.der")); + } + #[allow(unreachable_patterns)] + other => panic!("expected CertAndKey source, got {other:?}"), + } + } + + #[cfg(any(feature = "native-tls", feature = "vendored-openssl"))] + #[test] + fn client_certificate_pkcs12_sets_bundle_source() { + let mut config = Config::new(); + config.client_certificate_pkcs12("/tmp/identity.pfx", "s3cr3t"); + + match &config + .get_client_certificate() + .expect("client certificate should be set") + .source + { + ClientCertSource::Pkcs12 { path, password } => { + assert_eq!(path, &PathBuf::from("/tmp/identity.pfx")); + assert_eq!(password.expose_secret(), "s3cr3t"); + } + other => panic!("expected Pkcs12 source, got {other:?}"), + } + } + + #[cfg(any(feature = "native-tls", feature = "vendored-openssl"))] + #[test] + fn client_certificate_debug_redacts_pkcs12_password() { + let mut config = Config::new(); + config.client_certificate_pkcs12("/tmp/identity.pfx", "topsecret"); + + let dbg = format!("{:?}", config.get_client_certificate().unwrap()); + assert!(dbg.contains("REDACTED"), "password not redacted: {dbg}"); + assert!(!dbg.contains("topsecret"), "password leaked: {dbg}"); + } + + #[cfg(all(unix, feature = "sspi-rs"))] + #[test] + fn ado_integrated_security_sspi_with_credentials_uses_windows_ntlm() { + let config = Config::from_ado_string( + "server=tcp:localhost,1433;IntegratedSecurity=SSPI;uid=DOMAIN\\user;pwd=secret", + ) + .unwrap(); + + match config.auth { + AuthMethod::Windows(auth) => { + assert_eq!("user", auth.user); + assert_eq!(Some("DOMAIN"), auth.domain.as_deref()); + } + other => panic!("expected Windows NTLM auth, got {other:?}"), + } + } + + #[test] + fn config_direct_setters_populate_fields() { + let mut config = Config::new(); + config.database("northwind"); + config.instance_name("SQLEXPRESS"); + config.client_name("workstation-7"); + + assert_eq!(Some("northwind"), config.database.as_deref()); + assert_eq!(Some("SQLEXPRESS"), config.instance_name.as_deref()); + assert_eq!(Some("workstation-7"), config.client_name.as_deref()); + } + + #[test] + fn get_port_defaults_without_port_or_instance() { + // No explicit port and no instance -> default SQL Server port. + let config = Config::new(); + assert_eq!(1433, config.get_port()); + } + + #[test] + fn get_port_uses_sql_browser_port_for_named_instance() { + // A named instance without an explicit port -> SQL Browser port. + let mut config = Config::new(); + config.instance_name("SQLEXPRESS"); + assert_eq!(1434, config.get_port()); + } + + #[test] + #[should_panic(expected = "mutual exclusive")] + fn trust_cert_after_trust_cert_ca_panics() { + let mut config = Config::new(); + config.trust_cert_ca("/tmp/ca.crt"); + config.trust_cert(); + } + + #[test] + #[should_panic(expected = "mutual exclusive")] + fn trust_cert_ca_after_trust_cert_panics() { + let mut config = Config::new(); + config.trust_cert(); + config.trust_cert_ca("/tmp/ca.crt"); + } + + #[test] + fn trust_cert_ca_sets_ca_location() { + let mut config = Config::new(); + config.trust_cert_ca("/tmp/ca.crt"); + assert!(matches!( + config.trust.extra_cas.as_slice(), + [ExtraCa::File(_)] + )); + assert!(!config.trust.bypass); + assert!(matches!(config.trust.source, RootSource::Native)); + } + + #[test] + fn trust_cert_ca_accumulates_across_calls() { + // Repeated calls are additive (not replace-last-wins): both CAs must + // be trusted, across file + bundle sources. + let mut config = Config::new(); + config.trust_cert_ca("/tmp/a.crt"); + config.trust_cert_ca("/tmp/b.crt"); + config.trust_cert_ca_bundle(b"----- not really a cert -----".to_vec()); + + assert_eq!(config.trust.extra_cas.len(), 3); + match &config.trust.extra_cas[0] { + ExtraCa::File(p) => assert_eq!(p, &PathBuf::from("/tmp/a.crt")), + other => panic!("expected File, got {other:?}"), + } + match &config.trust.extra_cas[1] { + ExtraCa::File(p) => assert_eq!(p, &PathBuf::from("/tmp/b.crt")), + other => panic!("expected File, got {other:?}"), + } + assert!(matches!(config.trust.extra_cas[2], ExtraCa::Bundle(_))); + } + + #[test] + fn trust_cert_ca_bundle_pushes_bundle() { + let mut config = Config::new(); + config.trust_cert_ca_bundle(vec![1u8, 2, 3]); + match config.trust.extra_cas.as_slice() { + [ExtraCa::Bundle(bytes)] => assert_eq!(bytes, &[1, 2, 3]), + other => panic!("expected a single Bundle, got {other:?}"), + } + } + + #[test] + fn trust_config_debug_redacts_bundle_bytes() { + // The bundle's raw bytes must never appear in Debug; only a + // length summary. + let mut config = Config::new(); + config.trust_cert_ca_bundle(vec![0xDE, 0xAD, 0xBE, 0xEF]); + let dbg = format!("{:?}", config.trust); + assert!(dbg.contains("Bundle { len: 4 }"), "got: {dbg}"); + assert!( + !dbg.contains("222") && !dbg.contains("0xde"), + "bytes leaked: {dbg}" + ); + } + + #[cfg(feature = "rustls-webpki-roots")] + #[test] + fn trust_webpki_roots_sets_source_and_is_mutually_exclusive_with_extras() { + let mut config = Config::new(); + config.trust_webpki_roots(); + assert!(matches!(config.trust.source, RootSource::WebpkiRoots)); + // Extra CAs still layer on top of the webpki source. + config.trust_cert_ca("/tmp/extra.crt"); + assert_eq!(config.trust.extra_cas.len(), 1); + } + + #[cfg(feature = "rustls-webpki-roots")] + #[test] + #[should_panic(expected = "mutual exclusive")] + fn trust_cert_after_trust_webpki_roots_panics() { + let mut config = Config::new(); + config.trust_webpki_roots(); + config.trust_cert(); + } + + #[cfg(feature = "rustls-webpki-roots")] + #[test] + #[should_panic(expected = "mutual exclusive")] + fn trust_webpki_roots_after_trust_cert_panics() { + let mut config = Config::new(); + config.trust_cert(); + config.trust_webpki_roots(); + } + + #[test] + #[should_panic(expected = "mutual exclusive")] + fn trust_cert_after_trust_cert_ca_bundle_panics() { + let mut config = Config::new(); + config.trust_cert_ca_bundle(vec![1, 2, 3]); + config.trust_cert(); + } + + #[test] + #[should_panic(expected = "mutual exclusive")] + fn trust_cert_ca_bundle_after_trust_cert_panics() { + let mut config = Config::new(); + config.trust_cert(); + config.trust_cert_ca_bundle(vec![1, 2, 3]); + } + + #[test] + fn trust_cert_sets_bypass() { + // The default must be a validating config; `trust_cert()` is the explicit + // opt-in that switches to bypass validation. + let mut config = Config::new(); + assert!(config.trust.is_default()); + config.trust_cert(); + assert!(config.trust.bypass); + } + + #[test] + fn config_builder_covers_all_setters() { + let config = Config::builder() + .host("localhost") + .instance_name("SQLEXPRESS") + .encryption(EncryptionLevel::Off) + .trust_cert_ca("/tmp/ca.crt") + .build(); + + assert_eq!(Some("SQLEXPRESS"), config.instance_name.as_deref()); + assert!(matches!(config.encryption, EncryptionLevel::Off)); + assert!(matches!( + config.trust.extra_cas.as_slice(), + [ExtraCa::File(_)] + )); + } + + #[test] + fn config_builder_trust_cert_sets_bypass() { + let config = Config::builder().trust_cert().build(); + assert!(config.trust.bypass); + } + + #[test] + fn config_builder_trust_cert_ca_bundle_accumulates() { + let config = Config::builder() + .trust_cert_ca("/tmp/a.crt") + .trust_cert_ca_bundle(vec![1, 2, 3]) + .build(); + assert_eq!(config.trust.extra_cas.len(), 2); + } + + #[cfg(feature = "rustls-webpki-roots")] + #[test] + fn config_builder_trust_webpki_roots_sets_source() { + let config = Config::builder().trust_webpki_roots().build(); + assert!(matches!(config.trust.source, RootSource::WebpkiRoots)); + } + + #[test] + #[should_panic(expected = "mutual exclusive")] + fn config_builder_trust_cert_after_ca_panics() { + Config::builder().trust_cert_ca("/tmp/ca.crt").trust_cert(); + } + + #[test] + #[should_panic(expected = "mutual exclusive")] + fn config_builder_trust_cert_ca_after_trust_cert_panics() { + Config::builder().trust_cert().trust_cert_ca("/tmp/ca.crt"); + } + + #[test] + fn from_ado_string_populates_optional_fields() { + let config = Config::from_ado_string( + "server=tcp:my-server.com\\SQLEXPRESS;database=northwind;\ + HostNameInCertificate=cert.host;WorkstationID=ws-1", + ) + .expect("valid ado string"); + + assert_eq!("my-server.com", config.get_host()); + assert_eq!(Some("SQLEXPRESS"), config.instance_name.as_deref()); + assert_eq!(Some("northwind"), config.database.as_deref()); + assert_eq!(Some("cert.host"), config.hostname_in_certificate.as_deref()); + assert_eq!(Some("ws-1"), config.client_name.as_deref()); + } + + #[test] + fn from_ado_string_trust_cert_ca_populates_extra_cas() { + let config = Config::from_ado_string( + "server=tcp:localhost,1433;TrustServerCertificateCA=/tmp/ca.crt", + ) + .expect("valid ado string"); + match config.trust.extra_cas.as_slice() { + [ExtraCa::File(p)] => assert_eq!(p, &PathBuf::from("/tmp/ca.crt")), + other => panic!("expected a single File CA, got {other:?}"), + } + assert!(!config.trust.bypass); + } + + #[test] + fn from_ado_string_trust_cert_populates_bypass() { + let config = + Config::from_ado_string("server=tcp:localhost,1433;TrustServerCertificate=true") + .expect("valid ado string"); + assert!(config.trust.bypass); + } + + #[test] + fn from_ado_string_trust_cert_and_ca_conflict_errors() { + // The conflict must surface as a hard error, never silently let + // the bypass win. From a connection string it is an `Err`, not a panic. + let err = Config::from_ado_string( + "server=tcp:localhost,1433;TrustServerCertificate=true;\ + TrustServerCertificateCA=/tmp/ca.crt", + ) + .expect_err("conflicting trust settings must error"); + match err { + crate::Error::Conversion(msg) => { + assert!(msg.contains("mutually exclusive"), "got: {msg}") + } + other => panic!("expected Conversion error, got {other:?}"), + } + } + + #[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + ))] + #[test] + fn client_cert_source_debug_formats_cert_and_key() { + let mut config = Config::new(); + config.client_certificate("/tmp/client.pem", "/tmp/client.key"); + + let dbg = format!("{:?}", config.get_client_certificate().unwrap().source); + assert!(dbg.contains("CertAndKey")); + assert!(dbg.contains("client.pem")); + assert!(dbg.contains("client.key")); + } + + #[cfg(any(feature = "native-tls", feature = "vendored-openssl"))] + #[test] + fn config_builder_sets_pkcs12_client_certificate() { + let config = Config::builder() + .client_certificate_pkcs12("/tmp/identity.pfx", "s3cr3t") + .build(); + + match &config + .get_client_certificate() + .expect("client certificate should be set") + .source + { + ClientCertSource::Pkcs12 { path, password } => { + assert_eq!(path, &PathBuf::from("/tmp/identity.pfx")); + assert_eq!(password.expose_secret(), "s3cr3t"); + } + other => panic!("expected Pkcs12 source, got {other:?}"), + } + } + + #[cfg(all(unix, feature = "sspi-rs"))] + #[test] + fn ado_integrated_security_sspi_with_partial_credentials_uses_windows() { + // Only a username (no password) -> falls into the catch-all NTLM arm. + let config = Config::from_ado_string( + "server=tcp:localhost,1433;IntegratedSecurity=SSPI;uid=onlyuser", + ) + .unwrap(); + + match config.auth { + AuthMethod::Windows(auth) => { + assert_eq!("onlyuser", auth.user); + } + other => panic!("expected Windows auth, got {other:?}"), + } + } } diff --git a/src/client/config/ado_net.rs b/src/client/config/ado_net.rs index 94df9ca38..0c878016f 100644 --- a/src/client/config/ado_net.rs +++ b/src/client/config/ado_net.rs @@ -2,21 +2,29 @@ use super::{ConfigString, ServerDefinition}; use std::str::FromStr; pub(crate) struct AdoNetConfig { - dict: connection_string::AdoNetString, + dict: super::ado_parser::AdoNetString, } impl FromStr for AdoNetConfig { type Err = crate::error::Error; fn from_str(s: &str) -> crate::Result { - let dict = s.parse()?; + let dict = super::ado_parser::parse(s).map_err(|reason| { + super::hinted_conversion_error( + reason, + "wrap the value in single quotes (e.g. `password='p@ss;word'`), \ + double quotes, or braces (e.g. `password={p@ss;word}`), and \ + quote leading spaces too", + "Config::from_ado_string", + ) + })?; Ok(Self { dict }) } } impl ConfigString for AdoNetConfig { fn dict(&self) -> &std::collections::HashMap { - &self.dict + self.dict.pairs() } fn server(&self) -> crate::Result { @@ -55,8 +63,9 @@ impl ConfigString for AdoNetConfig { match self .dict + .pairs() .get("server") - .or_else(|| self.dict.get("data source")) + .or_else(|| self.dict.pairs().get("data source")) { Some(value) if value.starts_with("tcp:") => { parse_server(value[4..].split(',').collect()) @@ -230,6 +239,39 @@ mod tests { Ok(()) } + #[test] + fn server_parsing_too_many_parts_is_error() -> crate::Result<()> { + // The Server value must have at most two comma-separated parts + // (host[,port]). Three parts is invalid and must error. The guard is + // `parts.is_empty() || parts.len() >= 3`; a `&&` mutation would never + // trigger (a slice cannot be both empty and have >= 3 parts), so this + // three-part value would be wrongly accepted. + let ado: AdoNetConfig = "server=tcp:my-server.com,1433,extra".parse()?; + assert!(ado.server().is_err()); + + let ado: AdoNetConfig = "server=my-server.com,1433,extra".parse()?; + assert!(ado.server().is_err()); + + Ok(()) + } + + #[test] + fn server_parsing_missing_key() -> crate::Result<()> { + // No `server`/`data source` key at all -> an all-`None` definition. + let ado: AdoNetConfig = "database=Foo".parse()?; + let server = ado.server()?; + + assert_eq!(None, server.host); + assert_eq!(None, server.port); + assert_eq!(None, server.instance); + + // And the same path through the public constructor. + let config = crate::Config::from_ado_string("database=Foo")?; + assert_eq!("localhost", config.get_host()); + + Ok(()) + } + #[test] fn database_parsing() -> crate::Result<()> { let test_str = "database=Foo"; @@ -410,6 +452,262 @@ mod tests { Ok(()) } + #[test] + fn parsing_password_special_characters() -> crate::Result<()> { + // (value as written in the connection string, exact parsed password). + // The user and password must both round-trip exactly. + let cases: &[(&str, &str)] = &[ + // Single quotes: the most general escape; any ASCII char except a + // literal single quote can appear verbatim. + ("'a;b'", "a;b"), + ("'a=b'", "a=b"), + ("'a{b'", "a{b"), + ("'a}b'", "a}b"), + ("'a b'", "a b"), + ("'a$b'", "a$b"), + ("'a@b'", "a@b"), + ("'a!b'", "a!b"), + ("'a#b'", "a#b"), + ("'a%b'", "a%b"), + ("'a&b'", "a&b"), + ("'a,b'", "a,b"), + ("'a\"b'", "a\"b"), // a double quote inside single quotes + ("'p@ss;w{}rd=42'", "p@ss;w{}rd=42"), + // Double quotes: same idea, and lets a single quote appear. + ("\"a'b\"", "a'b"), + ("\"a;b=c\"", "a;b=c"), + // Braces (`{...}`): handy for `;` and `=`, but cannot hold a `}`. + ("{a;b}", "a;b"), + ("{a=b}", "a=b"), + ("{a$b}", "a$b"), + // Unquoted: non-structural characters pass through untouched. + ("a$b", "a$b"), + ("a@b", "a@b"), + ("a!b", "a!b"), + ("a#b", "a#b"), + ("a%b", "a%b"), + ("a&b", "a&b"), + ("a,b", "a,b"), + // A bare `}` is not structural: it passes through unquoted. + ("a}b", "a}b"), + // Percent-encoding is NOT decoded: `%24` stays literal, it does + // not become `$`. + ("a%24b", "a%24b"), + // An empty password is accepted by the ADO parser. + ("", ""), + ]; + + for (value, expected) in cases { + let s = format!("User Id=sa;Password={value};"); + let ado: AdoNetConfig = s.parse()?; + assert_eq!( + AuthMethod::sql_server("sa", *expected), + ado.authentication()?, + "value `{value}` should parse to password `{expected}`" + ); + } + + Ok(()) + } + + #[test] + fn parsing_password_whitespace_quirks() -> crate::Result<()> { + // Quoting preserves surrounding whitespace (matching ADO.NET): both a + // leading and a trailing space inside quotes survive. + let ado: AdoNetConfig = "User Id=sa;Password='a '".parse()?; + assert_eq!(AuthMethod::sql_server("sa", "a "), ado.authentication()?); + + let ado: AdoNetConfig = "User Id=sa;Password=' a'".parse()?; + assert_eq!(AuthMethod::sql_server("sa", " a"), ado.authentication()?); + + // An unquoted value is trimmed on both sides. + let ado: AdoNetConfig = "User Id=sa;Password= a ".parse()?; + assert_eq!(AuthMethod::sql_server("sa", "a"), ado.authentication()?); + + // Internal tabs survive. + let ado: AdoNetConfig = "User Id=sa;Password='a\tb'".parse()?; + assert_eq!(AuthMethod::sql_server("sa", "a\tb"), ado.authentication()?); + + Ok(()) + } + + #[test] + fn parsing_password_brace_cannot_hold_close_brace() -> crate::Result<()> { + // A `{…}` value closes at the first `}` and cannot contain one: any + // characters after that `}` are a syntax error rather than being + // silently folded into the value. + assert!("User Id=sa;Password={a}b}".parse::().is_err()); + + // Use single or double quotes for passwords that contain `}`. + let ado: AdoNetConfig = "User Id=sa;Password='a}b}'".parse()?; + assert_eq!(AuthMethod::sql_server("sa", "a}b}"), ado.authentication()?); + + Ok(()) + } + + #[test] + fn parsing_password_non_ascii_is_rejected() { + // The ADO parser only accepts ASCII; non-ASCII passwords must use the + // programmatic `Config` API instead. + assert!("User Id=sa;Password=café".parse::().is_err()); + assert!("User Id=sa;Password='café'" + .parse::() + .is_err()); + } + + // Every printable-ASCII special character must be usable in a password + // through single quotes and through double quotes, round-tripping exactly. + // The enclosing quote is embedded by doubling it. + #[test] + fn parsing_password_every_special_char_via_quotes() -> crate::Result<()> { + const SPECIALS: &str = "!@#$%^&*()-_+=[]{}|\\:;\"'<>,.?/~` "; + for c in SPECIALS.chars() { + let expected = format!("Pa{c}ss1"); + + // Single quotes: a literal `'` is written as `''`. + let inner = if c == '\'' { + "Pa''ss1".to_string() + } else { + format!("Pa{c}ss1") + }; + let s = format!("User Id=sa;Password='{inner}'"); + let ado: AdoNetConfig = s + .parse() + .unwrap_or_else(|e| panic!("char {c:?} single-quoted failed: {e}")); + assert_eq!( + AuthMethod::sql_server("sa", expected.as_str()), + ado.authentication()?, + "char {c:?} single-quoted" + ); + + // Double quotes: a literal `"` is written as `""`. + let inner = if c == '"' { + "Pa\"\"ss1".to_string() + } else { + format!("Pa{c}ss1") + }; + let s = format!("User Id=sa;Password=\"{inner}\""); + let ado: AdoNetConfig = s + .parse() + .unwrap_or_else(|e| panic!("char {c:?} double-quoted failed: {e}")); + assert_eq!( + AuthMethod::sql_server("sa", expected.as_str()), + ado.authentication()?, + "char {c:?} double-quoted" + ); + } + + Ok(()) + } + + // A doubled enclosing quote embeds that quote character. + #[test] + fn parsing_password_doubled_quotes_embed_the_quote() -> crate::Result<()> { + let ado: AdoNetConfig = "User Id=sa;Password='a''b'".parse()?; + assert_eq!(AuthMethod::sql_server("sa", "a'b"), ado.authentication()?); + + let ado: AdoNetConfig = "User Id=sa;Password=\"a\"\"b\"".parse()?; + assert_eq!(AuthMethod::sql_server("sa", "a\"b"), ado.authentication()?); + + Ok(()) + } + + // An unquoted value may contain `=` — each pair splits on the first `=` — + // so a base64/generated password with `=` padding parses (issue #313). + #[test] + fn issue_313_unquoted_equals_in_password_parses() -> crate::Result<()> { + let cases: &[(&str, &str)] = &[ + ("Zm9vYmFy==", "Zm9vYmFy=="), // base64 with `==` padding (the report) + ("aGVsbG8=", "aGVsbG8="), // base64 with a single `=` pad + ("Pa=ss1", "Pa=ss1"), + ("a=b=c", "a=b=c"), + ]; + for (raw, expected) in cases { + let s = + format!("Server=tcp:host.example.com,1433;Database=DB;User Id=sa;Password={raw}"); + let ado: AdoNetConfig = s + .parse() + .unwrap_or_else(|e| panic!("unquoted `{raw}` should parse: {e}")); + assert_eq!( + AuthMethod::sql_server("sa", *expected), + ado.authentication()?, + "unquoted `{raw}`" + ); + } + + Ok(()) + } + + // `==` in the key position is a literal `=`; the pair still splits on the + // first *single* `=`. + #[test] + fn parsing_key_double_equals_is_literal() -> crate::Result<()> { + let ado: AdoNetConfig = "User Id=sa;Password=p".parse()?; + assert_eq!(AuthMethod::sql_server("sa", "p"), ado.authentication()?); + // A hypothetical key containing `=` round-trips via `==`. + let cfg: AdoNetConfig = "a==b=c".parse()?; + assert_eq!(cfg.dict.pairs().get("a=b").map(String::as_str), Some("c")); + + Ok(()) + } + + // Characters that remain genuinely ambiguous unquoted must still error and + // point the caller at quoting — never silently mis-parse. + #[test] + fn parsing_password_ambiguous_unquoted_chars_error() { + // `;` is the separator; `Password=a;b` is read as a pair `a` then a + // second pair `b` with no `=`. + assert!("User Id=sa;Password=a;b".parse::().is_err()); + // A value that *starts* with a quote/brace but is not closed is an error. + assert!("User Id=sa;Password='ab".parse::().is_err()); + assert!("User Id=sa;Password=\"ab".parse::().is_err()); + assert!("User Id=sa;Password={ab".parse::().is_err()); + } + + #[test] + fn unquoted_special_char_password_error_has_hint() { + // An unquoted `;` in a password breaks parsing. The error must guide + // the user toward quoting or the programmatic API while preserving the + // underlying parser message. + let err = "User Id=sa;Password=a;b" + .parse::() + .err() + .unwrap(); + let msg = err.to_string(); + // Underlying terse parser message is preserved (not replaced). + assert!(msg.contains("must be joined"), "message was: {msg}"); + // Hint mentions the ADO quoting styles and the programmatic API. + assert!(msg.contains("single quotes"), "message was: {msg}"); + assert!(msg.contains("braces"), "message was: {msg}"); + assert!( + msg.contains("Config::from_ado_string"), + "message was: {msg}" + ); + // The `Conversion error:` prefix must not be doubled. + assert!( + !msg.contains("Conversion error: Conversion error:"), + "message was: {msg}" + ); + } + + #[test] + fn unclosed_quote_and_brace_errors_have_hint() { + // Unclosed quote/brace are quoting mistakes; the hint should attach. + for input in [ + "User Id=sa;Password='abc", // unclosed single quote + "User Id=sa;Password=\"abc", // unclosed double quote + "User Id=sa;Password={abc", // unclosed brace + ] { + let err = input.parse::().err().unwrap(); + let msg = err.to_string(); + assert!(msg.contains("must be quoted"), "message was: {msg}"); + assert!( + !msg.contains("Conversion error: Conversion error:"), + "message was: {msg}" + ); + } + } + #[test] #[cfg(any( feature = "rustls", @@ -465,7 +763,149 @@ mod tests { let test_str = ""; let ado: AdoNetConfig = test_str.parse()?; - assert_eq!(EncryptionLevel::Off, ado.encrypt()?); + assert_eq!(EncryptionLevel::Required, ado.encrypt()?); + + Ok(()) + } + + #[test] + #[cfg(feature = "tds80")] + fn encryption_parsing_strict() -> crate::Result<()> { + let test_str = "encrypt=strict"; + let ado: AdoNetConfig = test_str.parse()?; + + assert_eq!(EncryptionLevel::Strict, ado.encrypt()?); + + Ok(()) + } + + // No-TLS build: an explicit encryption request must error (not silently + // downgrade to plaintext, #305); opting out and an omitted keyword stay + // `NotSupported`. + + #[test] + #[cfg(not(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + )))] + fn encryption_parsing_on_errors_without_tls_backend() -> crate::Result<()> { + for test_str in ["encrypt=true", "encrypt=yes"] { + let ado: AdoNetConfig = test_str.parse()?; + let err = ado.encrypt().unwrap_err(); + assert!( + matches!(err, crate::Error::Tls(_)), + "expected Error::Tls for {test_str}, got {err:?}" + ); + let msg = err.to_string(); + assert!(msg.contains("without a TLS backend"), "message was: {msg}"); + } + + Ok(()) + } + + #[test] + #[cfg(not(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + )))] + fn encryption_parsing_strict_errors_without_tls_backend() { + let ado: AdoNetConfig = "encrypt=strict".parse().unwrap(); + let err = ado.encrypt().unwrap_err(); + assert!( + matches!(err, crate::Error::Tls(_)), + "expected Error::Tls, got {err:?}" + ); + } + + #[test] + #[cfg(not(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + )))] + fn encryption_parsing_mandatory_errors_without_tls_backend() { + // `mandatory` is not an accepted token in either build; the with-TLS + // parser rejects it as a bad boolean, so the no-TLS branch mirrors that + // (still an error, just not the TLS-missing one). + let ado: AdoNetConfig = "encrypt=mandatory".parse().unwrap(); + assert!(ado.encrypt().is_err()); + } + + #[test] + #[cfg(not(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + )))] + fn encryption_parsing_off_ok_without_tls_backend() -> crate::Result<()> { + for test_str in ["encrypt=false", "encrypt=no"] { + let ado: AdoNetConfig = test_str.parse()?; + assert_eq!(crate::EncryptionLevel::NotSupported, ado.encrypt()?); + } + + Ok(()) + } + + #[test] + #[cfg(not(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + )))] + fn encryption_parsing_plaintext_ok_without_tls_backend() -> crate::Result<()> { + let ado: AdoNetConfig = "encrypt=DANGER_PLAINTEXT".parse()?; + assert_eq!(crate::EncryptionLevel::NotSupported, ado.encrypt()?); + + Ok(()) + } + + #[test] + #[cfg(not(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + )))] + fn encryption_parsing_missing_ok_without_tls_backend() -> crate::Result<()> { + let ado: AdoNetConfig = "".parse()?; + assert_eq!(crate::EncryptionLevel::NotSupported, ado.encrypt()?); + + Ok(()) + } + + #[test] + fn client_name_parsing() -> crate::Result<()> { + let test_str = "workstationid=meow"; + let ado: AdoNetConfig = test_str.parse()?; + + assert_eq!(Some("meow".into()), ado.client_name()); + + let test_str = "Workstation ID=meow"; + let ado: AdoNetConfig = test_str.parse()?; + + assert_eq!(Some("meow".into()), ado.client_name()); + + Ok(()) + } + + #[test] + fn hostname_in_certificate_parsing() -> crate::Result<()> { + let test_str = "HostNameInCertificate=foo.example.com"; + let ado: AdoNetConfig = test_str.parse()?; + + assert_eq!( + Some("foo.example.com".into()), + ado.hostname_in_certificate() + ); + + let test_str = "HostName In Certificate=foo.example.com"; + let ado: AdoNetConfig = test_str.parse()?; + + assert_eq!( + Some("foo.example.com".into()), + ado.hostname_in_certificate() + ); Ok(()) } @@ -484,4 +924,64 @@ mod tests { Ok(()) } + + #[test] + fn application_intent_readonly_parsing() -> crate::Result<()> { + // Exact spelling from the ADO.NET connection string. + let ado: AdoNetConfig = "ApplicationIntent=ReadOnly".parse()?; + assert!(ado.readonly()); + + // ADO.NET treats the value case-insensitively. + let ado: AdoNetConfig = "applicationintent=readonly".parse()?; + assert!(ado.readonly()); + + // ReadWrite (the default) must not request read-only intent. + let ado: AdoNetConfig = "ApplicationIntent=ReadWrite".parse()?; + assert!(!ado.readonly()); + + // Absent altogether. + let ado: AdoNetConfig = "server=tcp:localhost,1433".parse()?; + assert!(!ado.readonly()); + + Ok(()) + } + + #[test] + fn multi_subnet_failover_parsing() -> crate::Result<()> { + let test_str = "MultiSubnetFailover=true"; + let ado: AdoNetConfig = test_str.parse()?; + assert!(ado.multi_subnet_failover()?); + + let test_str = "MultiSubnetFailover=yes"; + let ado: AdoNetConfig = test_str.parse()?; + assert!(ado.multi_subnet_failover()?); + + let test_str = "MultiSubnetFailover=false"; + let ado: AdoNetConfig = test_str.parse()?; + assert!(!ado.multi_subnet_failover()?); + + Ok(()) + } + + #[test] + fn multi_subnet_failover_parsing_missing() -> crate::Result<()> { + let test_str = ""; + let ado: AdoNetConfig = test_str.parse()?; + assert!(!ado.multi_subnet_failover()?); + + Ok(()) + } + + #[test] + fn multi_subnet_failover_from_ado_string() -> crate::Result<()> { + let config = crate::Config::from_ado_string( + "server=tcp:my-server.com,1433;MultiSubnetFailover=true", + )?; + assert!(config.get_multi_subnet_failover()); + + let config = crate::Config::from_ado_string("server=tcp:my-server.com,1433")?; + assert!(!config.get_multi_subnet_failover()); + + Ok(()) + } } diff --git a/src/client/config/ado_parser.rs b/src/client/config/ado_parser.rs new file mode 100644 index 000000000..ae3e7841f --- /dev/null +++ b/src/client/config/ado_parser.rs @@ -0,0 +1,426 @@ +//! In-house ADO.NET-style connection-string tokenizer. +//! +//! Tiberius parses ADO.NET connection strings as a **superset** of the rules +//! documented for .NET's `DbConnectionStringBuilder`: +//! +//! - pairs are separated by `;`; empty pairs (a stray `;`, trailing `;`) are +//! ignored; +//! - each pair is split on the **first** `=`, so a value may contain further +//! `=` characters (for example the `=` padding on a base64 password); a `==` +//! in the key position is a literal `=`; +//! - keys are trimmed and matched case-insensitively; +//! - an unquoted value is trimmed of surrounding whitespace; to keep `;`, a +//! quote, or leading/trailing spaces, quote the value with `'…'` or `"…"` and +//! double the enclosing quote to embed it (`"a""b"` → `a"b`); +//! - **extension (not ADO.NET):** `{…}` brace quoting is also accepted; the +//! value runs to the first `}` and therefore cannot itself contain `}`. +//! +//! Only ASCII input is accepted; build [`Config`] programmatically for anything +//! else. +//! +//! The returned error is a terse reason; the caller wraps it with the escaping +//! hint (see `hinted_conversion_error`), so the messages here stay short. +//! +//! [`Config`]: crate::Config + +use std::collections::HashMap; + +/// A parsed ADO.NET connection string: a map of lower-cased keys to their +/// unescaped values. +#[derive(Debug)] +pub(crate) struct AdoNetString { + pairs: HashMap, +} + +impl AdoNetString { + /// The parsed key/value pairs (keys lower-cased). + pub(crate) fn pairs(&self) -> &HashMap { + &self.pairs + } +} + +/// Parses an ADO.NET connection string. On failure returns a terse reason the +/// caller turns into a hinted [`crate::Error`]. +pub(crate) fn parse(input: &str) -> Result { + // The parser is byte-oriented for a single reason: every character it + // treats specially (`;`, `=`, `'`, `"`, `{`, `}`, whitespace) is ASCII, and + // rejecting non-ASCII up front lets the byte index and the char index stay + // identical without any multi-byte bookkeeping. + if !input.is_ascii() { + return Err("the connection string contains non-ASCII characters"); + } + + let b = input.as_bytes(); + let n = b.len(); + let mut i = 0; + let mut pairs = HashMap::new(); + + while i < n { + // Between pairs, skip separators and surrounding whitespace. + while i < n && (b[i] == b';' || b[i].is_ascii_whitespace()) { + i += 1; + } + if i >= n { + break; + } + + // Key: everything up to the first single `=`. A `==` is a literal `=`. + let mut key = String::new(); + loop { + if i >= n || b[i] == b';' { + return Err("key-value pairs must be joined by a `=`"); + } + if b[i] == b'=' { + if i + 1 < n && b[i + 1] == b'=' { + key.push('='); + i += 2; + continue; + } + i += 1; // consume the separating `=` + break; + } + key.push(b[i] as char); + i += 1; + } + let key = key + .trim_matches(|c: char| c.is_ascii_whitespace()) + .to_ascii_lowercase(); + if key.is_empty() { + return Err("a key in the connection string is empty"); + } + + // Value: leading whitespace before the value is not significant. + while i < n && b[i].is_ascii_whitespace() { + i += 1; + } + + let value = if i < n && (b[i] == b'\'' || b[i] == b'"') { + let quote = b[i]; + i += 1; + let mut v = String::new(); + loop { + if i >= n { + return Err(if quote == b'\'' { + "the connection string has an unclosed single quote" + } else { + "the connection string has an unclosed double quote" + }); + } + if b[i] == quote { + // A doubled quote is a literal quote; a lone quote closes. + if i + 1 < n && b[i + 1] == quote { + v.push(quote as char); + i += 2; + continue; + } + i += 1; + break; + } + v.push(b[i] as char); + i += 1; + } + skip_trailing(b, &mut i, "unexpected characters after a quoted value")?; + v + } else if i < n && b[i] == b'{' { + i += 1; + let mut v = String::new(); + loop { + if i >= n { + return Err("the connection string has an unclosed brace `{`"); + } + if b[i] == b'}' { + i += 1; + break; + } + v.push(b[i] as char); + i += 1; + } + skip_trailing( + b, + &mut i, + "unexpected characters after a `{…}` value; quote with '' or \"\" if the value contains `}`", + )?; + v + } else { + // Unquoted: runs to the next `;`, with surrounding whitespace + // trimmed (leading was already skipped above). Trimming uses the + // same ASCII-whitespace definition as the leading skip so a value + // is not truncated differently at its two ends. + let start = i; + while i < n && b[i] != b';' { + i += 1; + } + input[start..i] + .trim_end_matches(|c: char| c.is_ascii_whitespace()) + .to_string() + }; + + // Duplicate keys: last one wins, matching ADO.NET. + pairs.insert(key, value); + + if i < n && b[i] == b';' { + i += 1; + } + } + + Ok(AdoNetString { pairs }) +} + +/// After a closing quote or brace, only whitespace may precede the next `;` or +/// the end of the string; anything else is reported with `err`. +fn skip_trailing(b: &[u8], i: &mut usize, err: &'static str) -> Result<(), &'static str> { + while *i < b.len() && b[*i] != b';' { + if !b[*i].is_ascii_whitespace() { + return Err(err); + } + *i += 1; + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::parse; + + /// Parse and return the single value stored under `key`, panicking with the + /// parser error if parsing failed. + fn value(s: &str, key: &str) -> Option { + parse(s) + .unwrap_or_else(|e| panic!("`{s}` failed to parse: {e}")) + .pairs() + .get(key) + .cloned() + } + + #[test] + fn basic_pairs() { + let p = parse("Server=host;Database=db;User Id=sa;Password=pw").unwrap(); + assert_eq!(p.pairs().get("server").map(String::as_str), Some("host")); + assert_eq!(p.pairs().get("database").map(String::as_str), Some("db")); + assert_eq!(p.pairs().get("user id").map(String::as_str), Some("sa")); + assert_eq!(p.pairs().get("password").map(String::as_str), Some("pw")); + } + + #[test] + fn keys_are_lowercased_and_trimmed() { + assert_eq!(value(" SeRVeR =host", "server").as_deref(), Some("host")); + assert_eq!(value("USER ID=sa", "user id").as_deref(), Some("sa")); + } + + #[test] + fn splits_each_pair_on_the_first_equals() { + // A value may contain further `=` (base64 padding etc.); see #313. + assert_eq!( + value("Password=ab=cd", "password").as_deref(), + Some("ab=cd") + ); + assert_eq!( + value("Password=Zm9vYmFy==", "password").as_deref(), + Some("Zm9vYmFy==") + ); + assert_eq!(value("k=a=b=c=", "k").as_deref(), Some("a=b=c=")); + } + + #[test] + fn double_equals_in_key_is_a_literal_equals() { + // `k==v` reads as the key `k=v` with no separator, hence an error; a + // value that must *start* with `=` has to be quoted instead. + assert!(parse("k==v").is_err()); + assert_eq!(value("k='=v'", "k").as_deref(), Some("=v")); + assert_eq!(value("a==b=c", "a=b").as_deref(), Some("c")); + } + + #[test] + fn empty_pairs_are_ignored() { + assert_eq!(value(";;Server=host;;", "server").as_deref(), Some("host")); + assert_eq!(value("Server=host;", "server").as_deref(), Some("host")); + assert_eq!(value(";Server=host", "server").as_deref(), Some("host")); + assert!(parse("").unwrap().pairs().is_empty()); + assert!(parse(" ").unwrap().pairs().is_empty()); + assert!(parse(";;;").unwrap().pairs().is_empty()); + } + + #[test] + fn whitespace_between_pairs_is_ignored() { + let p = parse("Server=host ; Database=db").unwrap(); + assert_eq!(p.pairs().get("server").map(String::as_str), Some("host")); + assert_eq!(p.pairs().get("database").map(String::as_str), Some("db")); + } + + #[test] + fn unquoted_values_are_trimmed_both_sides() { + assert_eq!(value("k= a b ", "k").as_deref(), Some("a b")); + assert_eq!(value("k=", "k").as_deref(), Some("")); + assert_eq!(value("k= ", "k").as_deref(), Some("")); + } + + #[test] + fn quotes_preserve_surrounding_whitespace() { + assert_eq!(value("k=' a '", "k").as_deref(), Some(" a ")); + assert_eq!(value("k=\" a \"", "k").as_deref(), Some(" a ")); + assert_eq!(value("k=' '", "k").as_deref(), Some(" ")); + } + + #[test] + fn quotes_hold_structural_characters() { + assert_eq!(value("k='a;b=c{d'", "k").as_deref(), Some("a;b=c{d")); + assert_eq!(value("k=\"a;b=c{d\"", "k").as_deref(), Some("a;b=c{d")); + // The other quote character is a literal inside. + assert_eq!(value("k='a\"b'", "k").as_deref(), Some("a\"b")); + assert_eq!(value("k=\"a'b\"", "k").as_deref(), Some("a'b")); + } + + #[test] + fn doubled_quote_embeds_the_quote() { + assert_eq!(value("k='a''b'", "k").as_deref(), Some("a'b")); + assert_eq!(value("k=\"a\"\"b\"", "k").as_deref(), Some("a\"b")); + assert_eq!(value("k=''''", "k").as_deref(), Some("'")); // '' '' -> one ' + } + + #[test] + fn braces_hold_structural_characters() { + assert_eq!(value("k={a;b=c}", "k").as_deref(), Some("a;b=c")); + assert_eq!(value("k={a'b\"c}", "k").as_deref(), Some("a'b\"c")); + } + + #[test] + fn brace_cannot_contain_close_brace() { + // Closes at the first `}`; trailing content is an error, not folded in. + assert_eq!( + parse("k={a}b}").unwrap_err(), + "unexpected characters after a `{…}` value; quote with '' or \"\" if the value contains `}`" + ); + assert_eq!(value("k={a}", "k").as_deref(), Some("a")); + } + + #[test] + fn trailing_junk_after_quote_is_rejected() { + assert_eq!( + parse("k='ab'x").unwrap_err(), + "unexpected characters after a quoted value" + ); + assert_eq!( + parse("k=\"ab\"x").unwrap_err(), + "unexpected characters after a quoted value" + ); + // Trailing whitespace after a closing quote is fine. + assert_eq!(value("k='ab' ;n=1", "k").as_deref(), Some("ab")); + } + + #[test] + fn empty_quoted_value() { + assert_eq!(value("k=''", "k").as_deref(), Some("")); + assert_eq!(value("k=\"\"", "k").as_deref(), Some("")); + assert_eq!(value("k={}", "k").as_deref(), Some("")); + } + + #[test] + fn ascii_control_bytes_are_not_treated_as_trim_whitespace() { + // A trailing vertical tab (0x0B) is content, not whitespace: it must be + // preserved, consistently with the leading-whitespace skip which also + // ignores 0x0B. (str::trim_end would wrongly strip it.) + assert_eq!( + value("k=secret\u{000B}", "k").as_deref(), + Some("secret\u{000B}") + ); + // Ordinary ASCII whitespace is still trimmed when unquoted. + assert_eq!(value("k=secret\t ", "k").as_deref(), Some("secret")); + // Keys use the same ASCII-whitespace definition: ASCII spaces are + // trimmed, but a 0x0B is content (kept), consistent with values. + assert_eq!(value(" k =v", "k").as_deref(), Some("v")); + assert_eq!(value("k\u{000B}=v", "k\u{000B}").as_deref(), Some("v")); + } + + #[test] + fn duplicate_keys_last_wins() { + assert_eq!(value("k=a;k=b;k=c", "k").as_deref(), Some("c")); + } + + #[test] + fn server_value_passes_through_raw() { + assert_eq!( + value("Server=tcp:host.example.com,1433", "server").as_deref(), + Some("tcp:host.example.com,1433") + ); + assert_eq!( + value("Server=host\\instance", "server").as_deref(), + Some("host\\instance") + ); + } + + #[test] + fn error_reasons() { + assert_eq!( + parse("=value").unwrap_err(), + "a key in the connection string is empty" + ); + assert_eq!( + parse("k=a;b").unwrap_err(), + "key-value pairs must be joined by a `=`" + ); + assert_eq!( + parse("k='ab").unwrap_err(), + "the connection string has an unclosed single quote" + ); + assert_eq!( + parse("k=\"ab").unwrap_err(), + "the connection string has an unclosed double quote" + ); + assert_eq!( + parse("k={ab").unwrap_err(), + "the connection string has an unclosed brace `{`" + ); + assert!(parse("k=café").is_err()); + } + + // Every printable-ASCII character must round-trip verbatim inside single + // quotes (with `'` doubled) and inside double quotes (with `"` doubled). + #[test] + fn exhaustive_ascii_roundtrip_in_quotes() { + for byte in 0x20u8..=0x7e { + let c = byte as char; + + let inner = if c == '\'' { + "a''b".to_string() + } else { + format!("a{c}b") + }; + let expected = format!("a{c}b"); + assert_eq!( + value(&format!("k='{inner}'"), "k").as_deref(), + Some(expected.as_str()), + "char {c:?} (0x{byte:02x}) via single quotes" + ); + + let inner = if c == '"' { + "a\"\"b".to_string() + } else { + format!("a{c}b") + }; + assert_eq!( + value(&format!("k=\"{inner}\""), "k").as_deref(), + Some(expected.as_str()), + "char {c:?} (0x{byte:02x}) via double quotes" + ); + } + } + + // Unquoted values round-trip verbatim for every printable-ASCII character + // that is not structural in that position (`;`, and a leading `'`/`"`/`{`). + #[test] + fn exhaustive_ascii_roundtrip_unquoted_where_legal() { + for byte in 0x21u8..=0x7e { + let c = byte as char; + if c == ';' { + continue; // separator: cannot appear unquoted + } + // Placed after a leading letter so `'`/`"`/`{` are not value-leading. + let raw = format!("a{c}b"); + assert_eq!( + value(&format!("k={raw}"), "k").as_deref(), + Some(raw.as_str()), + "char {c:?} (0x{byte:02x}) unquoted" + ); + } + } +} diff --git a/src/client/config/jdbc.rs b/src/client/config/jdbc.rs index 4168cf975..d9577ee72 100644 --- a/src/client/config/jdbc.rs +++ b/src/client/config/jdbc.rs @@ -11,7 +11,14 @@ impl FromStr for JdbcConfig { type Err = Error; fn from_str(s: &str) -> crate::Result { - let config = s.parse()?; + let config = s.parse().map_err(|e| { + super::connection_string_error( + e, + "wrap the value in braces (e.g. `password={p@ss;word}`); single \ + and double quotes are ordinary characters in JDBC, not escapes", + "Config::from_jdbc_string", + ) + })?; Ok(Self { config }) } } @@ -259,6 +266,141 @@ mod tests { Ok(()) } + #[test] + fn parsing_password_special_characters() -> crate::Result<()> { + // (value as written after `Password=`, exact parsed password). Unlike + // ADO.NET, the JDBC parser only supports brace (`{...}`) quoting — not + // single or double quotes — so `'` and `"` are ordinary characters. + let cases: &[(&str, &str)] = &[ + // Braces are required for structural characters. + ("{a;b}", "a;b"), + ("{a=b}", "a=b"), + ("{a:b}", "a:b"), + ("{a b}", "a b"), + ("{a$b}", "a$b"), + ("{a'b}", "a'b"), + ("{a\"b}", "a\"b"), + // `\`, `/`, `[` and `]` are structural in JDBC values and need + // braces (they would otherwise terminate the value). + ("{a\\b}", "a\\b"), + ("{a/b}", "a/b"), + ("{a[b]}", "a[b]"), + // Non-structural characters pass through unquoted, including `}`, + // `'` and `"` (which are not special to the JDBC parser). + ("a$b", "a$b"), + ("a@b", "a@b"), + ("a!b", "a!b"), + ("a#b", "a#b"), + ("a%b", "a%b"), + ("a&b", "a&b"), + ("a,b", "a,b"), + ("a'b", "a'b"), + ("a\"b", "a\"b"), + ("a}b", "a}b"), + ("a b", "a b"), + // Percent-encoding is NOT decoded: `%24` stays literal. + ("a%24b", "a%24b"), + ]; + + for (value, expected) in cases { + let s = format!("jdbc:sqlserver://host.com:1433;User ID=sa;Password={value}"); + let jdbc: JdbcConfig = s.parse()?; + assert_eq!( + AuthMethod::sql_server("sa", *expected), + jdbc.authentication()?, + "value `{value}` should parse to password `{expected}`" + ); + } + + Ok(()) + } + + #[test] + fn parsing_password_preserves_surrounding_whitespace() -> crate::Result<()> { + // Unlike the ADO parser, JDBC does not trim; braced whitespace is kept. + let jdbc: JdbcConfig = + "jdbc:sqlserver://host.com:1433;User ID=sa;Password={ a }".parse()?; + assert_eq!(AuthMethod::sql_server("sa", " a "), jdbc.authentication()?); + + Ok(()) + } + + #[test] + fn parsing_password_brace_cannot_hold_close_brace() -> crate::Result<()> { + // Like ADO.NET, JDBC brace quoting has no `}}` doubling: the brace + // closes at the first `}`, so the inner `}` is consumed as the + // terminator. Since JDBC has no quote escaping, a value needing braces + // that also contains `}` cannot be represented — document the limit. + let jdbc: JdbcConfig = + "jdbc:sqlserver://host.com:1433;User ID=sa;Password={a}b}".parse()?; + assert_eq!(AuthMethod::sql_server("sa", "ab}"), jdbc.authentication()?); + + Ok(()) + } + + #[test] + fn parsing_password_non_ascii_is_rejected() { + // The JDBC parser only accepts ASCII, in every position; non-ASCII + // passwords must use the programmatic `Config` API instead. + assert!("jdbc:sqlserver://host.com:1433;User ID=sa;Password=café" + .parse::() + .is_err()); + // Braces are JDBC's only quoting mechanism; non-ASCII is rejected there + // too (mirrors the ADO quoted-value test). + assert!("jdbc:sqlserver://host.com:1433;User ID=sa;Password={café}" + .parse::() + .is_err()); + } + + #[test] + fn empty_password_is_rejected() { + // Unlike the ADO parser (which accepts an empty value), the JDBC parser + // requires a non-empty property value, so `Password=` fails to parse. + assert!("jdbc:sqlserver://host.com:1433;User ID=sa;Password=" + .parse::() + .is_err()); + } + + #[test] + fn unquoted_special_char_password_error_has_hint() { + // An unquoted `;` in a password breaks parsing. The error must guide + // the user toward brace quoting (JDBC's only escape) or the programmatic + // API, while preserving the underlying parser message. + let err = "jdbc:sqlserver://host.com:1433;User ID=sa;Password=a;b" + .parse::() + .err() + .unwrap(); + let msg = err.to_string(); + // Underlying terse parser message is preserved (not replaced). + assert!(msg.contains("must be joined"), "message was: {msg}"); + // JDBC guidance must steer toward braces and its own constructor docs. + assert!(msg.contains("braces"), "message was: {msg}"); + assert!( + msg.contains("Config::from_jdbc_string"), + "message was: {msg}" + ); + // The `Conversion error:` prefix must not be doubled. + assert!( + !msg.contains("Conversion error: Conversion error:"), + "message was: {msg}" + ); + } + + #[test] + fn unclosed_brace_error_has_hint() { + // An unclosed brace is a quoting mistake; the hint should attach. + let err = "jdbc:sqlserver://host.com:1433;User ID=sa;Password={abc" + .parse::() + .err() + .unwrap(); + let msg = err.to_string(); + assert!(msg.contains("must be quoted"), "message was: {msg}"); + assert!( + !msg.contains("Conversion error: Conversion error:"), + "message was: {msg}" + ); + } + #[test] #[cfg(any( feature = "rustls", @@ -314,7 +456,109 @@ mod tests { let test_str = "jdbc:sqlserver://my-server.com:4200;"; let jdbc: JdbcConfig = test_str.parse()?; - assert_eq!(EncryptionLevel::Off, jdbc.encrypt()?); + assert_eq!(EncryptionLevel::Required, jdbc.encrypt()?); + + Ok(()) + } + + // No-TLS build: an explicit encryption request must error (not silently + // downgrade to plaintext, #305); opting out and an omitted keyword stay + // `NotSupported`. + + #[test] + #[cfg(not(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + )))] + fn encryption_parsing_on_errors_without_tls_backend() -> crate::Result<()> { + for enc in ["true", "yes"] { + let test_str = format!("jdbc:sqlserver://my-server.com:4200;encrypt={enc};"); + let jdbc: JdbcConfig = test_str.parse()?; + let err = jdbc.encrypt().unwrap_err(); + assert!( + matches!(err, crate::Error::Tls(_)), + "expected Error::Tls for {test_str}, got {err:?}" + ); + let msg = err.to_string(); + assert!(msg.contains("without a TLS backend"), "message was: {msg}"); + } + + Ok(()) + } + + #[test] + #[cfg(not(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + )))] + fn encryption_parsing_strict_errors_without_tls_backend() { + let jdbc: JdbcConfig = "jdbc:sqlserver://my-server.com:4200;encrypt=strict;" + .parse() + .unwrap(); + let err = jdbc.encrypt().unwrap_err(); + assert!( + matches!(err, crate::Error::Tls(_)), + "expected Error::Tls, got {err:?}" + ); + } + + #[test] + #[cfg(not(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + )))] + fn encryption_parsing_mandatory_errors_without_tls_backend() { + // `mandatory` is not an accepted token in either build; the with-TLS + // parser rejects it as a bad boolean, so the no-TLS branch mirrors that + // (still an error, just not the TLS-missing one). + let jdbc: JdbcConfig = "jdbc:sqlserver://my-server.com:4200;encrypt=mandatory;" + .parse() + .unwrap(); + assert!(jdbc.encrypt().is_err()); + } + + #[test] + #[cfg(not(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + )))] + fn encryption_parsing_off_ok_without_tls_backend() -> crate::Result<()> { + for enc in ["false", "no"] { + let test_str = format!("jdbc:sqlserver://my-server.com:4200;encrypt={enc};"); + let jdbc: JdbcConfig = test_str.parse()?; + assert_eq!(crate::EncryptionLevel::NotSupported, jdbc.encrypt()?); + } + + Ok(()) + } + + #[test] + #[cfg(not(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + )))] + fn encryption_parsing_plaintext_ok_without_tls_backend() -> crate::Result<()> { + let jdbc: JdbcConfig = + "jdbc:sqlserver://my-server.com:4200;encrypt=DANGER_PLAINTEXT;".parse()?; + assert_eq!(crate::EncryptionLevel::NotSupported, jdbc.encrypt()?); + + Ok(()) + } + + #[test] + #[cfg(not(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + )))] + fn encryption_parsing_missing_ok_without_tls_backend() -> crate::Result<()> { + let jdbc: JdbcConfig = "jdbc:sqlserver://my-server.com:4200;".parse()?; + assert_eq!(crate::EncryptionLevel::NotSupported, jdbc.encrypt()?); Ok(()) } diff --git a/src/client/connection.rs b/src/client/connection.rs index 6467a8eec..707f94f34 100644 --- a/src/client/connection.rs +++ b/src/client/connection.rs @@ -18,9 +18,14 @@ use crate::{ }; use asynchronous_codec::Framed; use bytes::BytesMut; -#[cfg(any(windows, feature = "integrated-auth-gssapi", feature = "winauth"))] +#[cfg(any( + windows, + feature = "winauth", + feature = "integrated-auth-gssapi", + feature = "sspi-rs" +))] use codec::TokenSspi; -use futures_util::io::{AsyncRead, AsyncWrite}; +use futures_util::io::{AsyncRead, AsyncWrite, AsyncWriteExt}; use futures_util::ready; use futures_util::sink::SinkExt; use futures_util::stream::{Stream, TryStream, TryStreamExt}; @@ -32,15 +37,24 @@ use libgssapi::{ oid::{OidSet, GSS_MECH_KRB5, GSS_NT_KRB5_PRINCIPAL}, }; use pretty_hex::*; +use secrecy::ExposeSecret; +#[cfg(all(unix, feature = "sspi-rs"))] +use sspi::{ + AuthIdentity, BufferType, ClientRequestFlags, CredentialUse, DataRepresentation, Ntlm, + SecurityBuffer, Sspi, SspiImpl, Username, +}; #[cfg(all(unix, feature = "integrated-auth-gssapi"))] use std::ops::Deref; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::Arc; use std::{cmp, fmt::Debug, io, pin::Pin, task}; use task::Poll; use tracing::{event, Level}; #[cfg(all(windows, feature = "winauth"))] use winauth::windows::NtlmSspiBuilder; -#[cfg(feature = "winauth")] +#[cfg(all(feature = "winauth", not(all(unix, feature = "sspi-rs"))))] use winauth::NextBytes; +use zeroize::{Zeroize, Zeroizing}; /// A `Connection` is an abstraction between the [`Client`] and the server. It /// can be used as a `Stream` to fetch [`Packet`]s from and to `send` packets @@ -59,6 +73,23 @@ where flushed: bool, context: Context, buf: BytesMut, + /// Set for the duration of a multi-packet write. A message is only partly + /// on the wire while this is `true`; if the writing future is dropped + /// (a cancelled `query`/`execute`, a `select!` losing the race, a + /// `tokio::time::timeout` firing) the flag stays set, so the next write on + /// the same connection fails cleanly instead of appending a second message + /// after a half-sent one and silently desyncing the server. + poisoned: bool, + /// Set when a `command_timeout` fires mid-response. The tail of the + /// timed-out response is then still unread on the wire, so the connection is + /// out of sync with the server and must not be reused. This handle is shared + /// with the token stream's `RoundTripTimeout`, which flips it when its + /// deadline elapses; [`ensure_not_poisoned`] rejects every subsequent use so + /// a pool discards the connection instead of hanging in `flush_stream` on + /// the still-pending previous response. + /// + /// [`ensure_not_poisoned`]: Self::ensure_not_poisoned + command_desync: Arc, } impl Debug for Connection { @@ -73,14 +104,76 @@ impl Debug for Connection { } impl Connection { - /// Creates a new connection + /// Creates a new connection. + /// + /// Note: `tcp_stream` is a connected stream, so some parts of the + /// [`Config`] need to be handled outside of this method. + /// + /// The handshake performed here (prelogin, TLS negotiation and login) is + /// bounded by [`Config::handshake_timeout`]: if the server accepts the TCP + /// connection but then stalls mid-handshake (for example a TLS handshake + /// that never completes), the connect future fails with a + /// [`std::io::ErrorKind::TimedOut`] error instead of hanging forever. Enable + /// `tracing` at `DEBUG` to see which stage was last reached. pub(crate) async fn connect(config: Config, tcp_stream: S) -> crate::Result> { + let handshake_timeout = config.handshake_timeout; + with_optional_timeout(handshake_timeout, Self::establish(config, tcp_stream)).await + } + + /// Performs the full connection handshake (prelogin, TLS negotiation and + /// login) over an already-connected `tcp_stream`. + /// + /// Split out from [`connect`](Self::connect) so the whole handshake can be + /// wrapped in a single [`Config::handshake_timeout`] bound; on its own it + /// runs unbounded and would block forever if the server stalls. + async fn establish(config: Config, tcp_stream: S) -> crate::Result> { + // Captured before `config` is consumed below and applied to the + // `Context` only *after* the handshake completes (see the end of this + // method). Arming it up front would also bound the login-ack/SSPI drain + // that runs through the same token stream (`flush_done`/`flush_sspi`) + // during connect, which is wrong: that whole handshake is governed by + // `handshake_timeout` alone, so `handshake_timeout(None)` must genuinely + // wait indefinitely regardless of `command_timeout`. + let command_timeout = config.command_timeout; + let context = { let mut context = Context::new(); context.set_spn(config.get_host(), config.get_port()); + // Row-decode preference; unrelated to handshake timing, so set it up + // front alongside the SPN. + context.set_lossy_utf16(config.lossy_utf16_decoding); + context.set_lossy_codepage(config.lossy_codepage_decoding); context }; + // In TDS 8.0 "strict" mode the TLS handshake happens *before* the + // prelogin, so we wrap the stream in TLS up front. In every other mode + // the connection starts in the clear and TLS (if any) is negotiated + // during the prelogin. + #[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + ))] + let transport = match config.encryption { + EncryptionLevel::Strict => { + event!(Level::DEBUG, "Performing a TLS handshake (TDS 8.0 strict)"); + let mut pre_login_stream = TlsPreloginWrapper::new(tcp_stream); + // No prelogin framing is used for the strict handshake; pass the + // raw TLS bytes straight through. + pre_login_stream.handshake_complete(); + let stream = create_tls_stream(&config, pre_login_stream).await?; + event!(Level::DEBUG, "TLS handshake successful"); + Framed::new(MaybeTlsStream::Tls(stream), PacketCodec) + } + _ => Framed::new(MaybeTlsStream::Raw(tcp_stream), PacketCodec), + }; + + #[cfg(not(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + )))] let transport = Framed::new(MaybeTlsStream::Raw(tcp_stream), PacketCodec); let mut connection = Self { @@ -88,18 +181,31 @@ impl Connection { context, flushed: false, buf: BytesMut::new(), + poisoned: false, + command_desync: Arc::new(AtomicBool::new(false)), }; let fed_auth_required = matches!(config.auth, AuthMethod::AADToken(_)); + event!(Level::DEBUG, "Handshake stage: sending TDS prelogin"); let prelogin = connection - .prelogin(config.encryption, fed_auth_required) + .prelogin( + config.encryption, + fed_auth_required, + config.instance_name.clone(), + ) .await?; - let encryption = prelogin.negotiated_encryption(config.encryption); + let encryption = prelogin.negotiated_encryption(config.encryption)?; + event!( + Level::DEBUG, + "Handshake stage: prelogin complete, negotiated encryption = {:?}", + encryption + ); let connection = connection.tls_handshake(&config, encryption).await?; + event!(Level::DEBUG, "Handshake stage: sending login"); let mut connection = connection .login( config.auth, @@ -107,12 +213,21 @@ impl Connection { config.database, config.host, config.application_name, + config.client_name, config.readonly, + config.packet_size, prelogin, ) .await?; connection.flush_done().await?; + event!(Level::DEBUG, "Handshake stage: login complete"); + + // The handshake is done; arm the per-response command timeout now so it + // only ever bounds command result reads, never the handshake above. + connection + .context_mut() + .set_command_timeout(command_timeout); Ok(connection) } @@ -122,7 +237,12 @@ impl Connection { TokenStream::new(self).flush_done().await } - #[cfg(any(windows, feature = "integrated-auth-gssapi", feature = "winauth"))] + #[cfg(any( + windows, + feature = "winauth", + feature = "integrated-auth-gssapi", + feature = "sspi-rs" + ))] /// Flush the incoming token stream until receiving `SSPI` token. async fn flush_sspi(&mut self) -> crate::Result { TokenStream::new(self).flush_sspi().await @@ -168,12 +288,17 @@ impl Connection { where E: Sized + Encode, { + self.ensure_not_poisoned()?; self.flushed = false; let packet_size = (self.context.packet_size() as usize) - HEADER_BYTES; let mut payload = BytesMut::new(); item.encode(&mut payload)?; + // Mark the connection poisoned across the multi-packet write; a clean + // completion clears it below. A future dropped mid-loop leaves it set. + self.poisoned = true; + while !payload.is_empty() { let writable = cmp::min(payload.len(), packet_size); let split_payload = payload.split_to(writable); @@ -194,6 +319,95 @@ impl Connection { } self.flush_sink().await?; + self.poisoned = false; + + Ok(()) + } + + /// Returns an error if a previous multi-packet write on this connection was + /// interrupted (e.g. the query/execute future was cancelled), which would + /// have left a partial message on the wire. The connection cannot be safely + /// reused in that state and should be dropped. + pub(crate) fn ensure_not_poisoned(&self) -> crate::Result<()> { + if self.poisoned { + return Err(crate::Error::Protocol( + "connection was left in an inconsistent state by a cancelled write and can no longer be used; open a new connection" + .into(), + )); + } + if self.command_desync.load(Ordering::Acquire) { + return Err(crate::Error::Protocol( + "connection was left out of sync with the server by a command timeout (the previous response is only partly read) and can no longer be used; open a new connection" + .into(), + )); + } + Ok(()) + } + + /// A handle to the command-timeout desync flag, shared with the token stream + /// so it can mark the connection unusable when a `command_timeout` fires + /// mid-response. See the `command_desync` field. + pub(crate) fn command_desync_flag(&self) -> Arc { + Arc::clone(&self.command_desync) + } + + /// Marks the connection poisoned for the duration of a multi-packet write. + /// + /// Exposed to the bulk-load write path (`BulkLoadRequest`), which writes its + /// messages by calling [`write_to_wire`] directly rather than through + /// [`send`]. It must bracket those write loops with `poison`/[`unpoison`] so + /// that a dropped future (cancelled bulk insert, `select!`, `timeout`) leaves + /// the connection unusable instead of silently reused after a half-sent + /// message — the same guarantee `send`/`cancel_request` get from their inline + /// bracketing. This is deliberately *not* folded into `write_to_wire`: that + /// method is called once per packet inside `send`'s loop, so clearing the + /// flag there would drop the poison between packets of a multi-packet send + /// and destroy the very protection `send` relies on. + /// + /// [`write_to_wire`]: Self::write_to_wire + /// [`send`]: Self::send + /// [`unpoison`]: Self::unpoison + pub(crate) fn poison(&mut self) { + self.poisoned = true; + } + + /// Clears the poisoned flag after a multi-packet write completes cleanly. + /// The counterpart to [`poison`](Self::poison); see its docs. + pub(crate) fn unpoison(&mut self) { + self.poisoned = false; + } + + async fn send_sensitive_login( + &mut self, + header: PacketHeader, + mut payload: Zeroizing>, + ) -> crate::Result<()> { + self.ensure_not_poisoned()?; + self.flushed = false; + let packet_size = (self.context.packet_size() as usize) - HEADER_BYTES; + + // Frame the login into zeroizable packets off the shared `BytesMut` + // path. Each frame is a `Zeroizing` wiped on drop at the end of its + // loop iteration, so no explicit call is needed there; `payload` is + // zeroized right after framing to drop the plaintext copy before the + // network-write loop's awaits. + let frames = frame_sensitive_login(header, &payload[..], packet_size)?; + payload.zeroize(); + + // Mark the connection poisoned across the multi-frame write, exactly + // like `send`/`cancel_request`: if the writing future is dropped + // mid-loop the login is only partly on the wire and the connection must + // not be silently reused. A clean flush clears it below. + self.poisoned = true; + + for frame in frames { + event!(Level::TRACE, "Sending a packet ({} bytes)", frame.len(),); + + self.transport.write_all(&frame[..]).await?; + } + + (&mut *self.transport).flush().await?; + self.poisoned = false; Ok(()) } @@ -222,6 +436,43 @@ impl Connection { self.transport.flush().await } + /// Sends a TDS Attention signal (packet type `0x06`, MS-TDS section + /// 2.2.1.6) to request cancellation of the request currently in flight on + /// this connection, then drains the token stream until the acknowledging + /// DONE token (with the `DONE_ATTN` status bit set) is received. + /// + /// The Attention message carries no payload, so it is written to the wire + /// as a single end-of-message packet. Draining the acknowledgement leaves + /// the connection clean and ready to be reused for further queries. + pub(crate) async fn cancel_request(&mut self) -> crate::Result { + // Consistency with the other write paths (`send`, + // `send_sensitive_login`, `flush_stream`): if a previous multi-packet + // write was interrupted, the connection is already desynced with a + // partial message sitting on the wire. Appending an Attention packet + // after that half-sent message would corrupt the stream further and the + // acknowledging DONE could never be matched, so fail fast instead. A + // legitimate cancel targets an in-flight request, which means the prior + // `send` completed and cleared the flag, so this guard never rejects a + // valid cancellation. + self.ensure_not_poisoned()?; + + let id = self.context.next_packet_id(); + let header = PacketHeader::attention(id); + + // Mark the connection poisoned across the write, exactly like `send`: + // the Attention is a single end-of-message packet, but if the writing + // future is dropped mid-flight the header is only partly on the wire and + // the connection must not be silently reused. A clean flush clears it. + self.poisoned = true; + + // Attention has an empty payload; send just the 8-byte header. + self.write_to_wire(header, BytesMut::new()).await?; + self.flush_sink().await?; + self.poisoned = false; + + TokenStream::new(self).flush_done_attention().await + } + /// Cleans the packet stream from previous use. It is important to use the /// whole stream before using the connection again. Flushing the stream /// makes sure we don't have any old data causing undefined behaviour after @@ -230,22 +481,42 @@ impl Connection { /// Calling this will slow down the queries if stream is still dirty if all /// results are not handled. pub async fn flush_stream(&mut self) -> crate::Result<()> { + // If a previous write was cancelled mid-message the connection is + // already known-bad; fail fast rather than layering a new request on + // top of it. + self.ensure_not_poisoned()?; + + // Discard any partially-consumed packet payload, then drain whole + // packets up to the end-of-message marker. Truncating `buf` and + // re-reading on packet boundaries resynchronises the token stream even + // if a previous result stream was dropped part-way through a value + // (the lost bytes belonged to a packet we are discarding anyway). self.buf.truncate(0); if self.flushed { return Ok(()); } - while let Some(packet) = self.try_next().await? { - event!( - Level::WARN, - "Flushing unhandled packet from the wire. Please consume your streams!", - ); + loop { + match self.try_next().await { + Ok(Some(packet)) => { + event!( + Level::WARN, + "Flushing unhandled packet from the wire. Please consume your streams!", + ); - let is_last = packet.is_last(); - - if is_last { - break; + if packet.is_last() { + break; + } + } + Ok(None) => break, + // The stream could not be drained cleanly (e.g. it was + // abandoned at an unrecoverable offset). Poison the connection + // so it is not silently reused in an inconsistent state. + Err(e) => { + self.poisoned = true; + return Err(e); + } } } @@ -270,10 +541,12 @@ impl Connection { &mut self, encryption: EncryptionLevel, fed_auth_required: bool, + instance_name: Option, ) -> crate::Result { let mut msg = PreloginMessage::new(); msg.encryption = encryption; msg.fed_auth_required = fed_auth_required; + msg.instance_name = instance_name.clone(); let id = self.context.next_packet_id(); self.send(PacketHeader::pre_login(id), msg).await?; @@ -281,20 +554,24 @@ impl Connection { let response: PreloginMessage = codec::collect_from(self).await?; // threadid (should be empty when sent from server to client) debug_assert_eq!(response.thread_id, 0); + // ensure the server accepted the instance we asked it to validate + response.validate_instance(instance_name.as_deref())?; Ok(response) } /// Defines the login record rules with SQL Server. Authentication with /// connection options. #[allow(clippy::too_many_arguments)] - async fn login<'a>( + async fn login( mut self, auth: AuthMethod, encryption: EncryptionLevel, db: Option, server_name: Option, application_name: Option, + client_name: Option, readonly: bool, + packet_size: Option, prelogin: PreloginMessage, ) -> crate::Result { let mut login_message = LoginMessage::new(); @@ -311,8 +588,16 @@ impl Connection { login_message.app_name(app_name); } + if let Some(client_name) = client_name { + login_message.hostname(client_name); + } + login_message.readonly(readonly); + if let Some(size) = packet_size { + login_message.packet_size(size); + } + match auth { #[cfg(all(windows, feature = "winauth"))] AuthMethod::Integrated => { @@ -334,31 +619,37 @@ impl Connection { event!(Level::TRACE, sspi_response_len = sspi_response.len()); let id = self.context.next_packet_id(); - let header = PacketHeader::login(id); + let header = PacketHeader::sspi(id); let token = TokenSspi::new(sspi_response); self.send(header, token).await?; } - None => unreachable!(), + None => { + return Err(crate::Error::Protocol( + "NTLM handshake produced no response to the server challenge".into(), + )) + } } } #[cfg(all(unix, feature = "integrated-auth-gssapi"))] AuthMethod::Integrated => { - let mut s = OidSet::new()?; - s.add(&GSS_MECH_KRB5)?; + let mut s = OidSet::new(); + s.add(GSS_MECH_KRB5)?; let client_cred = Cred::acquire(None, None, CredUsage::Initiate, Some(&s))?; let mut ctx = ClientCtx::new( Some(client_cred), - Name::new(self.context.spn().as_bytes(), Some(&GSS_NT_KRB5_PRINCIPAL))?, + Name::new(self.context.spn().as_bytes(), Some(GSS_NT_KRB5_PRINCIPAL))?, CtxFlags::GSS_C_MUTUAL_FLAG | CtxFlags::GSS_C_SEQUENCE_FLAG, None, ); - let init_token = ctx.step(None, None)?; + let init_token = ctx.step(None, None)?.ok_or_else(|| { + crate::Error::Protocol("GSSAPI produced no initial token".into()) + })?; - login_message.integrated_security(Some(Vec::from(init_token.unwrap().deref()))); + login_message.integrated_security(Some(Vec::from(init_token.deref()))); let id = self.context.next_packet_id(); self.send(PacketHeader::login(id), login_message).await?; @@ -383,11 +674,114 @@ impl Connection { self.send(header, next_token).await?; } - #[cfg(feature = "winauth")] + #[cfg(all(unix, feature = "sspi-rs"))] + AuthMethod::Windows(auth) => { + let mut ntlm = Ntlm::new(); + + let username = + Username::new(&auth.user, auth.domain.as_deref()).map_err(sspi::Error::from)?; + + // `auth.password` is a `SecretString`; sspi requires a + // plaintext `String`, but that single copy is *moved* (not + // cloned again) into `AuthIdentity.password`, which is + // `sspi::Secret` — a `#[derive(ZeroizeOnDrop)]` wrapper. + // The plaintext is therefore wiped when `identity` (and the + // credentials handle derived from it) is dropped; no + // un-zeroized copy is left behind. Expose the secret once, for + // this single `to_string()`, so no extra plaintext lingers + // here (`auth.password` zeroizes when this arm's `auth` + // drops). + let identity = AuthIdentity { + username, + password: auth.password.expose_secret().to_string().into(), + }; + + let mut creds = ntlm + .acquire_credentials_handle() + .with_credential_use(CredentialUse::Outbound) + .with_auth_data(&identity) + .execute(&mut ntlm)?; + + let spn = self.context.spn().to_string(); + + // First leg of the NTLM handshake: produce the NEGOTIATE token + // and ship it in the login packet as integrated security data. + let mut input = vec![SecurityBuffer::new(Vec::new(), BufferType::Token)]; + let mut output = vec![SecurityBuffer::new(Vec::new(), BufferType::Token)]; + + let mut builder = ntlm + .initialize_security_context() + .with_credentials_handle(&mut creds.credentials_handle) + .with_context_requirements( + ClientRequestFlags::CONFIDENTIALITY | ClientRequestFlags::ALLOCATE_MEMORY, + ) + .with_target_data_representation(DataRepresentation::Native) + .with_target_name(&spn) + .with_input(&mut input) + .with_output(&mut output); + + ntlm.initialize_security_context_impl(&mut builder)? + .resolve_to_result()?; + + login_message.integrated_security(Some(output[0].buffer.clone())); + + let id = self.context.next_packet_id(); + self.send(PacketHeader::login(id), login_message).await?; + self = self.post_login_encryption(encryption); + + // Second leg: consume the server's CHALLENGE token and reply + // with the AUTHENTICATE token. + let sspi_bytes = self.flush_sspi().await?; + + let mut input = vec![SecurityBuffer::new( + sspi_bytes.as_ref().to_vec(), + BufferType::Token, + )]; + let mut output = vec![SecurityBuffer::new(Vec::new(), BufferType::Token)]; + + let mut builder = ntlm + .initialize_security_context() + .with_credentials_handle(&mut creds.credentials_handle) + .with_context_requirements( + ClientRequestFlags::CONFIDENTIALITY | ClientRequestFlags::ALLOCATE_MEMORY, + ) + .with_target_data_representation(DataRepresentation::Native) + .with_target_name(&spn) + .with_input(&mut input) + .with_output(&mut output); + + ntlm.initialize_security_context_impl(&mut builder)? + .resolve_to_result()?; + + event!(Level::TRACE, authenticate_len = output[0].buffer.len()); + + let id = self.context.next_packet_id(); + self.send( + PacketHeader::login(id), + TokenSspi::new(output[0].buffer.clone()), + ) + .await?; + } + // winauth's NTLMv2 client is pure Rust, so this arm serves every + // platform; on Unix the sspi-rs arm above wins when both are on. + #[cfg(all(feature = "winauth", not(all(unix, feature = "sspi-rs"))))] AuthMethod::Windows(auth) => { let spn = self.context.spn().to_string(); let builder = winauth::NtlmV2ClientBuilder::new().target_spn(spn); - let mut client = builder.build(auth.domain, auth.user, auth.password); + // `auth.password` is a `SecretString`, but `winauth` + // 0.0.5's `build` takes the password by value as a plain + // `String` and neither zeroizes it nor exposes it afterwards, so + // it cannot be wiped once handed over. Expose it once for a + // single `to_string()` copy (no additional retained plaintext + // here) and accept a residual: the plaintext lives inside + // the `NtlmV2Client` until that value is dropped, un-zeroized. + // Closing this fully requires zeroize support upstream in + // `winauth`. + let mut client = builder.build( + auth.domain, + auth.user, + auth.password.expose_secret().to_string(), + ); login_message.integrated_security(client.next_bytes(None)?); @@ -408,7 +802,11 @@ impl Connection { let token = TokenSspi::new(sspi_response); self.send(header, token).await?; } - None => unreachable!(), + None => { + return Err(crate::Error::Protocol( + "NTLM handshake produced no response to the server challenge".into(), + )) + } } } AuthMethod::None => { @@ -417,17 +815,38 @@ impl Connection { self = self.post_login_encryption(encryption); } AuthMethod::SqlServer(auth) => { - login_message.user_name(auth.user()); - login_message.password(auth.password()); + let (user, mut password) = auth.into_credentials(); + + login_message.user_name(user); + // Expose the password only to hand it to the login message, + // which stores its own `SecretString` copy. + login_message.password(password.expose_secret()); + let payload = login_message.encode_to_boxed_slice()?; + // Wipe the local copy immediately; `login_message` was consumed by + // `encode_to_boxed_slice` and its `SecretString` password was + // zeroized on drop there. + password.zeroize(); let id = self.context.next_packet_id(); - self.send(PacketHeader::login(id), login_message).await?; + self.send_sensitive_login(PacketHeader::login(id), payload) + .await?; self = self.post_login_encryption(encryption); } AuthMethod::AADToken(token) => { - login_message.aad_token(token, prelogin.fed_auth_required, prelogin.nonce); + // Expose the token only to hand it to the login message; the + // `SecretString` here is wiped on drop at the end of this arm, + // and the login message stores its own `SecretString` copy. + login_message.aad_token( + token.expose_secret(), + prelogin.fed_auth_required, + prelogin.nonce, + ); + // Encode into a zeroizing buffer and use the sensitive-login + // path so the bearer token does not linger in freed heap memory. + let payload = login_message.encode_to_boxed_slice()?; let id = self.context.next_packet_id(); - self.send(PacketHeader::login(id), login_message).await?; + self.send_sensitive_login(PacketHeader::login(id), payload) + .await?; self = self.post_login_encryption(encryption); } } @@ -446,37 +865,55 @@ impl Connection { config: &Config, encryption: EncryptionLevel, ) -> crate::Result { - if encryption != EncryptionLevel::NotSupported { - event!(Level::INFO, "Performing a TLS handshake"); - - let Self { - transport, context, .. - } = self; - let mut stream = match transport.into_inner() { - MaybeTlsStream::Raw(tcp) => { - create_tls_stream(config, TlsPreloginWrapper::new(tcp)).await? - } - _ => unreachable!(), - }; + match encryption { + EncryptionLevel::NotSupported => { + event!( + Level::WARN, + "TLS encryption is not enabled. All traffic including the login credentials are not encrypted." + ); - stream.get_mut().handshake_complete(); - event!(Level::INFO, "TLS handshake successful"); + Ok(self) + } + // In strict mode the handshake already happened before the prelogin, + // so the transport is already a TLS stream. Nothing to do here. + EncryptionLevel::Strict => { + event!( + Level::TRACE, + "Already in a TLS stream (TDS 8.0 strict), skipping handshake." + ); - let transport = Framed::new(MaybeTlsStream::Tls(stream), PacketCodec); + Ok(self) + } + EncryptionLevel::Off | EncryptionLevel::On | EncryptionLevel::Required => { + event!(Level::DEBUG, "Performing a TLS handshake"); + + let Self { + transport, + context, + command_desync, + .. + } = self; + let mut stream = match transport.into_inner() { + MaybeTlsStream::Raw(tcp) => { + create_tls_stream(config, TlsPreloginWrapper::new(tcp)).await? + } + _ => unreachable!(), + }; - Ok(Self { - transport, - context, - flushed: false, - buf: BytesMut::new(), - }) - } else { - event!( - Level::WARN, - "TLS encryption is not enabled. All traffic including the login credentials are not encrypted." - ); + stream.get_mut().handshake_complete(); + event!(Level::DEBUG, "TLS handshake successful"); + + let transport = Framed::new(MaybeTlsStream::Tls(stream), PacketCodec); - Ok(self) + Ok(Self { + transport, + context, + flushed: false, + buf: BytesMut::new(), + poisoned: false, + command_desync, + }) + } } } @@ -486,7 +923,9 @@ impl Connection { feature = "native-tls", feature = "vendored-openssl" )))] - async fn tls_handshake(self, _: &Config, _: EncryptionLevel) -> crate::Result { + async fn tls_handshake(self, config: &Config, _: EncryptionLevel) -> crate::Result { + check_tls_backend_available(config.encryption)?; + event!( Level::WARN, "TLS encryption is not enabled. All traffic including the login credentials are not encrypted." @@ -500,6 +939,284 @@ impl Connection { } } +#[cfg(test)] +impl Connection { + /// Builds a `Connection` over a caller-supplied mock `AsyncRead + AsyncWrite` + /// stream for server-free unit tests, bypassing the prelogin/login handshake. + /// + /// The buffer starts empty and `flushed` starts `false`, so the first read + /// pulls a packet from the mock transport; a default [`Context`] is used and + /// the `poisoned` flag is caller-controlled. This mirrors the real field + /// initialization in [`Connection::connect`] and is `#[cfg(test)]` only, so + /// it never ships. Shared by the poison-guard tests and the `TokenStream` + /// state-machine tests, which feed it canned TDS-framed bytes. + pub(crate) fn test_over(io: S, poisoned: bool) -> Connection { + Connection { + transport: Framed::new(MaybeTlsStream::Raw(io), PacketCodec), + flushed: false, + context: Context::new(), + buf: BytesMut::new(), + poisoned, + command_desync: Arc::new(AtomicBool::new(false)), + } + } +} + +/// Returns an error when the user requested encryption but no TLS backend was +/// compiled in. Without this check, a `Required`/`On` encryption request would +/// silently fall back to an unencrypted connection. +#[cfg(not(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" +)))] +fn check_tls_backend_available(encryption: EncryptionLevel) -> crate::Result<()> { + if let EncryptionLevel::On | EncryptionLevel::Required | EncryptionLevel::Strict = encryption { + return Err(crate::Error::Tls( + "TLS encryption was requested but the crate was compiled without a TLS backend. \ + Enable one of the `native-tls`, `rustls` or `vendored-openssl` features." + .to_string(), + )); + } + + Ok(()) +} + +/// Runs `fut` to completion, but fails with a clear [`crate::Error::Io`] of +/// kind [`io::ErrorKind::TimedOut`] if it does not finish within `timeout`. +/// +/// A `None` timeout runs the future unbounded (the pre-0.13 behaviour). The +/// timer comes from `futures-timer`, whose `Delay` is runtime-agnostic, so this +/// works under any executor driving the generic [`Connection::connect`] path — +/// it does not assume tokio or smol. Used to bound the connection handshake so +/// a server that accepts the TCP connection and then stalls yields +/// a diagnosable error instead of an indefinite hang. +async fn with_optional_timeout( + timeout: Option, + fut: F, +) -> crate::Result +where + F: std::future::Future>, +{ + match timeout { + None => fut.await, + Some(timeout) => { + // `pin!` makes the future `Unpin` so `select` can take it by value; + // `Delay` is already `Unpin`. Whichever finishes first wins the + // race and the loser is dropped. + let fut = std::pin::pin!(fut); + match futures_util::future::select(fut, futures_timer::Delay::new(timeout)).await { + futures_util::future::Either::Left((res, _)) => res, + futures_util::future::Either::Right(((), _)) => Err(crate::Error::Io { + kind: io::ErrorKind::TimedOut, + message: format!( + "the connection handshake (prelogin, TLS negotiation and login) did not \ + complete within {timeout:?}; the server accepted the TCP connection but \ + did not finish the handshake in time (it may have stalled, or be \ + responding too slowly for the configured bound). Enable `tracing` at \ + DEBUG to see which stage was last reached, or change the bound with \ + `Config::handshake_timeout`." + ), + }), + } + } + } +} + +/// Frame a login message into one or more login packets, each no larger than +/// `packet_size` payload bytes, ready to write to the wire. +/// +/// This is the packetization core of [`Connection::send_sensitive_login`], +/// factored out so it can be unit-tested without a live server. Every packet is +/// an 8-byte header (see [`PacketHeader::encode`]) followed by up to +/// `packet_size` payload bytes, with the big-endian total length written into +/// header bytes `[2..4]` — matching [`Packet::encode`]. All but the final packet +/// carry `NormalMessage`; the final one carries `EndOfMessage`. +/// +/// The returned frames are `Zeroizing` so the sensitive login bytes they contain +/// are wiped on drop, keeping them off the shared (non-zeroizing) `BytesMut` +/// path used by the ordinary `send`. +fn frame_sensitive_login( + mut header: PacketHeader, + payload: &[u8], + packet_size: usize, +) -> crate::Result>>> { + let mut frames = Vec::new(); + let mut offset = 0; + + while offset < payload.len() { + let end = cmp::min(payload.len(), offset + packet_size); + + if end == payload.len() { + header.set_status(PacketStatus::EndOfMessage); + } else { + header.set_status(PacketStatus::NormalMessage); + } + + // Build the frame at its exact final capacity. `header.encode` is the + // only fallible (`?`) operation, and it runs BEFORE the sensitive + // payload bytes are copied in, so the secret never lives in this plain + // `Vec` across a `?`/`.await`. Once the header (`HEADER_BYTES`) and the + // chunk (`end - offset`) are written, `len == capacity`, so + // `into_boxed_slice()` cannot shrink-reallocate and leak an un-zeroized + // copy. Boxing into `Zeroizing` at push means each finished frame is + // wiped on drop and can no longer be grown. + let mut frame = Vec::with_capacity(HEADER_BYTES + end - offset); + header.encode(&mut frame)?; + frame.extend_from_slice(&payload[offset..end]); + + let size = (frame.len() as u16).to_be_bytes(); + frame[2] = size[0]; + frame[3] = size[1]; + + debug_assert_eq!( + frame.len(), + frame.capacity(), + "login frame buffer would shrink-reallocate when boxed, leaking a copy" + ); + frames.push(Zeroizing::new(frame.into_boxed_slice())); + offset = end; + } + + Ok(frames) +} + +#[cfg(test)] +mod sensitive_login_tests { + use super::frame_sensitive_login; + use crate::tds::codec::{Decode, Packet, PacketHeader, PacketStatus}; + use crate::tds::HEADER_BYTES; + use bytes::BytesMut; + + // An oversized login payload must be split into >=2 packets, each framed + // exactly like `Packet::encode`: an 8-byte header whose `[2..4]` bytes hold + // the big-endian total length, contiguous payload chunks covering the whole + // input, `NormalMessage` on every packet but the last and `EndOfMessage` on + // the final one. + #[test] + fn oversized_login_is_split_into_multiple_framed_packets() { + let packet_size = 16; // payload bytes per packet (excludes the 8-byte header) + let payload: Vec = (0..50u16).map(|i| i as u8).collect(); // 50 bytes -> ceil(50/16)=4 + let header = PacketHeader::login(3); + + let frames = frame_sensitive_login(header, &payload, packet_size).unwrap(); + + assert!( + frames.len() >= 2, + "expected the oversized login to split into >=2 packets, got {}", + frames.len() + ); + assert_eq!( + frames.len(), + 4, + "50 bytes / 16 per packet should be 4 packets" + ); + + let mut reassembled = Vec::new(); + for (i, frame) in frames.iter().enumerate() { + let is_last = i == frames.len() - 1; + + // Decode the header back off the wire bytes and check framing. + let mut buf = BytesMut::from(&frame[..]); + let decoded = PacketHeader::decode(&mut buf).unwrap(); + + // Length field ([2..4], big-endian) must equal the whole frame length. + assert_eq!( + decoded.length() as usize, + frame.len(), + "packet {i} length field must match the framed size" + ); + // And the raw bytes must match `Packet::encode`'s placement exactly. + let expected_len = (frame.len() as u16).to_be_bytes(); + assert_eq!([frame[2], frame[3]], expected_len); + + // Payload chunk size: every packet but the last is full. + let payload_len = frame.len() - HEADER_BYTES; + if is_last { + assert_eq!(decoded.status(), PacketStatus::EndOfMessage); + assert!(payload_len <= packet_size && payload_len > 0); + } else { + assert_eq!(decoded.status(), PacketStatus::NormalMessage); + assert_eq!(payload_len, packet_size); + } + + reassembled.extend_from_slice(&frame[HEADER_BYTES..]); + } + + // The concatenated payloads must reconstruct the original login bytes. + assert_eq!(reassembled, payload); + } + + // The frames are now `Box<[u8]>` built at exact capacity (len==capacity + // before `into_boxed_slice`, guarded by a debug_assert in the builder). The + // total bytes across all frames must therefore be exactly one header per + // frame plus the whole payload — no slack from over-reserved/realloc'd + // buffers — which this test enforces. + #[test] + fn framed_login_has_no_slack_bytes() { + let packet_size = 16; + let payload: Vec = (0..50u16).map(|i| i as u8).collect(); + let header = PacketHeader::login(5); + + let frames = frame_sensitive_login(header, &payload, packet_size).unwrap(); + + let total: usize = frames.iter().map(|f| f.len()).sum(); + assert_eq!( + total, + frames.len() * HEADER_BYTES + payload.len(), + "framed bytes must equal one header per frame plus the exact payload" + ); + } + + // A cross-check that a single frame produced by `frame_sensitive_login` + // matches byte-for-byte what `Packet::encode` produces for the same + // header + payload, when it fits in one packet. + #[test] + fn single_packet_frame_matches_packet_encode() { + use crate::tds::codec::Encode; + + let payload = vec![0xABu8; 10]; + let header = PacketHeader::login(7); + + let frames = frame_sensitive_login(header, &payload, 100).unwrap(); + assert_eq!(frames.len(), 1); + + let mut expected = BytesMut::new(); + let mut eom_header = header; + eom_header.set_status(PacketStatus::EndOfMessage); + Packet::new(eom_header, BytesMut::from(&payload[..])) + .encode(&mut expected) + .unwrap(); + + assert_eq!(&frames[0][..], &expected[..]); + } +} + +#[cfg(all( + test, + not(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + )) +))] +mod tests { + use super::check_tls_backend_available; + use crate::EncryptionLevel; + + #[test] + fn requested_encryption_without_tls_backend_errors() { + assert!(check_tls_backend_available(EncryptionLevel::Required).is_err()); + assert!(check_tls_backend_available(EncryptionLevel::On).is_err()); + } + + #[test] + fn no_encryption_without_tls_backend_is_ok() { + assert!(check_tls_backend_available(EncryptionLevel::Off).is_ok()); + assert!(check_tls_backend_available(EncryptionLevel::NotSupported).is_ok()); + } +} + impl Stream for Connection { type Item = crate::Result; @@ -576,3 +1293,469 @@ impl SqlReadBytes for Connection { &mut self.context } } + +// Server-free tests for the poisoned-connection guard shared by every write +// path, including the Attention/cancel path (see `cancel_request`). These build +// a `Connection` over an in-memory stream that is never actually driven: the +// guard must reject before any I/O happens. +#[cfg(test)] +mod poison_tests { + use super::*; + use crate::tds::codec::{BulkLoadRequest, TokenRow}; + use std::pin::Pin; + + /// A stream that swallows writes and reports clean EOF on read. If the + /// poisoned guard is ever bypassed, `cancel_request` would reach here and + /// fail with an EOF/IO error instead of the poison `Protocol` error, which + /// is exactly what the assertions below distinguish. + struct NullIo; + + impl AsyncRead for NullIo { + fn poll_read( + self: Pin<&mut Self>, + _: &mut task::Context<'_>, + _: &mut [u8], + ) -> Poll> { + Poll::Ready(Ok(0)) + } + } + + impl AsyncWrite for NullIo { + fn poll_write( + self: Pin<&mut Self>, + _: &mut task::Context<'_>, + buf: &[u8], + ) -> Poll> { + Poll::Ready(Ok(buf.len())) + } + + fn poll_flush(self: Pin<&mut Self>, _: &mut task::Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn poll_close(self: Pin<&mut Self>, _: &mut task::Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + } + + /// A stream whose writes always fail, simulating a wire error / interrupted + /// write. Reads report clean EOF. + struct FailingIo; + + impl AsyncRead for FailingIo { + fn poll_read( + self: Pin<&mut Self>, + _: &mut task::Context<'_>, + _: &mut [u8], + ) -> Poll> { + Poll::Ready(Ok(0)) + } + } + + impl AsyncWrite for FailingIo { + fn poll_write( + self: Pin<&mut Self>, + _: &mut task::Context<'_>, + _: &[u8], + ) -> Poll> { + Poll::Ready(Err(io::Error::other("wire down"))) + } + + fn poll_flush(self: Pin<&mut Self>, _: &mut task::Context<'_>) -> Poll> { + Poll::Ready(Err(io::Error::other("wire down"))) + } + + fn poll_close(self: Pin<&mut Self>, _: &mut task::Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + } + + fn connection_over(io: S, poisoned: bool) -> Connection + where + S: AsyncRead + AsyncWrite + Unpin + Send, + { + Connection::test_over(io, poisoned) + } + + fn poisoned_connection() -> Connection { + connection_over(NullIo, true) + } + + fn is_poison_error(err: &crate::Error) -> bool { + matches!(err, crate::Error::Protocol(msg) if msg.contains("inconsistent state")) + } + + #[test] + fn ensure_not_poisoned_reports_poison() { + let conn = poisoned_connection(); + assert!(is_poison_error(&conn.ensure_not_poisoned().unwrap_err())); + } + + #[test] + fn ensure_not_poisoned_ok_when_clean() { + let mut conn = poisoned_connection(); + conn.poisoned = false; + assert!(conn.ensure_not_poisoned().is_ok()); + } + + #[tokio::test] + async fn cancel_request_rejects_poisoned_connection() { + let mut conn = poisoned_connection(); + let err = conn + .cancel_request() + .await + .expect_err("cancel on a poisoned connection must fail"); + // Must be the poison guard specifically, not an incidental IO/EOF error + // from reaching the wire — that is what proves the guard is in effect. + assert!( + is_poison_error(&err), + "expected poison Protocol error, got: {err:?}" + ); + // The guard rejected before touching the wire, so the flag is untouched. + assert!(conn.poisoned); + } + + // --- Bulk-load write path (`BulkLoadRequest`) --------------------------- + // The bulk path writes multi-packet messages via `write_to_wire` directly, + // so it must get the same poison guarantee as `send`. Empty columns + empty + // rows keep these server-free: each `send` appends a single Row-token byte + // to the buffer and a tiny `packet_size` forces the write loop to fire. + + #[tokio::test] + async fn bulk_send_rejects_poisoned_connection() { + let mut conn = poisoned_connection(); + let mut bulk = BulkLoadRequest::new(&mut conn, Vec::new()).unwrap(); + let err = bulk + .send(TokenRow::new()) + .await + .expect_err("a bulk send on a poisoned connection must fail"); + // Must be the poison guard specifically — not an incidental IO error + // from reaching the (swallowing) wire — which proves the guard is in + // effect before any bytes are buffered or written. + assert!(is_poison_error(&err), "expected poison error, got: {err:?}"); + } + + #[tokio::test] + async fn bulk_write_clears_poison_on_success() { + let mut conn = connection_over(NullIo, false); + // One usable payload byte per packet so buffered Row-token bytes spill + // into the multi-packet write loop. + conn.context.set_packet_size(HEADER_BYTES as u32 + 1); + + { + let mut bulk = BulkLoadRequest::new(&mut conn, Vec::new()).unwrap(); + for _ in 0..8 { + bulk.send(TokenRow::new()) + .await + .expect("bulk send over a clean wire must succeed"); + } + } + + // After the bulk writes complete cleanly the connection must be usable. + assert!(!conn.poisoned); + assert!(conn.ensure_not_poisoned().is_ok()); + } + + #[tokio::test] + async fn bulk_write_failure_leaves_connection_poisoned() { + let mut conn = connection_over(FailingIo, false); + conn.context.set_packet_size(HEADER_BYTES as u32 + 1); + + { + let mut bulk = BulkLoadRequest::new(&mut conn, Vec::new()).unwrap(); + let mut result = Ok(()); + for _ in 0..8 { + result = bulk.send(TokenRow::new()).await; + if result.is_err() { + break; + } + } + assert!( + result.is_err(), + "a failing wire must surface an error from the bulk write" + ); + } + + // The interrupted write must leave the connection poisoned so the next + // op fails cleanly instead of appending onto a half-sent message. + assert!( + conn.poisoned, + "a failed bulk write must poison the connection" + ); + assert!(is_poison_error(&conn.ensure_not_poisoned().unwrap_err())); + } +} + +// Server-free tests for the handshake timeout helper (`with_optional_timeout`) +// and its wiring into `Connection::connect`. A future that never completes must +// surface a `TimedOut` error rather than hang, a `None` bound must run +// unbounded, and inner results (both `Ok` and `Err`) must pass through +// untouched when the future wins the race. +#[cfg(test)] +mod timeout_tests { + use super::*; + use std::future; + use std::pin::Pin; + use std::time::{Duration, Instant}; + + fn is_timed_out(err: &crate::Error) -> bool { + matches!( + err, + crate::Error::Io { + kind: io::ErrorKind::TimedOut, + .. + } + ) + } + + #[tokio::test] + async fn none_timeout_runs_unbounded_and_returns_inner_ok() { + let out: crate::Result = + with_optional_timeout(None, async { Ok::<_, crate::Error>(7u8) }).await; + assert_eq!(out.unwrap(), 7); + } + + #[tokio::test] + async fn future_completing_before_bound_returns_its_ok() { + let out: crate::Result<&str> = + with_optional_timeout(Some(Duration::from_secs(30)), async { + Ok::<_, crate::Error>("done") + }) + .await; + assert_eq!(out.unwrap(), "done"); + } + + #[tokio::test] + async fn future_completing_before_bound_propagates_its_err() { + // A fast inner error must be surfaced as-is, not masked as a timeout. + let out: crate::Result<()> = with_optional_timeout(Some(Duration::from_secs(30)), async { + Err::<(), _>(crate::Error::Protocol("boom".into())) + }) + .await; + let err = out.unwrap_err(); + assert!( + !is_timed_out(&err), + "fast inner error must not be a timeout" + ); + assert!(matches!(err, crate::Error::Protocol(msg) if msg == "boom")); + } + + #[tokio::test] + async fn stalled_future_times_out_with_diagnosable_error() { + let started = Instant::now(); + // A future that never resolves models a server that accepts the TCP + // connection and then stops responding mid-handshake. + let never = future::pending::>(); + let out = with_optional_timeout(Some(Duration::from_millis(50)), never).await; + + let err = out.expect_err("a stalled handshake must not hang; it must error"); + assert!( + is_timed_out(&err), + "expected a TimedOut error, got: {err:?}" + ); + + let msg = err.to_string(); + assert!( + msg.contains("did not complete") && msg.contains("handshake"), + "timeout error should explain the stalled handshake, got: {msg}" + ); + // The bound is honoured: returns promptly rather than blocking. + assert!( + started.elapsed() < Duration::from_secs(5), + "timeout should fire promptly, took {:?}", + started.elapsed() + ); + } + + // A stream that accepts (swallows) every write but never yields any bytes on + // read, modelling a server that completes the TCP connection and then goes + // silent during the prelogin/TLS handshake. + struct SilentServer; + + impl AsyncRead for SilentServer { + fn poll_read( + self: Pin<&mut Self>, + _: &mut task::Context<'_>, + _: &mut [u8], + ) -> Poll> { + // Never ready: the read half hangs forever. The handshake timeout, + // not this stream, is what must unblock the connect future. + Poll::Pending + } + } + + impl AsyncWrite for SilentServer { + fn poll_write( + self: Pin<&mut Self>, + _: &mut task::Context<'_>, + buf: &[u8], + ) -> Poll> { + Poll::Ready(Ok(buf.len())) + } + + fn poll_flush(self: Pin<&mut Self>, _: &mut task::Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn poll_close(self: Pin<&mut Self>, _: &mut task::Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + } + + #[tokio::test] + async fn connect_times_out_when_server_stalls_after_tcp() { + // Full `Connection::connect` wiring: the prelogin is written (swallowed) + // and then the response read hangs. With a short handshake timeout the + // connect must return a TimedOut error instead of blocking forever. + let mut config = Config::new(); + config.host("stalled.example"); + config.handshake_timeout(Some(Duration::from_millis(100))); + + let started = Instant::now(); + let err = Connection::connect(config, SilentServer) + .await + .expect_err("a server that stalls after TCP connect must not hang connect"); + + assert!( + is_timed_out(&err), + "expected a TimedOut error, got: {err:?}" + ); + assert!( + started.elapsed() < Duration::from_secs(5), + "connect should give up promptly, took {:?}", + started.elapsed() + ); + } + + // The same stalled-peer + // scenario, pinned to `EncryptionLevel::NotSupported` so the prelogin read + // (not a TLS handshake) is what stalls, and a companion test proving that + // with the bound disabled the connect keeps waiting. + #[tokio::test] + async fn connect_honors_handshake_timeout_without_tls() { + let mut config = Config::new(); + config.host("stalled.example"); + config.encryption(EncryptionLevel::NotSupported); + config.handshake_timeout(Some(Duration::from_millis(50))); + + let started = Instant::now(); + let err = Connection::connect(config, SilentServer) + .await + .expect_err("a stalled handshake must fail rather than hang"); + + assert!(is_timed_out(&err), "expected TimedOut, got: {err:?}"); + assert!( + started.elapsed() < Duration::from_secs(5), + "connect must fail promptly under the configured timeout" + ); + } + + #[tokio::test] + async fn connect_without_handshake_timeout_does_not_time_out() { + // With the bound disabled the connect future keeps waiting on the + // stalled read; the outer guard elapsing proves no internal timeout + // fired. + let mut config = Config::new(); + config.host("stalled.example"); + config.encryption(EncryptionLevel::NotSupported); + config.handshake_timeout(None); + assert_eq!(config.get_handshake_timeout(), None); + + let outcome = tokio::time::timeout( + Duration::from_millis(100), + Connection::connect(config, SilentServer), + ) + .await; + assert!( + outcome.is_err(), + "connect without a handshake timeout should keep waiting on the stalled stream" + ); + } + + // Serves a canned prelogin response once, then parks every subsequent read + // forever — a server that completes the prelogin exchange and then stalls + // while the client waits for the login acknowledgement. + struct AnswerPreloginThenSilent { + data: Vec, + pos: usize, + } + + impl AsyncRead for AnswerPreloginThenSilent { + fn poll_read( + mut self: Pin<&mut Self>, + _: &mut task::Context<'_>, + buf: &mut [u8], + ) -> Poll> { + if self.pos >= self.data.len() { + return Poll::Pending; + } + let remaining = &self.data[self.pos..]; + let n = remaining.len().min(buf.len()); + buf[..n].copy_from_slice(&remaining[..n]); + self.pos += n; + Poll::Ready(Ok(n)) + } + } + + impl AsyncWrite for AnswerPreloginThenSilent { + fn poll_write( + self: Pin<&mut Self>, + _: &mut task::Context<'_>, + buf: &[u8], + ) -> Poll> { + Poll::Ready(Ok(buf.len())) + } + + fn poll_flush(self: Pin<&mut Self>, _: &mut task::Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn poll_close(self: Pin<&mut Self>, _: &mut task::Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + } + + // Frame a payload as a single final (EndOfMessage) prelogin packet. + fn final_prelogin_packet() -> Vec { + let mut payload = BytesMut::new(); + PreloginMessage::new() + .encode(&mut payload) + .expect("prelogin encode"); + let packet = Packet::new(PacketHeader::pre_login(1), payload); + let mut buf = BytesMut::new(); + packet.encode(&mut buf).expect("packet encode"); + buf.to_vec() + } + + #[tokio::test] + async fn command_timeout_does_not_bound_the_connect_handshake() { + // `command_timeout` must NOT bound the + // login-ack drain that runs through the token stream during connect. The + // whole handshake is governed by `handshake_timeout` alone, so with + // `handshake_timeout(None)` a server that answers prelogin and then + // stalls at the login acknowledgement must be waited on indefinitely, + // even though a short `command_timeout` is configured. Before the fix + // the 50ms command timeout fired during `flush_done`, so connect would + // (wrongly) return within the 300ms guard. + let mut config = Config::new(); + config.host("stalled.example"); + config.encryption(EncryptionLevel::NotSupported); + config.handshake_timeout(None); + config.command_timeout(Some(Duration::from_millis(50))); + + let server = AnswerPreloginThenSilent { + data: final_prelogin_packet(), + pos: 0, + }; + + let outcome = tokio::time::timeout( + Duration::from_millis(300), + Connection::connect(config, server), + ) + .await; + assert!( + outcome.is_err(), + "command_timeout must not cut off the connect handshake; connect returned {outcome:?}" + ); + } +} diff --git a/src/client/tls.rs b/src/client/tls.rs index 7a22d4333..44a133cbc 100644 --- a/src/client/tls.rs +++ b/src/client/tls.rs @@ -4,18 +4,44 @@ feature = "vendored-openssl" ))] use super::tls_stream::TlsStream; +#[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" +))] use crate::tds::{ codec::{Decode, Encode, PacketHeader, PacketStatus, PacketType}, HEADER_BYTES, }; +#[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" +))] use bytes::BytesMut; use futures_util::io::{AsyncRead, AsyncWrite}; +#[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" +))] use futures_util::ready; +#[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" +))] +use std::cmp; use std::{ - cmp, io, + io, pin::Pin, task::{self, Poll}, }; +#[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" +))] use tracing::{event, Level}; /// A wrapper to handle either TLS or bare connections. @@ -114,6 +140,11 @@ impl AsyncWrite for MaybeTlsStream /// /// What it does is it interferes on handshake for TDS packet handling, /// and when complete, just passes the calls to the underlying connection. +#[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" +))] pub(crate) struct TlsPreloginWrapper { stream: Option, pending_handshake: bool, @@ -150,6 +181,11 @@ impl TlsPreloginWrapper { } } +#[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" +))] impl AsyncRead for TlsPreloginWrapper { fn poll_read( mut self: Pin<&mut Self>, @@ -179,13 +215,27 @@ impl AsyncRead for TlsPreloginWrapper< } let header = PacketHeader::decode(&mut BytesMut::from(&inner.header_buf[..])) - .map_err(|err| io::Error::new(io::ErrorKind::Other, err))?; - - // We only get pre-login packets in the handshake process. - assert_eq!(header.r#type(), PacketType::PreLogin); + .map_err(io::Error::other)?; + + // We only get pre-login packets in the handshake process. This runs + // before any certificate has been validated, so the bytes are fully + // untrusted: reject anything unexpected instead of panicking. + if header.r#type() != PacketType::PreLogin { + return Poll::Ready(Err(io::Error::new( + io::ErrorKind::InvalidData, + "expected a pre-login packet during the TLS handshake", + ))); + } - // And we know from this point on how much data we should expect - inner.read_remaining = header.length() as usize - HEADER_BYTES; + // And we know from this point on how much data we should expect. + inner.read_remaining = (header.length() as usize) + .checked_sub(HEADER_BYTES) + .ok_or_else(|| { + io::Error::new( + io::ErrorKind::InvalidData, + "pre-login packet length shorter than its header", + ) + })?; event!( Level::TRACE, @@ -212,6 +262,11 @@ impl AsyncRead for TlsPreloginWrapper< } } +#[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" +))] impl AsyncWrite for TlsPreloginWrapper { fn poll_write( mut self: Pin<&mut Self>, @@ -275,3 +330,179 @@ impl AsyncWrite for TlsPreloginWrapper Pin::new(&mut self.stream.as_mut().unwrap()).poll_close(cx) } } + +#[cfg(all( + test, + any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + ) +))] +mod tests { + use super::*; + use futures_util::io::{AsyncRead, AsyncWrite}; + use std::task::{Context, Poll, Waker}; + + // A minimal in-memory stream that yields a fixed slice of bytes to the + // reader. Enough to drive `TlsPreloginWrapper::poll_read` through the + // prelogin header-parsing branch without a real socket / TLS handshake. + struct MockStream { + data: std::io::Cursor>, + } + + impl AsyncRead for MockStream { + fn poll_read( + mut self: Pin<&mut Self>, + _cx: &mut task::Context<'_>, + buf: &mut [u8], + ) -> Poll> { + let remaining = &self.data.get_ref()[self.data.position() as usize..]; + let n = remaining.len().min(buf.len()); + buf[..n].copy_from_slice(&remaining[..n]); + let pos = self.data.position(); + self.data.set_position(pos + n as u64); + Poll::Ready(Ok(n)) + } + } + + impl AsyncWrite for MockStream { + fn poll_write( + self: Pin<&mut Self>, + _cx: &mut task::Context<'_>, + buf: &[u8], + ) -> Poll> { + Poll::Ready(Ok(buf.len())) + } + fn poll_flush(self: Pin<&mut Self>, _cx: &mut task::Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + fn poll_close(self: Pin<&mut Self>, _cx: &mut task::Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + } + + // An in-memory stream that captures every byte written to it, so the + // framing produced by `poll_write` + `poll_flush` can be inspected. Reads + // report clean EOF; they are not exercised by the write-side tests. + struct CapturingStream { + written: Vec, + } + + impl AsyncRead for CapturingStream { + fn poll_read( + self: Pin<&mut Self>, + _cx: &mut task::Context<'_>, + _buf: &mut [u8], + ) -> Poll> { + Poll::Ready(Ok(0)) + } + } + + impl AsyncWrite for CapturingStream { + fn poll_write( + mut self: Pin<&mut Self>, + _cx: &mut task::Context<'_>, + buf: &[u8], + ) -> Poll> { + self.written.extend_from_slice(buf); + Poll::Ready(Ok(buf.len())) + } + fn poll_flush(self: Pin<&mut Self>, _cx: &mut task::Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + fn poll_close(self: Pin<&mut Self>, _cx: &mut task::Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + } + + // During the handshake, `poll_write` buffers the payload behind an 8-byte + // header placeholder and `poll_flush` back-fills that header and writes the + // whole frame to the wire. The result must be exactly one TDS packet: + // PreLogin type, EndOfMessage status, a total length covering header + + // payload, followed by the verbatim payload. This locks the wrapper's + // outbound framing (previously only the read side was tested). + #[test] + fn poll_write_then_flush_frames_a_prelogin_packet() { + let stream = CapturingStream { + written: Vec::new(), + }; + let mut wrapper = TlsPreloginWrapper::new(stream); + let waker = Waker::noop(); + let mut cx = Context::from_waker(waker); + + // A payload longer than the header, so the `wr_buf.len() > HEADER_BYTES` + // flush guard fires. Distinctive bytes make an off-by-one obvious. + let payload: Vec = (0u8..20).map(|i| 0xA0 ^ i).collect(); + assert!(payload.len() > HEADER_BYTES); + + match Pin::new(&mut wrapper).poll_write(&mut cx, &payload) { + Poll::Ready(Ok(n)) => assert_eq!(n, payload.len()), + other => panic!("expected the payload to be accepted, got {other:?}"), + } + + match Pin::new(&mut wrapper).poll_flush(&mut cx) { + Poll::Ready(Ok(())) => {} + other => panic!("expected a clean flush, got {other:?}"), + } + + let written = &wrapper.stream.as_ref().unwrap().written; + + // The frame is header + payload and nothing else. + assert_eq!( + written.len(), + HEADER_BYTES + payload.len(), + "framed packet must be exactly header + payload bytes" + ); + + // The 8-byte header decodes to a PreLogin end-of-message packet whose + // declared length matches the whole frame. + let header = PacketHeader::decode(&mut BytesMut::from(&written[..HEADER_BYTES])).unwrap(); + assert_eq!(header.r#type(), PacketType::PreLogin); + assert_eq!(header.status(), PacketStatus::EndOfMessage); + assert_eq!( + header.length() as usize, + HEADER_BYTES + payload.len(), + "header length field must cover header + payload" + ); + + // The payload follows the header verbatim. + assert_eq!(&written[HEADER_BYTES..], &payload[..]); + } + + // Drive one `poll_read` over a wrapper fed the given 8-byte header. + fn poll_header(header: [u8; HEADER_BYTES]) -> Poll> { + let stream = MockStream { + data: std::io::Cursor::new(header.to_vec()), + }; + let mut wrapper = TlsPreloginWrapper::new(stream); + let waker = Waker::noop(); + let mut cx = Context::from_waker(waker); + let mut buf = [0u8; 32]; + Pin::new(&mut wrapper).poll_read(&mut cx, &mut buf) + } + + // A non-PreLogin packet type during the (unauthenticated) handshake must be + // rejected with `InvalidData`, not passed through / panicked on. + #[test] + fn non_prelogin_packet_type_is_rejected() { + // ty=SQLBatch(1), status=0, length=20 (BE), spid=0, id=0, window=0 + let header = [1u8, 0, 0, 20, 0, 0, 0, 0]; + match poll_header(header) { + Poll::Ready(Err(e)) => assert_eq!(e.kind(), io::ErrorKind::InvalidData), + other => panic!("expected InvalidData error, got {other:?}"), + } + } + + // A PreLogin header whose declared length is shorter than the 8-byte header + // must be rejected via the `checked_sub` underflow guard, not panic. + #[test] + fn prelogin_length_shorter_than_header_is_rejected() { + // ty=PreLogin(18), status=0, length=4 (< HEADER_BYTES), rest 0 + let header = [18u8, 0, 0, 4, 0, 0, 0, 0]; + match poll_header(header) { + Poll::Ready(Err(e)) => assert_eq!(e.kind(), io::ErrorKind::InvalidData), + other => panic!("expected InvalidData error, got {other:?}"), + } + } +} diff --git a/src/client/tls_stream.rs b/src/client/tls_stream.rs index 9eba1060f..3c379df57 100644 --- a/src/client/tls_stream.rs +++ b/src/client/tls_stream.rs @@ -1,6 +1,25 @@ use crate::Config; use futures_util::io::{AsyncRead, AsyncWrite}; +/// ALPN protocol name advertised for TDS 8.0 ("strict") encryption. +#[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" +))] +// Used by the native-tls and rustls backends to advertise TDS 8.0 strict; the +// opentls (vendored-openssl) backend cannot set ALPN, so this is unused there. +#[allow(dead_code)] +pub(crate) const TDS_ALPN_PROTOCOL_NAME: &str = "tds/8.0"; + +// Backend-agnostic CA loader shared by all three TLS backends. +#[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" +))] +mod certs; + #[cfg(feature = "native-tls")] mod native_tls_stream; diff --git a/src/client/tls_stream/certs.rs b/src/client/tls_stream/certs.rs new file mode 100644 index 000000000..1eb224bf0 --- /dev/null +++ b/src/client/tls_stream/certs.rs @@ -0,0 +1,438 @@ +//! Backend-agnostic CA-certificate loading. +//! +//! All three TLS backends (`rustls`, `native-tls`, `vendored-openssl`) trust +//! extra CAs through this one module, so the PEM/DER parsing and validation +//! cannot drift between them. It produces `Vec>`; each +//! backend then hands those DER bytes to its own certificate type. +//! +//! Two input shapes are supported, mirroring [`ExtraCa`]: +//! +//! - A **file** ([`ExtraCa::File`]): the format is chosen from the extension — +//! `.pem`/`.crt` parse as (possibly multi-certificate) PEM, `.der` as a single +//! DER certificate, anything else is rejected. +//! - An in-memory **bundle** ([`ExtraCa::Bundle`]): the format is sniffed from +//! the bytes — a `-----BEGIN` marker at the start of a line means PEM (every +//! block is parsed), otherwise the bytes are treated as a single DER +//! certificate. +//! +//! A leading UTF-8 byte-order mark (as written by some Windows editors / +//! PowerShell) is stripped before parsing, so a BOM-prefixed PEM file is still +//! recognised rather than mis-sniffed or rejected. +//! +//! Either shape yielding **zero** usable certificates is a hard error naming the +//! source, so a mistyped path, an empty/zero-byte input, or a bundle with no +//! certificate blocks can never silently degrade trust to "base roots only". + +use crate::client::config::ExtraCa; +use crate::error::IoErrorKind; +use rustls_pki_types::{pem::PemObject, CertificateDer}; +use std::{fs, path::Path}; + +/// Load every certificate from an [`ExtraCa`], guaranteeing at least one. +/// +/// The returned certificates are additive trust anchors to be layered on top of +/// the configured [`RootSource`](crate::client::config::RootSource). +pub(crate) fn trust_anchors(extra: &ExtraCa) -> crate::Result>> { + let certs = match extra { + ExtraCa::File(path) => certs_from_file(path)?, + ExtraCa::Bundle(bytes) => certs_from_bundle(bytes)?, + }; + + // Zero usable certificates is fatal and names the source — never a + // silent degrade to base-roots-only. + if certs.is_empty() { + return Err(crate::Error::Io { + kind: IoErrorKind::InvalidData, + message: format!( + "the CA source `{}` contained no usable certificates", + source_name(extra) + ), + }); + } + + Ok(certs) +} + +/// A human-readable name for an [`ExtraCa`] source, used in error messages so +/// every failure (read, parse, zero-cert, and an invalid certificate rejected by +/// the backend) names which configured CA is at fault. +pub(crate) fn source_name(extra: &ExtraCa) -> String { + match extra { + ExtraCa::File(path) => path.to_string_lossy().into_owned(), + ExtraCa::Bundle(bytes) => format!("in-memory CA bundle ({} bytes)", bytes.len()), + } +} + +/// The error for a loaded DER blob a backend rejects as an invalid certificate, +/// naming the offending source so it matches the read/parse/zero-cert errors. +/// Shared by all three backends (via their extra-CA loaders) so the "every +/// failure names which configured CA is at fault" invariant can't drift and is +/// testable in one place. +pub(crate) fn invalid_cert_error(extra: &ExtraCa, err: E) -> crate::Error { + crate::Error::Io { + kind: IoErrorKind::InvalidData, + message: format!( + "the CA source `{}` contains an invalid certificate: {err}", + source_name(extra) + ), + } +} + +/// Strip a leading UTF-8 byte-order mark, if present. PEM is ASCII text and the +/// tokenizer requires `-----BEGIN` at the start of a line, so a BOM glued to the +/// first marker (common from Windows editors / PowerShell `Out-File`) would +/// otherwise cause an otherwise-valid certificate to parse as zero blocks. +fn strip_bom(bytes: &[u8]) -> &[u8] { + bytes.strip_prefix(b"\xEF\xBB\xBF").unwrap_or(bytes) +} + +/// Parse a certificate file into DER certificates, dispatching on the extension. +/// The underlying I/O error is preserved in the message so callers can tell +/// missing-file / permission / parse failures apart, and the path is always +/// named (error-message parity across backends). +pub(crate) fn certs_from_file(path: &Path) -> crate::Result>> { + let buf = fs::read(path).map_err(|e| crate::Error::Io { + kind: IoErrorKind::InvalidData, + message: format!("Could not read certificate {}: {e}", path.to_string_lossy()), + })?; + + match path.extension() { + Some(ext) if ext.eq_ignore_ascii_case("pem") || ext.eq_ignore_ascii_case("crt") => { + CertificateDer::pem_slice_iter(strip_bom(&buf)) + .collect::, _>>() + .map_err(|e| crate::Error::Io { + kind: IoErrorKind::InvalidData, + message: format!( + "Failed to parse PEM certificate {}: {e}", + path.to_string_lossy() + ), + }) + } + // An empty `.der` file yields zero certificates (caught by the + // zero-usable-certs check in `trust_anchors`), rather than a bogus + // zero-length "certificate" that only fails later at the backend. + Some(ext) if ext.eq_ignore_ascii_case("der") && buf.is_empty() => Ok(vec![]), + Some(ext) if ext.eq_ignore_ascii_case("der") => Ok(vec![CertificateDer::from(buf)]), + Some(_) | None => Err(crate::Error::Io { + kind: IoErrorKind::InvalidInput, + message: format!( + "Certificate {} has an unsupported file-extension! Supported types are pem, crt and der.", + path.to_string_lossy() + ), + }), + } +} + +/// Parse in-memory certificate bytes, sniffing the format: a `-----BEGIN` +/// marker at the start of a line means PEM (every certificate block is parsed), +/// otherwise the bytes are treated as a single DER certificate. A zero-byte +/// input yields zero certificates (a hard, source-naming error in +/// [`trust_anchors`]) rather than a bogus empty "certificate". +pub(crate) fn certs_from_bundle(bytes: &[u8]) -> crate::Result>> { + let bytes = strip_bom(bytes); + if looks_like_pem(bytes) { + CertificateDer::pem_slice_iter(bytes) + .collect::, _>>() + .map_err(|e| crate::Error::Io { + kind: IoErrorKind::InvalidData, + message: format!("Failed to parse PEM CA bundle: {e}"), + }) + } else if bytes.is_empty() { + Ok(vec![]) + } else { + Ok(vec![CertificateDer::from(bytes.to_vec())]) + } +} + +/// A bundle is PEM if it contains a `-----BEGIN` marker at the start of a line +/// (start of the buffer, or immediately after a `\n`/`\r`). This mirrors the PEM +/// tokenizer's own line-start requirement, so a raw DER certificate that merely +/// *contains* those bytes somewhere in its ASN.1 body is not misclassified as +/// PEM (which would reject a perfectly valid DER cert). Callers strip a leading +/// BOM before calling, so a BOM-prefixed first line still matches. +fn looks_like_pem(bytes: &[u8]) -> bool { + const MARKER: &[u8] = b"-----BEGIN"; + if bytes.starts_with(MARKER) { + return true; + } + bytes + .windows(MARKER.len() + 1) + .any(|w| matches!(w[0], b'\n' | b'\r') && &w[1..] == MARKER) +} + +#[cfg(test)] +mod tests { + use super::*; + use std::path::PathBuf; + + #[test] + fn certs_from_file_reads_single_pem() { + let chain = certs_from_file(Path::new("docker/certs/server.crt")).unwrap(); + assert_eq!(chain.len(), 1); + } + + #[test] + fn certs_from_file_reads_multi_pem_chain() { + // Multi-cert files must load ALL certs, never truncate. + let chain = certs_from_file(Path::new("docker/certs/server-full.crt")).unwrap(); + assert!( + chain.len() >= 2, + "server-full.crt is a multi-certificate chain" + ); + // Guard against a bug that duplicates the first block N times while + // still inflating the count: the blocks must be genuinely distinct. + assert_ne!( + chain[0].as_ref(), + chain[1].as_ref(), + "the two blocks must be distinct certificates, not a repeated first block" + ); + } + + #[test] + fn certs_from_file_missing_file_preserves_io_error() { + let err = certs_from_file(Path::new("docker/certs/does-not-exist.crt")).unwrap_err(); + let msg = format!("{err:?}"); + assert!( + msg.contains("Could not read certificate") && msg.contains("does-not-exist.crt"), + "error should name the read failure and path, got: {msg}" + ); + } + + #[test] + fn certs_from_file_unsupported_extension_errors() { + let err = certs_from_file(Path::new("docker/certs/README.md")).unwrap_err(); + assert!(format!("{err:?}").contains("unsupported file-extension")); + } + + #[test] + fn certs_from_file_reads_der() { + // Derive a DER from the PEM CA and round-trip it through the der branch. + let der = certs_from_file(Path::new("docker/certs/customCA.crt")) + .unwrap() + .into_iter() + .next() + .unwrap() + .as_ref() + .to_vec(); + let mut path = std::env::temp_dir(); + path.push(format!( + "tiberius_certs_from_file_{}.der", + std::process::id() + )); + std::fs::write(&path, &der).unwrap(); + let chain = certs_from_file(&path); + std::fs::remove_file(&path).ok(); + assert_eq!(chain.unwrap().len(), 1); + } + + #[test] + fn certs_from_file_zero_cert_pem_is_empty() { + let mut path = std::env::temp_dir(); + path.push(format!("tiberius_certs_zero_{}.pem", std::process::id())); + std::fs::write(&path, b"# no certificates here\n").unwrap(); + let chain = certs_from_file(&path); + std::fs::remove_file(&path).ok(); + assert_eq!(chain.unwrap().len(), 0); + } + + #[test] + fn certs_from_bundle_sniffs_pem_and_reads_all() { + // A PEM bundle sniffed by `-----BEGIN`, all blocks parsed. + let bytes = std::fs::read("docker/certs/server-full.crt").unwrap(); + let certs = certs_from_bundle(&bytes).unwrap(); + assert!(certs.len() >= 2, "PEM bundle must yield every block"); + assert_ne!( + certs[0].as_ref(), + certs[1].as_ref(), + "the blocks must be distinct certificates, not a repeated first block" + ); + } + + #[test] + fn certs_from_bundle_sniffs_der() { + // No `-----BEGIN` marker => treated as a single DER certificate. + let der = certs_from_file(Path::new("docker/certs/customCA.crt")) + .unwrap() + .into_iter() + .next() + .unwrap() + .as_ref() + .to_vec(); + let certs = certs_from_bundle(&der).unwrap(); + assert_eq!(certs.len(), 1); + } + + #[test] + fn certs_from_bundle_garbage_pem_errors() { + let bytes = + b"-----BEGIN CERTIFICATE-----\nnot valid base64!!!\n-----END CERTIFICATE-----\n"; + let err = certs_from_bundle(bytes).unwrap_err(); + assert!(format!("{err:?}").contains("Failed to parse PEM CA bundle")); + } + + #[test] + fn trust_anchors_file_zero_certs_errors_naming_source() { + // A zero-cert file is a hard error naming the source path. + let mut path = std::env::temp_dir(); + path.push(format!( + "tiberius_trust_anchors_zero_{}.pem", + std::process::id() + )); + std::fs::write(&path, b"# no certificates here\n").unwrap(); + let err = trust_anchors(&ExtraCa::File(path.clone())); + std::fs::remove_file(&path).ok(); + let msg = format!("{:?}", err.unwrap_err()); + assert!( + msg.contains("contained no usable certificates") && msg.contains(".pem"), + "error should name the empty source, got: {msg}" + ); + } + + #[test] + fn trust_anchors_bundle_no_certificate_blocks_errors_naming_source() { + // A PEM bundle that contains no CERTIFICATE blocks (here only a + // PRIVATE KEY) yields zero certs and must be a hard error naming the + // in-memory source, never a silent degrade to base-roots-only. + let bundle = b"-----BEGIN PRIVATE KEY-----\nMIIB\n-----END PRIVATE KEY-----\n".to_vec(); + let err = trust_anchors(&ExtraCa::Bundle(bundle)).unwrap_err(); + let msg = format!("{err:?}"); + assert!( + msg.contains("contained no usable certificates") && msg.contains("in-memory CA bundle"), + "error should name the empty in-memory source, got: {msg}" + ); + } + + #[test] + fn trust_anchors_bundle_multi_pem() { + let bytes = std::fs::read("docker/certs/server-full.crt").unwrap(); + let certs = trust_anchors(&ExtraCa::Bundle(bytes)).unwrap(); + assert!(certs.len() >= 2); + } + + #[test] + fn trust_anchors_file_loads_all() { + let certs = trust_anchors(&ExtraCa::File(PathBuf::from( + "docker/certs/server-full.crt", + ))) + .unwrap(); + assert!(certs.len() >= 2); + } + + #[test] + fn empty_der_bundle_is_zero_certs_naming_source() { + // A zero-byte DER-sniffed bundle must be a hard error naming the + // source, not a bogus 1-element vec of an empty "certificate" that only + // fails later at the backend with an anonymous error. + assert_eq!(certs_from_bundle(&[]).unwrap().len(), 0); + let err = trust_anchors(&ExtraCa::Bundle(Vec::new())).unwrap_err(); + let msg = format!("{err:?}"); + assert!( + msg.contains("contained no usable certificates") && msg.contains("in-memory CA bundle"), + "empty bundle must be a source-naming zero-cert error, got: {msg}" + ); + } + + #[test] + fn empty_der_file_is_zero_certs_naming_source() { + // A zero-byte `.der` file, same as above but the file shape. + let mut path = std::env::temp_dir(); + path.push(format!("tiberius_empty_{}.der", std::process::id())); + std::fs::write(&path, b"").unwrap(); + assert_eq!(certs_from_file(&path).unwrap().len(), 0); + let err = trust_anchors(&ExtraCa::File(path.clone())); + std::fs::remove_file(&path).ok(); + let msg = format!("{:?}", err.unwrap_err()); + assert!( + msg.contains("contained no usable certificates") && msg.contains(".der"), + "empty .der must be a source-naming zero-cert error, got: {msg}" + ); + } + + #[test] + fn bom_prefixed_pem_still_parses() { + // A UTF-8 BOM glued to the first `-----BEGIN` marker must not defeat + // parsing. Use a MULTI-cert fixture so both assertions are load-bearing: + // if `strip_bom` were dropped, the bundle would sniff as a single DER + // blob (len 1, or an error) and the file path would parse zero blocks — + // either way `>= 2` fails. + let pem = std::fs::read("docker/certs/server-full.crt").unwrap(); + let mut with_bom = vec![0xEF, 0xBB, 0xBF]; + with_bom.extend_from_slice(&pem); + + // Bundle path (content sniff). + assert!( + certs_from_bundle(&with_bom).unwrap().len() >= 2, + "BOM-prefixed PEM bundle must parse every block" + ); + + // File path (extension dispatch). + let mut path = std::env::temp_dir(); + path.push(format!("tiberius_bom_{}.pem", std::process::id())); + std::fs::write(&path, &with_bom).unwrap(); + let res = certs_from_file(&path); + std::fs::remove_file(&path).ok(); + assert!( + res.unwrap().len() >= 2, + "BOM-prefixed PEM file must parse every block" + ); + } + + #[test] + fn der_with_embedded_begin_marker_is_treated_as_der() { + // A raw DER cert whose body merely *contains* the `-----BEGIN` byte + // sequence (not at a line start) must be sniffed as DER, not PEM — so a + // valid DER cert isn't misclassified and rejected. + let der = certs_from_file(Path::new("docker/certs/customCA.crt")) + .unwrap() + .into_iter() + .next() + .unwrap() + .as_ref() + .to_vec(); + let mut crafted = der.clone(); + // Splice the marker into the middle of the DER body (not line-start). + let mid = crafted.len() / 2; + crafted.splice(mid..mid, b"-----BEGIN CERTIFICATE-----".iter().copied()); + assert!( + !looks_like_pem(&crafted), + "mid-body marker must not sniff PEM" + ); + // The genuine DER (no marker) round-trips as a single DER certificate. + assert!(!looks_like_pem(&der)); + assert_eq!(certs_from_bundle(&der).unwrap().len(), 1); + } + + #[test] + fn invalid_cert_error_names_source_for_every_shape() { + // The shared error used by all three backends when a DER blob is + // rejected must name the source (path / in-memory bundle) and carry the + // backend's underlying error text. + let file = invalid_cert_error(&ExtraCa::File(PathBuf::from("/tmp/ca.der")), "bad tag"); + let m = format!("{file:?}"); + assert!( + m.contains("/tmp/ca.der") + && m.contains("contains an invalid certificate") + && m.contains("bad tag"), + "file source must be named with the underlying error, got: {m}" + ); + let bundle = invalid_cert_error(&ExtraCa::Bundle(vec![0u8; 7]), "bad der"); + let m2 = format!("{bundle:?}"); + assert!( + m2.contains("in-memory CA bundle (7 bytes)") + && m2.contains("contains an invalid certificate"), + "bundle source must be named, got: {m2}" + ); + } + + #[test] + fn looks_like_pem_requires_line_start_marker() { + assert!(looks_like_pem(b"-----BEGIN CERTIFICATE-----\n")); + assert!(looks_like_pem( + b"# a comment\n-----BEGIN CERTIFICATE-----\n" + )); + assert!(looks_like_pem(b"lead\r-----BEGIN CERTIFICATE-----\r\n")); + assert!(!looks_like_pem(b"prefix -----BEGIN CERTIFICATE-----")); + assert!(!looks_like_pem(b"")); + assert!(!looks_like_pem(b"short")); + } +} diff --git a/src/client/tls_stream/native_tls_stream.rs b/src/client/tls_stream/native_tls_stream.rs index cf5591d80..11e4a303a 100644 --- a/src/client/tls_stream/native_tls_stream.rs +++ b/src/client/tls_stream/native_tls_stream.rs @@ -1,60 +1,251 @@ +use super::certs; use crate::{ - client::{config::Config, TrustConfig}, + client::config::{ClientCertSource, ClientCertificate, Config, ExtraCa}, error::{Error, IoErrorKind}, }; pub(crate) use async_native_tls::TlsStream; -use async_native_tls::{Certificate, TlsConnector}; +use async_native_tls::{Certificate, Identity, TlsConnector}; use futures_util::io::{AsyncRead, AsyncWrite}; +use secrecy::ExposeSecret; use std::fs; use tracing::{event, Level}; +/// Loads a client identity from the configured source for `native-tls`. +fn load_identity(cert: &ClientCertificate) -> crate::Result { + match &cert.source { + ClientCertSource::CertAndKey { cert, key } => { + // Accept only the extensions valid for each role: a certificate is + // `.pem`/`.crt`, a private key is `.pem`/`.key`. Sharing a single + // set previously let a `.key` file pass as the certificate (and a + // `.crt` as the key). + // + // `.pem` is legitimately ambiguous — it is valid for BOTH the cert + // and the key role — so it is accepted for either slot. A swapped + // `.pem`/`.pem` pair therefore passes this extension check and is + // only caught later by `Identity::from_pkcs8`, which fails to parse + // mismatched content. Rejecting `.pem` for either role would break + // valid usage, so it is intentionally left accepted for both. + let has_ext = |p: &std::path::Path, exts: &[&str]| { + matches!( + p.extension().and_then(|e| e.to_str()), + Some(ext) if exts.iter().any(|e| ext.eq_ignore_ascii_case(e)) + ) + }; + + if !has_ext(cert, &["pem", "crt"]) || !has_ext(key, &["pem", "key"]) { + return Err(Error::Tls( + "The native-tls backend requires PEM certificate and key files; \ + for a DER-bundled identity use `Config::client_certificate_pkcs12`." + .to_string(), + )); + } + + let cert_buf = fs::read(cert).map_err(|e| Error::Io { + kind: IoErrorKind::InvalidData, + message: format!( + "Could not read client certificate {}: {e}", + cert.to_string_lossy() + ), + })?; + let key_buf = fs::read(key).map_err(|e| Error::Io { + kind: IoErrorKind::InvalidData, + message: format!( + "Could not read client private key {}: {e}", + key.to_string_lossy() + ), + })?; + + Ok(Identity::from_pkcs8(&cert_buf, &key_buf)?) + } + ClientCertSource::Pkcs12 { path, password } => { + let buf = fs::read(path).map_err(|e| Error::Io { + kind: IoErrorKind::InvalidData, + message: format!( + "Could not read PKCS#12 identity {}: {e}", + path.to_string_lossy() + ), + })?; + // Expose the PKCS#12 password only for the decryption call itself. + Ok(Identity::from_pkcs12(&buf, password.expose_secret())?) + } + } +} + pub(crate) async fn create_tls_stream( config: &Config, stream: S, ) -> crate::Result> { let mut builder = TlsConnector::new(); - match &config.trust { - TrustConfig::CaCertificateLocation(path) => { - if let Ok(buf) = fs::read(path) { - let cert = match path.extension() { - Some(ext) - if ext.to_ascii_lowercase() == "pem" - || ext.to_ascii_lowercase() == "crt" => - { - Some(Certificate::from_pem(&buf)?) - } - Some(ext) if ext.to_ascii_lowercase() == "der" => { - Some(Certificate::from_der(&buf)?) - } - Some(_) | None => return Err(Error::Io { - kind: IoErrorKind::InvalidInput, - message: "Provided CA certificate with unsupported file-extension! Supported types are pem, crt and der.".to_string()}), - }; - if let Some(c) = cert { - builder = builder.add_root_certificate(c); - } - } else { - return Err(Error::Io { - kind: IoErrorKind::InvalidData, - message: "Could not read provided CA certificate!".to_string(), - }); - } + if matches!(config.encryption, crate::EncryptionLevel::Strict) { + builder = builder.request_alpns(&[super::TDS_ALPN_PROTOCOL_NAME]); + } + + if let Some(cert) = config.get_client_certificate() { + event!( + Level::DEBUG, + "Presenting a client certificate for mutual TLS." + ); + builder = builder.identity(load_identity(cert)?); + } + + if config.trust.bypass { + event!( + Level::WARN, + "Trusting the server certificate without validation." + ); + + builder = builder.danger_accept_invalid_certs(true); + builder = builder.danger_accept_invalid_hostnames(true); + builder = builder.use_sni(false); + } else { + // The base trust anchors are the platform trust store, which native-tls + // consults automatically (there is no `WebpkiRoots` source here — that + // variant only exists with the rustls backend). Layer every accumulated + // extra CA on top, loading ALL certificates from each multi-cert file or + // bundle via the shared, backend-agnostic loader. + for cert in load_extra_cas(&config.trust.extra_cas)? { + builder = builder.add_root_certificate(cert); } - TrustConfig::TrustAll => { - event!( - Level::WARN, - "Trusting the server certificate without validation." + } + + Ok(builder + .connect(config.get_hostname_in_certificate(), stream) + .await?) +} + +/// Load every accumulated extra CA into native-tls `Certificate`s, loading ALL +/// certificates from each multi-cert file/bundle and naming the offending source +/// if any DER blob is not a valid certificate. Factored out of the async connect +/// path so the load-and-name behaviour is unit-testable without a live server. +fn load_extra_cas(extras: &[ExtraCa]) -> crate::Result> { + let mut out = Vec::new(); + for extra in extras { + for cert in certs::trust_anchors(extra)? { + out.push( + Certificate::from_der(cert.as_ref()) + .map_err(|e| certs::invalid_cert_error(extra, e))?, ); + } + } + Ok(out) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::client::config::{ClientCertSource, ExtraCa}; + use std::path::PathBuf; - builder = builder.danger_accept_invalid_certs(true); - builder = builder.danger_accept_invalid_hostnames(true); - builder = builder.use_sni(false); + // Cross-backend loading: a multi-cert CA file must yield every + // certificate, and each must convert into a native-tls `Certificate`. This + // exercises the same shared loader + per-cert `Certificate::from_der` + // conversion the connect path uses, without needing a live server. + #[test] + fn multi_cert_ca_file_loads_all_certs() { + let ders = certs::trust_anchors(&ExtraCa::File(PathBuf::from( + "docker/certs/server-full.crt", + ))) + .expect("multi-cert CA file loads"); + assert!(ders.len() >= 2, "must not truncate a multi-cert file"); + for der in &ders { + Certificate::from_der(der.as_ref()).expect("each DER cert converts for native-tls"); } - TrustConfig::Default => { - event!(Level::INFO, "Using default trust configuration."); + } + + #[test] + fn multi_cert_ca_bundle_loads_all_certs() { + let bytes = std::fs::read("docker/certs/server-full.crt").unwrap(); + let ders = certs::trust_anchors(&ExtraCa::Bundle(bytes)).expect("multi-cert bundle loads"); + assert!(ders.len() >= 2, "must not truncate a multi-cert bundle"); + for der in &ders { + Certificate::from_der(der.as_ref()).expect("each DER cert converts for native-tls"); } } - Ok(builder.connect(config.get_host(), stream).await?) + #[test] + fn empty_ca_file_is_hard_error_naming_source() { + // A zero-cert CA never silently degrades to platform-only trust. + let mut path = std::env::temp_dir(); + path.push(format!("tiberius_nt_zero_{}.pem", std::process::id())); + std::fs::write(&path, b"# no certs\n").unwrap(); + let err = certs::trust_anchors(&ExtraCa::File(path.clone())); + std::fs::remove_file(&path).ok(); + assert!(format!("{:?}", err.unwrap_err()).contains("contained no usable certificates")); + } + + #[test] + fn invalid_der_extra_ca_names_source() { + // A DER blob that is not a valid certificate must fail the extra-CA load + // with a message naming the offending source (exercises the native-tls + // `Certificate::from_der` + source-naming path, not just the loader). + // `Certificate` isn't `Debug`, so match instead of `unwrap_err`. + let err = match load_extra_cas(&[ExtraCa::Bundle(vec![0x30, 0x03, 0x02, 0x01, 0x7f])]) { + Ok(_) => panic!("an invalid DER extra CA must be a hard error"), + Err(e) => e, + }; + let msg = format!("{err:?}"); + assert!( + msg.contains("in-memory CA bundle") && msg.contains("contains an invalid certificate"), + "invalid DER must name the source, got: {msg}" + ); + } + + #[test] + fn multi_cert_extra_cas_load_all_via_helper() { + // The helper loads every cert from a multi-cert file through the + // real connect-path code, converting each to a native-tls Certificate. + let certs = load_extra_cas(&[ExtraCa::File(PathBuf::from("docker/certs/server-full.crt"))]) + .expect("multi-cert file loads"); + assert!(certs.len() >= 2, "must not truncate a multi-cert file"); + } + + fn identity(cert: &str, key: &str) -> crate::Result { + load_identity(&ClientCertificate { + source: ClientCertSource::CertAndKey { + cert: PathBuf::from(cert), + key: PathBuf::from(key), + }, + }) + } + + // A `.key` file must not be accepted in the certificate role, nor a `.crt` + // in the key role. The extension check runs before any file read, so these + // paths need not exist. + #[test] + fn key_extension_rejected_as_certificate() { + assert!(matches!( + identity("client.key", "client.key"), + Err(Error::Tls(_)) + )); + } + + #[test] + fn crt_extension_rejected_as_key() { + assert!(matches!( + identity("client.crt", "client.crt"), + Err(Error::Tls(_)) + )); + } + + // The valid role combinations pass the extension check and fail only later, + // when the nonexistent files are read (an `Io` error, not a `Tls` one). + #[test] + fn valid_extensions_pass_extension_check() { + for (cert, key) in [ + ("client.crt", "client.key"), + ("client.pem", "client.pem"), + ("client.crt", "client.pem"), + ("client.pem", "client.key"), + ] { + match identity(cert, key) { + Err(Error::Tls(_)) => { + panic!("valid extensions {cert}/{key} were wrongly rejected by the check") + } + Err(Error::Io { .. }) => {} + Err(e) => panic!("expected an Io error from the missing file, got {e:?}"), + Ok(_) => panic!("nonexistent files should not yield an identity"), + } + } + } } diff --git a/src/client/tls_stream/opentls_tls_stream.rs b/src/client/tls_stream/opentls_tls_stream.rs index 76c8c9df0..f4516206b 100644 --- a/src/client/tls_stream/opentls_tls_stream.rs +++ b/src/client/tls_stream/opentls_tls_stream.rs @@ -1,101 +1,209 @@ +use super::certs; use crate::{ - client::{config::Config, TrustConfig}, + client::config::{ClientCertSource, ClientCertificate, Config, ExtraCa}, error::{Error, IoErrorKind}, }; use futures_util::io::{AsyncRead, AsyncWrite}; pub(crate) use opentls::async_io::{TlsConnector, TlsStream}; -use opentls::Certificate; +use opentls::{Certificate, Identity}; +use secrecy::ExposeSecret; use std::fs; use tracing::{event, Level}; +/// Loads a client identity from the configured source for the `opentls` +/// (vendored OpenSSL) backend. +/// +/// `opentls` only exposes `Identity::from_pkcs12`, so only a PKCS#12 / PFX +/// bundle (supplied via [`Config::client_certificate_pkcs12`]) is supported; +/// separate PEM/DER certificate and key files cannot be loaded by this backend. +fn load_identity(cert: &ClientCertificate) -> crate::Result { + match &cert.source { + ClientCertSource::Pkcs12 { path, password } => { + let buf = fs::read(path).map_err(|e| Error::Io { + kind: IoErrorKind::InvalidData, + message: format!( + "Could not read PKCS#12 identity {}: {e}", + path.to_string_lossy() + ), + })?; + // Expose the PKCS#12 password only for the decryption call itself. + Ok(Identity::from_pkcs12(&buf, password.expose_secret())?) + } + ClientCertSource::CertAndKey { .. } => Err(Error::Tls( + "The vendored-openssl (opentls) backend does not support separate \ + certificate/key files for client authentication; supply a PKCS#12 \ + bundle via `Config::client_certificate_pkcs12` instead." + .to_string(), + )), + } +} + pub(crate) async fn create_tls_stream( config: &Config, stream: S, ) -> crate::Result> { let mut builder = TlsConnector::new(); - match &config.trust { - TrustConfig::CaCertificateLocation(path) => { - if let Ok(buf) = fs::read(path) { - let cert = match path.extension() { - Some(ext) - if ext.to_ascii_lowercase() == "pem" - || ext.to_ascii_lowercase() == "crt" => - { - Some(Certificate::from_pem(&buf)?) - } - Some(ext) if ext.to_ascii_lowercase() == "der" => { - Some(Certificate::from_der(&buf)?) + if matches!(config.encryption, crate::EncryptionLevel::Strict) { + event!( + Level::WARN, + "OpenTLS does not support ALPN, so the TDS 8.0 ALPN protocol will not be requested. SQL Server will assume TDS 8.0." + ); + } + + if let Some(cert) = config.get_client_certificate() { + event!( + Level::DEBUG, + "Presenting a client certificate for mutual TLS." + ); + builder = builder.identity(load_identity(cert)?); + } + + if config.trust.bypass { + event!( + Level::WARN, + "Trusting the server certificate without validation." + ); + + builder = builder.danger_accept_invalid_certs(true); + builder = builder.danger_accept_invalid_hostnames(true); + builder = builder.use_sni(false); + } else { + // The base trust anchors are the platform trust store (there is no + // `WebpkiRoots` source here — that variant only exists with the rustls + // backend). On Unix opentls finds it automatically via openssl-probe. + // On Windows that probe finds nothing — its trust anchors live in + // registry-backed certificate stores — so load the ROOT store into the + // connector explicitly. The current-user view is a composite that + // includes the local-machine store, so certs installed via certlm.msc + // (machine) or GP-pushed (e.g. Zscaler) are picked up automatically. + // Individual certificates OpenSSL cannot parse are skipped, matching + // what rustls-native-certs does. + // + // NOTE: open_current_user("ROOT") is correct for user-session + // processes (e.g. a desktop app). A Windows service would need + // open_local_machine("ROOT") instead, since services run in session 0 + // and the current-user store may be empty there. + #[cfg(windows)] + match schannel::cert_store::CertStore::open_current_user("ROOT") { + Ok(store) => { + for windows_cert in store.certs() { + match Certificate::from_der(windows_cert.to_der()) { + Ok(root_cert) => { + builder = builder.add_root_certificate(root_cert); } - Some(_) | None => return Err(Error::Io { - kind: IoErrorKind::InvalidInput, - message: "Provided CA certificate with unsupported file-extension! Supported types are pem, crt and der.".to_string()}), - }; - if let Some(c) = cert { - builder = builder.add_root_certificate(c); + Err(e) => { + event!( + Level::WARN, + "Skipping an unparseable certificate from the Windows ROOT store: {}", + e + ); + } + } } - } else { - return Err(Error::Io { - kind: IoErrorKind::InvalidData, - message: "Could not read provided CA certificate!".to_string(), - }); } + Err(e) => { + event!( + Level::WARN, + "Could not open the Windows ROOT certificate store; certificate validation will have no trusted roots: {}", + e + ); + } + } + + // Layer every accumulated extra CA on top, loading ALL certificates + // from each multi-cert file or bundle via the shared, backend-agnostic + // loader. + for cert in load_extra_cas(&config.trust.extra_cas)? { + builder = builder.add_root_certificate(cert); } - TrustConfig::TrustAll => { - event!( - Level::WARN, - "Trusting the server certificate without validation." + } + + Ok(builder + .connect(config.get_hostname_in_certificate(), stream) + .await?) +} + +/// Load every accumulated extra CA into opentls `Certificate`s, loading ALL +/// certificates from each multi-cert file/bundle and naming the offending source +/// if any DER blob is not a valid certificate. Factored out of the async connect +/// path so the load-and-name behaviour is unit-testable without a live server. +fn load_extra_cas(extras: &[ExtraCa]) -> crate::Result> { + let mut out = Vec::new(); + for extra in extras { + for cert in certs::trust_anchors(extra)? { + out.push( + Certificate::from_der(cert.as_ref()) + .map_err(|e| certs::invalid_cert_error(extra, e))?, ); + } + } + Ok(out) +} - builder = builder.danger_accept_invalid_certs(true); - builder = builder.danger_accept_invalid_hostnames(true); - builder = builder.use_sni(false); +#[cfg(test)] +mod tests { + use super::*; + use std::path::PathBuf; + + // Cross-backend loading: a multi-cert CA file/bundle must yield + // every certificate, and each must convert into an opentls `Certificate`. + #[test] + fn multi_cert_ca_file_loads_all_certs() { + let ders = certs::trust_anchors(&ExtraCa::File(PathBuf::from( + "docker/certs/server-full.crt", + ))) + .expect("multi-cert CA file loads"); + assert!(ders.len() >= 2, "must not truncate a multi-cert file"); + for der in &ders { + Certificate::from_der(der.as_ref()).expect("each DER cert converts for opentls"); } - TrustConfig::Default => { - event!(Level::INFO, "Using default trust configuration."); + } - // The vendored OpenSSL discovers root certificates by probing - // Unix filesystem paths (openssl-probe), which finds nothing on - // Windows — its trust anchors live in registry-backed - // certificate stores. Load the ROOT store into the connector so - // certificate validation can succeed; the current-user view is - // a composite that includes the local-machine store, so certs - // installed via certlm.msc (machine) or GP-pushed (e.g. Zscaler) - // are picked up automatically. Individual certificates OpenSSL - // cannot parse are skipped, matching what rustls-native-certs does. - // - // NOTE: open_current_user("ROOT") is correct for user-session - // processes (e.g. the KeeperDB desktop app). A Windows service - // would need open_local_machine("ROOT") instead, since services - // run in session 0 and the current-user store may be empty there. - #[cfg(windows)] - match schannel::cert_store::CertStore::open_current_user("ROOT") { - Ok(store) => { - for windows_cert in store.certs() { - match Certificate::from_der(windows_cert.to_der()) { - Ok(root_cert) => { - builder = builder.add_root_certificate(root_cert); - } - Err(e) => { - event!( - Level::WARN, - "Skipping an unparseable certificate from the Windows ROOT store: {}", - e - ); - } - } - } - } - Err(e) => { - event!( - Level::WARN, - "Could not open the Windows ROOT certificate store; certificate validation will have no trusted roots: {}", - e - ); - } - } + #[test] + fn multi_cert_ca_bundle_loads_all_certs() { + let bytes = std::fs::read("docker/certs/server-full.crt").unwrap(); + let ders = certs::trust_anchors(&ExtraCa::Bundle(bytes)).expect("multi-cert bundle loads"); + assert!(ders.len() >= 2, "must not truncate a multi-cert bundle"); + for der in &ders { + Certificate::from_der(der.as_ref()).expect("each DER cert converts for opentls"); } } - Ok(builder.connect(config.get_host(), stream).await?) + #[test] + fn empty_ca_file_is_hard_error_naming_source() { + // A zero-cert CA never silently degrades to platform-only trust. + let mut path = std::env::temp_dir(); + path.push(format!("tiberius_ot_zero_{}.pem", std::process::id())); + std::fs::write(&path, b"# no certs\n").unwrap(); + let err = certs::trust_anchors(&ExtraCa::File(path.clone())); + std::fs::remove_file(&path).ok(); + assert!(format!("{:?}", err.unwrap_err()).contains("contained no usable certificates")); + } + + #[test] + fn invalid_der_extra_ca_names_source() { + // A DER blob that is not a valid certificate must fail the extra-CA load + // with a message naming the offending source (exercises the opentls + // `Certificate::from_der` + source-naming path, not just the loader). + // `Certificate` isn't `Debug`, so match instead of `unwrap_err`. + let err = match load_extra_cas(&[ExtraCa::Bundle(vec![0x30, 0x03, 0x02, 0x01, 0x7f])]) { + Ok(_) => panic!("an invalid DER extra CA must be a hard error"), + Err(e) => e, + }; + let msg = format!("{err:?}"); + assert!( + msg.contains("in-memory CA bundle") && msg.contains("contains an invalid certificate"), + "invalid DER must name the source, got: {msg}" + ); + } + + #[test] + fn multi_cert_extra_cas_load_all_via_helper() { + // The helper loads every cert from a multi-cert file through the + // real connect-path code, converting each to an opentls Certificate. + let certs = load_extra_cas(&[ExtraCa::File(PathBuf::from("docker/certs/server-full.crt"))]) + .expect("multi-cert file loads"); + assert!(certs.len() >= 2, "must not truncate a multi-cert file"); + } } diff --git a/src/client/tls_stream/rustls_tls_stream.rs b/src/client/tls_stream/rustls_tls_stream.rs index e417583a6..38d1a52a5 100644 --- a/src/client/tls_stream/rustls_tls_stream.rs +++ b/src/client/tls_stream/rustls_tls_stream.rs @@ -1,24 +1,30 @@ +use super::certs; use crate::{ - client::{config::Config, TrustConfig}, + client::{ + config::{ClientCertSource, ClientCertificate, Config, RootSource}, + TrustConfig, + }, error::IoErrorKind, Error, }; use futures_util::io::{AsyncRead, AsyncWrite}; use std::{ fs, io, + path::Path, pin::Pin, sync::Arc, task::{Context, Poll}, - time::SystemTime, }; use tokio_rustls::{ rustls::{ client::{ - HandshakeSignatureValid, ServerCertVerified, ServerCertVerifier, - WantsTransparencyPolicyOrClientCert, + danger::{HandshakeSignatureValid, ServerCertVerified, ServerCertVerifier}, + WantsClientCert, }, - Certificate, ClientConfig, ConfigBuilder, DigitallySignedStruct, Error as RustlsError, - RootCertStore, ServerName, WantsVerifier, + crypto::{aws_lc_rs, CryptoProvider}, + pki_types::{pem::PemObject, CertificateDer, PrivateKeyDer, ServerName, UnixTime}, + ClientConfig, ConfigBuilder, DigitallySignedStruct, Error as RustlsError, RootCertStore, + SignatureScheme, }, TlsConnector, }; @@ -35,17 +41,17 @@ pub(crate) struct TlsStream( Compat>>, ); +#[derive(Debug)] struct NoCertVerifier; impl ServerCertVerifier for NoCertVerifier { fn verify_server_cert( &self, - _end_entity: &Certificate, - _intermediates: &[Certificate], - _server_name: &ServerName, - _scts: &mut dyn Iterator, + _end_entity: &CertificateDer<'_>, + _intermediates: &[CertificateDer<'_>], + _server_name: &ServerName<'_>, _ocsp_response: &[u8], - _now: SystemTime, + _now: UnixTime, ) -> Result { Ok(ServerCertVerified::assertion()) } @@ -53,92 +59,171 @@ impl ServerCertVerifier for NoCertVerifier { fn verify_tls12_signature( &self, _message: &[u8], - _cert: &Certificate, + _cert: &CertificateDer<'_>, + _dss: &DigitallySignedStruct, + ) -> Result { + Ok(HandshakeSignatureValid::assertion()) + } + + fn verify_tls13_signature( + &self, + _message: &[u8], + _cert: &CertificateDer<'_>, _dss: &DigitallySignedStruct, ) -> Result { Ok(HandshakeSignatureValid::assertion()) } + + fn supported_verify_schemes(&self) -> Vec { + // Advertised only; the trust bypass stubs verification to always succeed. + vec![ + SignatureScheme::RSA_PKCS1_SHA256, + SignatureScheme::RSA_PKCS1_SHA384, + SignatureScheme::RSA_PKCS1_SHA512, + SignatureScheme::ECDSA_NISTP256_SHA256, + SignatureScheme::ECDSA_NISTP384_SHA384, + SignatureScheme::ECDSA_NISTP521_SHA512, + SignatureScheme::RSA_PSS_SHA256, + SignatureScheme::RSA_PSS_SHA384, + SignatureScheme::RSA_PSS_SHA512, + SignatureScheme::ED25519, + SignatureScheme::ED448, + ] + } +} + +fn get_server_name(config: &Config) -> crate::Result> { + match ( + ServerName::try_from(config.get_hostname_in_certificate()), + config.trust.bypass, + ) { + (Ok(sn), _) => Ok(sn.to_owned()), + // Under the trust bypass the certificate (and thus its name) is not + // validated, so the SNI value is irrelevant; use a syntactically-valid + // placeholder when the configured hostname can't be parsed as a + // `ServerName`. The literal is a valid DNS name, so + // `try_from(...).unwrap()` cannot panic. + (Err(_), true) => Ok(ServerName::try_from("placeholder.domain.com").unwrap()), + (Err(e), false) => Err(crate::Error::Tls(e.to_string())), + } } -fn get_server_name(config: &Config) -> crate::Result { - match (ServerName::try_from(config.get_host()), &config.trust) { - (Ok(sn), _) => Ok(sn), - (Err(_), TrustConfig::TrustAll) => { - Ok(ServerName::try_from("placeholder.domain.com").unwrap()) +/// Translate a TLS-handshake failure into an actionable [`crate::Error`]. +/// +/// `tokio-rustls` surfaces handshake failures as [`io::Error`]s that wrap the +/// underlying [`RustlsError`]. When that inner error is a *certificate +/// validation* failure (e.g. `UnsupportedCertVersion` from an older or +/// non-conformant SQL Server certificate), the raw +/// message ("invalid peer certificate ... UnsupportedCertVersion") is opaque +/// and never names the fix. This maps such failures to an [`Error::Tls`] that +/// spells out the remedies, while leaving every other I/O error untouched (it +/// still flows through `From` as an [`Error::Io`]). +/// +/// The message is only added for genuine certificate-validation failures under +/// a *validating* trust config; the trust bypass (`trust_cert`) already +/// short-circuits certificate checks in [`NoCertVerifier`], so a cert error +/// cannot originate there. +fn map_handshake_error(err: io::Error, trust: &TrustConfig) -> crate::Error { + // Only certificate-*validation* failures get the extra guidance. Under the + // trust bypass the verifier never rejects a cert, so any error there is + // genuinely transport-level and should pass through unchanged. + if !trust.bypass { + if let Some(RustlsError::InvalidCertificate(cert_err)) = + err.get_ref().and_then(|e| e.downcast_ref::()) + { + // `UnsupportedCertVersion` is the specific symptom. In + // this rustls version it arrives from webpki wrapped in + // `CertificateError::Other(..)` rather than as a named variant, so + // detect it from the rendered message. The same remedies apply to + // any validation rejection of a legacy/self-signed server cert. + let rendered = cert_err.to_string(); + let hint = if rendered.contains("UnsupportedCertVersion") { + " The SQL Server presented a certificate rustls considers too \ + old (an outdated X.509 version), which frequently happens with \ + the self-signed certificate SQL Server auto-generates." + } else { + "" + }; + + return Error::Tls(format!( + "the server's certificate was rejected during the TLS handshake: {cert_err}.{hint} \ + To connect anyway you can: (1) trust a specific CA certificate with \ + `Config::trust_cert_ca(path)`, or an in-memory CA bundle with \ + `Config::trust_cert_ca_bundle(bytes)`; (2) if the OS trust store is unusable, \ + base trust on the bundled Mozilla roots with `Config::trust_webpki_roots()` \ + (requires the `rustls-webpki-roots` feature); (3) skip certificate validation \ + entirely with `Config::trust_cert()` (accepts any certificate and also disables \ + hostname verification — only safe on a trusted network); or (4) build tiberius \ + with the `native-tls` backend, which is more lenient toward legacy certificates. \ + Note: SQL Server always performs a TLS handshake during login even when \ + `Encrypt=false`, so this can occur regardless of the encryption setting." + )); } - (Err(e), _) => Err(crate::Error::Tls(e.to_string())), } + + Error::from(err) } impl TlsStream { pub(super) async fn new(config: &Config, stream: S) -> crate::Result { - event!(Level::INFO, "Performing a TLS handshake"); - - let builder = ClientConfig::builder().with_safe_defaults(); - - let client_config = match &config.trust { - TrustConfig::CaCertificateLocation(path) => { - if let Ok(buf) = fs::read(path) { - let cert = match path.extension() { - Some(ext) - if ext.to_ascii_lowercase() == "pem" - || ext.to_ascii_lowercase() == "crt" => - { - let pem_cert = rustls_pemfile::certs(&mut buf.as_slice())?; - if pem_cert.len() != 1 { - return Err(crate::Error::Io { - kind: IoErrorKind::InvalidInput, - message: format!("Certificate file {} contain 0 or more than 1 certs", path.to_string_lossy()), - }); - } - - Certificate(pem_cert.into_iter().next().unwrap()) - } - Some(ext) if ext.to_ascii_lowercase() == "der" => { - Certificate(buf) - } - Some(_) | None => return Err(crate::Error::Io { - kind: IoErrorKind::InvalidInput, - message: "Provided CA certificate with unsupported file-extension! Supported types are pem, crt and der.".to_string(), - }), - }; - let mut cert_store = RootCertStore::empty(); - cert_store.add(&cert)?; - builder - .with_root_certificates(cert_store) - .with_no_client_auth() - } else { - return Err(Error::Io { - kind: IoErrorKind::InvalidData, - message: "Could not read provided CA certificate!".to_string(), - }); - } - } - TrustConfig::TrustAll => { + event!(Level::DEBUG, "Performing a TLS handshake"); + + let provider = resolve_crypto_provider(CryptoProvider::get_default().cloned()); + + // Negotiate the best available protocol version (TLS 1.2 or 1.3), the + // same policy as upstream's previous `with_safe_defaults()`. + let builder = ClientConfig::builder_with_provider(provider) + .with_safe_default_protocol_versions() + .map_err(|e| crate::Error::Tls(e.to_string()))?; + + // First select the server-certificate verification strategy, yielding a + // builder that still awaits the client-authentication decision. + let cc_builder: ConfigBuilder = if config.trust.bypass { + event!( + Level::WARN, + "Trusting the server certificate without validation." + ); + builder + .dangerous() + .with_custom_certificate_verifier(Arc::new(NoCertVerifier)) + } else { + // Build the root store from the configured base source plus any + // accumulated extra CAs, then validate against it. + let store = build_trust_store(&config.trust)?; + builder.with_root_certificates(store) + }; + + // Present a client certificate (mutual TLS / TDS 8.0 + // `ENCRYPT_CLIENT_CERT`) if one was configured, otherwise finalize + // without client authentication. + let mut client_config = match config.get_client_certificate() { + Some(cert) => { event!( - Level::WARN, - "Trusting the server certificate without validation." + Level::DEBUG, + "Presenting a client certificate for mutual TLS." ); - let mut config = builder - .with_root_certificates(RootCertStore::empty()) - .with_no_client_auth(); - config - .dangerous() - .set_certificate_verifier(Arc::new(NoCertVerifier {})); - // config.enable_sni = false; - config - } - TrustConfig::Default => { - event!(Level::INFO, "Using default trust configuration."); - builder.with_native_roots().with_no_client_auth() + let (chain, key) = load_client_auth(cert)?; + cc_builder + .with_client_auth_cert(chain, key) + .map_err(|e| crate::Error::Tls(e.to_string()))? } + None => cc_builder.with_no_client_auth(), }; + // TDS 8.0 "strict" mode advertises the `tds/8.0` ALPN protocol so the + // server knows to speak TDS directly over the TLS stream. + if matches!(config.encryption, crate::EncryptionLevel::Strict) { + client_config + .alpn_protocols + .push(super::TDS_ALPN_PROTOCOL_NAME.as_bytes().to_vec()); + } + let connector = TlsConnector::from(Arc::new(client_config)); let tls_stream = connector .connect(get_server_name(config)?, stream.compat()) - .await?; + .await + .map_err(|e| map_handshake_error(e, &config.trust))?; Ok(TlsStream(tls_stream.compat())) } @@ -180,36 +265,743 @@ impl AsyncWrite for TlsStream { } } -trait ConfigBuilderExt { - fn with_native_roots(self) -> ConfigBuilder; +/// Resolve the rustls `CryptoProvider`: honour a process-installed default +/// (`CryptoProvider::install_default`) if present, otherwise fall back to +/// aws-lc-rs. +fn resolve_crypto_provider(installed: Option>) -> Arc { + match installed { + Some(provider) => { + event!( + Level::DEBUG, + "Using process-installed rustls CryptoProvider" + ); + provider + } + None => { + event!( + Level::DEBUG, + "No process-installed CryptoProvider; using the aws-lc-rs default" + ); + Arc::new(aws_lc_rs::default_provider()) + } + } } -impl ConfigBuilderExt for ConfigBuilder { - fn with_native_roots(self) -> ConfigBuilder { - let mut roots = RootCertStore::empty(); - let mut valid_count = 0; - let mut invalid_count = 0; +/// Load the OS trust store's certificates into `roots`, returning +/// `(added, had_load_errors)`. +fn load_native_roots_into(roots: &mut RootCertStore) -> (usize, bool) { + let native = rustls_native_certs::load_native_certs(); + let had_load_errors = !native.errors.is_empty(); + if had_load_errors { + event!( + Level::DEBUG, + "loading platform certificates reported errors: {:?}", + native.errors + ); + } + let mut added = 0; + for cert in native.certs { + match roots.add(cert) { + Ok(_) => added += 1, + Err(err) => { + event!( + Level::DEBUG, + "skipping invalid platform certificate: {:?}", + err + ) + } + } + } + (added, had_load_errors) +} - for cert in rustls_native_certs::load_native_certs().expect("could not load platform certs") - { - let cert = Certificate(cert.0); - match roots.add(&cert) { - Ok(_) => valid_count += 1, - Err(err) => { - tracing::event!(Level::TRACE, "invalid cert der {:?}", cert.0); - tracing::event!(Level::DEBUG, "certificate parsing failed: {:?}", err); - invalid_count += 1 - } +/// Build the root-certificate store for a validating [`TrustConfig`]: the base +/// trust anchors (the OS store, or the bundled Mozilla roots) **plus** every +/// accumulated extra CA. +/// +/// Fail-closed invariant: when there are no extra CAs, an empty base +/// store is fatal — `Native` with an unusable OS store fails to connect rather +/// than trusting nothing. When extra CAs *are* supplied they provide trust on +/// their own, so an empty/best-effort base is tolerated (matching the additive +/// `trust_cert_ca` contract). Every extra source must yield at least one usable +/// certificate (enforced in [`certs::trust_anchors`]). +fn build_trust_store(trust: &TrustConfig) -> crate::Result { + let mut store = RootCertStore::empty(); + + // Base trust anchors from the configured source. + let (base_added, base_had_errors) = match &trust.source { + RootSource::Native => { + let (added, had_errors) = load_native_roots_into(&mut store); + event!(Level::TRACE, "native trust store added {added} certs"); + (added, had_errors) + } + #[cfg(feature = "rustls-webpki-roots")] + RootSource::WebpkiRoots => { + let before = store.roots.len(); + store + .roots + .extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned()); + let added = store.roots.len() - before; + event!(Level::TRACE, "webpki-roots added {added} certs"); + (added, false) + } + }; + + // Layer the accumulated extra CAs on top. A DER blob that parses as bytes + // but is not a valid certificate is rejected by `store.add`; name the + // offending source so the failure is as actionable as the read/parse/zero + // errors from `certs::trust_anchors`. + for extra in &trust.extra_cas { + for cert in certs::trust_anchors(extra)? { + store + .add(cert) + .map_err(|e| certs::invalid_cert_error(extra, e))?; + } + } + + // Fail closed only when nothing else provides trust. + ensure_base_or_extra(base_added, base_had_errors, !trust.extra_cas.is_empty())?; + + Ok(store) +} + +/// Fail-closed decision for a validating trust store: with no extra +/// CAs, an empty base store is fatal — `Native` with an unusable OS store must +/// fail to connect rather than silently trust nothing. When extra CAs *are* +/// present they provide trust on their own, so a best-effort/empty base is +/// tolerated. Factored out so the invariant is unit-testable without needing to +/// simulate an empty OS trust store. +fn ensure_base_or_extra( + base_added: usize, + base_had_errors: bool, + has_extras: bool, +) -> crate::Result<()> { + if base_added == 0 && !has_extras { + return Err(crate::Error::Io { + kind: IoErrorKind::NotFound, + message: if base_had_errors { + "could not load platform certificates".to_string() + } else { + "no usable CA certificates found in the platform trust store".to_string() + }, + }); + } + Ok(()) +} + +/// Read a private-key file, dispatching on the extension: `pem`/`key` parse as +/// PEM (PKCS#8, PKCS#1 or SEC1), `der` as DER (PKCS#8). Mirrors `certs::certs_from_file` +/// and preserves the underlying I/O error in the message. +fn read_private_key(path: &Path) -> crate::Result> { + let buf = fs::read(path).map_err(|e| crate::Error::Io { + kind: IoErrorKind::InvalidData, + message: format!("Could not read private key {}: {e}", path.to_string_lossy()), + })?; + + match path.extension() { + Some(ext) if ext.eq_ignore_ascii_case("pem") || ext.eq_ignore_ascii_case("key") => { + PrivateKeyDer::from_pem_slice(&buf).map_err(|e| crate::Error::Io { + kind: IoErrorKind::InvalidData, + message: format!("Failed to parse PEM private key {}: {e}", path.to_string_lossy()), + }) + } + Some(ext) if ext.eq_ignore_ascii_case("der") => { + PrivateKeyDer::try_from(buf).map_err(|e| crate::Error::Io { + kind: IoErrorKind::InvalidData, + message: format!("Failed to parse DER private key {}: {e}", path.to_string_lossy()), + }) + } + Some(_) | None => Err(crate::Error::Io { + kind: IoErrorKind::InvalidInput, + message: format!( + "Private key {} has an unsupported file-extension! Supported types are pem, key and der.", + path.to_string_lossy() + ), + }), + } +} + +/// Loads a client certificate chain and private key from the configured source +/// for use with rustls' `with_client_auth_cert`. +fn load_client_auth( + cert: &ClientCertificate, +) -> crate::Result<(Vec>, PrivateKeyDer<'static>)> { + match &cert.source { + ClientCertSource::CertAndKey { cert, key } => { + // Certificate chain: PEM (possibly multiple) or a single DER cert. + let chain = certs::certs_from_file(cert)?; + + if chain.is_empty() { + return Err(crate::Error::Io { + kind: IoErrorKind::InvalidInput, + message: format!( + "Client certificate file {} contains no certificates", + cert.to_string_lossy() + ), + }); + } + + let key = read_private_key(key)?; + + Ok((chain, key)) + } + #[cfg(any(feature = "native-tls", feature = "vendored-openssl"))] + ClientCertSource::Pkcs12 { .. } => Err(crate::Error::Tls( + "The rustls backend does not support PKCS#12 client certificates; \ + supply separate PEM/DER certificate and key files via \ + `Config::client_certificate` instead." + .to_string(), + )), + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::client::config::{ClientCertSource, ClientCertificate, Config}; + use std::path::PathBuf; + use tokio_rustls::rustls::CertificateError; + + use crate::client::config::ExtraCa; + #[cfg(feature = "rustls-webpki-roots")] + use crate::client::config::RootSource; + + fn make_config(host: Option<&str>, cert_host: Option<&str>, trust: TrustConfig) -> Config { + let mut c = Config::new(); + c.trust = trust; + if let Some(h) = host { + c.host = Some(h.to_string()); + } + if let Some(hc) = cert_host { + c.hostname_in_certificate = Some(hc.to_string()); + } + c + } + + /// The validating default trust config (OS store, no extras, no bypass). + fn default_trust() -> TrustConfig { + TrustConfig::default() + } + + /// The trust-everything bypass (`trust_cert`). + fn bypass_trust() -> TrustConfig { + TrustConfig { + bypass: true, + ..TrustConfig::default() + } + } + + /// Native source plus a single extra CA file. + fn ca_file_trust(path: &str) -> TrustConfig { + TrustConfig { + extra_cas: vec![ExtraCa::File(PathBuf::from(path))], + ..TrustConfig::default() + } + } + + #[test] + fn resolve_crypto_provider_honours_installed() { + let installed = Arc::new(aws_lc_rs::default_provider()); + let got = resolve_crypto_provider(Some(installed.clone())); + assert!( + Arc::ptr_eq(&got, &installed), + "an installed CryptoProvider must be used as-is" + ); + } + + #[test] + fn resolve_crypto_provider_falls_back_to_aws_lc_rs() { + let got = resolve_crypto_provider(None); + assert!( + !got.cipher_suites.is_empty(), + "the aws-lc-rs fallback must provide cipher suites" + ); + } + + #[test] + fn server_name_valid_host_is_ok() { + let c = make_config(Some("localhost"), None, default_trust()); + assert!(get_server_name(&c).is_ok()); + } + + #[test] + fn server_name_invalid_host_bypass_uses_placeholder() { + let c = make_config(None, Some("inv al id"), bypass_trust()); + let got = get_server_name(&c).expect("the bypass must fall back to the placeholder SNI"); + assert!(format!("{got:?}").contains("placeholder.domain.com")); + } + + #[test] + fn server_name_invalid_host_validating_errors() { + let c = make_config(None, Some("inv al id"), default_trust()); + assert!(get_server_name(&c).is_err()); + } + + #[test] + fn build_trust_store_augments_system_roots_with_custom_ca() { + // Independently measure this machine's native root count using the same + // loader `build_trust_store` uses, so the comparison holds on any host. + let mut native_only = RootCertStore::empty(); + let (native, _) = load_native_roots_into(&mut native_only); + + let store = build_trust_store(&ca_file_trust("docker/certs/customCA.crt")).unwrap(); + + // Custom CA must augment, not replace, the system roots (native + 1). + assert_eq!( + store.len(), + native + 1, + "custom CA must augment the system trust store, not replace it" + ); + } + + #[test] + fn build_trust_store_accepts_multi_cert_ca_file() { + // The pre-0.13 single-certificate restriction is relaxed: every cert in + // a multi-cert CA file is trusted. + let mut native_only = RootCertStore::empty(); + let (native, _) = load_native_roots_into(&mut native_only); + + let n_certs = certs::certs_from_file(Path::new("docker/certs/server-full.crt")) + .unwrap() + .len(); + assert!(n_certs >= 2); + + let store = build_trust_store(&ca_file_trust("docker/certs/server-full.crt")).unwrap(); + assert_eq!( + store.len(), + native + n_certs, + "all certs in a multi-cert CA file must be added" + ); + } + + #[test] + fn build_trust_store_accumulates_multiple_extra_cas() { + // Accumulate semantics at the store level: two extra CAs => both added. + let mut native_only = RootCertStore::empty(); + let (native, _) = load_native_roots_into(&mut native_only); + + let trust = TrustConfig { + extra_cas: vec![ + ExtraCa::File(PathBuf::from("docker/certs/customCA.crt")), + ExtraCa::Bundle(std::fs::read("docker/certs/customCA.crt").unwrap()), + ], + ..TrustConfig::default() + }; + let store = build_trust_store(&trust).unwrap(); + assert_eq!( + store.len(), + native + 2, + "both accumulated extra CAs must be trusted" + ); + } + + #[test] + fn build_trust_store_propagates_zero_cert_extra_ca_error() { + // An extra CA that yields zero usable certs is fatal, naming it. + let mut path = std::env::temp_dir(); + path.push(format!("tiberius_bts_zero_{}.pem", std::process::id())); + std::fs::write(&path, b"# no certificates here\n").unwrap(); + let trust = TrustConfig { + extra_cas: vec![ExtraCa::File(path.clone())], + ..TrustConfig::default() + }; + let err = build_trust_store(&trust); + std::fs::remove_file(&path).ok(); + let msg = format!("{:?}", err.unwrap_err()); + assert!( + msg.contains("contained no usable certificates"), + "error should name the empty source, got: {msg}" + ); + } + + #[test] + fn build_trust_store_names_source_on_invalid_der_cert() { + // A DER extra-CA whose bytes are not a valid certificate must fail with a + // message naming the offending source, not an anonymous backend error. + let trust = TrustConfig { + extra_cas: vec![ExtraCa::Bundle(vec![0x30, 0x03, 0x02, 0x01, 0x7f])], + ..TrustConfig::default() + }; + let msg = format!("{:?}", build_trust_store(&trust).unwrap_err()); + assert!( + msg.contains("in-memory CA bundle") && msg.contains("invalid certificate"), + "invalid DER must name the source, got: {msg}" + ); + } + + #[test] + fn ensure_base_or_extra_enforces_fail_closed_rule7() { + // Fail-closed empty-base behaviour, tested directly (an empty OS store can't be simulated in-process): + // no base + no extras => fail closed, with the message distinguishing a + // load error from a genuinely empty store. + let empty_store = ensure_base_or_extra(0, false, false).unwrap_err(); + assert!( + format!("{empty_store:?}").contains("no usable CA certificates found"), + "empty store must fail closed, got: {empty_store:?}" + ); + let load_err = ensure_base_or_extra(0, true, false).unwrap_err(); + assert!( + format!("{load_err:?}").contains("could not load platform certificates"), + "load failure must be reported distinctly, got: {load_err:?}" + ); + // Extras present => tolerate an empty/best-effort base (additive contract). + assert!(ensure_base_or_extra(0, true, true).is_ok()); + assert!(ensure_base_or_extra(0, false, true).is_ok()); + // A non-empty base is always fine. + assert!(ensure_base_or_extra(5, false, false).is_ok()); + } + + #[cfg(feature = "rustls-webpki-roots")] + #[test] + fn build_trust_store_uses_bundled_webpki_roots() { + // The webpki source seeds a non-empty store from the compiled-in Mozilla + // snapshot, independent of the OS trust store. + let trust = TrustConfig { + source: RootSource::WebpkiRoots, + ..TrustConfig::default() + }; + let store = build_trust_store(&trust).unwrap(); + assert_eq!( + store.len(), + webpki_roots::TLS_SERVER_ROOTS.len(), + "the store must be seeded from the bundled Mozilla roots" + ); + assert!(!store.is_empty()); + } + + #[cfg(feature = "rustls-webpki-roots")] + #[test] + fn build_trust_store_webpki_roots_plus_extra_ca() { + // Extras still layer on top of the webpki source. + let base = webpki_roots::TLS_SERVER_ROOTS.len(); + let trust = TrustConfig { + source: RootSource::WebpkiRoots, + extra_cas: vec![ExtraCa::File(PathBuf::from("docker/certs/customCA.crt"))], + ..TrustConfig::default() + }; + let store = build_trust_store(&trust).unwrap(); + assert_eq!(store.len(), base + 1); + } + + #[test] + fn load_client_auth_reads_pem_cert_and_key() { + let cert = ClientCertificate { + source: ClientCertSource::CertAndKey { + cert: PathBuf::from("docker/certs/server.crt"), + key: PathBuf::from("docker/certs/server.key"), + }, + }; + let (chain, _key) = load_client_auth(&cert).expect("valid PEM cert + key"); + assert_eq!(chain.len(), 1); + } + + #[test] + fn load_client_auth_missing_cert_errors() { + let cert = ClientCertificate { + source: ClientCertSource::CertAndKey { + cert: PathBuf::from("docker/certs/does-not-exist.crt"), + key: PathBuf::from("docker/certs/server.key"), + }, + }; + assert!(load_client_auth(&cert).is_err()); + } + + #[test] + fn read_private_key_missing_file_preserves_io_error() { + let err = read_private_key(Path::new("docker/certs/does-not-exist.key")).unwrap_err(); + let msg = format!("{err:?}"); + assert!( + msg.contains("Could not read private key"), + "error should name the read failure, got: {msg}" + ); + } + + #[test] + fn read_private_key_reads_pem() { + let key = read_private_key(Path::new("docker/certs/server.key")).unwrap(); + assert!(!key.secret_der().is_empty()); + } + + #[test] + fn read_private_key_reads_der() { + // No .der fixture is checked in, so derive one from the PEM key and write + // it to a temp file to exercise the `der` branch. + let der = read_private_key(Path::new("docker/certs/server.key")) + .unwrap() + .secret_der() + .to_vec(); + let mut path = std::env::temp_dir(); + path.push(format!( + "tiberius_read_private_key_{}.der", + std::process::id() + )); + std::fs::write(&path, &der).unwrap(); + let key = read_private_key(&path); + std::fs::remove_file(&path).ok(); + assert!(!key.unwrap().secret_der().is_empty()); + } + + #[test] + fn read_private_key_unsupported_extension_errors() { + // README.md exists under docker/certs but isn't a supported key type. + let err = read_private_key(Path::new("docker/certs/README.md")).unwrap_err(); + let msg = format!("{err:?}"); + assert!( + msg.contains("unsupported file-extension"), + "error should name the unsupported extension, got: {msg}" + ); + } + + #[test] + fn read_private_key_malformed_pem_errors() { + let mut path = std::env::temp_dir(); + path.push(format!( + "tiberius_read_private_key_malformed_{}.pem", + std::process::id() + )); + std::fs::write( + &path, + b"-----BEGIN PRIVATE KEY-----\nnot valid base64!!!\n-----END PRIVATE KEY-----\n", + ) + .unwrap(); + let err = read_private_key(&path); + std::fs::remove_file(&path).ok(); + let msg = format!("{:?}", err.unwrap_err()); + assert!( + msg.contains("Failed to parse PEM private key"), + "error should name the parse failure, got: {msg}" + ); + } + + fn cert_io_error(cert_err: CertificateError) -> io::Error { + // Mirror how tokio-rustls wraps a rustls handshake failure: the + // `RustlsError` is carried as the inner error of an `io::Error`. + io::Error::new( + io::ErrorKind::InvalidData, + RustlsError::InvalidCertificate(cert_err), + ) + } + + /// Reproduce the exact shape of the failure: webpki's `UnsupportedCertVersion` + /// reaches rustls as `CertificateError::Other(..)` whose message contains the + /// string "UnsupportedCertVersion". + fn unsupported_cert_version_error() -> CertificateError { + use tokio_rustls::rustls::OtherError; + + // `CertificateError`'s `Display` renders the `Other` variant via `{:?}` + // (Debug), so — matching real webpki, whose `UnsupportedCertVersion` + // Debugs to exactly that string — the fake must Debug to the same text. + struct WebpkiLike; + impl std::fmt::Debug for WebpkiLike { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("UnsupportedCertVersion") + } + } + impl std::fmt::Display for WebpkiLike { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("UnsupportedCertVersion") + } + } + impl std::error::Error for WebpkiLike {} + + CertificateError::Other(OtherError(Arc::new(WebpkiLike))) + } + + #[test] + fn map_handshake_error_annotates_unsupported_cert_version() { + let err = map_handshake_error( + cert_io_error(unsupported_cert_version_error()), + &default_trust(), + ); + let msg = format!("{err}"); + // Must be a TLS error (not the opaque Io variant) and name each remedy. + assert!( + matches!(err, Error::Tls(_)), + "expected Error::Tls, got: {msg}" + ); + assert!( + msg.contains("trust_cert_ca"), + "should mention trust_cert_ca: {msg}" + ); + assert!( + msg.contains("trust_cert()"), + "should mention trust_cert: {msg}" + ); + assert!( + msg.contains("native-tls"), + "should mention native-tls backend: {msg}" + ); + assert!( + msg.contains("Encrypt=false"), + "should explain the handshake still happens: {msg}" + ); + // Version-specific hint present for this variant. + assert!( + msg.contains("outdated X.509 version"), + "should add version hint: {msg}" + ); + } + + #[test] + fn map_handshake_error_annotates_other_cert_errors_without_version_hint() { + let err = map_handshake_error( + cert_io_error(CertificateError::NotValidForName), + &default_trust(), + ); + let msg = format!("{err}"); + assert!( + matches!(err, Error::Tls(_)), + "expected Error::Tls, got: {msg}" + ); + assert!( + msg.contains("trust_cert_ca"), + "should mention remedies: {msg}" + ); + // The version-specific hint is only for UnsupportedCertVersion. + assert!( + !msg.contains("outdated X.509 version"), + "no version hint here: {msg}" + ); + } + + #[test] + fn map_handshake_error_passes_through_non_cert_io_errors() { + let err = map_handshake_error( + io::Error::new(io::ErrorKind::UnexpectedEof, "connection reset"), + &default_trust(), + ); + // Non-certificate transport failures keep the generic Io mapping. + match err { + Error::Io { kind, message } => { + assert_eq!(kind, io::ErrorKind::UnexpectedEof); + assert!(message.contains("connection reset"), "got: {message}"); } + other => panic!("expected Error::Io, got: {other:?}"), } - tracing::event!( - Level::TRACE, - "with_native_roots processed {} valid and {} invalid certs", - valid_count, - invalid_count + } + + #[test] + fn map_handshake_error_passes_through_non_certificate_rustls_errors() { + // A rustls error that is *not* a certificate-validation failure (e.g. a + // transport/decrypt error) must keep the generic Io mapping — the cert + // remedies would be misleading for it. + let inner = RustlsError::General("handshake alert".to_string()); + let err = map_handshake_error( + io::Error::new(io::ErrorKind::InvalidData, inner), + &default_trust(), + ); + assert!( + matches!(err, Error::Io { .. }), + "non-certificate rustls errors must pass through as Io, got: {err:?}" + ); + } + + #[test] + fn no_cert_verifier_accepts_any_certificate_when_opted_in() { + // `trust_cert` (TrustAll) installs `NoCertVerifier`, which must accept a + // well-formed certificate that does NOT chain to any trusted root — this + // is the explicit opt-in bypass. This is the very same + // certificate the default verifier rejects in + // `default_verifier_rejects_untrusted_certificate` below, so the pair + // proves the bypass is opt-in rather than the default. + let leaf = certs::certs_from_file(Path::new("docker/certs/server.crt")).unwrap(); + let leaf = leaf.into_iter().next().unwrap(); + let verifier = NoCertVerifier; + let name = ServerName::try_from("legacy.sql.example.com").unwrap(); + assert!( + verifier + .verify_server_cert(&leaf, &[], &name, &[], UnixTime::now()) + .is_ok(), + "TrustAll must accept an untrusted server certificate" + ); + } + + #[test] + fn default_verifier_rejects_untrusted_certificate() { + // Security-preserving invariant: with no explicit opt-in, a *well-formed* + // certificate that does not chain to a trusted root MUST be rejected with + // a genuine trust-chain failure (`UnknownIssuer`) — not merely a DER-parse + // error. `NoCertVerifier` accepts this exact certificate under `trust_cert` + // (see `no_cert_verifier_accepts_any_certificate_when_opted_in`), so this + // proves the bypass is opt-in, not the default. + // + // `server.crt` is a well-formed leaf issued by `customCA`. Build a trust + // store that does NOT contain `customCA` (it trusts an unrelated anchor), + // so the leaf's real issuer is untrusted and path building must fail — this + // exercises chain-of-trust enforcement, which a malformed blob would never + // reach (it would be rejected at DER parsing instead). + let leaf = certs::certs_from_file(Path::new("docker/certs/server.crt")).unwrap(); + let leaf = leaf.into_iter().next().unwrap(); + + let mut roots = RootCertStore::empty(); + // Trust an unrelated anchor; `customCA` (the actual issuer of `leaf`) is + // deliberately absent, so `leaf` cannot chain to anything trusted. + roots.add(leaf.clone()).unwrap(); + + let provider = Arc::new(aws_lc_rs::default_provider()); + let verifier = tokio_rustls::rustls::client::WebPkiServerVerifier::builder_with_provider( + Arc::new(roots), + provider, + ) + .build() + .expect("verifier builds over a non-empty root store"); + + let name = ServerName::try_from("localhost").unwrap(); + // Pin verification time inside `server.crt`'s validity window + // (notBefore 2026-05-11, notAfter 2027-11-02). webpki checks certificate + // validity *before* issuer matching, so using the wall clock would make + // this assertion flip from `UnknownIssuer` to `Expired` once the fixture + // expires — a spurious failure unrelated to any code change. 2027-01-15. + let at = UnixTime::since_unix_epoch(std::time::Duration::from_secs(1_800_000_000)); + let result = verifier.verify_server_cert(&leaf, &[], &name, &[], at); + assert!( + matches!( + result, + Err(RustlsError::InvalidCertificate( + CertificateError::UnknownIssuer + )) + ), + "the default verifier MUST reject a cert that does not chain to a trusted \ + root with UnknownIssuer (a real trust failure, not a parse error), got: {result:?}" + ); + } + + #[test] + fn map_handshake_error_trustall_does_not_annotate() { + // Under TrustAll the verifier never rejects a cert, so even a cert-shaped + // error must pass through unchanged rather than gaining misleading advice. + let err = map_handshake_error( + cert_io_error(unsupported_cert_version_error()), + &bypass_trust(), + ); + assert!( + matches!(err, Error::Io { .. }), + "TrustAll must not rewrite the error into TLS guidance" ); - assert!(!roots.is_empty(), "no CA certificates found"); + } + + #[test] + fn map_handshake_error_trustall_passes_through_non_cert_io_errors() { + // Under TrustAll a plain transport error must also pass through as Io, + // preserving kind and message (the cert-detection block is skipped + // wholesale, regardless of the error's shape). + let err = map_handshake_error( + io::Error::new(io::ErrorKind::ConnectionReset, "reset by peer"), + &bypass_trust(), + ); + match err { + Error::Io { kind, message } => { + assert_eq!(kind, io::ErrorKind::ConnectionReset); + assert!(message.contains("reset by peer"), "got: {message}"); + } + other => panic!("expected Error::Io, got: {other:?}"), + } + } - self.with_root_certificates(roots) + #[test] + fn supported_verify_schemes_are_stable() { + let schemes = NoCertVerifier.supported_verify_schemes(); + assert_eq!(schemes.len(), 11); + assert!(schemes.contains(&SignatureScheme::ED25519)); } } diff --git a/src/command.rs b/src/command.rs new file mode 100644 index 000000000..ccc48f963 --- /dev/null +++ b/src/command.rs @@ -0,0 +1,469 @@ +use std::borrow::Cow; + +use enumflags2::BitFlags; +use futures_util::io::{AsyncRead, AsyncWrite}; + +use crate::{ + tds::{ + codec::{RpcParam, RpcStatus::ByRefValue, RpcValue, TypeInfoTvp}, + stream::{CommandStream, TokenStream}, + }, + Client, ColumnData, IntoSql, +}; + +#[doc(inline)] +pub use tiberius_macros::TableValueRow; + +/// A structure that represents a single row of a table-valued parameter (TVP) +/// implements this trait. +/// +/// It can be derived with `#[derive(TableValueRow)]` for structs with named +/// fields. +pub trait TableValueRow<'a> { + /// Binds this row's field values. Called by [`Command`] before making the + /// call to the server; implementations must call + /// [`SqlTableDataRow::add_field`] once per column, in column order. + fn bind_fields(&self, data_row: &mut SqlTableDataRow<'a>); + /// The database type name that represents this TVP, e.g. `dbo.MyType`. + fn get_db_type() -> &'static str; +} + +/// A collection of [`TableValueRow`] values that can be bound as a +/// table-valued parameter. Implemented for any `IntoIterator` of rows. +pub trait TableValue<'a> { + /// Converts this collection into the internal table data representation. + fn into_sql(self) -> SqlTableData<'a>; +} + +impl<'a, R, C> TableValue<'a> for C +where + R: TableValueRow<'a> + 'a, + C: IntoIterator, +{ + fn into_sql(self) -> SqlTableData<'a> { + let mut data = Vec::new(); + for row in self { + let mut data_row = SqlTableDataRow::new(); + row.bind_fields(&mut data_row); + data.push(data_row); + } + + SqlTableData { + rows: data, + db_type: R::get_db_type(), + } + } +} + +/// A remote command (stored procedure or user-defined function) with bound +/// parameters, executed by name via an RPC request. +#[derive(Debug)] +pub struct Command<'a> { + name: Cow<'a, str>, + // The server rejects repeated parameter names, so uniqueness is not checked here. + params: Vec>, +} + +#[derive(Debug)] +struct CommandParam<'a> { + name: Cow<'a, str>, + out: bool, + data: CommandParamData<'a>, +} + +#[derive(Debug)] +enum CommandParamData<'a> { + Scalar(ColumnData<'a>), + Table(SqlTableData<'a>), +} + +/// The internal representation of a table-valued parameter's data. +#[derive(Debug)] +pub struct SqlTableData<'a> { + rows: Vec>, + db_type: &'a str, +} + +/// A single row of a table-valued parameter, used by [`TableValueRow`] +/// implementations to bind column values. +#[derive(Debug)] +pub struct SqlTableDataRow<'a> { + col_data: Vec>, +} + +impl<'a> SqlTableDataRow<'a> { + fn new() -> SqlTableDataRow<'a> { + SqlTableDataRow { + col_data: Vec::new(), + } + } + + /// Adds a field value to this TVP row. Must be called once per column; the + /// values are sent to the server in call order. + pub fn add_field(&mut self, data: impl IntoSql<'a> + 'a) { + self.col_data.push(data.into_sql()); + } + + /// The column values bound to this row, in bind order. Exposed for testing + /// `TableValueRow` implementations without a live server connection. + #[doc(hidden)] + pub fn columns(&self) -> &[ColumnData<'a>] { + &self.col_data + } +} + +impl<'a> SqlTableData<'a> { + /// The rows collected for this table-valued parameter. Exposed for testing + /// `TableValue`/`TableValueRow` implementations without a live server. + #[doc(hidden)] + pub fn rows(&self) -> &[SqlTableDataRow<'a>] { + &self.rows + } + + /// The database type name for this table-valued parameter. Exposed for + /// testing. + #[doc(hidden)] + pub fn db_type(&self) -> &str { + self.db_type + } +} + +impl<'a> Command<'a> { + /// Constructs a new command with the given procedure or function name. + pub fn new(proc_name: impl Into>) -> Self { + Self { + name: proc_name.into(), + params: Vec::new(), + } + } + + /// Binds a scalar input parameter with the given name. + pub fn bind_param(&mut self, name: impl Into>, data: impl IntoSql<'a> + 'a) { + self.params.push(CommandParam { + name: name.into(), + out: false, + data: CommandParamData::Scalar(data.into_sql()), + }); + } + + /// Binds a by-ref (OUT) scalar parameter. The returned value can be found by + /// the same name in the [`CommandResult`] returned values. + /// + /// [`CommandResult`]: crate::CommandResult + pub fn bind_out_param(&mut self, name: impl Into>, data: impl IntoSql<'a> + 'a) { + self.params.push(CommandParam { + name: name.into(), + out: true, + data: CommandParamData::Scalar(data.into_sql()), + }); + } + + /// Binds a table-valued parameter. The provided argument must implement + /// [`TableValue`]. + /// + /// # Example + /// + /// ```no_run + /// # use std::env; + /// # use tiberius::Config; + /// # use tiberius::{numeric::Numeric, Command, TableValueRow}; + /// # use tokio_util::compat::TokioAsyncWriteCompatExt; + /// #[derive(TableValueRow)] + /// struct SomeGeoList { + /// eid: i32, + /// lat: Numeric, + /// lon: Numeric, + /// } + /// # #[tokio::main] + /// # async fn main() -> Result<(), Box> { + /// # let c_str = env::var("TIBERIUS_TEST_CONNECTION_STRING").unwrap_or( + /// # "server=tcp:localhost,1433;integratedSecurity=true;TrustServerCertificate=true".to_owned(), + /// # ); + /// # let config = Config::from_ado_string(&c_str)?; + /// # let tcp = tokio::net::TcpStream::connect(config.get_addr()).await?; + /// # tcp.set_nodelay(true)?; + /// # let client = tiberius::Client::connect(config, tcp.compat_write()).await?; + /// let r1 = SomeGeoList { + /// eid: 1, + /// lat: Numeric::new_with_scale(10, 6), + /// lon: Numeric::new_with_scale(14, 6), + /// }; + /// let r2 = SomeGeoList { + /// eid: 4, + /// lat: Numeric::new_with_scale(101, 6), + /// lon: Numeric::new_with_scale(142, 6), + /// }; + /// + /// let tbl = vec![r1, r2]; + /// + /// let mut cmd = Command::new("dbo.usp_TheGeoProcedure"); + /// cmd.bind_table("@table", tbl); + /// # Ok(()) + /// # } + /// ``` + pub fn bind_table(&mut self, name: impl Into>, data: impl TableValue<'a> + 'a) { + self.params.push(CommandParam { + name: name.into(), + out: false, + data: CommandParamData::Table(data.into_sql()), + }); + } + + /// The same as [`bind_table`](Self::bind_table), but overrides the database + /// type name used for the TVP. + /// + /// # Security + /// + /// `db_type` is interpolated directly into a SQL batch when the command is + /// executed (`DECLARE @P AS {db_type};SELECT TOP 0 * FROM @P`), because + /// T-SQL does not allow a type name to be parameterized. It must therefore + /// be a trusted identifier (a compile-time constant or a value from a + /// vetted allow-list), never untrusted or user-influenced input. A cheap + /// defense-in-depth guard rejects obviously-malformed identifiers (NUL/ASCII + /// control characters or an unbalanced `]` bracket) at [`exec`](Self::exec) + /// time, but that guard is not a substitute for passing a trusted value. + pub fn bind_table_with_dbtype( + &mut self, + name: impl Into>, + db_type: &'a str, + data: impl TableValue<'a> + 'a, + ) { + self.params.push(CommandParam { + name: name.into(), + out: false, + data: CommandParamData::Table(SqlTableData { + db_type, + ..data.into_sql() + }), + }); + } + + /// Executes the command on the server, returning a [`CommandStream`] that + /// can be collected into a [`CommandResult`] for convenience. + /// + /// [`CommandResult`]: crate::CommandResult + /// + /// # Example + /// + /// ```no_run + /// # use tiberius::{Config, Command}; + /// # use tokio_util::compat::TokioAsyncWriteCompatExt; + /// # use std::env; + /// # #[tokio::main] + /// # async fn main() -> Result<(), Box> { + /// # let c_str = env::var("TIBERIUS_TEST_CONNECTION_STRING").unwrap_or( + /// # "server=tcp:localhost,1433;integratedSecurity=true;TrustServerCertificate=true".to_owned(), + /// # ); + /// # let config = Config::from_ado_string(&c_str)?; + /// # let tcp = tokio::net::TcpStream::connect(config.get_addr()).await?; + /// # tcp.set_nodelay(true)?; + /// # let mut client = tiberius::Client::connect(config, tcp.compat_write()).await?; + /// let mut cmd = Command::new("dbo.usp_SomeStoredProc"); + /// + /// cmd.bind_param("@foo", 34i32); + /// cmd.bind_out_param("@bar", "bar"); + /// let res = cmd.exec(&mut client).await?.into_command_result().await?; + /// + /// let rv: Option<&str> = res.try_return_value("@bar")?; + /// let rc = res.return_code(); + /// # Ok(()) + /// # } + /// ``` + pub async fn exec<'b, S>(self, client: &'b mut Client) -> crate::Result> + where + S: AsyncRead + AsyncWrite + Unpin + Send, + { + let rpc_params = Command::build_rpc_params(self.params, client).await?; + + client.connection.flush_stream().await?; + client.rpc_run_command(self.name, rpc_params).await?; + + let ts = TokenStream::new(&mut client.connection); + let result = CommandStream::new(ts.try_unfold()); + + Ok(result) + } + + async fn build_rpc_params<'b, S>( + cmd_params: Vec>, + client: &'b mut Client, + ) -> crate::Result>> + where + S: AsyncRead + AsyncWrite + Unpin + Send, + { + let mut rpc_params = Vec::new(); + for p in cmd_params { + let rpc_val = match p.data { + CommandParamData::Scalar(col) => RpcValue::Scalar(col), + CommandParamData::Table(t) => { + // `db_type` is interpolated raw into the batch below (T-SQL + // cannot parameterize a type name), so reject obviously + // dangerous identifiers before building the SQL. + validate_db_type_identifier(t.db_type)?; + let type_info_tvp = TypeInfoTvp::new( + t.db_type, + t.rows.into_iter().map(|r| r.col_data).collect(), + ); + // Resolve the TVP column layout from the server. + let cols_metadata = client + .query_run_for_metadata(format!( + "DECLARE @P AS {};SELECT TOP 0 * FROM @P", + t.db_type + )) + .await?; + RpcValue::Table(if let Some(cm) = cols_metadata { + type_info_tvp.with_metadata(cm) + } else { + type_info_tvp + }) + } + }; + let rpc_param = RpcParam { + name: p.name, + flags: if p.out { + BitFlags::from_flag(ByRefValue) + } else { + BitFlags::empty() + }, + value: rpc_val, + }; + rpc_params.push(rpc_param); + } + Ok(rpc_params) + } +} + +/// Reject an obviously-malformed or dangerous TVP `db_type` identifier. +/// +/// The `db_type` of a table-valued parameter is interpolated directly into a +/// SQL batch (`DECLARE @P AS {db_type};SELECT TOP 0 * FROM @P`) because T-SQL +/// does not allow a type name to be parameterized. This guard is cheap +/// defense-in-depth — it does NOT make untrusted input safe. It rejects input +/// that cannot be a legitimate type identifier: +/// +/// - a NUL byte or any ASCII control character; +/// - **outside** a `[...]` bracket-quoted segment, any character that is not +/// identifier-safe. Only alphanumerics and the punctuation needed for real +/// type names are allowed — `_`, `.` (multi-part names like `dbo.MyType`), +/// `@`/`#` (variable/temp-style names), a plain space, and `(`, `)`, `,` +/// (parameterized types such as `decimal(10,2)` / `varchar(max)`). This +/// rejects statement-breaking characters such as `;`, quotes and `-`, so a +/// value like `int; DROP TABLE x--` cannot slip through; and +/// - **inside** a `[...]` bracket-quoted segment anything is allowed except an +/// unescaped `]` (per the T-SQL bracket-escaping rule a literal `]` must be +/// doubled as `]]`). A `]` seen outside any bracket is unbalanced and +/// rejected. +/// +/// It deliberately does NOT try to quote or rewrite the identifier, so +/// multi-part names (`dbo.MyType`) and already-bracketed names (`[my type]`) +/// keep working unchanged. This delegates to the shared +/// [`crate::client::validate_sql_identifier`], which the bulk-insert guards in +/// `src/client.rs` use too. +fn validate_db_type_identifier(db_type: &str) -> crate::Result<()> { + crate::client::validate_sql_identifier("TVP db_type", db_type) +} + +#[cfg(test)] +mod tests { + use super::*; + + struct TestRow; + + impl<'a> TableValueRow<'a> for TestRow { + fn bind_fields(&self, row: &mut SqlTableDataRow<'a>) { + row.add_field(1i32); + } + + fn get_db_type() -> &'static str { + "default.Type" + } + } + + #[test] + fn bind_table_with_dbtype_uses_the_explicit_db_type() { + // The explicit db_type argument must override the row's own get_db_type(). + let mut cmd = Command::new("proc"); + cmd.bind_table_with_dbtype("@tvp", "explicit.Type", vec![TestRow]); + + assert_eq!(cmd.params.len(), 1); + assert_eq!(cmd.params[0].name, "@tvp"); + match &cmd.params[0].data { + CommandParamData::Table(t) => assert_eq!(t.db_type, "explicit.Type"), + other => panic!("expected a table parameter, got {other:?}"), + } + } + + #[test] + fn bind_table_uses_the_rows_db_type() { + let mut cmd = Command::new("proc"); + cmd.bind_table("@tvp", vec![TestRow]); + + match &cmd.params[0].data { + CommandParamData::Table(t) => assert_eq!(t.db_type, "default.Type"), + other => panic!("expected a table parameter, got {other:?}"), + } + } + + #[test] + fn validate_db_type_accepts_normal_type_names() { + for db_type in [ + "int", + "dbo.MyType", + "[my type]", + "[dbo].[my type]", + "[weird]]type]", + "decimal(18,4)", + "decimal(10,2)", + "varchar(max)", + ] { + assert!( + validate_db_type_identifier(db_type).is_ok(), + "expected {db_type:?} to be accepted", + ); + } + } + + #[test] + fn validate_db_type_rejects_statement_breaking_characters() { + // An injected statement terminator / comment must not pass the guard. + for db_type in [ + "int; DROP TABLE x", + "int; DROP TABLE x--", + "int' OR '1'='1", + "int\" ", + "foo-- bar", + ] { + assert!( + matches!( + validate_db_type_identifier(db_type), + Err(crate::Error::BulkInput(_)) + ), + "expected {db_type:?} to be rejected", + ); + } + } + + #[test] + fn validate_db_type_rejects_control_characters() { + assert!(matches!( + validate_db_type_identifier("int\0"), + Err(crate::Error::BulkInput(_)) + )); + assert!(matches!( + validate_db_type_identifier("dbo.\nMyType"), + Err(crate::Error::BulkInput(_)) + )); + } + + #[test] + fn validate_db_type_rejects_unbalanced_closing_bracket() { + assert!(matches!( + validate_db_type_identifier("MyType]"), + Err(crate::Error::BulkInput(_)) + )); + assert!(matches!( + validate_db_type_identifier("a]b"), + Err(crate::Error::BulkInput(_)) + )); + } +} diff --git a/src/error.rs b/src/error.rs index 98bf01b58..23808b6ea 100644 --- a/src/error.rs +++ b/src/error.rs @@ -8,8 +8,8 @@ use thiserror::Error; /// the lifecycle of this driver #[derive(Debug, Clone, Error, PartialEq, Eq)] pub enum Error { - #[error("An error occured during the attempt of performing I/O: {}", message)] - /// An error occured when performing I/O to the server. + #[error("An error occurred during the attempt of performing I/O: {}", message)] + /// An error occurred when performing I/O to the server. Io { /// A list specifying general categories of I/O error. kind: IoErrorKind, @@ -27,9 +27,15 @@ pub enum Error { Conversion(Cow<'static, str>), #[error("UTF-8 error")] /// Tried to convert data to UTF-8 that was not valid. + /// + /// The originating [`std::str::Utf8Error`]/[`std::string::FromUtf8Error`] + /// source is intentionally not carried on this variant. Utf8, #[error("UTF-16 error")] /// Tried to convert data to UTF-16 that was not valid. + /// + /// The originating [`std::string::FromUtf16Error`] source is intentionally + /// not carried on this variant. Utf16, #[error("Error parsing an integer: {}", _0)] /// Tried to parse an integer that was not an integer. @@ -41,13 +47,15 @@ pub enum Error { /// An error in the TLS handshake. Tls(String), #[cfg(any(all(unix, feature = "integrated-auth-gssapi"), doc))] - #[cfg_attr( - feature = "docs", - doc(cfg(all(unix, feature = "integrated-auth-gssapi"))) - )] + #[cfg_attr(docsrs, doc(cfg(all(unix, feature = "integrated-auth-gssapi"))))] /// An error from the GSSAPI library. #[error("GSSAPI Error: {}", _0)] Gssapi(String), + #[cfg(any(all(unix, feature = "sspi-rs"), doc))] + #[cfg_attr(docsrs, doc(cfg(all(unix, feature = "sspi-rs"))))] + /// An error from the `sspi` (sspi-rs) library. + #[error("sspi-rs Error: {}", _0)] + SspiRs(String), #[error( "Server requested a connection to an alternative address: `{}:{}`", host, @@ -83,7 +91,7 @@ impl Error { impl From for Error { fn from(e: uuid::Error) -> Self { - Self::Conversion(format!("Error convertiong a Guid value {}", e).into()) + Self::Conversion(format!("Error converting a Guid value {}", e).into()) } } @@ -148,12 +156,150 @@ impl From for Error { } #[cfg(all(unix, feature = "integrated-auth-gssapi"))] -#[cfg_attr( - feature = "docs", - doc(cfg(all(unix, feature = "integrated-auth-gssapi"))) -)] +#[cfg_attr(docsrs, doc(cfg(all(unix, feature = "integrated-auth-gssapi"))))] impl From for Error { fn from(err: libgssapi::error::Error) -> Error { Error::Gssapi(format!("{}", err)) } } + +#[cfg(all(unix, feature = "sspi-rs"))] +#[cfg_attr(docsrs, doc(cfg(all(unix, feature = "sspi-rs"))))] +impl From for Error { + fn from(err: sspi::Error) -> Error { + Error::SspiRs(format!("{}", err)) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn token_error(code: u32) -> TokenError { + TokenError { + code, + state: 1, + class: 16, + message: "boom".to_string(), + server: "srv".to_string(), + procedure: "proc".to_string(), + line: 3, + } + } + + #[test] + fn code_and_is_deadlock() { + let deadlock = Error::Server(token_error(1205)); + assert_eq!(deadlock.code(), Some(1205)); + assert!(deadlock.is_deadlock()); + + let other = Error::Server(token_error(500)); + assert_eq!(other.code(), Some(500)); + assert!(!other.is_deadlock()); + + let non_server = Error::Utf8; + assert_eq!(non_server.code(), None); + assert!(!non_server.is_deadlock()); + } + + #[test] + fn display_variants() { + assert_eq!( + format!("{}", Error::Protocol("bad".into())), + "Protocol error: bad" + ); + assert_eq!( + format!("{}", Error::Encoding("bad".into())), + "Encoding error: bad" + ); + assert_eq!( + format!("{}", Error::Conversion("bad".into())), + "Conversion error: bad" + ); + assert_eq!(format!("{}", Error::Utf8), "UTF-8 error"); + assert_eq!(format!("{}", Error::Utf16), "UTF-16 error"); + assert_eq!( + format!("{}", Error::BulkInput("bad".into())), + "BULK UPLOAD input failure: bad" + ); + + let routing = Error::Routing { + host: "host".to_string(), + port: 1234, + }; + assert!(format!("{}", routing).contains("host:1234")); + } + + #[test] + fn from_io_error() { + let io_err = io::Error::new(io::ErrorKind::UnexpectedEof, "eof"); + let err: Error = io_err.into(); + match err { + Error::Io { kind, message } => { + assert_eq!(kind, io::ErrorKind::UnexpectedEof); + assert!(message.contains("eof")); + } + _ => panic!("expected Io"), + } + } + + #[test] + fn from_parse_int_error() { + let parse_err = "not-a-number".parse::().unwrap_err(); + let err: Error = parse_err.into(); + assert!(matches!(err, Error::ParseInt(_))); + } + + #[test] + #[allow(invalid_from_utf8)] // intentionally-invalid bytes to exercise the error path + fn from_utf8_and_utf16_errors() { + let utf8_err = String::from_utf8(vec![0xff, 0xfe]).unwrap_err(); + assert!(matches!(Error::from(utf8_err), Error::Utf8)); + + let invalid: &[u8] = &[0xff, 0xfe]; + let str_utf8 = std::str::from_utf8(invalid).unwrap_err(); + assert!(matches!(Error::from(str_utf8), Error::Utf8)); + + let utf16_err = String::from_utf16(&[0xd800]).unwrap_err(); + assert!(matches!(Error::from(utf16_err), Error::Utf16)); + } + + #[test] + fn from_uuid_error() { + let uuid_err = uuid::Uuid::parse_str("not-a-uuid").unwrap_err(); + assert!(matches!(Error::from(uuid_err), Error::Conversion(_))); + } + + #[test] + fn equality_between_errors() { + assert_eq!(Error::Utf8, Error::Utf8); + assert_ne!(Error::Utf8, Error::Utf16); + } + + #[test] + fn from_connection_string_error() { + let cs_err = connection_string::Error::new("bad connection string"); + let err: Error = cs_err.into(); + match err { + Error::Conversion(msg) => assert!(msg.contains("bad connection string")), + _ => panic!("expected Conversion"), + } + } + + #[cfg(all(unix, feature = "sspi-rs"))] + #[test] + fn from_sspi_error() { + let sspi_err = sspi::Error::new(sspi::ErrorKind::InternalError, "sspi boom"); + assert!(matches!(Error::from(sspi_err), Error::SspiRs(_))); + } + + #[cfg(all(unix, feature = "integrated-auth-gssapi"))] + #[test] + fn from_gssapi_error() { + let gss_err = libgssapi::error::Error { + major: libgssapi::error::MajorFlags::empty(), + minor: 0, + }; + assert!(matches!(Error::from(gss_err), Error::Gssapi(_))); + } +} diff --git a/src/from_sql.rs b/src/from_sql.rs index 8498fa01c..2769534dc 100644 --- a/src/from_sql.rs +++ b/src/from_sql.rs @@ -25,7 +25,8 @@ use uuid::Uuid; /// |[`NaiveDateTime`] (with feature flag `chrono`)|`datetime`/`datetime2`/`smalldatetime`| /// |[`NaiveDate`] (with feature flag `chrono`)|`date`| /// |[`NaiveTime`] (with feature flag `chrono`)|`time`| -/// |[`DateTime`] (with feature flag `chrono`)|`datetimeoffset`| +/// |[`DateTime`]`` (with feature flag `chrono`)|`datetimeoffset`/`datetime2`| +/// |[`DateTime`]`` (with feature flag `chrono`)|`datetimeoffset`| /// /// See the [`time`] module for more information about the date and time structs. /// @@ -60,7 +61,7 @@ where from_sql!(bool: ColumnData::Bit(val) => (*val, val)); from_sql!(u8: ColumnData::U8(val) => (*val, val), ColumnData::I32(None) => (None, None)); from_sql!(i16: ColumnData::I16(val) => (*val, val), ColumnData::U8(None) => (None, None), ColumnData::I32(None) => (None, None)); -from_sql!(i32: ColumnData::I32(val) => (*val, val), ColumnData::U8(None) => (None, None)); +from_sql!(i32: ColumnData::I32(val) => (*val, val), ColumnData::I16(val) => (val.map(i32::from), val.map(i32::from)), ColumnData::U8(None) => (None, None)); from_sql!(i64: ColumnData::I64(val) => (*val, val), ColumnData::U8(None) => (None, None), ColumnData::I32(None) => (None, None)); from_sql!(f32: ColumnData::F32(val) => (*val, val)); from_sql!(f64: ColumnData::F64(val) => (*val, val)); @@ -132,3 +133,161 @@ impl<'a> FromSql<'a> for &'a [u8] { } } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn i16_column_converts_to_i32() { + let data = ColumnData::I16(Some(8)); + assert_eq!(Some(8i32), i32::from_sql(&data).unwrap()); + assert_eq!(Some(8i32), i32::from_sql_owned(data).unwrap()); + } + + #[test] + fn null_i16_column_converts_to_i32() { + let data = ColumnData::I16(None); + assert_eq!(None, i32::from_sql(&data).unwrap()); + assert_eq!(None, i32::from_sql_owned(ColumnData::I16(None)).unwrap()); + } + + #[test] + fn bool_from_bit() { + let data = ColumnData::Bit(Some(true)); + assert_eq!(Some(true), bool::from_sql(&data).unwrap()); + assert_eq!(Some(true), bool::from_sql_owned(data).unwrap()); + } + + #[test] + fn u8_from_u8_and_null_i32() { + let data = ColumnData::U8(Some(5)); + assert_eq!(Some(5u8), u8::from_sql(&data).unwrap()); + assert_eq!(Some(5u8), u8::from_sql_owned(data).unwrap()); + + let null = ColumnData::I32(None); + assert_eq!(None, u8::from_sql(&null).unwrap()); + assert_eq!(None, u8::from_sql_owned(null).unwrap()); + } + + #[test] + fn i16_from_wrong_variant_errors() { + let data = ColumnData::F64(Some(1.0)); + let err = i16::from_sql(&data).unwrap_err(); + assert!(format!("{}", err).contains("cannot interpret")); + } + + #[test] + fn i64_from_i64_and_null() { + let data = ColumnData::I64(Some(42)); + assert_eq!(Some(42i64), i64::from_sql(&data).unwrap()); + assert_eq!(Some(42i64), i64::from_sql_owned(data).unwrap()); + + let null = ColumnData::U8(None); + assert_eq!(None, i64::from_sql_owned(null).unwrap()); + } + + #[test] + fn f32_and_f64_from_sql() { + let f32_data = ColumnData::F32(Some(1.5)); + assert_eq!(Some(1.5f32), f32::from_sql(&f32_data).unwrap()); + + let f64_data = ColumnData::F64(Some(2.5)); + assert_eq!(Some(2.5f64), f64::from_sql(&f64_data).unwrap()); + } + + #[test] + fn uuid_from_guid() { + let uuid = Uuid::new_v4(); + let data = ColumnData::Guid(Some(uuid)); + assert_eq!(Some(uuid), Uuid::from_sql(&data).unwrap()); + assert_eq!(Some(uuid), Uuid::from_sql_owned(data).unwrap()); + } + + #[test] + fn numeric_from_numeric() { + let numeric = crate::tds::Numeric::new_with_scale(1234, 2); + let data = ColumnData::Numeric(Some(numeric)); + assert_eq!(Some(numeric), Numeric::from_sql(&data).unwrap()); + assert_eq!(Some(numeric), Numeric::from_sql_owned(data).unwrap()); + } + + #[test] + fn xml_data_owned_and_borrowed() { + let xml = XmlData::new("".to_string()); + let data = ColumnData::Xml(Some(std::borrow::Cow::Owned(xml.clone()))); + + let borrowed = <&XmlData as FromSql>::from_sql(&data).unwrap().unwrap(); + assert_eq!(borrowed.to_string(), xml.to_string()); + + let owned = XmlData::from_sql_owned(data).unwrap().unwrap(); + assert_eq!(owned.to_string(), xml.to_string()); + } + + #[test] + fn xml_data_wrong_variant_errors() { + let data = ColumnData::I32(Some(1)); + let err = XmlData::from_sql_owned(data).unwrap_err(); + assert!(format!("{}", err).contains("cannot interpret")); + + let data = ColumnData::I32(Some(1)); + let err = <&XmlData as FromSql>::from_sql(&data).unwrap_err(); + assert!(format!("{}", err).contains("cannot interpret")); + } + + #[test] + fn string_owned_and_borrowed_str() { + let data = ColumnData::String(Some(std::borrow::Cow::Borrowed("hello"))); + let borrowed = <&str as FromSql>::from_sql(&data).unwrap(); + assert_eq!(Some("hello"), borrowed); + + let owned = String::from_sql_owned(data).unwrap(); + assert_eq!(Some("hello".to_string()), owned); + } + + #[test] + fn string_wrong_variant_errors() { + let data = ColumnData::I32(Some(1)); + let err = String::from_sql_owned(data).unwrap_err(); + assert!(format!("{}", err).contains("cannot interpret")); + + let data = ColumnData::I32(Some(1)); + let err = <&str as FromSql>::from_sql(&data).unwrap_err(); + assert!(format!("{}", err).contains("cannot interpret")); + } + + #[test] + fn binary_owned_and_borrowed_slice() { + let bytes = vec![1u8, 2, 3]; + let data = ColumnData::Binary(Some(std::borrow::Cow::Owned(bytes.clone()))); + + let borrowed = <&[u8] as FromSql>::from_sql(&data).unwrap(); + assert_eq!(Some(bytes.as_slice()), borrowed); + + let owned = Vec::::from_sql_owned(data).unwrap(); + assert_eq!(Some(bytes), owned); + } + + #[test] + fn binary_wrong_variant_errors() { + let data = ColumnData::I32(Some(1)); + let err = Vec::::from_sql_owned(data).unwrap_err(); + assert!(format!("{}", err).contains("cannot interpret")); + + let data = ColumnData::I32(Some(1)); + let err = <&[u8] as FromSql>::from_sql(&data).unwrap_err(); + assert!(format!("{}", err).contains("cannot interpret")); + } + + #[test] + fn null_string_and_binary_values() { + let data = ColumnData::String(None); + assert_eq!(None, String::from_sql_owned(data).unwrap()); + + let data = ColumnData::Binary(None); + assert_eq!(None, Vec::::from_sql_owned(data).unwrap()); + + let data = ColumnData::Xml(None); + assert_eq!(None, XmlData::from_sql_owned(data).unwrap()); + } +} diff --git a/src/lib.rs b/src/lib.rs index 882f5ad36..b8805fba1 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,61 +1,9 @@ //! An asynchronous, runtime-independent, pure-rust Tabular Data Stream (TDS) //! implementation for Microsoft SQL Server. //! -//! # Connecting with async-std -//! -//! Being not bound to any single runtime, a `TcpStream` must be created -//! separately and injected to the [`Client`]. -//! -//! ```no_run -//! use tiberius::{Client, Config, Query, AuthMethod}; -//! use async_std::net::TcpStream; -//! -//! #[async_std::main] -//! async fn main() -> anyhow::Result<()> { -//! // Using the builder method to construct the options. -//! let mut config = Config::new(); -//! -//! config.host("localhost"); -//! config.port(1433); -//! -//! // Using SQL Server authentication. -//! config.authentication(AuthMethod::sql_server("SA", "")); -//! -//! // on production, it is not a good idea to do this -//! config.trust_cert(); -//! -//! // Taking the address from the configuration, using async-std's -//! // TcpStream to connect to the server. -//! let tcp = TcpStream::connect(config.get_addr()).await?; -//! -//! // We'll disable the Nagle algorithm. Buffering is handled -//! // internally with a `Sink`. -//! tcp.set_nodelay(true)?; -//! -//! // Handling TLS, login and other details related to the SQL Server. -//! let mut client = Client::connect(config, tcp).await?; -//! -//! // Constructing a query object with one parameter annotated with `@P1`. -//! // This requires us to bind a parameter that will then be used in -//! // the statement. -//! let mut select = Query::new("SELECT @P1"); -//! select.bind(-4i32); -//! -//! // A response to a query is a stream of data, that must be -//! // polled to the end before querying again. Using streams allows -//! // fetching data in an asynchronous manner, if needed. -//! let stream = select.query(&mut client).await?; -//! -//! // In this case, we know we have only one query, returning one row -//! // and one column, so calling `into_row` will consume the stream -//! // and return us the first row of the first result. -//! let row = stream.into_row().await?; -//! -//! assert_eq!(Some(-4i32), row.unwrap().get(0)); -//! -//! Ok(()) -//! } -//! ``` +//! Tiberius is not bound to any single async runtime: a `TcpStream` is created +//! separately and injected into the [`Client`], so it works with Tokio, smol, +//! and other runtimes that provide `futures::io::{AsyncRead, AsyncWrite}`. //! //! # Connecting with Tokio //! @@ -156,11 +104,11 @@ //! Tiberius supports different [ways of authentication] to the SQL Server: //! //! - SQL Server authentication uses the facilities of the database to -//! authenticate the user. +//! authenticate the user. //! - On Windows, you can authenticate using the currently logged in user or -//! specified Windows credentials. +//! specified Windows credentials. //! - If enabling the `integrated-auth-gssapi` feature, it is possible to login -//! with the currently active Kerberos credentials. +//! with the currently active Kerberos credentials. //! //! ## AAD(Azure Active Directory) Authentication //! @@ -180,22 +128,24 @@ //! //! On Windows platforms, connecting to the SQL Server might require going through //! the SQL Browser service to get the correct port for the named instance. This -//! feature requires either the `sql-browser-async-std` or `sql-browser-tokio` feature -//! flag to be enabled and has a bit different way of connecting: +//! feature requires the `sql-browser-tokio` (or `sql-browser-smol`) feature flag +//! to be enabled and has a bit different way of connecting: //! //! ```no_run -//! # #[cfg(any(feature = "sql-browser-async-std", feature = "sql-browser-tokio"))] +//! # #[cfg(feature = "sql-browser-tokio")] //! use tiberius::{Client, Config, AuthMethod}; -//! # #[cfg(any(feature = "sql-browser-async-std", feature = "sql-browser-tokio"))] -//! use async_std::net::TcpStream; +//! # #[cfg(feature = "sql-browser-tokio")] +//! use tokio::net::TcpStream; +//! # #[cfg(feature = "sql-browser-tokio")] +//! use tokio_util::compat::TokioAsyncWriteCompatExt; //! //! // An extra trait that allows connecting to a named instance with the given //! // `TcpStream`. -//! # #[cfg(any(feature = "sql-browser-async-std", feature = "sql-browser-tokio"))] +//! # #[cfg(feature = "sql-browser-tokio")] //! use tiberius::SqlBrowser; //! -//! #[async_std::main] -//! # #[cfg(any(feature = "sql-browser-async-std", feature = "sql-browser-tokio"))] +//! # #[cfg(feature = "sql-browser-tokio")] +//! #[tokio::main] //! async fn main() -> anyhow::Result<()> { //! let mut config = Config::new(); //! @@ -211,16 +161,16 @@ //! // on production, it is not a good idea to do this //! config.trust_cert(); //! -//! // This will create a new `TcpStream` from `async-std`, connected to the -//! // right port of the named instance. +//! // This will create a new `TcpStream`, connected to the right port of the +//! // named instance. //! let tcp = TcpStream::connect_named(&config).await?; //! //! // And from here on continue the connection process in a normal way. -//! let mut client = Client::connect(config, tcp).await?; +//! let mut client = Client::connect(config, tcp.compat_write()).await?; //! # client.query("SELECT @P1", &[&-4i32]).await?; //! Ok(()) //! } -//! # #[cfg(any(not(feature = "sql-browser-async-std"), not(feature = "sql-browser-tokio")))] +//! # #[cfg(not(feature = "sql-browser-tokio"))] //! # fn main() {} //! ``` //! @@ -243,13 +193,23 @@ //! [`time`]: time/index.html //! [ways of authentication]: enum.AuthMethod.html //! [ADO.NET connection string]: https://docs.microsoft.com/en-us/dotnet/framework/data/adonet/connection-strings -#![cfg_attr(feature = "docs", feature(doc_cfg))] +#![cfg_attr(docsrs, feature(doc_cfg))] #![recursion_limit = "512"] #![warn(missing_docs)] #![warn(missing_debug_implementations, rust_2018_idioms)] #![doc(test(attr(deny(rust_2018_idioms, warnings))))] #![doc(test(attr(allow(unused_extern_crates, unused_variables))))] +#[cfg(all( + feature = "tds80", + not(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + )) +))] +compile_error!("The `tds80` feature requires one of the TLS features to be enabled."); + #[cfg(feature = "bigdecimal")] pub(crate) extern crate bigdecimal_ as bigdecimal; @@ -257,6 +217,7 @@ pub(crate) extern crate bigdecimal_ as bigdecimal; mod macros; mod client; +mod command; mod from_sql; mod query; mod sql_read_bytes; @@ -269,17 +230,26 @@ mod tds; mod sql_browser; -pub use client::{AuthMethod, Client, Config}; +mod bulk_options; + +pub use bulk_options::{ColumnOrderHint, SortOrder, SqlBulkCopyOption, SqlBulkCopyOptions}; +pub use client::{AuthMethod, Client, Config, ConfigBuilder}; +pub use command::{Command, SqlTableData, SqlTableDataRow, TableValue, TableValueRow}; pub(crate) use error::Error; pub use from_sql::{FromSql, FromSqlOwned}; pub use query::Query; pub use result::*; -pub use row::{Column, ColumnType, Row}; +pub use row::{Column, ColumnType, QueryIdx, Row, RowBuilder}; pub use sql_browser::SqlBrowser; pub use tds::{ - codec::{BulkLoadRequest, ColumnData, ColumnFlag, IntoRow, TokenRow, TypeLength}, + codec::{ + AltMetaDataColumn, BaseMetaDataColumn, BulkLoadRequest, ColumnData, ColumnFlag, + FixedLenType, IntoRow, IsolationLevel, MetaDataColumn, TokenAltMetaData, TokenAltRow, + TokenRow, TypeInfo, TypeLength, VarLenContext, VarLenType, + }, + collation::Collation, numeric, - stream::QueryStream, + stream::{CommandReturnValue, CommandStream, QueryStream}, time, xml, EncryptionLevel, }; pub use to_sql::{IntoSql, ToSql}; @@ -292,11 +262,61 @@ use tds::codec::*; pub type Result = std::result::Result; pub(crate) fn get_driver_version() -> u64 { - env!("CARGO_PKG_VERSION") + encode_driver_version(env!("CARGO_PKG_VERSION")) +} + +/// Packs a dotted version string into the little-endian byte layout the TDS +/// login record expects: the first component in the low byte, the next in bits +/// 8..16, and so on (up to six components). Non-numeric components contribute +/// zero. +fn encode_driver_version(version: &str) -> u64 { + version .splitn(6, '.') .enumerate() .fold(0u64, |acc, part| match part.1.parse::() { Ok(num) => acc | num << (part.0 * 8), - _ => acc | 0 << (part.0 * 8), + // A non-numeric component contributes nothing. + _ => acc, }) } + +#[cfg(test)] +mod driver_version_tests { + use super::encode_driver_version; + + #[test] + fn packs_each_component_into_its_own_byte() { + // Each component occupies its own byte, low component first. + assert_eq!(encode_driver_version("1.2.3"), 0x03_02_01); + assert_eq!(encode_driver_version("4.5.6.7"), 0x07_06_05_04); + } + + #[test] + fn shift_moves_components_left_not_right() { + // The minor version is shifted up by 8 bits, not down. + assert_eq!(encode_driver_version("34.17"), 34 | (17 << 8)); + assert_ne!(encode_driver_version("34.17"), 34); + } + + #[test] + fn components_are_combined_with_or_not_xor() { + // Overlapping bits are combined with OR, not XOR. + assert_eq!(encode_driver_version("257.1"), 0x101); + } + + #[test] + fn non_numeric_components_contribute_zero() { + assert_eq!(encode_driver_version("1.beta.3"), 1 | (3 << 16)); + assert_eq!(encode_driver_version("notaversion"), 0); + } + + #[test] + fn get_driver_version_encodes_the_crate_version() { + // The wrapper encodes the crate's own version and is non-zero. + assert_eq!( + super::get_driver_version(), + encode_driver_version(env!("CARGO_PKG_VERSION")) + ); + assert_ne!(super::get_driver_version(), 0); + } +} diff --git a/src/macros.rs b/src/macros.rs index 35f24228f..cbe0453e1 100644 --- a/src/macros.rs +++ b/src/macros.rs @@ -19,7 +19,10 @@ macro_rules! uint_enum { type Error = (); fn try_from(n: u8) -> ::std::result::Result<$ty, ()> { match n { - $( x if x == $ty::$variant as u8 => Ok($ty::$variant), )* + // Generic macro codegen: the `as u8` cast is compared against a `u8` + // input, so wider enum variants can never match here (they fall through + // to `Err`). The truncation is intentional and harmless. + $( #[allow(clippy::cast_enum_truncation)] x if x == $ty::$variant as u8 => Ok($ty::$variant), )* _ => Err(()), } } diff --git a/src/query.rs b/src/query.rs index 86e949996..352168747 100644 --- a/src/query.rs +++ b/src/query.rs @@ -35,6 +35,97 @@ impl<'a> Query<'a> { self.params.push(param.into_sql()); } + /// Bind every item of an iterator, in order. + /// + /// Equivalent to calling [`bind`] once per item. Pairs with + /// [`placeholders`] to build an `IN` list, where the number of + /// parameters is only known at runtime. + /// + /// # Example + /// + /// ``` + /// # use tiberius::Query; + /// let ids = vec![1i32, 2, 3]; + /// + /// let sql = format!( + /// "SELECT name FROM users WHERE id IN ({})", + /// Query::placeholders(1, ids.len()), + /// ); + /// + /// let mut query = Query::new(sql); + /// query.bind_iter(ids); + /// + /// assert_eq!(query.param_count(), 3); + /// ``` + /// + /// [`bind`]: #method.bind + /// [`placeholders`]: #method.placeholders + pub fn bind_iter(&mut self, params: impl IntoIterator + 'a>) { + for param in params { + self.bind(param); + } + } + + /// How many parameters have been bound so far. + /// + /// Useful for checking against [`MAX_PARAMETERS`] before executing a + /// statement whose parameter count is decided at runtime. + /// + /// [`MAX_PARAMETERS`]: #associatedconstant.MAX_PARAMETERS + pub fn param_count(&self) -> usize { + self.params.len() + } + + /// The 2100-parameter server-side limit for one statement; split larger + /// batches into chunks of at most `MAX_PARAMETERS / parameters_per_row` rows. + /// + /// # Example + /// + /// ``` + /// # use tiberius::Query; + /// // A three-column INSERT: three parameters per row. + /// let rows_per_statement = Query::MAX_PARAMETERS / 3; + /// assert_eq!(rows_per_statement, 700); + /// ``` + pub const MAX_PARAMETERS: usize = 2100; + + /// Builds `@P1, @P2, …` for `count` placeholders numbered from `first` + /// (1-based). + /// + /// # Example + /// + /// ``` + /// # use tiberius::Query; + /// assert_eq!(Query::placeholders(1, 3), "@P1, @P2, @P3"); + /// + /// // Continuing after parameters that are already bound. + /// assert_eq!(Query::placeholders(4, 2), "@P4, @P5"); + /// ``` + /// + /// A count of zero yields an empty string. `IN ()` is a syntax error, so + /// a caller with nothing to match on should skip the query rather than + /// build one: + /// + /// ``` + /// # use tiberius::Query; + /// let ids: Vec = Vec::new(); + /// assert!(Query::placeholders(1, ids.len()).is_empty()); + /// ``` + pub fn placeholders(first: usize, count: usize) -> String { + use std::fmt::Write; + + let mut out = String::with_capacity(count * 6); + + for index in 0..count { + if index > 0 { + out.push_str(", "); + } + let _ = write!(out, "@P{}", first + index); + } + + out + } + /// Executes SQL statements in the SQL Server, returning the number rows /// affected. Useful for `INSERT`, `UPDATE` and `DELETE` statements. See /// [`Client#execute`] for a simpler API if the parameters are statically @@ -69,7 +160,7 @@ impl<'a> Query<'a> { /// [`ToSql`]: trait.ToSql.html /// [`FromSql`]: trait.FromSql.html /// [`Client#execute`]: struct.Client.html#method.execute - pub async fn execute<'b, S>(self, client: &'b mut Client) -> crate::Result + pub async fn execute(self, client: &mut Client) -> crate::Result where S: AsyncRead + AsyncWrite + Unpin + Send, { @@ -136,3 +227,69 @@ impl<'a> Query<'a> { Ok(result) } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn placeholders_are_numbered_from_one() { + assert_eq!(Query::placeholders(1, 1), "@P1"); + assert_eq!(Query::placeholders(1, 3), "@P1, @P2, @P3"); + } + + #[test] + fn placeholders_can_continue_from_an_offset() { + assert_eq!(Query::placeholders(4, 2), "@P4, @P5"); + assert_eq!(Query::placeholders(10, 1), "@P10"); + } + + #[test] + fn no_placeholders_is_an_empty_string() { + assert_eq!(Query::placeholders(1, 0), ""); + assert_eq!(Query::placeholders(7, 0), ""); + } + + #[test] + fn placeholders_have_no_trailing_separator() { + let list = Query::placeholders(1, 5); + assert!(!list.ends_with(", ")); + assert_eq!(list.matches(',').count(), 4); + } + + #[test] + fn binding_an_iterator_counts_every_item() { + let mut query = Query::new("SELECT 1"); + assert_eq!(query.param_count(), 0); + + query.bind_iter(vec![1i32, 2, 3]); + assert_eq!(query.param_count(), 3); + + query.bind(4i32); + assert_eq!(query.param_count(), 4); + } + + #[test] + fn binding_an_empty_iterator_binds_nothing() { + let mut query = Query::new("SELECT 1"); + query.bind_iter(Vec::::new()); + assert_eq!(query.param_count(), 0); + } + + #[test] + fn a_generated_list_matches_the_number_of_bound_parameters() { + let ids = vec![10i32, 20, 30, 40]; + let list = Query::placeholders(1, ids.len()); + + let mut query = Query::new(format!("SELECT * FROM t WHERE id IN ({list})")); + query.bind_iter(ids); + + assert_eq!(list.matches("@P").count(), query.param_count()); + } + + #[test] + fn the_parameter_limit_is_the_documented_tds_maximum() { + assert_eq!(Query::MAX_PARAMETERS, 2100); + assert_eq!(Query::MAX_PARAMETERS / 3, 700); + } +} diff --git a/src/result.rs b/src/result.rs index 19ba6faf9..d4bafe64f 100644 --- a/src/result.rs +++ b/src/result.rs @@ -1,7 +1,9 @@ -pub use crate::tds::stream::{QueryItem, ResultMetadata}; +pub use crate::tds::stream::{CommandItem, QueryItem, ResultMetadata}; use crate::{ client::Connection, - tds::stream::{ReceivedToken, TokenStream}, + error::Error, + tds::stream::{CommandReturnValue, ReceivedToken, TokenStream}, + FromSql, Row, }; use futures_util::io::{AsyncRead, AsyncWrite}; use futures_util::stream::TryStreamExt; @@ -113,3 +115,176 @@ impl IntoIterator for ExecuteResult { self.rows_affected.into_iter() } } + +/// A materialized result from executing a [`Command`], carrying the number of +/// affected rows, the return code, the values of any OUT parameters and any +/// record sets returned by the command. +/// +/// [`Command`]: crate::Command +/// +/// # Example +/// +/// ```no_run +/// # use tiberius::{Config, Command}; +/// # use tokio_util::compat::TokioAsyncWriteCompatExt; +/// # use std::env; +/// # #[tokio::main] +/// # async fn main() -> Result<(), Box> { +/// # let c_str = env::var("TIBERIUS_TEST_CONNECTION_STRING").unwrap_or( +/// # "server=tcp:localhost,1433;integratedSecurity=true;TrustServerCertificate=true".to_owned(), +/// # ); +/// # let config = Config::from_ado_string(&c_str)?; +/// # let tcp = tokio::net::TcpStream::connect(config.get_addr()).await?; +/// # tcp.set_nodelay(true)?; +/// # let mut client = tiberius::Client::connect(config, tcp.compat_write()).await?; +/// let mut cmd = Command::new("dbo.usp_SomeStoredProc"); +/// +/// cmd.bind_param("@foo", 34i32); +/// cmd.bind_out_param("@bar", "bar"); +/// let res = cmd.exec(&mut client).await?.into_command_result().await?; +/// +/// let rv: Option<&str> = res.try_return_value("@bar")?; +/// let rc = res.return_code(); +/// let ra = res.rows_affected(); +/// +/// let rs0 = res.to_query_result(0); +/// # Ok(()) +/// # } +/// ``` +/// +#[derive(Debug)] +pub struct CommandResult { + pub(crate) rows_affected: Vec, + pub(crate) return_code: u32, + pub(crate) return_values: Vec, + pub(crate) query_results: Vec>, +} + +impl<'a> CommandResult { + /// A slice of the numbers of rows affected, in the same order as the + /// statements ran by the command. + pub fn rows_affected(&self) -> &[u64] { + self.rows_affected.as_slice() + } + + /// The return code of the command, as returned by the server. + pub fn return_code(&self) -> u32 { + self.return_code + } + + /// The number of returned values (OUT parameters) available. + pub fn return_values_len(&self) -> usize { + self.return_values.len() + } + + /// Gets a returned value by its OUT parameter name, converting it to `T`. + /// Returns `None` if the value is `NULL`, and an error if no OUT parameter + /// with the given name was returned. + pub fn try_return_value(&'a self, name: &str) -> crate::Result> + where + T: FromSql<'a>, + { + let col_data = self + .return_values + .iter() + .find(|p| p.name.eq(name)) + .ok_or_else(|| { + Error::Conversion(format!("Could not find return value {}", name).into()) + })?; + + T::from_sql(&col_data.data) + } + + /// Gets a returned record set by its zero-based index. Returns `None` if the + /// index is out of range. + pub fn to_query_result(&self, idx: usize) -> Option<&Vec> { + self.query_results.get(idx) + } +} + +impl IntoIterator for CommandResult { + type Item = Vec; + type IntoIter = std::vec::IntoIter; + + fn into_iter(self) -> Self::IntoIter { + self.query_results.into_iter() + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::tds::codec::ColumnData; + + impl ExecuteResult { + fn from_counts(counts: Vec) -> Self { + Self { + rows_affected: counts, + } + } + } + + #[test] + fn execute_result_rows_affected_preserves_order_and_values() { + let res = ExecuteResult::from_counts(vec![3, 0, 7]); + assert_eq!(res.rows_affected(), &[3, 0, 7]); + } + + #[test] + fn execute_result_total_sums_every_count() { + assert_eq!(ExecuteResult::from_counts(vec![3, 0, 7]).total(), 10); + } + + #[test] + fn execute_result_into_iter_yields_each_count() { + let counts: Vec = ExecuteResult::from_counts(vec![5, 9]).into_iter().collect(); + assert_eq!(counts, vec![5, 9]); + } + + fn return_value(name: &str, value: i32) -> CommandReturnValue { + CommandReturnValue { + name: name.to_string(), + ord: 0, + data: ColumnData::I32(Some(value)), + } + } + + fn command_result() -> CommandResult { + CommandResult { + rows_affected: vec![2, 4], + return_code: 7, + return_values: vec![return_value("@a", 1), return_value("@b", 42)], + // Two (empty) record sets so `to_query_result` has Some values to return. + query_results: vec![Vec::new(), Vec::new()], + } + } + + #[test] + fn command_result_scalar_accessors() { + let res = command_result(); + assert_eq!(res.rows_affected(), &[2, 4]); + assert_eq!(res.return_code(), 7); + assert_eq!(res.return_values_len(), 2); + } + + #[test] + fn command_result_to_query_result_indexes_record_sets() { + let res = command_result(); + assert!(res.to_query_result(0).is_some()); + assert!(res.to_query_result(1).is_some()); + assert!(res.to_query_result(2).is_none()); + } + + #[test] + fn command_result_try_return_value_reads_named_out_param() { + let res = command_result(); + let got: Option = res.try_return_value("@b").unwrap(); + assert_eq!(got, Some(42)); + assert!(res.try_return_value::("@missing").is_err()); + } + + #[test] + fn command_result_into_iter_yields_each_record_set() { + assert_eq!(command_result().into_iter().count(), 2); + } +} diff --git a/src/row.rs b/src/row.rs index 5441be700..64a6db8bb 100644 --- a/src/row.rs +++ b/src/row.rs @@ -7,6 +7,7 @@ use std::{fmt::Display, sync::Arc}; /// A column of data from a query. #[derive(Debug, Clone)] +#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] pub struct Column { pub(crate) name: String, pub(crate) column_type: ColumnType, @@ -30,6 +31,7 @@ impl Column { } #[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] /// The type of the column. pub enum ColumnType { /// The column doesn't have a specified type. @@ -192,6 +194,7 @@ impl From<&TypeInfo> for ColumnType { VarLenType::SSVariant => Self::SSVariant, }, TypeInfo::Xml { .. } => Self::Xml, + TypeInfo::Udt(_) => Self::Udt, } } } @@ -246,28 +249,44 @@ impl From<&TypeInfo> for ColumnType { /// [`try_get`]: #method.try_get /// [`IntoIterator`]: #impl-IntoIterator #[derive(Debug)] +#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] pub struct Row { pub(crate) columns: Arc>, pub(crate) data: TokenRow<'static>, pub(crate) result_index: usize, } +/// A type that can address a column within a [`Row`], either by its zero-based +/// position (`usize`) or by name (`&str`). +/// +/// Implement this for a custom column identifier (for example a generated +/// column-name enum) to index rows with it via [`Row::get`]/[`Row::try_get`]. pub trait QueryIdx where Self: Display, { + /// Resolves this index to the column's zero-based position in `row`, or + /// `None` if it does not name/point to a column in the row. fn idx(&self, row: &Row) -> Option; } impl QueryIdx for usize { - fn idx(&self, _row: &Row) -> Option { - Some(*self) + fn idx(&self, row: &Row) -> Option { + (*self < row.columns.len()).then_some(*self) } } impl QueryIdx for &str { fn idx(&self, row: &Row) -> Option { - row.columns.iter().position(|c| c.name() == *self) + // Prefer an exact column-name match so a column literally named `r#...` + // (or `type`) resolves to itself. + if let Some(p) = row.columns.iter().position(|c| c.name() == *self) { + return Some(p); + } + // Fallback: allow a Rust raw identifier (`r#type`) to match the plain SQL + // name (`type`) when no exact column exists. + self.strip_prefix("r#") + .and_then(|n| row.columns.iter().position(|c| c.name() == n)) } } @@ -399,19 +418,148 @@ impl Row { } /// Retrieve a column's value for a given column index. + /// + /// # Errors + /// + /// Returns an error if the given index does not name a column in the row, or + /// if the stored value cannot be converted into the requested Rust type `R`. #[track_caller] pub fn try_get<'a, R, I>(&'a self, idx: I) -> crate::Result> where R: FromSql<'a>, I: QueryIdx, + { + let data = self.get_column_data(idx)?; + + R::from_sql(data) + } + + /// Retrieve a column's data for a given column index. + /// + /// # Errors + /// + /// Returns an error if the given index does not name a column in the row, or + /// if the row carries fewer cells than columns (a malformed `ROW`/`NBCROW`). + #[track_caller] + pub fn get_column_data(&self, idx: I) -> crate::Result<&ColumnData<'static>> + where + I: QueryIdx, { let idx = idx.idx(self).ok_or_else(|| { Error::Conversion(format!("Could not find column with index {}", idx).into()) })?; - let data = self.data.get(idx).unwrap(); + // `idx` was validated against the column metadata; the cell should exist, + // but a malformed ROW/NBCROW with fewer cells than columns must not + // panic here — return an error instead of unwrapping. + self.data.get(idx).ok_or_else(|| { + Error::Protocol(format!("row has no data for column index {idx}").into()) + }) + } - R::from_sql(data) + /// Consumes the row, returning the underlying [`TokenRow`] holding the raw + /// column data as received from the server. + /// + /// This is useful when direct access to the raw [`ColumnData`] values is + /// needed instead of converting them through [`get`] or [`try_get`]. + /// + /// [`get`]: #method.get + /// [`try_get`]: #method.try_get + pub fn into_token_row(self) -> TokenRow<'static> { + self.data + } + + /// Creates a [`RowBuilder`] for assembling a `Row` outside of a query. + /// + /// A `Row` returned from a query is normally built by the driver from data + /// received over the wire, and its fields are not otherwise accessible. This + /// entry point exists so callers can construct a `Row` directly — primarily + /// to unit-test functions that take a [`Row`] or `&[Row]` without needing a + /// live database connection. + /// + /// See [`RowBuilder`] for a full example. + pub fn builder() -> RowBuilder { + RowBuilder::default() + } +} + +/// A builder for constructing a [`Row`] outside of a query. +/// +/// This is primarily intended for unit tests that exercise functions taking a +/// [`Row`] or `&[Row]` without a live database connection. Obtain one via +/// [`Row::builder`]. +/// +/// Each [`column`] call appends a column's metadata together with its value, so +/// the column count and the value count can never drift out of sync. The row is +/// finished with [`build`]. +/// +/// # Example +/// +/// ``` +/// use tiberius::{ColumnData, ColumnType, Row}; +/// +/// let row = Row::builder() +/// .column("id", ColumnType::Int4, ColumnData::I32(Some(1))) +/// .column( +/// "name", +/// ColumnType::NVarchar, +/// ColumnData::String(Some("Alice".into())), +/// ) +/// .column("age", ColumnType::Int4, ColumnData::I32(None)) +/// .build(); +/// +/// assert_eq!(Some(1), row.get::("id")); +/// assert_eq!(Some("Alice"), row.get::<&str, _>(1)); +/// assert_eq!(None, row.get::("age")); +/// ``` +/// +/// [`column`]: RowBuilder::column +/// [`build`]: RowBuilder::build +#[derive(Debug, Default)] +pub struct RowBuilder { + columns: Vec, + data: TokenRow<'static>, + result_index: usize, +} + +impl RowBuilder { + /// Appends a column and its value to the row being built. + /// + /// `name` is the column name used for by-name lookups, `column_type` is the + /// [`ColumnType`] reported by [`Column::column_type`], and `value` holds the + /// cell's [`ColumnData`]. Columns keep the order in which they are added. + /// + /// `column_type` is stored as-is and is *not* checked against `value`: the + /// value accessors ([`Row::get`]/[`Row::try_get`]) dispatch on the + /// [`ColumnData`] variant alone, so a mismatched `column_type` only affects + /// what [`Column::column_type`] reports. Keeping the two consistent is the + /// caller's responsibility. + pub fn column( + mut self, + name: impl Into, + column_type: ColumnType, + value: ColumnData<'static>, + ) -> Self { + self.columns.push(Column::new(name.into(), column_type)); + self.data.push(value); + self + } + + /// Sets the result-set index reported by [`Row::result_index`]. + /// + /// Defaults to `0` when not called. + pub fn result_index(mut self, result_index: usize) -> Self { + self.result_index = result_index; + self + } + + /// Consumes the builder and returns the assembled [`Row`]. + pub fn build(self) -> Row { + Row { + columns: Arc::new(self.columns), + data: self.data, + result_index: self.result_index, + } } } @@ -423,3 +571,846 @@ impl IntoIterator for Row { self.data.into_iter() } } + +#[cfg(test)] +mod tests { + use super::*; + use std::borrow::Cow; + + fn make_row() -> Row { + let columns = Arc::new(vec![ + Column::new("foo".to_string(), ColumnType::Int4), + Column::new("type".to_string(), ColumnType::Int4), + ]); + + let mut data = TokenRow::new(); + data.push(ColumnData::I32(Some(1))); + data.push(ColumnData::I32(Some(2))); + + Row { + columns, + data, + result_index: 0, + } + } + + #[test] + fn result_index_reflects_the_field() { + // `result_index()` returns the row's stored result index. + let columns = Arc::new(vec![Column::new("c".to_string(), ColumnType::Int4)]); + let mut data = TokenRow::new(); + data.push(ColumnData::I32(Some(1))); + let row = Row { + columns, + data, + result_index: 3, + }; + assert_eq!(row.result_index(), 3); + } + + // Regression test for #211: an out-of-range usize index must not panic. + #[test] + fn try_get_out_of_range_index_returns_none() { + let row = make_row(); + + assert_eq!(None, 2usize.idx(&row)); + assert_eq!(Some(0), 0usize.idx(&row)); + + let value: crate::Result> = row.try_get(5usize); + assert!(value.is_err()); + } + + // Regression test for #382: a raw-identifier column name (`r#type`) must + // match the plain SQL column name (`type`). + #[test] + fn raw_identifier_column_name_matches() { + let row = make_row(); + + assert_eq!(Some(1), "type".idx(&row)); + assert_eq!(Some(1), "r#type".idx(&row)); + assert_eq!(Some(0), "r#foo".idx(&row)); + assert_eq!(None, "r#missing".idx(&row)); + + assert_eq!(Some(2i32), row.get::("r#type")); + } + + // A column literally named `r#type` alongside a `type` column must each + // resolve to themselves: the exact match wins before the r# fallback. + #[test] + fn literal_raw_prefixed_column_wins_over_fallback() { + let columns = Arc::new(vec![ + Column::new("type".to_string(), ColumnType::Int4), + Column::new("r#type".to_string(), ColumnType::Int4), + ]); + + let mut data = TokenRow::new(); + data.push(ColumnData::I32(Some(1))); + data.push(ColumnData::I32(Some(2))); + + let row = Row { + columns, + data, + result_index: 0, + }; + + // Exact match: `type` -> the "type" column (index 0). + assert_eq!(Some(0), "type".idx(&row)); + // Exact match: `r#type` -> the literal "r#type" column (index 1), + // NOT the "type" column via the strip fallback. + assert_eq!(Some(1), "r#type".idx(&row)); + } + + // A column literally named `r#foo` (with no plain `foo` column) resolves to + // itself via the exact match; the fallback is never needed. + #[test] + fn literal_raw_prefixed_column_exact_match() { + let columns = Arc::new(vec![Column::new("r#foo".to_string(), ColumnType::Int4)]); + + let mut data = TokenRow::new(); + data.push(ColumnData::I32(Some(1))); + + let row = Row { + columns, + data, + result_index: 0, + }; + + assert_eq!(Some(0), "r#foo".idx(&row)); + // No plain "foo" column exists, so the fallback finds nothing. + assert_eq!(None, "foo".idx(&row)); + } + + #[test] + fn row_accessors() { + let row = make_row(); + + assert_eq!(2, row.columns().len()); + assert_eq!(2, row.len()); + assert_eq!(0, row.result_index()); + + let cells: Vec<_> = row.cells().collect(); + assert_eq!(2, cells.len()); + assert_eq!("foo", cells[0].0.name()); + + let token_row = row.into_token_row(); + assert_eq!(2, token_row.len()); + } + + #[test] + fn row_into_iterator_yields_column_data() { + let row = make_row(); + let values: Vec<_> = row.into_iter().collect(); + assert_eq!( + vec![ColumnData::I32(Some(1)), ColumnData::I32(Some(2))], + values + ); + } + + #[test] + fn get_column_data_missing_cell_errors() { + let columns = Arc::new(vec![ + Column::new("a".to_string(), ColumnType::Int4), + Column::new("b".to_string(), ColumnType::Int4), + ]); + + // Malformed row: metadata says 2 columns, but only 1 cell present. + let mut data = TokenRow::new(); + data.push(ColumnData::I32(Some(1))); + + let row = Row { + columns, + data, + result_index: 0, + }; + + let err = row.get_column_data(1usize).unwrap_err(); + assert!(format!("{}", err).contains("row has no data for column index")); + } + + // A `Row` can be built via the public builder and its values + // round-trip through the public accessors by index and by name, including + // nulls and mixed types. + #[test] + fn builder_round_trips_through_public_accessors() { + let row = Row::builder() + .column("id", ColumnType::Int4, ColumnData::I32(Some(1))) + .column( + "name", + ColumnType::NVarchar, + ColumnData::String(Some("Alice".into())), + ) + .column("flag", ColumnType::Bit, ColumnData::Bit(Some(true))) + .column("age", ColumnType::Int4, ColumnData::I32(None)) + .build(); + + // Shape. + assert_eq!(4, row.len()); + assert_eq!(4, row.columns().len()); + assert_eq!(0, row.result_index()); + assert_eq!("id", row.columns()[0].name()); + assert_eq!(ColumnType::NVarchar, row.columns()[1].column_type()); + + // By index. + assert_eq!(Some(1i32), row.get::(0)); + assert_eq!(Some("Alice"), row.get::<&str, _>(1)); + assert_eq!(Some(true), row.get::(2)); + + // By name. + assert_eq!(Some(1i32), row.get::("id")); + assert_eq!(Some("Alice"), row.get::<&str, _>("name")); + assert_eq!(Some(true), row.get::("flag")); + + // Nulls round-trip as `None` for both index and name. + assert_eq!(None, row.get::(3)); + assert_eq!(None, row.get::("age")); + + // try_get mirrors get and errors on an unknown column. + assert_eq!(Some(1i32), row.try_get::("id").unwrap()); + assert!(row.try_get::("nope").is_err()); + assert!(row.try_get::(9usize).is_err()); + } + + #[test] + fn builder_result_index_and_empty_row() { + let row = Row::builder().result_index(7).build(); + assert_eq!(7, row.result_index()); + assert_eq!(0, row.len()); + assert!(row.columns().is_empty()); + } + + #[test] + fn builder_empty_row_has_no_cells_or_columns() { + let row = Row::builder().build(); + assert_eq!(0, row.len()); + assert_eq!(0, row.result_index()); + assert!(row.columns().is_empty()); + assert_eq!(0, row.cells().count()); + // Any lookup on an empty row is a miss. + assert_eq!(None, 0usize.idx(&row)); + assert_eq!(None, "anything".idx(&row)); + assert!(row.try_get::(0usize).is_err()); + assert_eq!(0, row.into_iter().count()); + } + + #[test] + fn builder_single_column() { + let row = Row::builder() + .column("only", ColumnType::Int4, ColumnData::I32(Some(99))) + .build(); + + assert_eq!(1, row.len()); + assert_eq!(1, row.columns().len()); + assert_eq!("only", row.columns()[0].name()); + assert_eq!(Some(99i32), row.get::(0)); + assert_eq!(Some(99i32), row.get::("only")); + } + + #[test] + fn builder_many_columns_by_index_and_name() { + let mut builder = Row::builder(); + for i in 0..64i32 { + builder = builder.column(format!("c{i}"), ColumnType::Int4, ColumnData::I32(Some(i))); + } + let row = builder.build(); + + assert_eq!(64, row.len()); + assert_eq!(64, row.columns().len()); + for i in 0..64usize { + assert_eq!(Some(i as i32), row.get::(i)); + assert_eq!(Some(i as i32), row.get::(format!("c{i}").as_str())); + } + } + + #[test] + fn builder_int_variants_round_trip() { + let row = Row::builder() + .column("u8", ColumnType::Int1, ColumnData::U8(Some(200))) + .column("i16", ColumnType::Int2, ColumnData::I16(Some(-30000))) + .column("i32", ColumnType::Int4, ColumnData::I32(Some(123456))) + .column( + "i64", + ColumnType::Int8, + ColumnData::I64(Some(9_000_000_000)), + ) + .build(); + + assert_eq!(Some(200u8), row.get::("u8")); + assert_eq!(Some(200u8), row.get::(0)); + assert_eq!(Some(-30000i16), row.get::("i16")); + assert_eq!(Some(123456i32), row.get::("i32")); + assert_eq!(Some(9_000_000_000i64), row.get::("i64")); + } + + #[test] + fn builder_int_nulls_round_trip_as_none() { + let row = Row::builder() + .column("u8", ColumnType::Int1, ColumnData::U8(None)) + .column("i16", ColumnType::Int2, ColumnData::I16(None)) + .column("i32", ColumnType::Int4, ColumnData::I32(None)) + .column("i64", ColumnType::Int8, ColumnData::I64(None)) + .build(); + + assert_eq!(None, row.get::("u8")); + assert_eq!(None, row.get::("i16")); + assert_eq!(None, row.get::("i32")); + assert_eq!(None, row.get::("i64")); + } + + #[test] + fn builder_float_variants_round_trip() { + let row = Row::builder() + .column("f32", ColumnType::Float4, ColumnData::F32(Some(1.5))) + .column("f64", ColumnType::Float8, ColumnData::F64(Some(2.25))) + .column("f32_null", ColumnType::Float4, ColumnData::F32(None)) + .column("f64_null", ColumnType::Float8, ColumnData::F64(None)) + .build(); + + assert_eq!(Some(1.5f32), row.get::("f32")); + assert_eq!(Some(2.25f64), row.get::("f64")); + assert_eq!(None, row.get::("f32_null")); + assert_eq!(None, row.get::("f64_null")); + } + + #[test] + fn builder_bool_round_trip() { + let row = Row::builder() + .column("t", ColumnType::Bit, ColumnData::Bit(Some(true))) + .column("f", ColumnType::Bit, ColumnData::Bit(Some(false))) + .column("n", ColumnType::Bitn, ColumnData::Bit(None)) + .build(); + + assert_eq!(Some(true), row.get::("t")); + assert_eq!(Some(false), row.get::("f")); + assert_eq!(Some(true), row.get::(0)); + assert_eq!(None, row.get::("n")); + } + + #[test] + fn builder_string_owned_and_borrowed() { + let row = Row::builder() + .column( + "owned", + ColumnType::NVarchar, + ColumnData::String(Some(Cow::Owned("owned".to_string()))), + ) + .column( + "borrowed", + ColumnType::NVarchar, + ColumnData::String(Some(Cow::Borrowed("borrowed"))), + ) + .column("null", ColumnType::NVarchar, ColumnData::String(None)) + .build(); + + assert_eq!(Some("owned"), row.get::<&str, _>("owned")); + assert_eq!(Some("borrowed"), row.get::<&str, _>(1)); + // Owned `String` conversion goes through the by-value accessor. + use crate::FromSqlOwned; + let owned_cell = row.get_column_data("owned").unwrap().clone(); + assert_eq!( + Some("owned".to_string()), + String::from_sql_owned(owned_cell).unwrap() + ); + assert_eq!(None, row.get::<&str, _>("null")); + } + + #[test] + fn builder_binary_round_trip() { + let payload = vec![0u8, 1, 2, 3, 255]; + let row = Row::builder() + .column( + "bin", + ColumnType::BigVarBin, + ColumnData::Binary(Some(Cow::Owned(payload.clone()))), + ) + .column("null", ColumnType::BigVarBin, ColumnData::Binary(None)) + .build(); + + assert_eq!(Some(payload.as_slice()), row.get::<&[u8], _>("bin")); + assert_eq!(Some(payload.as_slice()), row.get::<&[u8], _>(0)); + assert_eq!(None, row.get::<&[u8], _>("null")); + } + + #[test] + fn builder_guid_round_trip() { + let uuid = crate::Uuid::from_u128(0x0123_4567_89ab_cdef_0123_4567_89ab_cdef); + let row = Row::builder() + .column("id", ColumnType::Guid, ColumnData::Guid(Some(uuid))) + .column("null", ColumnType::Guid, ColumnData::Guid(None)) + .build(); + + assert_eq!(Some(uuid), row.get::("id")); + assert_eq!(Some(uuid), row.get::(0)); + assert_eq!(None, row.get::("null")); + } + + #[test] + fn builder_numeric_round_trip() { + let numeric = crate::tds::Numeric::new_with_scale(1234, 2); + let row = Row::builder() + .column( + "n", + ColumnType::Numericn, + ColumnData::Numeric(Some(numeric)), + ) + .column("null", ColumnType::Numericn, ColumnData::Numeric(None)) + .build(); + + assert_eq!(Some(numeric), row.get::("n")); + assert_eq!(Some(numeric), row.get::(0)); + assert_eq!(None, row.get::("null")); + } + + #[test] + fn builder_datetime_variants_round_trip_raw_cells() { + // DateTime/SmallDateTime have no infallible Rust accessor without a + // date feature, so assert the raw `ColumnData` is preserved verbatim. + let dt = crate::time::DateTime::new(200, 3000); + let sdt = crate::time::SmallDateTime::new(100, 200); + + let row = Row::builder() + .column("dt", ColumnType::Datetime, ColumnData::DateTime(Some(dt))) + .column( + "sdt", + ColumnType::Datetime4, + ColumnData::SmallDateTime(Some(sdt)), + ) + .column("dt_null", ColumnType::Datetimen, ColumnData::DateTime(None)) + .build(); + + let cells: Vec<_> = row.into_iter().collect(); + assert_eq!(ColumnData::DateTime(Some(dt)), cells[0]); + assert_eq!(ColumnData::SmallDateTime(Some(sdt)), cells[1]); + assert_eq!(ColumnData::DateTime(None), cells[2]); + } + + #[cfg(feature = "tds73")] + #[test] + fn builder_tds73_date_variants_round_trip_raw_cells() { + let date = crate::time::Date::new(730119); + let time = crate::time::Time::new(1234, 5); + let dt2 = crate::time::DateTime2::new(date, time); + let dto = crate::time::DateTimeOffset::new(dt2, 60); + + let row = Row::builder() + .column("date", ColumnType::Daten, ColumnData::Date(Some(date))) + .column("time", ColumnType::Timen, ColumnData::Time(Some(time))) + .column( + "dt2", + ColumnType::Datetime2, + ColumnData::DateTime2(Some(dt2)), + ) + .column( + "dto", + ColumnType::DatetimeOffsetn, + ColumnData::DateTimeOffset(Some(dto)), + ) + .column("date_null", ColumnType::Daten, ColumnData::Date(None)) + .build(); + + let cells: Vec<_> = row.cells().map(|(_, d)| d.clone()).collect(); + assert_eq!(ColumnData::Date(Some(date)), cells[0]); + assert_eq!(ColumnData::Time(Some(time)), cells[1]); + assert_eq!(ColumnData::DateTime2(Some(dt2)), cells[2]); + assert_eq!(ColumnData::DateTimeOffset(Some(dto)), cells[3]); + assert_eq!(ColumnData::Date(None), cells[4]); + } + + #[test] + fn builder_mixed_types_and_cells_iteration() { + let row = Row::builder() + .column("i", ColumnType::Int4, ColumnData::I32(Some(7))) + .column( + "s", + ColumnType::NVarchar, + ColumnData::String(Some(Cow::Borrowed("mix"))), + ) + .column("b", ColumnType::Bit, ColumnData::Bit(Some(false))) + .build(); + + let cells: Vec<_> = row.cells().collect(); + assert_eq!(3, cells.len()); + assert_eq!("i", cells[0].0.name()); + assert_eq!(ColumnType::NVarchar, cells[1].0.column_type()); + assert_eq!(&ColumnData::Bit(Some(false)), cells[2].1); + } + + #[test] + fn builder_column_metadata_preserved() { + let row = Row::builder() + .column("a", ColumnType::Int8, ColumnData::I64(Some(1))) + .column("b", ColumnType::NVarchar, ColumnData::String(None)) + .column("c", ColumnType::Bit, ColumnData::Bit(Some(true))) + .build(); + + assert_eq!("a", row.columns()[0].name()); + assert_eq!(ColumnType::Int8, row.columns()[0].column_type()); + assert_eq!("b", row.columns()[1].name()); + assert_eq!(ColumnType::NVarchar, row.columns()[1].column_type()); + assert_eq!("c", row.columns()[2].name()); + assert_eq!(ColumnType::Bit, row.columns()[2].column_type()); + } + + #[test] + fn builder_duplicate_column_names_resolve_to_first_by_name() { + let row = Row::builder() + .column("dup", ColumnType::Int4, ColumnData::I32(Some(10))) + .column("dup", ColumnType::Int4, ColumnData::I32(Some(20))) + .build(); + + // By-name resolves to the first matching column... + assert_eq!(Some(0), "dup".idx(&row)); + assert_eq!(Some(10i32), row.get::("dup")); + // ...while both remain independently addressable by index. + assert_eq!(Some(10i32), row.get::(0)); + assert_eq!(Some(20i32), row.get::(1)); + assert_eq!(2, row.len()); + } + + #[test] + fn builder_unknown_name_and_out_of_range_index() { + let row = Row::builder() + .column("known", ColumnType::Int4, ColumnData::I32(Some(1))) + .build(); + + assert_eq!(None, "missing".idx(&row)); + assert_eq!(None, 1usize.idx(&row)); + + let by_name: crate::Result> = row.try_get("missing"); + assert!(by_name.is_err()); + assert!(format!("{}", by_name.unwrap_err()).contains("Could not find column")); + + let by_index: crate::Result> = row.try_get(5usize); + assert!(by_index.is_err()); + } + + #[test] + fn builder_wrong_type_conversion_errors() { + let row = Row::builder() + .column( + "s", + ColumnType::NVarchar, + ColumnData::String(Some("x".into())), + ) + .build(); + + // Requesting an incompatible Rust type surfaces a conversion error + // rather than panicking through `try_get`. + let wrong: crate::Result> = row.try_get("s"); + assert!(wrong.is_err()); + assert!(format!("{}", wrong.unwrap_err()).contains("cannot interpret")); + } + + #[test] + fn builder_by_index_and_by_name_agree() { + let row = Row::builder() + .column("x", ColumnType::Int4, ColumnData::I32(Some(11))) + .column("y", ColumnType::Int4, ColumnData::I32(Some(22))) + .build(); + + assert_eq!(row.get::(0), row.get::("x")); + assert_eq!(row.get::(1), row.get::("y")); + } + + #[test] + fn builder_result_index_variations() { + for idx in [0usize, 1, 5, 42, usize::MAX] { + let row = Row::builder() + .result_index(idx) + .column("c", ColumnType::Int4, ColumnData::I32(Some(0))) + .build(); + assert_eq!(idx, row.result_index()); + } + } + + #[test] + fn builder_raw_identifier_column_names() { + // A column named `type` is reachable via both `type` and `r#type`. + let row = Row::builder() + .column("foo", ColumnType::Int4, ColumnData::I32(Some(1))) + .column("type", ColumnType::Int4, ColumnData::I32(Some(2))) + .build(); + + assert_eq!(Some(1), "type".idx(&row)); + assert_eq!(Some(1), "r#type".idx(&row)); + assert_eq!(Some(0), "r#foo".idx(&row)); + assert_eq!(None, "r#missing".idx(&row)); + assert_eq!(Some(2i32), row.get::("r#type")); + } + + #[test] + fn builder_literal_raw_prefixed_column_wins_over_fallback() { + let row = Row::builder() + .column("type", ColumnType::Int4, ColumnData::I32(Some(1))) + .column("r#type", ColumnType::Int4, ColumnData::I32(Some(2))) + .build(); + + // Exact match wins before the `r#` strip fallback. + assert_eq!(Some(0), "type".idx(&row)); + assert_eq!(Some(1), "r#type".idx(&row)); + assert_eq!(Some(1i32), row.get::("type")); + assert_eq!(Some(2i32), row.get::("r#type")); + } + + #[test] + fn builder_into_token_row_preserves_data() { + let row = Row::builder() + .column("a", ColumnType::Int4, ColumnData::I32(Some(1))) + .column("b", ColumnType::Int4, ColumnData::I32(Some(2))) + .build(); + + let token_row = row.into_token_row(); + assert_eq!(2, token_row.len()); + assert_eq!(Some(&ColumnData::I32(Some(1))), token_row.get(0)); + assert_eq!(Some(&ColumnData::I32(Some(2))), token_row.get(1)); + } + + #[test] + fn builder_result_index_overwrites_and_is_order_independent() { + // Last write wins, and the setter is independent of column order: + // calling it twice and interleaving it with `column` still yields the + // final value with all columns intact. + let row = Row::builder() + .result_index(3) + .column("a", ColumnType::Int4, ColumnData::I32(Some(1))) + .result_index(9) + .column("b", ColumnType::Int4, ColumnData::I32(Some(2))) + .build(); + + assert_eq!(9, row.result_index()); + assert_eq!(2, row.len()); + assert_eq!(Some(1i32), row.get::("a")); + assert_eq!(Some(2i32), row.get::("b")); + } + + #[test] + fn builder_usize_max_index_lookup_is_a_miss() { + let row = Row::builder() + .column("only", ColumnType::Int4, ColumnData::I32(Some(1))) + .build(); + + // A wildly out-of-range index resolves to `None` rather than panicking. + assert_eq!(None, usize::MAX.idx(&row)); + assert!(row.try_get::(usize::MAX).is_err()); + } + + #[test] + fn builder_empty_column_name_and_raw_prefix_only_lookup() { + // Pins the interaction between an empty column name and the `r#` fallback + // in `<&str as QueryIdx>::idx`: `"r#"` strips to `""` and matches the + // empty-named column when no exact `"r#"` column exists. + let row = Row::builder() + .column("", ColumnType::Int4, ColumnData::I32(Some(1))) + .build(); + + assert_eq!(Some(0), "".idx(&row)); + assert_eq!(Some(0), "r#".idx(&row)); + assert_eq!(Some(1i32), row.get::("")); + assert_eq!(Some(1i32), row.get::("r#")); + // An unrelated name is still a miss. + assert_eq!(None, "other".idx(&row)); + } + + #[cfg(feature = "serde")] + #[test] + fn builder_row_serde_round_trips() { + // `RowBuilder` is the first way to construct a `Row` without a live + // connection, so it is also the first thing that can exercise the + // feature-gated serde impls on `Row`. Assert a full round-trip preserves + // the columns, the cell data, and the result index. + let row = Row::builder() + .result_index(2) + .column("id", ColumnType::Int4, ColumnData::I32(Some(42))) + .column( + "name", + ColumnType::NVarchar, + ColumnData::String(Some("Alice".into())), + ) + .column("age", ColumnType::Int4, ColumnData::I32(None)) + .build(); + + let json = serde_json::to_string(&row).unwrap(); + let back: Row = serde_json::from_str(&json).unwrap(); + + assert_eq!(2, back.result_index()); + assert_eq!(3, back.len()); + assert_eq!("id", back.columns()[0].name()); + assert_eq!(ColumnType::NVarchar, back.columns()[1].column_type()); + assert_eq!(Some(42i32), back.get::("id")); + assert_eq!(Some("Alice"), back.get::<&str, _>("name")); + assert_eq!(None, back.get::("age")); + } + + #[test] + fn column_new_and_accessors() { + let column = Column::new("id".to_string(), ColumnType::Int8); + assert_eq!("id", column.name()); + assert_eq!(ColumnType::Int8, column.column_type()); + } + + #[test] + fn column_type_from_fixed_len_type_info() { + use crate::tds::codec::FixedLenType; + + let cases = [ + (FixedLenType::Int1, ColumnType::Int1), + (FixedLenType::Bit, ColumnType::Bit), + (FixedLenType::Int2, ColumnType::Int2), + (FixedLenType::Int4, ColumnType::Int4), + (FixedLenType::Datetime4, ColumnType::Datetime4), + (FixedLenType::Float4, ColumnType::Float4), + (FixedLenType::Money, ColumnType::Money), + (FixedLenType::Datetime, ColumnType::Datetime), + (FixedLenType::Float8, ColumnType::Float8), + (FixedLenType::Money4, ColumnType::Money4), + (FixedLenType::Int8, ColumnType::Int8), + (FixedLenType::Null, ColumnType::Null), + ]; + + for (flt, expected) in cases { + let ti = TypeInfo::FixedLen(flt); + assert_eq!(ColumnType::from(&ti), expected); + } + } + + #[test] + fn column_type_from_var_len_sized_type_info() { + use crate::tds::codec::VarLenType; + use crate::VarLenContext; + + let cases = [ + (VarLenType::Guid, 16, ColumnType::Guid), + (VarLenType::Intn, 1, ColumnType::Int1), + (VarLenType::Intn, 2, ColumnType::Int2), + (VarLenType::Intn, 4, ColumnType::Int4), + (VarLenType::Intn, 8, ColumnType::Int8), + (VarLenType::Intn, 3, ColumnType::Intn), + (VarLenType::Bitn, 1, ColumnType::Bitn), + (VarLenType::Decimaln, 17, ColumnType::Decimaln), + (VarLenType::Numericn, 17, ColumnType::Numericn), + (VarLenType::Floatn, 4, ColumnType::Float4), + (VarLenType::Floatn, 8, ColumnType::Float8), + (VarLenType::Floatn, 2, ColumnType::Floatn), + (VarLenType::Money, 8, ColumnType::Money), + (VarLenType::Datetimen, 8, ColumnType::Datetimen), + (VarLenType::BigVarBin, 8000, ColumnType::BigVarBin), + (VarLenType::BigVarChar, 8000, ColumnType::BigVarChar), + (VarLenType::BigBinary, 8000, ColumnType::BigBinary), + (VarLenType::BigChar, 8000, ColumnType::BigChar), + (VarLenType::NVarchar, 4000, ColumnType::NVarchar), + (VarLenType::NChar, 4000, ColumnType::NChar), + (VarLenType::Xml, 0, ColumnType::Xml), + (VarLenType::Udt, 0, ColumnType::Udt), + (VarLenType::Text, 0, ColumnType::Text), + (VarLenType::Image, 0, ColumnType::Image), + (VarLenType::NText, 0, ColumnType::NText), + (VarLenType::SSVariant, 0, ColumnType::SSVariant), + ]; + + for (ty, len, expected) in cases { + let ti = TypeInfo::VarLenSized(VarLenContext::new(ty, len, None)); + assert_eq!(ColumnType::from(&ti), expected, "{:?} len {}", ty, len); + } + } + + #[test] + fn column_type_from_var_len_sized_precision_type_info() { + use crate::tds::codec::VarLenType; + + let cases = [ + (VarLenType::Guid, ColumnType::Guid), + (VarLenType::Intn, ColumnType::Intn), + (VarLenType::Bitn, ColumnType::Bitn), + (VarLenType::Decimaln, ColumnType::Decimaln), + (VarLenType::Numericn, ColumnType::Numericn), + (VarLenType::Floatn, ColumnType::Floatn), + (VarLenType::Money, ColumnType::Money), + (VarLenType::Datetimen, ColumnType::Datetimen), + (VarLenType::BigVarBin, ColumnType::BigVarBin), + (VarLenType::BigVarChar, ColumnType::BigVarChar), + (VarLenType::BigBinary, ColumnType::BigBinary), + (VarLenType::BigChar, ColumnType::BigChar), + (VarLenType::NVarchar, ColumnType::NVarchar), + (VarLenType::NChar, ColumnType::NChar), + (VarLenType::Xml, ColumnType::Xml), + (VarLenType::Udt, ColumnType::Udt), + (VarLenType::Text, ColumnType::Text), + (VarLenType::Image, ColumnType::Image), + (VarLenType::NText, ColumnType::NText), + (VarLenType::SSVariant, ColumnType::SSVariant), + ]; + + for (ty, expected) in cases { + let ti = TypeInfo::VarLenSizedPrecision { + ty, + size: 38, + precision: 38, + scale: 2, + }; + assert_eq!(ColumnType::from(&ti), expected, "{:?}", ty); + } + } + + #[test] + fn column_type_from_xml_and_udt_type_info() { + use crate::tds::codec::UdtInfo; + use crate::tds::xml::XmlSchema; + use std::sync::Arc as StdArc; + + let ti = TypeInfo::Xml { + schema: None::>, + size: 0, + }; + assert_eq!(ColumnType::from(&ti), ColumnType::Xml); + + let ti = TypeInfo::Udt(UdtInfo { + max_byte_size: 0xffff, + db_name: "db".to_string(), + schema_name: "dbo".to_string(), + type_name: "geometry".to_string(), + assembly_qualified_name: "asm".to_string(), + }); + assert_eq!(ColumnType::from(&ti), ColumnType::Udt); + } + + #[cfg(feature = "tds73")] + #[test] + fn column_type_from_var_len_sized_tds73_type_info() { + use crate::tds::codec::VarLenType; + use crate::VarLenContext; + + let cases = [ + (VarLenType::Daten, ColumnType::Daten), + (VarLenType::Timen, ColumnType::Timen), + (VarLenType::Datetime2, ColumnType::Datetime2), + (VarLenType::DatetimeOffsetn, ColumnType::DatetimeOffsetn), + ]; + + for (ty, expected) in cases { + let ti = TypeInfo::VarLenSized(VarLenContext::new(ty, 8, None)); + assert_eq!(ColumnType::from(&ti), expected, "{:?}", ty); + } + } + + #[cfg(feature = "tds73")] + #[test] + fn column_type_from_var_len_sized_precision_tds73_type_info() { + use crate::tds::codec::VarLenType; + + let cases = [ + (VarLenType::Daten, ColumnType::Daten), + (VarLenType::Timen, ColumnType::Timen), + (VarLenType::Datetime2, ColumnType::Datetime2), + (VarLenType::DatetimeOffsetn, ColumnType::DatetimeOffsetn), + ]; + + for (ty, expected) in cases { + let ti = TypeInfo::VarLenSizedPrecision { + ty, + size: 8, + precision: 0, + scale: 7, + }; + assert_eq!(ColumnType::from(&ti), expected, "{:?}", ty); + } + } +} diff --git a/src/sql_browser.rs b/src/sql_browser.rs index b07e8ee22..184691f98 100644 --- a/src/sql_browser.rs +++ b/src/sql_browser.rs @@ -1,9 +1,6 @@ #[cfg(feature = "sql-browser-tokio")] mod tokio; -#[cfg(feature = "sql-browser-async-std")] -mod async_std; - #[cfg(feature = "sql-browser-smol")] mod smol; @@ -27,11 +24,22 @@ pub trait SqlBrowser { Self: Sized + Send + Sync; } -#[cfg(any( - feature = "sql-browser-async-std", - feature = "sql-browser-tokio", - feature = "sql-browser-smol" -))] +/// SSRP `CLNT_UCAST_INST` opcode: a client unicast request for a specific +/// named instance (MS-SQLR §2.2.1). +#[cfg(any(feature = "sql-browser-tokio", feature = "sql-browser-smol"))] +pub(crate) const SSRP_CLIENT_UNICAST: u8 = 4; + +/// Size of the buffer used to receive an SSRP reply datagram. The protocol caps +/// a reply at 65535 bytes, but real replies are small; 4 KiB comfortably holds +/// any practical `tcp;` response. +#[cfg(any(feature = "sql-browser-tokio", feature = "sql-browser-smol"))] +pub(crate) const SSRP_REPLY_BUF_LEN: usize = 4096; + +/// How long to wait for an SSRP reply before giving up, in milliseconds. +#[cfg(any(feature = "sql-browser-tokio", feature = "sql-browser-smol"))] +pub(crate) const SSRP_TIMEOUT_MS: u64 = 1000; + +#[cfg(any(feature = "sql-browser-tokio", feature = "sql-browser-smol"))] fn get_port_from_sql_browser_reply( mut buf: Vec, len: usize, @@ -41,12 +49,23 @@ fn get_port_from_sql_browser_reply( buf.truncate(len); - let err = crate::Error::Conversion( - format!("Could not resolve SQL browser instance {}", instance_name).into(), - ); + // Built fresh on each failure path so the descriptive context (which + // instance failed to resolve) is preserved rather than being collapsed into + // a bare `Error::Utf8`/`Error::ParseInt` by `?`. + let err = || { + crate::Error::Conversion( + format!("Could not resolve SQL browser instance {}", instance_name).into(), + ) + }; - if len == 0 { - return Err(err); + // The SSRP reply is [SVR_RESP(1 byte)][RESP_SIZE(2 bytes, LE)][data...], so + // the instance data starts at offset 3. A reply shorter than that 3-byte + // header is malformed — and SSRP is unauthenticated UDP, so a spoofed or + // truncated datagram is fully attacker-controlled. Guard it explicitly: + // `&buf[3..len]` would otherwise panic ("slice index starts at 3 but ends + // at 1") for a 1- or 2-byte reply. + if len < 3 { + return Err(err()); } let rsp = &buf[3..len]; @@ -56,8 +75,41 @@ fn get_port_from_sql_browser_reply( .rev() .position(|window| window == DELIMITER) .and_then(|pos| rsp[(rsp.len() - pos)..].split(|item| *item == b';').next()) - .ok_or(err) - .and_then(|val| Ok(std::str::from_utf8(val)?.parse()?))?; + .and_then(|val| std::str::from_utf8(val).ok()) + .and_then(|val| val.parse().ok()) + .ok_or_else(err)?; Ok(port) } + +#[cfg(all(test, any(feature = "sql-browser-tokio", feature = "sql-browser-smol")))] +mod tests { + use super::*; + + // A truncated SSRP UDP reply (shorter than the 3-byte header) is fully + // attacker-controlled and must be rejected with a conversion error rather + // than panicking on the `&buf[3..len]` slice. + #[test] + fn truncated_reply_is_rejected_without_panic() { + for reply in [vec![], vec![0x05u8], vec![0x05u8, 0x10]] { + let len = reply.len(); + let err = get_port_from_sql_browser_reply(reply, len, "MSSQLSERVER") + .expect_err("a sub-3-byte reply must error, not panic"); + assert!( + matches!(err, crate::Error::Conversion(_)), + "expected a conversion error, got {err:?}" + ); + } + } + + // A well-formed reply advertising `tcp;1433` resolves to that port. + #[test] + fn well_formed_reply_resolves_port() { + let mut buf = vec![0x05, 0x00, 0x00]; // SVR_RESP + RESP_SIZE header + buf.extend_from_slice(b"ServerName;HOST;InstanceName;MSSQLSERVER;tcp;1433;"); + let len = buf.len(); + + let port = get_port_from_sql_browser_reply(buf, len, "MSSQLSERVER").unwrap(); + assert_eq!(port, 1433); + } +} diff --git a/src/sql_browser/async_std.rs b/src/sql_browser/async_std.rs deleted file mode 100644 index 14f55de57..000000000 --- a/src/sql_browser/async_std.rs +++ /dev/null @@ -1,72 +0,0 @@ -use super::SqlBrowser; -use async_std::{ - io, - net::{self, ToSocketAddrs}, -}; -use async_trait::async_trait; -use futures_util::future::TryFutureExt; -use std::time; -use tracing::Level; - -#[async_trait] -impl SqlBrowser for net::TcpStream { - /// This method can be used to connect to SQL Server named instances - /// when on a Windows platform with the `sql-browser-async-std` feature - /// enabled. Please see the crate examples for more detailed examples. - async fn connect_named(builder: &crate::client::Config) -> crate::Result { - let addrs = builder.get_addr().to_socket_addrs().await?; - - for mut addr in addrs { - if let Some(ref instance_name) = builder.instance_name { - // First resolve the instance to a port via the - // SSRP protocol/MS-SQLR protocol [1] - // [1] https://msdn.microsoft.com/en-us/library/cc219703.aspx - - let local_bind: std::net::SocketAddr = if addr.is_ipv4() { - "0.0.0.0:0".parse().unwrap() - } else { - "[::]:0".parse().unwrap() - }; - - tracing::event!( - Level::TRACE, - "Connecting to instance `{}` using SQL Browser in port `{}`", - instance_name, - builder.get_port() - ); - - let msg = [&[4u8], instance_name.as_bytes()].concat(); - let mut buf = vec![0u8; 4096]; - - let socket = net::UdpSocket::bind(&local_bind).await?; - socket.send_to(&msg, &addr).await?; - - let timeout = time::Duration::from_millis(1000); - - let len = io::timeout(timeout, socket.recv(&mut buf)) - .map_err(|_| { - crate::error::Error::Conversion( - format!( - "SQL browser timeout during resolving instance {}. Please check if browser is running in port {} and does the instance exist.", - instance_name, - builder.get_port(), - ) - .into(), - ) - }) - .await?; - - let port = super::get_port_from_sql_browser_reply(buf, len, instance_name)?; - tracing::event!(Level::TRACE, "Found port `{}` from SQL Browser", port); - addr.set_port(port); - }; - - if let Ok(stream) = net::TcpStream::connect(addr).await { - stream.set_nodelay(true)?; - return Ok(stream); - } - } - - Err(io::Error::new(io::ErrorKind::NotFound, "Could not resolve server host").into()) - } -} diff --git a/src/sql_browser/smol.rs b/src/sql_browser/smol.rs index 252b834e3..dc15ae01e 100644 --- a/src/sql_browser/smol.rs +++ b/src/sql_browser/smol.rs @@ -5,75 +5,101 @@ use async_net::{resolve, TcpStream, UdpSocket}; use async_trait::async_trait; use futures_lite::FutureExt; use futures_util::future::TryFutureExt; +use futures_util::stream::FuturesUnordered; +use futures_util::StreamExt; use std::io; +use std::net::SocketAddr; use std::time::Duration; use tracing::Level; #[async_trait] impl SqlBrowser for TcpStream { /// This method can be used to connect to SQL Server named instances - /// when on a Windows paltform with the `sql-browser-tokio` feature + /// when on a Windows platform with the `sql-browser-smol` feature /// enabled. Please see the crate examples for more detailed examples. async fn connect_named(builder: &Config) -> crate::Result { let addrs = resolve(builder.get_addr()).await?; + let mut first_error = None; - for mut addr in addrs { - if let Some(ref instance_name) = builder.instance_name { - // First resolve the instance to a port via the - // SSRP protocol/MS-SQLR protocol [1] - // [1] https://msdn.microsoft.com/en-us/library/cc219703.aspx - - let local_bind: std::net::SocketAddr = if addr.is_ipv4() { - "0.0.0.0:0".parse().unwrap() - } else { - "[::]:0".parse().unwrap() + if builder.multi_subnet_failover { + let mut futures = addrs + .into_iter() + .map(|addr| connect_addr(builder, addr)) + .collect::>(); + while let Some(connection) = futures.next().await { + match connection { + Ok(connection) => return Ok(connection), + Err(error) => first_error.get_or_insert(error), + }; + } + } else { + for addr in addrs { + match connect_addr(builder, addr).await { + Ok(connection) => return Ok(connection), + Err(error) => first_error.get_or_insert(error), }; + } + } + + // If we end up here, there was no successful connection. + Err(first_error.unwrap_or_else(|| { + io::Error::new(io::ErrorKind::NotFound, "Could not resolve server host").into() + })) + } +} - tracing::event!( - Level::TRACE, - "Connecting to instance `{}` using SQL Browser in port `{}`", - instance_name, - builder.get_port() - ); +async fn connect_addr(builder: &Config, mut addr: SocketAddr) -> crate::Result { + if let Some(ref instance_name) = builder.instance_name { + // First resolve the instance to a port via the + // SSRP protocol/MS-SQLR protocol [1] + // [1] https://msdn.microsoft.com/en-us/library/cc219703.aspx - let msg = [&[4u8], instance_name.as_bytes()].concat(); - let mut buf = vec![0u8; 4096]; + let local_bind: std::net::SocketAddr = if addr.is_ipv4() { + "0.0.0.0:0".parse().unwrap() + } else { + "[::]:0".parse().unwrap() + }; - let socket = UdpSocket::bind(&local_bind).await?; - socket.send_to(&msg, &addr).await?; + tracing::event!( + Level::TRACE, + "Connecting to instance `{}` using SQL Browser in port `{}`", + instance_name, + builder.get_port() + ); - let timeout = Duration::from_millis(1000); + let msg = [&[super::SSRP_CLIENT_UNICAST], instance_name.as_bytes()].concat(); + let mut buf = vec![0u8; super::SSRP_REPLY_BUF_LEN]; - let len = socket.recv(&mut buf).or(async { - Timer::after(timeout).await; - Err(std::io::ErrorKind::TimedOut.into()) - }) - .map_err(|e| { - if e.kind() == std::io::ErrorKind::TimedOut { - crate::error::Error::Conversion( - format!( - "SQL browser timeout during resolving instance {}. Please check if browser is running in port {} and does the instance exist.", - instance_name, - builder.get_port(), - ) - .into(), - ) - } else { - e.into() - } - }).await?; + let socket = UdpSocket::bind(&local_bind).await?; + socket.send_to(&msg, &addr).await?; - let port = super::get_port_from_sql_browser_reply(buf, len, instance_name)?; - tracing::event!(Level::TRACE, "Found port `{}` from SQL Browser", port); - addr.set_port(port); - }; + let timeout = Duration::from_millis(super::SSRP_TIMEOUT_MS); - if let Ok(stream) = TcpStream::connect(addr).await { - stream.set_nodelay(true)?; - return Ok(stream); - } - } + let len = socket.recv(&mut buf).or(async { + Timer::after(timeout).await; + Err(std::io::ErrorKind::TimedOut.into()) + }) + .map_err(|e| { + if e.kind() == std::io::ErrorKind::TimedOut { + crate::error::Error::Conversion( + format!( + "SQL browser timeout during resolving instance {}. Please check if browser is running in port {} and does the instance exist.", + instance_name, + builder.get_port(), + ) + .into(), + ) + } else { + e.into() + } + }).await?; - Err(io::Error::new(io::ErrorKind::NotFound, "Could not resolve server host").into()) - } + let port = super::get_port_from_sql_browser_reply(buf, len, instance_name)?; + tracing::event!(Level::TRACE, "Found port `{}` from SQL Browser", port); + addr.set_port(port); + }; + + let stream = TcpStream::connect(addr).await?; + stream.set_nodelay(true)?; + Ok(stream) } diff --git a/src/sql_browser/tokio.rs b/src/sql_browser/tokio.rs index 1fbf6e0e3..59838bb29 100644 --- a/src/sql_browser/tokio.rs +++ b/src/sql_browser/tokio.rs @@ -2,8 +2,10 @@ use super::SqlBrowser; use crate::client::Config; use async_trait::async_trait; use futures_util::future::TryFutureExt; +use futures_util::stream::FuturesUnordered; +use futures_util::StreamExt; use net::{TcpStream, UdpSocket}; -use std::io; +use std::{io, net::SocketAddr}; use tokio::{ net, time::{self, error::Elapsed, Duration}, @@ -13,62 +15,84 @@ use tracing::Level; #[async_trait] impl SqlBrowser for TcpStream { /// This method can be used to connect to SQL Server named instances - /// when on a Windows paltform with the `sql-browser-tokio` feature + /// when on a Windows platform with the `sql-browser-tokio` feature /// enabled. Please see the crate examples for more detailed examples. async fn connect_named(builder: &Config) -> crate::Result { let addrs = net::lookup_host(builder.get_addr()).await?; + let mut first_error = None; - for mut addr in addrs { - if let Some(ref instance_name) = builder.instance_name { - // First resolve the instance to a port via the - // SSRP protocol/MS-SQLR protocol [1] - // [1] https://msdn.microsoft.com/en-us/library/cc219703.aspx - - let local_bind: std::net::SocketAddr = if addr.is_ipv4() { - "0.0.0.0:0".parse().unwrap() - } else { - "[::]:0".parse().unwrap() + if builder.multi_subnet_failover { + let mut futures = addrs + .map(|addr| connect_addr(builder, addr)) + .collect::>(); + while let Some(connection) = futures.next().await { + match connection { + Ok(connection) => return Ok(connection), + Err(error) => first_error.get_or_insert(error), + }; + } + } else { + for addr in addrs { + match connect_addr(builder, addr).await { + Ok(connection) => return Ok(connection), + Err(error) => first_error.get_or_insert(error), }; + } + } + + // If we end up here, there was no successful connection. + Err(first_error.unwrap_or_else(|| { + io::Error::new(io::ErrorKind::NotFound, "Could not resolve server host").into() + })) + } +} - tracing::event!( - Level::TRACE, - "Connecting to instance `{}` using SQL Browser in port `{}`", - instance_name, - builder.get_port() - ); +async fn connect_addr(builder: &Config, mut addr: SocketAddr) -> crate::Result { + if let Some(ref instance_name) = builder.instance_name { + // First resolve the instance to a port via the + // SSRP protocol/MS-SQLR protocol [1] + // [1] https://msdn.microsoft.com/en-us/library/cc219703.aspx - let msg = [&[4u8], instance_name.as_bytes()].concat(); - let mut buf = vec![0u8; 4096]; + let local_bind: std::net::SocketAddr = if addr.is_ipv4() { + "0.0.0.0:0".parse().unwrap() + } else { + "[::]:0".parse().unwrap() + }; - let socket = UdpSocket::bind(&local_bind).await?; - socket.send_to(&msg, &addr).await?; + tracing::event!( + Level::TRACE, + "Connecting to instance `{}` using SQL Browser in port `{}`", + instance_name, + builder.get_port() + ); - let timeout = Duration::from_millis(1000); + let msg = [&[super::SSRP_CLIENT_UNICAST], instance_name.as_bytes()].concat(); + let mut buf = vec![0u8; super::SSRP_REPLY_BUF_LEN]; - let len = time::timeout(timeout, socket.recv(&mut buf)) - .map_err(|_: Elapsed| { - crate::error::Error::Conversion( - format!( - "SQL browser timeout during resolving instance {}. Please check if browser is running in port {} and does the instance exist.", - instance_name, - builder.get_port(), - ) - .into(), - ) - }) - .await??; + let socket = UdpSocket::bind(&local_bind).await?; + socket.send_to(&msg, &addr).await?; - let port = super::get_port_from_sql_browser_reply(buf, len, instance_name)?; - tracing::event!(Level::TRACE, "Found port `{}` from SQL Browser", port); - addr.set_port(port); - }; + let timeout = Duration::from_millis(super::SSRP_TIMEOUT_MS); - if let Ok(stream) = TcpStream::connect(addr).await { - stream.set_nodelay(true)?; - return Ok(stream); - } - } + let len = time::timeout(timeout, socket.recv(&mut buf)) + .map_err(|_: Elapsed| { + crate::error::Error::Conversion( + format!( + "SQL browser timeout during resolving instance {}. Please check if browser is running in port {} and does the instance exist.", + instance_name, + builder.get_port(), + ) + .into(), + ) + }) + .await??; - Err(io::Error::new(io::ErrorKind::NotFound, "Could not resolve server host").into()) - } + let port = super::get_port_from_sql_browser_reply(buf, len, instance_name)?; + tracing::event!(Level::TRACE, "Found port `{}` from SQL Browser", port); + addr.set_port(port); + }; + + let stream = TcpStream::connect(addr).await?; + stream.set_nodelay(true)?; + Ok(stream) } diff --git a/src/sql_read_bytes.rs b/src/sql_read_bytes.rs index 0455a1ce7..87edafa8e 100644 --- a/src/sql_read_bytes.rs +++ b/src/sql_read_bytes.rs @@ -6,6 +6,71 @@ use std::io::ErrorKind::UnexpectedEof; use std::{future::Future, io, mem::size_of, pin::Pin, task}; use task::Poll; +/// Reads exactly `len` bytes from a packet-spanning reader, appending them to +/// `dst`. +/// +/// This is the bulk counterpart to the per-byte `read_u8` loop used by the hot +/// decode paths. A single decoded value (a `varchar`/`nvarchar`/`varbinary` +/// column, a `text`/`ntext`/`image` blob, a `sql_variant` payload, …) can span +/// several TDS packets. The underlying [`SqlReadBytes`] reader hides those +/// packet boundaries: each `poll_read` transparently pulls and concatenates as +/// many packet payloads as needed to satisfy the request, so copying whatever a +/// `poll_read` returns crosses boundaries exactly the way the byte-by-byte loop +/// did — only O(packets) polls instead of O(bytes). +/// +/// `AsyncReadExt::read_exact` is deliberately avoided (see the module callers): +/// this helper drives `poll_read` directly and turns a clean `Ok(0)` (true +/// stream EOF) mid-value into an `UnexpectedEof`, matching the error behaviour +/// of the `read_u8` loop it replaces byte-for-byte. +/// +/// The up-front allocation is capped at `prealloc_cap` and the buffer grows in +/// windows of at most `prealloc_cap` bytes, preserving the `MAX_PREALLOC` +/// reservation-capping semantics of the call sites: an untrusted, over-large +/// `len` never triggers a huge allocation before the bytes actually arrive, yet +/// a genuine long value still reads in full. +pub(crate) async fn read_bytes_into( + src: &mut R, + dst: &mut Vec, + len: usize, + prealloc_cap: usize, +) -> io::Result<()> +where + R: AsyncRead + Unpin, +{ + let mut remaining = len; + + while remaining > 0 { + // Grow by at most one window so a lying `len` cannot force a large + // allocation up front; `want` never exceeds the value's own remaining + // bytes, so we never over-read into the following field. + let want = remaining.min(prealloc_cap); + let start = dst.len(); + let end = start + want; + dst.resize(end, 0); + + let mut filled = start; + while filled < end { + let n = + std::future::poll_fn(|cx| Pin::new(&mut *src).poll_read(cx, &mut dst[filled..end])) + .await?; + + // A clean EOF before the value is complete is an error, exactly as + // the per-byte `read_u8` loop treated a boundary `Ok(0)`. + if n == 0 { + // Drop the tail we reserved-but-never-filled. + dst.truncate(filled); + return Err(UnexpectedEof.into()); + } + + filled += n; + } + + remaining -= want; + } + + Ok(()) +} + macro_rules! varchar_reader { ($name:ident, $length_reader:ident) => { pin_project! { @@ -332,6 +397,402 @@ bytes_reader!(ReadF64, f64, get_f64); bytes_reader!(ReadF32Le, f32, get_f32_le); bytes_reader!(ReadF64Le, f64, get_f64_le); +#[cfg(test)] +mod tests { + use super::test_utils::IntoSqlReadBytes; + use crate::SqlReadBytes; + use bytes::{BufMut, BytesMut}; + + #[tokio::test] + async fn read_i8_value() { + let mut buf = BytesMut::new(); + buf.put_i8(-5); + assert_eq!(buf.into_sql_read_bytes().read_i8().await.unwrap(), -5); + } + + #[tokio::test] + async fn read_u32_big_endian() { + let mut buf = BytesMut::new(); + buf.put_u32(0x01020304); + assert_eq!( + buf.into_sql_read_bytes().read_u32().await.unwrap(), + 0x01020304 + ); + } + + #[tokio::test] + async fn read_f32_and_f64_big_endian() { + let mut buf = BytesMut::new(); + buf.put_f32(1.5); + assert_eq!(buf.into_sql_read_bytes().read_f32().await.unwrap(), 1.5); + + let mut buf = BytesMut::new(); + buf.put_f64(2.5); + assert_eq!(buf.into_sql_read_bytes().read_f64().await.unwrap(), 2.5); + } + + #[tokio::test] + async fn read_f32_and_f64_little_endian() { + let mut buf = BytesMut::new(); + buf.put_f32_le(1.5); + assert_eq!(buf.into_sql_read_bytes().read_f32_le().await.unwrap(), 1.5); + + let mut buf = BytesMut::new(); + buf.put_f64_le(2.5); + assert_eq!(buf.into_sql_read_bytes().read_f64_le().await.unwrap(), 2.5); + } + + #[tokio::test] + async fn read_u128_and_i128_le() { + let mut buf = BytesMut::new(); + buf.put_u128_le(12345); + assert_eq!( + buf.into_sql_read_bytes().read_u128_le().await.unwrap(), + 12345 + ); + + let mut buf = BytesMut::new(); + buf.put_i128_le(-12345); + assert_eq!( + buf.into_sql_read_bytes().read_i128_le().await.unwrap(), + -12345 + ); + } + + #[tokio::test] + async fn read_b_varchar_and_us_varchar() { + let mut buf = BytesMut::new(); + buf.put_u8(2); + buf.put_u16_le('h' as u16); + buf.put_u16_le('i' as u16); + assert_eq!( + buf.into_sql_read_bytes().read_b_varchar().await.unwrap(), + "hi" + ); + + let mut buf = BytesMut::new(); + buf.put_u16_le(2); + buf.put_u16_le('h' as u16); + buf.put_u16_le('i' as u16); + assert_eq!( + buf.into_sql_read_bytes().read_us_varchar().await.unwrap(), + "hi" + ); + } + + #[tokio::test] + async fn context_and_context_mut_accessible() { + let buf = BytesMut::new(); + let mut reader = buf.into_sql_read_bytes(); + assert_eq!(reader.context().packet_size(), 4096); + reader.context_mut().set_packet_size(8192); + assert_eq!(reader.context().packet_size(), 8192); + } + + // The length prefix cannot be read (empty wire) — exercises the error arm of + // the varchar length read (`Poll::Ready(Err(..))`). + #[tokio::test] + async fn b_varchar_length_read_error() { + let buf = BytesMut::new(); + assert!(buf.into_sql_read_bytes().read_b_varchar().await.is_err()); + } + + // The length is read but the character payload is truncated — exercises the + // error arm of the inner u16 read within the varchar loop. + #[tokio::test] + async fn b_varchar_data_read_error() { + let mut buf = BytesMut::new(); + buf.put_u8(1); // announce one u16 char... + buf.put_u8(0x41); // ...but supply only a single byte + assert!(buf.into_sql_read_bytes().read_b_varchar().await.is_err()); + } + + // A lone UTF-16 surrogate makes `String::from_utf16` fail — exercises the + // invalid-UTF-16 error mapping at the end of the varchar reader. + #[tokio::test] + async fn b_varchar_invalid_utf16_error() { + let mut buf = BytesMut::new(); + buf.put_u8(1); + buf.put_u16_le(0xD800); // unpaired high surrogate + assert!(buf.into_sql_read_bytes().read_b_varchar().await.is_err()); + } +} + +// Tests for the `Poll::Pending` / clean-EOF branches of the readers, which +// require an `AsyncRead` that can return `Pending` / `Ok(0)` on demand and a +// manually driven poll. +#[cfg(test)] +mod poll_branch_tests { + use crate::tds::Context; + use crate::SqlReadBytes; + use bytes::{BufMut, BytesMut}; + use futures_util::io::AsyncRead; + use std::future::Future; + use std::io; + use std::pin::Pin; + use std::task::{Context as TaskContext, Poll}; + + enum Then { + Pending, + Eof, + } + + // Hands out `data` while enough bytes remain, then switches to returning + // either `Poll::Pending` or a clean EOF (`Ok(0)`). + struct ScriptedReader { + data: BytesMut, + then: Then, + ctx: Context, + } + + impl ScriptedReader { + fn new(data: BytesMut, then: Then) -> Self { + Self { + data, + then, + ctx: Context::new(), + } + } + } + + impl AsyncRead for ScriptedReader { + fn poll_read( + self: Pin<&mut Self>, + _cx: &mut TaskContext<'_>, + buf: &mut [u8], + ) -> Poll> { + let this = self.get_mut(); + let size = buf.len(); + + if size > 0 && this.data.len() >= size { + buf.copy_from_slice(this.data.split_to(size).as_ref()); + return Poll::Ready(Ok(size)); + } + + match this.then { + Then::Pending => Poll::Pending, + Then::Eof => Poll::Ready(Ok(0)), + } + } + } + + impl SqlReadBytes for ScriptedReader { + fn debug_buffer(&self) {} + fn context(&self) -> &Context { + &self.ctx + } + fn context_mut(&mut self) -> &mut Context { + &mut self.ctx + } + } + + fn poll_once(fut: F) -> Poll { + let waker = std::task::Waker::noop(); + let mut cx = TaskContext::from_waker(waker); + let mut fut = std::pin::pin!(fut); + fut.as_mut().poll(&mut cx) + } + + // The varchar length read yields `Pending` (no bytes available yet). + #[test] + fn varchar_length_pending() { + let mut reader = ScriptedReader::new(BytesMut::new(), Then::Pending); + assert!(matches!(poll_once(reader.read_b_varchar()), Poll::Pending)); + } + + // The length is read, but the character payload read yields `Pending`. + #[test] + fn varchar_data_pending() { + let mut data = BytesMut::new(); + data.put_u8(1); // length available, character bytes are not + let mut reader = ScriptedReader::new(data, Then::Pending); + assert!(matches!(poll_once(reader.read_b_varchar()), Poll::Pending)); + } + + // A fixed-width numeric read yields `Pending` when no bytes are available. + #[test] + fn fixed_width_read_pending() { + let mut reader = ScriptedReader::new(BytesMut::new(), Then::Pending); + assert!(matches!(poll_once(reader.read_u32_le()), Poll::Pending)); + } + + // A clean EOF (`Ok(0)`) mid-read surfaces as an `UnexpectedEof` error. + #[test] + fn fixed_width_read_unexpected_eof() { + let mut reader = ScriptedReader::new(BytesMut::new(), Then::Eof); + match poll_once(reader.read_u8()) { + Poll::Ready(Err(e)) => assert_eq!(e.kind(), io::ErrorKind::UnexpectedEof), + other => panic!("expected UnexpectedEof, got {other:?}"), + } + } +} + +// Tests for the bulk `read_bytes_into` primitive: it must reproduce the exact +// bytes the old per-byte `read_u8` loop produced, including across simulated +// TDS packet boundaries, and must fail cleanly on a truncated value. +#[cfg(test)] +mod bulk_read_tests { + use super::read_bytes_into; + use crate::tds::Context; + use crate::SqlReadBytes; + use futures_util::io::AsyncRead; + use std::io; + use std::pin::Pin; + use std::task::{Context as TaskContext, Poll}; + + // Hands out at most `chunk` bytes per `poll_read`, simulating a value that + // is fragmented across several TDS packets. Returns a clean EOF (`Ok(0)`) + // once its data is exhausted. + struct ChunkedReader { + data: Vec, + pos: usize, + chunk: usize, + ctx: Context, + } + + impl ChunkedReader { + fn new(data: Vec, chunk: usize) -> Self { + Self { + data, + pos: 0, + chunk, + ctx: Context::new(), + } + } + } + + impl AsyncRead for ChunkedReader { + fn poll_read( + self: Pin<&mut Self>, + _cx: &mut TaskContext<'_>, + buf: &mut [u8], + ) -> Poll> { + let this = self.get_mut(); + let avail = this.data.len() - this.pos; + if avail == 0 { + return Poll::Ready(Ok(0)); + } + let n = buf.len().min(this.chunk).min(avail); + buf[..n].copy_from_slice(&this.data[this.pos..this.pos + n]); + this.pos += n; + Poll::Ready(Ok(n)) + } + } + + impl SqlReadBytes for ChunkedReader { + fn debug_buffer(&self) {} + fn context(&self) -> &Context { + &self.ctx + } + fn context_mut(&mut self) -> &mut Context { + &mut self.ctx + } + } + + // The byte-by-byte reference: what the old `for _ in 0..len { read_u8 }` + // loop would have produced. Since the reader is deterministic, the bulk + // read must equal the raw source bytes. + #[tokio::test] + async fn bulk_read_matches_source_across_packet_boundaries() { + let data: Vec = (0..5000u32).map(|i| (i % 251) as u8).collect(); + + // A 7-byte "packet" chunk forces many boundary crossings within one + // value. + let mut reader = ChunkedReader::new(data.clone(), 7); + let mut got = Vec::new(); + read_bytes_into(&mut reader, &mut got, data.len(), 8192) + .await + .unwrap(); + + assert_eq!(got, data); + } + + // Reading with a `prealloc_cap` far smaller than `len` must still read the + // whole value: exercises the windowed-growth outer loop. + #[tokio::test] + async fn bulk_read_windows_when_len_exceeds_cap() { + let data: Vec = (0..1000u32).map(|i| (i % 97) as u8).collect(); + + let mut reader = ChunkedReader::new(data.clone(), 13); + let mut got = Vec::new(); + // cap of 8 is much smaller than len (1000) and the chunk size (13). + read_bytes_into(&mut reader, &mut got, data.len(), 8) + .await + .unwrap(); + + assert_eq!(got, data); + } + + // Appending into a non-empty buffer preserves the existing prefix (the PLP + // chunk-accumulation case). + #[tokio::test] + async fn bulk_read_appends_to_existing_buffer() { + let mut got = vec![0xDE, 0xAD]; + let tail: Vec = (0..300u32).map(|i| i as u8).collect(); + + let mut reader = ChunkedReader::new(tail.clone(), 16); + read_bytes_into(&mut reader, &mut got, tail.len(), 8192) + .await + .unwrap(); + + let mut expected = vec![0xDE, 0xAD]; + expected.extend_from_slice(&tail); + assert_eq!(got, expected); + } + + // A value truncated by a clean EOF must surface as `UnexpectedEof`, exactly + // as the per-byte loop did, and must not leave reserved-but-unfilled tail + // bytes in the buffer. + #[tokio::test] + async fn bulk_read_truncated_value_is_unexpected_eof() { + let mut reader = ChunkedReader::new(vec![1, 2, 3], 2); + let mut got = Vec::new(); + let err = read_bytes_into(&mut reader, &mut got, 10, 8192) + .await + .expect_err("a truncated value must error"); + assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof); + // Only the bytes actually read are retained. + assert_eq!(got, vec![1, 2, 3]); + } + + // Anti-DoS windowing property: an attacker-controlled, absurdly large `len` + // must NOT trigger an up-front reservation of `len` bytes. The buffer only + // grows one `prealloc_cap` window at a time as bytes actually arrive, so + // before any byte is delivered the capacity cannot jump past a single + // `prealloc_cap` window — it never balloons toward `len`. + #[tokio::test] + async fn bulk_read_over_large_len_reserves_at_most_one_window() { + let prealloc_cap = 8192usize; + + // Immediate EOF: no bytes are ever delivered. + let mut reader = ChunkedReader::new(Vec::new(), 64); + let mut got = Vec::new(); + let err = read_bytes_into(&mut reader, &mut got, usize::MAX, prealloc_cap) + .await + .expect_err("EOF before any byte must error"); + assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof); + + // Only one window was reserved despite the near-`usize::MAX` `len`. + assert!( + got.capacity() <= prealloc_cap, + "over-large len reserved {} bytes, more than the {}-byte window", + got.capacity(), + prealloc_cap + ); + } + + // A zero-length read is a no-op that leaves the buffer untouched. + #[tokio::test] + async fn bulk_read_zero_len_is_noop() { + let mut reader = ChunkedReader::new(vec![9, 9, 9], 1); + let mut got = vec![0x11]; + read_bytes_into(&mut reader, &mut got, 0, 8192) + .await + .unwrap(); + assert_eq!(got, vec![0x11]); + } +} + #[cfg(test)] pub(crate) mod test_utils { use crate::tds::Context; @@ -352,12 +813,16 @@ pub(crate) mod test_utils { type T = BytesMutReader; fn into_sql_read_bytes(self) -> Self::T { - BytesMutReader { buf: self } + BytesMutReader { + buf: self, + ctx: Context::new(), + } } } pub(crate) struct BytesMutReader { buf: BytesMut, + ctx: Context, } impl AsyncRead for BytesMutReader { @@ -388,11 +853,11 @@ pub(crate) mod test_utils { } fn context(&self) -> &Context { - todo!() + &self.ctx } fn context_mut(&mut self) -> &mut Context { - todo!() + &mut self.ctx } } } diff --git a/src/tds.rs b/src/tds.rs index f4b6f9253..5c23b39d3 100644 --- a/src/tds.rs +++ b/src/tds.rs @@ -1,5 +1,5 @@ pub mod codec; -mod collation; +pub(crate) mod collation; mod context; pub mod numeric; pub mod stream; @@ -25,6 +25,37 @@ uint_enum! { NotSupported = 2, /// Encrypt everything and fail if not possible Required = 3, + /// Start encryption before the TDS prelogin (TDS 8.0 "strict" mode) and + /// encrypt everything, failing if not possible. + Strict = 4, } } + +impl EncryptionLevel { + /// The value sent on the wire in the prelogin `ENCRYPTION` option. + /// + /// `Strict` (TDS 8.0) is negotiated out-of-band via a TLS handshake before + /// the prelogin, so when a prelogin is emitted at all it advertises the + /// classic `Required` value. + pub(crate) fn as_wire_value(&self) -> u8 { + match self { + EncryptionLevel::Strict => EncryptionLevel::Required as u8, + other => *other as u8, + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn encryption_level_as_wire_value() { + assert_eq!(EncryptionLevel::Off.as_wire_value(), 0); + assert_eq!(EncryptionLevel::On.as_wire_value(), 1); + assert_eq!(EncryptionLevel::NotSupported.as_wire_value(), 2); + assert_eq!(EncryptionLevel::Required.as_wire_value(), 3); + assert_eq!(EncryptionLevel::Strict.as_wire_value(), 3); + } +} diff --git a/src/tds/codec.rs b/src/tds/codec.rs index 07f133106..44ec88db2 100644 --- a/src/tds/codec.rs +++ b/src/tds/codec.rs @@ -11,7 +11,9 @@ mod packet; mod pre_login; mod rpc_request; mod token; +mod transaction_manager; mod type_info; +mod type_info_tvp; pub use batch_request::*; pub use bulk_load::*; @@ -27,7 +29,9 @@ pub use packet::*; pub use pre_login::*; pub use rpc_request::*; pub use token::*; +pub use transaction_manager::*; pub use type_info::*; +pub use type_info_tvp::*; const HEADER_BYTES: usize = 8; const ALL_HEADERS_LEN_TX: usize = 22; diff --git a/src/tds/codec/batch_request.rs b/src/tds/codec/batch_request.rs index f68da0e2e..517e6cdb9 100644 --- a/src/tds/codec/batch_request.rs +++ b/src/tds/codec/batch_request.rs @@ -1,4 +1,4 @@ -use super::{AllHeaderTy, Encode, ALL_HEADERS_LEN_TX}; +use super::{encode_all_headers_tx, Encode}; use bytes::{BufMut, BytesMut}; use std::borrow::Cow; @@ -18,11 +18,7 @@ impl<'a> BatchRequest<'a> { impl<'a> Encode for BatchRequest<'a> { fn encode(self, dst: &mut BytesMut) -> crate::Result<()> { - dst.put_u32_le(ALL_HEADERS_LEN_TX as u32); - dst.put_u32_le(ALL_HEADERS_LEN_TX as u32 - 4); - dst.put_u16_le(AllHeaderTy::TransactionDescriptor as u16); - dst.put_slice(&self.transaction_descriptor); - dst.put_u32_le(1); + encode_all_headers_tx(dst, self.transaction_descriptor); for c in self.queries.encode_utf16() { dst.put_u16_le(c); @@ -31,3 +27,28 @@ impl<'a> Encode for BatchRequest<'a> { Ok(()) } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn encode_is_byte_exact() { + let td = [1u8, 2, 3, 4, 5, 6, 7, 8]; + let req = BatchRequest::new("Hi", td); + + let mut dst = BytesMut::new(); + req.encode(&mut dst).unwrap(); + + let mut expected = Vec::new(); + expected.extend_from_slice(&22u32.to_le_bytes()); // ALL_HEADERS_LEN_TX + expected.extend_from_slice(&18u32.to_le_bytes()); // header length (len - 4) + expected.extend_from_slice(&2u16.to_le_bytes()); // TransactionDescriptor type + expected.extend_from_slice(&td); // transaction descriptor + expected.extend_from_slice(&1u32.to_le_bytes()); // outstanding request count + expected.extend_from_slice(&('H' as u16).to_le_bytes()); + expected.extend_from_slice(&('i' as u16).to_le_bytes()); + + assert_eq!(&dst[..], &expected[..]); + } +} diff --git a/src/tds/codec/bulk_load.rs b/src/tds/codec/bulk_load.rs index 36a17d5cb..6d196ec00 100644 --- a/src/tds/codec/bulk_load.rs +++ b/src/tds/codec/bulk_load.rs @@ -58,10 +58,33 @@ where /// data and for the data to actually be available in the table. /// /// [`finalize`]: #method.finalize + /// + /// # Errors + /// + /// Returns an error if the connection is already poisoned by a prior + /// interrupted write, if the row cannot be encoded (e.g. a value that does + /// not fit its column), or if flushing the buffered packets to the wire + /// fails. On an encode failure the buffered bytes for the partial row are + /// rolled back so the stream stays in sync. pub async fn send(&mut self, row: TokenRow<'a>) -> crate::Result<()> { + // Fail fast if a previous write on this connection was interrupted: the + // stream is already desynced and buffering/writing more bulk data would + // corrupt it further. + self.connection.ensure_not_poisoned()?; + + // `row.encode` can now fail mid-row (e.g. an out-of-range money value). + // A failure would leave the Row-token byte plus any already-encoded + // columns in `self.buf`; those partial bytes would be flushed on the next + // successful `send`/`finalize` and desync the bulk stream. Snapshot the + // buffer length and roll back on error so the stream stays in sync. + let start = self.buf.len(); + let mut buf_with_columns = BytesMutWithDataColumns::new(&mut self.buf, &self.columns); + if let Err(e) = row.encode(&mut buf_with_columns) { + self.buf.truncate(start); + return Err(e); + } - row.encode(&mut buf_with_columns)?; self.write_packets().await?; Ok(()) @@ -72,7 +95,16 @@ where /// This method must be called after sending all the data to flush all /// pending data and to get the server actually to store the rows to the /// table. + /// + /// # Errors + /// + /// Returns an error if the connection is already poisoned by a prior + /// interrupted write, if flushing the remaining packets fails, or if the + /// server reports an error while committing the bulk load. pub async fn finalize(mut self) -> crate::Result { + // See `send`: never layer a finalize onto an already-desynced stream. + self.connection.ensure_not_poisoned()?; + TokenDone::default().encode(&mut self.buf)?; self.write_packets().await?; @@ -87,8 +119,14 @@ where data.len() + HEADER_BYTES, ); + // Bracket the final end-of-message write with the poison flag, mirroring + // `Connection::send`: if this future is dropped between the write and a + // clean flush the message is only partly on the wire and the connection + // must not be silently reused. A clean flush clears the flag. + self.connection.poison(); self.connection.write_to_wire(header, data).await?; self.connection.flush_sink().await?; + self.connection.unpoison(); ExecuteResult::new(self.connection).await } @@ -96,6 +134,18 @@ where async fn write_packets(&mut self) -> crate::Result<()> { let packet_size = (self.connection.context().packet_size() as usize) - HEADER_BYTES; + // Nothing to flush yet: the buffered data still fits in a single packet, + // so no partial message reaches the wire and there is nothing to guard. + if self.buf.len() <= packet_size { + return Ok(()); + } + + // Bracket the multi-packet write with the poison flag, mirroring + // `Connection::send`: while packets are going out the message is only + // partly on the wire, so a dropped future must leave the connection + // poisoned rather than reusable. A clean completion clears it. + self.connection.poison(); + while self.buf.len() > packet_size { let header = PacketHeader::bulk_load(self.packet_id); let data = self.buf.split_to(packet_size); @@ -109,6 +159,115 @@ where self.connection.write_to_wire(header, data).await?; } + self.connection.unpoison(); + Ok(()) } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::tds::codec::type_info::VarLenContext; + use crate::{BaseMetaDataColumn, ColumnData, ColumnFlag, TypeInfo, VarLenType}; + use std::pin::Pin; + use std::task::{Context as TaskContext, Poll}; + use std::{io, task}; + + // A do-nothing stream: writes are swallowed and reads report clean EOF. A + // single small bulk row fits one packet, so `write_packets` returns without + // ever touching this stream — the rollback test never needs the wire. + struct NullIo; + + impl AsyncRead for NullIo { + fn poll_read( + self: Pin<&mut Self>, + _: &mut TaskContext<'_>, + _: &mut [u8], + ) -> Poll> { + Poll::Ready(Ok(0)) + } + } + + impl AsyncWrite for NullIo { + fn poll_write( + self: Pin<&mut Self>, + _: &mut TaskContext<'_>, + buf: &[u8], + ) -> Poll> { + Poll::Ready(Ok(buf.len())) + } + + fn poll_flush(self: Pin<&mut Self>, _: &mut task::Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn poll_close(self: Pin<&mut Self>, _: &mut task::Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + } + + // A single nullable `money` (VarLenSized, len 8) column. An out-of-range + // `F64` value routes through `money::encode`, which rejects it with + // `Error::BulkInput`, making `TokenRow::encode` fail mid-row. + fn money_columns() -> Vec> { + vec![MetaDataColumn { + base: BaseMetaDataColumn { + flags: ColumnFlag::Nullable.into(), + ty: TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Money, 8, None)), + table_name: None, + }, + col_name: Default::default(), + }] + } + + // On an encode failure mid-row, `send` must roll the buffer back to the + // post-COLMETADATA baseline it had right after `new`, so the partial + // Row-token bytes never leak onto the wire and desync the bulk stream. + // + // Red-before-green evidence: commenting out the `self.buf.truncate(start)` + // line in `send` makes the final length assertion FAIL — the buffer stays + // longer than the baseline because the Row-token byte and any partial column + // bytes remain. Restoring the line makes it pass. Verified locally. + #[tokio::test] + async fn send_rolls_back_partial_row_on_encode_error() { + let mut conn = Connection::test_over(NullIo, false); + let columns = money_columns(); + + let mut req = BulkLoadRequest::new(&mut conn, columns).unwrap(); + + // The baseline is the COLMETADATA token `new` buffered; the rollback must + // restore exactly this length. + let baseline = req.buf.len(); + + // `1e18` is far past money's range; `money::encode` returns `BulkInput`. + let mut bad_row = TokenRow::new(); + bad_row.push(ColumnData::F64(Some(1e18))); + + let err = req + .send(bad_row) + .await + .expect_err("an out-of-range money value must fail to encode"); + assert!(matches!(err, crate::Error::BulkInput(_)), "got {err:?}"); + + // The partial Row-token bytes were rolled back: the buffer is exactly the + // post-COLMETADATA baseline again. + assert_eq!( + req.buf.len(), + baseline, + "partial row bytes were not rolled back" + ); + + // The stream stayed in sync: a subsequent valid row encodes cleanly and + // appends onto the untouched baseline. + let mut good_row = TokenRow::new(); + good_row.push(ColumnData::F64(Some(1.5))); + req.send(good_row) + .await + .expect("a valid row must send after a rolled-back failure"); + assert!( + req.buf.len() > baseline, + "a valid row should append to the buffer" + ); + } +} diff --git a/src/tds/codec/column_data.rs b/src/tds/codec/column_data.rs index fecd83f75..190423be3 100644 --- a/src/tds/codec/column_data.rs +++ b/src/tds/codec/column_data.rs @@ -13,21 +13,47 @@ mod float; mod guid; mod image; mod int; +#[cfg(test)] +mod legacy_codepages; +#[cfg(test)] +mod lossy_codepage; mod money; mod plp; +mod sql_variant; mod string; mod text; #[cfg(feature = "tds73")] mod time; +mod udt; mod var_len; mod xml; +/// Upper bound on how many bytes a value decoder will *pre-allocate* from a +/// server-supplied length field before it has read the corresponding data. +/// +/// The wire length is untrusted: a malformed or hostile server can claim a +/// value is up to `u32::MAX`/`u64::MAX` bytes long. Reserving that up front is a +/// memory-exhaustion vector (and, for `u64` lengths, can even exceed `Vec`'s +/// `isize::MAX` capacity limit and panic). Decoders therefore cap the initial +/// reservation to this value and let the buffer grow as bytes actually arrive; +/// a short/lying length still fails cleanly when the read runs out of input. +pub(crate) const MAX_PREALLOC: usize = 8192; // 8 KiB + +/// Absolute ceiling on the *total* size of a single PLP (partially +/// length-prefixed) value — `varchar(max)`, `nvarchar(max)`, `varbinary(max)`, +/// `xml`, and CLR UDTs. SQL Server's own MAX types top out at `2^31 - 1` bytes, +/// so any value that would grow past this is malformed. Without this bound the +/// "unknown length" PLP form (which streams an arbitrary number of chunks until +/// a zero-length terminator) lets a hostile server grow the accumulation buffer +/// without limit and OOM the client on a single column value. +pub(crate) const MAX_PLP_SIZE: usize = i32::MAX as usize; + use super::{Encode, FixedLenType, TypeInfo, VarLenType}; #[cfg(feature = "tds73")] use crate::tds::time::{Date, DateTime2, DateTimeOffset, Time}; use crate::{ tds::{time::DateTime, time::SmallDateTime, xml::XmlData, Numeric}, - SqlReadBytes, + FromSql, FromSqlOwned, IntoSql, SqlReadBytes, ToSql, }; use bytes::BufMut; pub(crate) use bytes_mut_with_type_info::BytesMutWithTypeInfo; @@ -36,7 +62,56 @@ use uuid::Uuid; const MAX_NVARCHAR_SIZE: usize = 1 << 30; +/// Number of days between `0001-01-01` (the `DateTime2`/`Date` epoch) and +/// `1900-01-01` (the `datetime`/`Datetimen` epoch). +#[cfg(feature = "tds73")] +const DAYS_YEAR_1_TO_1900: u32 = 693_595; + +/// Converts a [`DateTime2`] value into the legacy `datetime` ([`DateTime`]) +/// wire representation. +/// +/// This is used when bulk-inserting a `DateTime2`/`Date` value into a column +/// whose server-side type is `datetime` (`Datetimen`). The `datetime` type +/// counts days from `1900-01-01` and stores the time of day as 1/300-second +/// fragments, so the sub-second precision of the source value is degraded to +/// match. Returns a [`Conversion`] error if the date is earlier than +/// `1900-01-01`, which `datetime` cannot represent. +/// +/// [`Conversion`]: crate::Error::Conversion +#[cfg(feature = "tds73")] +fn datetime2_to_datetime(dt2: &DateTime2) -> crate::Result { + let dt2_days = dt2.date().days(); + + let days = dt2_days.checked_sub(DAYS_YEAR_1_TO_1900).ok_or_else(|| { + crate::Error::Conversion( + format!( + "invalid datetime, expecting a date not earlier than 1900-01-01 but got {} days after year 1", + dt2_days + ) + .into(), + ) + })? as i32; + + // `increments` are counted in 10^-scale seconds; convert to nanoseconds and + // then to the 1/300-second fragments used by `datetime`, degrading the + // sub-second precision in the process. + let time = dt2.time(); + // `Time::new` accepts any u8 scale, but a scale > 9 makes `9 - scale` + // underflow. Reject it rather than panic on hostile/garbage wire input. + let scale = time.scale(); + if scale > 9 { + return Err(crate::Error::Protocol( + format!("invalid datetime2 scale {scale}, expected 0..=9").into(), + )); + } + let nanos = time.increments() as u128 * 10u128.pow(9 - scale as u32); + let seconds_fragments = (nanos * 300 / 1_000_000_000) as u32; + + Ok(DateTime::new(days, seconds_fragments)) +} + #[derive(Clone, Debug, PartialEq)] +#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] /// A container of a value that can be represented as a TDS value. pub enum ColumnData<'a> { /// 8-bit integer, unsigned. @@ -68,19 +143,19 @@ pub enum ColumnData<'a> { /// A small DateTime value. SmallDateTime(Option), #[cfg(feature = "tds73")] - #[cfg_attr(feature = "docs", doc(cfg(feature = "tds73")))] + #[cfg_attr(docsrs, doc(cfg(feature = "tds73")))] /// Time value. Time(Option")))), + ) + .expect_err("xml must not encode as a sql_variant"); + + assert!(matches!(err, Error::Conversion(_))); + } + + // A `total_len` of exactly 2 is the smallest valid sql_variant (base type + + // prop-bytes count, zero property bytes, empty value). The guard is + // `total_len < 2`; mutating `<` to `<=`/`==` would reject this valid value. + #[tokio::test] + async fn decode_total_len_exactly_two() { + let mut buf = BytesMut::new(); + buf.put_u32_le(2); // total length + buf.put_u8(FixedLenType::Bit as u8); // base type + buf.put_u8(0); // prop bytes + buf.put_u8(1); // the Bit value (read by fixed_len::decode) + + let data = decode(&mut buf.into_sql_read_bytes()).await.unwrap(); + assert_eq!(data, ColumnData::Bit(Some(true))); + } + + // `data_len` of exactly MAX_VARIANT_PAYLOAD (8000) is allowed; the guard is + // `data_len > MAX_VARIANT_PAYLOAD`. Mutating `>` to `>=`/`==` would reject a + // value that is exactly at the limit. + #[tokio::test] + async fn decode_data_len_at_max_payload() { + let payload_bytes = vec![0x5Au8; MAX_VARIANT_PAYLOAD]; + + let mut payload = vec![VarLenType::BigVarBin as u8, 2]; + payload.extend_from_slice(&40u16.to_le_bytes()); // max length prop + payload.extend_from_slice(&payload_bytes); + + let data = decode(&mut variant_reader(&payload)).await.unwrap(); + assert_eq!(data, ColumnData::Binary(Some(payload_bytes.into()))); + } + + // The numeric-scale guard is `scale > 38`; a scale of exactly 38 is valid. + // Mutating `>` to `>=`/`==` would reject scale 38. + #[tokio::test] + async fn decode_numeric_scale_at_limit() { + let mut payload = vec![VarLenType::Numericn as u8, 2, 38, 38]; + payload.push(1); // positive sign + payload.extend_from_slice(&123u32.to_le_bytes()); + + let data = decode(&mut variant_reader(&payload)).await.unwrap(); + assert_eq!( + data, + ColumnData::Numeric(Some(Numeric::new_with_scale(123, 38))) + ); + } + + // A 12-byte magnitude is reconstructed as `low + high * (1 << 64)`. + // low = 5, high = 3 => 5 + 3 * 2^64. This kills the mutations of the three + // operators on that line: `+`->`-`/`*`, `*`->`+`/`/`, and `<<`->`>>` + // (which would give 5, 5*3*2^64, 5+(3+2^64), 5+3/2^64=5, and 5+3*1=8 + // respectively - all different from the true value). + #[tokio::test] + async fn decode_numeric_twelve_byte_magnitude() { + let mut payload = vec![VarLenType::Numericn as u8, 2, 38, 0]; + payload.push(1); // positive sign + payload.extend_from_slice(&5u64.to_le_bytes()); // low 8 bytes + payload.extend_from_slice(&3u32.to_le_bytes()); // high 4 bytes + + let expected = 5i128 + 3i128 * (1i128 << 64); + let data = decode(&mut variant_reader(&payload)).await.unwrap(); + assert_eq!( + data, + ColumnData::Numeric(Some(Numeric::new_with_scale(expected, 0))) + ); + } + + // A string whose UTF-16 encoding is exactly MAX_VARIANT_PAYLOAD (8000) + // bytes = 4000 BMP chars is allowed; the guard is `utf16.len() > MAX`. + // Mutating `>` to `>=`/`==` would reject a value exactly at the limit. + #[tokio::test] + async fn encode_string_at_max_payload() { + let s: String = "a".repeat(MAX_VARIANT_PAYLOAD / 2); + let mut buf = BytesMut::new(); + encode(&mut buf, ColumnData::String(Some(s.into()))) + .expect("a string exactly at the limit must encode"); + } + + // Binary of exactly MAX_VARIANT_PAYLOAD (8000) bytes is allowed; the guard + // is `bytes.len() > MAX`. Mutating `>` to `>=`/`==` would reject it. + #[tokio::test] + async fn encode_binary_at_max_payload() { + let bytes = vec![0u8; MAX_VARIANT_PAYLOAD]; + let mut buf = BytesMut::new(); + encode(&mut buf, ColumnData::Binary(Some(bytes.into()))) + .expect("binary exactly at the limit must encode"); + } + + /// Builds a reader from a raw buffer where `total_len` is set explicitly + /// (rather than derived from the payload), so the length-guard error arms + /// can be exercised. Covers lines 63-65, 72-78, 88-90, 98-104. + fn raw_reader(total_len: u32, rest: &[u8]) -> impl SqlReadBytes { + let mut buf = BytesMut::new(); + buf.put_u32_le(total_len); + buf.extend_from_slice(rest); + buf.into_sql_read_bytes() + } + + // total_len of 1 is non-zero but below the 2 byte minimum (base type + + // prop-bytes count). Covers 62-66. + #[tokio::test] + async fn decode_total_len_below_two_errors() { + let err = decode(&mut raw_reader(1, &[])).await.unwrap_err(); + assert!(matches!(err, Error::Protocol(_)), "got {err:?}"); + } + + // total_len smaller than 2 + prop_bytes is rejected. Covers 71-79. + #[tokio::test] + async fn decode_total_len_too_small_for_props_errors() { + // base type + prop count = 5, but total_len is only 2. + let err = decode(&mut raw_reader(2, &[FixedLenType::Int4 as u8, 5])) + .await + .unwrap_err(); + assert!(matches!(err, Error::Protocol(_)), "got {err:?}"); + } + + // data_len exceeding MAX_VARIANT_PAYLOAD is rejected before any allocation. + // Covers 87-91. + #[tokio::test] + async fn decode_data_len_over_max_payload_errors() { + // total_len 8005, prop_bytes 2 => data_len 8001 (> 8000). + let err = decode(&mut raw_reader(8005, &[VarLenType::BigVarBin as u8, 2])) + .await + .unwrap_err(); + assert!(matches!(err, Error::Protocol(_)), "got {err:?}"); + } + + // A fixed-length base type must not carry property bytes. Covers 97-105. + #[tokio::test] + async fn decode_fixed_type_with_props_errors() { + // base = Int4 (fixed), prop_bytes = 1, total_len 3 => data_len 0. + let err = decode(&mut raw_reader(3, &[FixedLenType::Int4 as u8, 1])) + .await + .unwrap_err(); + assert!(matches!(err, Error::Protocol(_)), "got {err:?}"); + } + + // A base type that is neither a known fixed nor var-len type is rejected. + // Covers 110-112. + #[tokio::test] + async fn decode_unknown_base_type_errors() { + // 0x00 is not a FixedLenType nor a VarLenType. + let err = decode(&mut variant_reader(&[0x00, 0])).await.unwrap_err(); + assert!(matches!(err, Error::Protocol(_)), "got {err:?}"); + } + + // An nchar/nvarchar value with an odd byte count cannot be a UTF-16 string. + // Covers 151-153. + #[tokio::test] + async fn decode_nchar_odd_length_errors() { + let mut payload = vec![VarLenType::NChar as u8, 7]; + payload.extend_from_slice(&0u32.to_le_bytes()); // collation info + payload.push(0); // sort id + payload.extend_from_slice(&40u16.to_le_bytes()); // max length + payload.extend_from_slice(&[1u8, 2, 3]); // 3 = odd value bytes + + let err = decode(&mut variant_reader(&payload)).await.unwrap_err(); + assert!(matches!(err, Error::Protocol(_)), "got {err:?}"); + } + + // A var-len base type with no sql_variant decode arm (here: Xml) is + // rejected via the `other` arm. Covers 203-207. + #[tokio::test] + async fn decode_unsupported_var_type_errors() { + let err = decode(&mut variant_reader(&[VarLenType::Xml as u8, 0])) + .await + .unwrap_err(); + assert!(matches!(err, Error::Protocol(_)), "got {err:?}"); + } + + // decode_numeric rejects a zero-length value. Covers 235-237. + #[tokio::test] + async fn decode_numeric_empty_value_errors() { + // prop_bytes 2 (precision + scale), no value bytes => data_len 0. + let payload = vec![VarLenType::Numericn as u8, 2, 18, 2]; + let err = decode(&mut variant_reader(&payload)).await.unwrap_err(); + assert!(matches!(err, Error::Protocol(_)), "got {err:?}"); + } + + // decode_numeric rejects a scale greater than 38. Covers 241-245. + #[tokio::test] + async fn decode_numeric_scale_too_large_errors() { + let mut payload = vec![VarLenType::Numericn as u8, 2, 18, 39]; // scale 39 + payload.push(1); // sign + payload.extend_from_slice(&123u32.to_le_bytes()); + + let err = decode(&mut variant_reader(&payload)).await.unwrap_err(); + assert!(matches!(err, Error::Protocol(_)), "got {err:?}"); + } + + // decode_numeric rejects a sign byte that is neither 0 nor 1. Covers 250. + #[tokio::test] + async fn decode_numeric_invalid_sign_errors() { + let mut payload = vec![VarLenType::Numericn as u8, 2, 18, 2]; + payload.push(5); // invalid sign + payload.extend_from_slice(&123u32.to_le_bytes()); + + let err = decode(&mut variant_reader(&payload)).await.unwrap_err(); + assert!(matches!(err, Error::Protocol(_)), "got {err:?}"); + } + + // An 8-byte magnitude is read as a little-endian u64. Covers 257. + #[tokio::test] + async fn decode_numeric_eight_byte_magnitude() { + let mut payload = vec![VarLenType::Numericn as u8, 2, 18, 0]; + payload.push(1); // positive sign + payload.extend_from_slice(&123_456_789_012u64.to_le_bytes()); + + let data = decode(&mut variant_reader(&payload)).await.unwrap(); + assert_eq!( + data, + ColumnData::Numeric(Some(Numeric::new_with_scale(123_456_789_012, 0))) + ); + } + + // A magnitude whose length is not 4/8/12/16 is rejected. Covers 268-272. + #[tokio::test] + async fn decode_numeric_bad_magnitude_length_errors() { + let mut payload = vec![VarLenType::Numericn as u8, 2, 18, 0]; + payload.push(1); // positive sign + payload.extend_from_slice(&[1u8, 2, 3, 4, 5]); // 5-byte magnitude + + let err = decode(&mut variant_reader(&payload)).await.unwrap_err(); + assert!(matches!(err, Error::Protocol(_)), "got {err:?}"); + } + + // A 16-byte magnitude whose reconstructed value (`low + high * 2^64`) + // exceeds `i128::MAX` must return a protocol error instead of overflowing + // (panic in debug / silent wraparound in release). Here `high` is + // `u64::MAX`, so `high * 2^64` alone is ~2^128, well past `i128::MAX`. + #[tokio::test] + async fn decode_numeric_sixteen_byte_magnitude_overflow_errors() { + let mut payload = vec![VarLenType::Numericn as u8, 2, 38, 0]; + payload.push(1); // positive sign + payload.extend_from_slice(&0u64.to_le_bytes()); // low 8 bytes + payload.extend_from_slice(&u64::MAX.to_le_bytes()); // high 8 bytes + + let err = decode(&mut variant_reader(&payload)).await.unwrap_err(); + assert!(matches!(err, Error::Protocol(_)), "got {err:?}"); + } + + // A 16-byte magnitude that fits in `i128` still decodes successfully, so the + // overflow guard does not reject valid large values. + #[tokio::test] + async fn decode_numeric_sixteen_byte_magnitude_in_range() { + let mut payload = vec![VarLenType::Numericn as u8, 2, 38, 0]; + payload.push(1); // positive sign + payload.extend_from_slice(&7u64.to_le_bytes()); // low 8 bytes + payload.extend_from_slice(&3u64.to_le_bytes()); // high 8 bytes + + let expected = 7i128 + 3i128 * (1i128 << 64); + let data = decode(&mut variant_reader(&payload)).await.unwrap(); + assert_eq!( + data, + ColumnData::Numeric(Some(Numeric::new_with_scale(expected, 0))) + ); + } + + // A string whose UTF-16 encoding exceeds MAX_VARIANT_PAYLOAD is rejected. + // Covers 381-390. + #[tokio::test] + async fn encode_string_over_max_payload_errors() { + // 4001 BMP chars => 8002 UTF-16 bytes (> 8000). + let s: String = "a".repeat(MAX_VARIANT_PAYLOAD / 2 + 1); + let mut buf = BytesMut::new(); + let err = encode(&mut buf, ColumnData::String(Some(s.into()))).unwrap_err(); + assert!(matches!(err, Error::Conversion(_)), "got {err:?}"); + } + + // Binary larger than MAX_VARIANT_PAYLOAD is rejected. Covers 403-412. + #[tokio::test] + async fn encode_binary_over_max_payload_errors() { + let bytes = vec![0u8; MAX_VARIANT_PAYLOAD + 1]; + let mut buf = BytesMut::new(); + let err = encode(&mut buf, ColumnData::Binary(Some(bytes.into()))).unwrap_err(); + assert!(matches!(err, Error::Conversion(_)), "got {err:?}"); + } + + // datetimeoffset value shorter than the mandatory 5 trailing bytes + // (datetime2 + 2 offset bytes) is rejected. Covers 192-195. + #[cfg(feature = "tds73")] + #[tokio::test] + async fn decode_datetimeoffset_too_short_errors() { + // prop_bytes 1 (scale), value = 4 bytes (< 5). + let mut payload = vec![VarLenType::DatetimeOffsetn as u8, 1, 0 /* scale */]; + payload.extend_from_slice(&[0u8; 4]); + + let err = decode(&mut variant_reader(&payload)).await.unwrap_err(); + assert!(matches!(err, Error::Protocol(_)), "got {err:?}"); + } + + // datetimeoffset whose computed time length overflows a u8 is rejected. + // Covers 196-198. + #[cfg(feature = "tds73")] + #[tokio::test] + async fn decode_datetimeoffset_time_len_too_large_errors() { + // prop_bytes 1 (scale), value = 261 bytes => time_len 256 (> u8::MAX). + let mut payload = vec![VarLenType::DatetimeOffsetn as u8, 1, 0 /* scale */]; + payload.extend_from_slice(&[0u8; 261]); + + let err = decode(&mut variant_reader(&payload)).await.unwrap_err(); + assert!(matches!(err, Error::Protocol(_)), "got {err:?}"); + } + + // A guid arm consumes exactly 16 value bytes; a peer declaring a different + // (here inflated) data_len must error rather than read 16 and desync the + // rest of the row. + #[tokio::test] + async fn decode_guid_wrong_data_len_errors() { + let uuid = uuid::Uuid::from_u128(0x0102030405060708090a0b0c0d0e0f10); + let mut wire = *uuid.as_bytes(); + guid::reorder_bytes(&mut wire); + + // total_len 22, base + prop count = 2, prop_bytes 0 => data_len 20 (!= 16). + let mut rest = vec![VarLenType::Guid as u8, 0]; + rest.extend_from_slice(&wire); + rest.extend_from_slice(&[0u8; 4]); // 4 extra bytes to match total_len 22 + + let err = decode(&mut raw_reader(22, &rest)).await.unwrap_err(); + assert!(matches!(err, Error::Protocol(_)), "got {err:?}"); + } + + // A guid must not declare any property bytes. + #[tokio::test] + async fn decode_guid_with_prop_bytes_errors() { + // total_len 19, prop_bytes 1 => data_len 16, but prop_bytes must be 0. + let mut rest = vec![VarLenType::Guid as u8, 1, 0xff /* stray prop */]; + rest.extend_from_slice(&[0u8; 16]); + + let err = decode(&mut raw_reader(19, &rest)).await.unwrap_err(); + assert!(matches!(err, Error::Protocol(_)), "got {err:?}"); + } + + // A char/varchar arm consumes exactly 7 property bytes (5 collation + 2 max + // length). A peer declaring a different propBytes count must error rather + // than read 7 and desync (data_len is derived from prop_bytes). + #[tokio::test] + async fn decode_bigvarchar_wrong_prop_bytes_errors() { + // prop_bytes declared as 5 (should be 7). + let mut payload = vec![VarLenType::BigVarChar as u8, 5]; + payload.extend_from_slice(&13632521u32.to_le_bytes()); + payload.push(52); + payload.extend_from_slice(b"abc"); + + let err = decode(&mut variant_reader(&payload)).await.unwrap_err(); + assert!(matches!(err, Error::Protocol(_)), "got {err:?}"); + } + + // A binary arm consumes exactly 2 property bytes (max length). A wrong + // propBytes count must error. + #[tokio::test] + async fn decode_binary_wrong_prop_bytes_errors() { + // prop_bytes declared as 7 (should be 2). + let mut payload = vec![VarLenType::BigVarBin as u8, 7]; + payload.extend_from_slice(&[0u8; 7]); + payload.extend_from_slice(&[1u8, 2, 3, 4]); + + let err = decode(&mut variant_reader(&payload)).await.unwrap_err(); + assert!(matches!(err, Error::Protocol(_)), "got {err:?}"); + } +} diff --git a/src/tds/codec/column_data/string.rs b/src/tds/codec/column_data/string.rs index 3a38794d4..70874f341 100644 --- a/src/tds/codec/column_data/string.rs +++ b/src/tds/codec/column_data/string.rs @@ -1,7 +1,5 @@ use std::borrow::Cow; -use byteorder::{ByteOrder, LittleEndian}; - use crate::{error::Error, sql_read_bytes::SqlReadBytes, tds::Collation, VarLenType}; pub(crate) async fn decode( @@ -20,15 +18,20 @@ where match (data, ty) { // Codepages other than UTF (Some(buf), BigChar) | (Some(buf), BigVarChar) => { - let collation = collation.as_ref().unwrap(); - let encoder = collation.encoding()?; + let collation = collation + .as_ref() + .ok_or_else(|| Error::Protocol("string column missing collation".into()))?; + let codec = collation.codec()?; - let s = encoder - .decode_without_bom_handling_and_without_replacement(buf.as_ref()) - .ok_or_else(|| Error::Encoding("invalid sequence".into()))? - .to_string(); + let value = if src.context().lossy_codepage() { + codec.decode_lossy(buf.as_ref()) + } else { + codec + .decode(buf.as_ref()) + .ok_or_else(|| Error::Encoding("invalid sequence".into()))? + }; - Ok(Some(s.into())) + Ok(Some(value.into())) } // UTF-16 (Some(buf), _) => { @@ -36,9 +39,99 @@ where return Err(Error::Protocol("nvarchar: invalid plp length".into())); } - let buf: Vec<_> = buf.chunks(2).map(LittleEndian::read_u16).collect(); - Ok(Some(String::from_utf16(&buf)?.into())) + // Decode UTF-16LE straight from the byte pairs, without first + // collecting an intermediate `Vec` (one fewer full-buffer + // allocation + copy per value). + let units = buf.chunks(2).map(|c| u16::from_le_bytes([c[0], c[1]])); + + // NVARCHAR/NCHAR may be decoded losslessly when the connection opted + // in (`Config::lossy_utf16_decoding`), replacing invalid surrogates + // with U+FFFD so legacy rows holding unchecked UCS-2 stay readable. + // XML (`ty == Xml`, routed here from `xml::decode`) is always strict, + // regardless of the flag. Either way the `buf.len() % 2` guard above + // already rejected desynced (odd) lengths. + if matches!(ty, NChar | NVarchar) && src.context().lossy_utf16() { + let s: String = char::decode_utf16(units) + .map(|r| r.unwrap_or(char::REPLACEMENT_CHARACTER)) + .collect(); + Ok(Some(s.into())) + } else { + // Strict: invalid surrogates error, matching the previous + // `String::from_utf16` behaviour. + let s = char::decode_utf16(units) + .collect::>() + .map_err(|_| Error::Protocol("nvarchar: invalid UTF-16 sequence".into()))?; + Ok(Some(s.into())) + } } _ => Ok(None), } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use bytes::{BufMut, BytesMut}; + + // Lossy NVARCHAR decoding replaces an unpaired surrogate with U+FFFD. + #[tokio::test] + async fn nvarchar_lossy_replaces_lone_surrogate() { + let mut buf = BytesMut::new(); + buf.put_u16_le(2); // fixed-size PLP length prefix + buf.put_u16_le(0xD800); // unpaired high surrogate + + let mut reader = buf.into_sql_read_bytes(); + reader.context_mut().set_lossy_utf16(true); + let value = decode(&mut reader, VarLenType::NVarchar, 40, None) + .await + .expect("lossy nvarchar must decode malformed UTF-16"); + assert_eq!(value.as_deref(), Some("\u{fffd}")); + } + + // Strict is the default: an unpaired surrogate is a protocol error. + #[tokio::test] + async fn nvarchar_strict_rejects_lone_surrogate() { + let mut buf = BytesMut::new(); + buf.put_u16_le(2); + buf.put_u16_le(0xD800); + + let err = decode( + &mut buf.into_sql_read_bytes(), + VarLenType::NVarchar, + 40, + None, + ) + .await + .expect_err("strict nvarchar must reject malformed UTF-16"); + assert!(matches!(err, Error::Protocol(_))); + } + + // A BigVarChar (non-UTF codepage) value with the collation omitted by the + // server must return a protocol error rather than panicking on `unwrap`. + #[tokio::test] + async fn decode_bigvarchar_missing_collation_errors() { + let mut buf = BytesMut::new(); + buf.put_u16_le(2); // fixed-size PLP length prefix + buf.put_slice(&[0x41, 0x42]); // "AB" in an 8-bit codepage + + let err = decode( + &mut buf.into_sql_read_bytes(), + VarLenType::BigVarChar, + 2, + None, + ) + .await + .expect_err("missing collation must error, not panic"); + + match err { + Error::Protocol(msg) => { + assert!( + msg.contains("missing collation"), + "unexpected protocol message: {msg}" + ); + } + other => panic!("expected a protocol error, got {other:?}"), + } + } +} diff --git a/src/tds/codec/column_data/text.rs b/src/tds/codec/column_data/text.rs index 0c454d251..34bcc2c2d 100644 --- a/src/tds/codec/column_data/text.rs +++ b/src/tds/codec/column_data/text.rs @@ -13,9 +13,9 @@ where return Ok(ColumnData::String(None)); } - for _ in 0..ptr_len { - src.read_u8().await?; - } + // Skip the text pointer (packet-aware bulk read into a throwaway buffer). + let mut ptr = Vec::new(); + crate::sql_read_bytes::read_bytes_into(src, &mut ptr, ptr_len, super::MAX_PREALLOC).await?; src.read_i32_le().await?; // days src.read_u32_le().await?; // second fractions @@ -23,31 +23,175 @@ where let text = match collation { // TEXT Some(collation) => { - let encoder = collation.encoding()?; + let codec = collation.codec()?; let text_len = src.read_u32_le().await? as usize; - let mut buf = Vec::with_capacity(text_len); + let mut buf = Vec::new(); + crate::sql_read_bytes::read_bytes_into(src, &mut buf, text_len, super::MAX_PREALLOC) + .await?; - for _ in 0..text_len { - buf.push(src.read_u8().await?); + if src.context().lossy_codepage() { + codec.decode_lossy(buf.as_ref()) + } else { + codec + .decode(buf.as_ref()) + .ok_or_else(|| Error::Encoding("invalid sequence".into()))? } - - encoder - .decode_without_bom_handling_and_without_replacement(buf.as_ref()) - .ok_or_else(|| Error::Encoding("invalid sequence".into()))? - .to_string() } // NTEXT None => { - let text_len = src.read_u32_le().await? as usize / 2; - let mut buf = Vec::with_capacity(text_len); - - for _ in 0..text_len { - buf.push(src.read_u16_le().await?); + let byte_len = src.read_u32_le().await? as usize; + // NTEXT is UTF-16: the byte length must be even. An odd length would + // desync the stream (the final `read_u16_le` would consume a byte + // from the next field), so reject it as a protocol error. + if !byte_len.is_multiple_of(2) { + return Err(Error::Protocol( + format!("ntext: odd byte length {byte_len} is invalid").into(), + )); } + // Bulk-read the raw UTF-16LE bytes, then decode them as u16 code + // units (like string.rs), instead of one packet-aware u16 per poll. + let mut raw = Vec::new(); + crate::sql_read_bytes::read_bytes_into(src, &mut raw, byte_len, super::MAX_PREALLOC) + .await?; + + // `byte_len` is guaranteed even (odd lengths are rejected above), + // so every 2-byte chunk is a complete UTF-16LE code unit. + let buf: Vec = raw + .chunks(2) + .map(|c| u16::from_le_bytes([c[0], c[1]])) + .collect(); - String::from_utf16(&buf[..])? + // NTEXT may be decoded losslessly when the connection opted in + // (`Config::lossy_utf16_decoding`), replacing invalid surrogates + // with U+FFFD so legacy rows holding unchecked UCS-2 stay readable. + // The odd-length guard above still fires in both modes; it is a + // framing (desync) error, not merely bad Unicode. + if src.context().lossy_utf16() { + String::from_utf16_lossy(&buf[..]) + } else { + String::from_utf16(&buf[..])? + } } }; Ok(ColumnData::String(Some(text.into()))) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use bytes::{BufMut, BytesMut}; + + #[tokio::test] + async fn decode_null_when_ptr_len_zero() { + let mut buf = BytesMut::new(); + buf.put_u8(0); + + let data = decode(&mut buf.into_sql_read_bytes(), None).await.unwrap(); + assert_eq!(data, ColumnData::String(None)); + } + + #[tokio::test] + async fn decode_ntext_reads_utf16_payload() { + let mut buf = BytesMut::new(); + buf.put_u8(1); // ptr_len + buf.put_u8(0xAA); // pointer byte (ignored) + buf.put_i32_le(0); // days + buf.put_u32_le(0); // second fractions + buf.put_u32_le(4); // byte length of the UTF-16 text (2 chars) + buf.put_u16_le('h' as u16); + buf.put_u16_le('i' as u16); + + let data = decode(&mut buf.into_sql_read_bytes(), None).await.unwrap(); + assert_eq!(data, ColumnData::String(Some("hi".into()))); + } + + #[tokio::test] + async fn decode_ntext_rejects_odd_byte_length() { + // An odd UTF-16 byte length must be a protocol error, not a desync. + let mut buf = BytesMut::new(); + buf.put_u8(1); // ptr_len + buf.put_u8(0xAA); // pointer byte (ignored) + buf.put_i32_le(0); // days + buf.put_u32_le(0); // second fractions + buf.put_u32_le(3); // odd byte length + buf.put_u16_le('h' as u16); + buf.put_u8(0); + + let err = decode(&mut buf.into_sql_read_bytes(), None) + .await + .expect_err("odd ntext length must be rejected"); + assert!(matches!(err, Error::Protocol(_))); + } + + #[tokio::test] + async fn decode_ntext_lossy_replaces_lone_surrogate() { + let mut buf = BytesMut::new(); + buf.put_u8(1); // ptr_len + buf.put_u8(0xAA); // pointer byte (ignored) + buf.put_i32_le(0); // days + buf.put_u32_le(0); // second fractions + buf.put_u32_le(2); // byte length (one lone surrogate) + buf.put_u16_le(0xD800); // unpaired high surrogate + + let mut reader = buf.into_sql_read_bytes(); + reader.context_mut().set_lossy_utf16(true); + let data = decode(&mut reader, None).await.unwrap(); + assert_eq!(data, ColumnData::String(Some("\u{fffd}".into()))); + } + + #[tokio::test] + async fn decode_ntext_strict_rejects_lone_surrogate() { + let mut buf = BytesMut::new(); + buf.put_u8(1); + buf.put_u8(0xAA); + buf.put_i32_le(0); + buf.put_u32_le(0); + buf.put_u32_le(2); + buf.put_u16_le(0xD800); + + // Default (strict) context: malformed UTF-16 must error. + let err = decode(&mut buf.into_sql_read_bytes(), None) + .await + .expect_err("strict ntext must reject malformed UTF-16"); + assert!(matches!(err, Error::Utf16)); + } + + #[tokio::test] + async fn decode_ntext_lossy_still_rejects_odd_length() { + // The odd-length framing guard fires even when lossy decoding is on. + let mut buf = BytesMut::new(); + buf.put_u8(1); + buf.put_u8(0xAA); + buf.put_i32_le(0); + buf.put_u32_le(0); + buf.put_u32_le(3); // odd byte length + buf.put_u16_le('h' as u16); + buf.put_u8(0); + + let mut reader = buf.into_sql_read_bytes(); + reader.context_mut().set_lossy_utf16(true); + let err = decode(&mut reader, None) + .await + .expect_err("odd ntext length must be rejected even when lossy"); + assert!(matches!(err, Error::Protocol(_))); + } + + #[tokio::test] + async fn decode_text_uses_collation_encoding() { + let mut buf = BytesMut::new(); + buf.put_u8(1); // ptr_len + buf.put_u8(0xAA); + buf.put_i32_le(0); + buf.put_u32_le(0); + buf.put_u32_le(2); // 2 raw bytes in codepage encoding + buf.put_slice(b"hi"); + + let collation = crate::tds::Collation::new(0x0409, 0); // WINDOWS_1252 + let data = decode(&mut buf.into_sql_read_bytes(), Some(collation)) + .await + .unwrap(); + assert_eq!(data, ColumnData::String(Some("hi".into()))); + } +} diff --git a/src/tds/codec/column_data/time.rs b/src/tds/codec/column_data/time.rs index 350fa9a6c..a52cccb16 100644 --- a/src/tds/codec/column_data/time.rs +++ b/src/tds/codec/column_data/time.rs @@ -8,11 +8,65 @@ where let time = match rlen { 0 => ColumnData::Time(None), - _ => { + // A `time` value is 3..=5 bytes on the wire (MS-TDS §2.2.5.5.1.2); the + // exact width depends on the scale. Bound the server-supplied `rlen` to + // that range here, mirroring the sibling temporal decoders, instead of + // deferring the check to `Time::decode`. + 3..=5 => { let time = Time::decode(src, len, rlen as usize).await?; ColumnData::Time(Some(time)) } + _ => { + return Err(crate::Error::Protocol( + format!("timen: invalid value length {rlen}").into(), + )) + } }; Ok(time) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use bytes::{BufMut, BytesMut}; + + #[tokio::test] + async fn zero_length_is_null() { + let mut buf = BytesMut::new(); + buf.put_u8(0); + let v = decode(&mut buf.into_sql_read_bytes(), 2).await.unwrap(); + assert!(matches!(v, ColumnData::Time(None))); + } + + #[tokio::test] + async fn rejects_out_of_range_rlen() { + // rlen of 7 is outside the legal 3..=5 range for a `time` value and must + // be a protocol error, not read as an oversized value. + let mut buf = BytesMut::new(); + buf.put_u8(7); + let err = decode(&mut buf.into_sql_read_bytes(), 2) + .await + .expect_err("rlen out of range must be rejected"); + assert!(matches!(err, crate::Error::Protocol(_))); + } + + #[tokio::test] + async fn decodes_normal_value() { + // scale (len) = 2 => 3-byte value: u16 low + u8 high. + let mut buf = BytesMut::new(); + buf.put_u8(3); // rlen + buf.put_u16_le(0x0102); + buf.put_u8(0x03); + + let v = decode(&mut buf.into_sql_read_bytes(), 2).await.unwrap(); + match v { + ColumnData::Time(Some(t)) => { + assert_eq!(t.increments(), 0x0102 | (0x03u64 << 16)); + assert_eq!(t.scale(), 2); + } + other => panic!("expected Time(Some(..)), got {other:?}"), + } + } +} diff --git a/src/tds/codec/column_data/udt.rs b/src/tds/codec/column_data/udt.rs new file mode 100644 index 000000000..c9d368b49 --- /dev/null +++ b/src/tds/codec/column_data/udt.rs @@ -0,0 +1,63 @@ +use std::borrow::Cow; + +use crate::{sql_read_bytes::SqlReadBytes, ColumnData}; + +/// Decode the value of a CLR user-defined type (UDT) column. +/// +/// UDT values are always transferred using the partially length-prefixed (PLP) +/// byte-stream format (MS-TDS §2.2.5.5.4). tiberius does not attempt to +/// deserialize the CLR representation; the raw serialized bytes are surfaced +/// verbatim as [`ColumnData::Binary`]. +pub(crate) async fn decode(src: &mut R) -> crate::Result> +where + R: SqlReadBytes + Unpin, +{ + // Force the PLP (u64-prefixed) code path, which is how UDT values are + // always encoded on the wire regardless of the declared max byte size. + let data = super::plp::decode(src, 0xffff_ffff).await?.map(Cow::Owned); + + Ok(ColumnData::Binary(data)) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use bytes::{BufMut, BytesMut}; + + #[tokio::test] + async fn decode_udt_plp_bytes() { + let payload: &[u8] = &[0xde, 0xad, 0xbe, 0xef]; + + let mut buf = BytesMut::new(); + // PLP: unknown total length sentinel. + buf.put_u64_le(0xfffffffffffffffe); + // One chunk carrying the payload. + buf.put_u32_le(payload.len() as u32); + buf.extend_from_slice(payload); + // PLP terminator. + buf.put_u32_le(0); + + let data = decode(&mut buf.into_sql_read_bytes()) + .await + .expect("decode must succeed"); + + match data { + ColumnData::Binary(Some(bytes)) => assert_eq!(bytes.as_ref(), payload), + other => panic!("expected Binary, got {:?}", other), + } + } + + #[tokio::test] + async fn decode_udt_null() { + let mut buf = BytesMut::new(); + // PLP NULL sentinel. + buf.put_u64_le(0xffffffffffffffff); + + let data = decode(&mut buf.into_sql_read_bytes()) + .await + .expect("decode must succeed"); + + assert!(matches!(data, ColumnData::Binary(None))); + } +} diff --git a/src/tds/codec/column_data/var_len.rs b/src/tds/codec/column_data/var_len.rs index 20f6a953a..c91e93819 100644 --- a/src/tds/codec/column_data/var_len.rs +++ b/src/tds/codec/column_data/var_len.rs @@ -1,4 +1,6 @@ -use crate::{sql_read_bytes::SqlReadBytes, tds::codec::VarLenContext, ColumnData, VarLenType}; +use crate::{ + sql_read_bytes::SqlReadBytes, tds::codec::VarLenContext, ColumnData, Error, VarLenType, +}; pub(crate) async fn decode( src: &mut R, @@ -41,8 +43,134 @@ where Text => super::text::decode(src, collation).await?, NText => super::text::decode(src, None).await?, Image => super::image::decode(src).await?, - t => unimplemented!("{:?}", t), + SSVariant => super::sql_variant::decode(src).await?, + t => { + return Err(Error::Protocol( + format!("unsupported column type: {:?}", t).into(), + )) + } }; Ok(res) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use bytes::{BufMut, BytesMut}; + + #[tokio::test] + async fn decode_bitn_true() { + let mut buf = BytesMut::new(); + buf.put_u8(1); // recv_len + buf.put_u8(1); // true + + let ctx = VarLenContext::new(VarLenType::Bitn, 1, None); + let data = decode(&mut buf.into_sql_read_bytes(), &ctx).await.unwrap(); + assert_eq!(data, ColumnData::Bit(Some(true))); + } + + #[tokio::test] + async fn decode_intn_null() { + let mut buf = BytesMut::new(); + buf.put_u8(0); // recv_len 0 -> null + + let ctx = VarLenContext::new(VarLenType::Intn, 4, None); + let data = decode(&mut buf.into_sql_read_bytes(), &ctx).await.unwrap(); + assert_eq!(data, ColumnData::I32(None)); + } + + #[tokio::test] + async fn decode_guid_null() { + let mut buf = BytesMut::new(); + buf.put_u8(0); + + let ctx = VarLenContext::new(VarLenType::Guid, 16, None); + let data = decode(&mut buf.into_sql_read_bytes(), &ctx).await.unwrap(); + assert_eq!(data, ColumnData::Guid(None)); + } + + #[tokio::test] + async fn decode_nvarchar_value() { + let mut buf = BytesMut::new(); + buf.put_u16_le(4); // 4 bytes of UTF-16 + buf.put_u16_le('a' as u16); + buf.put_u16_le('b' as u16); + + let ctx = VarLenContext::new(VarLenType::NVarchar, 100, None); + let data = decode(&mut buf.into_sql_read_bytes(), &ctx).await.unwrap(); + assert_eq!(data, ColumnData::String(Some("ab".into()))); + } + + #[tokio::test] + async fn decode_money_null() { + let mut buf = BytesMut::new(); + buf.put_u8(0); // len byte read inside decode() + + let ctx = VarLenContext::new(VarLenType::Money, 8, None); + let data = decode(&mut buf.into_sql_read_bytes(), &ctx).await.unwrap(); + assert_eq!(data, ColumnData::F64(None)); + } + + #[tokio::test] + async fn decode_datetimen_null_smalldatetime() { + let mut buf = BytesMut::new(); + buf.put_u8(0); // rlen == 0 + + let ctx = VarLenContext::new(VarLenType::Datetimen, 4, None); + let data = decode(&mut buf.into_sql_read_bytes(), &ctx).await.unwrap(); + assert_eq!(data, ColumnData::SmallDateTime(None)); + } + + #[tokio::test] + async fn decode_text_null() { + let mut buf = BytesMut::new(); + buf.put_u8(0); // ptr_len 0 -> null + + let ctx = VarLenContext::new(VarLenType::Text, 0, None); + let data = decode(&mut buf.into_sql_read_bytes(), &ctx).await.unwrap(); + assert_eq!(data, ColumnData::String(None)); + } + + #[tokio::test] + async fn decode_ntext_null() { + let mut buf = BytesMut::new(); + buf.put_u8(0); + + let ctx = VarLenContext::new(VarLenType::NText, 0, None); + let data = decode(&mut buf.into_sql_read_bytes(), &ctx).await.unwrap(); + assert_eq!(data, ColumnData::String(None)); + } + + #[tokio::test] + async fn decode_image_null() { + let mut buf = BytesMut::new(); + buf.put_u8(0); + + let ctx = VarLenContext::new(VarLenType::Image, 0, None); + let data = decode(&mut buf.into_sql_read_bytes(), &ctx).await.unwrap(); + assert_eq!(data, ColumnData::Binary(None)); + } + + #[tokio::test] + async fn decode_ssvariant_null() { + let mut buf = BytesMut::new(); + buf.put_u32_le(0); // total_len 0 -> null + + let ctx = VarLenContext::new(VarLenType::SSVariant, 0, None); + let data = decode(&mut buf.into_sql_read_bytes(), &ctx).await.unwrap(); + assert_eq!(data, ColumnData::String(None)); + } + + #[tokio::test] + async fn decode_unsupported_type_errors() { + let buf = BytesMut::new(); + + let ctx = VarLenContext::new(VarLenType::Udt, 0, None); + let err = decode(&mut buf.into_sql_read_bytes(), &ctx) + .await + .unwrap_err(); + assert!(format!("{}", err).contains("unsupported column type")); + } +} diff --git a/src/tds/codec/column_data/xml.rs b/src/tds/codec/column_data/xml.rs index 34b191390..2b44cac82 100644 --- a/src/tds/codec/column_data/xml.rs +++ b/src/tds/codec/column_data/xml.rs @@ -28,3 +28,39 @@ where Ok(ColumnData::Xml(xml)) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use bytes::{BufMut, BytesMut}; + + // A fixed-size PLP length of 0xffff marks a NULL value. + #[tokio::test] + async fn decode_null_xml() { + let mut buf = BytesMut::new(); + buf.put_u16_le(0xffff); + + let data = decode(&mut buf.into_sql_read_bytes(), 8, None) + .await + .unwrap(); + assert_eq!(data, ColumnData::Xml(None)); + } + + #[tokio::test] + async fn decode_present_xml() { + let mut buf = BytesMut::new(); + buf.put_u16_le(4); // 2 UTF-16 code units => 4 bytes + buf.put_u16_le('h' as u16); + buf.put_u16_le('i' as u16); + + let data = decode(&mut buf.into_sql_read_bytes(), 8, None) + .await + .unwrap(); + + match data { + ColumnData::Xml(Some(xml)) => assert_eq!(xml.to_string(), "hi"), + other => panic!("expected present XML, got {other:?}"), + } + } +} diff --git a/src/tds/codec/decode.rs b/src/tds/codec/decode.rs index d19fec0c9..30bfc9ec8 100644 --- a/src/tds/codec/decode.rs +++ b/src/tds/codec/decode.rs @@ -20,14 +20,28 @@ impl Decoder for PacketCodec { return Ok(None); } - let header = PacketHeader::decode(&mut BytesMut::from(&src[0..HEADER_BYTES]))?; - let length = header.length() as usize; + // Peek the packet length directly from the buffered header instead of + // allocating a throwaway `BytesMut` and fully decoding the header just + // to read one field. The length is a big-endian `u16` at offset 2..4 + // (see `PacketHeader::encode`); the full header is decoded below once we + // know the whole packet is buffered. + let length = u16::from_be_bytes([src[2], src[3]]) as usize; if src.len() < length { src.reserve(length); return Ok(None); } + // Reject a malformed short-length packet *before* `PacketHeader::decode` + // consumes the 8 header bytes. `length` was already peeked inline above, + // so this check needs no decoded header; performing it first leaves the + // buffer untouched for a bogus length instead of eating the header. + if length < HEADER_BYTES { + return Err(Error::Protocol("Invalid packet length".into())); + } + + let header = PacketHeader::decode(src)?; + event!( Level::TRACE, "Reading a {:?} ({} bytes)", @@ -35,12 +49,6 @@ impl Decoder for PacketCodec { length, ); - let header = PacketHeader::decode(src)?; - - if length < HEADER_BYTES { - return Err(Error::Protocol("Invalid packet length".into())); - } - let payload = src.split_to(length - HEADER_BYTES); Ok(Some(Packet::new(header, payload))) @@ -53,12 +61,131 @@ impl Decoder for PacketCodec { if buf.is_empty() { Ok(None) } else { - Err( - std::io::Error::new(std::io::ErrorKind::Other, "bytes remaining on stream") - .into(), - ) + Err(std::io::Error::other("bytes remaining on stream").into()) } } } } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::tds::codec::{Encode, PacketHeader, PacketType}; + + // A full SQLBatch packet: 8-byte header + "hello world" payload => length 19. + fn full_packet_bytes() -> BytesMut { + let payload = BytesMut::from(&b"hello world"[..]); + let packet = Packet::new(PacketHeader::batch(1), payload); + + let mut buf = BytesMut::new(); + packet.encode(&mut buf).unwrap(); + // Sanity: 8 header + 11 payload. + assert_eq!(buf.len(), 19); + buf + } + + #[test] + fn decode_partial_header_returns_none() { + // Fewer than HEADER_BYTES available: we must wait for more bytes and + // never index into `&src[0..HEADER_BYTES]`. + let mut src = BytesMut::from(&full_packet_bytes()[0..4]); + let mut codec = PacketCodec; + + let out = codec.decode(&mut src).unwrap(); + assert!(out.is_none()); + } + + #[test] + fn decode_complete_packet_returns_some() { + let mut src = full_packet_bytes(); + let mut codec = PacketCodec; + + let packet = codec + .decode(&mut src) + .unwrap() + .expect("a complete packet must decode to Some"); + + assert_eq!(packet.header.r#type() as u8, PacketType::SQLBatch as u8); + // payload length must be `length - HEADER_BYTES` = 19 - 8 = 11. + assert_eq!(&packet.payload[..], b"hello world"); + assert_eq!(packet.payload.len(), 11); + // The whole packet is consumed. + assert!(src.is_empty()); + } + + #[test] + fn decode_incomplete_body_returns_none() { + // Full header present (declares length 19) but only part of the body is + // buffered: must return None rather than splitting past the buffer end. + let mut src = BytesMut::from(&full_packet_bytes()[0..11]); + let mut codec = PacketCodec; + + let out = codec.decode(&mut src).unwrap(); + assert!(out.is_none()); + } + + #[test] + fn decode_minimal_packet_exactly_header_bytes() { + // An empty-payload packet is exactly HEADER_BYTES (8) long with a + // declared length of 8. This is the boundary for both length checks: + // `src.len() < HEADER_BYTES` and `length < HEADER_BYTES` must be false. + let packet = Packet::new(PacketHeader::attention(1), BytesMut::new()); + let mut src = BytesMut::new(); + packet.encode(&mut src).unwrap(); + assert_eq!(src.len(), 8); + + let mut codec = PacketCodec; + let packet = codec + .decode(&mut src) + .unwrap() + .expect("an 8-byte packet must decode to Some"); + + assert_eq!( + packet.header.r#type() as u8, + PacketType::AttentionSignal as u8 + ); + assert!(packet.payload.is_empty()); + } + + #[test] + fn decode_rejects_length_below_header() { + // Declare a total length smaller than the header itself. We have enough + // bytes buffered, so we reach the `length < HEADER_BYTES` guard, which + // must error rather than underflow `length - HEADER_BYTES`. + let packet = Packet::new(PacketHeader::attention(1), BytesMut::new()); + let mut src = BytesMut::new(); + packet.encode(&mut src).unwrap(); + // Overwrite the BE length field (bytes 2..4) with 5 (< HEADER_BYTES). + src[2] = 0; + src[3] = 5; + + let mut codec = PacketCodec; + let err = codec + .decode(&mut src) + .expect_err("length below header must error"); + assert!(matches!(err, Error::Protocol(_))); + } + + #[test] + fn decode_eof_returns_some_for_complete_packet() { + let mut src = full_packet_bytes(); + let mut codec = PacketCodec; + + let packet = codec + .decode_eof(&mut src) + .unwrap() + .expect("decode_eof must yield a complete packet"); + assert_eq!(packet.header.r#type() as u8, PacketType::SQLBatch as u8); + assert_eq!(&packet.payload[..], b"hello world"); + } + + #[test] + fn decode_eof_errors_on_trailing_partial_bytes() { + // A partial packet at EOF (no full frame, buffer not empty) is an error. + let mut src = BytesMut::from(&full_packet_bytes()[0..4]); + let mut codec = PacketCodec; + + assert!(codec.decode_eof(&mut src).is_err()); + } +} diff --git a/src/tds/codec/encode.rs b/src/tds/codec/encode.rs index 7c76420f1..dfdee0e3f 100644 --- a/src/tds/codec/encode.rs +++ b/src/tds/codec/encode.rs @@ -1,4 +1,4 @@ -use super::{Packet, PacketCodec}; +use super::{AllHeaderTy, Packet, PacketCodec, ALL_HEADERS_LEN_TX}; use asynchronous_codec::Encoder; use bytes::{BufMut, BytesMut}; @@ -6,8 +6,49 @@ pub(crate) trait Encode { fn encode(self, dst: &mut B) -> crate::Result<()>; } +/// Encodes a `B_VARCHAR` (MS-TDS 2.2.5.1.2): a single-byte count of UTF-16 code +/// units followed by the string encoded as little-endian UCS-2. +/// +/// The length prefix is a single byte, so a string longer than 255 UTF-16 code +/// units cannot be represented. Truncating the count with `as u8` while still +/// writing every unit would desync the wire, so an over-long string is rejected +/// with an [`Error::Protocol`](crate::Error::Protocol) instead. +pub(crate) fn encode_b_varchar(dst: &mut BytesMut, s: &str) -> crate::Result<()> { + let units: Vec = s.encode_utf16().collect(); + + if units.len() > u8::MAX as usize { + return Err(crate::Error::Protocol( + format!( + "string is too long for a B_VARCHAR ({} UTF-16 code units, max 255)", + units.len() + ) + .into(), + )); + } + + dst.put_u8(units.len() as u8); + + for unit in units { + dst.put_u16_le(unit); + } + + Ok(()) +} + +/// Writes the `ALL_HEADERS` block carrying the transaction descriptor that +/// precedes SQLBatch, RPC and Transaction Manager request payloads (MS-TDS +/// 2.2.5.3 / 2.2.5.3.1). The single header is the `TransactionDescriptor` +/// header with an outstanding-request count of 1. +pub(crate) fn encode_all_headers_tx(dst: &mut BytesMut, transaction_desc: [u8; 8]) { + dst.put_u32_le(ALL_HEADERS_LEN_TX as u32); + dst.put_u32_le(ALL_HEADERS_LEN_TX as u32 - 4); + dst.put_u16_le(AllHeaderTy::TransactionDescriptor as u16); + dst.put_slice(&transaction_desc); + dst.put_u32_le(1); +} + impl Encoder for PacketCodec { - type Item = Packet; + type Item<'a> = Packet; type Error = crate::Error; fn encode(&mut self, item: Packet, dst: &mut BytesMut) -> Result<(), Self::Error> { @@ -15,3 +56,85 @@ impl Encoder for PacketCodec { Ok(()) } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::tds::codec::{PacketHeader, PacketType}; + + #[test] + fn encode_writes_header_and_payload_to_dst() { + let payload = BytesMut::from(&b"abcd"[..]); + let packet = Packet::new(PacketHeader::batch(1), payload); + + let mut dst = BytesMut::new(); + let mut codec = PacketCodec; + codec.encode(packet, &mut dst).expect("encode must succeed"); + + // 8-byte header + 4-byte payload; a no-op encode would leave dst empty. + assert_eq!(dst.len(), 12); + assert_eq!(dst[0], PacketType::SQLBatch as u8); + // Total length is patched into the BE length field (bytes 2..4). + assert_eq!(&dst[2..4], &12u16.to_be_bytes()); + assert_eq!(&dst[8..], b"abcd"); + } + + #[test] + fn encode_b_varchar_is_byte_exact() { + let mut dst = BytesMut::new(); + encode_b_varchar(&mut dst, "Hi").unwrap(); + + let mut expected = Vec::new(); + expected.push(2u8); // two UTF-16 code units + expected.extend_from_slice(&('H' as u16).to_le_bytes()); + expected.extend_from_slice(&('i' as u16).to_le_bytes()); + + assert_eq!(&dst[..], &expected[..]); + } + + #[test] + fn encode_b_varchar_empty_writes_zero_length() { + let mut dst = BytesMut::new(); + encode_b_varchar(&mut dst, "").unwrap(); + assert_eq!(&dst[..], &[0u8]); + } + + // A B_VARCHAR length prefix is a single byte; exactly 255 code units is the + // boundary of what is representable and must still encode. + #[test] + fn encode_b_varchar_accepts_255_units() { + let name = "a".repeat(255); + let mut dst = BytesMut::new(); + encode_b_varchar(&mut dst, &name).expect("255-unit string must encode"); + + assert_eq!(dst[0], 255); + assert_eq!(dst.len(), 1 + 255 * 2); + } + + // 256+ code units cannot fit the u8 count; the encoder must error rather than + // truncate the count with `as u8` (which would keep writing every unit and + // desync the wire). + #[test] + fn encode_b_varchar_rejects_256_units() { + let name = "a".repeat(256); + let mut dst = BytesMut::new(); + let err = encode_b_varchar(&mut dst, &name).unwrap_err(); + assert!(matches!(err, crate::Error::Protocol(_)), "got {err:?}"); + } + + #[test] + fn encode_all_headers_tx_is_byte_exact() { + let td = [1u8, 2, 3, 4, 5, 6, 7, 8]; + let mut dst = BytesMut::new(); + encode_all_headers_tx(&mut dst, td); + + let mut expected = Vec::new(); + expected.extend_from_slice(&22u32.to_le_bytes()); // ALL_HEADERS_LEN_TX + expected.extend_from_slice(&18u32.to_le_bytes()); // header length (len - 4) + expected.extend_from_slice(&2u16.to_le_bytes()); // TransactionDescriptor type + expected.extend_from_slice(&td); // transaction descriptor + expected.extend_from_slice(&1u32.to_le_bytes()); // outstanding request count + + assert_eq!(&dst[..], &expected[..]); + } +} diff --git a/src/tds/codec/guid.rs b/src/tds/codec/guid.rs index 1298997ec..5d05a0f3e 100644 --- a/src/tds/codec/guid.rs +++ b/src/tds/codec/guid.rs @@ -8,3 +8,32 @@ pub(crate) fn reorder_bytes(bytes: &mut uuid::Bytes) { bytes.swap(4, 5); bytes.swap(6, 7); } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn reorder_bytes_swaps_the_guid_groups() { + // Swaps within the first three groups (0<->3, 1<->2, 4<->5, 6<->7); the + // trailing 8 bytes are left in place. + let mut bytes: uuid::Bytes = [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15]; + reorder_bytes(&mut bytes); + assert_eq!( + bytes, + [3, 2, 1, 0, 5, 4, 7, 6, 8, 9, 10, 11, 12, 13, 14, 15] + ); + } + + #[test] + fn reorder_bytes_is_its_own_inverse() { + let original: uuid::Bytes = [ + 10, 20, 30, 40, 50, 60, 70, 80, 90, 100, 110, 120, 130, 140, 150, 160, + ]; + let mut bytes = original; + reorder_bytes(&mut bytes); + assert_ne!(bytes, original); + reorder_bytes(&mut bytes); + assert_eq!(bytes, original); + } +} diff --git a/src/tds/codec/header.rs b/src/tds/codec/header.rs index 719fc158b..c9c0b62c1 100644 --- a/src/tds/codec/header.rs +++ b/src/tds/codec/header.rs @@ -56,8 +56,15 @@ pub(crate) struct PacketHeader { } impl PacketHeader { + /// Builds a packet header with the given wire `length` (including the 8 + /// header bytes) and packet `id`. + /// + /// # Panics + /// + /// Panics if `length` exceeds [`u16::MAX`], as the TDS length field is a + /// 16-bit value and cannot represent a larger packet. pub fn new(length: usize, id: u8) -> PacketHeader { - assert!(length <= u16::max_value() as usize); + assert!(length <= u16::MAX as usize); PacketHeader { ty: PacketType::TDSv7Login, status: PacketStatus::ResetConnection, @@ -86,7 +93,20 @@ impl PacketHeader { pub fn login(id: u8) -> Self { Self { - ty: PacketType::TDSv7Login, + // `ty` is inherited from `new()`, which already defaults to + // `TDSv7Login`; only the status differs here. + status: PacketStatus::EndOfMessage, + ..Self::new(0, id) + } + } + + // Only the Windows integrated-auth (winauth) login path sends a standalone + // SSPI packet; every other auth path (including unix GSSAPI) wraps the token + // in a login packet. Gate the constructor so it compiles only there. + #[cfg(all(windows, feature = "winauth"))] + pub fn sspi(id: u8) -> Self { + Self { + ty: PacketType::Sspi, status: PacketStatus::EndOfMessage, ..Self::new(0, id) } @@ -108,10 +128,35 @@ impl PacketHeader { } } + /// A client-to-server Attention Signal packet (packet type `0x06`, + /// MS-TDS section 2.2.1.6). The message carries no payload, so it is + /// always a single, end-of-message packet used to request cancellation + /// of the request currently in flight on the connection. + pub fn attention(id: u8) -> Self { + Self { + ty: PacketType::AttentionSignal, + status: PacketStatus::EndOfMessage, + ..Self::new(0, id) + } + } + + pub fn transaction_manager(id: u8) -> Self { + Self { + ty: PacketType::TransactionManagerReq, + status: PacketStatus::EndOfMessage, + ..Self::new(0, id) + } + } + pub fn set_status(&mut self, status: PacketStatus) { self.status = status; } + #[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + ))] pub fn set_type(&mut self, ty: PacketType) { self.ty = ty; } @@ -124,6 +169,17 @@ impl PacketHeader { self.ty } + // Outside of tests, the only caller is the TLS pre-login wrapper; a build + // with no TLS backend never reads it (the packet codec peeks the length + // field directly). Gate the allow so TLS builds still flag genuine disuse. + #[cfg_attr( + not(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + )), + allow(dead_code) + )] pub fn length(&self) -> u16 { self.length } @@ -171,3 +227,390 @@ impl Decode for PacketHeader { Ok(header) } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::tds::codec::Packet; + use bytes::BytesMut; + + // --- constructors ----------------------------------------------------- + + #[test] + fn new_sets_defaults() { + let header = PacketHeader::new(42, 7); + assert_eq!(header.ty, PacketType::TDSv7Login); + assert_eq!(header.status, PacketStatus::ResetConnection); + assert_eq!(header.length, 42); + assert_eq!(header.length(), 42); + assert_eq!(header.id, 7); + assert_eq!(header.spid, 0); + assert_eq!(header.window, 0); + } + + #[test] + fn new_accepts_max_length() { + let header = PacketHeader::new(u16::MAX as usize, 0); + assert_eq!(header.length, u16::MAX); + } + + #[test] + #[should_panic] + fn new_rejects_oversized_length() { + let _ = PacketHeader::new(u16::MAX as usize + 1, 0); + } + + #[test] + fn pre_login_constructor() { + let header = PacketHeader::pre_login(1); + assert_eq!(header.r#type(), PacketType::PreLogin); + assert_eq!(header.status(), PacketStatus::EndOfMessage); + assert_eq!(header.id, 1); + } + + #[test] + fn login_constructor() { + let header = PacketHeader::login(2); + assert_eq!(header.r#type(), PacketType::TDSv7Login); + assert_eq!(header.status(), PacketStatus::EndOfMessage); + assert_eq!(header.id, 2); + } + + #[test] + fn rpc_constructor() { + let header = PacketHeader::rpc(3); + assert_eq!(header.r#type(), PacketType::Rpc); + assert_eq!(header.status(), PacketStatus::NormalMessage); + assert_eq!(header.id, 3); + } + + #[test] + fn batch_constructor() { + let header = PacketHeader::batch(4); + assert_eq!(header.r#type(), PacketType::SQLBatch); + assert_eq!(header.status(), PacketStatus::NormalMessage); + assert_eq!(header.id, 4); + } + + #[test] + fn bulk_load_constructor() { + let header = PacketHeader::bulk_load(5); + assert_eq!(header.r#type(), PacketType::BulkLoad); + assert_eq!(header.status(), PacketStatus::NormalMessage); + assert_eq!(header.id, 5); + } + + // The `sspi` constructor only exists on the Windows integrated-auth path, so + // its test must be gated identically to the constructor itself. + #[cfg(all(windows, feature = "winauth"))] + #[test] + fn sspi_constructor() { + let header = PacketHeader::sspi(9); + assert_eq!(header.r#type(), PacketType::Sspi); + assert_eq!(header.status(), PacketStatus::EndOfMessage); + assert_eq!(header.length(), 0); + assert_eq!(header.id, 9); + } + + // --- mutators --------------------------------------------------------- + + #[test] + fn set_status_updates_status() { + let mut header = PacketHeader::batch(0); + assert_eq!(header.status(), PacketStatus::NormalMessage); + header.set_status(PacketStatus::EndOfMessage); + assert_eq!(header.status(), PacketStatus::EndOfMessage); + } + + #[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + ))] + #[test] + fn set_type_updates_type() { + let mut header = PacketHeader::login(0); + header.set_type(PacketType::PreLogin); + assert_eq!(header.r#type(), PacketType::PreLogin); + } + + // --- encode / decode round-trips ------------------------------------- + + fn round_trip(header: PacketHeader) -> PacketHeader { + let mut buf = BytesMut::new(); + header.encode(&mut buf).expect("encode"); + // The header is exactly 8 bytes on the wire [2.2.3.1]. + assert_eq!(buf.len(), 8); + PacketHeader::decode(&mut buf).expect("decode") + } + + fn assert_same(a: PacketHeader, b: PacketHeader) { + assert_eq!(a.ty, b.ty); + assert_eq!(a.status, b.status); + assert_eq!(a.length, b.length); + assert_eq!(a.spid, b.spid); + assert_eq!(a.id, b.id); + assert_eq!(a.window, b.window); + } + + #[test] + fn encode_writes_fields_big_endian() { + let mut header = PacketHeader::new(0x0102, 0xAB); + header.ty = PacketType::PreLogin; + header.status = PacketStatus::EndOfMessage; + header.spid = 0x0304; + header.window = 0xCD; + + let mut buf = BytesMut::new(); + header.encode(&mut buf).expect("encode"); + + assert_eq!( + &buf[..], + &[ + PacketType::PreLogin as u8, + PacketStatus::EndOfMessage as u8, + 0x01, + 0x02, // length, big-endian + 0x03, + 0x04, // spid, big-endian + 0xAB, // id + 0xCD, // window + ] + ); + } + + #[test] + fn round_trip_preserves_all_fields() { + let mut header = PacketHeader::new(1234, 200); + header.ty = PacketType::TabularResult; + header.status = PacketStatus::NormalMessage; + header.spid = 4321; + header.window = 99; + + assert_same(header, round_trip(header)); + } + + #[test] + fn round_trip_preserves_end_of_message_flag() { + let decoded = round_trip(PacketHeader::pre_login(1)); + assert_eq!(decoded.status(), PacketStatus::EndOfMessage); + assert_eq!(decoded.r#type(), PacketType::PreLogin); + } + + #[test] + fn round_trip_all_packet_types() { + for ty in [ + PacketType::SQLBatch, + PacketType::Rpc, + PacketType::TabularResult, + PacketType::AttentionSignal, + PacketType::BulkLoad, + PacketType::Fat, + PacketType::TransactionManagerReq, + PacketType::TDSv7Login, + PacketType::Sspi, + PacketType::PreLogin, + ] { + let mut header = PacketHeader::new(64, 1); + header.ty = ty; + assert_eq!(round_trip(header).ty, ty); + } + } + + #[test] + fn round_trip_all_status_flags() { + for status in [ + PacketStatus::NormalMessage, + PacketStatus::EndOfMessage, + PacketStatus::IgnoreEvent, + PacketStatus::ResetConnection, + PacketStatus::ResetConnectionSkipTran, + ] { + let mut header = PacketHeader::new(8, 0); + header.status = status; + assert_eq!(round_trip(header).status, status); + } + } + + #[test] + fn round_trip_length_boundaries() { + for length in [0usize, 1, 8, 255, 256, u16::MAX as usize] { + let header = PacketHeader::new(length, 0); + assert_eq!(round_trip(header).length(), length as u16); + } + } + + #[test] + fn round_trip_packet_id_range() { + for id in [0u8, 1, 127, 128, 254, 255] { + let header = PacketHeader::batch(id); + assert_eq!(round_trip(header).id, id); + } + } + + #[test] + fn decode_rejects_invalid_packet_type() { + let mut buf = BytesMut::from(&[0xFFu8, 1, 0, 8, 0, 0, 0, 0][..]); + assert!(PacketHeader::decode(&mut buf).is_err()); + } + + #[test] + fn decode_rejects_invalid_status() { + // 0x02 is not a valid PacketStatus. + let mut buf = BytesMut::from(&[PacketType::SQLBatch as u8, 0x02, 0, 8, 0, 0, 0, 0][..]); + assert!(PacketHeader::decode(&mut buf).is_err()); + } + + // --- attention signal ------------------------------------------------- + + #[test] + fn attention_packet_header_fields() { + let header = PacketHeader::attention(42); + + assert_eq!(header.r#type() as u8, PacketType::AttentionSignal as u8); + assert_eq!(header.r#type() as u8, 0x06); + // An attention message is always a single end-of-message packet. + assert_eq!(header.status(), PacketStatus::EndOfMessage); + } + + #[test] + fn attention_packet_header_encodes_to_eight_bytes() { + let header = PacketHeader::attention(42); + + let mut buf = BytesMut::new(); + header.encode(&mut buf).unwrap(); + + // 8-byte fixed header, no payload. + assert_eq!(buf.len(), 8); + assert_eq!( + &buf[..], + &[ + 0x06, // type: attention signal + 0x01, // status: end of message + 0x00, 0x00, // length (patched by Packet::encode) + 0x00, 0x00, // spid + 42, // packet id + 0x00, // window + ] + ); + } + + #[test] + fn attention_packet_encodes_with_length_of_header() { + // A full attention packet has an empty payload, so the wire length + // is exactly the 8 header bytes. + let packet = Packet::new(PacketHeader::attention(1), BytesMut::new()); + + let mut buf = BytesMut::new(); + packet.encode(&mut buf).unwrap(); + + assert_eq!(&buf[..], &[0x06, 0x01, 0x00, 0x08, 0x00, 0x00, 0x01, 0x00]); + } + + #[test] + fn new_sets_login_type_and_reset_connection_status() { + let header = PacketHeader::new(123, 5); + + assert_eq!(header.r#type() as u8, PacketType::TDSv7Login as u8); + assert_eq!(header.status(), PacketStatus::ResetConnection); + assert_eq!(header.length(), 123); + } + + #[test] + #[should_panic(expected = "length <= u16::MAX as usize")] + fn new_panics_on_length_overflow() { + PacketHeader::new(usize::from(u16::MAX) + 1, 0); + } + + #[test] + fn rpc_header_type_and_status() { + let header = PacketHeader::rpc(7); + assert_eq!(header.r#type() as u8, PacketType::Rpc as u8); + assert_eq!(header.status(), PacketStatus::NormalMessage); + } + + #[test] + fn pre_login_header_type_and_status() { + let header = PacketHeader::pre_login(7); + assert_eq!(header.r#type() as u8, PacketType::PreLogin as u8); + assert_eq!(header.status(), PacketStatus::EndOfMessage); + } + + #[test] + fn login_header_type_and_status() { + let header = PacketHeader::login(7); + assert_eq!(header.r#type() as u8, PacketType::TDSv7Login as u8); + assert_eq!(header.status(), PacketStatus::EndOfMessage); + } + + // `PacketHeader::sspi` only exists on the Windows integrated-auth path. + #[cfg(all(windows, feature = "winauth"))] + #[test] + fn sspi_header_type_and_status() { + let header = PacketHeader::sspi(7); + // Both fields differ from the `PacketHeader::new` defaults + // (TDSv7Login / ResetConnection), so a deleted field would be caught. + assert_eq!(header.r#type() as u8, PacketType::Sspi as u8); + assert_eq!(header.status(), PacketStatus::EndOfMessage); + } + + #[test] + fn batch_header_type_and_status() { + let header = PacketHeader::batch(7); + assert_eq!(header.r#type() as u8, PacketType::SQLBatch as u8); + assert_eq!(header.status(), PacketStatus::NormalMessage); + } + + #[test] + fn bulk_load_header_type_and_status() { + let header = PacketHeader::bulk_load(7); + assert_eq!(header.r#type() as u8, PacketType::BulkLoad as u8); + assert_eq!(header.status(), PacketStatus::NormalMessage); + } + + #[test] + fn transaction_manager_header_type_and_status() { + let header = PacketHeader::transaction_manager(7); + assert_eq!( + header.r#type() as u8, + PacketType::TransactionManagerReq as u8 + ); + assert_eq!(header.status(), PacketStatus::EndOfMessage); + } + + #[test] + fn set_status_mutates_header() { + let mut header = PacketHeader::batch(1); + assert_eq!(header.status(), PacketStatus::NormalMessage); + + header.set_status(PacketStatus::IgnoreEvent); + assert_eq!(header.status(), PacketStatus::IgnoreEvent); + } + + #[test] + fn decode_round_trips_header_fields() { + let header = PacketHeader::rpc(9); + + let mut buf = BytesMut::new(); + header.encode(&mut buf).unwrap(); + + let decoded = PacketHeader::decode(&mut buf).unwrap(); + assert_eq!(decoded.r#type() as u8, PacketType::Rpc as u8); + assert_eq!(decoded.status(), PacketStatus::NormalMessage); + assert_eq!(decoded.length(), 0); + } + + #[test] + fn decode_invalid_packet_type_errors() { + let mut buf = BytesMut::from(&[0xffu8, 0x01, 0x00, 0x08, 0x00, 0x00, 0x01, 0x00][..]); + let err = PacketHeader::decode(&mut buf).unwrap_err(); + assert!(format!("{}", err).contains("invalid packet type")); + } + + #[test] + fn decode_invalid_packet_status_errors() { + let mut buf = BytesMut::from(&[0x01u8, 0xff, 0x00, 0x08, 0x00, 0x00, 0x01, 0x00][..]); + let err = PacketHeader::decode(&mut buf).unwrap_err(); + assert!(format!("{}", err).contains("invalid packet status")); + } +} diff --git a/src/tds/codec/iterator_ext.rs b/src/tds/codec/iterator_ext.rs index aecdd6d5a..b8a160fab 100644 --- a/src/tds/codec/iterator_ext.rs +++ b/src/tds/codec/iterator_ext.rs @@ -25,3 +25,20 @@ where out } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn join_interleaves_separator_between_elements() { + let joined = [1, 2, 3].into_iter().join(", "); + assert_eq!(joined, "1, 2, 3"); + } + + #[test] + fn join_single_element_has_no_separator() { + let joined = std::iter::once("only").join(", "); + assert_eq!(joined, "only"); + } +} diff --git a/src/tds/codec/login.rs b/src/tds/codec/login.rs index 0ecc3d2fd..f842d14da 100644 --- a/src/tds/codec/login.rs +++ b/src/tds/codec/login.rs @@ -3,12 +3,14 @@ use byteorder::{LittleEndian, WriteBytesExt}; use bytes::BytesMut; use enumflags2::{bitflags, BitFlags}; use io::{Cursor, Write}; +use secrecy::{ExposeSecret, SecretString}; use std::fmt::Debug; use std::{borrow::Cow, io}; +use zeroize::{Zeroize, Zeroizing}; uint_enum! { #[repr(u32)] - #[derive(PartialOrd)] + #[derive(PartialOrd, Default)] pub enum FeatureLevel { SqlServerV7 = 0x70000000, SqlServer2000 = 0x71000000, @@ -17,16 +19,11 @@ uint_enum! { SqlServer2008 = 0x730A0003, SqlServer2008R2 = 0x730B0003, /// 2012, 2014, 2016 + #[default] SqlServerN = 0x74000004, } } -impl Default for FeatureLevel { - fn default() -> Self { - Self::SqlServerN - } -} - impl FeatureLevel { pub fn done_row_count_bytes(self) -> u8 { if self as u32 >= FeatureLevel::SqlServer2005 as u32 { @@ -134,17 +131,43 @@ pub(crate) const FEA_EXT_TERMINATOR: u8 = 0xFFu8; pub(crate) const FED_AUTH_LIBRARYSECURITYTOKEN: u8 = 0x01; /// https://docs.microsoft.com/en-us/openspecs/windows_protocols/ms-tds/773a62b6-ee89-4c02-9e5e-344882630aac -#[derive(Debug, Clone, Default)] -#[cfg_attr(test, derive(PartialEq, Eq))] -struct FedAuthExt<'a> { +#[derive(Clone, Default)] +struct FedAuthExt { fed_auth_echo: bool, - fed_auth_token: Cow<'a, str>, + fed_auth_token: SecretString, nonce: Option<[u8; 32]>, } +// `SecretString` has no `PartialEq`; this test-only impl compares the exposed +// token so the round-trip tests keep working. +#[cfg(test)] +impl PartialEq for FedAuthExt { + fn eq(&self, other: &Self) -> bool { + self.fed_auth_echo == other.fed_auth_echo + && self.nonce == other.nonce + && self.fed_auth_token.expose_secret() == other.fed_auth_token.expose_secret() + } +} + +#[cfg(test)] +impl Eq for FedAuthExt {} + +// The `fed_auth_token` is a `secrecy::SecretString`, so it self-redacts in +// `Debug` output; a derive would keep the AAD bearer token safe on that field. +// This manual impl is retained only to summarize `nonce` as ``/`None` +// rather than dumping the raw nonce bytes a derive would print. +impl std::fmt::Debug for FedAuthExt { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("FedAuthExt") + .field("fed_auth_echo", &self.fed_auth_echo) + .field("fed_auth_token", &self.fed_auth_token) + .field("nonce", &self.nonce.map(|_| "")) + .finish() + } +} + /// the login packet -#[derive(Debug, Clone, Default)] -#[cfg_attr(test, derive(PartialEq, Eq))] +#[derive(Clone, Default)] pub struct LoginMessage<'a> { /// the highest TDS version the client supports tds_version: FeatureLevel, @@ -167,12 +190,75 @@ pub struct LoginMessage<'a> { client_lcid: u32, hostname: Cow<'a, str>, username: Cow<'a, str>, - password: Cow<'a, str>, + // Credentials are stored as `SecretString`: zeroized on drop and redacted + // from `Debug`. The plaintext is exposed only at the point its bytes are + // written into the LOGIN7 buffer (see `encode_to_boxed_slice`). + password: SecretString, app_name: Cow<'a, str>, server_name: Cow<'a, str>, /// the default database to connect to db_name: Cow<'a, str>, - fed_auth_ext: Option>, + fed_auth_ext: Option, +} + +// `SecretString` has no `PartialEq`; this test-only impl compares the exposed +// password so the encode/decode round-trip tests keep working. +#[cfg(test)] +impl PartialEq for LoginMessage<'_> { + fn eq(&self, other: &Self) -> bool { + self.tds_version == other.tds_version + && self.packet_size == other.packet_size + && self.client_prog_ver == other.client_prog_ver + && self.client_pid == other.client_pid + && self.connection_id == other.connection_id + && self.option_flags_1 == other.option_flags_1 + && self.option_flags_2 == other.option_flags_2 + && self.integrated_security == other.integrated_security + && self.type_flags == other.type_flags + && self.option_flags_3 == other.option_flags_3 + && self.client_timezone == other.client_timezone + && self.client_lcid == other.client_lcid + && self.hostname == other.hostname + && self.username == other.username + && self.app_name == other.app_name + && self.server_name == other.server_name + && self.db_name == other.db_name + && self.fed_auth_ext == other.fed_auth_ext + && self.password.expose_secret() == other.password.expose_secret() + } +} + +#[cfg(test)] +impl Eq for LoginMessage<'_> {} + +// Kept as a hand-written impl for defense-in-depth even though `password` +// self-redacts via `SecretString` (`[REDACTED]`) and the output now matches a +// derive. Enumerating every field explicitly forces any future field to be +// consciously handled here; a derive would silently print a newly-added secret. +impl std::fmt::Debug for LoginMessage<'_> { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("LoginMessage") + .field("tds_version", &self.tds_version) + .field("packet_size", &self.packet_size) + .field("client_prog_ver", &self.client_prog_ver) + .field("client_pid", &self.client_pid) + .field("connection_id", &self.connection_id) + .field("option_flags_1", &self.option_flags_1) + .field("option_flags_2", &self.option_flags_2) + .field("integrated_security", &self.integrated_security) + .field("type_flags", &self.type_flags) + .field("option_flags_3", &self.option_flags_3) + .field("client_timezone", &self.client_timezone) + .field("client_lcid", &self.client_lcid) + .field("hostname", &self.hostname) + .field("username", &self.username) + .field("password", &self.password) + .field("app_name", &self.app_name) + .field("server_name", &self.server_name) + .field("db_name", &self.db_name) + .field("fed_auth_ext", &self.fed_auth_ext) + .finish() + } } impl<'a> LoginMessage<'a> { @@ -183,11 +269,67 @@ impl<'a> LoginMessage<'a> { option_flags_2: OptionFlag2::InitLangFatal | OptionFlag2::OdbcDriver, option_flags_3: BitFlags::from_flag(OptionFlag3::UnknownCollationHandling), app_name: "tiberius".into(), + hostname: Self::get_hostname(), ..Default::default() } } - #[cfg(any(all(unix, feature = "integrated-auth-gssapi"), windows, feature = "winauth"))] + /// Best-effort local workstation id (machine hostname), used as the default + /// login `hostname`. Returns an empty string if it cannot be determined. + fn get_hostname() -> Cow<'static, str> { + #[cfg(windows)] + fn get_computer_name() -> io::Result { + extern "system" { + // https://learn.microsoft.com/en-us/windows/win32/api/winbase/nf-winbase-getcomputernamew + fn GetComputerNameW(lpBuffer: *mut u16, nSize: *mut u32) -> i32; + } + + // MAX_COMPUTERNAME_LENGTH is 15, plus 1 for the null terminator. + let mut buffer = [0u16; 15 + 1]; + let mut size = buffer.len() as u32; + let result = unsafe { GetComputerNameW(buffer.as_mut_ptr(), &mut size) }; + if result == 0 { + let lerr = io::Error::last_os_error(); + tracing::error!("GetComputerNameW failed: {lerr}"); + Err(lerr) + } else { + Ok(String::from_utf16_lossy(&buffer[..size as usize])) + } + } + + #[cfg(target_family = "unix")] + fn get_computer_name() -> io::Result { + // POSIX gethostname() may or may not null-terminate on truncation, + // so we split on the first NUL (falling back to the whole buffer). + let mut buffer = [0u8; 255 + 1]; + let result = unsafe { + libc::gethostname(buffer.as_mut_ptr() as *mut _, buffer.len() as libc::size_t) + }; + if result != 0 { + let lerr = io::Error::last_os_error(); + tracing::error!("gethostname failed: {lerr}"); + Err(lerr) + } else { + match buffer.split(|b| *b == 0).next() { + Some(hostname) => Ok(String::from_utf8_lossy(hostname).into_owned()), + None => Ok(String::from_utf8_lossy(&buffer).into_owned()), + } + } + } + + #[cfg(not(any(windows, target_family = "unix")))] + fn get_computer_name() -> io::Result { + Ok(String::new()) + } + + get_computer_name().map(Cow::Owned).unwrap_or_default() + } + + #[cfg(any( + all(unix, any(feature = "integrated-auth-gssapi", feature = "sspi-rs")), + windows, + feature = "winauth" + ))] pub fn integrated_security(&mut self, bytes: Option>) { if bytes.is_some() { self.option_flags_2.insert(OptionFlag2::IntegratedSecurity); @@ -210,17 +352,22 @@ impl<'a> LoginMessage<'a> { self.server_name = server_name.into(); } + /// Sets the client / workstation name reported to the server. + pub fn hostname(&mut self, hostname: impl Into>) { + self.hostname = hostname.into(); + } + pub fn user_name(&mut self, user_name: impl Into>) { self.username = user_name.into(); } - pub fn password(&mut self, password: impl Into>) { - self.password = password.into(); + pub fn password(&mut self, password: impl Into) { + self.password = crate::client::auth::secret_from_string(password.into()); } pub fn aad_token( &mut self, - token: impl Into>, + token: impl Into, fed_auth_echo: bool, nonce: Option<[u8; 32]>, ) { @@ -228,7 +375,7 @@ impl<'a> LoginMessage<'a> { self.fed_auth_ext = Some(FedAuthExt { fed_auth_echo, - fed_auth_token: token.into(), + fed_auth_token: crate::client::auth::secret_from_string(token.into()), nonce, }) } @@ -240,11 +387,78 @@ impl<'a> LoginMessage<'a> { self.type_flags.remove(LoginTypeFlag::ReadOnlyIntent); } } -} -impl<'a> Encode for LoginMessage<'a> { - fn encode(self, dst: &mut BytesMut) -> crate::Result<()> { - let mut cursor = Cursor::new(Vec::with_capacity(512)); + /// Sets the requested TDS packet size. + pub fn packet_size(&mut self, size: u32) { + self.packet_size = size; + } + + /// Exact number of bytes [`Self::encode_to_boxed_slice`] will write, i.e. the final + /// length of the LOGIN7 buffer. + /// + /// This is used to reserve the whole buffer up front so it never reallocates + /// while the (obfuscated) password lives inside it — see the security note + /// in `encode_to_boxed_slice`. The layout mirrors the writes in `encode_to_boxed_slice` + /// exactly; if that layout changes this must change with it (the capacity + /// assertion at the end of `encode_to_boxed_slice` guards against drift). + fn encoded_len(&self) -> usize { + // Fixed prefix written before any variable-length data. This equals the + // initial `data_offset` computed in `encode_to_boxed_slice`: + // 4 (length) + 5 * 4 (header u32s) + 4 (flag bytes) + 2 * 4 (tz + lcid) + // = 36 bytes of fixed header, then + // var_data.len() (13) * 2 * 2 offset/length table entries + 6 + // (2 extra ClientId bytes + 4-byte cbSSPILong) + // = 36 + 52 + 6 = 94. + const FIXED_OVERHEAD: usize = 94; + + // Every variable-length string is encoded as UTF-16 (2 bytes/unit). + fn utf16_bytes(s: &str) -> usize { + s.encode_utf16().count() * 2 + } + + let mut len = FIXED_OVERHEAD; + len += utf16_bytes(&self.hostname); + len += utf16_bytes(&self.username); + len += utf16_bytes(self.password.expose_secret()); + len += utf16_bytes(&self.app_name); + len += utf16_bytes(&self.server_name); + len += utf16_bytes(&self.db_name); + + if let Some(ref bytes) = self.integrated_security { + len += bytes.len(); + } + + if let Some(ref ext) = self.fed_auth_ext { + // 4 (FeatureExt data offset) + 1 (FEA_EXT_FEDAUTH) + 4 (feature ext + // length) + 1 (options) + 4 (token length) + token bytes + nonce + // + 1 (FEA_EXT_TERMINATOR). + len += 15 + utf16_bytes(ext.fed_auth_token.expose_secret()); + if ext.nonce.is_some() { + len += 32; + } + } + + len + } + + pub(crate) fn encode_to_boxed_slice(self) -> crate::Result>> { + // SECURITY (password zeroization): the password is written into this + // buffer (only lightly obfuscated with a trivially reversible transform) + // and the returned `Vec` is wrapped in `Zeroizing` so it is wiped on + // drop. That wipe only covers the buffer's *current* heap allocation. If + // the `Vec` were to reallocate *after* the password bytes were written + // (e.g. because a later field such as db_name or the fed-auth token grew + // it past its capacity), the old allocation would be freed WITHOUT being + // zeroized, leaving a recoverable plaintext-equivalent copy of the + // password in freed heap. + // + // To make that impossible we reserve the exact final size up front, + // before writing any variable-length data, so no reallocation can occur + // during encoding. The `assert_eq!` at the end verifies the capacity + // never changed, so any future change that breaks this invariant fails + // loudly instead of silently leaking a password copy. + let mut cursor = Cursor::new(Vec::with_capacity(self.encoded_len())); + let reserved_capacity = cursor.get_ref().capacity(); // Space for the length cursor.write_u32::(0)?; @@ -263,21 +477,27 @@ impl<'a> Encode for LoginMessage<'a> { cursor.write_u32::(self.client_timezone as u32)?; cursor.write_u32::(self.client_lcid)?; - // variable length data (OffsetLength) - let var_data = [ + // variable length data (OffsetLength). Expose the password only here, as + // a short-lived `&str` borrowed for the duration of this encode; the + // resulting bytes land in the `Zeroizing` buffer returned below. + let password = self.password.expose_secret(); + // `13` is the number of TDS LOGIN7 variable-length `OffsetLength` fields + // assembled in `var_data`; it must equal the element count of the + // literal below. The annotation is kept for the mixed-element coercion. + let var_data: [&str; 13] = [ &self.hostname, &self.username, - &self.password, + password, &self.app_name, &self.server_name, - &"".into(), // 5. ibExtension - &"".into(), // ibCltIntName - &"".into(), // ibLanguage + "", // 5. ibExtension + "", // ibCltIntName + "", // ibLanguage &self.db_name, - &"".into(), // 9. ClientId (6 bytes); this is included in var_data so we don't lack the bytes of cbSspiLong (4=2*2) and can insert it at the correct position - &"".into(), // 10. ibSSPI - &"".into(), // ibAtchDBFile - &"".into(), // ibChangePassword + "", // 9. ClientId (6 bytes); this is included in var_data so we don't lack the bytes of cbSspiLong (4=2*2) and can insert it at the correct position + "", // 10. ibSSPI + "", // ibAtchDBFile + "", // ibChangePassword ]; let mut data_offset = cursor.position() as usize + var_data.len() * 2 * 2 + 6; @@ -289,10 +509,11 @@ impl<'a> Encode for LoginMessage<'a> { fea_ext_offset = cursor.position(); } - // write the client ID (created from the MAC address) + // Client ID field: a fixed placeholder (not derived from the + // MAC address). SQL Server does not require a real value here. if i == 9 { - cursor.write_u32::(0)?; //TODO: - cursor.write_u16::(42)?; //TODO: generate real client id + cursor.write_u32::(0)?; + cursor.write_u16::(42)?; continue; } @@ -362,11 +583,30 @@ impl<'a> Encode for LoginMessage<'a> { cursor.write_u8(FEA_EXT_FEDAUTH)?; - let mut token = Cursor::new(Vec::new()); - for codepoint in fed_auth_ext.fed_auth_token.encode_utf16() { + // SECURITY (fed-auth token): the token is a bearer credential. Like + // the password buffer above, reserve its exact final size up front + // (one UTF-16 code unit is 2 bytes) so this temporary buffer never + // reallocates while holding the token — a realloc would free the old + // allocation without zeroizing it, leaking a recoverable copy. It is + // wrapped in `Zeroizing` so it is wiped on drop as well. + // Expose the fed-auth token only to size and write its bytes; they + // go into the `Zeroizing` buffer and the local `token` Vec is + // wrapped in `Zeroizing` below. + let fed_auth_token = fed_auth_ext.fed_auth_token.expose_secret(); + let token_capacity = fed_auth_token.encode_utf16().count() * 2; + let mut token = Cursor::new(Vec::with_capacity(token_capacity)); + for codepoint in fed_auth_token.encode_utf16() { token.write_u16::(codepoint)?; } - let token = token.into_inner(); + // Wrap in `Zeroizing` so the finished token buffer is also wiped on + // drop. (`Cursor` cannot be built over a `Zeroizing` inner because + // `Write` is only implemented for a fixed set of inner types.) + let mut token = Zeroizing::new(token.into_inner()); + debug_assert_eq!( + token.capacity(), + token_capacity, + "fed-auth token buffer reallocated: a copy may remain in freed heap" + ); // options (1) + TokenLength(4) + Token.length + nonce.length let feature_ext_length = @@ -383,6 +623,7 @@ impl<'a> Encode for LoginMessage<'a> { cursor.write_u32::(token.len() as u32)?; cursor.write_all(token.as_slice())?; + token.zeroize(); if let Some(nonce) = fed_auth_ext.nonce { cursor.write_all(nonce.as_ref())?; @@ -394,7 +635,33 @@ impl<'a> Encode for LoginMessage<'a> { cursor.set_position(0); cursor.write_u32::(cursor.get_ref().len() as u32)?; - dst.extend(cursor.into_inner()); + // The password lived in this buffer; a reallocation here would have + // leaked an un-zeroized copy into freed heap (see the security note + // above). Reserving `encoded_len()` up front must have prevented any + // growth — verify the invariant held. + assert_eq!( + cursor.get_ref().capacity(), + reserved_capacity, + "LOGIN7 encode buffer reallocated during encoding: a plaintext-equivalent \ + password copy may have been left in freed heap. `encoded_len()` under-reserved." + ); + + // The capacity now provably equals the length (asserted above), so + // `into_boxed_slice()` will NOT shrink-reallocate — a shrink realloc + // would free the current allocation without zeroizing it, leaking the + // very password copy this buffer is protecting. Returning a boxed slice + // also means the finished buffer can no longer be grown by a caller. + Ok(Zeroizing::new(cursor.into_inner().into_boxed_slice())) + } +} + +impl<'a> Encode for LoginMessage<'a> { + fn encode(self, dst: &mut BytesMut) -> crate::Result<()> { + // `encoded` is `Zeroizing>`; it is wiped on drop at the end of + // this function, immediately after the copy into `dst`, so no explicit + // `zeroize()` is needed here. + let encoded = self.encode_to_boxed_slice()?; + dst.extend_from_slice(&encoded[..]); Ok(()) } @@ -558,6 +825,44 @@ mod tests { } } + #[test] + fn readonly_intent_sets_type_flag_bit() { + // The TypeFlags byte is the third of the four flag bytes, which follow + // the length + five u32 header fields: + // 4 (length) + 5 * 4 (header) = 24, then OptionFlags1, OptionFlags2, + // TypeFlags at byte offset 26. + const TYPE_FLAGS_OFFSET: usize = 26; + + let mut payload = BytesMut::new(); + let mut login = LoginMessage::new(); + login.readonly(true); + login + .clone() + .encode(&mut payload) + .expect("encode should succeed"); + + assert_eq!( + payload[TYPE_FLAGS_OFFSET] & LoginTypeFlag::ReadOnlyIntent as u8, + LoginTypeFlag::ReadOnlyIntent as u8, + "fReadOnlyIntent bit must be set in the encoded LOGIN7 TypeFlags byte" + ); + + // Round-trips back into the decoded message. + let decoded = LoginMessage::decode(&mut payload).expect("decode should succeed"); + assert!(decoded.type_flags.contains(LoginTypeFlag::ReadOnlyIntent)); + + // And when not requested, the bit stays clear. + let mut payload = BytesMut::new(); + let mut login = LoginMessage::new(); + login.readonly(false); + login.encode(&mut payload).expect("encode should succeed"); + assert_eq!( + payload[TYPE_FLAGS_OFFSET] & LoginTypeFlag::ReadOnlyIntent as u8, + 0, + "fReadOnlyIntent bit must be clear when read-only intent is not requested" + ); + } + #[test] fn login_message_round_trip() { let mut payload = BytesMut::new(); @@ -595,6 +900,120 @@ mod tests { ) } + #[test] + fn encoded_len_matches_actual_output_length() { + // The reserved capacity must equal the bytes actually produced, so no + // reallocation can occur while the password is in the buffer. + let mut login = LoginMessage::new(); + login.db_name("some-database"); + login.user_name("some-user"); + login.password("hunter2"); + login.server_name("some-server"); + + let expected = login.encoded_len(); + let encoded = login + .encode_to_boxed_slice() + .expect("encode should succeed"); + assert_eq!(encoded.len(), expected); + } + + #[test] + fn large_fields_do_not_reallocate_encode_buffer() { + // Fields far larger than the old fixed 512-byte capacity: with the old + // code the Vec would reallocate after the password was written, leaking + // an un-zeroized copy. `encode_to_boxed_slice` now asserts the capacity never + // changed, so this both exercises and enforces the fix. + let mut login = LoginMessage::new(); + login.user_name("u".repeat(200)); + login.password("p".repeat(400)); + login.db_name("d".repeat(400)); + login.server_name("s".repeat(200)); + login.app_name("a".repeat(200)); + + let expected = login.encoded_len(); + let encoded = login + .encode_to_boxed_slice() + .expect("encode should succeed"); + assert_eq!(encoded.len(), expected); + } + + #[test] + fn large_fed_auth_token_does_not_reallocate() { + // Same invariant on the fed-auth path, whose token/nonce are written + // after the password. + let mut login = LoginMessage::new(); + login.password("p".repeat(300)); + login.db_name("d".repeat(300)); + login.aad_token("t".repeat(500), true, Some([7u8; 32])); + + let expected = login.encoded_len(); + let encoded = login + .encode_to_boxed_slice() + .expect("encode should succeed"); + assert_eq!(encoded.len(), expected); + } + + #[test] + fn encode_to_boxed_slice_returns_exact_len_that_round_trips() { + // The buffer is now a `Box<[u8]>` produced via `into_boxed_slice()` from + // a Vec whose len equals its (reserved) capacity, so the boxed slice + // must be exactly `encoded_len()` bytes — no shrink-realloc leak — and it + // must still decode back into an equivalent message. + let mut login = LoginMessage::new(); + login.db_name("some-database"); + login.user_name("some-user"); + login.password("hunter2"); + login.server_name("some-server"); + + let expected_len = login.encoded_len(); + let encoded: Zeroizing> = login + .clone() + .encode_to_boxed_slice() + .expect("encode should succeed"); + assert_eq!( + encoded.len(), + expected_len, + "boxed login buffer must be exactly encoded_len() bytes (no shrink-realloc)" + ); + + let mut buf = BytesMut::from(&encoded[..]); + let decoded = LoginMessage::decode(&mut buf).expect("decode should succeed"); + assert_eq!(login, decoded); + } + + #[test] + fn fed_auth_token_encode_path_produces_correct_output() { + // Exercises the fed-auth token buffer specifically: a non-empty token + // and nonce force the `encode_to_boxed_slice` fed-auth branch to build and copy + // the token temp buffer (whose capacity==len invariant is checked by an + // internal debug_assert, so this test would panic on realloc). Assert + // the encoded bytes decode back to exactly the token/echo/nonce we set. + let token = "a-fake-security-token-value"; + let nonce = [9u8; 32]; + + let mut login = LoginMessage::new(); + login.password("hunter2"); + login.aad_token(token, true, Some(nonce)); + + let expected_len = login.encoded_len(); + let encoded = login + .encode_to_boxed_slice() + .expect("encode should succeed"); + assert_eq!( + encoded.len(), + expected_len, + "fed-auth login buffer must be exactly encoded_len() bytes (no realloc)" + ); + + let mut buf = BytesMut::from(&encoded[..]); + let decoded = LoginMessage::decode(&mut buf).expect("decode should succeed"); + + let ext = decoded.fed_auth_ext.expect("fed_auth_ext must be present"); + assert_eq!(ext.fed_auth_token.expose_secret(), token); + assert!(ext.fed_auth_echo); + assert_eq!(ext.nonce, Some(nonce)); + } + #[test] fn login_message_with_fed_auth_round_trip() { let mut payload = BytesMut::new(); @@ -610,4 +1029,119 @@ mod tests { assert_eq!(login, decoded); } + + #[test] + fn hostname_and_packet_size_setters_apply() { + let mut login = LoginMessage::new(); + login.hostname("my-workstation"); + login.packet_size(8192); + + assert_eq!(login.hostname, "my-workstation"); + assert_eq!(login.packet_size, 8192); + } + + #[cfg(any( + all(unix, any(feature = "integrated-auth-gssapi", feature = "sspi-rs")), + windows, + feature = "winauth" + ))] + #[test] + fn integrated_security_setter_toggles_flag() { + let mut login = LoginMessage::new(); + + login.integrated_security(Some(vec![1, 2, 3, 4])); + assert!(login + .option_flags_2 + .contains(OptionFlag2::IntegratedSecurity)); + assert_eq!( + login.integrated_security.as_deref(), + Some(&[1, 2, 3, 4][..]) + ); + + login.integrated_security(None); + assert!(!login + .option_flags_2 + .contains(OptionFlag2::IntegratedSecurity)); + assert!(login.integrated_security.is_none()); + } + + #[test] + fn encode_round_trips_integrated_security_bytes() { + let mut payload = BytesMut::new(); + let mut login = LoginMessage::new(); + // Set the field directly to exercise the ibSSPI encode branch without + // depending on the platform-gated setter. + login.integrated_security = Some(vec![9, 8, 7, 6, 5]); + login + .clone() + .encode(&mut payload) + .expect("encode should succeed"); + + let decoded = LoginMessage::decode(&mut payload).expect("decode should succeed"); + assert_eq!(decoded.integrated_security, Some(vec![9, 8, 7, 6, 5])); + } + + #[test] + fn fed_auth_without_nonce_round_trips() { + let mut payload = BytesMut::new(); + let mut login = LoginMessage::new(); + login.aad_token("fake-aad-token", true, None); + login + .clone() + .encode(&mut payload) + .expect("encode should succeed"); + + let decoded = LoginMessage::decode(&mut payload).expect("decode should succeed"); + assert_eq!(login, decoded); + assert_eq!( + decoded.fed_auth_ext.expect("fed auth ext present").nonce, + None + ); + } + + #[test] + fn debug_redacts_fed_auth_token() { + let mut login = LoginMessage::new(); + // Distinctive plaintext so the assertion is load-bearing: a generic + // "REDACTED" check is vacuous because the always-present `password: + // SecretString` field prints "REDACTED" regardless of the fed-auth + // token. Asserting this exact plaintext is absent fails iff the token + // leaks. + let token = "fed-auth-token-PLAINTEXT-XYZ"; + login.aad_token(token, true, Some([9u8; 32])); + assert!( + login.fed_auth_ext.is_some(), + "test must exercise a LoginMessage whose fed_auth_ext is Some" + ); + + let dbg = format!("{login:?}"); + assert!( + !dbg.contains(token), + "AAD token leaked in Debug output: {dbg}" + ); + } + + #[test] + fn debug_redacts_login_password() { + let mut login = LoginMessage::new(); + login.user_name("some-user"); + login.password("super-secret-login-pw"); + + let dbg = format!("{login:?}"); + assert!( + !dbg.contains("super-secret-login-pw"), + "password leaked in Debug output: {dbg}" + ); + assert!(dbg.contains("REDACTED"), "password not redacted: {dbg}"); + // Non-secret fields remain visible for diagnostics. + assert!(dbg.contains("some-user"), "username should be shown: {dbg}"); + } + + #[test] + fn password_setter_stores_exposable_secret() { + let mut login = LoginMessage::new(); + login.password("hunter2"); + // The stored `SecretString` exposes exactly the plaintext provided. + assert_eq!(login.password.expose_secret(), "hunter2"); + } } diff --git a/src/tds/codec/packet.rs b/src/tds/codec/packet.rs index 9927ed35d..9c4e9e34a 100644 --- a/src/tds/codec/packet.rs +++ b/src/tds/codec/packet.rs @@ -55,3 +55,65 @@ impl<'a> Extend<&'a u8> for Packet { self.payload.extend(iter) } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::tds::codec::PacketHeader; + + #[test] + fn is_last_reflects_end_of_message_status() { + let mut packet = Packet::new(PacketHeader::batch(1), BytesMut::new()); + assert!(!packet.is_last()); + + packet.header.set_status(PacketStatus::EndOfMessage); + assert!(packet.is_last()); + } + + #[test] + fn into_parts_returns_header_and_payload() { + let payload = BytesMut::from(&b"hello"[..]); + let packet = Packet::new(PacketHeader::batch(3), payload.clone()); + + let (header, parts_payload) = packet.into_parts(); + assert_eq!(header.r#type() as u8, PacketHeader::batch(3).r#type() as u8); + assert_eq!(parts_payload, payload); + } + + #[test] + fn encode_patches_total_length_into_header() { + let payload = BytesMut::from(&b"abcd"[..]); + let packet = Packet::new(PacketHeader::batch(1), payload); + + let mut buf = BytesMut::new(); + packet.encode(&mut buf).unwrap(); + + // 8 header bytes + 4 payload bytes. + assert_eq!(&buf[2..4], &12u16.to_be_bytes()); + assert_eq!(&buf[8..], b"abcd"); + } + + #[test] + fn decode_splits_header_and_remaining_payload() { + let payload = BytesMut::from(&b"xyz"[..]); + let packet = Packet::new(PacketHeader::batch(1), payload); + + let mut buf = BytesMut::new(); + packet.encode(&mut buf).unwrap(); + + let decoded = Packet::decode(&mut buf).unwrap(); + assert_eq!(&decoded.payload[..], b"xyz"); + assert!(buf.is_empty()); + } + + #[test] + fn extend_by_value_and_by_ref_append_to_payload() { + let mut packet = Packet::new(PacketHeader::batch(1), BytesMut::new()); + packet.extend(vec![1u8, 2, 3]); + assert_eq!(&packet.payload[..], &[1, 2, 3]); + + let more = [4u8, 5]; + packet.extend(more.iter()); + assert_eq!(&packet.payload[..], &[1, 2, 3, 4, 5]); + } +} diff --git a/src/tds/codec/pre_login.rs b/src/tds/codec/pre_login.rs index 0374756d6..eb96b784b 100644 --- a/src/tds/codec/pre_login.rs +++ b/src/tds/codec/pre_login.rs @@ -4,13 +4,12 @@ use crate::{tds, Error, Result}; use byteorder::{BigEndian, LittleEndian, ReadBytesExt, WriteBytesExt}; use bytes::{BufMut, BytesMut}; use std::convert::TryFrom; -use std::io::{Cursor, Read}; +use std::io::{Cursor, Read, Write}; use tds::EncryptionLevel; use uuid::Uuid; /// Client application activity id token used for debugging purposes introduced /// in TDS 7.4. -#[allow(unused)] #[derive(Debug, Clone)] #[cfg_attr(test, derive(PartialEq))] pub struct ActivityId { @@ -18,6 +17,17 @@ pub struct ActivityId { sequence: u32, } +impl ActivityId { + /// Creates a new activity id (`TRACEID`, MS-TDS §2.2.6.5) from a client + /// activity [`Uuid`] and a monotonically increasing sequence number. The + /// value is emitted in the PRELOGIN packet so a server administrator can + /// correlate the connection in server-side traces. + #[cfg(test)] + pub fn new(id: Uuid, sequence: u32) -> Self { + Self { id, sequence } + } +} + /// The prelogin packet used to initialize a connection #[derive(Debug, Clone)] #[cfg_attr(test, derive(PartialEq))] @@ -57,13 +67,40 @@ impl PreloginMessage { } } + /// Validates the server's answer to the `INSTOPT` prelogin option + /// (MS-TDS §2.2.6.5). + /// + /// When the client sends an instance name, the server replies with a single + /// `0x00` byte if the instance the connection landed on is the one that was + /// requested. Any other payload means the server considers the instance + /// invalid, in which case this returns a protocol error. + pub fn validate_instance(&self, requested: Option<&str>) -> Result<()> { + // Nothing to validate if the client never asked for a named instance. + if requested.is_none() { + return Ok(()); + } + + match self.instance_name.as_deref() { + // `0x00` terminator only -> decoded as `None`: instance is valid. + None => Ok(()), + Some(other) => Err(Error::Protocol( + format!( + "server rejected the requested instance {:?} (INSTOPT validity byte was non-zero, got {:?})", + requested.unwrap_or_default(), + other, + ) + .into(), + )), + } + } + #[cfg(any( feature = "rustls", feature = "native-tls", feature = "vendored-openssl" ))] - pub fn negotiated_encryption(&self, expected: EncryptionLevel) -> EncryptionLevel { - match (expected, self.encryption) { + pub fn negotiated_encryption(&self, expected: EncryptionLevel) -> Result { + let level = match (expected, self.encryption) { (EncryptionLevel::NotSupported, EncryptionLevel::NotSupported) => { EncryptionLevel::NotSupported } @@ -74,10 +111,25 @@ impl PreloginMessage { (EncryptionLevel::Off, EncryptionLevel::NotSupported) => EncryptionLevel::NotSupported, (EncryptionLevel::On, EncryptionLevel::Off) | (EncryptionLevel::On, EncryptionLevel::NotSupported) => { - panic!("Server does not allow the requested encryption level.") + return Err(Error::Protocol( + "Server does not allow the requested encryption level.".into(), + )) + } + // The client required encryption but the server declined it: this is + // a hard failure, not a silent downgrade to `On`. + (EncryptionLevel::Required, EncryptionLevel::Off) + | (EncryptionLevel::Required, EncryptionLevel::NotSupported) => { + return Err(Error::Protocol( + "Server does not allow the requested encryption level.".into(), + )) } + // In TDS 8.0 "strict" mode encryption is established before the + // prelogin, so there is nothing to negotiate here. + (EncryptionLevel::Strict, _) => EncryptionLevel::Strict, (_, _) => EncryptionLevel::On, - } + }; + + Ok(level) } #[cfg(not(any( @@ -85,8 +137,8 @@ impl PreloginMessage { feature = "native-tls", feature = "vendored-openssl" )))] - pub fn negotiated_encryption(&self, _: EncryptionLevel) -> EncryptionLevel { - EncryptionLevel::NotSupported + pub fn negotiated_encryption(&self, _: EncryptionLevel) -> Result { + Ok(EncryptionLevel::NotSupported) } } @@ -114,7 +166,16 @@ impl Encode for PreloginMessage { // encryption fields.push((PRELOGIN_ENCRYPTION, 0x01)); // encryption - data_cursor.write_u8(self.encryption as u8)?; + data_cursor.write_u8(self.encryption.as_wire_value())?; + + // instance name (INSTOPT): a null-terminated MBCS string naming the + // instance the client wants the server to validate. An empty name is + // encoded as a lone `0x00` terminator. + let instance = self.instance_name.as_deref().unwrap_or_default(); + let instance_bytes = instance.as_bytes(); + fields.push((PRELOGIN_INSTOPT, (instance_bytes.len() + 1) as u16)); + data_cursor.write_all(instance_bytes)?; + data_cursor.write_u8(0x00)?; // null terminator // threadid fields.push((PRELOGIN_THREADID, 0x04)); // thread id @@ -124,6 +185,17 @@ impl Encode for PreloginMessage { fields.push((PRELOGIN_MARS, 0x01)); // MARS data_cursor.write_u8(self.mars as u8)?; + // activity id (TRACEID): a client GUID plus a sequence number, emitted + // only when the client supplies one for server-side trace correlation. + if let Some(activity_id) = self.activity_id.as_ref() { + fields.push((PRELOGIN_TRACEID, 0x14)); // 16-byte GUID + 4-byte sequence + + let mut data = *activity_id.id.as_bytes(); + reorder_bytes(&mut data); + data_cursor.write_all(&data)?; + data_cursor.write_u32::(activity_id.sequence)?; + } + // fed auth if self.fed_auth_required { fields.push((PRELOGIN_FEDAUTHREQUIRED, 0x01)); @@ -163,7 +235,7 @@ impl Decode for PreloginMessage { let token = cursor.read_u8()?; // read until terminator - if token == 0xff { + if token == PRELOGIN_TERMINATOR { break; } @@ -175,7 +247,7 @@ impl Decode for PreloginMessage { // verify whether the server acts in accordance to what we requested // and if we can handle on what we seemingly agreed to - // TODO: support parsing more + // Unrecognized pre-login option tokens are rejected as a protocol error. match token { // version PRELOGIN_VERSION => { @@ -209,7 +281,9 @@ impl Decode for PreloginMessage { } else if length == 4 { cursor.read_u32::()? } else { - panic!("should never happen") + return Err(Error::Protocol( + format!("prelogin: invalid threadid length: {}", length).into(), + )); } } // mars @@ -244,7 +318,11 @@ impl Decode for PreloginMessage { ret.nonce = Some(data); } - _ => panic!("unsupported prelogin token: {}", token), + _ => { + return Err(Error::Protocol( + format!("unsupported prelogin token: {}", token).into(), + )) + } } cursor.set_position(old_pos); @@ -258,6 +336,261 @@ impl Decode for PreloginMessage { mod tests { use super::*; + /// Parses the PRELOGIN option-offset table and returns the option tokens in + /// the order they were emitted, stopping at the terminator. + fn option_tokens(bytes: &[u8]) -> Vec { + let mut tokens = Vec::new(); + let mut pos = 0; + + while pos < bytes.len() { + let token = bytes[pos]; + + if token == PRELOGIN_TERMINATOR { + break; + } + + tokens.push(token); + // token (1) + offset (2) + length (2) + pos += 5; + } + + tokens + } + + /// Reads the raw payload bytes for a given option token from an encoded + /// PRELOGIN packet using its offset-table entry. + fn option_payload(bytes: &[u8], token: u8) -> Option> { + let mut pos = 0; + + while pos < bytes.len() && bytes[pos] != PRELOGIN_TERMINATOR { + if bytes[pos] == token { + let offset = u16::from_be_bytes([bytes[pos + 1], bytes[pos + 2]]) as usize; + let length = u16::from_be_bytes([bytes[pos + 3], bytes[pos + 4]]) as usize; + + return Some(bytes[offset..offset + length].to_vec()); + } + + pos += 5; + } + + None + } + + #[test] + fn prelogin_always_emits_instopt() { + let mut payload = BytesMut::new(); + PreloginMessage::new() + .encode(&mut payload) + .expect("encode should succeed"); + + assert!(option_tokens(&payload).contains(&PRELOGIN_INSTOPT)); + // An empty instance name is a single null terminator byte. + assert_eq!(option_payload(&payload, PRELOGIN_INSTOPT), Some(vec![0x00])); + } + + #[test] + fn prelogin_emits_named_instance() { + let mut payload = BytesMut::new(); + let mut prelogin = PreloginMessage::new(); + prelogin.instance_name = Some("MSSQLServer".to_string()); + prelogin + .clone() + .encode(&mut payload) + .expect("encode should succeed"); + + let expected: Vec = b"MSSQLServer\0".to_vec(); + assert_eq!(option_payload(&payload, PRELOGIN_INSTOPT), Some(expected)); + + let decoded = PreloginMessage::decode(&mut payload).expect("decode should succeed"); + assert_eq!(decoded.instance_name.as_deref(), Some("MSSQLServer")); + } + + #[test] + fn prelogin_emits_traceid_only_when_present() { + let mut without = BytesMut::new(); + PreloginMessage::new() + .encode(&mut without) + .expect("encode should succeed"); + assert!(!option_tokens(&without).contains(&PRELOGIN_TRACEID)); + + let mut with = BytesMut::new(); + let mut prelogin = PreloginMessage::new(); + prelogin.activity_id = Some(ActivityId::new( + Uuid::parse_str("6f9619ff-8b86-d011-b42d-00c04fc964ff").unwrap(), + 42, + )); + prelogin + .clone() + .encode(&mut with) + .expect("encode should succeed"); + + assert!(option_tokens(&with).contains(&PRELOGIN_TRACEID)); + // 16-byte GUID + 4-byte sequence. + assert_eq!( + option_payload(&with, PRELOGIN_TRACEID).map(|p| p.len()), + Some(20) + ); + + let decoded = PreloginMessage::decode(&mut with).expect("decode should succeed"); + assert_eq!(decoded.activity_id, prelogin.activity_id); + } + + #[test] + fn decode_accepts_zero_length_threadid() { + // A THREADID option with length 0 must decode to thread_id 0, not error. + // Table entry (token, offset=6, length=0) then the terminator. + let mut buf = BytesMut::from( + &[ + PRELOGIN_THREADID, + 0x00, + 0x06, + 0x00, + 0x00, + PRELOGIN_TERMINATOR, + ][..], + ); + let decoded = PreloginMessage::decode(&mut buf).expect("zero-length threadid must decode"); + assert_eq!(decoded.thread_id, 0); + } + + #[test] + fn decode_reads_nonce_option() { + // Table entry for NONCEOPT (offset 6, length 32) + terminator + 32 bytes. + let mut bytes = vec![ + PRELOGIN_NONCEOPT, + 0x00, + 0x06, + 0x00, + 0x20, + PRELOGIN_TERMINATOR, + ]; + bytes.extend_from_slice(&[0xAB; 32]); + let mut buf = BytesMut::from(&bytes[..]); + + let decoded = PreloginMessage::decode(&mut buf).expect("nonce option must decode"); + assert_eq!(decoded.nonce, Some([0xAB; 32])); + } + + #[test] + fn option_payload_returns_none_for_absent_token() { + let mut payload = BytesMut::new(); + PreloginMessage::new() + .encode(&mut payload) + .expect("encode should succeed"); + + // A fresh message emits no TRACEID option. + assert_eq!(option_payload(&payload, PRELOGIN_TRACEID), None); + } + + #[test] + fn decode_rejects_invalid_encryption_value() { + // ENCRYPTION option (offset 6, length 1) + terminator + an out-of-range + // encryption byte. + let mut buf = BytesMut::from( + &[ + PRELOGIN_ENCRYPTION, + 0x00, + 0x06, + 0x00, + 0x01, + PRELOGIN_TERMINATOR, + 0x63, // 99: not a valid EncryptionLevel + ][..], + ); + + match PreloginMessage::decode(&mut buf) { + Err(Error::Protocol(_)) => {} + other => panic!("expected protocol error, got {other:?}"), + } + } + + #[test] + fn decode_rejects_invalid_threadid_length() { + // THREADID option with an unsupported length (2) must error. + let mut buf = BytesMut::from( + &[ + PRELOGIN_THREADID, + 0x00, + 0x06, + 0x00, + 0x02, + PRELOGIN_TERMINATOR, + 0x00, + 0x00, + ][..], + ); + + match PreloginMessage::decode(&mut buf) { + Err(Error::Protocol(_)) => {} + other => panic!("expected protocol error, got {other:?}"), + } + } + + #[test] + fn decode_rejects_unsupported_token() { + // An unknown option token must produce a protocol error. + let mut buf = BytesMut::from( + &[ + 0x50, // unsupported token + 0x00, + 0x06, + 0x00, + 0x01, + PRELOGIN_TERMINATOR, + 0x00, + ][..], + ); + + match PreloginMessage::decode(&mut buf) { + Err(Error::Protocol(_)) => {} + other => panic!("expected protocol error, got {other:?}"), + } + } + + #[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + ))] + #[test] + fn negotiated_encryption_off_and_strict() { + let mut prelogin = PreloginMessage::new(); + + // Both sides Off -> Off. + prelogin.encryption = EncryptionLevel::Off; + assert_eq!( + prelogin + .negotiated_encryption(EncryptionLevel::Off) + .unwrap(), + EncryptionLevel::Off + ); + + // Strict is negotiated out-of-band; it stays Strict regardless of the + // server's advertised level. + assert_eq!( + prelogin + .negotiated_encryption(EncryptionLevel::Strict) + .unwrap(), + EncryptionLevel::Strict + ); + } + + #[test] + fn validate_instance_accepts_valid_response() { + // Server valid response = lone 0x00 -> decoded as `None`. + let msg = PreloginMessage::new(); + assert!(msg.validate_instance(Some("MSSQLServer")).is_ok()); + // No requested instance -> nothing to validate. + assert!(msg.validate_instance(None).is_ok()); + } + + #[test] + fn validate_instance_rejects_invalid_response() { + let mut msg = PreloginMessage::new(); + msg.instance_name = Some("otherinstance".to_string()); + assert!(msg.validate_instance(Some("MSSQLServer")).is_err()); + } + #[test] #[cfg(any( feature = "rustls", @@ -269,7 +602,9 @@ mod tests { response.encryption = EncryptionLevel::NotSupported; assert_eq!( - response.negotiated_encryption(EncryptionLevel::Off), + response + .negotiated_encryption(EncryptionLevel::Off) + .unwrap(), EncryptionLevel::NotSupported ); } @@ -290,7 +625,9 @@ mod tests { response.encryption = server_encryption; assert_eq!( - response.negotiated_encryption(EncryptionLevel::Off), + response + .negotiated_encryption(EncryptionLevel::Off) + .unwrap(), expected ); } @@ -324,4 +661,53 @@ mod tests { assert_eq!(prelogin, decoded); } + + // #425: a server declining the requested encryption level must yield a + // catchable protocol error instead of panicking. + #[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + ))] + #[test] + fn negotiated_encryption_rejects_declined_level() { + let mut prelogin = PreloginMessage::new(); + // Server responds with an encryption level the client did not offer / + // that is weaker than the required `On`. + prelogin.encryption = EncryptionLevel::Off; + + let result = prelogin.negotiated_encryption(EncryptionLevel::On); + + match result { + Err(Error::Protocol(_)) => {} + other => panic!("expected Err(Error::Protocol), got {other:?}"), + } + + // A matching, valid negotiation still succeeds. + prelogin.encryption = EncryptionLevel::On; + assert_eq!( + prelogin.negotiated_encryption(EncryptionLevel::On).unwrap(), + EncryptionLevel::On + ); + } + + #[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" + ))] + #[test] + fn negotiated_encryption_required_rejects_declined_level() { + // The client required encryption; a server that responds `Off` or + // `NotSupported` must be a hard protocol error, not a silent downgrade. + for server in [EncryptionLevel::Off, EncryptionLevel::NotSupported] { + let mut prelogin = PreloginMessage::new(); + prelogin.encryption = server; + + match prelogin.negotiated_encryption(EncryptionLevel::Required) { + Err(Error::Protocol(_)) => {} + other => panic!("expected Err(Error::Protocol) for {server:?}, got {other:?}"), + } + } + } } diff --git a/src/tds/codec/rpc_request.rs b/src/tds/codec/rpc_request.rs index 368cbd3c8..414e1c02a 100644 --- a/src/tds/codec/rpc_request.rs +++ b/src/tds/codec/rpc_request.rs @@ -1,4 +1,5 @@ -use super::{AllHeaderTy, Encode, ALL_HEADERS_LEN_TX}; +use super::TypeInfoTvp; +use super::{encode_all_headers_tx, encode_b_varchar, Encode}; use crate::{tds::codec::ColumnData, BytesMutWithTypeInfo, Result}; use bytes::{BufMut, BytesMut}; use enumflags2::{bitflags, BitFlags}; @@ -46,11 +47,22 @@ impl<'a> TokenRpcRequest<'a> { } } +/// The value carried by an [`RpcParam`]. A scalar column value, or a +/// table-valued parameter (TVP). +#[derive(Debug)] +pub enum RpcValue<'a> { + /// An ordinary scalar parameter value. + Scalar(ColumnData<'a>), + /// A table-valued parameter. As per the TDS grammar, `TYPE_INFO_TVP` + /// carries both the type metadata and the data rows. + Table(TypeInfoTvp<'a>), +} + #[derive(Debug)] pub struct RpcParam<'a> { pub name: Cow<'a, str>, pub flags: BitFlags, - pub value: ColumnData<'a>, + pub value: RpcValue<'a>, } /// 2.2.6.6 RPC Request @@ -92,21 +104,40 @@ impl<'a> From for RpcProcIdValue<'a> { impl<'a> Encode for TokenRpcRequest<'a> { fn encode(self, dst: &mut BytesMut) -> Result<()> { - dst.put_u32_le(ALL_HEADERS_LEN_TX as u32); - dst.put_u32_le(ALL_HEADERS_LEN_TX as u32 - 4); - dst.put_u16_le(AllHeaderTy::TransactionDescriptor as u16); - dst.put_slice(&self.transaction_desc); - dst.put_u32_le(1); + encode_all_headers_tx(dst, self.transaction_desc); match self.proc_id { RpcProcIdValue::Id(ref id) => { let val = (0xffff_u32) | ((*id as u16) as u32) << 16; dst.put_u32_le(val); } - RpcProcIdValue::Name(ref _name) => { - //let (left_bytes, _) = try!(write_varchar::(&mut cursor, name, 0)); - //assert_eq!(left_bytes, 0); - todo!() + RpcProcIdValue::Name(ref name) => { + // ProcName is a US_VARCHAR: a u16 little-endian character count + // followed by that many UTF-16 code units. The count is a u16, so + // a name longer than 65535 code units cannot be represented; + // reject it rather than overflow the counter and desync the wire. + let units = name.encode_utf16().count(); + if units > u16::MAX as usize { + return Err(crate::Error::Protocol( + format!( + "RPC procedure name is too long ({units} UTF-16 code units, max 65535)" + ) + .into(), + )); + } + + let len_pos = dst.len(); + dst.put_u16_le(0u16); + let mut length = 0_u16; + + for chr in name.encode_utf16() { + dst.put_u16_le(chr); + length += 1; + } + + let dst: &mut [u8] = dst.borrow_mut(); + let mut dst = &mut dst[len_pos..]; + dst.put_u16_le(length); } } @@ -122,24 +153,154 @@ impl<'a> Encode for TokenRpcRequest<'a> { impl<'a> Encode for RpcParam<'a> { fn encode(self, dst: &mut BytesMut) -> Result<()> { - let len_pos = dst.len(); - let mut length = 0u8; + // ParamName is a B_VARCHAR (u8 code-unit count); reject an over-long + // name rather than wrap the counter and desync the wire. + encode_b_varchar(dst, &self.name)?; - dst.put_u8(length); + dst.put_u8(self.flags.bits()); - for codepoint in self.name.encode_utf16() { - length += 1; - dst.put_u16_le(codepoint); + match self.value { + RpcValue::Scalar(value) => { + let mut dst_ti = BytesMutWithTypeInfo::new(dst); + value.encode(&mut dst_ti)?; + } + RpcValue::Table(value) => value.encode(dst)?, } - dst.put_u8(self.flags.bits()); + Ok(()) + } +} - let mut dst_fi = BytesMutWithTypeInfo::new(dst); - self.value.encode(&mut dst_fi)?; +#[cfg(test)] +mod tests { + use super::super::ALL_HEADERS_LEN_TX; + use super::*; + use crate::tds::codec::ColumnData; - let dst: &mut [u8] = dst.borrow_mut(); - dst[len_pos] = length; + fn scalar(value: ColumnData<'static>) -> RpcValue<'static> { + RpcValue::Scalar(value) + } - Ok(()) + #[test] + fn encodes_named_proc_header() { + let req = TokenRpcRequest::new( + "dbo.usp_MyProc", + vec![RpcParam { + name: Cow::Borrowed("@id"), + flags: BitFlags::empty(), + value: scalar(ColumnData::I32(Some(1))), + }], + [0u8; 8], + ); + + let mut buf = BytesMut::new(); + req.encode(&mut buf).unwrap(); + + // Skip the ALL_HEADERS block, positioned right at the ProcName. + let name_pos = ALL_HEADERS_LEN_TX; + let len = u16::from_le_bytes([buf[name_pos], buf[name_pos + 1]]); + assert_eq!(len as usize, "dbo.usp_MyProc".encode_utf16().count()); + + // Verify the UTF-16 payload matches the proc name. + let mut chars = Vec::new(); + let mut off = name_pos + 2; + for _ in 0..len { + chars.push(u16::from_le_bytes([buf[off], buf[off + 1]])); + off += 2; + } + assert_eq!(String::from_utf16(&chars).unwrap(), "dbo.usp_MyProc"); + + // Option flags (u16) follow the name. + let flags = u16::from_le_bytes([buf[off], buf[off + 1]]); + assert_eq!(flags, 0); + } + + #[test] + fn named_and_by_id_differ_only_in_proc_slot() { + let by_id = { + let req = TokenRpcRequest::new(RpcProcId::ExecuteSQL, vec![], [0u8; 8]); + let mut buf = BytesMut::new(); + req.encode(&mut buf).unwrap(); + buf + }; + + // By-id encodes 0xFFFF followed by the proc id in the high word. + let val = u32::from_le_bytes([ + by_id[ALL_HEADERS_LEN_TX], + by_id[ALL_HEADERS_LEN_TX + 1], + by_id[ALL_HEADERS_LEN_TX + 2], + by_id[ALL_HEADERS_LEN_TX + 3], + ]); + assert_eq!(val & 0xffff, 0xffff); + assert_eq!((val >> 16) as u16, RpcProcId::ExecuteSQL as u16); + } + + #[test] + fn encodes_param_name_and_by_ref_flag() { + let param = RpcParam { + name: Cow::Borrowed("@out"), + flags: BitFlags::from_flag(RpcStatus::ByRefValue), + value: scalar(ColumnData::I32(Some(7))), + }; + + let mut buf = BytesMut::new(); + param.encode(&mut buf).unwrap(); + + // First byte is the param-name length (in UTF-16 code units). + assert_eq!(buf[0] as usize, "@out".encode_utf16().count()); + + let mut chars = Vec::new(); + let mut off = 1usize; + for _ in 0..buf[0] { + chars.push(u16::from_le_bytes([buf[off], buf[off + 1]])); + off += 2; + } + assert_eq!(String::from_utf16(&chars).unwrap(), "@out"); + + // Status flags byte carries the ByRefValue bit. + assert_eq!(buf[off], RpcStatus::ByRefValue as u8); + } + + // ProcName length is a u16; a name longer than 65535 UTF-16 code units must + // error rather than overflow the counter (and desync the wire). + #[test] + fn rejects_over_long_proc_name() { + let long = "a".repeat(u16::MAX as usize + 1); + let req = TokenRpcRequest::new(long, vec![], [0u8; 8]); + + let mut buf = BytesMut::new(); + let err = req.encode(&mut buf).unwrap_err(); + assert!(matches!(err, crate::Error::Protocol(_)), "got {err:?}"); + } + + // A param name is a B_VARCHAR (u8 length); a name longer than 255 UTF-16 code + // units must error rather than wrap the counter and desync the wire. + #[test] + fn rejects_over_long_param_name() { + let param = RpcParam { + name: Cow::Owned("@".to_string() + &"a".repeat(255)), + flags: BitFlags::empty(), + value: scalar(ColumnData::I32(Some(1))), + }; + + let mut buf = BytesMut::new(); + let err = param.encode(&mut buf).unwrap_err(); + assert!(matches!(err, crate::Error::Protocol(_)), "got {err:?}"); + } + + // A param name of exactly 255 units is the boundary and must still encode, + // writing the correct u8 length prefix. + #[test] + fn encodes_max_length_param_name() { + let name = "a".repeat(255); + let param = RpcParam { + name: Cow::Owned(name.clone()), + flags: BitFlags::empty(), + value: scalar(ColumnData::I32(Some(1))), + }; + + let mut buf = BytesMut::new(); + param.encode(&mut buf).expect("255-unit name must encode"); + assert_eq!(buf[0], 255); } } diff --git a/src/tds/codec/token.rs b/src/tds/codec/token.rs index 706d03710..3c8ad9805 100644 --- a/src/tds/codec/token.rs +++ b/src/tds/codec/token.rs @@ -1,25 +1,55 @@ +mod token_alt_meta_data; +mod token_alt_row; +mod token_col_info; mod token_col_metadata; mod token_done; mod token_env_change; mod token_error; mod token_feature_ext_ack; +mod token_fed_auth_info; mod token_info; mod token_login_ack; mod token_order; mod token_return_value; mod token_row; +mod token_session_state; mod token_sspi; +mod token_tab_name; mod token_type; +pub use token_alt_meta_data::*; +pub use token_alt_row::*; +pub use token_col_info::*; pub use token_col_metadata::*; pub use token_done::*; pub use token_env_change::*; pub use token_error::*; pub use token_feature_ext_ack::*; +pub use token_fed_auth_info::*; pub use token_info::*; pub use token_login_ack::*; pub use token_order::*; pub use token_return_value::*; pub use token_row::*; +pub use token_session_state::*; pub use token_sspi::*; +pub use token_tab_name::*; pub use token_type::*; + +/// Upper bound on the length a variable-length token declares for its body +/// before we allocate for it. The length is server-controlled; without a cap a +/// single 4-byte field could force a multi-gigabyte allocation. Chosen well +/// above any realistic FEDAUTHINFO / SESSIONSTATE payload. +pub(crate) const MAX_TOKEN_BODY: usize = 16 * 1024 * 1024; + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn max_token_body_is_sixteen_mebibytes() { + // 16 MiB. Guards the `16 * 1024 * 1024` computation against arithmetic + // mutation (e.g. `+`/`/` would yield a wildly different cap). + assert_eq!(MAX_TOKEN_BODY, 16_777_216); + } +} diff --git a/src/tds/codec/token/token_alt_meta_data.rs b/src/tds/codec/token/token_alt_meta_data.rs new file mode 100644 index 000000000..5537b81d8 --- /dev/null +++ b/src/tds/codec/token/token_alt_meta_data.rs @@ -0,0 +1,145 @@ +use std::borrow::Cow; + +use crate::{tds::codec::BaseMetaDataColumn, SqlReadBytes}; + +/// A column produced by a COMPUTE clause, as described by an +/// [`TokenAltMetaData`] (`ALTMETADATA`, token `0x88`) stream. +/// +/// In addition to the regular column metadata, each computed column carries the +/// aggregate operator (`op`) that produced it (for example `SUM`, `AVG`, +/// `COUNT`) and the operand column number (`operand`) the operator was applied +/// to. See [MS-TDS] section 2.2.7.1. +/// +/// [MS-TDS]: https://learn.microsoft.com/en-us/openspecs/windows_protocols/ms-tds/ +#[derive(Debug, Clone)] +pub struct AltMetaDataColumn<'a> { + /// The aggregate operator that produced this column (`Op` in [MS-TDS]). + pub op: u8, + /// The column number in the originating result set that the aggregate + /// operator was applied to (`Operand` in [MS-TDS]). + pub operand: u16, + /// The regular column metadata (flags, type information). + pub base: BaseMetaDataColumn, + /// The name of the computed column. + pub col_name: Cow<'a, str>, +} + +/// The token describing the layout of a COMPUTE (BY) result set +/// (`ALTMETADATA`, token `0x88`). +/// +/// A single query can contain more than one COMPUTE clause; each is uniquely +/// identified by [`id`](Self::id), which the matching [`TokenAltRow`] rows refer +/// back to. See [MS-TDS] section 2.2.7.1. +/// +/// [`TokenAltRow`]: crate::tds::codec::TokenAltRow +/// [MS-TDS]: https://learn.microsoft.com/en-us/openspecs/windows_protocols/ms-tds/ +#[derive(Debug, Clone)] +pub struct TokenAltMetaData<'a> { + /// Identifies the COMPUTE clause this metadata describes. The associated + /// [`TokenAltRow`](crate::tds::codec::TokenAltRow) rows carry the same id. + pub id: u16, + /// The column numbers (from the originating result set) listed in the + /// COMPUTE `BY` clause, in order. + pub by_columns: Vec, + /// The computed columns, one per aggregate operator in the COMPUTE clause. + pub columns: Vec>, +} + +impl TokenAltMetaData<'static> { + pub(crate) async fn decode(src: &mut R) -> crate::Result + where + R: SqlReadBytes + Unpin, + { + // Number of computed columns, e.g. `COMPUTE SUM(x), AVG(x)` -> 2. + let column_count = src.read_u16_le().await?; + + // Identifies the COMPUTE clause; referenced by the ALTROW token. + let id = src.read_u16_le().await?; + + // Number of grouping columns in the `BY` list. + let by_cols = src.read_u8().await?; + + let mut by_columns = Vec::with_capacity(by_cols as usize); + for _ in 0..by_cols { + by_columns.push(src.read_u16_le().await?); + } + + // `column_count` is an untrusted u16 (up to 65535); cap the up-front + // reservation so a hostile ALTMETADATA token can't force a large + // transient allocation before the column data has arrived. The Vec + // still grows as real columns are decoded. + let mut columns = Vec::with_capacity( + (column_count as usize).min(crate::tds::codec::column_data::MAX_PREALLOC), + ); + for _ in 0..column_count { + let op = src.read_u8().await?; + let operand = src.read_u16_le().await?; + + let base = BaseMetaDataColumn::decode(src).await?; + let col_name = Cow::from(src.read_b_varchar().await?); + + columns.push(AltMetaDataColumn { + op, + operand, + base, + col_name, + }); + } + + Ok(TokenAltMetaData { + id, + by_columns, + columns, + }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::{sql_read_bytes::test_utils::IntoSqlReadBytes, tds::codec::TypeInfo, FixedLenType}; + use bytes::{BufMut, BytesMut}; + + #[tokio::test] + async fn decode_alt_meta_data_single_sum_column() { + // `SELECT ... COMPUTE SUM(x) BY y` style metadata for one Int4 column. + let mut buf = BytesMut::new(); + + buf.put_u16_le(1); // column count (one aggregate) + buf.put_u16_le(7); // compute id + buf.put_u8(1); // by_cols + buf.put_u16_le(2); // BY column number + + // ComputeData for the single column: + buf.put_u8(0x4f); // Op = SUM + buf.put_u16_le(1); // Operand column number + + // BaseMetaDataColumn: user type (u32), flags (u16), TYPE_INFO + buf.put_u32_le(0); // user type + buf.put_u16_le(0x0001); // flags (Nullable) + buf.put_u8(FixedLenType::Int4 as u8); // TYPE_INFO: INT4TYPE + + // ColName as B_VARCHAR (length in chars, then UTF-16LE) + let name: Vec = "sum".encode_utf16().collect(); + buf.put_u8(name.len() as u8); + for c in name { + buf.put_u16_le(c); + } + + let mut reader = buf.into_sql_read_bytes(); + let meta = TokenAltMetaData::decode(&mut reader).await.unwrap(); + + assert_eq!(7, meta.id); + assert_eq!(vec![2], meta.by_columns); + assert_eq!(1, meta.columns.len()); + + let col = &meta.columns[0]; + assert_eq!(0x4f, col.op); + assert_eq!(1, col.operand); + assert_eq!("sum", col.col_name); + assert!(matches!( + col.base.ty, + TypeInfo::FixedLen(FixedLenType::Int4) + )); + } +} diff --git a/src/tds/codec/token/token_alt_row.rs b/src/tds/codec/token/token_alt_row.rs new file mode 100644 index 000000000..4d5714aec --- /dev/null +++ b/src/tds/codec/token/token_alt_row.rs @@ -0,0 +1,140 @@ +use crate::{ + tds::codec::{ColumnData, TokenAltMetaData}, + SqlReadBytes, +}; + +/// A row of computed data produced by a COMPUTE (BY) clause (`ALTROW`, token +/// `0xD3`). +/// +/// The row refers back, through [`id`](Self::id), to the +/// [`TokenAltMetaData`](crate::tds::codec::TokenAltMetaData) that describes the +/// type of each value. See [MS-TDS] section 2.2.7.2. +/// +/// [MS-TDS]: https://learn.microsoft.com/en-us/openspecs/windows_protocols/ms-tds/ +#[derive(Debug, Clone)] +pub struct TokenAltRow<'a> { + /// Identifies the COMPUTE clause (and thus the `ALTMETADATA`) this row + /// belongs to. + pub id: u16, + data: Vec>, +} + +impl<'a> TokenAltRow<'a> { + /// The id of the COMPUTE clause this row belongs to. + pub fn id(&self) -> u16 { + self.id + } + + /// The number of computed columns in the row. + pub fn len(&self) -> usize { + self.data.len() + } + + /// True if the row has no columns. + pub fn is_empty(&self) -> bool { + self.data.is_empty() + } + + /// Returns an iterator over the computed column values. + pub fn iter(&self) -> std::slice::Iter<'_, ColumnData<'a>> { + self.data.iter() + } + + /// Gets the computed value at the given index, `None` if out of bounds. + pub fn get(&self, index: usize) -> Option<&ColumnData<'a>> { + self.data.get(index) + } +} + +impl TokenAltRow<'static> { + /// Decodes the column values of an `ALTROW` for the COMPUTE clause `id`, + /// using the previously received [`TokenAltMetaData`] that describes it. + /// + /// The `id` is read from the wire separately (by the token stream) so that + /// the correct metadata can be looked up before the values are parsed. + pub(crate) async fn decode( + src: &mut R, + id: u16, + meta: &TokenAltMetaData<'static>, + ) -> crate::Result + where + R: SqlReadBytes + Unpin, + { + let mut data = Vec::with_capacity(meta.columns.len()); + + for column in meta.columns.iter() { + data.push(ColumnData::decode(src, &column.base.ty).await?); + } + + Ok(TokenAltRow { id, data }) + } +} + +#[cfg(test)] +mod tests { + use std::borrow::Cow; + + use super::*; + use crate::{ + sql_read_bytes::test_utils::IntoSqlReadBytes, AltMetaDataColumn, BaseMetaDataColumn, + ColumnFlag, FixedLenType, TypeInfo, + }; + use bytes::{BufMut, BytesMut}; + + fn int_alt_meta() -> TokenAltMetaData<'static> { + TokenAltMetaData { + id: 1, + by_columns: vec![], + columns: vec![AltMetaDataColumn { + op: 0x4f, // SUM + operand: 1, + base: BaseMetaDataColumn { + flags: ColumnFlag::Nullable.into(), + ty: TypeInfo::FixedLen(FixedLenType::Int4), + table_name: None, + }, + col_name: Cow::from("sum"), + }], + } + } + + #[tokio::test] + async fn decode_alt_row_reads_values() { + let meta = int_alt_meta(); + + let mut buf = BytesMut::new(); + buf.put_i32_le(42); // the single Int4 computed value + + let mut reader = buf.into_sql_read_bytes(); + let row = TokenAltRow::decode(&mut reader, meta.id, &meta) + .await + .unwrap(); + + assert_eq!(1, row.id()); + assert_eq!(1, row.len()); + assert!(matches!(row.get(0), Some(ColumnData::I32(Some(42))))); + } + + #[test] + fn accessors_reflect_id_and_columns() { + // Empty row: id must be the stored value (not a hardcoded 1), len 0, + // is_empty true. + let empty = TokenAltRow { + id: 5, + data: vec![], + }; + assert_eq!(empty.id(), 5); + assert_eq!(empty.len(), 0); + assert!(empty.is_empty()); + + // Non-empty row with a distinct id and two columns: len 2, is_empty + // false. + let filled = TokenAltRow { + id: 9, + data: vec![ColumnData::I32(Some(10)), ColumnData::I32(Some(20))], + }; + assert_eq!(filled.id(), 9); + assert_eq!(filled.len(), 2); + assert!(!filled.is_empty()); + } +} diff --git a/src/tds/codec/token/token_col_info.rs b/src/tds/codec/token/token_col_info.rs new file mode 100644 index 000000000..a8526f3f6 --- /dev/null +++ b/src/tds/codec/token/token_col_info.rs @@ -0,0 +1,331 @@ +use crate::SqlReadBytes; + +/// The column is the result of an expression rather than a direct reference to +/// a base-table column (`fExpression`). +const STATUS_EXPRESSION: u8 = 0x04; +/// The column is part of a key for the associated table (`fKey`). +const STATUS_KEY: u8 = 0x08; +/// The column was not requested, but was added because it is part of a key for +/// the associated table (`fHidden`). +const STATUS_HIDDEN: u8 = 0x10; +/// The column name is different from the name of the base-table column it was +/// derived from (`fDifferentName`). When set, the entry carries a `ColName`. +const STATUS_DIFFERENT_NAME: u8 = 0x20; + +/// A single column description within a [`TokenColInfo`] token. +#[allow(dead_code)] // informational: exposed for debugging/browse-mode consumers +#[derive(Debug, Clone)] +pub struct ColInfo { + /// The column number in the result set (1-based). + pub(crate) col_num: u8, + /// The number of the base table the column was derived from, as an index + /// into the table names carried by a preceding `TABNAME` token. Zero when + /// the column is not derived from a table column. + pub(crate) table_num: u8, + /// The raw status bitmap for this column. + pub(crate) status: u8, + /// The base-table column name, present only when + /// [`ColInfo::has_different_name`] is `true`. + pub(crate) col_name: Option, +} + +#[allow(dead_code)] // informational accessors for browse-mode consumers +impl ColInfo { + /// Whether the column is the result of an expression (`fExpression`). + pub(crate) fn is_expression(&self) -> bool { + self.status & STATUS_EXPRESSION != 0 + } + + /// Whether the column is part of a key (`fKey`). + pub(crate) fn is_key(&self) -> bool { + self.status & STATUS_KEY != 0 + } + + /// Whether the column was added implicitly because it is part of a key + /// (`fHidden`). + pub(crate) fn is_hidden(&self) -> bool { + self.status & STATUS_HIDDEN != 0 + } + + /// Whether the column carries a differing base-table name (`fDifferentName`). + pub(crate) fn has_different_name(&self) -> bool { + self.status & STATUS_DIFFERENT_NAME != 0 + } +} + +/// The `COLINFO` token (`0xA5`), sent by the server in browse mode to describe +/// the origin of each column in the result set. +/// +/// See MS-TDS §2.2.7.4. This token is informational; it is decoded so the +/// token stream does not error out when browse-mode metadata is returned. +#[allow(dead_code)] // informational: consumed by the token stream, exposed for debugging +#[derive(Debug, Clone)] +pub struct TokenColInfo { + /// One entry per column in the current result set. + pub(crate) columns: Vec, +} + +impl TokenColInfo { + pub(crate) async fn decode(src: &mut R) -> crate::Result + where + R: SqlReadBytes + Unpin, + { + // Total length in bytes of the column-info entries that follow. + let length = src.read_u16_le().await? as usize; + + let mut consumed = 0usize; + let mut columns = Vec::new(); + + while consumed < length { + // Each entry has a fixed 3-byte header (ColNum, TableNum, Status). + // A truncated header would read past the declared `length`. + if length - consumed < 3 { + return Err(crate::error::Error::Protocol( + "COLINFO entry header extends past the token length".into(), + )); + } + + let col_num = src.read_u8().await?; + let table_num = src.read_u8().await?; + let status = src.read_u8().await?; + consumed += 3; + + let col_name = if status & STATUS_DIFFERENT_NAME != 0 { + // ColName is a B_VARCHAR: a byte length (in UCS-2 characters) + // followed by that many little-endian UTF-16 code units. Both the + // length byte and the code units must stay within `length`. + if length - consumed < 1 { + return Err(crate::error::Error::Protocol( + "COLINFO ColName length byte extends past the token length".into(), + )); + } + let char_len = src.read_u8().await? as usize; + consumed += 1; + + let name_bytes = char_len * 2; + if length - consumed < name_bytes { + return Err(crate::error::Error::Protocol( + "COLINFO ColName extends past the token length".into(), + )); + } + + let mut units = Vec::with_capacity(char_len); + for _ in 0..char_len { + units.push(src.read_u16_le().await?); + } + consumed += name_bytes; + + // Strict decode with a descriptive protocol error, matching the + // sibling string decoders (TABNAME, the varchar readers) rather + // than silently substituting replacement characters. + Some(String::from_utf16(&units).map_err(|_| { + crate::error::Error::Protocol("COLINFO ColName is not valid UTF-16".into()) + })?) + } else { + None + }; + + columns.push(ColInfo { + col_num, + table_num, + status, + col_name, + }); + } + + Ok(TokenColInfo { columns }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use bytes::{BufMut, BytesMut}; + + #[tokio::test] + async fn decode_col_info_with_and_without_names() { + let mut buf = BytesMut::new(); + + // Placeholder for the u16 length prefix; filled in after the body. + let mut body = BytesMut::new(); + + // Column 1: derived directly from base table 1, no different name. + body.put_u8(1); // ColNum + body.put_u8(1); // TableNum + body.put_u8(STATUS_KEY); // Status: part of a key + + // Column 2: expression + different name "Id". + body.put_u8(2); // ColNum + body.put_u8(0); // TableNum (not from a table) + body.put_u8(STATUS_EXPRESSION | STATUS_DIFFERENT_NAME); // Status + body.put_u8(2); // ColName length in characters + body.put_u16_le(u16::from(b'I')); + body.put_u16_le(u16::from(b'd')); + + buf.put_u16_le(body.len() as u16); + buf.extend_from_slice(&body); + + let mut reader = buf.into_sql_read_bytes(); + let token = TokenColInfo::decode(&mut reader).await.unwrap(); + + assert_eq!(token.columns.len(), 2); + + let first = &token.columns[0]; + assert_eq!(first.col_num, 1); + assert_eq!(first.table_num, 1); + assert!(first.is_key()); + assert!(!first.has_different_name()); + assert_eq!(first.col_name, None); + + let second = &token.columns[1]; + assert_eq!(second.col_num, 2); + assert_eq!(second.table_num, 0); + assert!(second.is_expression()); + assert!(second.has_different_name()); + assert!(!second.is_hidden()); + assert_eq!(second.col_name.as_deref(), Some("Id")); + } + + #[test] + fn colinfo_status_predicates() { + // Expression-only: is_expression true, the others false. + let expr = ColInfo { + col_num: 1, + table_num: 0, + status: STATUS_EXPRESSION, + col_name: None, + }; + assert!(expr.is_expression()); + assert!(!expr.is_key()); + assert!(!expr.is_hidden()); + + // Key-only: is_key true, the others false. + let key = ColInfo { + col_num: 1, + table_num: 0, + status: STATUS_KEY, + col_name: None, + }; + assert!(!key.is_expression()); + assert!(key.is_key()); + assert!(!key.is_hidden()); + + // Hidden-only: is_hidden true, the others false. + let hidden = ColInfo { + col_num: 1, + table_num: 0, + status: STATUS_HIDDEN, + col_name: None, + }; + assert!(!hidden.is_expression()); + assert!(!hidden.is_key()); + assert!(hidden.is_hidden()); + } + + #[tokio::test] + async fn decode_col_info_advances_consumed_by_name_bytes() { + // A different-name column (with a multi-character name) followed by a + // plain column. The `consumed += char_len * 2` update must be exact for + // the loop to read both columns. + let mut body = BytesMut::new(); + + // Column 1: expression + different name "abc" (3 chars => 6 bytes). + body.put_u8(1); // ColNum + body.put_u8(0); // TableNum + body.put_u8(STATUS_EXPRESSION | STATUS_DIFFERENT_NAME); // Status + body.put_u8(3); // ColName length in characters + body.put_u16_le(u16::from(b'a')); + body.put_u16_le(u16::from(b'b')); + body.put_u16_le(u16::from(b'c')); + + // Column 2: plain key column, no different name. + body.put_u8(2); // ColNum + body.put_u8(1); // TableNum + body.put_u8(STATUS_KEY); // Status + + let mut buf = BytesMut::new(); + buf.put_u16_le(body.len() as u16); + buf.extend_from_slice(&body); + + let mut reader = buf.into_sql_read_bytes(); + let token = TokenColInfo::decode(&mut reader).await.unwrap(); + + assert_eq!(token.columns.len(), 2); + assert_eq!(token.columns[0].col_num, 1); + assert_eq!(token.columns[0].col_name.as_deref(), Some("abc")); + assert_eq!(token.columns[1].col_num, 2); + assert!(token.columns[1].is_key()); + } + + #[tokio::test] + async fn decode_rejects_colname_overrunning_length() { + // A different-name entry claims a 3-char (6-byte) name, but the declared + // token length only covers the header + length byte. Decoding must fail + // with a protocol error rather than reading past `length`. + let mut body = BytesMut::new(); + body.put_u8(1); // ColNum + body.put_u8(0); // TableNum + body.put_u8(STATUS_EXPRESSION | STATUS_DIFFERENT_NAME); // Status + body.put_u8(3); // ColName length in characters (claims 6 bytes) + + // Declared length stops right after the length byte (4 bytes total), + // even though real name bytes follow on the wire. + let declared_len = body.len() as u16; + body.put_u16_le(u16::from(b'a')); + body.put_u16_le(u16::from(b'b')); + body.put_u16_le(u16::from(b'c')); + + let mut buf = BytesMut::new(); + buf.put_u16_le(declared_len); + buf.extend_from_slice(&body); + + let mut reader = buf.into_sql_read_bytes(); + let err = TokenColInfo::decode(&mut reader) + .await + .expect_err("ColName overrun must be rejected"); + assert!(matches!(err, crate::error::Error::Protocol(_))); + } + + #[tokio::test] + async fn decode_rejects_invalid_utf16_colname() { + // A different-name column whose ColName is a lone high surrogate is not + // valid UTF-16. Strict decoding must reject it with a protocol error + // rather than lossily substituting a replacement character. + let mut body = BytesMut::new(); + body.put_u8(1); // ColNum + body.put_u8(0); // TableNum + body.put_u8(STATUS_EXPRESSION | STATUS_DIFFERENT_NAME); // Status + body.put_u8(1); // ColName length in characters + body.put_u16_le(0xD800); // unpaired high surrogate + + let mut buf = BytesMut::new(); + buf.put_u16_le(body.len() as u16); + buf.extend_from_slice(&body); + + let mut reader = buf.into_sql_read_bytes(); + let err = TokenColInfo::decode(&mut reader) + .await + .expect_err("invalid UTF-16 ColName must be rejected"); + assert!(matches!(err, crate::error::Error::Protocol(_))); + } + + #[tokio::test] + async fn decode_rejects_truncated_entry_header() { + // Declared length of 2 bytes cannot hold a full 3-byte entry header. + let mut body = BytesMut::new(); + body.put_u8(1); // ColNum + body.put_u8(0); // TableNum (only 2 of 3 header bytes fit in `length`) + + let mut buf = BytesMut::new(); + buf.put_u16_le(2); + buf.extend_from_slice(&body); + buf.put_u8(0); // trailing status byte present on the wire but past length + + let mut reader = buf.into_sql_read_bytes(); + let err = TokenColInfo::decode(&mut reader) + .await + .expect_err("truncated header must be rejected"); + assert!(matches!(err, crate::error::Error::Protocol(_))); + } +} diff --git a/src/tds/codec/token/token_col_metadata.rs b/src/tds/codec/token/token_col_metadata.rs index 53ffdf1c6..f01e2c5b6 100644 --- a/src/tds/codec/token/token_col_metadata.rs +++ b/src/tds/codec/token/token_col_metadata.rs @@ -1,12 +1,8 @@ -use std::{ - borrow::{BorrowMut, Cow}, - fmt::Display, -}; +use std::{borrow::Cow, fmt::Display}; use crate::{ - error::Error, - tds::codec::{Encode, FixedLenType, TokenType, TypeInfo, VarLenType}, - Column, ColumnData, ColumnType, SqlReadBytes, + tds::codec::{encode_b_varchar, Encode, FixedLenType, TokenType, TypeInfo, VarLenType}, + Column, ColumnData, ColumnType, Error, SqlReadBytes, }; use asynchronous_codec::BytesMut; use bytes::BufMut; @@ -17,15 +13,36 @@ pub struct TokenColMetaData<'a> { pub columns: Vec>, } +/// Metadata for a single result/table column: its name plus the +/// [`BaseMetaDataColumn`] describing its type, size and flags. #[derive(Debug, Clone)] pub struct MetaDataColumn<'a> { + /// The type and flag metadata for the column. pub base: BaseMetaDataColumn, + /// The name of the column. pub col_name: Cow<'a, str>, } +impl<'a> MetaDataColumn<'a> { + /// The name of the column. + pub fn col_name(&self) -> &str { + self.col_name.as_ref() + } + + /// The [`BaseMetaDataColumn`] describing the column's type and flags + /// (nullability, identity, etc.). + pub fn base(&self) -> &BaseMetaDataColumn { + &self.base + } +} + impl<'a> Display for MetaDataColumn<'a> { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - write!(f, "{} ", self.col_name)?; + // Bracket-quote the identifier, escaping any literal `]` by doubling it + // (`]]`) per the T-SQL rule. Without this a column name containing `]` + // (e.g. `my]col`) would emit a malformed identifier `[my]col]`, breaking + // the `INSERT BULK (...)` column list this Display feeds into. + write!(f, "[{}] ", self.col_name.replace(']', "]]"))?; match &self.base.ty { TypeInfo::FixedLen(fixed) => match fixed { @@ -40,7 +57,10 @@ impl<'a> Display for MetaDataColumn<'a> { FixedLenType::Float8 => write!(f, "float")?, FixedLenType::Money4 => write!(f, "smallmoney")?, FixedLenType::Int8 => write!(f, "bigint")?, - FixedLenType::Null => unreachable!(), + // The TDS "null" fixed type carries no value; surface it as the + // int it decodes to rather than panicking (a bare `SELECT NULL` + // produces such a column). + FixedLenType::Null => write!(f, "int")?, }, TypeInfo::VarLenSized(ctx) => match ctx.r#type() { VarLenType::Bitn => write!(f, "bit")?, @@ -52,6 +72,10 @@ impl<'a> Display for MetaDataColumn<'a> { #[cfg(feature = "tds73")] VarLenType::Datetime2 => write!(f, "datetime2({})", ctx.len())?, VarLenType::Datetimen => write!(f, "datetime")?, + VarLenType::Money => match ctx.len() { + 4 => write!(f, "smallmoney")?, + _ => write!(f, "money")?, + }, #[cfg(feature = "tds73")] VarLenType::DatetimeOffsetn => write!(f, "datetimeoffset")?, VarLenType::BigVarBin => { @@ -85,15 +109,22 @@ impl<'a> Display for MetaDataColumn<'a> { 1 => write!(f, "tinyint")?, 2 => write!(f, "smallint")?, 4 => write!(f, "int")?, - 8 => write!(f, "bigint")?, - _ => unreachable!(), + _ => write!(f, "bigint")?, }, VarLenType::Floatn => match ctx.len() { 4 => write!(f, "real")?, - 8 => write!(f, "float")?, - _ => unreachable!(), + _ => write!(f, "float")?, }, - _ => unreachable!(), + VarLenType::SSVariant => write!(f, "sql_variant")?, + // Any other var-len type (e.g. Decimaln/Numericn arriving + // without the precision/scale they need, or Xml/Udt appearing + // in a sized context) has no valid SQL type name we can emit + // here. Emitting a bogus name (such as the Debug name) would + // produce an invalid `INSERT BULK` statement, so return a + // formatting error instead. This keeps the library from ever + // panicking on server-supplied metadata while refusing to emit + // invalid SQL. + _ => return Err(std::fmt::Error), }, TypeInfo::VarLenSizedPrecision { ty, @@ -101,24 +132,71 @@ impl<'a> Display for MetaDataColumn<'a> { precision, scale, } => match ty { - VarLenType::Decimaln => write!(f, "decimal({},{})", precision, scale)?, VarLenType::Numericn => write!(f, "numeric({},{})", precision, scale)?, - _ => unreachable!(), + // Decimaln, and any other precision-carrying type. + _ => write!(f, "decimal({},{})", precision, scale)?, }, TypeInfo::Xml { .. } => write!(f, "xml")?, + TypeInfo::Udt(info) => write!(f, "{}.{}", info.schema_name, info.type_name)?, } Ok(()) } } +/// Describes the type and flags of a column, exposing metadata such as the +/// column type (including size, precision and scale), whether the column is +/// nullable and whether it is an identity column. #[derive(Debug, Clone)] pub struct BaseMetaDataColumn { + /// The set of [`ColumnFlag`]s describing the column (nullability, identity, + /// updateability, and so on). pub flags: BitFlags, + /// The type of the column, including its size, precision and scale where + /// applicable. pub ty: TypeInfo, + /// Destination table name, used only on the *encode* (bulk-load) path. Per + /// MS-TDS §2.2.7.4, `text`/`ntext`/`image` columns carry a `TableName` + /// element in COLMETADATA; the bulk-insert code sets this to the target + /// table so [`BaseMetaDataColumn::encode`] can emit it. It is always `None` + /// on the decode (server→client) path, which never needs it. + pub table_name: Option, } impl BaseMetaDataColumn { + /// The type of the column, including its size, precision and scale where + /// applicable. + pub fn ty(&self) -> &TypeInfo { + &self.ty + } + + /// The set of flags describing the column. + pub fn flags(&self) -> BitFlags { + self.flags + } + + /// `true` if the column accepts `NULL` values. + pub fn is_nullable(&self) -> bool { + self.flags.contains(ColumnFlag::Nullable) + } + + /// `true` if the column is an identity column. + pub fn is_identity(&self) -> bool { + self.flags.contains(ColumnFlag::Identity) + } + + /// `true` if the column is writeable (e.g. usable as a bulk-insert target). + /// + /// The COLMETADATA `usUpdateable` sub-field (MS-TDS §2.2.7.4) is one of + /// read-only (0), read/write (1), or **unknown** (2). SQL Server reports + /// `unknown` for many legitimate bulk-target columns rather than an explicit + /// read/write, so both read/write and unknown are treated as writeable here; + /// only an explicit read-only column (neither flag set) is excluded. + pub fn is_updateable(&self) -> bool { + self.flags.contains(ColumnFlag::Updateable) + || self.flags.contains(ColumnFlag::UpdateableUnknown) + } + pub(crate) fn null_value(&self) -> ColumnData<'static> { match &self.ty { TypeInfo::FixedLen(ty) => match ty { @@ -167,11 +245,16 @@ impl BaseMetaDataColumn { VarLenType::NVarchar => ColumnData::String(None), VarLenType::NChar => ColumnData::String(None), VarLenType::Xml => ColumnData::Xml(None), - VarLenType::Udt => todo!("User-defined types not supported"), + // A null CLR UDT carries no payload; surface it as a null + // binary, matching `udt::decode` (which yields + // `ColumnData::Binary`). + VarLenType::Udt => ColumnData::Binary(None), VarLenType::Text => ColumnData::String(None), VarLenType::Image => ColumnData::Binary(None), VarLenType::NText => ColumnData::String(None), - VarLenType::SSVariant => todo!(), + // A null `sql_variant` carries no base type, so surface a + // generic null value. + VarLenType::SSVariant => ColumnData::String(None), }, TypeInfo::VarLenSizedPrecision { ty, .. } => match ty { VarLenType::Guid => ColumnData::Guid(None), @@ -197,13 +280,19 @@ impl BaseMetaDataColumn { VarLenType::NVarchar => ColumnData::String(None), VarLenType::NChar => ColumnData::String(None), VarLenType::Xml => ColumnData::Xml(None), - VarLenType::Udt => todo!("User-defined types not supported"), + // A null CLR UDT carries no payload; surface it as a null + // binary, matching `udt::decode` (which yields + // `ColumnData::Binary`). + VarLenType::Udt => ColumnData::Binary(None), VarLenType::Text => ColumnData::String(None), VarLenType::Image => ColumnData::Binary(None), VarLenType::NText => ColumnData::String(None), - VarLenType::SSVariant => todo!(), + // A null `sql_variant` carries no base type, so surface a + // generic null value. + VarLenType::SSVariant => ColumnData::String(None), }, TypeInfo::Xml { .. } => ColumnData::Xml(None), + TypeInfo::Udt(_) => ColumnData::Binary(None), } } } @@ -226,18 +315,9 @@ impl<'a> Encode for MetaDataColumn<'a> { dst.put_u32_le(0); self.base.encode(dst)?; - let len_pos = dst.len(); - let mut length = 0u8; - - dst.put_u8(length); - - for chr in self.col_name.encode_utf16() { - length += 1; - dst.put_u16_le(chr); - } - - let dst: &mut [u8] = dst.borrow_mut(); - dst[len_pos] = length; + // ColName is a B_VARCHAR (u8 code-unit count); reject an over-long name + // rather than wrap the counter and desync the wire. + encode_b_varchar(dst, &self.col_name)?; Ok(()) } @@ -246,12 +326,63 @@ impl<'a> Encode for MetaDataColumn<'a> { impl Encode for BaseMetaDataColumn { fn encode(self, dst: &mut BytesMut) -> crate::Result<()> { dst.put_u16_le(BitFlags::bits(self.flags)); + + // `text`/`ntext`/`image` columns carry a TableName element after the + // TYPE_INFO. The client->server INSERT BULK COLMETADATA uses a DIFFERENT + // TableName shape than the server->client result COLMETADATA the decode + // side reads: + // + // * READ (server->client): NumParts BYTE, then NumParts US_VARCHARs + // (see `BaseMetaDataColumn::decode` below and MS-TDS §2.2.7.4). + // * WRITE (client->server bulk): a single bare US_VARCHAR carrying the + // whole destination table name, with NO NumParts byte and no + // splitting on `.`. + // + // This asymmetry matches Microsoft's own go-mssqldb bulk-copy path + // (`createColMetadata` in bulkcopy.go), verified against real SQL Server: + // a `uint16` code-unit count + UTF-16LE bytes, no NumParts. Emitting the + // spec/result-style `NumParts` byte here inserts one extra leading 0x01 + // the server does not expect, mis-parsing the column and overrunning the + // row stream (TDS error 4804). Only these three types carry a TableName, + // so the desync hits text/ntext/image bulk exclusively. + // + // Capture whether this is such a column before `encode` consumes `self.ty`. + let emits_table_name = matches!( + &self.ty, + TypeInfo::VarLenSized(cx) + if matches!(cx.r#type(), VarLenType::Text | VarLenType::NText | VarLenType::Image) + ); + self.ty.encode(dst)?; + if emits_table_name { + let table_name = self.table_name.as_deref().unwrap_or(""); + encode_us_varchar(dst, table_name)?; + } + Ok(()) } } +/// Encode a US_VARCHAR: a `u16` length in UTF-16 code units followed by the +/// UTF-16LE characters. The length field is a `u16`, so a part longer than +/// `u16::MAX` code units cannot be represented and is rejected. +fn encode_us_varchar(dst: &mut BytesMut, s: &str) -> crate::Result<()> { + let units = s.encode_utf16().count(); + if units > u16::MAX as usize { + return Err(Error::BulkInput( + format!("table name is too long ({units} UTF-16 code units, max 65535)").into(), + )); + } + + dst.put_u16_le(units as u16); + for chr in s.encode_utf16() { + dst.put_u16_le(chr); + } + + Ok(()) +} + /// A setting a column can hold. #[bitflags] #[repr(u16)] @@ -262,23 +393,33 @@ pub enum ColumnFlag { /// Set for string columns with binary collation and always for the XML data /// type. CaseSensitive = 1 << 1, - /// If column is writeable. - Updateable = 1 << 3, - /// Column modification status unknown. - UpdateableUnknown = 1 << 4, + /// If column is writeable. This is value 1 (0b01) of the 2-bit + /// `usUpdateable` sub-field (MS-TDS §2.2.7.4), i.e. the low bit at position + /// 2 (0x04): `0 = read-only, 1 = read/write, 2 = unknown`. + Updateable = 1 << 2, + /// Column modification status unknown. This is value 2 (0b10) of the 2-bit + /// `usUpdateable` sub-field, i.e. the high bit at position 3 (0x08). + UpdateableUnknown = 1 << 3, /// Column is an identity. - Identity = 1 << 5, - /// Coulumn is computed. - Computed = 1 << 7, + Identity = 1 << 4, + /// Column is computed. Per MS-TDS §2.2.7.4 `fComputed` is bit 5 (0x0020), + /// immediately after `fIdentity` and before the 2-bit `usReservedODBC` + /// field (bits 6-7). Introduced in TDS 7.2. + Computed = 1 << 5, /// Column is a fixed-length common language runtime user-defined type (CLR - /// UDT). - FixedLenClrType = 1 << 10, - /// Column is the special XML column for the sparse column set. - SparseColumnSet = 1 << 11, + /// UDT). Per MS-TDS §2.2.7.4 `fFixedLenCLRType` is bit 8 (0x0100), directly + /// after `usReservedODBC` (bits 6-7). Introduced in TDS 7.2. + FixedLenClrType = 1 << 8, + /// Column is the special XML column for the sparse column set. Per + /// MS-TDS §2.2.7.4 `fSparseColumnSet` is bit 10 (0x0400): bit 9 is a + /// reserved bit (`FRESERVEDBIT`). Introduced in TDS 7.3.B. + SparseColumnSet = 1 << 10, /// Column is encrypted transparently and has to be decrypted to view the /// plaintext value. This flag is valid when the column encryption feature - /// is negotiated between client and server and is turned on. - Encrypted = 1 << 12, + /// is negotiated between client and server and is turned on. Per + /// MS-TDS §2.2.7.4 `fEncrypted` is bit 11 (0x0800), directly after + /// `fSparseColumnSet`. Introduced in TDS 7.4. + Encrypted = 1 << 11, /// Column is part of a hidden primary key created to support a T-SQL SELECT /// statement containing FOR BROWSE. Hidden = 1 << 13, @@ -295,9 +436,17 @@ impl TokenColMetaData<'static> { R: SqlReadBytes + Unpin, { let column_count = src.read_u16_le().await?; - let mut columns = Vec::with_capacity(column_count as usize); + // `column_count` is an untrusted u16 (up to 65535); cap the up-front + // reservation so a hostile COLMETADATA token can't force a large + // transient allocation before the column bodies arrive. The Vec still + // grows as real columns are decoded. + let mut columns = Vec::with_capacity( + (column_count as usize).min(crate::tds::codec::column_data::MAX_PREALLOC), + ); - if column_count > 0 && column_count < 0xffff { + // `0xffff` is the "no metadata" sentinel; any other count drives the + // loop directly (a count of 0 simply iterates zero times). + if column_count < 0xffff { for _ in 0..column_count { let base = BaseMetaDataColumn::decode(src).await?; let col_name = Cow::from(src.read_b_varchar().await?); @@ -328,8 +477,11 @@ impl BaseMetaDataColumn { let _user_ty = src.read_u32_le().await?; - let flags = BitFlags::from_bits(src.read_u16_le().await?) - .map_err(|_| Error::Protocol("column metadata: invalid flags".into()))?; + // The COLMETADATA `Flags` field (MS-TDS §2.2.7.4) is a 16-bit field that + // includes reserved / ODBC bits the server may set and which future + // protocol revisions may extend. Truncate to the flags we model rather + // than rejecting the whole token on an unrecognized bit. + let flags = BitFlags::from_bits_truncate(src.read_u16_le().await?); let ty = TypeInfo::decode(src).await?; @@ -344,6 +496,1007 @@ impl BaseMetaDataColumn { }; }; - Ok(BaseMetaDataColumn { flags, ty }) + Ok(BaseMetaDataColumn { + flags, + ty, + table_name: None, + }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use crate::tds::codec::type_info::VarLenContext; + use crate::tds::Collation; + use std::fmt::Write as _; + + // Build the on-wire bytes a US_VARCHAR should produce for `s`. + fn us_varchar_bytes(s: &str) -> Vec { + let mut out = Vec::new(); + let units: Vec = s.encode_utf16().collect(); + out.extend_from_slice(&(units.len() as u16).to_le_bytes()); + for u in units { + out.extend_from_slice(&u.to_le_bytes()); + } + out + } + + fn text_column(table_name: Option<&str>) -> BaseMetaDataColumn { + BaseMetaDataColumn { + flags: ColumnFlag::Nullable.into(), + ty: TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Text, 2147483647, None)), + table_name: table_name.map(str::to_string), + } + } + + // On the client->server INSERT BULK path a `text` column's metadata must end + // with the TableName as a SINGLE bare US_VARCHAR carrying the whole name -- + // NO NumParts byte and no splitting on `.` -- matching go-mssqldb's verified + // bulkcopy `createColMetadata`. A dotted name is sent verbatim as one part. + #[test] + fn text_column_emits_dotted_name_as_bare_us_varchar() { + let mut buf = BytesMut::new(); + text_column(Some("dbo.MyTable")).encode(&mut buf).unwrap(); + + // Whole name as one US_VARCHAR, with no leading NumParts byte. + let expected_tail = us_varchar_bytes("dbo.MyTable"); + + assert!( + buf.ends_with(&expected_tail), + "buffer {:02x?} did not end with bare US_VARCHAR TableName {:02x?}", + &buf[..], + expected_tail + ); + } + + // A single, unqualified name is likewise a single bare US_VARCHAR (no + // NumParts byte). + #[test] + fn text_column_single_name_is_bare_us_varchar() { + let mut buf = BytesMut::new(); + text_column(Some("##bulk_test")).encode(&mut buf).unwrap(); + + let expected_tail = us_varchar_bytes("##bulk_test"); + assert!(buf.ends_with(&expected_tail), "got {:02x?}", &buf[..]); + } + + // Non-text columns must not emit any TableName, even if one is set: the bytes + // must be identical with and without a table name. + #[test] + fn non_text_column_never_emits_table_name() { + let with = { + let mut buf = BytesMut::new(); + BaseMetaDataColumn { + flags: ColumnFlag::Nullable.into(), + ty: TypeInfo::FixedLen(FixedLenType::Int4), + table_name: Some("dbo.MyTable".to_string()), + } + .encode(&mut buf) + .unwrap(); + buf.to_vec() + }; + let without = { + let mut buf = BytesMut::new(); + BaseMetaDataColumn { + flags: ColumnFlag::Nullable.into(), + ty: TypeInfo::FixedLen(FixedLenType::Int4), + table_name: None, + } + .encode(&mut buf) + .unwrap(); + buf.to_vec() + }; + assert_eq!(with, without); + } + + // A US_VARCHAR length field is a u16; an over-long name must error rather + // than wrap the length (which would desync the wire). + #[test] + fn text_column_rejects_over_long_table_name() { + let long = "a".repeat(u16::MAX as usize + 1); + let mut buf = BytesMut::new(); + let err = text_column(Some(&long)).encode(&mut buf).unwrap_err(); + assert!(matches!(err, Error::BulkInput(_)), "got {err:?}"); + } + + fn meta(ty: TypeInfo, name: &'static str) -> MetaDataColumn<'static> { + MetaDataColumn { + base: BaseMetaDataColumn { + flags: ColumnFlag::Nullable.into(), + ty, + table_name: None, + }, + col_name: Cow::Borrowed(name), + } + } + + #[test] + fn display_var_len_unknown_type_yields_err_not_panic() { + // A VarLenSized carrying a type with no valid sized SQL representation + // (e.g. Decimaln/Numericn without precision/scale) must NOT panic and + // must NOT emit a bogus SQL type name. Formatting it returns a + // `std::fmt::Error` so the caller gets an `Err`, never a panic and + // never invalid SQL. + for ty in [VarLenType::Decimaln, VarLenType::Numericn] { + let col = meta(TypeInfo::VarLenSized(VarLenContext::new(ty, 17, None)), "c"); + let mut out = String::new(); + let result = write!(out, "{col}"); + assert!( + result.is_err(), + "expected Err for unhandled var-len type {ty:?}, got Ok({out:?})" + ); + } + } + + #[test] + fn display_money_columns_never_panic_and_render_expected() { + // MONEY/SMALLMONEY reach the Display impl both as FixedLen (Money / + // Money4) and as VarLenSized (Money with len 8 / 4). A money column + // must never fall through to a catch-all; each renders its exact SQL + // type name for the bulk `INSERT` column list. Assert the rendered type + // token (with a leading space so `money` can't match `smallmoney`) so + // the test is agnostic to how the column-name prefix is quoted. + let cases = vec![ + (TypeInfo::FixedLen(FixedLenType::Money), " money"), + (TypeInfo::FixedLen(FixedLenType::Money4), " smallmoney"), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Money, 8, None)), + " money", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Money, 4, None)), + " smallmoney", + ), + ]; + + for (ty, expected_suffix) in cases { + // Must not panic, and must end with the expected SQL type token. + let rendered = format!("{}", meta(ty, "c")); + assert!( + rendered.ends_with(expected_suffix), + "expected {rendered:?} to end with {expected_suffix:?}" + ); + } + } + + #[test] + fn display_text_ntext_image_render_expected_type_names() { + // The LOB var-len types feed the bulk `INSERT` column list and must + // render their exact SQL type names (not a sized form), complementing + // the encode-path `TableName` tests above. Assert on a leading-space + // suffix so the column-name prefix quoting is irrelevant. + let cases = [ + (VarLenType::Text, " text"), + (VarLenType::NText, " ntext"), + (VarLenType::Image, " image"), + ]; + + for (ty, expected_suffix) in cases { + let rendered = format!( + "{}", + meta( + TypeInfo::VarLenSized(VarLenContext::new(ty, 2147483647, None)), + "c", + ) + ); + assert!( + rendered.ends_with(expected_suffix), + "expected {rendered:?} to end with {expected_suffix:?}", + ); + } + } + + fn column(name: &'static str) -> MetaDataColumn<'static> { + MetaDataColumn { + base: BaseMetaDataColumn { + flags: BitFlags::empty(), + ty: TypeInfo::FixedLen(FixedLenType::Int4), + table_name: None, + }, + col_name: Cow::Borrowed(name), + } + } + + // Bit-exact layout of the COLMETADATA `Flags` field per MS-TDS §2.2.7.4, + // in least-significant-bit order: + // bit 0 fNullable + // bit 1 fCaseSen + // bits 2-3 usUpdateable (0=read-only, 1=read/write, 2=unknown) + // bit 4 fIdentity + // bit 5 fComputed + // bits 6-7 usReservedODBC + // bit 8 fFixedLenCLRType + // ... + // bit 13 fHidden + // bit 14 fKey + // bit 15 fNullableUnknown + // The `usUpdateable` sub-field is a 2-bit value, so read/write (value 1) is + // the low bit 0x04 and unknown (value 2) is the high bit 0x08. A server + // reporting a genuinely writeable column sets 0x04; `is_updateable()` must + // therefore key off 0x04, not 0x08. + #[test] + fn column_flag_bit_positions_match_ms_tds() { + assert_eq!(BitFlags::bits(BitFlags::from(ColumnFlag::Nullable)), 0x0001); + assert_eq!( + BitFlags::bits(BitFlags::from(ColumnFlag::CaseSensitive)), + 0x0002 + ); + // usUpdateable: read/write = 0x04 (bit 2), unknown = 0x08 (bit 3). + assert_eq!( + BitFlags::bits(BitFlags::from(ColumnFlag::Updateable)), + 0x0004 + ); + assert_eq!( + BitFlags::bits(BitFlags::from(ColumnFlag::UpdateableUnknown)), + 0x0008 + ); + assert_eq!(BitFlags::bits(BitFlags::from(ColumnFlag::Identity)), 0x0010); + } + + // `is_updateable()` means "not read-only", i.e. a valid bulk-insert target. + // usUpdateable read/write (1, wire bit 0x04) and unknown (2, wire bit 0x08) + // are both writeable; only explicit read-only (0, neither bit) is excluded. + // SQL Server reports `unknown` for many real bulk-target columns, so treating + // it as writeable is required (keying off the read/write bit alone drops + // every column and breaks bulk insert entirely). + #[test] + fn is_updateable_true_for_read_write_and_unknown_false_for_read_only() { + let read_write = BaseMetaDataColumn { + flags: BitFlags::from_bits_truncate(0x0004), + ty: TypeInfo::FixedLen(FixedLenType::Int4), + table_name: None, + }; + assert!(read_write.is_updateable()); + + let unknown = BaseMetaDataColumn { + flags: BitFlags::from_bits_truncate(0x0008), + ty: TypeInfo::FixedLen(FixedLenType::Int4), + table_name: None, + }; + assert!(unknown.is_updateable()); + + let read_only = BaseMetaDataColumn { + flags: BitFlags::from_bits_truncate(0x0000), + ty: TypeInfo::FixedLen(FixedLenType::Int4), + table_name: None, + }; + assert!(!read_only.is_updateable()); + } + + // Byte-exact decode of a COLMETADATA `Flags` USHORT with a known set of + // sub-fields, constructed per the MS-TDS §2.2.7.4 LSB-first layout: + // bit 0 fNullable bit 5 fComputed + // bit 1 fCaseSen bits 6-7 usReservedODBC + // bits 2-3 usUpdateable bit 8 fFixedLenCLRType + // bit 4 fIdentity bit 9 (reserved) + // bit 10 fSparseColumnSet bit 13 fHidden + // bit 11 fEncrypted bit 14 fKey + // bit 12 (usReserved3) bit 15 fNullableUnknown + // This locks the wire layout: a regression that shifts any bit changes the + // decoded flag set and fails here. + #[test] + fn column_flags_decode_byte_exact_low_and_mid_bits() { + // fNullable(0) | usUpdateable=Read/Write(bit2) | fIdentity(4) + // | fComputed(5) | fFixedLenCLRType(8) | fSparseColumnSet(10) + // | fEncrypted(11) ==> 0x0D35, little-endian on the wire. + let wire: [u8; 2] = [0x35, 0x0D]; + let flags = BitFlags::::from_bits_truncate(u16::from_le_bytes(wire)); + + let col = BaseMetaDataColumn { + flags, + ty: TypeInfo::FixedLen(FixedLenType::Int4), + table_name: None, + }; + + assert!(col.is_nullable()); + assert!(col.is_updateable()); // usUpdateable == 1 (Read/Write, bit 2) + assert!(col.is_identity()); + assert!(flags.contains(ColumnFlag::Computed)); + assert!(flags.contains(ColumnFlag::FixedLenClrType)); + assert!(flags.contains(ColumnFlag::SparseColumnSet)); + assert!(flags.contains(ColumnFlag::Encrypted)); + + // Bits that were NOT set must not read back as set. + assert!(!flags.contains(ColumnFlag::CaseSensitive)); + assert!(!flags.contains(ColumnFlag::UpdateableUnknown)); + assert!(!flags.contains(ColumnFlag::Hidden)); + assert!(!flags.contains(ColumnFlag::Key)); + assert!(!flags.contains(ColumnFlag::NullableUnknown)); + } + + // Companion to the above, exercising the CaseSen bit, the *high* bit of + // usUpdateable (value 2 = unknown) and the top three flags introduced in + // TDS 7.2 (fHidden=13, fKey=14, fNullableUnknown=15). + #[test] + fn column_flags_decode_byte_exact_high_bits() { + // fCaseSen(1) | usUpdateable=Unknown(bit3) | fHidden(13) | fKey(14) + // | fNullableUnknown(15) ==> 0xE00A, little-endian on the wire. + let wire: [u8; 2] = [0x0A, 0xE0]; + let flags = BitFlags::::from_bits_truncate(u16::from_le_bytes(wire)); + + assert!(flags.contains(ColumnFlag::CaseSensitive)); + assert!(flags.contains(ColumnFlag::UpdateableUnknown)); + assert!(flags.contains(ColumnFlag::Hidden)); + assert!(flags.contains(ColumnFlag::Key)); + assert!(flags.contains(ColumnFlag::NullableUnknown)); + + // usUpdateable == 2 (unknown) means NOT read/write. + assert!(!flags.contains(ColumnFlag::Updateable)); + assert!(!flags.contains(ColumnFlag::Nullable)); + assert!(!flags.contains(ColumnFlag::Identity)); + assert!(!flags.contains(ColumnFlag::Computed)); + } + + // ColName is a B_VARCHAR (u8 length); a name longer than 255 UTF-16 code + // units must error rather than wrap the counter and desync the wire. + #[test] + fn metadata_column_rejects_over_long_col_name() { + let col = MetaDataColumn { + base: BaseMetaDataColumn { + flags: BitFlags::empty(), + ty: TypeInfo::FixedLen(FixedLenType::Int4), + table_name: None, + }, + col_name: Cow::Owned("a".repeat(256)), + }; + + let mut buf = BytesMut::new(); + let err = col.encode(&mut buf).unwrap_err(); + assert!(matches!(err, Error::Protocol(_)), "got {err:?}"); + } + + // A col_name of exactly 255 units is the boundary and must still encode with + // the correct u8 length prefix following the u32 user-type and TYPE_INFO. + #[test] + fn metadata_column_accepts_max_length_col_name() { + let col = MetaDataColumn { + base: BaseMetaDataColumn { + flags: BitFlags::empty(), + ty: TypeInfo::FixedLen(FixedLenType::Int4), + table_name: None, + }, + col_name: Cow::Owned("a".repeat(255)), + }; + + let mut buf = BytesMut::new(); + col.encode(&mut buf).expect("255-unit col_name must encode"); + // Layout: u32 user-type (4) + flags u16 (2) + FixedLen TYPE_INFO (1) + + // B_VARCHAR length prefix. + assert_eq!(buf[7], 255); + } + + #[test] + fn display_escapes_closing_bracket_in_column_name() { + // A `]` in the column name must be doubled so the bracket-quoted + // identifier stays well-formed for the `INSERT BULK (...)` column list. + assert_eq!(format!("{}", column("my]col")), "[my]]col] int"); + } + + #[test] + fn display_leaves_plain_column_name_unchanged() { + assert_eq!(format!("{}", column("foo")), "[foo] int"); + } + + #[tokio::test] + async fn round_trip_via_encode_decode() { + let cmd = TokenColMetaData { + columns: vec![ + meta(TypeInfo::FixedLen(FixedLenType::Int4), "id"), + meta( + TypeInfo::VarLenSized(VarLenContext::new( + VarLenType::NVarchar, + 4000, + Some(Collation::new(13632521, 52)), + )), + "name", + ), + ], + }; + + // Build a decodable buffer: column count followed by each column. The + // MetaDataColumn encoder writes the leading user-type u32 that the + // decoder expects. + let mut buf = BytesMut::new(); + buf.put_u16_le(cmd.columns.len() as u16); + for col in cmd.columns.iter().cloned() { + col.encode(&mut buf).unwrap(); + } + + let decoded = TokenColMetaData::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + assert_eq!(decoded.columns.len(), 2); + assert_eq!(decoded.columns[0].col_name, "id"); + assert_eq!(decoded.columns[1].col_name, "name"); + + let columns: Vec<_> = decoded.columns().collect(); + assert_eq!(columns.len(), 2); + assert_eq!(columns[0].name(), "id"); + } + + #[test] + fn encode_writes_token_header_and_column_count() { + let cmd = TokenColMetaData { + columns: vec![ + meta(TypeInfo::FixedLen(FixedLenType::Int4), "id"), + meta(TypeInfo::FixedLen(FixedLenType::Bit), "flag"), + ], + }; + + let mut buf = BytesMut::new(); + cmd.encode(&mut buf).unwrap(); + + // First the ColMetaData token byte, then the little-endian column count. + assert_eq!(buf[0], TokenType::ColMetaData as u8); + assert_eq!(u16::from_le_bytes([buf[1], buf[2]]), 2); + // The two column bodies follow the 3-byte header. + assert!(buf.len() > 3); + } + + #[tokio::test] + async fn zero_columns_yields_empty() { + let mut buf = BytesMut::new(); + buf.put_u16_le(0); + + let decoded = TokenColMetaData::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + assert!(decoded.columns.is_empty()); + } + + #[tokio::test] + async fn text_column_reads_table_name_parts() { + let mut buf = BytesMut::new(); + buf.put_u16_le(1); // one column + + // user_ty + flags + buf.put_u32_le(0); + buf.put_u16_le(BitFlags::bits(BitFlags::from(ColumnFlag::Nullable))); + + // type info for a text column with collation + let ti = TypeInfo::VarLenSized(VarLenContext::new( + VarLenType::Text, + 2147483647, + Some(Collation::new(13632521, 52)), + )); + ti.encode(&mut buf).unwrap(); + + // table name: one part, us_varchar "dbo" + buf.put_u8(1); + let part: Vec = "dbo".encode_utf16().collect(); + buf.put_u16_le(part.len() as u16); + for c in part { + buf.put_u16_le(c); + } + + // column name (b_varchar) + let name: Vec = "body".encode_utf16().collect(); + buf.put_u8(name.len() as u8); + for c in name { + buf.put_u16_le(c); + } + + let decoded = TokenColMetaData::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + assert_eq!(decoded.columns.len(), 1); + assert_eq!(decoded.columns[0].col_name, "body"); + } + + #[test] + fn display_formats_various_types() { + let cases = vec![ + (TypeInfo::FixedLen(FixedLenType::Int4), "c int"), + (TypeInfo::FixedLen(FixedLenType::Bit), "c bit"), + (TypeInfo::FixedLen(FixedLenType::Float8), "c float"), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Intn, 1, None)), + "c tinyint", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Intn, 4, None)), + "c int", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Floatn, 4, None)), + "c real", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Guid, 16, None)), + "c uniqueidentifier", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::BigVarBin, 100, None)), + "c varbinary(100)", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::BigVarBin, 100000, None)), + "c varbinary(max)", + ), + ( + TypeInfo::VarLenSizedPrecision { + ty: VarLenType::Decimaln, + size: 17, + precision: 18, + scale: 2, + }, + "c decimal(18,2)", + ), + // Numericn must render as `numeric(...)`, distinct from the + // `decimal(...)` fallback that every other precision type uses. + ( + TypeInfo::VarLenSizedPrecision { + ty: VarLenType::Numericn, + size: 9, + precision: 10, + scale: 4, + }, + "c numeric(10,4)", + ), + ( + TypeInfo::Xml { + schema: None, + size: 0, + }, + "c xml", + ), + ]; + + for (ty, expected) in cases { + // Display brackets the column name for use in bulk `INSERT` statements. + let expected = expected.replacen("c ", "[c] ", 1); + assert_eq!(format!("{}", meta(ty, "c")), expected); + } + } + + #[test] + fn null_value_maps_types() { + let fixed = BaseMetaDataColumn { + flags: BitFlags::empty(), + ty: TypeInfo::FixedLen(FixedLenType::Int4), + table_name: None, + }; + assert_eq!(fixed.null_value(), ColumnData::I32(None)); + + let varlen = BaseMetaDataColumn { + flags: BitFlags::empty(), + ty: TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Intn, 2, None)), + table_name: None, + }; + assert_eq!(varlen.null_value(), ColumnData::I16(None)); + + // Each Intn width maps to a distinct integer column; 1 and 4 sit either + // side of the `_ => I64` fallback and pin their own arms. + let tinyint = BaseMetaDataColumn { + flags: BitFlags::empty(), + ty: TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Intn, 1, None)), + table_name: None, + }; + assert_eq!(tinyint.null_value(), ColumnData::U8(None)); + + let int4 = BaseMetaDataColumn { + flags: BitFlags::empty(), + ty: TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Intn, 4, None)), + table_name: None, + }; + assert_eq!(int4.null_value(), ColumnData::I32(None)); + + let int8 = BaseMetaDataColumn { + flags: BitFlags::empty(), + ty: TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Intn, 8, None)), + table_name: None, + }; + assert_eq!(int8.null_value(), ColumnData::I64(None)); + + // Floatn splits on width too: 4 bytes is F32, anything else F64. + let real = BaseMetaDataColumn { + flags: BitFlags::empty(), + ty: TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Floatn, 4, None)), + table_name: None, + }; + assert_eq!(real.null_value(), ColumnData::F32(None)); + + let double = BaseMetaDataColumn { + flags: BitFlags::empty(), + ty: TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Floatn, 8, None)), + table_name: None, + }; + assert_eq!(double.null_value(), ColumnData::F64(None)); + + let guid = BaseMetaDataColumn { + flags: BitFlags::empty(), + ty: TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Guid, 16, None)), + table_name: None, + }; + assert_eq!(guid.null_value(), ColumnData::Guid(None)); + } + + #[test] + fn null_value_maps_precision_and_xml_and_udt() { + let precision = BaseMetaDataColumn { + flags: BitFlags::empty(), + ty: TypeInfo::VarLenSizedPrecision { + ty: VarLenType::Numericn, + size: 17, + precision: 18, + scale: 2, + }, + table_name: None, + }; + assert_eq!(precision.null_value(), ColumnData::Numeric(None)); + + let xml = BaseMetaDataColumn { + flags: BitFlags::empty(), + ty: TypeInfo::Xml { + schema: None, + size: 0, + }, + table_name: None, + }; + assert_eq!(xml.null_value(), ColumnData::Xml(None)); + + let udt = BaseMetaDataColumn { + flags: BitFlags::empty(), + ty: TypeInfo::Udt(crate::tds::codec::type_info::UdtInfo { + max_byte_size: 0, + db_name: "db".into(), + schema_name: "dbo".into(), + type_name: "T".into(), + assembly_qualified_name: "A".into(), + }), + table_name: None, + }; + assert_eq!(udt.null_value(), ColumnData::Binary(None)); + } + + #[test] + fn base_meta_data_column_flag_accessors() { + let base = BaseMetaDataColumn { + flags: ColumnFlag::Nullable | ColumnFlag::Identity | ColumnFlag::Updateable, + ty: TypeInfo::FixedLen(FixedLenType::Int4), + table_name: None, + }; + + assert!(base.is_nullable()); + assert!(base.is_identity()); + assert!(base.is_updateable()); + assert_eq!(base.ty(), &TypeInfo::FixedLen(FixedLenType::Int4)); + assert_eq!(base.flags(), base.flags); + + let base2 = BaseMetaDataColumn { + flags: BitFlags::empty(), + ty: TypeInfo::FixedLen(FixedLenType::Int4), + table_name: None, + }; + assert!(!base2.is_nullable()); + assert!(!base2.is_identity()); + assert!(!base2.is_updateable()); + } + + #[test] + fn meta_data_column_accessors() { + let m = meta(TypeInfo::FixedLen(FixedLenType::Int4), "id"); + assert_eq!(m.col_name(), "id"); + assert!(m.base().is_nullable()); + } + + #[test] + fn display_formats_more_types() { + let cases = vec![ + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Bitn, 1, None)), + "c bit", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Datetimen, 8, None)), + "c datetime", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Money, 4, None)), + "c smallmoney", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Money, 8, None)), + "c money", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::BigVarChar, 100, None)), + "c varchar(100)", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::BigVarChar, 100000, None)), + "c varchar(max)", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::BigBinary, 10, None)), + "c binary(10)", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::BigChar, 10, None)), + "c char(10)", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::NVarchar, 100, None)), + "c nvarchar(100)", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::NVarchar, 100000, None)), + "c nvarchar(max)", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::NChar, 10, None)), + "c nchar(10)", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Text, 0, None)), + "c text", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Image, 0, None)), + "c image", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::NText, 0, None)), + "c ntext", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Intn, 2, None)), + "c smallint", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Intn, 8, None)), + "c bigint", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Floatn, 8, None)), + "c float", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::SSVariant, 0, None)), + "c sql_variant", + ), + ]; + + for (ty, expected) in cases { + let expected = expected.replacen("c ", "[c] ", 1); + assert_eq!(format!("{}", meta(ty, "c")), expected); + } + } + + #[test] + fn display_formats_udt_and_decimaln() { + let udt = TypeInfo::Udt(crate::tds::codec::type_info::UdtInfo { + max_byte_size: 0, + db_name: "db".into(), + schema_name: "dbo".into(), + type_name: "MyType".into(), + assembly_qualified_name: "asm".into(), + }); + assert_eq!(format!("{}", meta(udt, "c")), "[c] dbo.MyType"); + + let decimaln = TypeInfo::VarLenSizedPrecision { + ty: VarLenType::Decimaln, + size: 17, + precision: 10, + scale: 4, + }; + assert_eq!(format!("{}", meta(decimaln, "c")), "[c] decimal(10,4)"); + } + + #[tokio::test] + async fn decode_all_ones_column_count_yields_empty() { + // column_count == 0xffff is treated as "no columns" (guards against a + // sentinel/placeholder value rather than a real column list). + let mut buf = BytesMut::new(); + buf.put_u16_le(0xffff); + + let decoded = TokenColMetaData::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + assert!(decoded.columns.is_empty()); + } + + #[test] + fn display_formats_fixed_len_money_and_datetime_types() { + // Covers the FixedLenType Display arms not exercised elsewhere: + // tinyint/smallint/smalldatetime/real/money/datetime/smallmoney/bigint + // and the `Null` sentinel (which surfaces as `int`). + let cases = vec![ + (TypeInfo::FixedLen(FixedLenType::Int1), "c tinyint"), + (TypeInfo::FixedLen(FixedLenType::Int2), "c smallint"), + ( + TypeInfo::FixedLen(FixedLenType::Datetime4), + "c smalldatetime", + ), + (TypeInfo::FixedLen(FixedLenType::Float4), "c real"), + (TypeInfo::FixedLen(FixedLenType::Money), "c money"), + (TypeInfo::FixedLen(FixedLenType::Datetime), "c datetime"), + (TypeInfo::FixedLen(FixedLenType::Money4), "c smallmoney"), + (TypeInfo::FixedLen(FixedLenType::Int8), "c bigint"), + (TypeInfo::FixedLen(FixedLenType::Null), "c int"), + ]; + + for (ty, expected) in cases { + let expected = expected.replacen("c ", "[c] ", 1); + assert_eq!(format!("{}", meta(ty, "c")), expected); + } + } + + #[cfg(feature = "tds73")] + #[test] + fn display_formats_tds73_date_time_types() { + // date/time/datetime2/datetimeoffset Display arms (tds73-only). + let cases = vec![ + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Daten, 3, None)), + "c date", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Timen, 7, None)), + "c time", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Datetime2, 7, None)), + "c datetime2(7)", + ), + ( + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::DatetimeOffsetn, 7, None)), + "c datetimeoffset", + ), + ]; + + for (ty, expected) in cases { + let expected = expected.replacen("c ", "[c] ", 1); + assert_eq!(format!("{}", meta(ty, "c")), expected); + } + } + + #[test] + fn null_value_all_fixed_len() { + use FixedLenType::*; + let cases = [ + (Null, ColumnData::I32(None)), + (Int1, ColumnData::U8(None)), + (Bit, ColumnData::Bit(None)), + (Int2, ColumnData::I16(None)), + (Int4, ColumnData::I32(None)), + (Datetime4, ColumnData::SmallDateTime(None)), + (Float4, ColumnData::F32(None)), + (Money, ColumnData::F64(None)), + (Datetime, ColumnData::DateTime(None)), + (Float8, ColumnData::F64(None)), + (Money4, ColumnData::F32(None)), + (Int8, ColumnData::I64(None)), + ]; + + for (ty, expected) in cases { + let base = BaseMetaDataColumn { + flags: BitFlags::empty(), + ty: TypeInfo::FixedLen(ty), + table_name: None, + }; + assert_eq!(base.null_value(), expected); + } + } + + fn vsize_null(ty: VarLenType, len: usize) -> ColumnData<'static> { + BaseMetaDataColumn { + flags: BitFlags::empty(), + ty: TypeInfo::VarLenSized(VarLenContext::new(ty, len, None)), + table_name: None, + } + .null_value() + } + + #[test] + fn null_value_all_var_len_sized() { + use VarLenType::*; + assert_eq!(vsize_null(Guid, 16), ColumnData::Guid(None)); + assert_eq!(vsize_null(Bitn, 1), ColumnData::Bit(None)); + assert_eq!(vsize_null(Decimaln, 17), ColumnData::Numeric(None)); + assert_eq!(vsize_null(Numericn, 17), ColumnData::Numeric(None)); + assert_eq!(vsize_null(Money, 8), ColumnData::F64(None)); + assert_eq!(vsize_null(Datetimen, 8), ColumnData::DateTime(None)); + assert_eq!(vsize_null(BigVarBin, 100), ColumnData::Binary(None)); + assert_eq!(vsize_null(BigVarChar, 100), ColumnData::String(None)); + assert_eq!(vsize_null(BigBinary, 10), ColumnData::Binary(None)); + assert_eq!(vsize_null(BigChar, 10), ColumnData::String(None)); + assert_eq!(vsize_null(NVarchar, 100), ColumnData::String(None)); + assert_eq!(vsize_null(NChar, 10), ColumnData::String(None)); + assert_eq!(vsize_null(Xml, 0), ColumnData::Xml(None)); + assert_eq!(vsize_null(Udt, 0), ColumnData::Binary(None)); + assert_eq!(vsize_null(Text, 0), ColumnData::String(None)); + assert_eq!(vsize_null(Image, 0), ColumnData::Binary(None)); + assert_eq!(vsize_null(NText, 0), ColumnData::String(None)); + assert_eq!(vsize_null(SSVariant, 0), ColumnData::String(None)); + } + + #[cfg(feature = "tds73")] + #[test] + fn null_value_var_len_sized_tds73() { + use VarLenType::*; + assert_eq!(vsize_null(Daten, 3), ColumnData::Date(None)); + assert_eq!(vsize_null(Timen, 7), ColumnData::Time(None)); + assert_eq!(vsize_null(Datetime2, 7), ColumnData::DateTime2(None)); + assert_eq!( + vsize_null(DatetimeOffsetn, 7), + ColumnData::DateTimeOffset(None) + ); + } + + fn vprec_null(ty: VarLenType) -> ColumnData<'static> { + BaseMetaDataColumn { + flags: BitFlags::empty(), + ty: TypeInfo::VarLenSizedPrecision { + ty, + size: 8, + precision: 18, + scale: 2, + }, + table_name: None, + } + .null_value() + } + + #[test] + fn null_value_all_var_len_precision() { + use VarLenType::*; + assert_eq!(vprec_null(Guid), ColumnData::Guid(None)); + assert_eq!(vprec_null(Intn), ColumnData::I32(None)); + assert_eq!(vprec_null(Bitn), ColumnData::Bit(None)); + assert_eq!(vprec_null(Decimaln), ColumnData::Numeric(None)); + assert_eq!(vprec_null(Numericn), ColumnData::Numeric(None)); + assert_eq!(vprec_null(Floatn), ColumnData::F32(None)); + assert_eq!(vprec_null(Money), ColumnData::F64(None)); + assert_eq!(vprec_null(Datetimen), ColumnData::DateTime(None)); + assert_eq!(vprec_null(BigVarBin), ColumnData::Binary(None)); + assert_eq!(vprec_null(BigVarChar), ColumnData::String(None)); + assert_eq!(vprec_null(BigBinary), ColumnData::Binary(None)); + assert_eq!(vprec_null(BigChar), ColumnData::String(None)); + assert_eq!(vprec_null(NVarchar), ColumnData::String(None)); + assert_eq!(vprec_null(NChar), ColumnData::String(None)); + assert_eq!(vprec_null(Xml), ColumnData::Xml(None)); + assert_eq!(vprec_null(Udt), ColumnData::Binary(None)); + assert_eq!(vprec_null(Text), ColumnData::String(None)); + assert_eq!(vprec_null(Image), ColumnData::Binary(None)); + assert_eq!(vprec_null(NText), ColumnData::String(None)); + assert_eq!(vprec_null(SSVariant), ColumnData::String(None)); + } + + #[cfg(feature = "tds73")] + #[test] + fn null_value_var_len_precision_tds73() { + use VarLenType::*; + assert_eq!(vprec_null(Daten), ColumnData::Date(None)); + assert_eq!(vprec_null(Timen), ColumnData::Time(None)); + assert_eq!(vprec_null(Datetime2), ColumnData::DateTime2(None)); + assert_eq!( + vprec_null(DatetimeOffsetn), + ColumnData::DateTimeOffset(None) + ); + } + + #[test] + fn column_flag_bits_are_distinct() { + let all = ColumnFlag::Nullable + | ColumnFlag::CaseSensitive + | ColumnFlag::Updateable + | ColumnFlag::UpdateableUnknown + | ColumnFlag::Identity + | ColumnFlag::Computed + | ColumnFlag::FixedLenClrType + | ColumnFlag::SparseColumnSet + | ColumnFlag::Encrypted + | ColumnFlag::Hidden + | ColumnFlag::Key + | ColumnFlag::NullableUnknown; + + assert!(all.contains(ColumnFlag::Nullable)); + assert!(all.contains(ColumnFlag::NullableUnknown)); + assert_eq!(BitFlags::bits(all).count_ones(), 12); } } diff --git a/src/tds/codec/token/token_done.rs b/src/tds/codec/token/token_done.rs index bb45c34ac..a7a84cfd5 100644 --- a/src/tds/codec/token/token_done.rs +++ b/src/tds/codec/token/token_done.rs @@ -1,4 +1,4 @@ -use crate::{tds::codec::Encode, Error, SqlReadBytes, TokenType}; +use crate::{tds::codec::Encode, SqlReadBytes, TokenType}; use asynchronous_codec::BytesMut; use bytes::BufMut; use enumflags2::{bitflags, BitFlags}; @@ -31,8 +31,11 @@ impl TokenDone { where R: SqlReadBytes + Unpin, { - let status = BitFlags::from_bits(src.read_u16_le().await?) - .map_err(|_| Error::Protocol("done(variant): invalid status".into()))?; + // The DONE Status (MS-TDS §2.2.7.6) is a 2-byte bitmask with reserved + // bits that a server (or a future SQL Server / Azure build) may set. + // Truncate to the flags we model rather than erroring, matching how + // COLMETADATA flags are handled. + let status = BitFlags::from_bits_truncate(src.read_u16_le().await?); let cur_cmd = src.read_u16_le().await?; let done_row_count_bytes = src.context().version().done_row_count_bytes(); @@ -40,7 +43,14 @@ impl TokenDone { let done_rows = match done_row_count_bytes { 8 => src.read_u64_le().await?, 4 => src.read_u32_le().await? as u64, - _ => unreachable!(), + // `done_row_count_bytes()` only ever yields 4 or 8, so this arm is + // currently unreachable. Return a protocol error rather than panic + // so a future/out-of-range width can never crash the decoder. + other => { + return Err(crate::Error::Protocol( + format!("DONE token: unexpected row-count width {other} bytes").into(), + )) + } }; Ok(TokenDone { @@ -54,8 +64,21 @@ impl TokenDone { self.status.is_empty() } + /// `true` when the server has set the `DONE_ATTN` status bit, indicating + /// this DONE token acknowledges a client Attention signal (MS-TDS + /// section 2.2.7.6). + pub(crate) fn is_attention(&self) -> bool { + self.status.contains(DoneStatus::Attention) + } + pub(crate) fn rows(&self) -> u64 { - self.done_rows + // The row count is only meaningful when the DONE_COUNT status bit is + // set (MS-TDS §2.2.7.6); otherwise the field is not a valid count. + if self.status.contains(DoneStatus::Count) { + self.done_rows + } else { + 0 + } } } @@ -86,3 +109,130 @@ impl fmt::Display for TokenDone { } } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use bytes::BytesMut; + + #[tokio::test] + async fn decode_final_done() { + let mut buf = BytesMut::new(); + buf.put_u16_le(0); // status: empty => final + buf.put_u16_le(0); // cur_cmd + buf.put_u64_le(0); // done_rows (SqlServerN => 8 bytes) + + let done = TokenDone::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + assert!(done.is_final()); + assert_eq!(done.rows(), 0); + assert!(format!("{}", done).starts_with("Done with status")); + } + + #[tokio::test] + async fn decode_with_count_and_rows() { + let mut buf = BytesMut::new(); + buf.put_u16_le(DoneStatus::Count as u16); + buf.put_u16_le(0); + buf.put_u64_le(5); + + let done = TokenDone::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + assert!(!done.is_final()); + assert_eq!(done.rows(), 5); + assert!(format!("{}", done).contains("5 rows left")); + } + + #[tokio::test] + async fn decode_reads_four_byte_rowcount_on_pre_2005_versions() { + // Pre-2005 servers encode the DONE rowcount in 4 bytes; the decoder must + // pick the 4-byte arm from the negotiated version. + use crate::sql_read_bytes::SqlReadBytes; + use crate::tds::codec::login::FeatureLevel; + + let mut buf = BytesMut::new(); + buf.put_u16_le(DoneStatus::Count as u16); + buf.put_u16_le(0); + buf.put_u32_le(7); // 4-byte rowcount + + let mut reader = buf.into_sql_read_bytes(); + reader + .context_mut() + .set_version(FeatureLevel::SqlServer2000); + + let done = TokenDone::decode(&mut reader).await.unwrap(); + assert_eq!(done.rows(), 7); + } + + #[tokio::test] + async fn decode_single_row_display() { + let mut buf = BytesMut::new(); + buf.put_u16_le(DoneStatus::Count as u16); + buf.put_u16_le(0); + buf.put_u64_le(1); + + let done = TokenDone::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + assert!(format!("{}", done).contains("1 row left")); + } + + #[tokio::test] + async fn decode_tolerates_reserved_status_bits() { + let mut buf = BytesMut::new(); + // bit 3 (0b1000 = 8) is reserved/undefined; combined with a real bit + // (More = 0b1). We tolerate the reserved bit and keep the modeled one. + buf.put_u16_le(0b1001); + buf.put_u16_le(0); + buf.put_u64_le(0); + + let done = TokenDone::decode(&mut buf.into_sql_read_bytes()) + .await + .expect("reserved status bits must be tolerated"); + + assert!(done.status.contains(DoneStatus::More)); + } + + #[tokio::test] + async fn is_attention_reflects_attention_status_bit() { + // With the Attention bit set, is_attention() must be true. + let mut buf = BytesMut::new(); + buf.put_u16_le(DoneStatus::Attention as u16); + buf.put_u16_le(0); + buf.put_u64_le(0); + + let done = TokenDone::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + assert!(done.is_attention()); + + // Without the Attention bit (a different, non-attention bit set), + // is_attention() must be false. + let mut buf = BytesMut::new(); + buf.put_u16_le(DoneStatus::More as u16); + buf.put_u16_le(0); + buf.put_u64_le(0); + + let done = TokenDone::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + assert!(!done.is_attention()); + } + + #[test] + fn encode_writes_token_type_and_fields() { + let done = TokenDone::default(); + let mut buf = BytesMut::new(); + done.encode(&mut buf).unwrap(); + + assert_eq!(buf[0], TokenType::Done as u8); + // status(2) + cur_cmd(2) + done_rows(8) after the 1-byte token type + assert_eq!(buf.len(), 1 + 2 + 2 + 8); + } +} diff --git a/src/tds/codec/token/token_env_change.rs b/src/tds/codec/token/token_env_change.rs index ecbb9612f..8a3924a59 100644 --- a/src/tds/codec/token/token_env_change.rs +++ b/src/tds/codec/token/token_env_change.rs @@ -63,8 +63,14 @@ impl fmt::Display for EnvChangeTy { #[derive(Debug)] pub enum TokenEnvChange { - Database(String, String), - PacketSize(u32, u32), + Database { + old: String, + new: String, + }, + PacketSize { + old: u32, + new: u32, + }, SqlCollation { old: Option, new: Option, @@ -84,10 +90,10 @@ pub enum TokenEnvChange { impl fmt::Display for TokenEnvChange { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { match self { - Self::Database(ref old, ref new) => { + Self::Database { old, new } => { write!(f, "Database change from '{}' to '{}'", old, new) } - Self::PacketSize(old, new) => { + Self::PacketSize { old, new } => { write!(f, "Packet size change from '{}' to '{}'", old, new) } Self::SqlCollation { old, new } => match (old, new) { @@ -110,6 +116,21 @@ impl fmt::Display for TokenEnvChange { } } +/// Read a `B_VARCHAR`: a single `u8` UTF-16 code-unit count followed by that +/// many little-endian `u16` code units, decoded strictly as UTF-16. Kept +/// behavior-identical to the inlined copies it replaces (a `FromUtf16Error` +/// surfaces as `Error::Utf16` via `?`). +fn read_b_varchar(buf: &mut R) -> crate::Result { + let len = buf.read_u8()? as usize; + let mut units = vec![0u16; len]; + + for unit in units.iter_mut() { + *unit = buf.read_u16::()?; + } + + Ok(String::from_utf16(&units[..])?) +} + impl TokenEnvChange { pub(crate) async fn decode(src: &mut R) -> crate::Result where @@ -119,7 +140,8 @@ impl TokenEnvChange { // We read all the bytes now, due to whatever environment change tokens // we read, they might contain padding zeroes in the end we must - // discard. + // discard. `len` is bounded by the u16 length field (<= 64 KiB), so no + // named allocation cap is required here. let mut bytes = vec![0; len]; src.read_exact(&mut bytes[0..len]).await?; @@ -131,46 +153,22 @@ impl TokenEnvChange { let token = match ty { EnvChangeTy::Database => { - let len = buf.read_u8()? as usize; - let mut bytes = vec![0; len]; - - for item in bytes.iter_mut().take(len) { - *item = buf.read_u16::()?; - } - - let new_value = String::from_utf16(&bytes[..])?; - - let len = buf.read_u8()? as usize; - let mut bytes = vec![0; len]; + let new_value = read_b_varchar(&mut buf)?; + let old_value = read_b_varchar(&mut buf)?; - for item in bytes.iter_mut().take(len) { - *item = buf.read_u16::()?; + TokenEnvChange::Database { + new: new_value, + old: old_value, } - - let old_value = String::from_utf16(&bytes[..])?; - - TokenEnvChange::Database(new_value, old_value) } EnvChangeTy::PacketSize => { - let len = buf.read_u8()? as usize; - let mut bytes = vec![0; len]; + let new_value = read_b_varchar(&mut buf)?; + let old_value = read_b_varchar(&mut buf)?; - for item in bytes.iter_mut().take(len) { - *item = buf.read_u16::()?; + TokenEnvChange::PacketSize { + new: new_value.parse()?, + old: old_value.parse()?, } - - let new_value = String::from_utf16(&bytes[..])?; - - let len = buf.read_u8()? as usize; - let mut bytes = vec![0; len]; - - for item in bytes.iter_mut().take(len) { - *item = buf.read_u16::()?; - } - - let old_value = String::from_utf16(&bytes[..])?; - - TokenEnvChange::PacketSize(new_value.parse()?, old_value.parse()?) } EnvChangeTy::SqlCollation => { let len = buf.read_u8()? as usize; @@ -213,7 +211,11 @@ impl TokenEnvChange { } EnvChangeTy::BeginTransaction | EnvChangeTy::EnlistDTCTransaction => { let len = buf.read_u8()?; - assert!(len == 8); + if len != 8 { + return Err(Error::Protocol( + format!("ENVCHANGE transaction descriptor length {len}, expected 8").into(), + )); + } let mut desc = [0; 8]; buf.read_exact(&mut desc)?; @@ -243,15 +245,7 @@ impl TokenEnvChange { TokenEnvChange::Routing { host, port } } EnvChangeTy::Rtls => { - let len = buf.read_u8()? as usize; - let mut bytes = vec![0; len]; - - for item in bytes.iter_mut().take(len) { - *item = buf.read_u16::()?; - } - - let mirror_name = String::from_utf16(&bytes[..])?; - + let mirror_name = read_b_varchar(&mut buf)?; TokenEnvChange::ChangeMirror(mirror_name) } ty => TokenEnvChange::Ignored(ty), @@ -260,3 +254,359 @@ impl TokenEnvChange { Ok(token) } } + +#[cfg(test)] +mod tests { + use super::{EnvChangeTy, TokenEnvChange}; + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use byteorder::{LittleEndian, WriteBytesExt}; + use bytes::{BufMut, BytesMut}; + + #[test] + fn database_display_uses_old_then_new() { + // Fields are stored (new, old); Display must print "from old to new". + let change = TokenEnvChange::Database { + new: "newdb".to_string(), + old: "olddb".to_string(), + }; + assert_eq!( + format!("{}", change), + "Database change from 'olddb' to 'newdb'" + ); + } + + #[test] + fn packet_size_display_uses_old_then_new() { + // Fields are stored (new, old); Display must print "from old to new". + let change = TokenEnvChange::PacketSize { + new: 8192, + old: 4096, + }; + assert_eq!( + format!("{}", change), + "Packet size change from '4096' to '8192'" + ); + } + + fn write_utf16_str(body: &mut Vec, s: &str) { + body.push(s.encode_utf16().count() as u8); + for unit in s.encode_utf16() { + body.write_u16::(unit).unwrap(); + } + } + + fn envchange_buf(ty: u8, payload: &[u8]) -> BytesMut { + let mut body = Vec::new(); + body.push(ty); + body.extend_from_slice(payload); + + let mut buf = BytesMut::new(); + buf.put_u16_le(body.len() as u16); + buf.put_slice(&body); + buf + } + + #[tokio::test] + async fn decode_database_roundtrip() { + let mut payload = Vec::new(); + write_utf16_str(&mut payload, "newdb"); + write_utf16_str(&mut payload, "olddb"); + + let buf = envchange_buf(EnvChangeTy::Database as u8, &payload); + let decoded = TokenEnvChange::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + match decoded { + TokenEnvChange::Database { new, old } => { + assert_eq!(new, "newdb"); + assert_eq!(old, "olddb"); + } + other => panic!("unexpected variant: {:?}", other), + } + } + + #[tokio::test] + async fn decode_packet_size_parses_numbers() { + let mut payload = Vec::new(); + write_utf16_str(&mut payload, "8192"); + write_utf16_str(&mut payload, "4096"); + + let buf = envchange_buf(EnvChangeTy::PacketSize as u8, &payload); + let decoded = TokenEnvChange::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + match decoded { + TokenEnvChange::PacketSize { new, old } => { + assert_eq!(new, 8192); + assert_eq!(old, 4096); + } + other => panic!("unexpected variant: {:?}", other), + } + } + + #[tokio::test] + async fn decode_sql_collation_with_both_present() { + let mut payload = Vec::new(); + payload.push(5u8); + payload.write_u32::(13632521).unwrap(); + payload.push(52); + payload.push(5u8); + payload.write_u32::(13632521).unwrap(); + payload.push(52); + + let buf = envchange_buf(EnvChangeTy::SqlCollation as u8, &payload); + let decoded = TokenEnvChange::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + match decoded { + TokenEnvChange::SqlCollation { old, new } => { + assert!(old.is_some()); + assert!(new.is_some()); + } + other => panic!("unexpected variant: {:?}", other), + } + } + + #[tokio::test] + async fn decode_sql_collation_none_when_length_not_five() { + let payload = vec![0u8, 0u8]; + + let buf = envchange_buf(EnvChangeTy::SqlCollation as u8, &payload); + let decoded = TokenEnvChange::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + match decoded { + TokenEnvChange::SqlCollation { old, new } => { + assert!(old.is_none()); + assert!(new.is_none()); + } + other => panic!("unexpected variant: {:?}", other), + } + } + + #[tokio::test] + async fn decode_begin_transaction_reads_descriptor() { + let mut payload = vec![8u8]; + payload.extend_from_slice(&[1, 2, 3, 4, 5, 6, 7, 8]); + + let buf = envchange_buf(EnvChangeTy::BeginTransaction as u8, &payload); + let decoded = TokenEnvChange::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + match decoded { + TokenEnvChange::BeginTransaction(desc) => { + assert_eq!(desc, [1, 2, 3, 4, 5, 6, 7, 8]); + } + other => panic!("unexpected variant: {:?}", other), + } + } + + #[tokio::test] + async fn decode_begin_transaction_wrong_length_errors() { + let payload = vec![3u8, 1, 2, 3]; + + let buf = envchange_buf(EnvChangeTy::BeginTransaction as u8, &payload); + let err = TokenEnvChange::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap_err(); + + assert!(format!("{}", err).contains("expected 8")); + } + + #[tokio::test] + async fn decode_commit_rollback_defect_transaction() { + for (ty, is_match) in [ + ( + EnvChangeTy::CommitTransaction, + (|t: &TokenEnvChange| matches!(t, TokenEnvChange::CommitTransaction)) + as fn(&TokenEnvChange) -> bool, + ), + (EnvChangeTy::RollbackTransaction, |t: &TokenEnvChange| { + matches!(t, TokenEnvChange::RollbackTransaction) + }), + (EnvChangeTy::DefectTransaction, |t: &TokenEnvChange| { + matches!(t, TokenEnvChange::DefectTransaction) + }), + ] { + let buf = envchange_buf(ty as u8, &[]); + let decoded = TokenEnvChange::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + assert!(is_match(&decoded)); + } + } + + #[tokio::test] + async fn decode_routing_reads_host_and_port() { + let mut payload = Vec::new(); + payload.write_u16::(0).unwrap(); // routing data value length (unused) + payload.push(0); // protocol, always 0 + payload.write_u16::(1433).unwrap(); // port + + let host = "sql.example.com"; + payload + .write_u16::(host.encode_utf16().count() as u16) + .unwrap(); + for unit in host.encode_utf16() { + payload.write_u16::(unit).unwrap(); + } + + let buf = envchange_buf(EnvChangeTy::Routing as u8, &payload); + let decoded = TokenEnvChange::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + match decoded { + TokenEnvChange::Routing { host: h, port } => { + assert_eq!(h, host); + assert_eq!(port, 1433); + } + other => panic!("unexpected variant: {:?}", other), + } + } + + #[tokio::test] + async fn decode_rtls_yields_change_mirror() { + let mut payload = Vec::new(); + write_utf16_str(&mut payload, "mirror.example.com"); + + let buf = envchange_buf(EnvChangeTy::Rtls as u8, &payload); + let decoded = TokenEnvChange::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + match decoded { + TokenEnvChange::ChangeMirror(name) => { + assert_eq!(name, "mirror.example.com"); + } + other => panic!("unexpected variant: {:?}", other), + } + } + + #[tokio::test] + async fn decode_unhandled_type_is_ignored() { + let buf = envchange_buf(EnvChangeTy::Language as u8, &[]); + let decoded = TokenEnvChange::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + match decoded { + TokenEnvChange::Ignored(EnvChangeTy::Language) => {} + other => panic!("unexpected variant: {:?}", other), + } + } + + #[tokio::test] + async fn decode_invalid_type_byte_errors() { + let buf = envchange_buf(0x63, &[]); + let err = TokenEnvChange::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap_err(); + + assert!(format!("{}", err).contains("invalid envchange type")); + } + + #[test] + fn env_change_ty_display_all_variants() { + let cases: &[(EnvChangeTy, &str)] = &[ + (EnvChangeTy::Database, "Database"), + (EnvChangeTy::Language, "Language"), + (EnvChangeTy::CharacterSet, "CharacterSet"), + (EnvChangeTy::PacketSize, "PacketSize"), + (EnvChangeTy::UnicodeDataSortingLID, "UnicodeDataSortingLID"), + (EnvChangeTy::UnicodeDataSortingCFL, "UnicodeDataSortingCFL"), + (EnvChangeTy::SqlCollation, "SqlCollation"), + (EnvChangeTy::BeginTransaction, "BeginTransaction"), + (EnvChangeTy::CommitTransaction, "CommitTransaction"), + (EnvChangeTy::RollbackTransaction, "RollbackTransaction"), + (EnvChangeTy::EnlistDTCTransaction, "EnlistDTCTransaction"), + (EnvChangeTy::DefectTransaction, "DefectTransaction"), + (EnvChangeTy::Rtls, "RTLS"), + (EnvChangeTy::PromoteTransaction, "PromoteTransaction"), + ( + EnvChangeTy::TransactionManagerAddress, + "TransactionManagerAddress", + ), + (EnvChangeTy::TransactionEnded, "TransactionEnded"), + (EnvChangeTy::ResetConnection, "ResetConnection"), + (EnvChangeTy::UserName, "UserName"), + (EnvChangeTy::Routing, "Routing"), + ]; + + for (variant, expected) in cases { + assert_eq!(format!("{}", variant), *expected); + } + } + + #[test] + fn sql_collation_display_both_and_new_only() { + use crate::tds::Collation; + + // Both old and new present: "from {old} to {new}". + let both = TokenEnvChange::SqlCollation { + old: Some(Collation::new(13632521, 52)), + new: Some(Collation::new(13632521, 52)), + }; + assert!(format!("{}", both).starts_with("SQL collation change from ")); + + // Only new present: "changed to {new}". + let new_only = TokenEnvChange::SqlCollation { + old: None, + new: Some(Collation::new(13632521, 52)), + }; + assert!(format!("{}", new_only).starts_with("SQL collation changed to ")); + } + + #[test] + fn token_env_change_display_variants() { + assert_eq!( + format!("{}", TokenEnvChange::CommitTransaction), + "Commit transaction" + ); + assert_eq!( + format!("{}", TokenEnvChange::RollbackTransaction), + "Rollback transaction" + ); + assert_eq!( + format!("{}", TokenEnvChange::DefectTransaction), + "Defect transaction" + ); + assert_eq!( + format!("{}", TokenEnvChange::BeginTransaction([0; 8])), + "Begin transaction" + ); + assert_eq!( + format!( + "{}", + TokenEnvChange::Routing { + host: "host".into(), + port: 1433 + } + ), + "Server requested routing to a new address: host:1433" + ); + assert_eq!( + format!("{}", TokenEnvChange::ChangeMirror("mirror".into())), + "Fallback mirror server: `mirror`" + ); + assert_eq!( + format!("{}", TokenEnvChange::Ignored(EnvChangeTy::Language)), + "Ignored env change: `Language`" + ); + assert_eq!( + format!( + "{}", + TokenEnvChange::SqlCollation { + old: None, + new: None + } + ), + "SQL collation change" + ); + } +} diff --git a/src/tds/codec/token/token_error.rs b/src/tds/codec/token/token_error.rs index d1e435a77..37cf4fc45 100644 --- a/src/tds/codec/token/token_error.rs +++ b/src/tds/codec/token/token_error.rs @@ -1,5 +1,8 @@ -use crate::{tds::codec::FeatureLevel, SqlReadBytes}; +use crate::{tds::codec::FeatureLevel, Error, SqlReadBytes}; +use byteorder::{LittleEndian, ReadBytesExt}; +use futures_util::io::AsyncReadExt; use std::fmt; +use std::io::Cursor; #[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)] /// An error token returned from the server. @@ -17,38 +20,99 @@ pub struct TokenError { pub(crate) line: u32, } +/// The body shared by the ERROR (`0xAA`) and INFO (`0xAB`) tokens, which have +/// byte-for-byte identical layouts (MS-TDS §2.2.7.10 / §2.2.7.13). +pub(crate) struct ErrorInfoBody { + pub(crate) code: u32, + pub(crate) state: u8, + pub(crate) class: u8, + pub(crate) message: String, + pub(crate) server: String, + pub(crate) procedure: String, + pub(crate) line: u32, +} + +fn read_utf16(cur: &mut Cursor<&[u8]>, char_count: usize) -> crate::Result { + let mut units = Vec::with_capacity(char_count); + for _ in 0..char_count { + units.push(cur.read_u16::()?); + } + String::from_utf16(&units) + .map_err(|_| Error::Protocol("ERROR/INFO token string is not valid UTF-16".into())) +} + +fn read_us_varchar(cur: &mut Cursor<&[u8]>) -> crate::Result { + let char_count = cur.read_u16::()? as usize; + read_utf16(cur, char_count) +} + +fn read_b_varchar(cur: &mut Cursor<&[u8]>) -> crate::Result { + let char_count = cur.read_u8()? as usize; + read_utf16(cur, char_count) +} + +/// Decode the common ERROR/INFO token body. +/// +/// The token's declared `Length` is honored exactly: the whole body is read +/// into a bounded buffer and parsed from there, so a `Length`/content mismatch +/// surfaces as a clean protocol error instead of over- or under-reading the +/// wire and desyncing the token stream. Any trailing bytes within `Length` that +/// the fields do not consume are discarded. +pub(crate) async fn decode_error_info_body(src: &mut R) -> crate::Result +where + R: SqlReadBytes + Unpin, +{ + // MS-TDS §2.2.7.10/§2.2.7.13: LineNumber is a 4-byte LONG for TDS 7.2 (SQL + // Server 2005) and later, and a 2-byte USHORT before that. The boundary is + // inclusive of 7.2, so use `>=`. Capture it before consuming the body. + let four_byte_line = src.context().version() >= FeatureLevel::SqlServer2005; + + let length = src.read_u16_le().await? as usize; + + let mut data = vec![0u8; length]; + src.read_exact(&mut data).await?; + + let mut cur = Cursor::new(&data[..]); + + let code = cur.read_u32::()?; + let state = cur.read_u8()?; + let class = cur.read_u8()?; + let message = read_us_varchar(&mut cur)?; + let server = read_b_varchar(&mut cur)?; + let procedure = read_b_varchar(&mut cur)?; + let line = if four_byte_line { + cur.read_u32::()? + } else { + cur.read_u16::()? as u32 + }; + + Ok(ErrorInfoBody { + code, + state, + class, + message, + server, + procedure, + line, + }) +} + impl TokenError { pub(crate) async fn decode(src: &mut R) -> crate::Result where R: SqlReadBytes + Unpin, { - let _length = src.read_u16_le().await? as usize; - - let code = src.read_u32_le().await?; - let state = src.read_u8().await?; - let class = src.read_u8().await?; + let body = decode_error_info_body(src).await?; - let message = src.read_us_varchar().await?; - let server = src.read_b_varchar().await?; - let procedure = src.read_b_varchar().await?; - - let line = if src.context().version() > FeatureLevel::SqlServer2005 { - src.read_u32_le().await? - } else { - src.read_u16_le().await? as u32 - }; - - let token = TokenError { - code, - state, - class, - message, - server, - procedure, - line, - }; - - Ok(token) + Ok(TokenError { + code: body.code, + state: body.state, + class: body.class, + message: body.message, + server: body.server, + procedure: body.procedure, + line: body.line, + }) } /// The error code, see descriptions from [the manual]. @@ -101,3 +165,183 @@ impl fmt::Display for TokenError { ) } } + +#[cfg(test)] +mod tests { + use super::*; + + fn sample() -> TokenError { + TokenError { + code: 1205, + state: 2, + class: 13, + message: "deadlocked".to_string(), + server: "myserver".to_string(), + procedure: "myproc".to_string(), + line: 42, + } + } + + #[test] + fn accessors() { + let e = sample(); + assert_eq!(e.code(), 1205); + assert_eq!(e.state(), 2); + assert_eq!(e.class(), 13); + assert_eq!(e.message(), "deadlocked"); + assert_eq!(e.server(), "myserver"); + assert_eq!(e.procedure(), "myproc"); + assert_eq!(e.line(), 42); + } + + #[test] + fn display_contains_all_fields() { + let rendered = format!("{}", sample()); + assert_eq!( + rendered, + "'deadlocked' on server myserver executing myproc on line 42 (code: 1205, state: 2, class: 13)" + ); + } + + #[tokio::test] + async fn decode_reads_all_fields_with_four_byte_line_number() { + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use byteorder::{LittleEndian, WriteBytesExt}; + use bytes::{BufMut, BytesMut}; + + fn write_us_varchar(buf: &mut Vec, s: &str) { + buf.write_u16::(s.encode_utf16().count() as u16) + .unwrap(); + for u in s.encode_utf16() { + buf.write_u16::(u).unwrap(); + } + } + + fn write_b_varchar(buf: &mut Vec, s: &str) { + buf.push(s.encode_utf16().count() as u8); + for u in s.encode_utf16() { + buf.write_u16::(u).unwrap(); + } + } + + let mut body = Vec::new(); + body.write_u32::(1205).unwrap(); // code + body.push(2); // state + body.push(13); // class + write_us_varchar(&mut body, "deadlocked"); + write_b_varchar(&mut body, "myserver"); + write_b_varchar(&mut body, "myproc"); + body.write_u32::(42).unwrap(); // line, TDS >= 7.2 (default context) + + let mut buf = BytesMut::new(); + buf.put_u16_le(body.len() as u16); // length prefix, ignored by decode + buf.put_slice(&body); + + let decoded = TokenError::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + assert_eq!(decoded, sample()); + } + + #[tokio::test] + async fn decode_reads_full_four_byte_line_number_on_tds72_plus() { + // The default test context reports SqlServerN (>= TDS 7.2), so the + // LineNumber must be read as a 4-byte LONG. 0x0001_0001 (65537) has + // distinct low-16-bit and full-32-bit values, so a 2-byte read yields 1 + // while the correct 4-byte read yields 65537. + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use bytes::{BufMut, BytesMut}; + + let mut body = BytesMut::new(); + body.put_u32_le(1205); // code + body.put_u8(2); // state + body.put_u8(13); // class + body.put_u16_le(0); // message: us_varchar, length 0 + body.put_u8(0); // server: b_varchar, length 0 + body.put_u8(0); // procedure: b_varchar, length 0 + body.put_u32_le(0x0001_0001); // line number, 4 bytes + + let mut buf = BytesMut::new(); + buf.put_u16_le(body.len() as u16); // length prefix, ignored by decode + buf.put_slice(&body); + + let decoded = TokenError::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + assert_eq!(decoded.line(), 0x0001_0001); + } + + #[tokio::test] + async fn decode_consumes_exactly_length_and_does_not_desync() { + // Declared Length (6) covers only code+state+class; the message length + // prefix falls outside it. Decoding must fail with a clean error AND + // consume exactly `Length` body bytes, so the following byte on the wire + // (a sentinel here) is still readable — i.e. the stream is not desynced. + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use crate::SqlReadBytes; + use bytes::{BufMut, BytesMut}; + + // The trailing bytes below belong to the *next* token on the wire. A + // decoder that ignores `Length` (the old behaviour) would greedily read + // them as this token's message/server/procedure/line and succeed, + // desyncing the stream. They form a valid trailing field sequence so the + // old code returns Ok, making `expect_err` a genuine red. + let mut buf = BytesMut::new(); + buf.put_u16_le(6); // declared Length (too small for the full body) + buf.put_u32_le(0x1122_3344); // code + buf.put_u8(0x55); // state + buf.put_u8(0x66); // class + buf.put_u16_le(0); // next token: (mis)read as message length 0 + buf.put_u8(0); // next token: (mis)read as server length 0 + buf.put_u8(0); // next token: (mis)read as procedure length 0 + buf.put_u32_le(0x2A); // next token: (mis)read as line + buf.put_u8(0xEF); // sentinel further along + + let mut reader = buf.into_sql_read_bytes(); + + let err = TokenError::decode(&mut reader) + .await + .expect_err("under-declared length must be a clean error"); + assert!(matches!(err, Error::Protocol(_) | Error::Io { .. })); + + // Exactly the 6 declared body bytes were consumed, so the first byte of + // the next token region is still on the wire (no desync). + assert_eq!(reader.read_u8().await.unwrap(), 0x00); + } + + #[tokio::test] + async fn decode_ignores_trailing_bytes_within_declared_length() { + // Declared Length is larger than the fields consume; the extra bytes are + // discarded and the next token (sentinel) is read cleanly afterwards. + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use crate::SqlReadBytes; + use bytes::{BufMut, BytesMut}; + + let mut body = BytesMut::new(); + body.put_u32_le(1205); // code + body.put_u8(2); // state + body.put_u8(13); // class + body.put_u16_le(0); // message (empty) + body.put_u8(0); // server (empty) + body.put_u8(0); // procedure (empty) + body.put_u32_le(42); // line + body.put_u8(0x99); // trailing padding within the declared length + + let mut buf = BytesMut::new(); + buf.put_u16_le(body.len() as u16); + buf.put_slice(&body); + buf.put_u8(0xCD); // sentinel next token + + let mut reader = buf.into_sql_read_bytes(); + + let decoded = TokenError::decode(&mut reader) + .await + .expect("trailing padding within length must decode"); + assert_eq!(decoded.code(), 1205); + assert_eq!(decoded.line(), 42); + + assert_eq!(reader.read_u8().await.unwrap(), 0xCD); + } +} diff --git a/src/tds/codec/token/token_feature_ext_ack.rs b/src/tds/codec/token/token_feature_ext_ack.rs index 1ba108f99..9d7fe7457 100644 --- a/src/tds/codec/token/token_feature_ext_ack.rs +++ b/src/tds/codec/token/token_feature_ext_ack.rs @@ -1,4 +1,4 @@ -use crate::{SqlReadBytes, FEA_EXT_FEDAUTH, FEA_EXT_TERMINATOR}; +use crate::{Error, SqlReadBytes, FEA_EXT_FEDAUTH, FEA_EXT_TERMINATOR}; use futures_util::AsyncReadExt; #[derive(Debug)] @@ -40,15 +40,106 @@ impl TokenFeatureExtAck { } else if data_len == 0 { None } else { - panic!("invalid Feature_Ext_Ack token"); + return Err(Error::Protocol( + format!( + "invalid Feature_Ext_Ack token: invalid data length {}", + data_len + ) + .into(), + )); }; features.push(FeatureAck::FedAuth(FedAuthAck::SecurityToken { nonce })) } else { - unimplemented!("unsupported feature {}", feature_id) + return Err(Error::Protocol( + format!("unsupported feature {}", feature_id).into(), + )); } } Ok(TokenFeatureExtAck { features }) } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use bytes::{BufMut, BytesMut}; + + #[tokio::test] + async fn decodes_fedauth_with_nonce() { + let mut buf = BytesMut::new(); + buf.put_u8(FEA_EXT_FEDAUTH); + buf.put_u32_le(32); + buf.extend_from_slice(&[7u8; 32]); + buf.put_u8(FEA_EXT_TERMINATOR); + + let ack = TokenFeatureExtAck::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + assert_eq!(ack.features.len(), 1); + match &ack.features[0] { + FeatureAck::FedAuth(FedAuthAck::SecurityToken { nonce }) => { + assert_eq!(*nonce, Some([7u8; 32])); + } + } + } + + #[tokio::test] + async fn decodes_fedauth_without_nonce() { + let mut buf = BytesMut::new(); + buf.put_u8(FEA_EXT_FEDAUTH); + buf.put_u32_le(0); + buf.put_u8(FEA_EXT_TERMINATOR); + + let ack = TokenFeatureExtAck::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + match &ack.features[0] { + FeatureAck::FedAuth(FedAuthAck::SecurityToken { nonce }) => { + assert!(nonce.is_none()); + } + } + } + + #[tokio::test] + async fn decode_rejects_invalid_data_length() { + // A FEDAUTH ack with a data length that is neither 0 nor 32 is invalid. + let mut buf = BytesMut::new(); + buf.put_u8(FEA_EXT_FEDAUTH); + buf.put_u32_le(5); + buf.extend_from_slice(&[0u8; 5]); + + let err = TokenFeatureExtAck::decode(&mut buf.into_sql_read_bytes()) + .await + .expect_err("invalid data length must error"); + assert!(matches!(err, Error::Protocol(_))); + } + + #[tokio::test] + async fn decode_rejects_unsupported_feature() { + // A feature id that is neither the terminator nor FEDAUTH is unsupported. + let mut buf = BytesMut::new(); + buf.put_u8(0x99); + + let err = TokenFeatureExtAck::decode(&mut buf.into_sql_read_bytes()) + .await + .expect_err("unsupported feature must error"); + assert!(matches!(err, Error::Protocol(_))); + } + + #[tokio::test] + async fn empty_feature_list() { + let mut buf = BytesMut::new(); + buf.put_u8(FEA_EXT_TERMINATOR); + + let ack = TokenFeatureExtAck::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + assert!(ack.features.is_empty()); + } +} diff --git a/src/tds/codec/token/token_fed_auth_info.rs b/src/tds/codec/token/token_fed_auth_info.rs new file mode 100644 index 000000000..90605d055 --- /dev/null +++ b/src/tds/codec/token/token_fed_auth_info.rs @@ -0,0 +1,361 @@ +use crate::{Error, SqlReadBytes}; +use futures_util::io::AsyncReadExt; + +/// `FedAuthInfoId` for the STS URL, an Active Directory Security Token Service +/// endpoint that the client contacts to acquire an access token. +const FED_AUTH_INFO_ID_STSURL: u8 = 0x01; + +/// `FedAuthInfoId` for the Service Principal Name the token is requested for. +const FED_AUTH_INFO_ID_SPN: u8 = 0x02; + +/// A `FEDAUTHINFO` token (`0xEE`), returned by the server during the federated +/// authentication handshake to describe how the client should acquire a +/// federated access token. +/// +/// The server sends this token when the client requested a library-driven +/// federated authentication flow (for example, the ADAL/MSAL interactive or +/// integrated Azure Active Directory flows) in the `FEDAUTH` `FeatureExt` +/// option of the login request. The token carries the information elements the +/// client needs to talk to the Active Directory Security Token Service (STS): +/// the STS URL to authenticate against and the Service Principal Name (SPN) the +/// token is requested for. +/// +/// See [MS-TDS] §2.2.7.12 (`FEDAUTHINFO`). +#[derive(Debug, Clone, PartialEq, Eq, Default)] +pub struct TokenFedAuthInfo { + /// The URL of the Active Directory Security Token Service the client should + /// authenticate against (`FedAuthInfoId` `STSURL`, `0x01`), if the server + /// provided one. + pub sts_url: Option, + /// The Service Principal Name the federated access token is requested for + /// (`FedAuthInfoId` `SPN`, `0x02`), if the server provided one. + pub spn: Option, +} + +impl TokenFedAuthInfo { + pub(crate) async fn decode(src: &mut R) -> crate::Result + where + R: SqlReadBytes + Unpin, + { + // TokenLength: the length, in bytes, of the token value that follows, + // starting at (and including) `CountOfInfoIDs`. + let token_length = src.read_u32_le().await? as usize; + + if token_length > super::MAX_TOKEN_BODY { + return Err(Error::Protocol( + format!("FEDAUTHINFO token length {token_length} exceeds the maximum").into(), + )); + } + + let mut body = vec![0u8; token_length]; + src.read_exact(&mut body).await?; + + Self::parse(&body) + } + + /// Parses the body of a `FEDAUTHINFO` token: the bytes that follow the + /// `TokenType` and `TokenLength` fields, starting at `CountOfInfoIDs`. + /// + /// The information-data offsets carried by each option are measured from the + /// start of this body (the `CountOfInfoIDs` field), matching the on-the-wire + /// layout described in [MS-TDS] §2.2.7.12. + fn parse(body: &[u8]) -> crate::Result { + let read_u32 = |buf: &[u8], at: usize| -> crate::Result { + buf.get(at..at + 4) + .map(|b| u32::from_le_bytes([b[0], b[1], b[2], b[3]])) + .ok_or_else(|| { + Error::Protocol("FEDAUTHINFO token truncated while reading a DWORD".into()) + }) + }; + + // CountOfInfoIDs: the number of `FedAuthInfoOpt` options that follow. + let count = read_u32(body, 0)? as usize; + + let mut info = TokenFedAuthInfo::default(); + + for i in 0..count { + // Each `FedAuthInfoOpt` is 9 bytes: a 1-byte id followed by two DWORDs. + let opt = 4 + i * 9; + + let id = *body.get(opt).ok_or_else(|| { + Error::Protocol("FEDAUTHINFO token truncated while reading an option id".into()) + })?; + + let data_len = read_u32(body, opt + 1)? as usize; + let data_offset = read_u32(body, opt + 5)? as usize; + + let data = body + .get(data_offset..data_offset + data_len) + .ok_or_else(|| { + Error::Protocol("FEDAUTHINFO token data offset out of bounds".into()) + })?; + + // The info data is a Unicode (UCS-2/UTF-16LE) string, so it must be + // an even number of bytes. + if data_len & 1 != 0 { + return Err(Error::Protocol( + "FEDAUTHINFO token data is not valid UTF-16".into(), + )); + } + + let mut utf16 = Vec::with_capacity(data_len / 2); + let mut idx = 0; + while idx < data_len { + utf16.push(u16::from_le_bytes([data[idx], data[idx + 1]])); + idx += 2; + } + + let value = String::from_utf16(&utf16).map_err(|_| { + Error::Protocol("FEDAUTHINFO token data is not valid UTF-16".into()) + })?; + + match id { + FED_AUTH_INFO_ID_STSURL => info.sts_url = Some(value), + FED_AUTH_INFO_ID_SPN => info.spn = Some(value), + // Unknown info ids are ignored for forward compatibility, as + // required by [MS-TDS] §2.2.7.12. + _ => (), + } + } + + Ok(info) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn utf16le(s: &str) -> Vec { + s.encode_utf16().flat_map(|u| u.to_le_bytes()).collect() + } + + #[test] + fn parses_stsurl_and_spn() { + let sts = utf16le("https://login.microsoftonline.com/"); + let spn = utf16le("https://database.windows.net/"); + + // Body layout: CountOfInfoIDs, then two FedAuthInfoOpt, then the data. + let count: u32 = 2; + let header_len = 4 + 2 * 9; // count + two options + let sts_offset = header_len; + let spn_offset = header_len + sts.len(); + + let mut body = Vec::new(); + body.extend_from_slice(&count.to_le_bytes()); + + // Option 1: STSURL + body.push(FED_AUTH_INFO_ID_STSURL); + body.extend_from_slice(&(sts.len() as u32).to_le_bytes()); + body.extend_from_slice(&(sts_offset as u32).to_le_bytes()); + + // Option 2: SPN + body.push(FED_AUTH_INFO_ID_SPN); + body.extend_from_slice(&(spn.len() as u32).to_le_bytes()); + body.extend_from_slice(&(spn_offset as u32).to_le_bytes()); + + body.extend_from_slice(&sts); + body.extend_from_slice(&spn); + + let info = TokenFedAuthInfo::parse(&body).unwrap(); + + assert_eq!( + info.sts_url.as_deref(), + Some("https://login.microsoftonline.com/") + ); + assert_eq!(info.spn.as_deref(), Some("https://database.windows.net/")); + } + + #[test] + fn ignores_unknown_info_id() { + let count: u32 = 1; + let mut body = Vec::new(); + body.extend_from_slice(&count.to_le_bytes()); + body.push(0x7F); // unknown id + body.extend_from_slice(&0u32.to_le_bytes()); // data len + body.extend_from_slice(&13u32.to_le_bytes()); // offset (past header, no data) + + let info = TokenFedAuthInfo::parse(&body).unwrap(); + assert_eq!(info, TokenFedAuthInfo::default()); + } + + #[test] + fn rejects_out_of_bounds_offset() { + let count: u32 = 1; + let mut body = Vec::new(); + body.extend_from_slice(&count.to_le_bytes()); + body.push(FED_AUTH_INFO_ID_STSURL); + body.extend_from_slice(&8u32.to_le_bytes()); // data len + body.extend_from_slice(&1000u32.to_le_bytes()); // bogus offset + + assert!(TokenFedAuthInfo::parse(&body).is_err()); + } + + #[tokio::test] + async fn decode_reads_length_prefix_and_parses_body() { + // Exercises the full `decode` path: reading the 4-byte TokenLength, the + // length bound check, reading the body, and parsing it into the STSURL. + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use bytes::{BufMut, BytesMut}; + + let sts = utf16le("https://sts.example/"); + + let count: u32 = 1; + let header_len = 4 + 9; // count + one FedAuthInfoOpt + let sts_offset = header_len; + + let mut body = Vec::new(); + body.extend_from_slice(&count.to_le_bytes()); + body.push(FED_AUTH_INFO_ID_STSURL); + body.extend_from_slice(&(sts.len() as u32).to_le_bytes()); + body.extend_from_slice(&(sts_offset as u32).to_le_bytes()); + body.extend_from_slice(&sts); + + let mut buf = BytesMut::new(); + buf.put_u32_le(body.len() as u32); // TokenLength + buf.put_slice(&body); + + let info = TokenFedAuthInfo::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + assert_eq!(info.sts_url.as_deref(), Some("https://sts.example/")); + assert_eq!(info.spn, None); + } + + #[tokio::test] + async fn decode_rejects_oversized_token_length() { + // A TokenLength above MAX_TOKEN_BODY must be rejected before any body is + // read (the `token_length > MAX_TOKEN_BODY` guard). + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use bytes::{BufMut, BytesMut}; + + let mut buf = BytesMut::new(); + buf.put_u32_le((super::super::MAX_TOKEN_BODY + 1) as u32); + + let err = TokenFedAuthInfo::decode(&mut buf.into_sql_read_bytes()) + .await + .expect_err("oversized token length must error"); + assert!(matches!(err, Error::Protocol(_))); + } + + #[test] + fn parse_rejects_truncated_dword() { + // A body too short to even read CountOfInfoIDs (a DWORD) trips the + // `read_u32` truncation guard. + let err = TokenFedAuthInfo::parse(&[0u8, 0u8]).expect_err("truncated body must error"); + assert!(matches!(err, Error::Protocol(_))); + } + + #[test] + fn parse_rejects_missing_option_id() { + // CountOfInfoIDs claims one option, but the body ends right after the + // count, so reading the option id is out of bounds. + let mut body = Vec::new(); + body.extend_from_slice(&1u32.to_le_bytes()); + + let err = TokenFedAuthInfo::parse(&body).expect_err("missing option id must error"); + assert!(matches!(err, Error::Protocol(_))); + } + + #[test] + fn parse_rejects_odd_data_length() { + // A data length that is not a multiple of two cannot be valid UTF-16. + let count: u32 = 1; + let header_len = 4 + 9; + let mut body = Vec::new(); + body.extend_from_slice(&count.to_le_bytes()); + body.push(FED_AUTH_INFO_ID_STSURL); + body.extend_from_slice(&3u32.to_le_bytes()); // odd data len + body.extend_from_slice(&(header_len as u32).to_le_bytes()); // offset + body.extend_from_slice(&[0u8, 0u8, 0u8]); // 3 data bytes + + let err = TokenFedAuthInfo::parse(&body).expect_err("odd data length must error"); + assert!(matches!(err, Error::Protocol(_))); + } + + #[test] + fn parse_rejects_invalid_utf16() { + // Even-length but not valid UTF-16 (a lone high surrogate) must error. + let count: u32 = 1; + let header_len = 4 + 9; + let mut body = Vec::new(); + body.extend_from_slice(&count.to_le_bytes()); + body.push(FED_AUTH_INFO_ID_STSURL); + body.extend_from_slice(&2u32.to_le_bytes()); // data len + body.extend_from_slice(&(header_len as u32).to_le_bytes()); // offset + body.extend_from_slice(&0xD800u16.to_le_bytes()); // lone high surrogate + + let err = TokenFedAuthInfo::parse(&body).expect_err("invalid UTF-16 must error"); + assert!(matches!(err, Error::Protocol(_))); + } + + #[tokio::test] + async fn decode_accepts_token_length_at_maximum() { + // A token whose length is exactly MAX_TOKEN_BODY must be accepted. The + // body is a valid, empty (CountOfInfoIDs == 0) token padded to the maximum. + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use bytes::{BufMut, BytesMut}; + + let token_length = super::super::MAX_TOKEN_BODY; + + let mut buf = BytesMut::new(); + buf.put_u32_le(token_length as u32); // TokenLength == MAX_TOKEN_BODY + buf.put_slice(&vec![0u8; token_length]); // count = 0, rest padding + + let info = TokenFedAuthInfo::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + assert_eq!(info, TokenFedAuthInfo::default()); + } + + #[tokio::test] + async fn decode_consumes_exactly_token_length() { + // `decode` must read exactly the `TokenLength` bytes of body it is told + // to and stop, leaving whatever follows for the next token. A trailing + // sentinel that reads back unchanged proves the length-prefixed body was + // consumed exactly and the stream was not desynced. + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use bytes::{BufMut, BytesMut}; + + // Minimal valid body: CountOfInfoIDs == 0, no options. + let body = 0u32.to_le_bytes(); + + let mut buf = BytesMut::new(); + buf.put_u32_le(body.len() as u32); // TokenLength + buf.put_slice(&body); + buf.put_u32_le(0xDEAD_BEEF); // sentinel: belongs to the next token + + let mut reader = buf.into_sql_read_bytes(); + let info = TokenFedAuthInfo::decode(&mut reader).await.unwrap(); + assert_eq!(info, TokenFedAuthInfo::default()); + + let sentinel = reader + .read_u32_le() + .await + .expect("the sentinel following the FEDAUTHINFO token must still be readable"); + assert_eq!( + sentinel, 0xDEAD_BEEF, + "decode consumed past its declared TokenLength and desynced the stream" + ); + } + + #[tokio::test] + async fn decode_rejects_body_shorter_than_token_length() { + // `TokenLength` declares more body bytes than the stream actually holds. + // The `read_exact` for the body must surface a clean error rather than + // hanging or panicking, so a truncated/inconsistent frame fails safely. + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use bytes::{BufMut, BytesMut}; + + let mut buf = BytesMut::new(); + buf.put_u32_le(64); // claims 64 body bytes + buf.put_slice(&[0u8; 8]); // but only 8 are present + + let err = TokenFedAuthInfo::decode(&mut buf.into_sql_read_bytes()) + .await + .expect_err("a body shorter than the declared TokenLength must error"); + assert!(matches!(err, Error::Io { .. })); + } +} diff --git a/src/tds/codec/token/token_info.rs b/src/tds/codec/token/token_info.rs index 96c4eca53..b36ed245d 100644 --- a/src/tds/codec/token/token_info.rs +++ b/src/tds/codec/token/token_info.rs @@ -1,6 +1,10 @@ +use super::token_error::decode_error_info_body; use crate::SqlReadBytes; -#[allow(dead_code)] // we might want to debug the values +// Fields are decoded from the INFO token and logged at DEBUG on receipt, but +// are not otherwise read; retained for `Debug` diagnostics and future surfacing +// of server informational messages to callers. +#[allow(dead_code)] #[derive(Debug)] pub struct TokenInfo { /// info number @@ -20,24 +24,95 @@ impl TokenInfo { where R: SqlReadBytes + Unpin, { - let _length = src.read_u16_le().await?; - - let number = src.read_u32_le().await?; - let state = src.read_u8().await?; - let class = src.read_u8().await?; - let message = src.read_us_varchar().await?; - let server = src.read_b_varchar().await?; - let procedure = src.read_b_varchar().await?; - let line = src.read_u32_le().await?; + // INFO and ERROR share an identical body layout (MS-TDS §2.2.7.13 / + // §2.2.7.10); reuse the shared, length-bounded decoder. `number` is the + // INFO spelling of ERROR's `code` field. + let body = decode_error_info_body(src).await?; Ok(TokenInfo { - number, - state, - class, - message, - server, - procedure, - line, + number: body.code, + state: body.state, + class: body.class, + message: body.message, + server: body.server, + procedure: body.procedure, + line: body.line, }) } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use bytes::{BufMut, BytesMut}; + + fn put_b_varchar(buf: &mut BytesMut, s: &str) { + let utf16: Vec = s.encode_utf16().collect(); + buf.put_u8(utf16.len() as u8); + for c in utf16 { + buf.put_u16_le(c); + } + } + + fn put_us_varchar(buf: &mut BytesMut, s: &str) { + let utf16: Vec = s.encode_utf16().collect(); + buf.put_u16_le(utf16.len() as u16); + for c in utf16 { + buf.put_u16_le(c); + } + } + + #[tokio::test] + async fn decodes_all_fields() { + let mut body = BytesMut::new(); + body.put_u32_le(4711); + body.put_u8(2); + body.put_u8(9); + put_us_varchar(&mut body, "informational"); + put_b_varchar(&mut body, "server"); + put_b_varchar(&mut body, "proc"); + body.put_u32_le(123); + + let mut buf = BytesMut::new(); + buf.put_u16_le(body.len() as u16); // declared Length + buf.put_slice(&body); + + let info = TokenInfo::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + assert_eq!(info.number, 4711); + assert_eq!(info.state, 2); + assert_eq!(info.class, 9); + assert_eq!(info.message, "informational"); + assert_eq!(info.server, "server"); + assert_eq!(info.procedure, "proc"); + assert_eq!(info.line, 123); + } + + #[tokio::test] + async fn decode_reads_full_four_byte_line_number_on_tds72_plus() { + // The default test context reports SqlServerN (>= TDS 7.2), so the + // LineNumber must be read as a 4-byte LONG. 0x0001_0001 (65537) reads as 1 + // when truncated to 2 bytes but as 65537 when read correctly as 4 bytes. + let mut body = BytesMut::new(); + body.put_u32_le(4711); + body.put_u8(2); + body.put_u8(9); + put_us_varchar(&mut body, "informational"); + put_b_varchar(&mut body, "server"); + put_b_varchar(&mut body, "proc"); + body.put_u32_le(0x0001_0001); + + let mut buf = BytesMut::new(); + buf.put_u16_le(body.len() as u16); // declared Length + buf.put_slice(&body); + + let info = TokenInfo::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + assert_eq!(info.line, 0x0001_0001); + } +} diff --git a/src/tds/codec/token/token_login_ack.rs b/src/tds/codec/token/token_login_ack.rs index 28a4bfc40..043b87bd1 100644 --- a/src/tds/codec/token/token_login_ack.rs +++ b/src/tds/codec/token/token_login_ack.rs @@ -1,13 +1,20 @@ use crate::{Error, FeatureLevel, SqlReadBytes}; +use byteorder::{LittleEndian, ReadBytesExt}; +use futures_util::io::AsyncReadExt; use std::convert::TryFrom; +use std::io::Cursor; -#[allow(dead_code)] // we might want to debug the values +// `tds_version` is applied to the connection `Context` (see +// `TokenStream::get_login_ack`), and `prog_name`/`version` are logged via +// `event!(Level::DEBUG, ...)` there. Only `interface` is otherwise unread; +// retained for `Debug` diagnostics. #[derive(Debug)] pub struct TokenLoginAck { /// The type of interface with which the server will accept client requests /// 0: SQL_DFLT (server confirms that whatever is sent by the client is acceptable. If the client /// requested SQL_DFLT, SQL_TSQL will be used) /// 1: SQL_TSQL (TSQL is accepted) + #[allow(dead_code)] pub(crate) interface: u8, pub(crate) tds_version: FeatureLevel, pub(crate) prog_name: String, @@ -15,20 +22,60 @@ pub struct TokenLoginAck { pub(crate) version: u32, } +/// Read a `B_VARCHAR` (a single `u8` UTF-16 code-unit count followed by that +/// many little-endian `u16` code units) from an in-memory cursor. +fn read_b_varchar(cur: &mut Cursor<&[u8]>) -> crate::Result { + let char_count = cur.read_u8()? as usize; + let mut units = Vec::with_capacity(char_count); + + for _ in 0..char_count { + units.push(cur.read_u16::()?); + } + + String::from_utf16(&units) + .map_err(|_| Error::Protocol("Login ACK: ProgName is not valid UTF-16".into())) +} + impl TokenLoginAck { pub(crate) async fn decode(src: &mut R) -> crate::Result where R: SqlReadBytes + Unpin, { - let _length = src.read_u16_le().await?; + let length = src.read_u16_le().await? as usize; + + // `Interface` (1) and `TDSVersion` (4) are fixed-width, so they cannot + // desync on their own; read them directly. The rest of the body + // (`ProgName`, a greedy B_VARCHAR, plus the 4-byte build version) is what + // could over-/under-read on a Length/content mismatch, so it is read + // into a bounded buffer sized by the *remaining* declared Length and + // parsed from there. Total consumption is therefore exactly `Length`, + // keeping the token stream aligned. + const FIXED_PREFIX: usize = 5; // interface (1) + tds_version (4) + + if length < FIXED_PREFIX { + // Consume whatever the (too-small) Length declared so the stream + // stays aligned to the next token boundary, then fail cleanly. + let mut discard = vec![0u8; length]; + src.read_exact(&mut discard).await?; + + return Err(Error::Protocol( + "Login ACK: token length shorter than fixed header".into(), + )); + } let interface = src.read_u8().await?; + // TDS version is a 4-byte big-endian value. let tds_version = FeatureLevel::try_from(src.read_u32().await?) .map_err(|_| Error::Protocol("Login ACK: Invalid TDS version".into()))?; - let prog_name = src.read_b_varchar().await?; - let version = src.read_u32_le().await?; + let mut data = vec![0u8; length - FIXED_PREFIX]; + src.read_exact(&mut data).await?; + + let mut cur = Cursor::new(&data[..]); + + let prog_name = read_b_varchar(&mut cur)?; + let version = cur.read_u32::()?; Ok(TokenLoginAck { interface, @@ -38,3 +85,159 @@ impl TokenLoginAck { }) } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use bytes::{BufMut, BytesMut}; + + fn put_b_varchar(buf: &mut BytesMut, s: &str) { + let utf16: Vec = s.encode_utf16().collect(); + buf.put_u8(utf16.len() as u8); + for c in utf16 { + buf.put_u16_le(c); + } + } + + /// Build a LOGINACK wire buffer: the token body prefixed by its exact + /// `Length` (u16, little-endian). + fn ack_body(interface: u8, tds_version: u32, prog_name: &str, version: u32) -> BytesMut { + let mut body = BytesMut::new(); + body.put_u8(interface); + body.put_u32(tds_version); // big-endian tds version + put_b_varchar(&mut body, prog_name); + body.put_u32_le(version); + + let mut buf = BytesMut::new(); + buf.put_u16_le(body.len() as u16); + buf.put_slice(&body); + buf + } + + #[tokio::test] + async fn decodes_valid_ack() { + let buf = ack_body( + 1, + FeatureLevel::SqlServerN as u32, + "Microsoft SQL Server", + 0x0F00_0FA0, + ); + + let ack = TokenLoginAck::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + assert_eq!(ack.interface, 1); + assert_eq!(ack.tds_version, FeatureLevel::SqlServerN); + assert_eq!(ack.prog_name, "Microsoft SQL Server"); + assert_eq!(ack.version, 0x0F00_0FA0); + } + + #[tokio::test] + async fn negotiated_version_is_recorded_in_context() { + // Mirrors the wiring in `TokenStream::get_login_ack`: after decoding a + // LOGINACK the negotiated `tds_version` must be pushed into the + // connection context so version-dependent decoders use the right + // layout. A pre-2005 version differs from the default (SqlServerN), so + // this proves the value is actually applied and not left at the default. + use crate::SqlReadBytes; + + let buf = ack_body( + 1, + FeatureLevel::SqlServer2000 as u32, + "Microsoft SQL Server", + 0x0800_0000, + ); + + let mut reader = buf.into_sql_read_bytes(); + assert_eq!(reader.context().version(), FeatureLevel::SqlServerN); + + let ack = TokenLoginAck::decode(&mut reader).await.unwrap(); + assert_eq!(ack.tds_version, FeatureLevel::SqlServer2000); + + // The exact call `get_login_ack` performs. + reader.context_mut().set_version(ack.tds_version); + + assert_eq!(reader.context().version(), FeatureLevel::SqlServer2000); + } + + #[tokio::test] + async fn invalid_tds_version_errors() { + let mut body = BytesMut::new(); + body.put_u8(1); // interface + body.put_u32(0xDEAD_BEEF); // not a valid FeatureLevel + + let mut buf = BytesMut::new(); + buf.put_u16_le(body.len() as u16); + buf.put_slice(&body); + + let err = TokenLoginAck::decode(&mut buf.into_sql_read_bytes()) + .await + .expect_err("must fail on invalid version"); + + assert!(matches!(err, Error::Protocol(_))); + } + + #[tokio::test] + async fn under_declared_length_is_clean_error_and_does_not_desync() { + // Declared Length (5) covers only interface + tds_version; the ProgName + // length byte and everything after it fall outside it. Decoding must + // fail cleanly AND consume exactly `Length` body bytes, leaving the + // following byte on the wire readable — i.e. no desync. A decoder that + // ignored `Length` (the old behaviour) would greedily read the trailing + // bytes as ProgName/version and succeed, desyncing the stream. + use crate::SqlReadBytes; + + let mut buf = BytesMut::new(); + buf.put_u16_le(5); // declared Length (too small for the full body) + buf.put_u8(1); // interface + buf.put_u32(FeatureLevel::SqlServerN as u32); // big-endian tds version + + // The bytes below belong to the *next* token on the wire; a length- + // ignoring decoder would (mis)read them as ProgName + version. + buf.put_u8(0); // (mis)read as ProgName length 0 + buf.put_u32_le(0x1234_5678); // (mis)read as version + buf.put_u8(0xEF); // sentinel further along + + let mut reader = buf.into_sql_read_bytes(); + + let err = TokenLoginAck::decode(&mut reader) + .await + .expect_err("under-declared length must be a clean error"); + assert!(matches!(err, Error::Protocol(_) | Error::Io { .. })); + + // Exactly the 5 declared body bytes were consumed, so the first byte of + // the next token region is still on the wire (no desync). + assert_eq!(reader.read_u8().await.unwrap(), 0x00); + } + + #[tokio::test] + async fn trailing_bytes_within_declared_length_are_ignored() { + // Declared Length is larger than the fields consume; the extra padding + // is discarded and the next token (sentinel) is read cleanly afterwards. + use crate::SqlReadBytes; + + let mut body = BytesMut::new(); + body.put_u8(1); // interface + body.put_u32(FeatureLevel::SqlServerN as u32); // tds version + put_b_varchar(&mut body, "srv"); + body.put_u32_le(0x0F00_0FA0); // version + body.put_u8(0x99); // trailing padding within the declared length + + let mut buf = BytesMut::new(); + buf.put_u16_le(body.len() as u16); + buf.put_slice(&body); + buf.put_u8(0xCD); // sentinel next token + + let mut reader = buf.into_sql_read_bytes(); + + let ack = TokenLoginAck::decode(&mut reader) + .await + .expect("trailing padding within length must decode"); + assert_eq!(ack.prog_name, "srv"); + assert_eq!(ack.version, 0x0F00_0FA0); + + assert_eq!(reader.read_u8().await.unwrap(), 0xCD); + } +} diff --git a/src/tds/codec/token/token_order.rs b/src/tds/codec/token/token_order.rs index d39dfdbb2..8d90bd685 100644 --- a/src/tds/codec/token/token_order.rs +++ b/src/tds/codec/token/token_order.rs @@ -1,6 +1,9 @@ use crate::SqlReadBytes; -#[allow(dead_code)] // we might want to debug the values +// The ordered column indexes are decoded and traced at TRACE on receipt but are +// not otherwise consumed; retained for `Debug` diagnostics and future surfacing +// of result ordering to callers. +#[allow(dead_code)] #[derive(Debug)] pub struct TokenOrder { pub(crate) column_indexes: Vec, @@ -11,9 +14,21 @@ impl TokenOrder { where R: SqlReadBytes + Unpin, { - let len = src.read_u16_le().await? / 2; + // `Length` is the byte length of the column-index list; each index is a + // 2-byte USHORT, so it must be even. An odd length would truncate on the + // `/ 2` below and leave a stray byte unconsumed, desyncing the stream. + let raw_len = src.read_u16_le().await?; + if raw_len % 2 != 0 { + return Err(crate::Error::Protocol( + format!("ORDER token length {raw_len} is not a multiple of 2").into(), + )); + } + let len = raw_len / 2; - let mut column_indexes = Vec::with_capacity(len as usize); + // `len` is derived from an untrusted u16; cap the up-front reservation + // (the Vec still grows as indexes are actually read). + let mut column_indexes = + Vec::with_capacity((len as usize).min(crate::tds::codec::column_data::MAX_PREALLOC)); for _ in 0..len { column_indexes.push(src.read_u16_le().await?); @@ -22,3 +37,96 @@ impl TokenOrder { Ok(TokenOrder { column_indexes }) } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use bytes::{BufMut, BytesMut}; + + #[tokio::test] + async fn decodes_column_indexes() { + let mut buf = BytesMut::new(); + // length is in bytes; three u16 indexes => 6 bytes + buf.put_u16_le(6); + buf.put_u16_le(1); + buf.put_u16_le(2); + buf.put_u16_le(3); + + let order = TokenOrder::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + assert_eq!(order.column_indexes, vec![1, 2, 3]); + } + + #[tokio::test] + async fn decodes_empty() { + let mut buf = BytesMut::new(); + buf.put_u16_le(0); + + let order = TokenOrder::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + + assert!(order.column_indexes.is_empty()); + } + + #[tokio::test] + async fn rejects_odd_length() { + // An odd byte length cannot hold a whole number of 2-byte indexes; it + // must be a protocol error rather than silently truncating and leaving a + // stray byte on the wire. + let mut buf = BytesMut::new(); + buf.put_u16_le(3); // odd + buf.put_u16_le(1); + + let err = TokenOrder::decode(&mut buf.into_sql_read_bytes()) + .await + .expect_err("odd length must be rejected"); + + assert!(matches!(err, crate::Error::Protocol(_))); + } + + #[tokio::test] + async fn consumes_exactly_declared_length() { + // The decoder must read exactly `Length` bytes of index data and stop, + // leaving anything that follows untouched. A trailing sentinel that + // reads back cleanly proves the token did not over- or under-consume + // and so did not desync the stream for the next token. + let mut buf = BytesMut::new(); + buf.put_u16_le(4); // two u16 indexes + buf.put_u16_le(10); + buf.put_u16_le(20); + buf.put_u16_le(0xBEEF); // sentinel: belongs to the *next* token + + let mut reader = buf.into_sql_read_bytes(); + let order = TokenOrder::decode(&mut reader).await.unwrap(); + assert_eq!(order.column_indexes, vec![10, 20]); + + let sentinel = reader + .read_u16_le() + .await + .expect("the sentinel following the ORDER token must still be readable"); + assert_eq!( + sentinel, 0xBEEF, + "decode consumed past its declared length and desynced the stream" + ); + } + + #[tokio::test] + async fn rejects_length_longer_than_content() { + // `Length` claims three indexes (6 bytes) but only one is present. The + // decoder must surface a clean error when the stream runs out rather + // than panicking or looping. + let mut buf = BytesMut::new(); + buf.put_u16_le(6); // three indexes declared + buf.put_u16_le(1); // only one provided + + let err = TokenOrder::decode(&mut buf.into_sql_read_bytes()) + .await + .expect_err("a length longer than the content must error"); + // An out-of-data read surfaces as an IO error through the reader. + assert!(matches!(err, crate::Error::Io { .. })); + } +} diff --git a/src/tds/codec/token/token_return_value.rs b/src/tds/codec/token/token_return_value.rs index 183e46be0..864d122b1 100644 --- a/src/tds/codec/token/token_return_value.rs +++ b/src/tds/codec/token/token_return_value.rs @@ -1,13 +1,18 @@ use super::BaseMetaDataColumn; use crate::{tds::codec::ColumnData, Error, SqlReadBytes}; +// Decoded from the RETURNVALUE token. `param_ordinal`, `param_name`, and +// `value` are forwarded to callers via `CommandReturnValue` in +// `src/tds/stream/command.rs`; `udf` and `meta` are not otherwise read and are +// retained for `Debug` diagnostics and future surfacing to callers. #[derive(Debug)] -#[allow(dead_code)] pub struct TokenReturnValue { pub param_ordinal: u16, pub param_name: String, /// return value of user defined function + #[allow(dead_code)] pub udf: bool, + #[allow(dead_code)] pub meta: BaseMetaDataColumn, pub value: ColumnData<'static>, } @@ -40,3 +45,65 @@ impl TokenReturnValue { Ok(token) } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use crate::tds::codec::{Encode, FixedLenType, TypeInfo}; + use bytes::{BufMut, BytesMut}; + + fn put_b_varchar(buf: &mut BytesMut, s: &str) { + let utf16: Vec = s.encode_utf16().collect(); + buf.put_u8(utf16.len() as u8); + for c in utf16 { + buf.put_u16_le(c); + } + } + + fn build(status: u8) -> BytesMut { + let mut buf = BytesMut::new(); + buf.put_u16_le(1); // param ordinal + put_b_varchar(&mut buf, "@out"); + buf.put_u8(status); + + // BaseMetaDataColumn: user_ty, flags, type info + buf.put_u32_le(0); + buf.put_u16_le(0); + TypeInfo::FixedLen(FixedLenType::Int4) + .encode(&mut buf) + .unwrap(); + + // value payload (i32) + buf.put_i32_le(42); + buf + } + + #[tokio::test] + async fn decodes_non_udf_value() { + let token = TokenReturnValue::decode(&mut build(0x01).into_sql_read_bytes()) + .await + .unwrap(); + + assert_eq!(token.param_ordinal, 1); + assert_eq!(token.param_name, "@out"); + assert!(!token.udf); + assert_eq!(token.value, ColumnData::I32(Some(42))); + } + + #[tokio::test] + async fn decodes_udf_flag() { + let token = TokenReturnValue::decode(&mut build(0x02).into_sql_read_bytes()) + .await + .unwrap(); + assert!(token.udf); + } + + #[tokio::test] + async fn invalid_status_errors() { + let err = TokenReturnValue::decode(&mut build(0x00).into_sql_read_bytes()) + .await + .expect_err("invalid status must fail"); + assert!(matches!(err, Error::Protocol(_))); + } +} diff --git a/src/tds/codec/token/token_row.rs b/src/tds/codec/token/token_row.rs index b1ff16b6c..3c0cc15e8 100644 --- a/src/tds/codec/token/token_row.rs +++ b/src/tds/codec/token/token_row.rs @@ -9,6 +9,7 @@ pub use into_row::IntoRow; /// A row of data. #[derive(Debug, Default, Clone)] +#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] pub struct TokenRow<'a> { data: Vec>, } @@ -100,7 +101,9 @@ impl TokenRow<'static> { where R: SqlReadBytes + Unpin, { - let col_meta = src.context().last_meta().unwrap(); + let col_meta = src.context().last_meta().ok_or_else(|| { + crate::Error::Protocol("ROW token arrived before any COLMETADATA".into()) + })?; let mut row = Self { data: Vec::with_capacity(col_meta.columns.len()), @@ -120,7 +123,9 @@ impl TokenRow<'static> { where R: SqlReadBytes + Unpin, { - let col_meta = src.context().last_meta().unwrap(); + let col_meta = src.context().last_meta().ok_or_else(|| { + crate::Error::Protocol("NBCROW token arrived before any COLMETADATA".into()) + })?; let row_bitmap = RowBitmap::decode(src, col_meta.columns.len()).await?; let mut row = Self { @@ -177,7 +182,7 @@ impl RowBitmap { where R: SqlReadBytes + Unpin, { - let size = (columns + 8 - 1) / 8; + let size = columns.div_ceil(8); let mut data = vec![0; size]; src.read_exact(&mut data[0..size]).await?; @@ -198,6 +203,7 @@ mod tests { base: BaseMetaDataColumn { flags: ColumnFlag::Nullable.into(), ty: TypeInfo::FixedLen(FixedLenType::Bit), + table_name: None, }, col_name: Default::default(), }]; @@ -207,4 +213,234 @@ mod tests { row.encode(&mut buf_with_columns) .expect_err("wrong number of columns"); } + + // A row whose encoding fails partway (here: an out-of-range money value in + // the second column, after the Row token byte and first column are already + // written) must be rolled back by the caller so the bulk stream stays in + // sync. This mirrors `BulkLoadRequest::send`'s snapshot-and-truncate logic. + #[tokio::test] + async fn partial_row_can_be_rolled_back_on_encode_error() { + use crate::tds::codec::type_info::VarLenContext; + use crate::{ColumnData, VarLenType}; + + let columns = vec![ + MetaDataColumn { + base: BaseMetaDataColumn { + flags: ColumnFlag::Nullable.into(), + ty: TypeInfo::FixedLen(FixedLenType::Int4), + table_name: None, + }, + col_name: Default::default(), + }, + MetaDataColumn { + base: BaseMetaDataColumn { + flags: ColumnFlag::Nullable.into(), + ty: TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Money, 8, None)), + table_name: None, + }, + col_name: Default::default(), + }, + ]; + + // Pretend some earlier, fully-encoded rows already sit in the buffer. + let mut buf = BytesMut::new(); + buf.extend_from_slice(&[0xde, 0xad, 0xbe, 0xef]); + let snapshot = buf.to_vec(); + let start = buf.len(); + + let mut row = TokenRow::new(); + row.push(ColumnData::I32(Some(1))); + row.push(ColumnData::F64(Some(1e18))); // out of range for money + + let mut buf_with_columns = BytesMutWithDataColumns::new(&mut buf, &columns); + let err = row.encode(&mut buf_with_columns).unwrap_err(); + assert!(matches!(err, crate::Error::BulkInput(_)), "got {err:?}"); + + // Partial bytes (Row token + first column) were written... + assert!(buf.len() > start, "expected a partial row to be present"); + // ...and truncating back to the snapshot restores the buffer exactly. + buf.truncate(start); + assert_eq!(buf.to_vec(), snapshot); + } + + #[tokio::test] + async fn row_before_colmetadata_is_protocol_error() { + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + // No COLMETADATA has been seen, so last_meta() is None: decoding a ROW + // must be a protocol error rather than an unwrap() panic. + let buf = BytesMut::new(); + let err = TokenRow::decode(&mut buf.into_sql_read_bytes()) + .await + .expect_err("ROW before COLMETADATA must error"); + assert!(matches!(err, crate::Error::Protocol(_))); + } + + #[test] + fn basic_container_operations() { + let mut row = TokenRow::new(); + assert!(row.is_empty()); + assert_eq!(row.len(), 0); + assert_eq!(row.get(0), None); + + row.push(ColumnData::I32(Some(1))); + row.push(ColumnData::I32(Some(2))); + assert_eq!(row.len(), 2); + assert!(!row.is_empty()); + assert_eq!(row.get(0), Some(&ColumnData::I32(Some(1)))); + assert_eq!(row.get(5), None); + + let collected: Vec<_> = row.iter().collect(); + assert_eq!(collected.len(), 2); + + row.clear(); + assert!(row.is_empty()); + + let with_cap = TokenRow::with_capacity(4); + assert!(with_cap.is_empty()); + } + + #[test] + fn with_capacity_preallocates() { + // with_capacity must actually reserve room; Default::default() would + // give a zero-capacity vec. + let row = TokenRow::with_capacity(16); + assert!(row.is_empty()); + assert!(row.data.capacity() >= 16); + } + + #[test] + fn row_bitmap_is_null_checks_correct_bit() { + // Only bit 3 is set in the single bitmap byte; is_null must consult that + // exact bit. + let bitmap = RowBitmap { + data: vec![0b0000_1000], + }; + + assert!(bitmap.is_null(3)); + assert!(!bitmap.is_null(0)); + assert!(!bitmap.is_null(1)); + assert!(!bitmap.is_null(2)); + assert!(!bitmap.is_null(4)); + } + + #[test] + fn into_iter_yields_owned_values() { + let mut row = TokenRow::new(); + row.push(ColumnData::I32(Some(1))); + row.push(ColumnData::I32(Some(2))); + + let values: Vec<_> = row.into_iter().collect(); + assert_eq!( + values, + vec![ColumnData::I32(Some(1)), ColumnData::I32(Some(2))] + ); + } + + #[tokio::test] + async fn encode_matching_columns_round_trip() { + let row = (true, 5i32).into_row(); + let columns = vec![ + MetaDataColumn { + base: BaseMetaDataColumn { + flags: ColumnFlag::Nullable.into(), + ty: TypeInfo::FixedLen(FixedLenType::Bit), + table_name: None, + }, + col_name: Default::default(), + }, + MetaDataColumn { + base: BaseMetaDataColumn { + flags: ColumnFlag::Nullable.into(), + ty: TypeInfo::FixedLen(FixedLenType::Int4), + table_name: None, + }, + col_name: Default::default(), + }, + ]; + let mut buf = BytesMut::new(); + let mut buf_with_columns = BytesMutWithDataColumns::new(&mut buf, &columns); + + row.encode(&mut buf_with_columns).unwrap(); + assert!(!buf.is_empty()); + } + + #[tokio::test] + async fn decode_reads_columns_from_cached_meta() { + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use crate::tds::codec::TokenColMetaData; + use std::sync::Arc; + + let col_meta = TokenColMetaData { + columns: vec![MetaDataColumn { + base: BaseMetaDataColumn { + flags: ColumnFlag::Nullable.into(), + ty: TypeInfo::FixedLen(FixedLenType::Int4), + table_name: None, + }, + col_name: Default::default(), + }], + }; + + let mut buf = BytesMut::new(); + buf.put_i32_le(42); + + let mut reader = buf.into_sql_read_bytes(); + reader.context_mut().set_last_meta(Arc::new(col_meta)); + + let row = TokenRow::decode(&mut reader).await.unwrap(); + assert_eq!(row.len(), 1); + assert_eq!(row.get(0), Some(&ColumnData::I32(Some(42)))); + } + + #[tokio::test] + async fn decode_nbc_before_colmetadata_is_protocol_error() { + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + + let buf = BytesMut::new(); + let err = TokenRow::decode_nbc(&mut buf.into_sql_read_bytes()) + .await + .expect_err("NBCROW before COLMETADATA must error"); + assert!(matches!(err, crate::Error::Protocol(_))); + } + + #[tokio::test] + async fn decode_nbc_uses_bitmap_for_nulls() { + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use crate::tds::codec::TokenColMetaData; + use std::sync::Arc; + + // Two int columns: first null (bit 0 set), second present (value 7). + let col_meta = TokenColMetaData { + columns: vec![ + MetaDataColumn { + base: BaseMetaDataColumn { + flags: ColumnFlag::Nullable.into(), + ty: TypeInfo::FixedLen(FixedLenType::Int4), + table_name: None, + }, + col_name: Default::default(), + }, + MetaDataColumn { + base: BaseMetaDataColumn { + flags: ColumnFlag::Nullable.into(), + ty: TypeInfo::FixedLen(FixedLenType::Int4), + table_name: None, + }, + col_name: Default::default(), + }, + ], + }; + + let mut buf = BytesMut::new(); + buf.put_u8(0b0000_0001); // bitmap: column 0 is null + buf.put_i32_le(7); // column 1's value + + let mut reader = buf.into_sql_read_bytes(); + reader.context_mut().set_last_meta(Arc::new(col_meta)); + + let row = TokenRow::decode_nbc(&mut reader).await.unwrap(); + assert_eq!(row.len(), 2); + assert_eq!(row.get(0), Some(&ColumnData::I32(None))); + assert_eq!(row.get(1), Some(&ColumnData::I32(Some(7)))); + } } diff --git a/src/tds/codec/token/token_row/into_row.rs b/src/tds/codec/token/token_row/into_row.rs index 8bee0dcd4..c0615d351 100644 --- a/src/tds/codec/token/token_row/into_row.rs +++ b/src/tds/codec/token/token_row/into_row.rs @@ -205,3 +205,58 @@ where row } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::tds::codec::ColumnData; + + #[test] + fn single_value_into_row() { + let row = 42i32.into_row(); + assert_eq!(row.len(), 1); + assert_eq!(row.get(0), Some(&ColumnData::I32(Some(42)))); + } + + #[test] + fn tuple_arities_produce_expected_lengths_and_order() { + assert_eq!((1i32, 2i32).into_row().len(), 2); + assert_eq!((1i32, 2i32, 3i32).into_row().len(), 3); + assert_eq!((1i32, 2i32, 3i32, 4i32).into_row().len(), 4); + assert_eq!((1i32, 2i32, 3i32, 4i32, 5i32).into_row().len(), 5); + assert_eq!((1i32, 2i32, 3i32, 4i32, 5i32, 6i32).into_row().len(), 6); + assert_eq!( + (1i32, 2i32, 3i32, 4i32, 5i32, 6i32, 7i32).into_row().len(), + 7 + ); + assert_eq!( + (1i32, 2i32, 3i32, 4i32, 5i32, 6i32, 7i32, 8i32) + .into_row() + .len(), + 8 + ); + assert_eq!( + (1i32, 2i32, 3i32, 4i32, 5i32, 6i32, 7i32, 8i32, 9i32) + .into_row() + .len(), + 9 + ); + + let row = (1i32, 2i32, 3i32, 4i32, 5i32, 6i32, 7i32, 8i32, 9i32, 10i32).into_row(); + assert_eq!(row.len(), 10); + + // Values are pushed in tuple order. + for (i, value) in row.iter().enumerate() { + assert_eq!(value, &ColumnData::I32(Some(i as i32 + 1))); + } + } + + #[test] + fn mixed_types_preserve_positions() { + let row = (true, 7u8, "hello", 3.5f64).into_row(); + assert_eq!(row.len(), 4); + assert_eq!(row.get(0), Some(&ColumnData::Bit(Some(true)))); + assert_eq!(row.get(1), Some(&ColumnData::U8(Some(7)))); + assert_eq!(row.get(3), Some(&ColumnData::F64(Some(3.5)))); + } +} diff --git a/src/tds/codec/token/token_session_state.rs b/src/tds/codec/token/token_session_state.rs new file mode 100644 index 000000000..a34a3fd7c --- /dev/null +++ b/src/tds/codec/token/token_session_state.rs @@ -0,0 +1,286 @@ +use crate::{Error, SqlReadBytes}; +use byteorder::{LittleEndian, ReadBytesExt}; +use futures_util::io::AsyncReadExt; +use std::io::{Cursor, Read}; + +/// A single session state value carried by a [`TokenSessionState`] token. +/// +/// Each entry is identified by a `state_id` and carries an opaque, driver +/// server-defined payload. The client is expected to retain these values and +/// replay them when transparently re-establishing a broken connection during +/// session recovery. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SessionStateValue { + /// The identifier of the state slot (`StateId`). + pub id: u8, + /// The opaque state payload (`StateValue`). + pub value: Vec, +} + +/// The `SESSIONSTATE` token (`0xE4`). +/// +/// Sent by the server as part of the connection-resiliency / session-recovery +/// feature (MS-TDS §2.2.7.22). It informs the client about the current session +/// state so that the client can transparently reconnect and restore the +/// session after an idle connection has been broken. Tiberius does not yet +/// initiate transparent reconnects, so the token is decoded and retained. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct TokenSessionState { + /// Sequence number of this session-state update (`SeqNo`). A value of + /// `0xFFFFFFFF` indicates that the state cannot be reset and the connection + /// is not recoverable. + pub seq_no: u32, + /// The `Status` byte. Bit 0 (`fRecoverable`) indicates whether the session + /// is currently in a recoverable state. + pub status: u8, + /// The individual state values carried by this token. + pub states: Vec, +} + +impl TokenSessionState { + /// Whether the session is currently recoverable (`fRecoverable`, bit 0 of + /// the `Status` byte). + pub fn is_recoverable(&self) -> bool { + self.status & 0x01 != 0 + } + + /// Parse the token body (everything following the `TokenType` and `Length` + /// fields) from an in-memory buffer. + fn parse(bytes: Vec) -> crate::Result { + let mut buf = Cursor::new(bytes); + + let seq_no = buf.read_u32::()?; + let status = buf.read_u8()?; + + let mut states = Vec::new(); + + // The remaining bytes are a sequence of SessionStateData entries that + // fill exactly the token length. Keep decoding until the buffer is + // exhausted. + let total = buf.get_ref().len() as u64; + + while buf.position() < total { + let id = buf.read_u8()?; + + // StateLen is a single byte, unless it is 0xFF, in which case a + // 4-byte (LONG) length follows. + let short_len = buf.read_u8()?; + let state_len = if short_len == 0xFF { + buf.read_u32::()? as usize + } else { + short_len as usize + }; + + // `state_len` (up to a full u32 via the 0xFF LONG escape) is + // untrusted. Even though the outer token body is capped at + // MAX_TOKEN_BODY, a single entry could still declare ~4GiB while the + // token itself is only a few bytes on the wire. Reject any length + // that cannot possibly fit in the remaining buffered bytes before + // allocating, so `vec![0u8; state_len]` can't be used for + // memory exhaustion. + let remaining = total - buf.position(); + if state_len as u64 > remaining { + return Err(Error::Protocol( + format!( + "SESSIONSTATE entry length {state_len} exceeds the {remaining} bytes remaining in the token" + ) + .into(), + )); + } + + let mut value = vec![0u8; state_len]; + buf.read_exact(&mut value)?; + + states.push(SessionStateValue { id, value }); + } + + Ok(TokenSessionState { + seq_no, + status, + states, + }) + } + + pub(crate) async fn decode(src: &mut R) -> crate::Result + where + R: SqlReadBytes + Unpin, + { + // Length (ULONG) of the token stream that follows, covering SeqNo, + // Status and all SessionStateData entries. + let len = src.read_u32_le().await? as usize; + + if len > super::MAX_TOKEN_BODY { + return Err(Error::Protocol( + format!("SESSIONSTATE token length {len} exceeds the maximum").into(), + )); + } + + let mut bytes = vec![0u8; len]; + src.read_exact(&mut bytes[0..len]).await?; + + if bytes.len() < 5 { + return Err(Error::Protocol( + "SESSIONSTATE token too short to contain SeqNo and Status".into(), + )); + } + + Self::parse(bytes) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn parse_two_state_values() { + // SeqNo = 1, Status = 0x01 (fRecoverable), then two SessionStateData + // entries with short lengths. + let mut body = Vec::new(); + body.extend_from_slice(&1u32.to_le_bytes()); // SeqNo + body.push(0x01); // Status: recoverable + + // State 1: id = 0, len = 3, value = [0xAA, 0xBB, 0xCC] + body.push(0x00); + body.push(0x03); + body.extend_from_slice(&[0xAA, 0xBB, 0xCC]); + + // State 2: id = 7, len = 1, value = [0x42] + body.push(0x07); + body.push(0x01); + body.push(0x42); + + let token = TokenSessionState::parse(body).unwrap(); + + assert_eq!(token.seq_no, 1); + assert_eq!(token.status, 0x01); + assert!(token.is_recoverable()); + assert_eq!(token.states.len(), 2); + + assert_eq!(token.states[0].id, 0); + assert_eq!(token.states[0].value, vec![0xAA, 0xBB, 0xCC]); + + assert_eq!(token.states[1].id, 7); + assert_eq!(token.states[1].value, vec![0x42]); + } + + #[test] + fn parse_long_state_length() { + // A single state whose length is encoded with the 0xFF escape followed + // by a 4-byte length. + let mut body = Vec::new(); + body.extend_from_slice(&0xFFFF_FFFFu32.to_le_bytes()); // SeqNo (not recoverable) + body.push(0x00); // Status: not recoverable + + body.push(0x02); // StateId + body.push(0xFF); // long-length escape + body.extend_from_slice(&300u32.to_le_bytes()); // StateLen = 300 + body.extend_from_slice(&vec![0x5A; 300]); + + let token = TokenSessionState::parse(body).unwrap(); + + assert_eq!(token.seq_no, 0xFFFF_FFFF); + assert!(!token.is_recoverable()); + assert_eq!(token.states.len(), 1); + assert_eq!(token.states[0].id, 2); + assert_eq!(token.states[0].value.len(), 300); + assert!(token.states[0].value.iter().all(|&b| b == 0x5A)); + } + + #[test] + fn parse_rejects_oversized_state_len() { + // A single entry whose declared StateLen (~4GiB via the 0xFF escape) far + // exceeds the bytes actually present must error, not attempt the + // allocation. + let mut body = Vec::new(); + body.extend_from_slice(&1u32.to_le_bytes()); // SeqNo + body.push(0x00); // Status + body.push(0x01); // StateId + body.push(0xFF); // long-length escape + body.extend_from_slice(&0xFFFF_FFF0u32.to_le_bytes()); // StateLen ~4GiB + // ...but no value bytes follow. + + let err = TokenSessionState::parse(body).expect_err("oversized StateLen must be rejected"); + assert!(matches!(err, Error::Protocol(_))); + } + + #[test] + fn parse_rejects_state_len_exceeding_remaining() { + // StateLen (5) is larger than the bytes actually remaining after the + // header (3), so it must be rejected as a protocol error. + let mut body = Vec::new(); + body.extend_from_slice(&1u32.to_le_bytes()); // SeqNo + body.push(0x00); // Status + body.push(0x00); // StateId + body.push(0x05); // StateLen = 5 + body.extend_from_slice(&[0xAA, 0xBB, 0xCC]); // only 3 value bytes present + + let err = TokenSessionState::parse(body) + .expect_err("StateLen exceeding remaining bytes must be rejected"); + assert!(matches!(err, Error::Protocol(_))); + } + + #[tokio::test] + async fn decode_accepts_minimum_and_larger_lengths() { + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use bytes::{BufMut, BytesMut}; + + // len == 5: exactly SeqNo + Status, no states. This boundary must decode. + let mut buf = BytesMut::new(); + buf.put_u32_le(5); + buf.put_u32_le(1); // SeqNo + buf.put_u8(0x01); // Status + + let token = TokenSessionState::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + assert_eq!(token.seq_no, 1); + assert_eq!(token.status, 0x01); + assert!(token.states.is_empty()); + + // len == 8: a full token with one state value. Must decode fine. + let mut buf = BytesMut::new(); + buf.put_u32_le(8); + buf.put_u32_le(2); // SeqNo + buf.put_u8(0x00); // Status + buf.put_u8(0x07); // StateId + buf.put_u8(0x01); // StateLen = 1 + buf.put_u8(0x42); // value + + let token = TokenSessionState::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + assert_eq!(token.seq_no, 2); + assert_eq!(token.states.len(), 1); + assert_eq!(token.states[0].id, 7); + assert_eq!(token.states[0].value, vec![0x42]); + } + + #[tokio::test] + async fn decode_length_boundary_against_max_token_body() { + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use bytes::{BufMut, BytesMut}; + + // len == MAX_TOKEN_BODY + 1: over the cap, so a protocol error is + // returned immediately. + let mut buf = BytesMut::new(); + buf.put_u32_le((super::super::MAX_TOKEN_BODY + 1) as u32); + buf.put_u32_le(0); // a few bytes so the read gets that far + + let err = TokenSessionState::decode(&mut buf.into_sql_read_bytes()) + .await + .expect_err("length over MAX_TOKEN_BODY must be a protocol error"); + assert!(matches!(err, Error::Protocol(_))); + + // len == MAX_TOKEN_BODY exactly: at the boundary the length check must + // NOT fire. The buffer is truncated, so decoding fails with an I/O error. + let mut buf = BytesMut::new(); + buf.put_u32_le(super::super::MAX_TOKEN_BODY as u32); + buf.put_u32_le(0); // far fewer than MAX bytes follow + + let err = TokenSessionState::decode(&mut buf.into_sql_read_bytes()) + .await + .expect_err("truncated body must fail after the length check"); + assert!(matches!(err, Error::Io { .. })); + } +} diff --git a/src/tds/codec/token/token_sspi.rs b/src/tds/codec/token/token_sspi.rs index ccb078bf2..45a756f08 100644 --- a/src/tds/codec/token/token_sspi.rs +++ b/src/tds/codec/token/token_sspi.rs @@ -12,7 +12,11 @@ impl AsRef<[u8]> for TokenSspi { } impl TokenSspi { - #[cfg(any(windows, feature = "winauth", all(unix, feature = "integrated-auth-gssapi")))] + #[cfg(any( + windows, + feature = "winauth", + all(unix, any(feature = "integrated-auth-gssapi", feature = "sspi-rs")) + ))] pub fn new(bytes: Vec) -> Self { Self(bytes) } @@ -21,6 +25,8 @@ impl TokenSspi { where R: SqlReadBytes + Unpin, { + // `len` is bounded by the u16 length field (<= 64 KiB), so no named + // allocation cap is required here. let len = src.read_u16_le().await? as usize; let mut bytes = vec![0; len]; src.read_exact(&mut bytes[0..len]).await?; @@ -35,3 +41,39 @@ impl Encode for TokenSspi { Ok(()) } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use crate::Error; + use bytes::{BufMut, BytesMut}; + + #[tokio::test] + async fn decode_reads_declared_bytes() { + let mut buf = BytesMut::new(); + buf.put_u16_le(4); // declared length + buf.put_slice(&[0xDE, 0xAD, 0xBE, 0xEF]); + + let token = TokenSspi::decode_async(&mut buf.into_sql_read_bytes()) + .await + .expect("must decode"); + + assert_eq!(token.as_ref(), &[0xDE, 0xAD, 0xBE, 0xEF]); + } + + #[tokio::test] + async fn decode_truncated_body_is_clean_error() { + // Declared length is 8 but only 2 bytes follow: the `read_exact` must + // surface a clean IO/EOF error rather than panic. + let mut buf = BytesMut::new(); + buf.put_u16_le(8); // declared length + buf.put_slice(&[0x01, 0x02]); // short body + + let err = TokenSspi::decode_async(&mut buf.into_sql_read_bytes()) + .await + .expect_err("truncated body must error"); + + assert!(matches!(err, Error::Io { .. })); + } +} diff --git a/src/tds/codec/token/token_tab_name.rs b/src/tds/codec/token/token_tab_name.rs new file mode 100644 index 000000000..1b41147d2 --- /dev/null +++ b/src/tds/codec/token/token_tab_name.rs @@ -0,0 +1,259 @@ +use crate::{Error, SqlReadBytes}; + +/// A multi-part table name as sent inside a [`TokenTabName`]. +/// +/// Each name is composed of one or more parts, ordered from the most +/// significant to the least significant, for example +/// `[database].[schema].[table]`. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct TableName { + parts: Vec, +} + +impl TableName { + /// The individual parts of the table name, ordered from the most + /// significant (for example the database name) to the least significant + /// (the table name itself). + #[allow(dead_code)] + pub fn parts(&self) -> &[String] { + &self.parts + } +} + +/// The `TABNAME` token (`0xA4`, MS-TDS §2.2.7.21). +/// +/// Sent by the server to convey the table name(s) that back a result set. It +/// is only produced in browse mode (a `SELECT ... FOR BROWSE` query or a +/// connection with `SET NO_BROWSETABLE ON`) and is used together with the +/// [`ColInfo`](crate::tds::codec::TokenType::ColInfo) token, whose entries +/// reference tables by their one-based index in this token. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct TokenTabName { + tables: Vec, +} + +impl TokenTabName { + /// The table names carried by this token, in the order the server sent + /// them. `ColInfo` table indexes are one-based positions into this slice. + #[allow(dead_code)] + pub fn tables(&self) -> &[TableName] { + &self.tables + } + + pub(crate) async fn decode(src: &mut R) -> crate::Result + where + R: SqlReadBytes + Unpin, + { + // `Length` is the number of bytes of token data that follow. Read the + // whole payload up front and parse it in memory so that the exact + // number of bytes is always consumed, regardless of how many table + // names are packed into the token. + // `len` is bounded by the u16 length field (<= 64 KiB), so no named + // allocation cap is required here. + let len = src.read_u16_le().await? as usize; + + // Bulk-read the payload in one packet-aware pass instead of `len` + // separate `read_u8().await` calls. `len` is u16-bounded (<= 64 KiB); + // cap the up-front reservation like the sibling token decoders. + let mut data = Vec::new(); + crate::sql_read_bytes::read_bytes_into( + src, + &mut data, + len, + crate::tds::codec::column_data::MAX_PREALLOC, + ) + .await?; + + Self::parse(&data) + } + + /// Parse the `TABNAME` token payload (the bytes following the `Length` + /// field). + /// + /// Each table name is encoded as a `NumParts` byte followed by that many + /// `US_VARCHAR` parts (a `USHORT` UTF-16 code-unit count followed by the + /// UTF-16LE characters). + fn parse(data: &[u8]) -> crate::Result { + let mut tables = Vec::new(); + let mut pos = 0; + + while pos < data.len() { + let num_parts = data[pos]; + pos += 1; + + let mut parts = Vec::with_capacity(num_parts as usize); + + for _ in 0..num_parts { + if pos + 2 > data.len() { + return Err(Error::Protocol( + "TABNAME token truncated while reading part length".into(), + )); + } + + let char_count = u16::from_le_bytes([data[pos], data[pos + 1]]) as usize; + pos += 2; + + let byte_count = char_count * 2; + + if pos + byte_count > data.len() { + return Err(Error::Protocol( + "TABNAME token truncated while reading part name".into(), + )); + } + + let mut units = Vec::with_capacity(char_count); + for _ in 0..char_count { + units.push(u16::from_le_bytes([data[pos], data[pos + 1]])); + pos += 2; + } + + let part = String::from_utf16(&units).map_err(|_| { + Error::Protocol("TABNAME token part is not valid UTF-16".into()) + })?; + + parts.push(part); + } + + tables.push(TableName { parts }); + } + + Ok(TokenTabName { tables }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn us_varchar(s: &str) -> Vec { + let units: Vec = s.encode_utf16().collect(); + let mut out = Vec::new(); + out.extend_from_slice(&(units.len() as u16).to_le_bytes()); + for u in units { + out.extend_from_slice(&u.to_le_bytes()); + } + out + } + + #[test] + fn parse_single_multipart_table() { + let mut data = Vec::new(); + // NumParts = 3 + data.push(3u8); + data.extend_from_slice(&us_varchar("mydb")); + data.extend_from_slice(&us_varchar("dbo")); + data.extend_from_slice(&us_varchar("Customers")); + + let token = TokenTabName::parse(&data).expect("must parse"); + + assert_eq!(token.tables().len(), 1); + assert_eq!( + token.tables()[0].parts(), + &[ + "mydb".to_string(), + "dbo".to_string(), + "Customers".to_string() + ] + ); + } + + #[test] + fn parse_multiple_tables() { + let mut data = Vec::new(); + // First table: single part. + data.push(1u8); + data.extend_from_slice(&us_varchar("Orders")); + // Second table: two parts. + data.push(2u8); + data.extend_from_slice(&us_varchar("dbo")); + data.extend_from_slice(&us_varchar("Products")); + + let token = TokenTabName::parse(&data).expect("must parse"); + + assert_eq!(token.tables().len(), 2); + assert_eq!(token.tables()[0].parts(), &["Orders".to_string()]); + assert_eq!( + token.tables()[1].parts(), + &["dbo".to_string(), "Products".to_string()] + ); + } + + #[test] + fn parse_empty_payload() { + let token = TokenTabName::parse(&[]).expect("must parse"); + assert!(token.tables().is_empty()); + } + + #[test] + fn parse_non_ascii_name_reads_both_bytes_of_each_unit() { + // A code point with a non-zero high byte (U+20AC EURO SIGN => 0xAC 0x20) + // only decodes correctly if both bytes of the UTF-16 unit are read; a + // one-off in the low/high byte index would corrupt it. + let mut data = vec![1u8]; + data.extend_from_slice(&us_varchar("€uro")); + + let token = TokenTabName::parse(&data).expect("must parse"); + assert_eq!(token.tables()[0].parts(), &["€uro".to_string()]); + } + + #[test] + fn parse_zero_length_part_at_buffer_end() { + // NumParts = 1 followed by a zero-length part that ends exactly at the + // buffer boundary: the `pos + 2 > len` check must accept (not reject) an + // exact fit. + let data = vec![1u8, 0u8, 0u8]; + let token = TokenTabName::parse(&data).expect("exact-fit length must parse"); + assert_eq!(token.tables()[0].parts(), &[String::new()]); + } + + #[test] + fn parse_rejects_name_length_exceeding_payload() { + // NumParts = 1, part claims 4 code units (8 bytes) but only 6 follow. + // The `char_count * 2` byte check must reject this; a wrong multiplier + // would under-count and read past the buffer. + let mut data = vec![1u8]; + data.extend_from_slice(&4u16.to_le_bytes()); + data.extend_from_slice(&[0xAB; 6]); + assert!(TokenTabName::parse(&data).is_err()); + } + + #[test] + fn parse_truncated_length_fails() { + // NumParts says 1 part but no length bytes follow. + let data = vec![1u8]; + assert!(TokenTabName::parse(&data).is_err()); + } + + #[test] + fn parse_truncated_name_fails() { + // NumParts = 1, claims a 4-code-unit name but provides no bytes. + let mut data = vec![1u8]; + data.extend_from_slice(&4u16.to_le_bytes()); + assert!(TokenTabName::parse(&data).is_err()); + } + + #[tokio::test] + async fn decode_reads_length_prefixed_payload() { + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + use bytes::BytesMut; + + let mut payload = Vec::new(); + payload.push(2u8); + payload.extend_from_slice(&us_varchar("dbo")); + payload.extend_from_slice(&us_varchar("Invoices")); + + let mut wire = BytesMut::new(); + wire.extend_from_slice(&(payload.len() as u16).to_le_bytes()); + wire.extend_from_slice(&payload); + + let token = TokenTabName::decode(&mut wire.into_sql_read_bytes()) + .await + .expect("decode must succeed"); + + assert_eq!(token.tables().len(), 1); + assert_eq!( + token.tables()[0].parts(), + &["dbo".to_string(), "Invoices".to_string()] + ); + } +} diff --git a/src/tds/codec/token/token_type.rs b/src/tds/codec/token/token_type.rs index e4b516bc0..959281d04 100644 --- a/src/tds/codec/token/token_type.rs +++ b/src/tds/codec/token/token_type.rs @@ -9,7 +9,12 @@ uint_enum! { /// the server. ReturnStatus = 0x79, - /// Describes the result setfor interpretation of following ROW data + /// Describes the data type, length, and name of column data that + /// result from a COMPUTE clause (`ALTMETADATA`). This token describes + /// the format of the following `ALTROW` data streams. + AltMetaData = 0x88, + + /// Describes the result set for interpretation of following ROW data /// streams ColMetaData = 0x81, @@ -25,7 +30,12 @@ uint_enum! { /// Describes the column information in browse mode. ColInfo = 0xA5, - /// Used to send the return value of an RPCto the client. When an RPC is + /// Used to send the table name to the client in browse mode (for + /// example `SELECT ... FOR BROWSE`). Paired with the COLINFO token, + /// whose entries reference the tables carried here by index. + TabName = 0xA4, + + /// Used to send the return value of an RPC to the client. When an RPC is /// executed, the associated parameters may be defined as input or /// output (or "return") parameters. /// @@ -45,9 +55,25 @@ uint_enum! { /// COLMETADATA token. NbcRow = 0xD2, + /// Used to send a complete row of computed data, as defined by the + /// `ALTMETADATA` token, to the client. This is the row produced by a + /// COMPUTE or COMPUTE BY clause. + AltRow = 0xD3, + /// The SSPI token returned during the login process. Sspi = 0xED, + /// Used to inform the client about the current session state so the + /// session can be transparently recovered after a broken connection + /// (connection resiliency). Sent only when session recovery is enabled. + SessionState = 0xE4, + + /// Carries the information the client needs to acquire a federated + /// authentication (Azure Active Directory) access token, such as the + /// Security Token Service URL and the Service Principal Name. Sent by + /// the server during a library-driven federated authentication flow. + FedAuthInfo = 0xEE, + /// A notification of an environment change (such as database and /// language). EnvChange = 0xE3, @@ -80,3 +106,50 @@ uint_enum! { FeatureExtAck = 0xAE, } } + +#[cfg(test)] +mod tests { + use super::TokenType; + use std::convert::TryFrom; + + #[test] + fn known_values_map_to_variants() { + let cases: &[(u8, TokenType)] = &[ + (0x79, TokenType::ReturnStatus), + (0x88, TokenType::AltMetaData), + (0x81, TokenType::ColMetaData), + (0xAA, TokenType::Error), + (0xAB, TokenType::Info), + (0xA9, TokenType::Order), + (0xA5, TokenType::ColInfo), + (0xA4, TokenType::TabName), + (0xAC, TokenType::ReturnValue), + (0xAD, TokenType::LoginAck), + (0xD1, TokenType::Row), + (0xD2, TokenType::NbcRow), + (0xD3, TokenType::AltRow), + (0xED, TokenType::Sspi), + (0xE4, TokenType::SessionState), + (0xEE, TokenType::FedAuthInfo), + (0xE3, TokenType::EnvChange), + (0xFD, TokenType::Done), + (0xFE, TokenType::DoneProc), + (0xFF, TokenType::DoneInProc), + (0xAE, TokenType::FeatureExtAck), + ]; + + for &(byte, variant) in cases { + // Byte -> variant. + assert_eq!(TokenType::try_from(byte).unwrap(), variant); + // Round-trip: variant -> byte -> variant. + assert_eq!(variant as u8, byte); + assert_eq!(TokenType::try_from(variant as u8).unwrap(), variant); + } + } + + #[test] + fn unknown_byte_is_rejected() { + // 0x00 is not a defined token type. + assert!(TokenType::try_from(0x00u8).is_err()); + } +} diff --git a/src/tds/codec/transaction_manager.rs b/src/tds/codec/transaction_manager.rs new file mode 100644 index 000000000..49d842f10 --- /dev/null +++ b/src/tds/codec/transaction_manager.rs @@ -0,0 +1,294 @@ +use super::{encode_all_headers_tx, encode_b_varchar, Encode}; +use bytes::{BufMut, BytesMut}; +use std::borrow::Cow; + +uint_enum! { + /// The request type of a Transaction Manager request, as defined in + /// MS-TDS 2.2.6.8 (`TM_*` request kinds). + #[repr(u16)] + pub enum TransactionManagerRequestType { + /// Get the address of the Distributed Transaction Coordinator. + GetDtcAddress = 0, + /// Import an existing distributed transaction (propagate). + Propagate = 1, + /// Begin a new transaction (`TM_BEGIN_XACT`). + Begin = 5, + /// Promote a local transaction to a distributed one (`TM_PROMOTE_XACT`). + Promote = 6, + /// Commit the active transaction (`TM_COMMIT_XACT`). + Commit = 7, + /// Roll back the active transaction (`TM_ROLLBACK_XACT`). + Rollback = 8, + /// Create a savepoint in the active transaction (`TM_SAVE_XACT`). + Save = 9, + } +} + +uint_enum! { + /// The transaction isolation level requested when beginning a transaction + /// through a Transaction Manager request (MS-TDS 2.2.6.8). + #[repr(u8)] + pub enum IsolationLevel { + /// Use the server's default isolation level. + Unspecified = 0x00, + /// `READ UNCOMMITTED`. + ReadUncommitted = 0x01, + /// `READ COMMITTED`. + ReadCommitted = 0x02, + /// `REPEATABLE READ`. + RepeatableRead = 0x03, + /// `SERIALIZABLE`. + Serializable = 0x04, + /// `SNAPSHOT`. + Snapshot = 0x05, + } +} + +/// A Transaction Manager request (packet type `0x14`, MS-TDS 2.2.6.8). +/// +/// These requests let the client begin, commit, roll back or create a +/// savepoint in a transaction directly through the TDS protocol instead of +/// issuing the equivalent T-SQL batch (`BEGIN TRAN`, `COMMIT`, ...). +/// +/// Every request carries the current transaction descriptor in the request's +/// `ALL_HEADERS` block so the server can associate it with the correct +/// transaction. +#[derive(Debug, Clone)] +pub struct TransactionManagerRequest<'a> { + transaction_desc: [u8; 8], + body: TransactionRequestBody<'a>, +} + +#[derive(Debug, Clone)] +enum TransactionRequestBody<'a> { + Begin { + isolation_level: IsolationLevel, + name: Cow<'a, str>, + }, + Commit { + name: Cow<'a, str>, + }, + Rollback { + name: Cow<'a, str>, + }, + Save { + name: Cow<'a, str>, + }, +} + +impl<'a> TransactionManagerRequest<'a> { + /// Build a `TM_BEGIN_XACT` request that begins a new transaction with the + /// given isolation level. The (usually empty) transaction name is sent as + /// a `B_VARCHAR`. + pub fn begin( + transaction_desc: [u8; 8], + isolation_level: IsolationLevel, + name: impl Into>, + ) -> Self { + Self { + transaction_desc, + body: TransactionRequestBody::Begin { + isolation_level, + name: name.into(), + }, + } + } + + /// Build a `TM_COMMIT_XACT` request that commits the active transaction. + pub fn commit(transaction_desc: [u8; 8], name: impl Into>) -> Self { + Self { + transaction_desc, + body: TransactionRequestBody::Commit { name: name.into() }, + } + } + + /// Build a `TM_ROLLBACK_XACT` request that rolls back the active + /// transaction (or to a savepoint of the given name). + pub fn rollback(transaction_desc: [u8; 8], name: impl Into>) -> Self { + Self { + transaction_desc, + body: TransactionRequestBody::Rollback { name: name.into() }, + } + } + + /// Build a `TM_SAVE_XACT` request that creates a savepoint with the given + /// name in the active transaction. + pub fn save(transaction_desc: [u8; 8], name: impl Into>) -> Self { + Self { + transaction_desc, + body: TransactionRequestBody::Save { name: name.into() }, + } + } + + fn request_type(&self) -> TransactionManagerRequestType { + match self.body { + TransactionRequestBody::Begin { .. } => TransactionManagerRequestType::Begin, + TransactionRequestBody::Commit { .. } => TransactionManagerRequestType::Commit, + TransactionRequestBody::Rollback { .. } => TransactionManagerRequestType::Rollback, + TransactionRequestBody::Save { .. } => TransactionManagerRequestType::Save, + } + } +} + +impl<'a> Encode for TransactionManagerRequest<'a> { + fn encode(self, dst: &mut BytesMut) -> crate::Result<()> { + // ALL_HEADERS block carrying the transaction descriptor. + encode_all_headers_tx(dst, self.transaction_desc); + + // Request type (USHORT). + dst.put_u16_le(self.request_type() as u16); + + match self.body { + TransactionRequestBody::Begin { + isolation_level, + name, + } => { + dst.put_u8(isolation_level as u8); + encode_b_varchar(dst, &name)?; + } + TransactionRequestBody::Commit { name } | TransactionRequestBody::Rollback { name } => { + encode_b_varchar(dst, &name)?; + // Flags byte: bit 0 (`fBeginXact`) unset — do not begin a new + // transaction after commit/rollback. + dst.put_u8(0); + } + TransactionRequestBody::Save { name } => { + encode_b_varchar(dst, &name)?; + } + } + + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use super::super::{AllHeaderTy, ALL_HEADERS_LEN_TX}; + use super::*; + + fn all_headers() -> Vec { + let mut v = Vec::new(); + v.extend_from_slice(&(ALL_HEADERS_LEN_TX as u32).to_le_bytes()); + v.extend_from_slice(&(ALL_HEADERS_LEN_TX as u32 - 4).to_le_bytes()); + v.extend_from_slice(&(AllHeaderTy::TransactionDescriptor as u16).to_le_bytes()); + v.extend_from_slice(&[1, 2, 3, 4, 5, 6, 7, 8]); + v.extend_from_slice(&1u32.to_le_bytes()); + v + } + + #[test] + fn encodes_begin_request() { + let desc = [1, 2, 3, 4, 5, 6, 7, 8]; + let req = TransactionManagerRequest::begin(desc, IsolationLevel::ReadCommitted, ""); + + let mut buf = BytesMut::new(); + req.encode(&mut buf).unwrap(); + + let mut expected = all_headers(); + expected.extend_from_slice(&(TransactionManagerRequestType::Begin as u16).to_le_bytes()); + expected.push(IsolationLevel::ReadCommitted as u8); // isolation level + expected.push(0); // B_VARCHAR length (empty name) + + assert_eq!(&buf[..], &expected[..]); + } + + #[test] + fn encodes_begin_request_with_name() { + let desc = [1, 2, 3, 4, 5, 6, 7, 8]; + let req = TransactionManagerRequest::begin(desc, IsolationLevel::Serializable, "tx"); + + let mut buf = BytesMut::new(); + req.encode(&mut buf).unwrap(); + + let mut expected = all_headers(); + expected.extend_from_slice(&(TransactionManagerRequestType::Begin as u16).to_le_bytes()); + expected.push(IsolationLevel::Serializable as u8); + expected.push(2); // two UTF-16 code units + expected.extend_from_slice(&b't'.to_le_bytes()); + expected.push(0); + expected.extend_from_slice(&b'x'.to_le_bytes()); + expected.push(0); + + assert_eq!(&buf[..], &expected[..]); + } + + #[test] + fn encodes_commit_request() { + let desc = [1, 2, 3, 4, 5, 6, 7, 8]; + let req = TransactionManagerRequest::commit(desc, ""); + + let mut buf = BytesMut::new(); + req.encode(&mut buf).unwrap(); + + let mut expected = all_headers(); + expected.extend_from_slice(&(TransactionManagerRequestType::Commit as u16).to_le_bytes()); + expected.push(0); // B_VARCHAR length (empty name) + expected.push(0); // flags: no new transaction + + assert_eq!(&buf[..], &expected[..]); + } + + #[test] + fn encodes_rollback_request() { + let desc = [8, 7, 6, 5, 4, 3, 2, 1]; + let req = TransactionManagerRequest::rollback(desc, ""); + + let mut buf = BytesMut::new(); + req.encode(&mut buf).unwrap(); + + let mut expected = Vec::new(); + expected.extend_from_slice(&(ALL_HEADERS_LEN_TX as u32).to_le_bytes()); + expected.extend_from_slice(&(ALL_HEADERS_LEN_TX as u32 - 4).to_le_bytes()); + expected.extend_from_slice(&(AllHeaderTy::TransactionDescriptor as u16).to_le_bytes()); + expected.extend_from_slice(&desc); + expected.extend_from_slice(&1u32.to_le_bytes()); + expected.extend_from_slice(&(TransactionManagerRequestType::Rollback as u16).to_le_bytes()); + expected.push(0); // B_VARCHAR length + expected.push(0); // flags + + assert_eq!(&buf[..], &expected[..]); + } + + #[test] + fn encodes_save_request() { + let desc = [1, 2, 3, 4, 5, 6, 7, 8]; + let req = TransactionManagerRequest::save(desc, "sp1"); + + let mut buf = BytesMut::new(); + req.encode(&mut buf).unwrap(); + + let mut expected = all_headers(); + expected.extend_from_slice(&(TransactionManagerRequestType::Save as u16).to_le_bytes()); + expected.push(3); // three UTF-16 code units + for unit in "sp1".encode_utf16() { + expected.extend_from_slice(&unit.to_le_bytes()); + } + + assert_eq!(&buf[..], &expected[..]); + } + + // A B_VARCHAR length prefix is a single byte; a name longer than 255 UTF-16 + // code units must error rather than truncate the count with `as u8` (which + // would keep writing every unit and corrupt the stream). + #[test] + fn encode_rejects_over_long_save_name() { + let desc = [1, 2, 3, 4, 5, 6, 7, 8]; + let long = "a".repeat(256); + let req = TransactionManagerRequest::save(desc, long); + + let mut buf = BytesMut::new(); + let err = req.encode(&mut buf).unwrap_err(); + assert!(matches!(err, crate::Error::Protocol(_)), "got {err:?}"); + } + + // A name of exactly 255 units is still allowed (boundary of the guard). + #[test] + fn encode_accepts_max_length_save_name() { + let desc = [1, 2, 3, 4, 5, 6, 7, 8]; + let name = "a".repeat(255); + let req = TransactionManagerRequest::save(desc, name); + + let mut buf = BytesMut::new(); + req.encode(&mut buf).expect("255-unit name must encode"); + } +} diff --git a/src/tds/codec/type_info.rs b/src/tds/codec/type_info.rs index 20647d70a..3026a4c1f 100644 --- a/src/tds/codec/type_info.rs +++ b/src/tds/codec/type_info.rs @@ -2,9 +2,9 @@ use asynchronous_codec::BytesMut; use bytes::BufMut; use crate::{tds::Collation, xml::XmlSchema, Error, SqlReadBytes}; -use std::{convert::TryFrom, sync::Arc, usize}; +use std::{convert::TryFrom, sync::Arc}; -use super::Encode; +use super::{encode_b_varchar, Encode}; /// A length of a column in bytes or characters. #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -18,20 +18,58 @@ pub enum TypeLength { /// Describes a type of a column. #[derive(Debug, Clone, PartialEq, Eq)] pub enum TypeInfo { + /// A fixed-length type, whose size is fully determined by the type itself. FixedLen(FixedLenType), + /// A variable-length type with an explicit size (and optional collation). VarLenSized(VarLenContext), + /// A variable-length type carrying a precision and scale, such as `decimal` + /// and `numeric`. VarLenSizedPrecision { + /// The underlying variable-length type. ty: VarLenType, + /// The reserved size of the column in bytes. size: usize, + /// The total number of digits. precision: u8, + /// The number of digits to the right of the decimal point. scale: u8, }, + /// The `xml` type, with an optional associated schema. Xml { + /// The XML schema associated with the column, if any. schema: Option>, + /// The reserved size of the column in bytes. size: usize, }, + /// A CLR user-defined type (UDT), MS-TDS §2.2.5.5.4. + Udt(UdtInfo), } +/// Metadata describing a CLR user-defined type (UDT) column, as defined by the +/// `UDT_INFO` rule in MS-TDS §2.2.5.5.4. +/// +/// This carries only the identifying metadata of the type. The value bytes are +/// surfaced verbatim (see [`ColumnData::Binary`]); tiberius does not attempt to +/// deserialize the CLR representation. +/// +/// [`ColumnData::Binary`]: crate::ColumnData::Binary +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct UdtInfo { + /// Maximum size of the UDT value in bytes. A value of `0xFFFF` indicates a + /// large (`MAX`) UDT with no fixed upper bound. + pub max_byte_size: u16, + /// Name of the database in which the UDT is defined. + pub db_name: String, + /// Name of the schema that owns the UDT. + pub schema_name: String, + /// Name of the UDT. + pub type_name: String, + /// Assembly-qualified name of the CLR type that implements the UDT. + pub assembly_qualified_name: String, +} + +/// The context of a variable-length column: its underlying type, size and +/// optional collation. #[derive(Clone, Debug, Copy, PartialEq, Eq)] pub struct VarLenContext { r#type: VarLenType, @@ -40,6 +78,8 @@ pub struct VarLenContext { } impl VarLenContext { + /// Create a new variable-length context from a type, length and optional + /// collation. pub fn new(r#type: VarLenType, len: usize, collation: Option) -> Self { Self { r#type, @@ -58,6 +98,11 @@ impl VarLenContext { self.len } + /// `true` if the column reserves no length. + pub fn is_empty(&self) -> bool { + self.len == 0 + } + /// Get the var len context's collation. pub fn collation(&self) -> Option { self.collation @@ -70,11 +115,16 @@ impl Encode for VarLenContext { // length match self.r#type { + // DATE (0x28) carries NO scale byte in TYPE_INFO (MS-TDS + // §2.2.5.4.2 / §2.2.5.5.1.2), unlike TIME/DATETIME2/DATETIMEOFFSET + // which each carry a SCALE byte. The decoder already special-cases + // this (`Daten => 3`, reading no byte); emitting a byte here would + // desync every field after a `date` column in a TYPE_INFO stream + // (bulk-load column metadata / TVP). #[cfg(feature = "tds73")] - VarLenType::Daten - | VarLenType::Timen - | VarLenType::DatetimeOffsetn - | VarLenType::Datetime2 => { + VarLenType::Daten => {} + #[cfg(feature = "tds73")] + VarLenType::Timen | VarLenType::DatetimeOffsetn | VarLenType::Datetime2 => { dst.put_u8(self.len() as u8); } VarLenType::Bitn @@ -95,11 +145,15 @@ impl Encode for VarLenContext { | VarLenType::BigVarBin => { dst.put_u16_le(self.len() as u16); } - VarLenType::Image | VarLenType::Text | VarLenType::NText => { + VarLenType::Image | VarLenType::Text | VarLenType::NText | VarLenType::SSVariant => { dst.put_u32_le(self.len() as u32); } VarLenType::Xml => (), - typ => todo!("encoding {:?} is not supported yet", typ), + typ => { + return Err(Error::Protocol( + format!("encoding a {typ:?} var-len context is not supported").into(), + )) + } } if let Some(collation) = self.collation() { @@ -149,12 +203,12 @@ uint_enum! { NVarchar = 0xE7, NChar = 0xEF, Xml = 0xF1, - // not supported yet + // CLR user-defined type; decoded as raw PLP bytes (see column_data/udt.rs). Udt = 0xF0, Text = 0x23, Image = 0x22, NText = 0x63, - // not supported yet + // sql_variant; fully decoded/encoded (see column_data/sql_variant.rs). SSVariant = 0x62, // legacy types (not supported since post-7.2): // Char = 0x2F, // Binary = 0x2D, @@ -189,12 +243,12 @@ uint_enum! { NVarchar = 0xE7, NChar = 0xEF, Xml = 0xF1, - // not supported yet + // CLR user-defined type; decoded as raw PLP bytes (see column_data/udt.rs). Udt = 0xF0, Text = 0x23, Image = 0x22, NText = 0x63, - // not supported yet + // sql_variant; fully decoded/encoded (see column_data/sql_variant.rs). SSVariant = 0x62, // legacy types (not supported since post-7.2): // Char = 0x2F, // Binary = 0x2D, @@ -229,33 +283,65 @@ impl Encode for TypeInfo { if let Some(xs) = schema { dst.put_u8(1); - let db_name_encoded: Vec = xs.db_name().encode_utf16().collect(); - dst.put_u8(db_name_encoded.len() as u8); - for chr in db_name_encoded { - dst.put_u16_le(chr); - } - - let owner_encoded: Vec = xs.owner().encode_utf16().collect(); - dst.put_u8(owner_encoded.len() as u8); - for chr in owner_encoded { - dst.put_u16_le(chr); - } - - let collection_encoded: Vec = xs.collection().encode_utf16().collect(); - dst.put_u16_le(collection_encoded.len() as u16); - for chr in collection_encoded { - dst.put_u16_le(chr); - } + // db_name and owner are B_VARCHARs (u8 code-unit count); the + // collection name is a US_VARCHAR (u16). Each length prefix + // must bound-check rather than truncate with `as`, which would + // write the full payload behind a wrong length and desync the + // wire. + encode_b_varchar(dst, xs.db_name())?; + encode_b_varchar(dst, xs.owner())?; + encode_us_varchar(dst, xs.collection())?; } else { dst.put_u8(0); } } + TypeInfo::Udt(info) => { + dst.put_u8(VarLenType::Udt as u8); + dst.put_u16_le(info.max_byte_size); + + // db_name, schema_name and type_name are B_VARCHARs (u8 count); + // the assembly-qualified name is a US_VARCHAR (u16). Bound-check + // each length prefix rather than truncating with `as`. + encode_b_varchar(dst, &info.db_name)?; + encode_b_varchar(dst, &info.schema_name)?; + encode_b_varchar(dst, &info.type_name)?; + encode_us_varchar(dst, &info.assembly_qualified_name)?; + } } Ok(()) } } +/// Encodes a `US_VARCHAR` (MS-TDS 2.2.5.1.4): a `u16` code-unit count followed +/// by that many little-endian UTF-16 code units. +/// +/// The length prefix is a `u16`, so a string longer than 65535 UTF-16 code +/// units cannot be represented. Truncating the count with `as u16` while still +/// writing every unit would desync the wire, so an over-long string is rejected +/// with an [`Error::Protocol`] instead. +fn encode_us_varchar(dst: &mut BytesMut, s: &str) -> crate::Result<()> { + let units: Vec = s.encode_utf16().collect(); + + if units.len() > u16::MAX as usize { + return Err(Error::Protocol( + format!( + "string is too long for a US_VARCHAR ({} UTF-16 code units, max 65535)", + units.len() + ) + .into(), + )); + } + + dst.put_u16_le(units.len() as u16); + + for unit in units { + dst.put_u16_le(unit); + } + + Ok(()) +} + impl TypeInfo { pub(crate) async fn decode(src: &mut R) -> crate::Result where @@ -289,6 +375,22 @@ impl TypeInfo { size: 0xfffffffffffffffe_usize, }) } + Ok(VarLenType::Udt) => { + // UDT_INFO, MS-TDS §2.2.5.5.4 + let max_byte_size = src.read_u16_le().await?; + let db_name = src.read_b_varchar().await?; + let schema_name = src.read_b_varchar().await?; + let type_name = src.read_b_varchar().await?; + let assembly_qualified_name = src.read_us_varchar().await?; + + Ok(TypeInfo::Udt(UdtInfo { + max_byte_size, + db_name, + schema_name, + type_name, + assembly_qualified_name, + })) + } Ok(ty) => { let len = match ty { #[cfg(feature = "tds73")] @@ -311,10 +413,15 @@ impl TypeInfo { | VarLenType::BigVarChar | VarLenType::BigBinary | VarLenType::BigVarBin => src.read_u16_le().await? as usize, - VarLenType::Image | VarLenType::Text | VarLenType::NText => { - src.read_u32_le().await? as usize + VarLenType::Image + | VarLenType::Text + | VarLenType::NText + | VarLenType::SSVariant => src.read_u32_le().await? as usize, + _ => { + return Err(Error::Protocol( + format!("unsupported column type in COLMETADATA: {:?}", ty).into(), + )) } - _ => todo!("not yet implemented for {:?}", ty), }; let collation = match ty { @@ -337,6 +444,18 @@ impl TypeInfo { let precision = src.read_u8().await?; let scale = src.read_u8().await?; + // MS-TDS: precision is 1..=38 and scale 0..=precision. + // Reject out-of-range server values here so downstream + // (Numeric decode/Display) never sees an impossible scale. + if precision > 38 || scale > precision { + return Err(Error::Protocol( + format!( + "decimal/numeric: invalid precision {precision} / scale {scale}" + ) + .into(), + )); + } + TypeInfo::VarLenSizedPrecision { size: len, ty, @@ -380,6 +499,14 @@ mod tests { 40, Some(Collation::new(13632521, 52)), )), + TypeInfo::Udt(UdtInfo { + max_byte_size: 0xffff, + db_name: "fake-db".to_string(), + schema_name: "dbo".to_string(), + type_name: "geometry".to_string(), + assembly_qualified_name: + "Microsoft.SqlServer.Types.SqlGeometry, Microsoft.SqlServer.Types".to_string(), + }), ]; for ti in types { @@ -396,4 +523,285 @@ mod tests { assert_eq!(nti, ti) } } + + // The XML schema db_name/owner length prefixes are B_VARCHARs (u8). The + // length prefix and the payload must agree: encode writes an exact u8 count + // then that many UTF-16 code units, and the collection name a u16 count. + #[test] + fn xml_schema_encodes_byte_exact_length_prefixes() { + let ti = TypeInfo::Xml { + schema: Some(XmlSchema::new("ab", "c", "de").into()), + size: 0, + }; + + let mut buf = BytesMut::new(); + ti.encode(&mut buf).unwrap(); + + let mut expected = Vec::new(); + expected.push(VarLenType::Xml as u8); + expected.push(1u8); // has_schema + expected.push(2u8); // db_name B_VARCHAR length + expected.extend_from_slice(&('a' as u16).to_le_bytes()); + expected.extend_from_slice(&('b' as u16).to_le_bytes()); + expected.push(1u8); // owner B_VARCHAR length + expected.extend_from_slice(&('c' as u16).to_le_bytes()); + expected.extend_from_slice(&2u16.to_le_bytes()); // collection US_VARCHAR length + expected.extend_from_slice(&('d' as u16).to_le_bytes()); + expected.extend_from_slice(&('e' as u16).to_le_bytes()); + + assert_eq!(&buf[..], &expected[..]); + } + + // db_name/owner are B_VARCHARs (u8); a name over 255 code units must error + // rather than truncate the length prefix while writing the full payload. + #[test] + fn xml_schema_rejects_over_long_db_name() { + let ti = TypeInfo::Xml { + schema: Some(XmlSchema::new("a".repeat(256), "owner", "coll").into()), + size: 0, + }; + + let mut buf = BytesMut::new(); + let err = ti.encode(&mut buf).unwrap_err(); + assert!(matches!(err, Error::Protocol(_)), "got {err:?}"); + } + + // The collection name is a US_VARCHAR (u16); over 65535 code units must error. + #[test] + fn xml_schema_rejects_over_long_collection() { + let ti = TypeInfo::Xml { + schema: Some(XmlSchema::new("db", "owner", "a".repeat(u16::MAX as usize + 1)).into()), + size: 0, + }; + + let mut buf = BytesMut::new(); + let err = ti.encode(&mut buf).unwrap_err(); + assert!(matches!(err, Error::Protocol(_)), "got {err:?}"); + } + + // The UDT db_name/schema_name/type_name are B_VARCHARs (u8); an over-long one + // must error rather than truncate its length prefix. + #[test] + fn udt_rejects_over_long_type_name() { + let ti = TypeInfo::Udt(UdtInfo { + max_byte_size: 0xffff, + db_name: "db".to_string(), + schema_name: "dbo".to_string(), + type_name: "a".repeat(256), + assembly_qualified_name: "asm".to_string(), + }); + + let mut buf = BytesMut::new(); + let err = ti.encode(&mut buf).unwrap_err(); + assert!(matches!(err, Error::Protocol(_)), "got {err:?}"); + } + + // The UDT assembly-qualified name is a US_VARCHAR (u16); over 65535 code + // units must error. + #[test] + fn udt_rejects_over_long_assembly_qualified_name() { + let ti = TypeInfo::Udt(UdtInfo { + max_byte_size: 0xffff, + db_name: "db".to_string(), + schema_name: "dbo".to_string(), + type_name: "T".to_string(), + assembly_qualified_name: "a".repeat(u16::MAX as usize + 1), + }); + + let mut buf = BytesMut::new(); + let err = ti.encode(&mut buf).unwrap_err(); + assert!(matches!(err, Error::Protocol(_)), "got {err:?}"); + } + + #[cfg(feature = "tds73")] + #[tokio::test] + async fn date_typeinfo_round_trips_without_scale_byte() { + // DATE (0x28) has no scale byte in TYPE_INFO: encode must emit only the + // type token, and it must round-trip through decode. + let ti = TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Daten, 3, None)); + let mut buf = BytesMut::new(); + ti.clone().encode(&mut buf).expect("encode must succeed"); + + assert_eq!(buf.as_ref(), &[VarLenType::Daten as u8]); + + let nti = TypeInfo::decode(&mut buf.into_sql_read_bytes()) + .await + .expect("decode must succeed"); + assert_eq!(nti, ti); + } + + #[tokio::test] + async fn decode_rejects_out_of_range_precision_scale() { + // Decimaln TYPE_INFO: [type][size][precision][scale]. A precision > 38 + // from an untrusted server must be rejected rather than flowing into + // Numeric decoding (which would later panic on an impossible scale). + let mut buf = BytesMut::new(); + buf.put_u8(VarLenType::Decimaln as u8); + buf.put_u8(17); // size + buf.put_u8(200); // precision (invalid, > 38) + buf.put_u8(2); // scale + + let err = TypeInfo::decode(&mut buf.into_sql_read_bytes()) + .await + .expect_err("out-of-range precision must error"); + assert!(matches!(err, Error::Protocol(_))); + } + + #[test] + fn var_len_context_is_empty() { + assert!(VarLenContext::new(VarLenType::Intn, 0, None).is_empty()); + assert!(!VarLenContext::new(VarLenType::Intn, 4, None).is_empty()); + } + + #[tokio::test] + async fn decode_intn_reads_one_byte_length() { + // Covers the Bitn|Intn|Floatn|... match arm: the length is a single u8. + let mut buf = BytesMut::new(); + buf.put_u8(VarLenType::Intn as u8); + buf.put_u8(4); // length in bytes + + let ti = TypeInfo::decode(&mut buf.into_sql_read_bytes()) + .await + .expect("decode must succeed"); + assert_eq!( + ti, + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Intn, 4, None)) + ); + } + + #[cfg(feature = "tds73")] + #[tokio::test] + async fn decode_timen_reads_one_byte_scale() { + // Covers the Timen|DatetimeOffsetn|Datetime2 match arm: reads a u8 scale + // as the length. Deleting the arm would make this an error. + let mut buf = BytesMut::new(); + buf.put_u8(VarLenType::Timen as u8); + buf.put_u8(7); // scale + + let ti = TypeInfo::decode(&mut buf.into_sql_read_bytes()) + .await + .expect("decode must succeed"); + assert_eq!( + ti, + TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Timen, 7, None)) + ); + } + + #[tokio::test] + async fn decode_accepts_precision_38_and_scale_below_precision() { + // Boundary: precision == 38 is the maximum valid precision and must be + // accepted; scale (10) is below precision. + let mut buf = BytesMut::new(); + buf.put_u8(VarLenType::Decimaln as u8); + buf.put_u8(17); // size + buf.put_u8(38); // precision (max valid) + buf.put_u8(10); // scale + + let ti = TypeInfo::decode(&mut buf.into_sql_read_bytes()) + .await + .expect("precision 38 must be accepted"); + assert_eq!( + ti, + TypeInfo::VarLenSizedPrecision { + ty: VarLenType::Decimaln, + size: 17, + precision: 38, + scale: 10, + } + ); + } + + #[tokio::test] + async fn decode_accepts_scale_equal_to_precision() { + // Boundary: scale == precision is valid (scale may equal precision). + let mut buf = BytesMut::new(); + buf.put_u8(VarLenType::Numericn as u8); + buf.put_u8(17); // size + buf.put_u8(20); // precision + buf.put_u8(20); // scale == precision + + let ti = TypeInfo::decode(&mut buf.into_sql_read_bytes()) + .await + .expect("scale == precision must be accepted"); + assert_eq!( + ti, + TypeInfo::VarLenSizedPrecision { + ty: VarLenType::Numericn, + size: 17, + precision: 20, + scale: 20, + } + ); + } + + #[test] + fn var_len_context_encode_xml_emits_only_type_byte() { + // Xml in a VarLenContext carries no length bytes: encode must emit only + // the type token (the `VarLenType::Xml => ()` arm). + let mut buf = BytesMut::new(); + VarLenContext::new(VarLenType::Xml, 0, None) + .encode(&mut buf) + .expect("encode must succeed"); + assert_eq!(buf.as_ref(), &[VarLenType::Xml as u8]); + } + + #[test] + fn var_len_context_encode_unsupported_type_errors() { + // Udt is not encodable through VarLenContext (it has its own TypeInfo + // arm), so it hits the `typ => Err(..)` fallback. + let mut buf = BytesMut::new(); + let err = VarLenContext::new(VarLenType::Udt, 0, None) + .encode(&mut buf) + .expect_err("encoding a Udt var-len context must error"); + assert!(matches!(err, Error::Protocol(_))); + } + + #[tokio::test] + async fn decode_rejects_invalid_type_byte() { + // A leading byte that is neither a FixedLenType nor a VarLenType must be + // rejected (`Err(())` arm of the VarLenType match). + let mut buf = BytesMut::new(); + buf.put_u8(0x00); + + let err = TypeInfo::decode(&mut buf.into_sql_read_bytes()) + .await + .expect_err("invalid type byte must error"); + assert!(matches!(err, Error::Protocol(_))); + } + + #[tokio::test] + async fn decode_udt_info_round_trips() { + // Exercises the UDT_INFO decode arm: max_byte_size + three b_varchars + + // a us_varchar assembly-qualified name. + let ti = TypeInfo::Udt(UdtInfo { + max_byte_size: 0xffff, + db_name: "db".to_string(), + schema_name: "dbo".to_string(), + type_name: "geometry".to_string(), + assembly_qualified_name: "asm".to_string(), + }); + + let mut buf = BytesMut::new(); + ti.clone().encode(&mut buf).expect("encode must succeed"); + + let decoded = TypeInfo::decode(&mut buf.into_sql_read_bytes()) + .await + .expect("decode must succeed"); + assert_eq!(decoded, ti); + } + + #[tokio::test] + async fn decode_rejects_scale_greater_than_precision() { + // scale > precision must be rejected. + let mut buf = BytesMut::new(); + buf.put_u8(VarLenType::Decimaln as u8); + buf.put_u8(17); // size + buf.put_u8(10); // precision + buf.put_u8(20); // scale > precision + + let err = TypeInfo::decode(&mut buf.into_sql_read_bytes()) + .await + .expect_err("scale > precision must error"); + assert!(matches!(err, Error::Protocol(_))); + } } diff --git a/src/tds/codec/type_info_tvp.rs b/src/tds/codec/type_info_tvp.rs new file mode 100644 index 000000000..fb760497f --- /dev/null +++ b/src/tds/codec/type_info_tvp.rs @@ -0,0 +1,290 @@ +use asynchronous_codec::BytesMut; +use bytes::BufMut; + +use crate::ColumnData; + +use super::{ + encode_b_varchar, BytesMutWithTypeInfo, Encode, FixedLenType, MetaDataColumn, TypeInfo, + VarLenContext, +}; + +const TVPTYPE: u8 = 0xF3; + +/// A table-valued parameter (TVP), as described by section 2.2.5.5.5 of MS-TDS. +/// +/// A TVP is passed to a stored procedure as an `RpcParam` and carries both the +/// column metadata (optionally resolved from the server) and the data rows. +#[derive(Debug)] +pub struct TypeInfoTvp<'a> { + schema_name: &'a str, + db_type_name: &'a str, + columns: Option>>, + data: Vec>>, +} + +impl<'a> Encode for TypeInfoTvp<'a> { + fn encode(self, dst: &mut BytesMut) -> crate::Result<()> { + // TVPTYPE = %xF3 + // TVP_TYPE_INFO = TVPTYPE + // TVP_TYPENAME + // TVP_COLMETADATA + // [TVP_ORDER_UNIQUE] + // [TVP_COLUMN_ORDERING] + // TVP_END_TOKEN + + dst.put_u8(TVPTYPE); + encode_b_varchar(dst, "")?; // DB name (unused) + encode_b_varchar(dst, self.schema_name)?; + encode_b_varchar(dst, self.db_type_name)?; + + if let Some(ref columns_metadata) = self.columns { + dst.put_u16_le(columns_metadata.len() as u16); + for col in columns_metadata { + // TvpColumnMetaData = UserType + // Flags + // TYPE_INFO + // ColName + dst.put_u32_le(0_u32); // UserType + col.base.clone().encode(dst)?; // Flags + TYPE_INFO + // 2.2.5.5.5.1: ColName MUST be a zero-length string in the TVP. + encode_b_varchar(dst, "")?; + } + } else { + // TVP_NULL_TOKEN: the server is expected to know the type. + dst.put_u16_le(0xFFFF_u16); + } + + dst.put_u8(0_u8); // end of TVP_COLMETADATA / optional metadata + + // TVP_ROW_TOKEN = %x01 ; A row as defined by TVP_COLMETADATA follows + // TvpColumnData = TYPE_VARBYTE ; Actual value must match column metadata + // AllColumnData = *TvpColumnData + // TVP_ROW = TVP_ROW_TOKEN AllColumnData + for row in self.data.into_iter() { + dst.put_u8(0x01u8); // TVP_ROW_TOKEN + for (i, col) in row.into_iter().enumerate() { + let mut dst_ti = BytesMutWithTypeInfo::new(dst); + if let Some(ref metadata) = self.columns { + dst_ti = dst_ti.with_type_info(&metadata[i].base.ty); + } + col.encode(&mut dst_ti)?; + } + } + + dst.put_u8(0_u8); // TVP_END_TOKEN + + Ok(()) + } +} + +impl<'a> TypeInfoTvp<'a> { + /// Creates a new TVP for the given database type name and data rows. The + /// type name may be qualified with a schema (e.g. `dbo.MyType`). + pub fn new(type_name: &'a str, rows: Vec>>) -> TypeInfoTvp<'a> { + let (schema_name, db_type_name) = if let Some((s, t)) = type_name.split_once('.') { + (s, t) + } else { + ("", type_name) + }; + TypeInfoTvp { + schema_name, + db_type_name, + columns: None, + data: rows, + } + } + + /// Attaches column metadata resolved from the server. Fixed-length column + /// types are rewritten to their nullable variable-length equivalents, as + /// required for TVP column metadata (2.2.5.5.5.3). + pub fn with_metadata(self, metadata: Vec>) -> TypeInfoTvp<'a> { + let mut metadata = metadata; + for mdc in metadata.iter_mut() { + let ty_replace = match mdc.base.ty { + TypeInfo::FixedLen(ref ty) => fixed_to_var_len(*ty), + _ => None, + }; + if let Some(ty) = ty_replace { + mdc.base.ty = ty; + } + } + TypeInfoTvp { + columns: Some(metadata), + ..self + } + } +} + +/// Maps a fixed-length column type to the nullable variable-length equivalent +/// that must be used in TVP column metadata. +fn fixed_to_var_len(ty: FixedLenType) -> Option { + use super::VarLenType; + + let ctx = match ty { + FixedLenType::Int1 => VarLenContext::new(VarLenType::Intn, 1, None), + FixedLenType::Bit => VarLenContext::new(VarLenType::Bitn, 1, None), + FixedLenType::Int2 => VarLenContext::new(VarLenType::Intn, 2, None), + FixedLenType::Int4 => VarLenContext::new(VarLenType::Intn, 4, None), + FixedLenType::Datetime4 => VarLenContext::new(VarLenType::Datetimen, 4, None), + FixedLenType::Float4 => VarLenContext::new(VarLenType::Floatn, 4, None), + FixedLenType::Money => VarLenContext::new(VarLenType::Money, 8, None), + FixedLenType::Datetime => VarLenContext::new(VarLenType::Datetimen, 8, None), + FixedLenType::Float8 => VarLenContext::new(VarLenType::Floatn, 8, None), + FixedLenType::Money4 => VarLenContext::new(VarLenType::Money, 4, None), + FixedLenType::Int8 => VarLenContext::new(VarLenType::Intn, 8, None), + _ => return None, + }; + Some(TypeInfo::VarLenSized(ctx)) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn splits_schema_qualified_type_name() { + let tvp = TypeInfoTvp::new("dbo.MyType", Vec::new()); + assert_eq!(tvp.schema_name, "dbo"); + assert_eq!(tvp.db_type_name, "MyType"); + } + + #[test] + fn unqualified_type_name_has_empty_schema() { + let tvp = TypeInfoTvp::new("MyType", Vec::new()); + assert_eq!(tvp.schema_name, ""); + assert_eq!(tvp.db_type_name, "MyType"); + } + + #[test] + fn encodes_tvp_header_and_null_metadata() { + let tvp = TypeInfoTvp::new("dbo.MyType", Vec::new()); + let mut buf = BytesMut::new(); + tvp.encode(&mut buf).unwrap(); + + // TVPTYPE marker. + assert_eq!(buf[0], TVPTYPE); + // DB name is a zero-length B_VARCHAR. + assert_eq!(buf[1], 0); + // Schema name "dbo": length 3 followed by 3 UTF-16 code units. + assert_eq!(buf[2], 3); + + // With no metadata, the column count must be the TVP_NULL_TOKEN. + // Layout: F3, 00 (db), 03 + 6 bytes (schema), 06 + 12 bytes (type), + // then the u16 null token. + let null_token_pos = 1 + 1 + (1 + 6) + (1 + 12); + let token = u16::from_le_bytes([buf[null_token_pos], buf[null_token_pos + 1]]); + assert_eq!(token, 0xFFFF); + } + + // The TVP type name is written as a B_VARCHAR (u8 length); a name longer than + // 255 UTF-16 code units must error rather than wrap the counter and desync + // the wire. + #[test] + fn encode_rejects_over_long_type_name() { + let long = "a".repeat(256); + let tvp = TypeInfoTvp::new(&long, Vec::new()); + + let mut buf = BytesMut::new(); + let err = tvp.encode(&mut buf).unwrap_err(); + assert!(matches!(err, crate::Error::Protocol(_)), "got {err:?}"); + } + + // A type name of exactly 255 units is the boundary and must still encode. + #[test] + fn encode_accepts_max_length_type_name() { + let name = "a".repeat(255); + let tvp = TypeInfoTvp::new(&name, Vec::new()); + + let mut buf = BytesMut::new(); + tvp.encode(&mut buf) + .expect("255-unit type name must encode"); + } + + #[test] + fn rewrites_fixed_len_to_nullable_var_len() { + use super::super::VarLenType; + assert!(matches!( + fixed_to_var_len(FixedLenType::Int4), + Some(TypeInfo::VarLenSized(ctx)) if ctx.r#type() == VarLenType::Intn + )); + assert!(matches!( + fixed_to_var_len(FixedLenType::Bit), + Some(TypeInfo::VarLenSized(ctx)) if ctx.r#type() == VarLenType::Bitn + )); + assert!(fixed_to_var_len(FixedLenType::Null).is_none()); + } + + #[test] + fn rewrites_all_fixed_len_variants() { + use super::super::VarLenType; + + let cases = [ + (FixedLenType::Int1, VarLenType::Intn, 1), + (FixedLenType::Bit, VarLenType::Bitn, 1), + (FixedLenType::Int2, VarLenType::Intn, 2), + (FixedLenType::Int4, VarLenType::Intn, 4), + (FixedLenType::Datetime4, VarLenType::Datetimen, 4), + (FixedLenType::Float4, VarLenType::Floatn, 4), + (FixedLenType::Money, VarLenType::Money, 8), + (FixedLenType::Datetime, VarLenType::Datetimen, 8), + (FixedLenType::Float8, VarLenType::Floatn, 8), + (FixedLenType::Money4, VarLenType::Money, 4), + (FixedLenType::Int8, VarLenType::Intn, 8), + ]; + + for (fixed, expected_ty, expected_len) in cases { + match fixed_to_var_len(fixed) { + Some(TypeInfo::VarLenSized(ctx)) => { + assert_eq!(ctx.r#type(), expected_ty, "{:?}", fixed); + assert_eq!(ctx.len(), expected_len, "{:?}", fixed); + } + other => panic!("unexpected result for {:?}: {:?}", fixed, other), + } + } + } + + #[test] + fn with_metadata_rewrites_fixed_len_columns() { + use crate::{BaseMetaDataColumn, ColumnFlag}; + use enumflags2::BitFlags; + + let metadata = vec![MetaDataColumn { + base: BaseMetaDataColumn { + flags: BitFlags::from(ColumnFlag::Nullable), + ty: TypeInfo::FixedLen(FixedLenType::Int4), + table_name: None, + }, + col_name: Default::default(), + }]; + + let tvp = TypeInfoTvp::new("MyType", Vec::new()).with_metadata(metadata); + let columns = tvp.columns.as_ref().unwrap(); + assert!(matches!(columns[0].base.ty, TypeInfo::VarLenSized(_))); + } + + #[test] + fn encodes_metadata_and_row_data() { + use crate::{BaseMetaDataColumn, ColumnData, ColumnFlag, VarLenContext, VarLenType}; + use enumflags2::BitFlags; + + let metadata = vec![MetaDataColumn { + base: BaseMetaDataColumn { + flags: BitFlags::from(ColumnFlag::Nullable), + ty: TypeInfo::VarLenSized(VarLenContext::new(VarLenType::Intn, 4, None)), + table_name: None, + }, + col_name: Default::default(), + }]; + + let rows = vec![vec![ColumnData::I32(Some(7))]]; + let tvp = TypeInfoTvp::new("dbo.MyType", rows).with_metadata(metadata); + + let mut buf = BytesMut::new(); + tvp.encode(&mut buf).unwrap(); + + // Ends with the TVP_END_TOKEN. + assert_eq!(*buf.last().unwrap(), 0); + // A single TVP_ROW_TOKEN (0x01) must appear before the row data. + assert!(buf.contains(&0x01u8)); + } +} diff --git a/src/tds/collation.rs b/src/tds/collation.rs index 20367728a..72f4702ef 100644 --- a/src/tds/collation.rs +++ b/src/tds/collation.rs @@ -3,14 +3,21 @@ //! directly from microsoft //! [2] is helpful to map CP1234 to the appropriate encoding //! -//! [1] https://github.com/Microsoft/mssql-jdbc/blob/eb14f63077c47ef1fc1c690deb8cfab602baeb85/src/main/java/com/microsoft/sqlserver/jdbc/SQLCollation.java -//! [2] https://github.com/lifthrasiir/rust-encoding/blob/496823171f15d9b9446b2ec3fb7765f22346256b/src/label.rs#L282 +//! [1] +//! [2] +//! +//! The CP437/CP850 high-byte tables follow the Unicode Consortium mappings: +//! and +//! . use encoding_rs::Encoding; +use std::borrow::Cow; use std::fmt; use crate::error::Error; +/// The collation of a character column, describing its locale (LCID), sort +/// order and code page. #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct Collation { /// LCID ColFlags Version @@ -20,6 +27,7 @@ pub struct Collation { } impl Collation { + /// Create a new collation from a raw LCID/flags/version word and sort id. pub fn new(info: u32, sort_id: u8) -> Self { Self { info, sort_id } } @@ -29,15 +37,21 @@ impl Collation { (self.info & 0xffff) as u16 } + /// The sort id of the collation. pub fn sort_id(&self) -> u8 { self.sort_id } + /// The raw LCID/flags/version word of the collation. pub fn info(&self) -> u32 { self.info } - /// return an encoding for a given collation + /// Returns an `encoding_rs` encoding for a given collation. + /// + /// CP437 and CP850 are supported by internal row/bulk codecs, but cannot + /// be represented by `encoding_rs::Encoding`; this method still returns + /// an error for those code pages. pub fn encoding(&self) -> crate::Result<&'static Encoding> { let res = if self.sort_id == 0 { lcid_to_encoding(self.lcid()) @@ -48,7 +62,7 @@ impl Collation { res.ok_or_else(|| { Error::Encoding( format!( - "encoding: unspported encoding (LCID: {:#02x}, sort ID: {})", + "encoding: unspported encoding (LCID: {:#04x}, sort ID: {})", self.lcid(), self.sort_id(), ) @@ -56,17 +70,137 @@ impl Collation { ) }) } + + pub(crate) fn codec(&self) -> crate::Result { + match self.sort_id { + 30..=35 => Ok(CollationCodec::Cp437), + 40..=45 | 49 | 55..=61 => Ok(CollationCodec::Cp850), + _ => self.encoding().map(CollationCodec::BuiltIn), + } + } } impl fmt::Display for Collation { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - match self.encoding() { - Ok(encoding) => write!(f, "{}", encoding.name()), + match self.codec() { + Ok(codec) => write!(f, "{}", codec.name()), _ => write!(f, "None"), } } } +const CP437_HIGH: [char; 128] = [ + '\u{c7}', '\u{fc}', '\u{e9}', '\u{e2}', '\u{e4}', '\u{e0}', '\u{e5}', '\u{e7}', '\u{ea}', + '\u{eb}', '\u{e8}', '\u{ef}', '\u{ee}', '\u{ec}', '\u{c4}', '\u{c5}', '\u{c9}', '\u{e6}', + '\u{c6}', '\u{f4}', '\u{f6}', '\u{f2}', '\u{fb}', '\u{f9}', '\u{ff}', '\u{d6}', '\u{dc}', + '\u{a2}', '\u{a3}', '\u{a5}', '\u{20a7}', '\u{192}', '\u{e1}', '\u{ed}', '\u{f3}', '\u{fa}', + '\u{f1}', '\u{d1}', '\u{aa}', '\u{ba}', '\u{bf}', '\u{2310}', '\u{ac}', '\u{bd}', '\u{bc}', + '\u{a1}', '\u{ab}', '\u{bb}', '\u{2591}', '\u{2592}', '\u{2593}', '\u{2502}', '\u{2524}', + '\u{2561}', '\u{2562}', '\u{2556}', '\u{2555}', '\u{2563}', '\u{2551}', '\u{2557}', '\u{255d}', + '\u{255c}', '\u{255b}', '\u{2510}', '\u{2514}', '\u{2534}', '\u{252c}', '\u{251c}', '\u{2500}', + '\u{253c}', '\u{255e}', '\u{255f}', '\u{255a}', '\u{2554}', '\u{2569}', '\u{2566}', '\u{2560}', + '\u{2550}', '\u{256c}', '\u{2567}', '\u{2568}', '\u{2564}', '\u{2565}', '\u{2559}', '\u{2558}', + '\u{2552}', '\u{2553}', '\u{256b}', '\u{256a}', '\u{2518}', '\u{250c}', '\u{2588}', '\u{2584}', + '\u{258c}', '\u{2590}', '\u{2580}', '\u{3b1}', '\u{df}', '\u{393}', '\u{3c0}', '\u{3a3}', + '\u{3c3}', '\u{b5}', '\u{3c4}', '\u{3a6}', '\u{398}', '\u{3a9}', '\u{3b4}', '\u{221e}', + '\u{3c6}', '\u{3b5}', '\u{2229}', '\u{2261}', '\u{b1}', '\u{2265}', '\u{2264}', '\u{2320}', + '\u{2321}', '\u{f7}', '\u{2248}', '\u{b0}', '\u{2219}', '\u{b7}', '\u{221a}', '\u{207f}', + '\u{b2}', '\u{25a0}', '\u{a0}', +]; + +const CP850_HIGH: [char; 128] = [ + '\u{c7}', '\u{fc}', '\u{e9}', '\u{e2}', '\u{e4}', '\u{e0}', '\u{e5}', '\u{e7}', '\u{ea}', + '\u{eb}', '\u{e8}', '\u{ef}', '\u{ee}', '\u{ec}', '\u{c4}', '\u{c5}', '\u{c9}', '\u{e6}', + '\u{c6}', '\u{f4}', '\u{f6}', '\u{f2}', '\u{fb}', '\u{f9}', '\u{ff}', '\u{d6}', '\u{dc}', + '\u{f8}', '\u{a3}', '\u{d8}', '\u{d7}', '\u{192}', '\u{e1}', '\u{ed}', '\u{f3}', '\u{fa}', + '\u{f1}', '\u{d1}', '\u{aa}', '\u{ba}', '\u{bf}', '\u{ae}', '\u{ac}', '\u{bd}', '\u{bc}', + '\u{a1}', '\u{ab}', '\u{bb}', '\u{2591}', '\u{2592}', '\u{2593}', '\u{2502}', '\u{2524}', + '\u{c1}', '\u{c2}', '\u{c0}', '\u{a9}', '\u{2563}', '\u{2551}', '\u{2557}', '\u{255d}', + '\u{a2}', '\u{a5}', '\u{2510}', '\u{2514}', '\u{2534}', '\u{252c}', '\u{251c}', '\u{2500}', + '\u{253c}', '\u{e3}', '\u{c3}', '\u{255a}', '\u{2554}', '\u{2569}', '\u{2566}', '\u{2560}', + '\u{2550}', '\u{256c}', '\u{a4}', '\u{f0}', '\u{d0}', '\u{ca}', '\u{cb}', '\u{c8}', '\u{131}', + '\u{cd}', '\u{ce}', '\u{cf}', '\u{2518}', '\u{250c}', '\u{2588}', '\u{2584}', '\u{a6}', + '\u{cc}', '\u{2580}', '\u{d3}', '\u{df}', '\u{d4}', '\u{d2}', '\u{f5}', '\u{d5}', '\u{b5}', + '\u{fe}', '\u{de}', '\u{da}', '\u{db}', '\u{d9}', '\u{fd}', '\u{dd}', '\u{af}', '\u{b4}', + '\u{ad}', '\u{b1}', '\u{2017}', '\u{be}', '\u{b6}', '\u{a7}', '\u{f7}', '\u{b8}', '\u{b0}', + '\u{a8}', '\u{b7}', '\u{b9}', '\u{b3}', '\u{b2}', '\u{25a0}', '\u{a0}', +]; + +pub(crate) enum CollationCodec { + BuiltIn(&'static Encoding), + Cp437, + Cp850, +} + +impl CollationCodec { + fn name(&self) -> &'static str { + match self { + Self::BuiltIn(encoding) => encoding.name(), + Self::Cp437 => "CP437", + Self::Cp850 => "CP850", + } + } + + pub(crate) fn decode(&self, bytes: &[u8]) -> Option { + match self { + Self::BuiltIn(encoding) => encoding + .decode_without_bom_handling_and_without_replacement(bytes) + .map(Cow::into_owned), + Self::Cp437 => Some(decode_single_byte(bytes, &CP437_HIGH)), + Self::Cp850 => Some(decode_single_byte(bytes, &CP850_HIGH)), + } + } + + pub(crate) fn decode_lossy(&self, bytes: &[u8]) -> String { + match self { + Self::BuiltIn(encoding) => encoding.decode_without_bom_handling(bytes).0.into_owned(), + Self::Cp437 => decode_single_byte(bytes, &CP437_HIGH), + Self::Cp850 => decode_single_byte(bytes, &CP850_HIGH), + } + } + + pub(crate) fn encode(&self, text: &str) -> Option> { + match self { + Self::BuiltIn(encoding) => { + let mut encoder = encoding.new_encoder(); + let capacity = + encoder.max_buffer_length_from_utf8_without_replacement(text.len())?; + let mut bytes = Vec::with_capacity(capacity); + let (result, _) = + encoder.encode_from_utf8_to_vec_without_replacement(text, &mut bytes, true); + match result { + encoding_rs::EncoderResult::InputEmpty => Some(bytes), + _ => None, + } + } + Self::Cp437 => encode_single_byte(text, &CP437_HIGH), + Self::Cp850 => encode_single_byte(text, &CP850_HIGH), + } + } +} + +fn decode_single_byte(bytes: &[u8], high: &[char; 128]) -> String { + bytes + .iter() + .map(|&byte| match byte { + 0..=0x7f => char::from(byte), + _ => high[usize::from(byte) - 0x80], + }) + .collect() +} + +fn encode_single_byte(text: &str, high: &[char; 128]) -> Option> { + let mut bytes = Vec::with_capacity(text.len()); + for character in text.chars() { + let byte = match character { + '\0'..='\u{7f}' => character as u8, + _ => high.iter().position(|&mapped| mapped == character)? as u8 + 0x80, + }; + bytes.push(byte); + } + Some(bytes) +} + /// https://github.com/Microsoft/mssql-jdbc/blob/eb14f63077c47ef1fc1c690deb8cfab602baeb85/src/main/java/com/microsoft/sqlserver/jdbc/SQLCollation.java#L102-L310 /// maps an LCID (it's locale part which is only 2 bytes) to a codepage /// @@ -74,7 +208,7 @@ impl fmt::Display for Collation { /// 1. (regex)replace: (.*?)\((.*?),(.*?)\) with $2 => $3 /// 2. replace: Encoding.CP(.*?) with encoding::all::WINDOWS_$1 /// 3. replace: Encoding.UNICODE with encoding::all::UTF16_LE -// +/// /// the unimplemented!() one's are not supported by rust-encoding pub fn lcid_to_encoding(locale: u16) -> Option<&'static Encoding> { match locale { @@ -389,30 +523,415 @@ pub fn sortid_to_encoding(sort_id: u8) -> Option<&'static Encoding> { } } -/* TODO #[cfg(test)] mod tests { - use futures_state_stream::StateStream; - use tokio::executor::current_thread; - use crate::tests::new_connection; + use super::*; #[test] - fn select_nvarchar_collation_test() { - let c1 = new_connection(); - let query = c1.simple_query( - "select cast(cast(N'cześć' as nvarchar(5)) collate Polish_CI_AI as varchar(5))", - ); - let mut i = 0; - { - let future = query.for_each(|x| { - let val: &str = x.get(0); - assert_eq!(val, "cześć"); - i += 1; - Ok(()) - }); - current_thread::block_on_all(future).unwrap(); + fn accessors_split_info_and_sort_id() { + // info holds LCID in the low 16 bits plus flags/version in the high bits. + let collation = Collation::new(0x0020_0409, 52); + assert_eq!(collation.info(), 0x0020_0409); + assert_eq!(collation.lcid(), 0x0409); + assert_eq!(collation.sort_id(), 52); + } + + #[test] + fn encoding_from_lcid_when_sort_id_zero() { + // sort_id == 0 -> resolve via the LCID. + let collation = Collation::new(0x0409, 0); + let encoding = collation.encoding().expect("known LCID must resolve"); + assert_eq!(encoding, encoding_rs::WINDOWS_1252); + } + + #[test] + fn encoding_from_sort_id_when_present() { + // A non-zero sort_id takes precedence over the LCID. + let collation = Collation::new(0x0405, 80); + let encoding = collation.encoding().expect("known sort id must resolve"); + assert_eq!(encoding, encoding_rs::WINDOWS_1250); + } + + #[test] + fn encoding_unknown_lcid_errors() { + let collation = Collation::new(0xFFFF, 0); + let err = collation.encoding().unwrap_err(); + assert!(matches!(err, Error::Encoding(_))); + } + + #[test] + fn encoding_unknown_sort_id_errors() { + let collation = Collation::new(0x0409, 250); + assert!(collation.encoding().is_err()); + } + + #[test] + fn display_uses_encoding_name() { + let collation = Collation::new(0x0409, 0); + assert_eq!(format!("{}", collation), encoding_rs::WINDOWS_1252.name()); + } + + #[test] + fn display_falls_back_to_none_on_unknown() { + let collation = Collation::new(0xFFFF, 0); + assert_eq!(format!("{}", collation), "None"); + } + + #[test] + fn lcid_to_encoding_known_and_unknown() { + assert_eq!(lcid_to_encoding(0x0401), Some(encoding_rs::WINDOWS_1256)); + assert_eq!(lcid_to_encoding(0x0404), Some(encoding_rs::BIG5)); + assert_eq!(lcid_to_encoding(0x0411), Some(encoding_rs::SHIFT_JIS)); + assert_eq!(lcid_to_encoding(0x0412), Some(encoding_rs::EUC_KR)); + assert_eq!(lcid_to_encoding(0x0804), Some(encoding_rs::GB18030)); + assert_eq!(lcid_to_encoding(0x0439), Some(encoding_rs::UTF_16LE)); + assert_eq!(lcid_to_encoding(0x041e), Some(encoding_rs::WINDOWS_874)); + assert_eq!(lcid_to_encoding(0x0000), None); + } + + #[test] + fn sortid_to_encoding_known_and_unknown() { + assert_eq!(sortid_to_encoding(50), Some(encoding_rs::WINDOWS_1252)); + assert_eq!(sortid_to_encoding(80), Some(encoding_rs::WINDOWS_1250)); + assert_eq!(sortid_to_encoding(104), Some(encoding_rs::WINDOWS_1251)); + assert_eq!(sortid_to_encoding(112), Some(encoding_rs::WINDOWS_1253)); + assert_eq!(sortid_to_encoding(192), Some(encoding_rs::SHIFT_JIS)); + assert_eq!(sortid_to_encoding(194), Some(encoding_rs::EUC_KR)); + assert_eq!(sortid_to_encoding(196), Some(encoding_rs::BIG5)); + assert_eq!(sortid_to_encoding(198), Some(encoding_rs::GB18030)); + assert_eq!(sortid_to_encoding(204), Some(encoding_rs::WINDOWS_874)); + assert_eq!(sortid_to_encoding(0), None); + assert_eq!(sortid_to_encoding(255), None); + } + + #[test] + fn lcid_to_encoding_covers_all_documented_locales() { + let cases: &[(u16, &encoding_rs::Encoding)] = &[ + (0x0401, encoding_rs::WINDOWS_1256), + (0x0402, encoding_rs::WINDOWS_1251), + (0x0403, encoding_rs::WINDOWS_1252), + (0x0404, encoding_rs::BIG5), + (0x0c04, encoding_rs::BIG5), + (0x1404, encoding_rs::BIG5), + (0x0405, encoding_rs::WINDOWS_1250), + (0x0406, encoding_rs::WINDOWS_1252), + (0x0407, encoding_rs::WINDOWS_1252), + (0x0408, encoding_rs::WINDOWS_1253), + (0x0409, encoding_rs::WINDOWS_1252), + (0x040a, encoding_rs::WINDOWS_1252), + (0x040b, encoding_rs::WINDOWS_1252), + (0x040c, encoding_rs::WINDOWS_1252), + (0x040d, encoding_rs::WINDOWS_1255), + (0x040e, encoding_rs::WINDOWS_1250), + (0x040f, encoding_rs::WINDOWS_1252), + (0x0410, encoding_rs::WINDOWS_1252), + (0x0411, encoding_rs::SHIFT_JIS), + (0x0412, encoding_rs::EUC_KR), + (0x0413, encoding_rs::WINDOWS_1252), + (0x0414, encoding_rs::WINDOWS_1252), + (0x0415, encoding_rs::WINDOWS_1250), + (0x0416, encoding_rs::WINDOWS_1252), + (0x0417, encoding_rs::WINDOWS_1252), + (0x0418, encoding_rs::WINDOWS_1250), + (0x0419, encoding_rs::WINDOWS_1251), + (0x041a, encoding_rs::WINDOWS_1250), + (0x041b, encoding_rs::WINDOWS_1250), + (0x041c, encoding_rs::WINDOWS_1250), + (0x041d, encoding_rs::WINDOWS_1252), + (0x041e, encoding_rs::WINDOWS_874), + (0x041f, encoding_rs::WINDOWS_1254), + (0x0420, encoding_rs::WINDOWS_1256), + (0x0421, encoding_rs::WINDOWS_1252), + (0x0422, encoding_rs::WINDOWS_1251), + (0x0423, encoding_rs::WINDOWS_1251), + (0x0424, encoding_rs::WINDOWS_1250), + (0x0425, encoding_rs::WINDOWS_1257), + (0x0426, encoding_rs::WINDOWS_1257), + (0x0427, encoding_rs::WINDOWS_1257), + (0x0428, encoding_rs::WINDOWS_1251), + (0x0429, encoding_rs::WINDOWS_1256), + (0x042a, encoding_rs::WINDOWS_1258), + (0x042b, encoding_rs::WINDOWS_1252), + (0x042c, encoding_rs::WINDOWS_1254), + (0x042d, encoding_rs::WINDOWS_1252), + (0x042e, encoding_rs::WINDOWS_1252), + (0x042f, encoding_rs::WINDOWS_1251), + (0x0432, encoding_rs::WINDOWS_1252), + (0x0434, encoding_rs::WINDOWS_1252), + (0x0435, encoding_rs::WINDOWS_1252), + (0x0436, encoding_rs::WINDOWS_1252), + (0x0437, encoding_rs::WINDOWS_1252), + (0x0438, encoding_rs::WINDOWS_1252), + (0x0439, encoding_rs::UTF_16LE), + (0x043a, encoding_rs::UTF_16LE), + (0x043b, encoding_rs::WINDOWS_1252), + (0x043e, encoding_rs::WINDOWS_1252), + (0x043f, encoding_rs::WINDOWS_1251), + (0x0440, encoding_rs::WINDOWS_1251), + (0x0441, encoding_rs::WINDOWS_1252), + (0x0442, encoding_rs::WINDOWS_1250), + (0x0443, encoding_rs::WINDOWS_1254), + (0x0444, encoding_rs::WINDOWS_1251), + (0x0445, encoding_rs::UTF_16LE), + (0x0446, encoding_rs::UTF_16LE), + (0x0447, encoding_rs::UTF_16LE), + (0x0448, encoding_rs::UTF_16LE), + (0x0449, encoding_rs::UTF_16LE), + (0x044a, encoding_rs::UTF_16LE), + (0x044b, encoding_rs::UTF_16LE), + (0x044c, encoding_rs::UTF_16LE), + (0x044d, encoding_rs::UTF_16LE), + (0x044e, encoding_rs::UTF_16LE), + (0x044f, encoding_rs::UTF_16LE), + (0x0450, encoding_rs::WINDOWS_1251), + (0x0451, encoding_rs::UTF_16LE), + (0x0452, encoding_rs::WINDOWS_1252), + (0x0453, encoding_rs::UTF_16LE), + (0x0454, encoding_rs::UTF_16LE), + (0x0456, encoding_rs::WINDOWS_1252), + (0x0457, encoding_rs::UTF_16LE), + (0x045a, encoding_rs::UTF_16LE), + (0x045b, encoding_rs::UTF_16LE), + (0x045d, encoding_rs::WINDOWS_1252), + (0x045e, encoding_rs::WINDOWS_1252), + (0x0461, encoding_rs::UTF_16LE), + (0x0462, encoding_rs::WINDOWS_1252), + (0x0463, encoding_rs::UTF_16LE), + (0x0464, encoding_rs::WINDOWS_1252), + (0x0465, encoding_rs::UTF_16LE), + (0x0468, encoding_rs::WINDOWS_1252), + (0x046a, encoding_rs::WINDOWS_1252), + (0x046b, encoding_rs::WINDOWS_1252), + (0x046c, encoding_rs::WINDOWS_1252), + (0x046d, encoding_rs::WINDOWS_1251), + (0x046e, encoding_rs::WINDOWS_1252), + (0x046f, encoding_rs::WINDOWS_1252), + (0x0470, encoding_rs::WINDOWS_1252), + (0x0478, encoding_rs::WINDOWS_1252), + (0x047a, encoding_rs::WINDOWS_1252), + (0x047c, encoding_rs::WINDOWS_1252), + (0x047e, encoding_rs::WINDOWS_1252), + (0x0480, encoding_rs::WINDOWS_1256), + (0x0481, encoding_rs::UTF_16LE), + (0x0482, encoding_rs::WINDOWS_1252), + (0x0483, encoding_rs::WINDOWS_1252), + (0x0484, encoding_rs::WINDOWS_1252), + (0x0485, encoding_rs::WINDOWS_1251), + (0x0486, encoding_rs::WINDOWS_1252), + (0x0487, encoding_rs::WINDOWS_1252), + (0x0488, encoding_rs::WINDOWS_1252), + (0x048c, encoding_rs::WINDOWS_1256), + (0x0801, encoding_rs::WINDOWS_1256), + (0x0804, encoding_rs::GB18030), + (0x1004, encoding_rs::GB18030), + (0x0807, encoding_rs::WINDOWS_1252), + (0x0809, encoding_rs::WINDOWS_1252), + (0x080a, encoding_rs::WINDOWS_1252), + (0x080c, encoding_rs::WINDOWS_1252), + (0x0810, encoding_rs::WINDOWS_1252), + (0x0813, encoding_rs::WINDOWS_1252), + (0x0814, encoding_rs::WINDOWS_1252), + (0x0816, encoding_rs::WINDOWS_1252), + (0x081a, encoding_rs::WINDOWS_1250), + (0x081d, encoding_rs::WINDOWS_1252), + (0x0827, encoding_rs::WINDOWS_1257), + (0x082c, encoding_rs::WINDOWS_1251), + (0x082e, encoding_rs::WINDOWS_1252), + (0x083b, encoding_rs::WINDOWS_1252), + (0x083c, encoding_rs::WINDOWS_1252), + (0x083e, encoding_rs::WINDOWS_1252), + (0x0843, encoding_rs::WINDOWS_1251), + (0x0845, encoding_rs::UTF_16LE), + (0x0850, encoding_rs::WINDOWS_1251), + (0x085d, encoding_rs::WINDOWS_1252), + (0x085f, encoding_rs::WINDOWS_1252), + (0x086b, encoding_rs::WINDOWS_1252), + (0x0c01, encoding_rs::WINDOWS_1256), + (0x0c07, encoding_rs::WINDOWS_1252), + (0x0c09, encoding_rs::WINDOWS_1252), + (0x0c0a, encoding_rs::WINDOWS_1252), + (0x0c0c, encoding_rs::WINDOWS_1252), + (0x0c1a, encoding_rs::WINDOWS_1251), + (0x0c3b, encoding_rs::WINDOWS_1252), + (0x0c6b, encoding_rs::WINDOWS_1252), + (0x1001, encoding_rs::WINDOWS_1256), + (0x1007, encoding_rs::WINDOWS_1252), + (0x1009, encoding_rs::WINDOWS_1252), + (0x100a, encoding_rs::WINDOWS_1252), + (0x100c, encoding_rs::WINDOWS_1252), + (0x101a, encoding_rs::WINDOWS_1250), + (0x103b, encoding_rs::WINDOWS_1252), + (0x1401, encoding_rs::WINDOWS_1256), + (0x1407, encoding_rs::WINDOWS_1252), + (0x1409, encoding_rs::WINDOWS_1252), + (0x140a, encoding_rs::WINDOWS_1252), + (0x140c, encoding_rs::WINDOWS_1252), + (0x141a, encoding_rs::WINDOWS_1250), + (0x143b, encoding_rs::WINDOWS_1252), + (0x1801, encoding_rs::WINDOWS_1256), + (0x1809, encoding_rs::WINDOWS_1252), + (0x180a, encoding_rs::WINDOWS_1252), + (0x180c, encoding_rs::WINDOWS_1252), + (0x181a, encoding_rs::WINDOWS_1250), + (0x183b, encoding_rs::WINDOWS_1252), + (0x1c01, encoding_rs::WINDOWS_1256), + (0x1c09, encoding_rs::WINDOWS_1252), + (0x1c0a, encoding_rs::WINDOWS_1252), + (0x1c1a, encoding_rs::WINDOWS_1251), + (0x1c3b, encoding_rs::WINDOWS_1252), + (0x2001, encoding_rs::WINDOWS_1256), + (0x2009, encoding_rs::WINDOWS_1252), + (0x200a, encoding_rs::WINDOWS_1252), + (0x201a, encoding_rs::WINDOWS_1251), + (0x203b, encoding_rs::WINDOWS_1252), + (0x2401, encoding_rs::WINDOWS_1256), + (0x2409, encoding_rs::WINDOWS_1252), + (0x240a, encoding_rs::WINDOWS_1252), + (0x243b, encoding_rs::WINDOWS_1252), + (0x2801, encoding_rs::WINDOWS_1256), + (0x2809, encoding_rs::WINDOWS_1252), + (0x280a, encoding_rs::WINDOWS_1252), + (0x2c01, encoding_rs::WINDOWS_1256), + (0x2c09, encoding_rs::WINDOWS_1252), + (0x2c0a, encoding_rs::WINDOWS_1252), + (0x3001, encoding_rs::WINDOWS_1256), + (0x3009, encoding_rs::WINDOWS_1252), + (0x300a, encoding_rs::WINDOWS_1252), + (0x3401, encoding_rs::WINDOWS_1256), + (0x3409, encoding_rs::WINDOWS_1252), + (0x340a, encoding_rs::WINDOWS_1252), + (0x3801, encoding_rs::WINDOWS_1256), + (0x380a, encoding_rs::WINDOWS_1252), + (0x3c01, encoding_rs::WINDOWS_1256), + (0x3c0a, encoding_rs::WINDOWS_1252), + (0x4001, encoding_rs::WINDOWS_1256), + (0x4009, encoding_rs::WINDOWS_1252), + (0x400a, encoding_rs::WINDOWS_1252), + (0x4409, encoding_rs::WINDOWS_1252), + (0x440a, encoding_rs::WINDOWS_1252), + (0x4809, encoding_rs::WINDOWS_1252), + (0x480a, encoding_rs::WINDOWS_1252), + (0x4c0a, encoding_rs::WINDOWS_1252), + (0x500a, encoding_rs::WINDOWS_1252), + (0x540a, encoding_rs::WINDOWS_1252), + ]; + + for (locale, expected) in cases { + assert_eq!( + lcid_to_encoding(*locale), + Some(*expected), + "locale {:#06x}", + locale + ); } - assert_eq!(i, 1); + } + + #[test] + fn sortid_to_encoding_covers_all_documented_sort_ids() { + let cases: &[(u8, &encoding_rs::Encoding)] = &[ + (50, encoding_rs::WINDOWS_1252), + (51, encoding_rs::WINDOWS_1252), + (52, encoding_rs::WINDOWS_1252), + (53, encoding_rs::WINDOWS_1252), + (54, encoding_rs::WINDOWS_1252), + (71, encoding_rs::WINDOWS_1252), + (72, encoding_rs::WINDOWS_1252), + (73, encoding_rs::WINDOWS_1252), + (74, encoding_rs::WINDOWS_1252), + (75, encoding_rs::WINDOWS_1252), + (80, encoding_rs::WINDOWS_1250), + (81, encoding_rs::WINDOWS_1250), + (82, encoding_rs::WINDOWS_1250), + (83, encoding_rs::WINDOWS_1250), + (84, encoding_rs::WINDOWS_1250), + (85, encoding_rs::WINDOWS_1250), + (86, encoding_rs::WINDOWS_1250), + (87, encoding_rs::WINDOWS_1250), + (88, encoding_rs::WINDOWS_1250), + (89, encoding_rs::WINDOWS_1250), + (90, encoding_rs::WINDOWS_1250), + (91, encoding_rs::WINDOWS_1250), + (92, encoding_rs::WINDOWS_1250), + (93, encoding_rs::WINDOWS_1250), + (94, encoding_rs::WINDOWS_1250), + (95, encoding_rs::WINDOWS_1250), + (96, encoding_rs::WINDOWS_1250), + (97, encoding_rs::WINDOWS_1250), + (98, encoding_rs::WINDOWS_1250), + (104, encoding_rs::WINDOWS_1251), + (105, encoding_rs::WINDOWS_1251), + (106, encoding_rs::WINDOWS_1251), + (107, encoding_rs::WINDOWS_1251), + (108, encoding_rs::WINDOWS_1251), + (112, encoding_rs::WINDOWS_1253), + (113, encoding_rs::WINDOWS_1253), + (114, encoding_rs::WINDOWS_1253), + (120, encoding_rs::WINDOWS_1253), + (121, encoding_rs::WINDOWS_1253), + (122, encoding_rs::WINDOWS_1253), + (124, encoding_rs::WINDOWS_1253), + (128, encoding_rs::WINDOWS_1254), + (129, encoding_rs::WINDOWS_1254), + (130, encoding_rs::WINDOWS_1254), + (136, encoding_rs::WINDOWS_1255), + (137, encoding_rs::WINDOWS_1255), + (138, encoding_rs::WINDOWS_1255), + (144, encoding_rs::WINDOWS_1256), + (145, encoding_rs::WINDOWS_1256), + (146, encoding_rs::WINDOWS_1256), + (152, encoding_rs::WINDOWS_1257), + (153, encoding_rs::WINDOWS_1257), + (154, encoding_rs::WINDOWS_1257), + (155, encoding_rs::WINDOWS_1257), + (156, encoding_rs::WINDOWS_1257), + (157, encoding_rs::WINDOWS_1257), + (158, encoding_rs::WINDOWS_1257), + (159, encoding_rs::WINDOWS_1257), + (160, encoding_rs::WINDOWS_1257), + (183, encoding_rs::WINDOWS_1252), + (184, encoding_rs::WINDOWS_1252), + (185, encoding_rs::WINDOWS_1252), + (186, encoding_rs::WINDOWS_1252), + (192, encoding_rs::SHIFT_JIS), + (193, encoding_rs::SHIFT_JIS), + (200, encoding_rs::SHIFT_JIS), + (194, encoding_rs::EUC_KR), + (195, encoding_rs::EUC_KR), + (196, encoding_rs::BIG5), + (197, encoding_rs::BIG5), + (202, encoding_rs::BIG5), + (198, encoding_rs::GB18030), + (199, encoding_rs::GB18030), + (203, encoding_rs::GB18030), + (201, encoding_rs::BIG5), + (204, encoding_rs::WINDOWS_874), + (205, encoding_rs::WINDOWS_874), + (206, encoding_rs::WINDOWS_874), + (210, encoding_rs::WINDOWS_1252), + (211, encoding_rs::WINDOWS_1252), + (212, encoding_rs::WINDOWS_1252), + (213, encoding_rs::WINDOWS_1252), + (214, encoding_rs::WINDOWS_1252), + (215, encoding_rs::WINDOWS_1252), + (216, encoding_rs::WINDOWS_1252), + (217, encoding_rs::WINDOWS_1252), + ]; + + for (sort_id, expected) in cases { + assert_eq!( + sortid_to_encoding(*sort_id), + Some(*expected), + "sort_id {}", + sort_id + ); + } + } + + #[test] + fn derives_eq_and_copy() { + let a = Collation::new(0x0409, 0); + let b = a; + assert_eq!(a, b); + assert_ne!(a, Collation::new(0x0409, 1)); } } -*/ diff --git a/src/tds/context.rs b/src/tds/context.rs index 55797e649..5031b55dd 100644 --- a/src/tds/context.rs +++ b/src/tds/context.rs @@ -1,5 +1,7 @@ use super::codec::*; +use std::collections::HashMap; use std::sync::Arc; +use std::time::Duration; /// Context, that might be required to make sure we understand and are understood by the server #[derive(Debug)] @@ -9,18 +11,50 @@ pub(crate) struct Context { packet_id: u8, transaction_desc: [u8; 8], last_meta: Option>>, + /// Metadata for COMPUTE (BY) result sets (`ALTMETADATA`), keyed by the + /// COMPUTE clause id that the matching `ALTROW` rows refer back to. + alt_metas: HashMap>>, spn: Option, + /// Per-response deadline for reading command results, propagated from + /// [`Config::command_timeout`](crate::Config::command_timeout) at connect + /// time. `None` means unbounded. Read by the token stream to bound each + /// server round-trip (see `TokenStream::try_unfold`). + command_timeout: Option, + /// When `true`, NVARCHAR/NTEXT row values that contain malformed UTF-16 + /// (e.g. unpaired surrogates that SQL Server stored as unchecked UCS-2) + /// are decoded losslessly by replacing each invalid sequence with the + /// Unicode replacement character (U+FFFD) instead of aborting the row + /// stream with a protocol error. Propagated from + /// [`Config::lossy_utf16_decoding`](crate::Config::lossy_utf16_decoding) at + /// connect time. Defaults to `false` (strict decoding). + lossy_utf16: bool, + /// Replaces invalid code-page row sequences with U+FFFD when enabled by + /// [`Config::lossy_codepage_decoding`](crate::Config::lossy_codepage_decoding) + /// at connect time. Defaults to `false`, independently of `lossy_utf16`. + lossy_codepage: bool, } impl Context { pub fn new() -> Context { Context { + // Overwritten with the server's negotiated TDS version from LOGINACK + // on a successful login (`set_version`, called from + // `TokenStream::get_login_ack` in stream/token.rs). tiberius targets + // TDS 7.2+ (SQL Server 2005+), so this default's 7.2+ field widths + // (4-byte ERROR/INFO LineNumber, 8-byte DONE row-count) are correct + // for all supported servers even on the login-failure path, where no + // LOGINACK ever arrives to update this value. A pre-7.2 server that + // rejects login is out of supported scope. version: FeatureLevel::SqlServerN, packet_size: 4096, packet_id: 0, transaction_desc: [0; 8], last_meta: None, + alt_metas: HashMap::new(), spn: None, + command_timeout: None, + lossy_utf16: false, + lossy_codepage: false, } } @@ -38,6 +72,17 @@ impl Context { self.last_meta.clone() } + /// Stores the metadata for a COMPUTE (BY) result set, keyed by its id, so + /// that a following `ALTROW` token can be decoded. + pub fn set_alt_meta(&mut self, meta: Arc>) { + self.alt_metas.insert(meta.id, meta); + } + + /// Retrieves previously seen COMPUTE (BY) metadata by its id. + pub fn alt_meta(&self, id: u16) -> Option>> { + self.alt_metas.get(&id).cloned() + } + pub fn packet_size(&self) -> u32 { self.packet_size } @@ -46,6 +91,44 @@ impl Context { self.packet_size = new_size; } + /// The per-response command timeout, if any. See + /// [`Config::command_timeout`](crate::Config::command_timeout). + pub(crate) fn command_timeout(&self) -> Option { + self.command_timeout + } + + /// Records the per-response command timeout negotiated from the [`Config`] + /// at connect time. + /// + /// [`Config`]: crate::Config + pub(crate) fn set_command_timeout(&mut self, timeout: Option) { + self.command_timeout = timeout; + } + + /// Whether malformed UTF-16 in NVARCHAR/NTEXT row values should be decoded + /// losslessly (replacing invalid sequences with U+FFFD) rather than + /// erroring. See + /// [`Config::lossy_utf16_decoding`](crate::Config::lossy_utf16_decoding). + pub(crate) fn lossy_utf16(&self) -> bool { + self.lossy_utf16 + } + + /// Records whether lossy UTF-16 decoding is enabled, negotiated from the + /// [`Config`](crate::Config) at connect time. + pub(crate) fn set_lossy_utf16(&mut self, lossy: bool) { + self.lossy_utf16 = lossy; + } + + /// Whether malformed code-page row values use replacement decoding. + pub(crate) fn lossy_codepage(&self) -> bool { + self.lossy_codepage + } + + /// Records the code-page decoding preference from the connection config. + pub(crate) fn set_lossy_codepage(&mut self, lossy: bool) { + self.lossy_codepage = lossy; + } + pub fn transaction_descriptor(&self) -> [u8; 8] { self.transaction_desc } @@ -54,6 +137,13 @@ impl Context { self.transaction_desc = desc; } + /// Records the protocol version negotiated with the server (from the + /// LOGINACK token). This drives version-dependent decode paths such as the + /// pre-2005 4-byte DONE rowcount and the 2-byte ERROR/INFO LineNumber. + pub(crate) fn set_version(&mut self, version: FeatureLevel) { + self.version = version; + } + pub fn version(&self) -> FeatureLevel { self.version } @@ -62,8 +152,136 @@ impl Context { self.spn = Some(format!("MSSQLSvc/{}:{}", host.as_ref(), port)); } - #[cfg(any(windows, feature = "winauth", all(unix, feature = "integrated-auth-gssapi")))] + #[cfg(any( + windows, + feature = "winauth", + all(unix, any(feature = "integrated-auth-gssapi", feature = "sspi-rs")) + ))] pub fn spn(&self) -> &str { self.spn.as_deref().unwrap_or("") } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn new_has_expected_defaults() { + let ctx = Context::new(); + assert_eq!(ctx.packet_size(), 4096); + assert_eq!(ctx.transaction_descriptor(), [0; 8]); + assert!(ctx.last_meta().is_none()); + assert!(ctx.alt_meta(0).is_none()); + } + + #[test] + fn next_packet_id_increments_and_wraps() { + let mut ctx = Context::new(); + assert_eq!(ctx.next_packet_id(), 0); + assert_eq!(ctx.next_packet_id(), 1); + assert_eq!(ctx.next_packet_id(), 2); + + // Force a wraparound to make sure it doesn't panic on overflow. + for _ in 0..252 { + ctx.next_packet_id(); + } + assert_eq!(ctx.next_packet_id(), 255); + assert_eq!(ctx.next_packet_id(), 0); + } + + #[test] + fn set_and_get_packet_size() { + let mut ctx = Context::new(); + ctx.set_packet_size(8192); + assert_eq!(ctx.packet_size(), 8192); + } + + #[test] + fn command_timeout_defaults_to_none_and_roundtrips() { + let mut ctx = Context::new(); + assert_eq!(ctx.command_timeout(), None); + + ctx.set_command_timeout(Some(Duration::from_secs(5))); + assert_eq!(ctx.command_timeout(), Some(Duration::from_secs(5))); + + ctx.set_command_timeout(None); + assert_eq!(ctx.command_timeout(), None); + } + + #[test] + fn lossy_utf16_defaults_to_false_and_roundtrips() { + let mut ctx = Context::new(); + assert!(!ctx.lossy_utf16()); + + ctx.set_lossy_utf16(true); + assert!(ctx.lossy_utf16()); + + ctx.set_lossy_utf16(false); + assert!(!ctx.lossy_utf16()); + } + + #[test] + fn set_and_get_transaction_descriptor() { + let mut ctx = Context::new(); + let desc = [1, 2, 3, 4, 5, 6, 7, 8]; + ctx.set_transaction_descriptor(desc); + assert_eq!(ctx.transaction_descriptor(), desc); + } + + #[test] + fn set_and_get_last_meta() { + let mut ctx = Context::new(); + let meta = Arc::new(TokenColMetaData { columns: vec![] }); + ctx.set_last_meta(meta.clone()); + + let got = ctx.last_meta().unwrap(); + assert_eq!(got.columns.len(), meta.columns.len()); + } + + #[test] + fn set_and_get_alt_meta_by_id() { + let mut ctx = Context::new(); + let meta = Arc::new(TokenAltMetaData { + id: 7, + by_columns: vec![1, 2], + columns: vec![], + }); + ctx.set_alt_meta(meta.clone()); + + let got = ctx.alt_meta(7).unwrap(); + assert_eq!(got.id, 7); + assert_eq!(got.by_columns, vec![1, 2]); + + // A different id should still be absent. + assert!(ctx.alt_meta(8).is_none()); + } + + #[test] + fn version_defaults_to_sql_server_n() { + let ctx = Context::new(); + assert_eq!(ctx.version(), FeatureLevel::SqlServerN); + } + + #[test] + fn set_spn_formats_service_principal_name() { + let mut ctx = Context::new(); + ctx.set_spn("dbhost", 1433); + + #[cfg(any( + windows, + feature = "winauth", + all(unix, any(feature = "integrated-auth-gssapi", feature = "sspi-rs")) + ))] + assert_eq!(ctx.spn(), "MSSQLSvc/dbhost:1433"); + + // On platforms without an spn() accessor, at least make sure setting + // it doesn't panic. + #[cfg(not(any( + windows, + feature = "winauth", + all(unix, any(feature = "integrated-auth-gssapi", feature = "sspi-rs")) + )))] + let _ = ctx; + } +} diff --git a/src/tds/numeric.rs b/src/tds/numeric.rs index 4f856bebb..36791b474 100644 --- a/src/tds/numeric.rs +++ b/src/tds/numeric.rs @@ -3,12 +3,12 @@ use super::codec::Encode; use crate::{sql_read_bytes::SqlReadBytes, Error}; #[cfg(feature = "bigdecimal")] -#[cfg_attr(feature = "docs", doc(cfg(feature = "bigdecimal")))] +#[cfg_attr(docsrs, doc(cfg(feature = "bigdecimal")))] pub use bigdecimal::{num_bigint::BigInt, BigDecimal}; use byteorder::{ByteOrder, LittleEndian}; use bytes::{BufMut, BytesMut}; #[cfg(feature = "rust_decimal")] -#[cfg_attr(feature = "docs", doc(cfg(feature = "rust_decimal")))] +#[cfg_attr(docsrs, doc(cfg(feature = "rust_decimal")))] pub use rust_decimal::Decimal; use std::cmp::{Ordering, PartialEq}; use std::fmt::{self, Debug, Display, Formatter}; @@ -19,20 +19,25 @@ use std::fmt::{self, Debug, Display, Formatter}; /// A recommended way of dealing with numeric values is by enabling the /// `rust_decimal` feature and using its `Decimal` type instead. #[derive(Copy, Clone)] +#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] pub struct Numeric { value: i128, scale: u8, } impl Numeric { + /// The maximum scale SQL Server supports for `NUMERIC`/`DECIMAL`: the + /// precision tops out at 38 digits and the scale may equal the precision + /// (e.g. `decimal(38, 38)`), so 38 is the largest valid scale. `10^38` still + /// fits in an `i128`. + pub const MAX_NUMERIC_SCALE: u8 = 38; + /// Creates a new Numeric value. /// /// # Panic - /// It will panic if the scale exceed 37. + /// It will panic if the scale exceeds [`Numeric::MAX_NUMERIC_SCALE`] (38). pub fn new_with_scale(value: i128, scale: u8) -> Self { - // scale cannot exceed 37 since a - // max precision of 38 is possible here. - assert!(scale < 38); + assert!(scale <= Self::MAX_NUMERIC_SCALE); Numeric { value, scale } } @@ -108,11 +113,10 @@ impl Numeric { _ => unreachable!(), }; - // swap high&low for big endian - #[cfg(target_endian = "big")] - let (low_part, high_part) = (high_part, low_part); - - let high_part = high_part * (u64::max_value() as u128 + 1); + // `byteorder::LittleEndian` already yields the correct host-native + // integer regardless of target endianness, so `low_part`/`high_part` + // need no further swapping. + let high_part = high_part * (u64::MAX as u128 + 1); low_part + high_part } @@ -131,18 +135,32 @@ impl Numeric { 5 => src.read_u32_le().await? as i128 * sign, 9 => src.read_u64_le().await? as i128 * sign, 13 => { + // Bulk-read the 12 magnitude bytes (u96) in one packet-aware + // pass instead of 12 separate `read_u8().await` calls. + let mut buf = Vec::new(); + crate::sql_read_bytes::read_bytes_into(src, &mut buf, 12, 12).await?; let mut bytes = [0u8; 12]; //u96 - for item in &mut bytes { - *item = src.read_u8().await?; - } + bytes.copy_from_slice(&buf); decode_d128(&bytes) as i128 * sign } 17 => { + // Bulk-read the 16 magnitude bytes in one packet-aware pass + // instead of 16 separate `read_u8().await` calls. + let mut buf = Vec::new(); + crate::sql_read_bytes::read_bytes_into(src, &mut buf, 16, 16).await?; let mut bytes = [0u8; 16]; - for item in &mut bytes { - *item = src.read_u8().await?; + bytes.copy_from_slice(&buf); + let magnitude = decode_d128(&bytes); + // A legal `decimal(38, s)` magnitude is < 10^38 < i128::MAX, + // so any 16-byte magnitude that does not fit in i128 is + // malformed. Reject it rather than letting `as i128` wrap to + // a negative value (and `i128::MIN * -1` overflow-panic). + if magnitude > i128::MAX as u128 { + return Err(Error::Protocol( + "decimal/numeric: magnitude exceeds the representable range".into(), + )); } - decode_d128(&bytes) as i128 * sign + magnitude as i128 * sign } x => { return Err(Error::Protocol( @@ -158,7 +176,9 @@ impl Numeric { impl Encode for Numeric { fn encode(self, dst: &mut BytesMut) -> crate::Result<()> { - dst.put_u8(self.len()); + // `len()` recomputes `precision()` via a division loop; compute it once. + let len = self.len(); + dst.put_u8(len); if self.value < 0 { dst.put_u8(0); @@ -166,16 +186,20 @@ impl Encode for Numeric { dst.put_u8(1); } - let value = self.value().abs(); + // The sign is written above; use `unsigned_abs()` for the magnitude so + // `i128::MIN` (whose two's-complement negation overflows) does not panic + // via `.abs()`. `i128::MIN` is reachable through the pub + // `new_with_scale` constructor and via an adversarial wire magnitude. + let value = self.value().unsigned_abs(); - match self.len() { + match len { 5 => dst.put_u32_le(value as u32), 9 => dst.put_u64_le(value as u64), 13 => { dst.put_u64_le(value as u64); dst.put_u32_le((value >> 64) as u32) } - _ => dst.put_u128_le(value as u128), + _ => dst.put_u128_le(value), } Ok(()) @@ -184,11 +208,18 @@ impl Encode for Numeric { impl Debug for Numeric { fn fmt(&self, f: &mut Formatter<'_>) -> Result<(), fmt::Error> { + // Use `unsigned_abs()` rather than `.abs()`: a server may send an + // adversarial magnitude that decodes to `i128::MIN` (or any value whose + // negation overflows), and `i128::abs()` panics ("attempt to negate with + // overflow") for `i128::MIN`. `unsigned_abs()` returns a `u128` and never + // overflows, so `Debug`-formatting is total for all i128 inputs while + // preserving the output for every in-range value. write!( f, - "{}.{:0pad$}", - self.int_part(), - self.dec_part(), + "{}{}.{:0pad$}", + if self.value() < 0 { "-" } else { "" }, + self.int_part().unsigned_abs(), + self.dec_part().unsigned_abs(), pad = self.scale as usize ) } @@ -263,8 +294,40 @@ mod decimal { Numeric::new_with_scale(value, self_.scale() as u8) }); ); + + #[cfg(feature = "tds73")] + into_sql!(self_, + Decimal: (ColumnData::Numeric, { + let unpacked = self_.unpack(); + + let mut value = (((unpacked.hi as u128) << 64) + + ((unpacked.mid as u128) << 32) + + unpacked.lo as u128) as i128; + + if self_.is_sign_negative() { + value = -value; + } + + Numeric::new_with_scale(value, self_.scale() as u8) + }); + ); } +/// `ToSql`/`IntoSql` conversions for [`bigdecimal::BigDecimal`]. +/// +/// # Limitation +/// +/// SQL Server's `NUMERIC`/`DECIMAL` mantissa fits in an `i128` with a scale of +/// at most [`Numeric::MAX_NUMERIC_SCALE`] (38). The `ToSql` and `IntoSql` +/// implementations below therefore **panic** if the `BigDecimal` mantissa +/// overflows `i128`, or if its scale exceeds 38 (via the internal `expect(..)` +/// calls, kept consistent with the `Numeric::new_with_scale` bound). This is a +/// documented +/// limitation of the current trait signatures, which return the column value +/// directly rather than a `Result`; callers holding arbitrarily large +/// `BigDecimal` values should range-check them before binding. Values produced +/// by round-tripping data that originated from SQL Server always fit and never +/// hit this path. #[cfg(feature = "bigdecimal")] mod bigdecimal_ { use super::{BigDecimal, BigInt, Numeric}; @@ -298,7 +361,9 @@ mod bigdecimal_ { let value = int.to_i128().expect("Given BigDecimal overflowing the maximum accepted value."); let scale = u8::try_from(std::cmp::max(exp, 0)) - .expect("Given BigDecimal exponent overflowing the maximum accepted scale (255)."); + .ok() + .filter(|s| *s <= Numeric::MAX_NUMERIC_SCALE) + .expect("Given BigDecimal exponent overflowing the maximum accepted scale (38)."); Numeric::new_with_scale(value, scale) }); @@ -322,7 +387,9 @@ mod bigdecimal_ { let value = int.to_i128().expect("Given BigDecimal overflowing the maximum accepted value."); let scale = u8::try_from(std::cmp::max(exp, 0)) - .expect("Given BigDecimal exponent overflowing the maximum accepted scale (255)."); + .ok() + .filter(|s| *s <= Numeric::MAX_NUMERIC_SCALE) + .expect("Given BigDecimal exponent overflowing the maximum accepted scale (38)."); Numeric::new_with_scale(value, scale) }); @@ -356,6 +423,94 @@ mod tests { ); } + #[test] + fn numeric_eq_normalizes_across_a_scale_gap() { + // 1.23 at scale 5 (123000) equals 1.23 at scale 2 (123). A scale gap of + // 3 is chosen so the `self.scale - other.scale` exponent (3) differs from + // both `+` (7) and `/` (1) — pinning the subtraction — and the + // `10^gap * v` multiply differs from `+`/`/`. Both comparison directions + // exercise the Greater and Less arms. + let wide = Numeric { + value: 123_000, + scale: 5, + }; + let narrow = Numeric { + value: 123, + scale: 2, + }; + assert_eq!(wide, narrow); // Greater arm (self.scale > other.scale) + assert_eq!(narrow, wide); // Less arm + assert!( + narrow + != Numeric { + value: 124, + scale: 2 + } + ); + } + + #[test] + fn encode_byte_layout_matches_length_bucket() { + // The encoder writes 1 length byte + 1 sign byte + (len-1) magnitude + // bytes. This pins the per-length arms (deleting the 9- or 13-byte arm + // would change the byte count) and the sign byte for zero. + for value in [1i128, 10i128.pow(12), 10i128.pow(20), 10i128.pow(30)] { + let n = Numeric::new_with_scale(value, 0); + let expected = n.len() as usize + 1; + let mut buf = BytesMut::new(); + n.encode(&mut buf).unwrap(); + assert_eq!(buf.len(), expected, "byte count for {value}"); + } + + // Zero is encoded as positive (sign byte 1), not negative. + let mut zero = BytesMut::new(); + Numeric::new_with_scale(0, 0).encode(&mut zero).unwrap(); + assert_eq!(zero[1], 1, "zero must carry the positive sign byte"); + } + + #[tokio::test] + async fn decode_d128_keeps_high_and_low_words() { + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + + // A magnitude whose high bytes are all non-zero: if decode_d128 wrongly + // short-circuited on "all high bytes non-zero" it would drop the high + // word and mis-decode. Positive (high byte 0x01 < i128::MAX high bit). + let value = 0x0101_0101_0101_0101_0101_0101_0101_0101i128; + let n = Numeric::new_with_scale(value, 0); + let mut buf = BytesMut::new(); + n.encode(&mut buf).unwrap(); + let decoded = Numeric::decode(&mut buf.into_sql_read_bytes(), 0) + .await + .unwrap() + .unwrap(); + assert_eq!(decoded.value(), value); + } + + #[tokio::test] + async fn decode_accepts_magnitude_at_i128_max_but_rejects_beyond() { + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + + // 17-byte form: len, sign(1 = positive), then 16 magnitude bytes. + let mut at_max = BytesMut::new(); + at_max.put_u8(17); + at_max.put_u8(1); + at_max.put_i128_le(i128::MAX); // magnitude exactly i128::MAX + let decoded = Numeric::decode(&mut at_max.into_sql_read_bytes(), 0) + .await + .expect("i128::MAX magnitude is representable") + .unwrap(); + assert_eq!(decoded.value(), i128::MAX); + + // One past i128::MAX (high bit set) must be rejected, not wrapped. + let mut beyond = BytesMut::new(); + beyond.put_u8(17); + beyond.put_u8(1); + beyond.put_u128_le((i128::MAX as u128) + 1); + assert!(Numeric::decode(&mut beyond.into_sql_read_bytes(), 0) + .await + .is_err()); + } + #[test] fn numeric_to_f64() { assert_eq!(f64::from(Numeric::new_with_scale(57705, 2)), 577.05); @@ -368,12 +523,257 @@ mod tests { assert_eq!(n.dec_part(), 5); } + #[test] + fn numeric_to_string() { + assert_eq!(Numeric::new_with_scale(123, 0).to_string(), "123.0"); + assert_eq!(Numeric::new_with_scale(123, 1).to_string(), "12.3"); + assert_eq!(Numeric::new_with_scale(123, 2).to_string(), "1.23"); + assert_eq!(Numeric::new_with_scale(123, 3).to_string(), "0.123"); + assert_eq!(Numeric::new_with_scale(123, 4).to_string(), "0.0123"); + assert_eq!( + Numeric::new_with_scale(123, 36).to_string(), + "0.000000000000000000000000000000000123" + ); + assert_eq!( + Numeric::new_with_scale(123, 37).to_string(), + "0.0000000000000000000000000000000000123" + ); + assert_eq!(Numeric::new_with_scale(-123, 0).to_string(), "-123.0"); + assert_eq!(Numeric::new_with_scale(-123, 1).to_string(), "-12.3"); + assert_eq!(Numeric::new_with_scale(-123, 2).to_string(), "-1.23"); + assert_eq!(Numeric::new_with_scale(-123, 3).to_string(), "-0.123"); + assert_eq!(Numeric::new_with_scale(-123, 4).to_string(), "-0.0123"); + assert_eq!( + Numeric::new_with_scale(-123, 36).to_string(), + "-0.000000000000000000000000000000000123" + ); + assert_eq!( + Numeric::new_with_scale(-123, 37).to_string(), + "-0.0000000000000000000000000000000000123" + ); + } + + // An adversarial server can send a 17-byte NUMERIC magnitude that + // `Numeric::decode` casts `as i128` into `i128::MIN` (whose two's-complement + // negation overflows). `Debug`/`Display` must not panic on such a value. + #[test] + fn debug_does_not_panic_on_i128_min() { + for scale in [0u8, 2, 37] { + let n = Numeric { + value: i128::MIN, + scale, + }; + // Both must produce *some* string without panicking on `.abs()`. + let _ = format!("{:?}", n); + let _ = format!("{}", n); + } + } + + // A value just below 2^127 also wraps negative when cast `as i128`; formatting + // it must likewise be total. + #[test] + fn debug_does_not_panic_near_2_pow_127() { + // (2^127 - 1) reinterpreted as i128 is i128::MAX; (2^127) wraps to i128::MIN. + // Exercise a spread of large-magnitude values around the boundary. + for value in [i128::MAX, i128::MIN, i128::MIN + 1, i128::MAX - 1] { + for scale in [0u8, 5, 37] { + let n = Numeric { value, scale }; + let _ = format!("{:?}", n); + } + } + } + + #[test] + fn encode_does_not_panic_on_i128_min() { + // `i128::MIN.abs()` overflows ("attempt to negate with overflow"), so + // `encode` must use `unsigned_abs()`. A `Numeric` holding `i128::MIN` is + // reachable via the pub `new_with_scale` constructor and via an + // adversarial 17-byte wire magnitude. + let mut buf = BytesMut::new(); + Numeric::new_with_scale(i128::MIN, 0) + .encode(&mut buf) + .expect("encode of i128::MIN must not panic"); + + // Sign byte is negative (0) and the magnitude is 2^127 little-endian. + assert_eq!(buf[0], 17, "i128::MIN needs the 17-byte length bucket"); + assert_eq!(buf[1], 0, "i128::MIN must carry the negative sign byte"); + let mut expected = BytesMut::new(); + expected.put_u128_le(i128::MIN.unsigned_abs()); + assert_eq!(&buf[2..], &expected[..]); + } + + #[test] + fn max_numeric_scale_is_38() { + // The public scale ceiling must stay consistent between the + // `new_with_scale` assert and the bigdecimal bound check. + assert_eq!(Numeric::MAX_NUMERIC_SCALE, 38); + assert_eq!( + Numeric::new_with_scale(1, Numeric::MAX_NUMERIC_SCALE).scale(), + 38 + ); + } + #[test] fn calculates_precision_correctly() { let n = Numeric::new_with_scale(57705, 2); assert_eq!(5, n.precision()); } + #[test] + fn new_with_scale_accessors() { + let n = Numeric::new_with_scale(12345, 3); + assert_eq!(n.value(), 12345); + assert_eq!(n.scale(), 3); + assert_eq!(n.int_part(), 12); + assert_eq!(n.dec_part(), 345); + } + + #[test] + fn new_with_scale_allows_max_scale() { + // decimal(38, 38) is valid in SQL Server, so scale 38 must be accepted. + assert_eq!(Numeric::new_with_scale(1, 38).scale(), 38); + } + + #[test] + #[should_panic(expected = "scale <= Self::MAX_NUMERIC_SCALE")] + fn new_with_scale_panics_on_too_large_scale() { + Numeric::new_with_scale(1, 39); + } + + #[test] + fn precision_with_zero_int_part() { + // int_part == 0 -> precision is 1 + scale. + let n = Numeric::new_with_scale(5, 2); + assert_eq!(n.int_part(), 0); + assert_eq!(n.precision(), 3); + } + + #[test] + fn precision_scaling_by_length_buckets() { + assert_eq!(Numeric::new_with_scale(1, 0).len(), 5); + assert_eq!(Numeric::new_with_scale(1_000_000_000, 0).len(), 9); + assert_eq!(Numeric::new_with_scale(10i128.pow(19), 0).len(), 13); + assert_eq!(Numeric::new_with_scale(10i128.pow(28), 0).len(), 17); + } + + #[test] + fn display_and_debug() { + let n = Numeric::new_with_scale(57705, 2); + assert_eq!(format!("{:?}", n), "577.05"); + assert_eq!(format!("{}", n), "577.05"); + + // Negative values format with a single leading sign and an unsigned + // fractional part (see #390). + let n = Numeric::new_with_scale(-57705, 3); + assert_eq!(format!("{}", n), "-57.705"); + + // Zero-padded fractional part for small decimals. + let n = Numeric::new_with_scale(102, 4); + assert_eq!(format!("{}", n), "0.0102"); + } + + #[test] + fn from_numeric_conversions() { + let n = Numeric::new_with_scale(57705, 2); + assert_eq!(i128::from(n), 577); + assert_eq!(u128::from(n), 577); + assert!((f64::from(n) - 577.05).abs() < f64::EPSILON); + } + + #[test] + fn eq_across_scales_negative() { + assert_eq!( + Numeric::new_with_scale(-100501, 2), + Numeric::new_with_scale(-1005010, 3), + ); + } + + async fn round_trip(value: i128, scale: u8) { + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + + let n = Numeric::new_with_scale(value, scale); + let mut buf = BytesMut::new(); + n.encode(&mut buf).expect("encode must succeed"); + + let decoded = Numeric::decode(&mut buf.into_sql_read_bytes(), scale) + .await + .expect("decode must succeed") + .expect("value must be present"); + + assert_eq!(decoded, n); + assert_eq!(decoded.value(), value); + } + + #[tokio::test] + async fn encode_decode_round_trip() { + round_trip(0, 0).await; // len 5 + round_trip(42, 0).await; // len 5 + round_trip(-42, 2).await; // negative, len 5 + round_trip(10i128.pow(12), 0).await; // len 9 + round_trip(10i128.pow(20), 0).await; // len 13 + round_trip(-(10i128.pow(20)), 3).await; // negative, len 13 + round_trip(10i128.pow(30), 0).await; // len 17 + } + + #[tokio::test] + async fn decode_zero_length_is_none() { + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + + let mut buf = BytesMut::new(); + buf.put_u8(0); + + let decoded = Numeric::decode(&mut buf.into_sql_read_bytes(), 0) + .await + .expect("decode must succeed"); + + assert!(decoded.is_none()); + } + + #[tokio::test] + async fn decode_rejects_len17_magnitude_over_i128_max() { + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + + // len = 17, sign = 1 (positive), magnitude = 2^127 (byte[15] = 0x80), + // which exceeds i128::MAX. Must return a protocol error rather than + // wrapping to a negative value (or panicking on i128::MIN * -1). + let mut buf = BytesMut::new(); + buf.put_u8(17); + buf.put_u8(1); + let mut mag = [0u8; 16]; + mag[15] = 0x80; + buf.extend_from_slice(&mag); + + let err = Numeric::decode(&mut buf.into_sql_read_bytes(), 0) + .await + .expect_err("out-of-range magnitude must error"); + assert!(matches!(err, Error::Protocol(_))); + } + + #[tokio::test] + async fn decode_rejects_invalid_sign_and_length() { + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + + // Invalid sign byte (2 is neither 0 nor 1). + let mut buf = BytesMut::new(); + buf.put_u8(5); + buf.put_u8(2); + buf.put_u32_le(1); + let err = Numeric::decode(&mut buf.into_sql_read_bytes(), 0) + .await + .expect_err("invalid sign must error"); + assert!(matches!(err, Error::Protocol(_))); + + // Invalid length byte (6 is not one of 0/5/9/13/17). + let mut buf = BytesMut::new(); + buf.put_u8(6); + buf.put_u8(1); + buf.extend_from_slice(&[0u8; 4]); + let err = Numeric::decode(&mut buf.into_sql_read_bytes(), 0) + .await + .expect_err("invalid length must error"); + assert!(matches!(err, Error::Protocol(_))); + } + #[test] #[cfg(feature = "bigdecimal")] fn no_overflowing_pow() { diff --git a/src/tds/stream.rs b/src/tds/stream.rs index e6454876b..d01e6b9b9 100644 --- a/src/tds/stream.rs +++ b/src/tds/stream.rs @@ -1,5 +1,7 @@ +mod command; mod query; mod token; +pub use command::*; pub use query::*; pub use token::*; diff --git a/src/tds/stream/command.rs b/src/tds/stream/command.rs new file mode 100644 index 000000000..6ea2c3378 --- /dev/null +++ b/src/tds/stream/command.rs @@ -0,0 +1,398 @@ +use crate::tds::stream::ReceivedToken; +use crate::{row::ColumnType, Column, Row}; +use crate::{ColumnData, CommandResult, ResultMetadata}; +use futures_util::{ + ready, + stream::{BoxStream, Peekable, Stream, StreamExt, TryStreamExt}, +}; +use std::{ + fmt::Debug, + pin::Pin, + sync::Arc, + task::{self, Poll}, +}; + +/// A `Stream` of [`CommandItem`] values produced by executing a [`Command`]. +/// +/// Items can be result metadata, rows, a return status, return values (OUT +/// parameters) or a rows-affected count. +/// +/// [`Command`]: crate::Command +/// +/// # Example +/// +/// ```no_run +/// # use std::env; +/// # use tiberius::Config; +/// # use tiberius::{Command, CommandItem}; +/// # use futures_util::TryStreamExt; +/// # use tokio_util::compat::TokioAsyncWriteCompatExt; +/// # #[tokio::main] +/// # async fn main() -> Result<(), Box> { +/// # let c_str = env::var("TIBERIUS_TEST_CONNECTION_STRING").unwrap_or( +/// # "server=tcp:localhost,1433;integratedSecurity=true;TrustServerCertificate=true".to_owned(), +/// # ); +/// # let config = Config::from_ado_string(&c_str)?; +/// # let tcp = tokio::net::TcpStream::connect(config.get_addr()).await?; +/// # tcp.set_nodelay(true)?; +/// # let mut client = tiberius::Client::connect(config, tcp.compat_write()).await?; +/// let mut cmd = Command::new("dbo.usp_SomeStoredProc"); +/// +/// cmd.bind_param("@foo", 34i32); +/// cmd.bind_param("@zoo", "the zoo string prm"); +/// cmd.bind_out_param("@bar", "bar"); +/// let mut stream = cmd.exec(&mut client).await?; +/// +/// while let Some(item) = stream.try_next().await? { +/// match item { +/// // our first item is the column data always +/// CommandItem::Metadata(meta) if meta.result_index() == 0 => { +/// // the first result column info can be handled here +/// } +/// // ... and from there on from 0..N rows +/// CommandItem::Row(row) if row.result_index() == 0 => { +/// let var: Option = row.get(0); +/// } +/// // the second result set returns first another metadata item +/// CommandItem::Metadata(meta) => { +/// // .. handling +/// } +/// // ...and, again, we get rows from the second resultset +/// CommandItem::Row(row) => { +/// let var: Option = row.get(0); +/// } +/// // check return status (returned always) +/// CommandItem::ReturnStatus(rs) => { +/// // .... do something +/// } +/// // collect OUT parameter values +/// CommandItem::ReturnValue(rv) => { +/// // .... do something, like push to a collection +/// } +/// // get affected row count +/// CommandItem::RowsAffected(ra) => { +/// // .... do something, like push to a collection +/// } +/// } +/// } +/// # Ok(()) +/// # } +/// ``` +/// +pub struct CommandStream<'a> { + token_stream: Peekable>>, + columns: Option>>, + result_set_index: Option, +} + +impl<'a> Debug for CommandStream<'a> { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("CommandStream") + .field( + "token_stream", + &"BoxStream<'a, crate::Result>", + ) + .finish() + } +} + +impl<'a> CommandStream<'a> { + pub(crate) fn new(token_stream: BoxStream<'a, crate::Result>) -> Self { + Self { + token_stream: token_stream.peekable(), + columns: None, + result_set_index: None, + } + } + + /// Collects all results from the command into memory, in the order they + /// were produced by the server. + pub async fn into_command_result(mut self) -> crate::Result { + let mut results: Vec> = Vec::new(); + let mut result: Option> = None; + let mut return_status = 0; + let mut return_values = Vec::new(); + let mut rows_affected = Vec::new(); + + while let Some(item) = self.try_next().await? { + match (item, &mut result) { + (CommandItem::Row(row), None) => { + result = Some(vec![row]); + } + (CommandItem::Row(row), Some(ref mut result)) => result.push(row), + (CommandItem::Metadata(_), previous_result) => { + // A new result set begins. Flush the previous one (if any) + // and open a fresh, empty set so a trailing zero-row result + // still produces an (empty) `Vec` at the correct position + // rather than being silently dropped (matching the sibling + // `stream::query::collect_results`). + if let Some(previous) = previous_result.take() { + results.push(previous); + } + *previous_result = Some(Vec::new()); + } + (CommandItem::ReturnStatus(rs), _) => return_status = rs, + (CommandItem::ReturnValue(rv), _) => return_values.push(rv), + (CommandItem::RowsAffected(rows), _) => rows_affected.push(rows), + } + } + + if let Some(result) = result { + results.push(result); + } + + Ok(CommandResult { + return_code: return_status, + return_values, + query_results: results, + rows_affected, + }) + } + + /// Converts the stream into a stream of rows, dropping all other items. + pub fn into_row_stream(self) -> BoxStream<'a, crate::Result> { + let s = self.try_filter_map(|item| async { + match item { + CommandItem::Row(row) => Ok(Some(row)), + _ => Ok(None), + } + }); + + Box::pin(s) + } +} + +/// A single OUT parameter value returned by a [`Command`]. +/// +/// [`Command`]: crate::Command +#[derive(Debug)] +pub struct CommandReturnValue { + pub(crate) name: String, + pub(crate) ord: u16, + pub(crate) data: ColumnData<'static>, +} + +impl CommandReturnValue { + /// The name of the OUT parameter this value corresponds to. + pub fn name(&self) -> &str { + &self.name + } + + /// The ordinal of the OUT parameter as returned by the server. + pub fn ordinal(&self) -> u16 { + self.ord + } + + /// A reference to the raw column data of the returned value. + pub fn data(&self) -> &ColumnData<'static> { + &self.data + } +} + +/// An item produced by a [`CommandStream`]. +#[derive(Debug)] +pub enum CommandItem { + /// A single row of data. + Row(Row), + /// Metadata describing the upcoming rows. + Metadata(ResultMetadata), + /// The return status from the server. + ReturnStatus(u32), + /// A return value, matching an OUT parameter. + ReturnValue(CommandReturnValue), + /// The number of rows affected by one of the statements ran on the server. + RowsAffected(u64), +} + +impl CommandItem { + pub(crate) fn metadata(columns: Arc>, result_index: usize) -> Self { + Self::Metadata(ResultMetadata { + columns, + result_index, + }) + } + + /// Returns a reference to the metadata, if the item is of a correct variant. + pub fn as_metadata(&self) -> Option<&ResultMetadata> { + match self { + CommandItem::Metadata(ref metadata) => Some(metadata), + _ => None, + } + } + + /// Returns a reference to the row, if the item is of a correct variant. + pub fn as_row(&self) -> Option<&Row> { + match self { + CommandItem::Row(ref row) => Some(row), + _ => None, + } + } + + /// Returns the metadata, if the item is of a correct variant. + pub fn into_metadata(self) -> Option { + match self { + CommandItem::Metadata(metadata) => Some(metadata), + _ => None, + } + } + + /// Returns the row, if the item is of a correct variant. + pub fn into_row(self) -> Option { + match self { + CommandItem::Row(row) => Some(row), + _ => None, + } + } + + /// Returns the return status, if the item is of a correct variant. + pub fn as_return_status(&self) -> Option { + match self { + CommandItem::ReturnStatus(rs) => Some(*rs), + _ => None, + } + } + + /// Returns a reference to the return value, if the item is of a correct variant. + pub fn as_return_value(&self) -> Option<&CommandReturnValue> { + match self { + CommandItem::ReturnValue(rv) => Some(rv), + _ => None, + } + } + + /// Returns the return value, if the item is of a correct variant. + pub fn into_return_value(self) -> Option { + match self { + CommandItem::ReturnValue(rv) => Some(rv), + _ => None, + } + } +} + +impl<'a> Stream for CommandStream<'a> { + type Item = crate::Result; + + fn poll_next(self: Pin<&mut Self>, cx: &mut task::Context<'_>) -> Poll> { + let this = self.get_mut(); + + loop { + let token = match ready!(this.token_stream.poll_next_unpin(cx)) { + Some(res) => res?, + None => return Poll::Ready(None), + }; + + return match token { + ReceivedToken::NewResultset(meta) => { + let column_meta = meta + .columns + .iter() + .map(|x| Column { + name: x.col_name.to_string(), + column_type: ColumnType::from(&x.base.ty), + }) + .collect::>(); + + let column_meta = Arc::new(column_meta); + this.columns = Some(column_meta.clone()); + + this.result_set_index = this.result_set_index.map(|i| i + 1); + + let query_item = + CommandItem::metadata(column_meta, *this.result_set_index.get_or_insert(0)); + + Poll::Ready(Some(Ok(query_item))) + } + ReceivedToken::Row(data) => { + let Some(columns) = this.columns.as_ref() else { + return Poll::Ready(Some(Err(crate::Error::Protocol( + "ROW token arrived before any column metadata".into(), + )))); + }; + let columns = columns.clone(); + let result_index = this.result_set_index.unwrap_or(0); + + let row = Row { + columns, + data, + result_index, + }; + + Poll::Ready(Some(Ok(CommandItem::Row(row)))) + } + ReceivedToken::ReturnStatus(rs) => { + Poll::Ready(Some(Ok(CommandItem::ReturnStatus(rs)))) + } + ReceivedToken::ReturnValue(rv) => { + Poll::Ready(Some(Ok(CommandItem::ReturnValue(CommandReturnValue { + name: rv.param_name, + ord: rv.param_ordinal, + data: rv.value, + })))) + } + ReceivedToken::DoneProc(done) if done.is_final() => continue, + ReceivedToken::DoneProc(done) => { + Poll::Ready(Some(Ok(CommandItem::RowsAffected(done.rows())))) + } + ReceivedToken::DoneInProc(done) => { + Poll::Ready(Some(Ok(CommandItem::RowsAffected(done.rows())))) + } + ReceivedToken::Done(done) => { + Poll::Ready(Some(Ok(CommandItem::RowsAffected(done.rows())))) + } + _ => continue, + }; + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::tds::codec::{ + BaseMetaDataColumn, FixedLenType, MetaDataColumn, TokenColMetaData, TokenDone, TypeInfo, + }; + use crate::tds::stream::ReceivedToken; + use futures_util::stream::{self, StreamExt}; + use std::borrow::Cow; + + fn token_meta(name: &'static str) -> ReceivedToken { + let col = MetaDataColumn { + base: BaseMetaDataColumn { + flags: enumflags2::BitFlags::empty(), + ty: TypeInfo::FixedLen(FixedLenType::Int4), + table_name: None, + }, + col_name: Cow::Borrowed(name), + }; + ReceivedToken::NewResultset(Arc::new(TokenColMetaData { columns: vec![col] })) + } + + fn token_row() -> ReceivedToken { + ReceivedToken::Row(crate::tds::codec::TokenRow::new()) + } + + fn command_stream(tokens: Vec) -> CommandStream<'static> { + let s = stream::iter(tokens.into_iter().map(Ok::<_, crate::Error>)); + CommandStream::new(s.boxed()) + } + + #[tokio::test] + async fn trailing_empty_result_set_is_preserved() { + // `SELECT 1; SELECT TOP 0 * FROM t` shape: a first set with one row, + // then a second metadata token with no rows before DONE. The trailing + // empty result set must survive as an empty `Vec`, not be dropped. + let result = command_stream(vec![ + token_meta("first"), + token_row(), + token_meta("second"), + ReceivedToken::Done(TokenDone::default()), + ]) + .into_command_result() + .await + .expect("into_command_result"); + + assert_eq!(result.query_results.len(), 2); + assert_eq!(result.query_results[0].len(), 1); + assert!(result.query_results[1].is_empty()); + } +} diff --git a/src/tds/stream/query.rs b/src/tds/stream/query.rs index 0dc694749..5e88e8de0 100644 --- a/src/tds/stream/query.rs +++ b/src/tds/stream/query.rs @@ -220,31 +220,8 @@ impl<'a> QueryStream<'a> { /// Collects results from all queries in the stream into memory in the order /// of querying. - pub async fn into_results(mut self) -> crate::Result>> { - let mut results: Vec> = Vec::new(); - let mut result: Option> = None; - - while let Some(item) = self.try_next().await? { - match (item, &mut result) { - (QueryItem::Row(row), None) => { - result = Some(vec![row]); - } - (QueryItem::Row(row), Some(ref mut result)) => result.push(row), - (QueryItem::Metadata(_), None) => { - result = Some(Vec::new()); - } - (QueryItem::Metadata(_), ref mut previous_result) => { - results.push(previous_result.take().unwrap()); - result = None; - } - } - } - - if let Some(result) = result { - results.push(result); - } - - Ok(results) + pub async fn into_results(self) -> crate::Result>> { + collect_results(self).await } /// Collects the output of the first query, dropping any further @@ -277,11 +254,59 @@ impl<'a> QueryStream<'a> { } } +/// Collect a stream of [`QueryItem`]s into one `Vec` per result set, in +/// stream order. +/// +/// Each result set is delimited by a [`QueryItem::Metadata`] item: a metadata +/// item opens a new (initially empty) result set and every subsequent +/// [`QueryItem::Row`] is appended to it. Nothing is discarded on the assumption +/// that the first item is metadata — if a row were ever to arrive before any +/// metadata it is still captured into a result set rather than silently dropped. +/// +/// Behaviour preserved from the original implementation: +/// - an empty stream yields `Ok(vec![])`; +/// - a result set with zero rows still yields an (empty) inner `Vec` in the +/// right position (metadata with no following rows -> one empty result set); +/// - ordering across multiple result sets is preserved. +async fn collect_results(mut stream: S) -> crate::Result>> +where + S: Stream> + Unpin, +{ + let mut results: Vec> = Vec::new(); + let mut current: Option> = None; + + while let Some(item) = stream.try_next().await? { + match item { + QueryItem::Metadata(_) => { + // A new result set begins. Flush the previous one (if any) and + // open a fresh, empty set so a zero-row result still produces an + // inner Vec at the correct position. + if let Some(previous) = current.take() { + results.push(previous); + } + current = Some(Vec::new()); + } + QueryItem::Row(row) => { + // Rows are always preceded by their metadata in a well-formed + // stream, so `current` is normally `Some`. Be defensive and open + // a result set on the fly rather than discard a leading row. + current.get_or_insert_with(Vec::new).push(row); + } + } + } + + if let Some(last) = current.take() { + results.push(last); + } + + Ok(results) +} + /// Info about the following stream of rows. #[derive(Debug, Clone)] pub struct ResultMetadata { - columns: Arc>, - result_index: usize, + pub(crate) columns: Arc>, + pub(crate) result_index: usize, } impl ResultMetadata { @@ -382,8 +407,13 @@ impl<'a> Stream for QueryStream<'a> { return Poll::Ready(Some(Ok(query_item))); } ReceivedToken::Row(data) => { - let columns = this.columns.as_ref().unwrap().clone(); - let result_index = this.result_set_index.unwrap(); + let Some(columns) = this.columns.as_ref() else { + return Poll::Ready(Some(Err(crate::Error::Protocol( + "ROW token arrived before any column metadata".into(), + )))); + }; + let columns = columns.clone(); + let result_index = this.result_set_index.unwrap_or(0); let row = Row { columns, @@ -398,3 +428,277 @@ impl<'a> Stream for QueryStream<'a> { } } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::row::ColumnType; + use crate::tds::codec::TokenRow; + use crate::Column; + use futures_util::stream; + + fn columns() -> Arc> { + Arc::new(vec![Column::new("c".to_string(), ColumnType::Int4)]) + } + + fn meta(cols: &Arc>, result_index: usize) -> QueryItem { + QueryItem::metadata(cols.clone(), result_index) + } + + fn row(cols: &Arc>, result_index: usize) -> QueryItem { + QueryItem::Row(Row { + columns: cols.clone(), + data: TokenRow::new(), + result_index, + }) + } + + async fn collect(items: Vec) -> Vec> { + let stream = stream::iter(items.into_iter().map(Ok::<_, crate::error::Error>)); + collect_results(stream).await.expect("collect_results") + } + + #[tokio::test] + async fn empty_stream_yields_no_result_sets() { + assert!(collect(vec![]).await.is_empty()); + } + + #[tokio::test] + async fn metadata_without_rows_yields_one_empty_result_set() { + let cols = columns(); + let results = collect(vec![meta(&cols, 0)]).await; + assert_eq!(results.len(), 1); + assert!(results[0].is_empty()); + } + + #[tokio::test] + async fn multiple_empty_result_sets_are_preserved() { + let cols = columns(); + let results = collect(vec![meta(&cols, 0), meta(&cols, 1)]).await; + assert_eq!(results.len(), 2); + assert!(results[0].is_empty()); + assert!(results[1].is_empty()); + } + + #[tokio::test] + async fn rows_are_grouped_by_result_set_in_order() { + let cols = columns(); + let results = collect(vec![ + meta(&cols, 0), + row(&cols, 0), + row(&cols, 0), + meta(&cols, 1), + row(&cols, 1), + ]) + .await; + + assert_eq!(results.len(), 2); + assert_eq!(results[0].len(), 2); + assert_eq!(results[1].len(), 1); + } + + #[tokio::test] + async fn leading_row_is_not_discarded() { + // The old implementation consumed and threw away the first stream item, + // assuming it was metadata. If a row led, it was silently lost. The + // robust implementation keeps every row: a stream of two leading rows + // must produce one result set containing BOTH rows (the old logic would + // have produced a single-row set). + let cols = columns(); + let results = collect(vec![row(&cols, 0), row(&cols, 0)]).await; + assert_eq!(results.len(), 1); + assert_eq!(results[0].len(), 2); + } + + #[tokio::test] + async fn single_metadata_with_rows_yields_one_result_set_of_n_rows() { + // One metadata item followed by N rows must collapse into exactly one + // result set holding all N rows. + let cols = columns(); + let results = collect(vec![ + meta(&cols, 0), + row(&cols, 0), + row(&cols, 0), + row(&cols, 0), + ]) + .await; + assert_eq!(results.len(), 1); + assert_eq!(results[0].len(), 3); + } + + #[tokio::test] + async fn collect_propagates_errors() { + // An error partway through the stream must surface rather than be + // silently swallowed by the collection loop. + let cols = columns(); + let items = vec![ + Ok(meta(&cols, 0)), + Ok(row(&cols, 0)), + Err(crate::Error::Protocol("boom".into())), + ]; + let stream = stream::iter(items); + let err = collect_results(stream).await.expect_err("expected error"); + assert!(matches!(err, crate::Error::Protocol(_))); + } + + // --- QueryItem accessor helpers ------------------------------------- + + #[test] + fn query_item_accessors_on_metadata() { + let cols = columns(); + let item = meta(&cols, 7); + + assert!(item.as_metadata().is_some()); + assert!(item.as_row().is_none()); + assert_eq!(item.as_metadata().unwrap().result_index(), 7); + + let md = item.into_metadata().expect("metadata"); + assert_eq!(md.result_index(), 7); + assert_eq!(md.columns().len(), 1); + } + + #[test] + fn query_item_accessors_on_row() { + let cols = columns(); + let item = row(&cols, 3); + + assert!(item.as_row().is_some()); + assert!(item.as_metadata().is_none()); + assert_eq!(item.as_row().unwrap().result_index, 3); + + let r = item.into_row().expect("row"); + assert_eq!(r.result_index, 3); + } + + #[test] + fn query_item_into_wrong_variant_is_none() { + let cols = columns(); + assert!(meta(&cols, 0).into_row().is_none()); + assert!(row(&cols, 0).into_metadata().is_none()); + } + + // --- Full QueryStream state machine (poll_next + into_* helpers) ----- + // + // These build a `QueryStream` from an in-memory stream of `ReceivedToken`s + // (no server), exercising `poll_next` (result-set indexing, the + // ROW-before-metadata guard) and the `into_results`/`into_first_result`/ + // `into_row`/`into_row_stream` helpers end to end. + + use crate::tds::codec::{ + BaseMetaDataColumn, FixedLenType, MetaDataColumn, TokenColMetaData, TypeInfo, + }; + use crate::tds::stream::ReceivedToken; + use futures_util::stream::{StreamExt, TryStreamExt}; + use std::borrow::Cow; + + fn token_meta(name: &'static str) -> ReceivedToken { + let col = MetaDataColumn { + base: BaseMetaDataColumn { + flags: enumflags2::BitFlags::empty(), + ty: TypeInfo::FixedLen(FixedLenType::Int4), + table_name: None, + }, + col_name: Cow::Borrowed(name), + }; + ReceivedToken::NewResultset(Arc::new(TokenColMetaData { columns: vec![col] })) + } + + fn token_row() -> ReceivedToken { + ReceivedToken::Row(TokenRow::new()) + } + + fn query_stream(tokens: Vec) -> QueryStream<'static> { + let s = stream::iter(tokens.into_iter().map(Ok::<_, crate::Error>)); + QueryStream::new(s.boxed()) + } + + #[tokio::test] + async fn query_stream_into_results_groups_and_indexes_result_sets() { + let stream = query_stream(vec![ + token_meta("first"), + token_row(), + token_row(), + token_meta("second"), + token_row(), + ]); + + let results = stream.into_results().await.expect("into_results"); + assert_eq!(results.len(), 2); + assert_eq!(results[0].len(), 2); + assert_eq!(results[1].len(), 1); + + // result_index must increment across result sets. + assert!(results[0].iter().all(|r| r.result_index == 0)); + assert!(results[1].iter().all(|r| r.result_index == 1)); + } + + #[tokio::test] + async fn query_stream_empty_yields_no_results() { + let results = query_stream(vec![]).into_results().await.expect("empty"); + assert!(results.is_empty()); + } + + #[tokio::test] + async fn query_stream_metadata_only_yields_one_empty_result_set() { + let results = query_stream(vec![token_meta("first")]) + .into_results() + .await + .expect("metadata only"); + assert_eq!(results.len(), 1); + assert!(results[0].is_empty()); + } + + #[tokio::test] + async fn query_stream_row_before_metadata_is_protocol_error() { + // The rewritten poll_next must reject a ROW token that arrives before + // any column metadata rather than panic or silently drop it. + let err = query_stream(vec![token_row()]) + .into_results() + .await + .expect_err("expected protocol error"); + assert!(matches!(err, crate::Error::Protocol(_))); + } + + #[tokio::test] + async fn query_stream_into_first_result_and_into_row() { + let first = query_stream(vec![ + token_meta("first"), + token_row(), + token_row(), + token_meta("second"), + token_row(), + ]) + .into_first_result() + .await + .expect("into_first_result"); + assert_eq!(first.len(), 2); + + let one = query_stream(vec![token_meta("first"), token_row(), token_row()]) + .into_row() + .await + .expect("into_row"); + assert!(one.is_some()); + + let none = query_stream(vec![token_meta("first")]) + .into_row() + .await + .expect("into_row empty"); + assert!(none.is_none()); + } + + #[tokio::test] + async fn query_stream_into_row_stream_skips_metadata() { + let rows: Vec = query_stream(vec![ + token_meta("first"), + token_row(), + token_meta("second"), + token_row(), + token_row(), + ]) + .into_row_stream() + .try_collect() + .await + .expect("into_row_stream"); + assert_eq!(rows.len(), 3); + } +} diff --git a/src/tds/stream/token.rs b/src/tds/stream/token.rs index 87c343174..1176b4b2d 100644 --- a/src/tds/stream/token.rs +++ b/src/tds/stream/token.rs @@ -2,34 +2,177 @@ use crate::tds::codec::TokenSspi; use crate::{ client::Connection, tds::codec::{ - TokenColMetaData, TokenDone, TokenEnvChange, TokenError, TokenFeatureExtAck, TokenInfo, - TokenLoginAck, TokenOrder, TokenReturnValue, TokenRow, + TokenAltMetaData, TokenAltRow, TokenColInfo, TokenColMetaData, TokenDone, TokenEnvChange, + TokenError, TokenFeatureExtAck, TokenFedAuthInfo, TokenInfo, TokenLoginAck, TokenOrder, + TokenReturnValue, TokenRow, TokenSessionState, TokenTabName, }, Error, SqlReadBytes, TokenType, }; use futures_util::{ io::{AsyncRead, AsyncWrite}, - stream::{BoxStream, TryStreamExt}, + stream::{BoxStream, Stream, StreamExt, TryStreamExt}, }; +use std::future::Future; +use std::pin::Pin; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::task::{Context as TaskContext, Poll}; +use std::time::Duration; use std::{convert::TryFrom, sync::Arc}; use tracing::{event, Level}; +/// Wraps a token stream so each server round-trip is bounded by a deadline. +/// +/// The timer only runs while an item is *in flight* — i.e. while the inner +/// stream is `Pending` waiting on the server — and is reset the moment a token +/// is delivered. A slow consumer (which simply stops polling between items) +/// therefore never trips it; only a stalled server does. This backs +/// [`Config::command_timeout`](crate::Config::command_timeout), whose rustdoc +/// documents the exact semantics. +struct RoundTripTimeout<'a> { + inner: BoxStream<'a, crate::Result>, + /// `None` disables the bound entirely (unbounded reads). + timeout: Option, + /// The deadline for the token currently being awaited. Created lazily when + /// the inner stream first parks on the server and cleared on every delivery, + /// so it measures per-round-trip stall rather than consumer pace. + delay: Option, + /// Shared with the owning [`Connection`], which flips it when this deadline + /// fires so the now-desynced connection is rejected on any subsequent use + /// (its `flush_stream` would otherwise block forever on the still-pending + /// previous response). `None` in the combinator's own unit tests, which run + /// without a connection. + /// + /// [`Connection`]: crate::client::Connection + poison: Option>, +} + +impl<'a> RoundTripTimeout<'a> { + fn new( + timeout: Option, + poison: Option>, + inner: BoxStream<'a, crate::Result>, + ) -> Self { + Self { + inner, + timeout, + delay: None, + poison, + } + } +} + +impl Stream for RoundTripTimeout<'_> { + type Item = crate::Result; + + fn poll_next(self: Pin<&mut Self>, cx: &mut TaskContext<'_>) -> Poll> { + // Both `inner` (a `BoxStream`) and `Delay` are `Unpin`, so they can be + // polled through `&mut` without structural pin projection. + let this = self.get_mut(); + + match this.inner.poll_next_unpin(cx) { + // A token arrived (or the stream ended / errored): reset the deadline + // so the next round-trip is timed afresh. + Poll::Ready(item) => { + this.delay = None; + Poll::Ready(item) + } + Poll::Pending => { + let Some(timeout) = this.timeout else { + // Unbounded: never start a timer. + return Poll::Pending; + }; + + let delay = this + .delay + .get_or_insert_with(|| futures_timer::Delay::new(timeout)); + + match Pin::new(delay).poll(cx) { + Poll::Ready(()) => { + // Clear the fired timer and mark the connection desynced: + // the tail of this response is still unread on the wire, + // so it must not be reused. Flipping the shared flag makes + // `Connection::ensure_not_poisoned` reject the next use + // (fast, deterministic) instead of `flush_stream` blocking + // forever on the still-pending previous response. + this.delay = None; + if let Some(poison) = &this.poison { + poison.store(true, Ordering::Release); + } + Poll::Ready(Some(Err(crate::Error::Io { + kind: std::io::ErrorKind::TimedOut, + message: format!( + "the server did not send the next part of the command result \ + within {timeout:?} (it may have stalled, or be responding too \ + slowly for the configured bound); the connection is now out of \ + sync and must be dropped. Adjust the bound with \ + `Config::command_timeout`, or pass `None` to wait indefinitely." + ), + }))) + } + Poll::Pending => Poll::Pending, + } + } + } + } +} + #[derive(Debug)] -#[allow(dead_code)] pub enum ReceivedToken { NewResultset(Arc>), + // The ALTMETADATA is stashed in the connection context on decode (see + // `get_alt_col_metadata`); this payload copy is not otherwise read. + #[allow(dead_code)] + NewAltResultset(Arc>), Row(TokenRow<'static>), + // COMPUTE-clause rows are decoded and traced but not surfaced through the + // public result API. TODO(surface): expose ALTROW data to callers. + #[allow(dead_code)] + AltRow(TokenAltRow<'static>), Done(TokenDone), DoneInProc(TokenDone), DoneProc(TokenDone), ReturnStatus(u32), ReturnValue(TokenReturnValue), + // Parsed and traced but not yet surfaced through the public result API. + // TODO(surface): expose ordering metadata to callers. + #[allow(dead_code)] Order(TokenOrder), + // TODO(surface): expose browse-mode column info to callers. + #[allow(dead_code)] + ColInfo(TokenColInfo), + // TODO(surface): expose browse-mode table names to callers. + #[allow(dead_code)] + TabName(TokenTabName), EnvChange(TokenEnvChange), + // Informational messages are logged at DEBUG on decode; the payload is not + // otherwise consumed. TODO(surface): expose INFO messages to callers. + #[allow(dead_code)] Info(TokenInfo), + // Consumed for its `tds_version` during login (see `get_login_ack`); the + // remaining fields are retained for debugging but not otherwise read. + #[allow(dead_code)] LoginAck(TokenLoginAck), + // Consumed by `flush_sspi`, which is only compiled with an integrated-auth + // backend; without one the payload is decoded but never read. + #[cfg_attr( + not(any( + windows, + feature = "winauth", + feature = "integrated-auth-gssapi", + feature = "sspi-rs" + )), + allow(dead_code) + )] Sspi(TokenSspi), + // TODO(surface): act on negotiated feature acknowledgements. + #[allow(dead_code)] FeatureExtAck(TokenFeatureExtAck), + // Connection-resiliency state; parsed but recovery is not yet implemented. + #[allow(dead_code)] + SessionState(TokenSessionState), + // Federated-auth info; parsed for the AAD flow but not read from this enum. + #[allow(dead_code)] + FedAuthInfo(TokenFedAuthInfo), Error(TokenError), } @@ -75,7 +218,39 @@ where } } - #[cfg(any(windows, feature = "integrated-auth-gssapi", feature = "winauth"))] + /// Drain the token stream after a client Attention signal has been sent, + /// discarding every remaining token of the cancelled request until the + /// acknowledging DONE token (with the `DONE_ATTN` status bit set) is + /// received, per MS-TDS section 2.2.1.6. Returns that DONE token so the + /// connection is left clean and ready for reuse. + pub(crate) async fn flush_done_attention(self) -> crate::Result { + let mut stream = self.try_unfold(); + + loop { + match stream.try_next().await? { + Some(ReceivedToken::Done(token)) + | Some(ReceivedToken::DoneProc(token)) + | Some(ReceivedToken::DoneInProc(token)) + if token.is_attention() => + { + return Ok(token); + } + Some(_) => (), + None => { + return Err(crate::Error::Protocol( + "Never got a DONE token acknowledging the Attention signal.".into(), + )) + } + } + } + } + + #[cfg(any( + windows, + feature = "winauth", + feature = "integrated-auth-gssapi", + feature = "sspi-rs" + ))] pub(crate) async fn flush_sspi(self) -> crate::Result { let mut stream = self.try_unfold(); let mut last_error = None; @@ -106,6 +281,30 @@ where Ok(ReceivedToken::NewResultset(meta)) } + async fn get_alt_col_metadata(&mut self) -> crate::Result { + let meta = Arc::new(TokenAltMetaData::decode(self.conn).await?); + self.conn.context_mut().set_alt_meta(meta.clone()); + + event!(Level::TRACE, ?meta); + + Ok(ReceivedToken::NewAltResultset(meta)) + } + + async fn get_alt_row(&mut self) -> crate::Result { + // The id is read first so the matching ALTMETADATA can be looked up + // before the column values are parsed. + let id = self.conn.read_u16_le().await?; + + let meta = self.conn.context().alt_meta(id).ok_or_else(|| { + Error::Protocol(format!("ALTROW for unknown compute id {}", id).into()) + })?; + + let row = TokenAltRow::decode(self.conn, id, &meta).await?; + + event!(Level::TRACE, message = ?row); + Ok(ReceivedToken::AltRow(row)) + } + async fn get_row(&mut self) -> crate::Result { let return_value = TokenRow::decode(self.conn).await?; @@ -148,6 +347,18 @@ where Ok(ReceivedToken::Order(order)) } + async fn get_col_info(&mut self) -> crate::Result { + let col_info = TokenColInfo::decode(self.conn).await?; + event!(Level::TRACE, message = ?col_info); + Ok(ReceivedToken::ColInfo(col_info)) + } + + async fn get_tab_name(&mut self) -> crate::Result { + let tab_name = TokenTabName::decode(self.conn).await?; + event!(Level::TRACE, message = ?tab_name); + Ok(ReceivedToken::TabName(tab_name)) + } + async fn get_done_value(&mut self) -> crate::Result { let done = TokenDone::decode(self.conn).await?; event!(Level::TRACE, "{}", done); @@ -170,7 +381,17 @@ where let change = TokenEnvChange::decode(self.conn).await?; match change { - TokenEnvChange::PacketSize(new_size, _) => { + TokenEnvChange::PacketSize { new: new_size, .. } => { + // MS-TDS: a negotiated packet size must be within 512..=32767. + // Reject an out-of-range value from the server: the subsequent + // `packet_size - HEADER_BYTES` in the send paths would otherwise + // underflow (debug panic) or, for a value of exactly 8, produce + // a zero chunk size and spin `split_to(0)` forever. + if !(512..=32767).contains(&new_size) { + return Err(crate::Error::Protocol( + format!("server requested an invalid packet size of {new_size}").into(), + )); + } self.conn.context_mut().set_packet_size(new_size); } TokenEnvChange::BeginTransaction(desc) => { @@ -184,33 +405,60 @@ where _ => (), } - event!(Level::INFO, "{}", change); + event!(Level::DEBUG, "{}", change); Ok(ReceivedToken::EnvChange(change)) } async fn get_info(&mut self) -> crate::Result { let info = TokenInfo::decode(self.conn).await?; - event!(Level::INFO, "{}", info.message); + event!(Level::DEBUG, "{}", info.message); Ok(ReceivedToken::Info(info)) } async fn get_login_ack(&mut self) -> crate::Result { let ack = TokenLoginAck::decode(self.conn).await?; - event!(Level::INFO, "{} version {}", ack.prog_name, ack.version); + event!(Level::DEBUG, "{} version {}", ack.prog_name, ack.version); + + // Record the TDS version the server negotiated so version-dependent + // decoders (DONE rowcount width, ERROR/INFO LineNumber width) use the + // correct layout instead of the hardcoded default. For a modern server + // (TDS 7.2+) this resolves to the same widths as the default, so there + // is no behavior change; it only corrects decoding against pre-2005 + // servers that negotiate an earlier TDS version. + self.conn.context_mut().set_version(ack.tds_version); + Ok(ReceivedToken::LoginAck(ack)) } async fn get_feature_ext_ack(&mut self) -> crate::Result { let ack = TokenFeatureExtAck::decode(self.conn).await?; event!( - Level::INFO, + Level::DEBUG, "FeatureExtAck with {} features", ack.features.len() ); Ok(ReceivedToken::FeatureExtAck(ack)) } + async fn get_session_state(&mut self) -> crate::Result { + let state = TokenSessionState::decode(self.conn).await?; + event!( + Level::TRACE, + "SessionState seq_no={} recoverable={} states={}", + state.seq_no, + state.is_recoverable(), + state.states.len() + ); + Ok(ReceivedToken::SessionState(state)) + } + + async fn get_fed_auth_info(&mut self) -> crate::Result { + let info = TokenFedAuthInfo::decode(self.conn).await?; + event!(Level::TRACE, message = ?info); + Ok(ReceivedToken::FedAuthInfo(info)) + } + async fn get_sspi(&mut self) -> crate::Result { let sspi = TokenSspi::decode_async(self.conn).await?; event!(Level::TRACE, "SSPI response"); @@ -218,6 +466,17 @@ where } pub fn try_unfold(self) -> BoxStream<'a, crate::Result> { + // Read the per-response command timeout before `self` is moved into the + // unfold closure. Every command result-read path funnels through here + // (query/execute/simple_query result streams, the bulk-insert server + // acknowledgement and column-metadata), so bounding each round-trip once + // at this choke point covers them all. + let command_timeout = self.conn.context().command_timeout(); + // Shared handle so the timeout, when it fires, can mark the underlying + // connection desynced (it is moved into the unfold closure below and is + // otherwise unreachable from `RoundTripTimeout`). + let poison = self.conn.command_desync_flag(); + let stream = futures_util::stream::try_unfold(self, |mut this| async move { if this.conn.is_eof() { match this.last_error { @@ -234,7 +493,9 @@ where let token = match ty { TokenType::ReturnStatus => this.get_return_status().await?, TokenType::ColMetaData => this.get_col_metadata().await?, + TokenType::AltMetaData => this.get_alt_col_metadata().await?, TokenType::Row => this.get_row().await?, + TokenType::AltRow => this.get_alt_row().await?, TokenType::NbcRow => this.get_nbc_row().await?, TokenType::Done => this.get_done_value().await?, TokenType::DoneProc => this.get_done_proc_value().await?, @@ -242,17 +503,674 @@ where TokenType::ReturnValue => this.get_return_value().await?, TokenType::Error => this.get_error().await?, TokenType::Order => this.get_order().await?, + TokenType::ColInfo => this.get_col_info().await?, + TokenType::TabName => this.get_tab_name().await?, TokenType::EnvChange => this.get_env_change().await?, TokenType::Info => this.get_info().await?, TokenType::LoginAck => this.get_login_ack().await?, TokenType::Sspi => this.get_sspi().await?, + TokenType::SessionState => this.get_session_state().await?, + TokenType::FedAuthInfo => this.get_fed_auth_info().await?, TokenType::FeatureExtAck => this.get_feature_ext_ack().await?, - _ => panic!("Token {:?} unimplemented!", ty), + // NOTE: every `TokenType` variant is handled above. This match + // is intentionally exhaustive (no wildcard arm) so that adding a + // new `TokenType` fails to compile until a handler is wired up, + // rather than silently falling through. Unknown token *bytes* + // are already rejected by the `TokenType::try_from` above. }; Ok(Some((token, this))) }); - Box::pin(stream) + Box::pin(RoundTripTimeout::new( + command_timeout, + Some(poison), + Box::pin(stream), + )) + } +} + +#[cfg(test)] +mod tests { + use super::TokenStream; + use crate::client::Connection; + use crate::tds::codec::{Encode, Packet, PacketHeader, PacketStatus}; + use crate::{Error, SqlReadBytes}; + use bytes::BytesMut; + use futures_util::io::{AsyncRead, AsyncWrite}; + use std::io; + use std::pin::Pin; + use std::task::{Context as TaskContext, Poll}; + + // A mock stream that hands out a fixed byte script to `poll_read` (the raw + // TDS packet bytes a server would have written) and swallows every write. + // Once the script is exhausted it reports clean EOF (`Ok(0)`), exactly like + // a closed connection with no more packets on the wire. + struct MockStream { + data: Vec, + pos: usize, + } + + impl MockStream { + fn new(data: Vec) -> Self { + Self { data, pos: 0 } + } + } + + impl AsyncRead for MockStream { + fn poll_read( + mut self: Pin<&mut Self>, + _: &mut TaskContext<'_>, + buf: &mut [u8], + ) -> Poll> { + let remaining = &self.data[self.pos..]; + let n = remaining.len().min(buf.len()); + buf[..n].copy_from_slice(&remaining[..n]); + self.pos += n; + Poll::Ready(Ok(n)) + } + } + + impl AsyncWrite for MockStream { + fn poll_write( + self: Pin<&mut Self>, + _: &mut TaskContext<'_>, + buf: &[u8], + ) -> Poll> { + Poll::Ready(Ok(buf.len())) + } + + fn poll_flush(self: Pin<&mut Self>, _: &mut TaskContext<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn poll_close(self: Pin<&mut Self>, _: &mut TaskContext<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + } + + // --- Wire-byte builders ------------------------------------------------- + // Each returns a token exactly as it appears on the wire (leading token-type + // byte + body), built from the same layouts the decoders parse. See + // `token_done.rs`, `token_error.rs`, and `token_env_change.rs`. + + fn write_b_varchar(buf: &mut Vec, s: &str) { + buf.push(s.encode_utf16().count() as u8); + for unit in s.encode_utf16() { + buf.extend_from_slice(&unit.to_le_bytes()); + } + } + + fn write_us_varchar(buf: &mut Vec, s: &str) { + buf.extend_from_slice(&(s.encode_utf16().count() as u16).to_le_bytes()); + for unit in s.encode_utf16() { + buf.extend_from_slice(&unit.to_le_bytes()); + } + } + + // DONE token (0xFD): status u16, cur_cmd u16, done_rows u64 (8-byte width + // for the default TDS 7.2+ context). + fn done_token(status: u16, cur_cmd: u16, rows: u64) -> Vec { + let mut v = vec![0xFDu8]; + v.extend_from_slice(&status.to_le_bytes()); + v.extend_from_slice(&cur_cmd.to_le_bytes()); + v.extend_from_slice(&rows.to_le_bytes()); + v + } + + // ERROR token (0xAA): u16 length prefix, then the ERROR/INFO body. The + // default test context is TDS 7.2+, so LineNumber is a 4-byte LONG. + fn error_token(code: u32, message: &str, server: &str, procedure: &str, line: u32) -> Vec { + let mut body = Vec::new(); + body.extend_from_slice(&code.to_le_bytes()); + body.push(1); // state + body.push(16); // class + write_us_varchar(&mut body, message); + write_b_varchar(&mut body, server); + write_b_varchar(&mut body, procedure); + body.extend_from_slice(&line.to_le_bytes()); + + let mut v = vec![0xAAu8]; + v.extend_from_slice(&(body.len() as u16).to_le_bytes()); + v.extend_from_slice(&body); + v + } + + // ENVCHANGE token (0xE3) of type Routing (20): u16 length prefix over the + // type byte + payload. + fn envchange_routing(port: u16, host: &str) -> Vec { + let mut payload = Vec::new(); + payload.extend_from_slice(&0u16.to_le_bytes()); // routing data value length (unused) + payload.push(0); // protocol, always 0 (tcp) + payload.extend_from_slice(&port.to_le_bytes()); + payload.extend_from_slice(&(host.encode_utf16().count() as u16).to_le_bytes()); + for unit in host.encode_utf16() { + payload.extend_from_slice(&unit.to_le_bytes()); + } + + let mut body = vec![20u8]; + body.extend_from_slice(&payload); + + let mut v = vec![0xE3u8]; + v.extend_from_slice(&(body.len() as u16).to_le_bytes()); + v.extend_from_slice(&body); + v + } + + // ENVCHANGE token (0xE3) of type PacketSize (4): new/old sizes as B_VARCHAR + // decimal strings that the decoder parses into u32. + fn envchange_packet_size(new: &str, old: &str) -> Vec { + let mut payload = Vec::new(); + write_b_varchar(&mut payload, new); + write_b_varchar(&mut payload, old); + + let mut body = vec![4u8]; + body.extend_from_slice(&payload); + + let mut v = vec![0xE3u8]; + v.extend_from_slice(&(body.len() as u16).to_le_bytes()); + v.extend_from_slice(&body); + v + } + + // Wrap a token-stream payload in a single end-of-message TDS packet, framed + // exactly like `Packet::encode` (8-byte header with the big-endian total + // length patched into bytes [2..4]). This is what the mock stream serves. + fn packet_bytes(payload: &[u8]) -> Vec { + let mut header = PacketHeader::batch(1); + header.set_status(PacketStatus::EndOfMessage); + let packet = Packet::new(header, BytesMut::from(payload)); + + let mut buf = BytesMut::new(); + packet.encode(&mut buf).unwrap(); + buf.to_vec() + } + + fn conn_over(payload: &[u8]) -> Connection { + Connection::test_over(MockStream::new(packet_bytes(payload)), false) + } + + // --- (a) flush_done precedence ----------------------------------------- + + #[tokio::test] + async fn flush_done_error_beats_routing_beats_done() { + // [ERROR][ENVCHANGE Routing][DONE]: the ERROR must win over the routing + // redirect, which in turn would win over a plain DONE. Locks the + // precedence in `flush_done`'s match on (last_error, routing). + let mut payload = Vec::new(); + payload.extend_from_slice(&error_token(1205, "boom", "srv", "proc", 7)); + payload.extend_from_slice(&envchange_routing(1433, "other.example.com")); + payload.extend_from_slice(&done_token(0, 0, 0)); + + let mut conn = conn_over(&payload); + let err = TokenStream::new(&mut conn) + .flush_done() + .await + .expect_err("an ERROR token must surface as an error, not Ok(DONE)"); + + match err { + Error::Server(e) => assert_eq!(e.code(), 1205), + other => panic!("expected Error::Server, got {other:?}"), + } + } + + #[tokio::test] + async fn flush_done_routing_beats_done() { + // [ENVCHANGE Routing][DONE] with no ERROR: the routing redirect wins + // over a plain DONE. + let mut payload = Vec::new(); + payload.extend_from_slice(&envchange_routing(1433, "other.example.com")); + payload.extend_from_slice(&done_token(0, 0, 0)); + + let mut conn = conn_over(&payload); + let err = TokenStream::new(&mut conn) + .flush_done() + .await + .expect_err("a routing env-change must surface as a Routing error"); + + match err { + Error::Routing { host, port } => { + assert_eq!(host, "other.example.com"); + assert_eq!(port, 1433); + } + other => panic!("expected Error::Routing, got {other:?}"), + } + } + + #[tokio::test] + async fn flush_done_plain_done_is_ok() { + // A lone [DONE] returns Ok(TokenDone). + let mut conn = conn_over(&done_token(0, 0, 0)); + let done = TokenStream::new(&mut conn) + .flush_done() + .await + .expect("a plain DONE must succeed"); + assert!(done.is_final()); + } + + // --- (b) get_env_change packet-size bound ------------------------------ + + #[tokio::test] + async fn env_change_packet_size_below_512_is_protocol_error() { + // A negotiated packet size below the 512..=32767 range must be rejected + // (would otherwise underflow / spin the send path), not panic. + let mut conn = conn_over(&envchange_packet_size("8", "4096")); + let err = TokenStream::new(&mut conn) + .flush_done() + .await + .expect_err("packet size 8 is below the 512 floor"); + match err { + Error::Protocol(msg) => assert!(msg.contains("invalid packet size")), + other => panic!("expected Error::Protocol, got {other:?}"), + } + } + + #[tokio::test] + async fn env_change_packet_size_above_32767_is_protocol_error() { + let mut conn = conn_over(&envchange_packet_size("40000", "4096")); + let err = TokenStream::new(&mut conn) + .flush_done() + .await + .expect_err("packet size 40000 is above the 32767 ceiling"); + match err { + Error::Protocol(msg) => assert!(msg.contains("invalid packet size")), + other => panic!("expected Error::Protocol, got {other:?}"), + } + } + + #[tokio::test] + async fn env_change_packet_size_in_range_is_applied() { + // A valid packet size (8192) is accepted and stored in the context. + let mut payload = Vec::new(); + payload.extend_from_slice(&envchange_packet_size("8192", "4096")); + payload.extend_from_slice(&done_token(0, 0, 0)); + + let mut conn = conn_over(&payload); + TokenStream::new(&mut conn) + .flush_done() + .await + .expect("an in-range packet size must be accepted"); + assert_eq!(conn.context().packet_size(), 8192); + } + + // --- (c) try_unfold dispatch ------------------------------------------- + + #[tokio::test] + async fn unknown_token_type_byte_is_protocol_error() { + // 0x00 is not a defined token type; `TokenType::try_from` must reject it. + let mut conn = conn_over(&[0x00u8]); + let err = TokenStream::new(&mut conn) + .flush_done() + .await + .expect_err("an unknown token byte must be rejected"); + match err { + Error::Protocol(msg) => assert!(msg.contains("invalid token type")), + other => panic!("expected Error::Protocol, got {other:?}"), + } + } + + #[tokio::test] + async fn row_before_col_metadata_is_protocol_error() { + // A ROW token (0xD1) arriving before any COLMETADATA has no cached + // metadata to parse against and must be a clean protocol error. + let mut conn = conn_over(&[0xD1u8]); + let err = TokenStream::new(&mut conn) + .flush_done() + .await + .expect_err("a ROW before COLMETADATA must error"); + match err { + Error::Protocol(msg) => assert!(msg.contains("before any COLMETADATA")), + other => panic!("expected Error::Protocol, got {other:?}"), + } + } + + // --- (d) command timeout end-to-end ------------------------------------- + // + // A server that answers with a first packet and then goes silent mid-result + // must surface a prompt TimedOut error through the whole + // Connection -> TokenStream -> RoundTripTimeout read path, rather than + // hanging (command-timeout arm). + + // Frame a token payload into a single *non-final* (NormalMessage) TDS + // packet, so the connection expects at least one more packet afterwards. + fn normal_packet_bytes(payload: &[u8]) -> Vec { + let mut header = PacketHeader::batch(1); + header.set_status(PacketStatus::NormalMessage); + let packet = Packet::new(header, BytesMut::from(payload)); + + let mut buf = BytesMut::new(); + packet.encode(&mut buf).unwrap(); + buf.to_vec() + } + + // Serves a fixed script of bytes once, then parks every subsequent read + // forever — a server that answers and then stalls mid-response. Distinct + // from `MockStream`, which reports clean EOF once its script is exhausted. + struct AnswerThenSilent { + data: Vec, + pos: usize, + } + + impl AsyncRead for AnswerThenSilent { + fn poll_read( + mut self: Pin<&mut Self>, + _: &mut TaskContext<'_>, + buf: &mut [u8], + ) -> Poll> { + if self.pos >= self.data.len() { + // Script exhausted: the server has gone silent. Only the command + // timeout can unblock the reader now. + return Poll::Pending; + } + let remaining = &self.data[self.pos..]; + let n = remaining.len().min(buf.len()); + buf[..n].copy_from_slice(&remaining[..n]); + self.pos += n; + Poll::Ready(Ok(n)) + } + } + + impl AsyncWrite for AnswerThenSilent { + fn poll_write( + self: Pin<&mut Self>, + _: &mut TaskContext<'_>, + buf: &[u8], + ) -> Poll> { + Poll::Ready(Ok(buf.len())) + } + + fn poll_flush(self: Pin<&mut Self>, _: &mut TaskContext<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn poll_close(self: Pin<&mut Self>, _: &mut TaskContext<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + } + + fn is_timed_out(err: &Error) -> bool { + matches!(err, Error::Io { kind, .. } if *kind == io::ErrorKind::TimedOut) + } + + #[tokio::test] + async fn command_timeout_fires_when_server_stalls_mid_stream() { + use crate::SqlReadBytes; + use std::time::{Duration, Instant}; + + // First (non-final) packet carries a valid ENVCHANGE but no DONE, then + // the stream goes silent — so the reader must wait for a packet that + // never comes. + let script = normal_packet_bytes(&envchange_packet_size("8192", "4096")); + let mut conn = Connection::test_over( + AnswerThenSilent { + data: script, + pos: 0, + }, + false, + ); + conn.context_mut() + .set_command_timeout(Some(Duration::from_millis(100))); + + let started = Instant::now(); + let err = TokenStream::new(&mut conn) + .flush_done() + .await + .expect_err("a server that stalls mid-result must not hang the read"); + + assert!( + is_timed_out(&err), + "expected a TimedOut error, got: {err:?}" + ); + assert!( + err.to_string().contains("did not send the next part") + && err.to_string().contains("out of sync"), + "command timeout error should be diagnosable, got: {err}" + ); + assert!( + started.elapsed() < Duration::from_secs(5), + "the command timeout must fire promptly, took {:?}", + started.elapsed() + ); + // The token that *did* arrive before the stall was still processed. + assert_eq!(conn.context().packet_size(), 8192); + // The connection is now desynced (the tail of the timed-out response is + // still unread): it must be poisoned so a pool cannot silently reuse it + // and re-hang in `flush_stream` on the still-pending response. + let reuse = conn.ensure_not_poisoned(); + assert!( + reuse.is_err(), + "a command timeout must leave the connection unusable, got: {reuse:?}" + ); + assert!( + reuse.unwrap_err().to_string().contains("out of sync"), + "reuse should be rejected with a command-timeout desync error" + ); + } + + #[tokio::test] + async fn no_command_timeout_keeps_waiting_on_a_stalled_server() { + // With no command timeout configured, a mid-result stall keeps the read + // pending; an outer guard proves no internal timeout fires. + let script = normal_packet_bytes(&envchange_packet_size("8192", "4096")); + let mut conn = Connection::test_over( + AnswerThenSilent { + data: script, + pos: 0, + }, + false, + ); + // command_timeout defaults to None on a bare Context. + + let outcome = tokio::time::timeout(std::time::Duration::from_millis(150), async { + TokenStream::new(&mut conn).flush_done().await + }) + .await; + + assert!( + outcome.is_err(), + "with no command timeout the read should stay pending on the stalled server" + ); + } +} + +// Server-free tests for the `RoundTripTimeout` stream combinator that backs +// `Config::command_timeout`. They assert the per-round-trip semantics directly: +// a stalled server times out promptly, responsive reads pass through, `None` +// disables the bound, and — critically — a slow *consumer* never trips it. +#[cfg(test)] +mod round_trip_timeout_tests { + use super::{ReceivedToken, RoundTripTimeout}; + use futures_util::stream::{self, BoxStream, StreamExt}; + use std::time::{Duration, Instant}; + + fn token() -> ReceivedToken { + ReceivedToken::ReturnStatus(0) + } + + fn is_timed_out(err: &crate::Error) -> bool { + matches!(err, crate::Error::Io { kind, .. } if *kind == std::io::ErrorKind::TimedOut) + } + + #[tokio::test] + async fn stalled_server_times_out_promptly() { + // Inner stream never yields: a server that stops responding mid-result. + let inner: BoxStream<'static, crate::Result> = stream::pending().boxed(); + let mut s = RoundTripTimeout::new(Some(Duration::from_millis(50)), None, inner); + + let started = Instant::now(); + let item = s.next().await.expect("must yield a timeout error, not end"); + let err = item.expect_err("a stalled server must produce an error"); + assert!(is_timed_out(&err), "expected TimedOut, got: {err:?}"); + assert!( + started.elapsed() < Duration::from_secs(5), + "must fire promptly, took {:?}", + started.elapsed() + ); + } + + #[tokio::test] + async fn responsive_reads_pass_through_under_the_bound() { + let inner: BoxStream<'static, crate::Result> = + stream::iter(vec![Ok(token()), Ok(token()), Ok(token())]).boxed(); + let mut s = RoundTripTimeout::new(Some(Duration::from_secs(30)), None, inner); + + let mut count = 0; + while let Some(item) = s.next().await { + item.expect("responsive reads must not error"); + count += 1; + } + assert_eq!(count, 3); + } + + #[tokio::test] + async fn none_timeout_never_injects_a_deadline() { + // A permanently pending inner stream with no bound must stay pending — + // the outer guard elapsing proves no internal timeout fired. + let inner: BoxStream<'static, crate::Result> = stream::pending().boxed(); + let mut s = RoundTripTimeout::new(None, None, inner); + + let outcome = tokio::time::timeout(Duration::from_millis(100), s.next()).await; + assert!( + outcome.is_err(), + "with no bound the stream must keep waiting, not time out" + ); + } + + #[tokio::test] + async fn slow_consumer_does_not_trip_the_timeout() { + // The server answers each round-trip with a genuine `Pending` gap that + // stays *under* the bound (so the timer really is armed and reset on + // each delivery), but the consumer then idles far *longer* than the + // bound between pulls. Because the timer only runs while a read is in + // flight — never between the consumer's polls — this must NOT time out, + // proving the bound measures server latency, not consumer pace. Using a + // real slow server (not an always-ready `stream::iter`) means the + // `Pending`/timer arm is actually exercised, so the test would fail if + // the timer wrongly counted consumer idle time. + let mut s = RoundTripTimeout::new( + Some(Duration::from_millis(100)), + None, + slow_server(Duration::from_millis(30), 3), + ); + + for _ in 0..3 { + s.next() + .await + .expect("item expected") + .expect("a responsive server must not time out"); + // Idle far longer than the 100ms bound between consuming items. + tokio::time::sleep(Duration::from_millis(200)).await; + } + assert!(s.next().await.is_none(), "stream should end cleanly"); + } + + // An inner stream whose every item is delivered only after a genuine + // `Pending` gap of `gap`, driven by a real timer. Unlike `stream::iter` + // (always immediately `Ready`), this forces `RoundTripTimeout` down its + // `Pending` arm on every round-trip, so the lazy-create / reset-on-delivery + // logic is actually exercised. + fn slow_server( + gap: Duration, + count: usize, + ) -> BoxStream<'static, crate::Result> { + stream::unfold(0usize, move |i| async move { + if i >= count { + return None; + } + tokio::time::sleep(gap).await; + Some((Ok(token()), i + 1)) + }) + .boxed() + } + + #[tokio::test] + async fn bounds_each_round_trip_not_total_enumeration() { + // Five genuine round-trips, each 50ms (well under the 200ms bound), for a + // total enumeration of ~250ms — longer than the bound. A correct + // per-round-trip timer resets on each delivery and never trips; a broken + // one that measured total stream lifetime (delay created once, never + // reset) would fire partway through. All five tokens must arrive. + let mut s = RoundTripTimeout::new( + Some(Duration::from_millis(200)), + None, + slow_server(Duration::from_millis(50), 5), + ); + + let mut count = 0; + while let Some(item) = s.next().await { + item.expect("no round-trip exceeds the bound, so none must time out"); + count += 1; + } + assert_eq!(count, 5, "all tokens must be delivered across genuine gaps"); + } + + #[tokio::test] + async fn deadline_resets_after_delivery_then_a_later_stall_trips() { + // One token arrives after a short gap (under the bound), then the server + // goes silent forever. The first item must be delivered (proving the + // pre-stall token survives), and the *next* poll must arm a fresh timer + // and trip (proving the deadline re-arms after a delivery rather than + // being consumed once). + let inner = stream::unfold(0usize, |i| async move { + match i { + 0 => { + tokio::time::sleep(Duration::from_millis(20)).await; + Some((Ok(token()), 1usize)) + } + _ => std::future::pending().await, + } + }) + .boxed(); + let mut s = RoundTripTimeout::new(Some(Duration::from_millis(80)), None, inner); + + s.next() + .await + .expect("first token expected") + .expect("the first round-trip is under the bound and must not time out"); + + let err = s + .next() + .await + .expect("a second item (the timeout error) is expected") + .expect_err("the second, stalled round-trip must trip the re-armed deadline"); + assert!(is_timed_out(&err), "expected TimedOut, got: {err:?}"); + } + + #[tokio::test] + async fn poison_flag_is_flipped_when_the_deadline_fires() { + // The shared poison handle must be set when the timer fires, so the + // owning connection can reject reuse. + let poison = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)); + let inner: BoxStream<'static, crate::Result> = stream::pending().boxed(); + let mut s = + RoundTripTimeout::new(Some(Duration::from_millis(50)), Some(poison.clone()), inner); + + assert!(!poison.load(std::sync::atomic::Ordering::Acquire)); + let err = s + .next() + .await + .expect("a timeout error is expected") + .expect_err("a stalled server must produce an error"); + assert!(is_timed_out(&err)); + assert!( + poison.load(std::sync::atomic::Ordering::Acquire), + "the poison flag must be set when the command timeout fires" + ); + } + + #[tokio::test] + async fn zero_duration_trips_on_the_first_stall() { + // A zero-duration bound is a degenerate but valid setting: the first time + // the server is not immediately ready, the deadline is already elapsed, + // so the read fails fast rather than hanging. + let inner: BoxStream<'static, crate::Result> = stream::pending().boxed(); + let mut s = RoundTripTimeout::new(Some(Duration::ZERO), None, inner); + + let err = s + .next() + .await + .expect("a timeout error is expected") + .expect_err("a zero bound must trip on the first stall"); + assert!(is_timed_out(&err), "expected TimedOut, got: {err:?}"); } } diff --git a/src/tds/time.rs b/src/tds/time.rs index 05a1c053c..30e9efd7d 100644 --- a/src/tds/time.rs +++ b/src/tds/time.rs @@ -22,11 +22,13 @@ //! [`OffsetDateTime`]: time/struct.OffsetDateTime.html #[cfg(feature = "chrono")] -#[cfg_attr(feature = "docs", doc(cfg(feature = "chrono")))] +#[cfg_attr(docsrs, doc(cfg(feature = "chrono")))] pub mod chrono; #[cfg(feature = "time")] -#[cfg_attr(feature = "docs", doc(cfg(feature = "time")))] +#[cfg_attr(docsrs, doc(cfg(feature = "time")))] +// Submodule intentionally shares the name of the `time` feature/crate it wraps. +#[allow(clippy::module_inception)] pub mod time; use crate::{tds::codec::Encode, SqlReadBytes}; @@ -43,6 +45,7 @@ use futures_util::io::AsyncReadExt; /// It isn't recommended to use this type directly. For dealing with `datetime`, /// use the `time` feature of this crate and its `PrimitiveDateTime` type. #[derive(Copy, Clone, Debug, Eq, PartialEq)] +#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] pub struct DateTime { days: i32, seconds_fragments: u32, @@ -99,6 +102,7 @@ impl Encode for DateTime { /// `smalldatetime`, use the `time` feature of this crate and its /// `PrimitiveDateTime` type. #[derive(Copy, Clone, Debug, Eq, PartialEq)] +#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] pub struct SmallDateTime { days: u16, seconds_fragments: u16, @@ -152,12 +156,13 @@ impl Encode for SmallDateTime { /// It isn't recommended to use this type directly. If you want to deal with /// `date`, use the `time` feature of this crate and its `Date` type. #[derive(Copy, Clone, Debug, Eq, PartialEq)] +#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] #[cfg(feature = "tds73")] -#[cfg_attr(feature = "docs", doc(cfg(feature = "tds73")))] +#[cfg_attr(docsrs, doc(cfg(feature = "tds73")))] pub struct Date(u32); #[cfg(feature = "tds73")] -#[cfg_attr(feature = "docs", doc(cfg(feature = "tds73")))] +#[cfg_attr(docsrs, doc(cfg(feature = "tds73")))] impl Date { #[inline] /// Construct a new `Date` @@ -169,6 +174,18 @@ impl Date { Date(days) } + /// Construct a `Date` from a raw day count *without* the 3-byte range check, + /// deferring validation to [`Encode::encode`]. Used by the chrono + /// conversion path, whose `ToSql`/`IntoSql` impls return a `ColumnData` + /// directly (not a `Result`) and so cannot reject an out-of-range + /// `NaiveDate` at conversion time; the out-of-range value instead surfaces + /// as an `Err` when the value is encoded. + #[inline] + #[cfg(feature = "chrono")] + pub(crate) fn new_unchecked(days: u32) -> Date { + Date(days) + } + #[inline] /// The number of days from 1st of January, year 1. pub fn days(self) -> u32 { @@ -186,12 +203,19 @@ impl Date { } #[cfg(feature = "tds73")] -#[cfg_attr(feature = "docs", doc(cfg(feature = "tds73")))] +#[cfg_attr(docsrs, doc(cfg(feature = "tds73")))] impl Encode for Date { fn encode(self, dst: &mut BytesMut) -> crate::Result<()> { let mut tmp = [0u8; 4]; LittleEndian::write_u32(&mut tmp, self.days()); - assert_eq!(tmp[3], 0); + // A `date` is a 3-byte value; a day count that does not fit (e.g. a + // chrono `NaiveDate` outside the SQL Server range) is rejected here + // rather than panicking. + if tmp[3] != 0 { + return Err(crate::Error::Protocol( + format!("date day count {} is out of the 3-byte range", self.days()).into(), + )); + } dst.extend_from_slice(&tmp[0..3]); Ok(()) @@ -205,15 +229,16 @@ impl Encode for Date { /// It isn't recommended to use this type directly. If you want to deal with /// `time`, use the `time` feature of this crate and its `Time` type. #[derive(Copy, Clone, Debug)] +#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] #[cfg(feature = "tds73")] -#[cfg_attr(feature = "docs", doc(cfg(feature = "tds73")))] +#[cfg_attr(docsrs, doc(cfg(feature = "tds73")))] pub struct Time { increments: u64, scale: u8, } #[cfg(feature = "tds73")] -#[cfg_attr(feature = "docs", doc(cfg(feature = "tds73")))] +#[cfg_attr(docsrs, doc(cfg(feature = "tds73")))] impl PartialEq for Time { fn eq(&self, t: &Time) -> bool { self.increments as f64 / 10f64.powi(self.scale as i32) @@ -222,7 +247,7 @@ impl PartialEq for Time { } #[cfg(feature = "tds73")] -#[cfg_attr(feature = "docs", doc(cfg(feature = "tds73")))] +#[cfg_attr(docsrs, doc(cfg(feature = "tds73")))] impl Time { /// Construct a new `Time` pub fn new(increments: u64, scale: u8) -> Self { @@ -292,21 +317,39 @@ impl Time { } #[cfg(feature = "tds73")] -#[cfg_attr(feature = "docs", doc(cfg(feature = "tds73")))] +#[cfg_attr(docsrs, doc(cfg(feature = "tds73")))] impl Encode for Time { fn encode(self, dst: &mut BytesMut) -> crate::Result<()> { + // The field width is determined by the scale; an `increments` value + // that does not fit that width (e.g. a `Time` built directly via the + // pub `Time::new`) is rejected here rather than panicking. + let width_bits = match self.len()? { + 3 => 24, + 4 => 32, + 5 => 40, + _ => unreachable!(), + }; + if self.increments >> width_bits != 0 { + return Err(crate::Error::Protocol( + format!( + "time increments {} do not fit the {}-byte field for scale {}", + self.increments, + width_bits / 8, + self.scale + ) + .into(), + )); + } + match self.len()? { 3 => { - assert_eq!(self.increments >> 24, 0); dst.put_u16_le(self.increments as u16); dst.put_u8((self.increments >> 16) as u8); } 4 => { - assert_eq!(self.increments >> 32, 0); dst.put_u32_le(self.increments as u32); } 5 => { - assert_eq!(self.increments >> 40, 0); dst.put_u32_le(self.increments as u32); dst.put_u8((self.increments >> 32) as u8); } @@ -318,8 +361,9 @@ impl Encode for Time { } #[derive(Copy, Clone, Debug, PartialEq)] +#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] #[cfg(feature = "tds73")] -#[cfg_attr(feature = "docs", doc(cfg(feature = "tds73")))] +#[cfg_attr(docsrs, doc(cfg(feature = "tds73")))] /// A presentation of `datetime2` type in the server. /// /// # Warning @@ -333,7 +377,7 @@ pub struct DateTime2 { } #[cfg(feature = "tds73")] -#[cfg_attr(feature = "docs", doc(cfg(feature = "tds73")))] +#[cfg_attr(docsrs, doc(cfg(feature = "tds73")))] impl DateTime2 { /// Construct a new `DateTime2` from the date and time components. pub fn new(date: Date, time: Time) -> Self { @@ -365,23 +409,22 @@ impl DateTime2 { } #[cfg(feature = "tds73")] -#[cfg_attr(feature = "docs", doc(cfg(feature = "tds73")))] +#[cfg_attr(docsrs, doc(cfg(feature = "tds73")))] impl Encode for DateTime2 { fn encode(self, dst: &mut BytesMut) -> crate::Result<()> { self.time.encode(dst)?; - - let mut tmp = [0u8; 4]; - LittleEndian::write_u32(&mut tmp, self.date.days()); - assert_eq!(tmp[3], 0); - dst.extend_from_slice(&tmp[0..3]); + // Reuse `Date::encode` so an out-of-range date surfaces as an `Err` + // (same 3-byte wire layout as the previous inline encoding). + self.date.encode(dst)?; Ok(()) } } #[derive(Copy, Clone, Debug, PartialEq)] +#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] #[cfg(feature = "tds73")] -#[cfg_attr(feature = "docs", doc(cfg(feature = "tds73")))] +#[cfg_attr(docsrs, doc(cfg(feature = "tds73")))] /// A presentation of `datetimeoffset` type in the server. /// /// # Warning @@ -395,7 +438,7 @@ pub struct DateTimeOffset { } #[cfg(feature = "tds73")] -#[cfg_attr(feature = "docs", doc(cfg(feature = "tds73")))] +#[cfg_attr(docsrs, doc(cfg(feature = "tds73")))] impl DateTimeOffset { /// Construct a new `DateTimeOffset` from a `datetime2`, offset marking /// number of minutes from UTC. @@ -425,7 +468,7 @@ impl DateTimeOffset { } #[cfg(feature = "tds73")] -#[cfg_attr(feature = "docs", doc(cfg(feature = "tds73")))] +#[cfg_attr(docsrs, doc(cfg(feature = "tds73")))] impl Encode for DateTimeOffset { fn encode(self, dst: &mut BytesMut) -> crate::Result<()> { self.datetime2.encode(dst)?; @@ -434,3 +477,193 @@ impl Encode for DateTimeOffset { Ok(()) } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::sql_read_bytes::test_utils::IntoSqlReadBytes; + + #[test] + fn datetime_accessors() { + let dt = DateTime::new(-100, 12345); + assert_eq!(dt.days(), -100); + assert_eq!(dt.seconds_fragments(), 12345); + } + + #[tokio::test] + async fn datetime_round_trip_including_pre_1900() { + for dt in [ + DateTime::new(0, 0), + DateTime::new(200, 3000), + DateTime::new(-53690, 25920000), + ] { + let mut buf = BytesMut::new(); + dt.encode(&mut buf).unwrap(); + let decoded = DateTime::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + assert_eq!(decoded, dt); + } + } + + #[test] + fn smalldatetime_accessors() { + let dt = SmallDateTime::new(100, 200); + assert_eq!(dt.days(), 100); + assert_eq!(dt.seconds_fragments(), 200); + } + + #[tokio::test] + async fn smalldatetime_round_trip() { + let dt = SmallDateTime::new(65535, 1439); + let mut buf = BytesMut::new(); + dt.encode(&mut buf).unwrap(); + let decoded = SmallDateTime::decode(&mut buf.into_sql_read_bytes()) + .await + .unwrap(); + assert_eq!(decoded, dt); + } + + #[cfg(feature = "tds73")] + #[test] + fn date_accessor_and_new() { + let date = Date::new(730119); + assert_eq!(date.days(), 730119); + } + + #[cfg(feature = "tds73")] + #[test] + #[should_panic(expected = "left == right")] + fn date_new_panics_on_overflow() { + // Anything not representable in three bytes must panic. `Date::new` + // asserts `days >> 24 == 0` via `assert_eq!`, whose panic message + // contains "assertion `left == right` failed". + Date::new(0x0100_0000); + } + + #[cfg(feature = "tds73")] + #[tokio::test] + async fn date_round_trip() { + for days in [0u32, 1, 730119, 0x00ff_ffff] { + let date = Date::new(days); + let mut buf = BytesMut::new(); + date.encode(&mut buf).unwrap(); + assert_eq!(buf.len(), 3); + let decoded = Date::decode(&mut buf.into_sql_read_bytes()).await.unwrap(); + assert_eq!(decoded, date); + } + } + + #[cfg(feature = "tds73")] + #[test] + fn time_accessors_and_len() { + let time = Time::new(1234, 5); + assert_eq!(time.increments(), 1234); + assert_eq!(time.scale(), 5); + assert_eq!(time.len().unwrap(), 5); + + assert_eq!(Time::new(0, 0).len().unwrap(), 3); + assert_eq!(Time::new(0, 3).len().unwrap(), 4); + assert!(Time::new(0, 8).len().is_err()); + } + + #[cfg(feature = "tds73")] + #[test] + fn time_partial_eq_across_scales() { + // 1 second expressed at two different scales must compare equal. + assert_eq!(Time::new(100, 2), Time::new(10_000_000, 7)); + assert_ne!(Time::new(100, 2), Time::new(200, 2)); + } + + #[cfg(feature = "tds73")] + #[tokio::test] + async fn time_round_trip_all_len_buckets() { + for (increments, scale) in [(255u64, 2u8), (65535, 4), (16_777_215, 7)] { + let time = Time::new(increments, scale); + let rlen = time.len().unwrap(); + let mut buf = BytesMut::new(); + time.encode(&mut buf).unwrap(); + let decoded = Time::decode( + &mut buf.into_sql_read_bytes(), + scale as usize, + rlen as usize, + ) + .await + .unwrap(); + assert_eq!(decoded, time); + } + } + + #[cfg(feature = "tds73")] + #[tokio::test] + async fn time_round_trip_high_bytes_set() { + // Values whose most-significant byte (the byte handled by the + // `lo << 16` / `lo << 32` shift in `decode` and the `>> 16` / `>> 32` + // shift in `encode`) is non-zero. This distinguishes: + // * decode `<< N` from `>> N` (the latter zeroes an `u8`), and + // * encode `>> N` from `<< N` (the latter zeroes the byte written). + // The 16-bit / 32-bit low halves and the shifted high byte occupy + // disjoint bit ranges, so `|` vs `^` cannot be distinguished here. + for (increments, scale) in [(0x00FF_1234u64, 2u8), (0x00AB_1234_5678u64, 7)] { + let time = Time::new(increments, scale); + let rlen = time.len().unwrap(); + + let mut buf = BytesMut::new(); + time.encode(&mut buf).unwrap(); + + let decoded = Time::decode( + &mut buf.into_sql_read_bytes(), + scale as usize, + rlen as usize, + ) + .await + .unwrap(); + + assert_eq!(decoded, time); + assert_eq!(decoded.increments(), increments); + } + } + + #[cfg(feature = "tds73")] + #[tokio::test] + async fn time_decode_invalid_length_errors() { + let mut buf = BytesMut::new(); + buf.put_u8(0); + // scale/length combination not one of the accepted pairs. + let err = Time::decode(&mut buf.into_sql_read_bytes(), 0, 4).await; + assert!(err.is_err()); + } + + #[cfg(feature = "tds73")] + #[tokio::test] + async fn datetime2_round_trip_and_accessors() { + let dt2 = DateTime2::new(Date::new(730119), Time::new(222, 7)); + assert_eq!(dt2.date(), Date::new(730119)); + assert_eq!(dt2.time(), Time::new(222, 7)); + + let rlen = dt2.time().len().unwrap(); + let mut buf = BytesMut::new(); + dt2.encode(&mut buf).unwrap(); + let decoded = DateTime2::decode(&mut buf.into_sql_read_bytes(), 7, rlen as usize) + .await + .unwrap(); + assert_eq!(decoded, dt2); + } + + #[cfg(feature = "tds73")] + #[tokio::test] + async fn datetimeoffset_round_trip_and_accessors() { + let dt2 = DateTime2::new(Date::new(730119), Time::new(222, 7)); + let dto = DateTimeOffset::new(dt2, -120); + assert_eq!(dto.datetime2(), dt2); + assert_eq!(dto.offset(), -120); + + let rlen = dto.datetime2().time().len().unwrap(); + let mut buf = BytesMut::new(); + dto.encode(&mut buf).unwrap(); + let decoded = DateTimeOffset::decode(&mut buf.into_sql_read_bytes(), 7, rlen) + .await + .unwrap(); + assert_eq!(decoded, dto); + } +} diff --git a/src/tds/time/chrono.rs b/src/tds/time/chrono.rs index f6de50012..da3ba277a 100644 --- a/src/tds/time/chrono.rs +++ b/src/tds/time/chrono.rs @@ -11,15 +11,54 @@ use super::DateTime as DateTime1; use super::{Date, DateTime2, DateTimeOffset, Time}; use crate::tds::codec::ColumnData; #[cfg(feature = "tds73")] -#[cfg_attr(feature = "docs", doc(cfg(feature = "tds73")))] +#[cfg_attr(docsrs, doc(cfg(feature = "tds73")))] pub use chrono::offset::{FixedOffset, Utc}; pub use chrono::{DateTime, NaiveDate, NaiveDateTime, NaiveTime}; + +#[inline] +fn from_days(days: i64, start_year: i32) -> crate::Result { + // `days` derives from untrusted server bytes. Every valid SQL date fits + // within `NaiveDate`; a genuinely out-of-range/malformed day offset is + // rejected as a protocol error rather than silently clamped to MIN/MAX + // (which would decode a malformed value to a plausible-but-wrong date). + let base = NaiveDate::from_ymd_opt(start_year, 1, 1).unwrap(); + base.checked_add_signed(chrono::Duration::days(days)) + .ok_or_else(|| { + crate::Error::Protocol( + format!("date day offset {days} is out of the representable range").into(), + ) + }) +} + +/// Validate a server-supplied UTC offset (in whole minutes). SQL Server's +/// `datetimeoffset` is only valid for -14:00..=+14:00; a malformed offset +/// outside that range is rejected as a protocol error rather than silently +/// falling back to UTC (which would shift the represented instant). +#[inline] #[cfg(feature = "tds73")] -use std::ops::Sub; +fn validate_offset_minutes(minutes: i16) -> crate::Result { + if !(-840..=840).contains(&minutes) { + return Err(crate::Error::Protocol( + format!( + "datetimeoffset offset {minutes} minutes is outside the valid -14:00..=+14:00 range" + ) + .into(), + )); + } + + Ok(minutes as i32) +} +/// Convert a server-supplied fractional-seconds `increments` at the given +/// `scale` into nanoseconds without panicking (`scale > 9` would underflow +/// `9 - scale`; a large `increments` would overflow the multiply). #[inline] -fn from_days(days: i64, start_year: i32) -> NaiveDate { - NaiveDate::from_ymd_opt(start_year, 1, 1).unwrap() + chrono::Duration::days(days) +#[cfg(feature = "tds73")] +fn nanos_from_increments(increments: u64, scale: u8) -> i64 { + let pow = 9u32.saturating_sub(scale as u32); + increments + .saturating_mul(10u64.saturating_pow(pow)) + .min(i64::MAX as u64) as i64 } #[inline] @@ -30,8 +69,18 @@ fn from_sec_fragments(sec_fragments: i64) -> NaiveTime { #[inline] #[cfg(feature = "tds73")] -fn from_mins(mins: u32) -> NaiveTime { - NaiveTime::from_num_seconds_from_midnight_opt(mins, 0).unwrap() +fn from_mins(mins: u32) -> crate::Result { + // `mins` is the seconds-from-midnight value derived from a `SmallDateTime` + // minute field (`seconds_fragments * 60`). A valid minute field is + // 0..=1439, so a valid `mins` is < 86400; a hostile/buggy server can send a + // larger `u16` whose seconds overflow the day and make + // `from_num_seconds_from_midnight_opt` return `None`. Reject it as a + // protocol error rather than panicking on the `unwrap`. + NaiveTime::from_num_seconds_from_midnight_opt(mins, 0).ok_or_else(|| { + crate::Error::Protocol( + format!("smalldatetime seconds-of-day {mins} is out of range").into(), + ) + }) } #[inline] @@ -53,59 +102,94 @@ fn to_sec_fragments(time: NaiveTime) -> i64 { #[cfg(feature = "tds73")] from_sql!( NaiveDateTime: - ColumnData::SmallDateTime(ref dt) => dt.map(|dt| NaiveDateTime::new( - from_days(dt.days as i64, 1900), - from_mins(dt.seconds_fragments as u32 * 60), - )), - ColumnData::DateTime2(ref dt) => dt.map(|dt| NaiveDateTime::new( - from_days(dt.date.days() as i64, 1), - NaiveTime::from_hms_opt(0,0,0).unwrap() + chrono::Duration::nanoseconds(dt.time.increments as i64 * 10i64.pow(9 - dt.time.scale as u32)) - )), - ColumnData::DateTime(ref dt) => dt.map(|dt| NaiveDateTime::new( - from_days(dt.days as i64, 1900), - from_sec_fragments(dt.seconds_fragments as i64) - )); + ColumnData::SmallDateTime(ref dt) => match *dt { + Some(dt) => Some(NaiveDateTime::new( + from_days(dt.days as i64, 1900)?, + from_mins(dt.seconds_fragments as u32 * 60)?, + )), + None => None, + }, + ColumnData::DateTime2(ref dt) => match *dt { + Some(dt) => Some(NaiveDateTime::new( + from_days(dt.date.days() as i64, 1)?, + NaiveTime::from_hms_opt(0,0,0).unwrap() + chrono::Duration::nanoseconds(nanos_from_increments(dt.time.increments, dt.time.scale)) + )), + None => None, + }, + ColumnData::DateTime(ref dt) => match *dt { + Some(dt) => Some(NaiveDateTime::new( + from_days(dt.days as i64, 1900)?, + from_sec_fragments(dt.seconds_fragments as i64) + )), + None => None, + }; NaiveTime: - ColumnData::Time(ref time) => time.map(|time| { - let ns = time.increments as i64 * 10i64.pow(9 - time.scale as u32); - NaiveTime::from_hms_opt(0,0,0).unwrap() + chrono::Duration::nanoseconds(ns) - }); + ColumnData::Time(ref time) => match *time { + Some(time) => { + let ns = nanos_from_increments(time.increments, time.scale); + Some(NaiveTime::from_hms_opt(0,0,0).unwrap() + chrono::Duration::nanoseconds(ns)) + } + None => None, + }; NaiveDate: - ColumnData::Date(ref date) => date.map(|date| from_days(date.days() as i64, 1)); + ColumnData::Date(ref date) => match *date { + Some(date) => Some(from_days(date.days() as i64, 1)?), + None => None, + }; chrono::DateTime: - ColumnData::DateTimeOffset(ref dto) => dto.map(|dto| { - let date = from_days(dto.datetime2.date.days() as i64, 1); - let ns = dto.datetime2.time.increments as i64 * 10i64.pow(9 - dto.datetime2.time.scale as u32); + ColumnData::DateTimeOffset(ref dto) => match *dto { + Some(dto) => { + let date = from_days(dto.datetime2.date.days() as i64, 1)?; + let ns = nanos_from_increments(dto.datetime2.time.increments, dto.datetime2.time.scale); + let time = NaiveTime::from_hms_opt(0,0,0).unwrap() + chrono::Duration::nanoseconds(ns); + + let minutes = validate_offset_minutes(dto.offset)?; + let offset = chrono::Duration::minutes(minutes as i64); + let base = NaiveDateTime::new(date, time); + // A valid offset keeps the instant representable; a malformed one + // that pushes the value out of range is rejected as a protocol error. + let naive = base.checked_sub_signed(offset).ok_or_else(|| { + crate::Error::Protocol( + "datetimeoffset value is out of the representable range".into(), + ) + })?; + + Some(chrono::DateTime::from_naive_utc_and_offset(naive, Utc)) + } + None => None, + }, + ColumnData::DateTime2(ref dt2) => match *dt2 { + Some(dt2) => { + let date = from_days(dt2.date.days() as i64, 1)?; + let ns = nanos_from_increments(dt2.time.increments, dt2.time.scale); + let time = NaiveTime::from_hms_opt(0,0,0).unwrap() + chrono::Duration::nanoseconds(ns); + let naive = NaiveDateTime::new(date, time); + + Some(chrono::DateTime::from_naive_utc_and_offset(naive, Utc)) + } + None => None, + }; + chrono::DateTime: ColumnData::DateTimeOffset(ref dto) => match *dto { + Some(dto) => { + let date = from_days(dto.datetime2.date.days() as i64, 1)?; + let ns = nanos_from_increments(dto.datetime2.time.increments, dto.datetime2.time.scale); let time = NaiveTime::from_hms_opt(0,0,0).unwrap() + chrono::Duration::nanoseconds(ns); - let offset = chrono::Duration::minutes(dto.offset as i64); - let naive = NaiveDateTime::new(date, time).sub(offset); - - chrono::DateTime::from_naive_utc_and_offset(naive, Utc) - }), - ColumnData::DateTime2(ref dt2) => dt2.map(|dt2| { - let date = from_days(dt2.date.days() as i64, 1); - let ns = dt2.time.increments as i64 * 10i64.pow(9 - dt2.time.scale as u32); - let time = NaiveTime::from_hms_opt(0,0,0).unwrap() + chrono::Duration::nanoseconds(ns); + let minutes = validate_offset_minutes(dto.offset)?; + let offset = FixedOffset::east_opt(minutes * 60).ok_or_else(|| { + crate::Error::Protocol("datetimeoffset offset is not representable".into()) + })?; let naive = NaiveDateTime::new(date, time); - chrono::DateTime::from_naive_utc_and_offset(naive, Utc) - }); - chrono::DateTime: ColumnData::DateTimeOffset(ref dto) => dto.map(|dto| { - let date = from_days(dto.datetime2.date.days() as i64, 1); - let ns = dto.datetime2.time.increments as i64 * 10i64.pow(9 - dto.datetime2.time.scale as u32); - let time = NaiveTime::from_hms_opt(0,0,0).unwrap() + chrono::Duration::nanoseconds(ns); - - let offset = FixedOffset::east_opt((dto.offset as i32) * 60).unwrap(); - let naive = NaiveDateTime::new(date, time); - - chrono::DateTime::from_naive_utc_and_offset(naive, offset) - }) + Some(chrono::DateTime::from_naive_utc_and_offset(naive, offset)) + } + None => None, + } ); #[cfg(feature = "tds73")] to_sql!(self_, - NaiveDate: (ColumnData::Date, Date::new(to_days(*self_, 1) as u32)); + NaiveDate: (ColumnData::Date, Date::new_unchecked(to_days(*self_, 1) as u32)); NaiveTime: (ColumnData::Time, { use chrono::Timelike; @@ -121,7 +205,7 @@ to_sql!(self_, let nanos = time.num_seconds_from_midnight() as u64 * 1e9 as u64 + time.nanosecond() as u64; let increments = nanos / 100; - let date = Date::new(to_days(self_.date(), 1) as u32); + let date = Date::new_unchecked(to_days(self_.date(), 1) as u32); let time = Time {increments, scale: 7}; DateTime2::new(date, time) @@ -133,7 +217,7 @@ to_sql!(self_, let time = naive.time(); let nanos = time.num_seconds_from_midnight() as u64 * 1e9 as u64 + time.nanosecond() as u64; - let date = Date::new(to_days(naive.date(), 1) as u32); + let date = Date::new_unchecked(to_days(naive.date(), 1) as u32); let time = Time {increments: nanos / 100, scale: 7}; DateTime2::new(date, time) @@ -145,7 +229,7 @@ to_sql!(self_, let time = naive.time(); let nanos = time.num_seconds_from_midnight() as u64 * 1e9 as u64 + time.nanosecond() as u64; - let date = Date::new(to_days(naive.date(), 1) as u32); + let date = Date::new_unchecked(to_days(naive.date(), 1) as u32); let time = Time { increments: nanos / 100, scale: 7 }; let tz = self_.timezone(); @@ -157,7 +241,7 @@ to_sql!(self_, #[cfg(feature = "tds73")] into_sql!(self_, - NaiveDate: (ColumnData::Date, Date::new(to_days(self_, 1) as u32)); + NaiveDate: (ColumnData::Date, Date::new_unchecked(to_days(self_, 1) as u32)); NaiveTime: (ColumnData::Time, { use chrono::Timelike; @@ -173,7 +257,7 @@ into_sql!(self_, let nanos = time.num_seconds_from_midnight() as u64 * 1e9 as u64 + time.nanosecond() as u64; let increments = nanos / 100; - let date = Date::new(to_days(self_.date(), 1) as u32); + let date = Date::new_unchecked(to_days(self_.date(), 1) as u32); let time = Time {increments, scale: 7}; DateTime2::new(date, time) @@ -185,7 +269,7 @@ into_sql!(self_, let time = naive.time(); let nanos = time.num_seconds_from_midnight() as u64 * 1e9 as u64 + time.nanosecond() as u64; - let date = Date::new(to_days(naive.date(), 1) as u32); + let date = Date::new_unchecked(to_days(naive.date(), 1) as u32); let time = Time {increments: nanos / 100, scale: 7}; DateTime2::new(date, time) @@ -197,7 +281,7 @@ into_sql!(self_, let time = naive.time(); let nanos = time.num_seconds_from_midnight() as u64 * 1e9 as u64 + time.nanosecond() as u64; - let date = Date::new(to_days(naive.date(), 1) as u32); + let date = Date::new_unchecked(to_days(naive.date(), 1) as u32); let time = Time { increments: nanos / 100, scale: 7 }; let tz = self_.timezone(); @@ -236,8 +320,294 @@ into_sql!(self_, #[cfg(not(feature = "tds73"))] from_sql!( NaiveDateTime: - ColumnData::DateTime(ref dt) => dt.map(|dt| NaiveDateTime::new( - from_days(dt.days as i64, 1900), - from_sec_fragments(dt.seconds_fragments as i64) - )) + ColumnData::DateTime(ref dt) => match *dt { + Some(dt) => Some(NaiveDateTime::new( + from_days(dt.days as i64, 1900)?, + from_sec_fragments(dt.seconds_fragments as i64) + )), + None => None, + } ); + +#[cfg(test)] +mod tests { + use super::*; + use crate::{FromSql, IntoSql}; + + #[test] + fn from_days_out_of_range_errors() { + // A day offset far outside the representable `NaiveDate` range must + // return a protocol error rather than silently clamping to MIN/MAX + // (which would decode a malformed value to a plausible-but-wrong date). + for days in [200_000_000_i64, -200_000_000_i64] { + let err = from_days(days, 1).expect_err("out-of-range day offset must error"); + assert!( + matches!(err, crate::Error::Protocol(_)), + "expected a protocol error, got {err:?}" + ); + } + + // A valid in-range date still decodes correctly (happy path unchanged). + assert_eq!( + from_days(0, 1).unwrap(), + NaiveDate::from_ymd_opt(1, 1, 1).unwrap() + ); + } + + #[cfg(feature = "tds73")] + #[test] + fn validate_offset_minutes_rejects_out_of_range() { + // Valid SQL Server range -14:00..=+14:00 (±840 minutes) succeeds. + assert_eq!(validate_offset_minutes(60).unwrap(), 60); + assert_eq!(validate_offset_minutes(-840).unwrap(), -840); + + // A malformed offset beyond ±14h must error, not silently fall back to UTC. + for minutes in [841_i16, -841, 5000, -5000] { + let err = validate_offset_minutes(minutes).expect_err("out-of-range offset must error"); + assert!( + matches!(err, crate::Error::Protocol(_)), + "expected a protocol error, got {err:?}" + ); + } + } + + #[cfg(feature = "tds73")] + #[test] + fn datetimeoffset_out_of_range_offset_errors() { + // Build a DateTimeOffset with an offset well outside ±14h. Both the + // `DateTime` and `DateTime` decode arms must error. + let dt2 = DateTime2::new(Date::new(0), Time::new(0, 7)); + let dto = DateTimeOffset::new(dt2, 5000); + let data = ColumnData::DateTimeOffset(Some(dto)); + + let err = chrono::DateTime::::from_sql(&data) + .expect_err("out-of-range offset must error, not silently fall back"); + assert!( + matches!(err, crate::Error::Protocol(_)), + "expected a protocol error, got {err:?}" + ); + + let err = chrono::DateTime::::from_sql(&data) + .expect_err("out-of-range offset must error, not silently fall back"); + assert!( + matches!(err, crate::Error::Protocol(_)), + "expected a protocol error, got {err:?}" + ); + } + + #[test] + fn from_sec_fragments_converts() { + // 300 sec-fragments (1/300 s units) == exactly one second. + assert_eq!( + from_sec_fragments(300), + NaiveTime::from_hms_opt(0, 0, 1).unwrap() + ); + } + + #[cfg(feature = "tds73")] + #[test] + fn from_mins_converts() { + // `from_mins` takes seconds-from-midnight; 3600 s == 01:00:00. + assert_eq!( + from_mins(3600).unwrap(), + NaiveTime::from_hms_opt(1, 0, 0).unwrap() + ); + } + + #[cfg(feature = "tds73")] + #[test] + fn smalldatetime_out_of_range_minute_field_errors() { + // A `SmallDateTime` minute field is spec'd 0..=1439. A hostile/buggy + // server can send a larger `u16`; `seconds_fragments * 60` then + // overflows the day and previously panicked on the `unwrap` inside + // `from_mins`. It must now surface as a protocol error, and a valid + // value must still decode. + for minute_field in [1440u16, 65535] { + let sdt = crate::tds::time::SmallDateTime::new(0, minute_field); + let data = ColumnData::SmallDateTime(Some(sdt)); + let err = NaiveDateTime::from_sql(&data) + .expect_err("out-of-range minute field must error, not panic"); + assert!( + matches!(err, crate::Error::Protocol(_)), + "expected a protocol error, got {err:?}" + ); + } + + // 1439 (the maximum valid minute-of-day) still decodes to 23:59:00. + let sdt = crate::tds::time::SmallDateTime::new(0, 1439); + let data = ColumnData::SmallDateTime(Some(sdt)); + let decoded = NaiveDateTime::from_sql(&data).unwrap().unwrap(); + assert_eq!(decoded.time(), NaiveTime::from_hms_opt(23, 59, 0).unwrap()); + } + + #[cfg(not(feature = "tds73"))] + #[test] + fn to_sec_fragments_converts() { + // One second == 300 sec-fragments (1/300 s units). + assert_eq!( + to_sec_fragments(NaiveTime::from_hms_opt(0, 0, 1).unwrap()), + 300 + ); + } + + #[cfg(feature = "tds73")] + #[test] + fn naive_date_round_trip() { + let date = NaiveDate::from_ymd_opt(2021, 6, 15).unwrap(); + let cd: ColumnData<'static> = date.into_sql(); + assert!(matches!(cd, ColumnData::Date(Some(_)))); + assert_eq!(NaiveDate::from_sql(&cd).unwrap(), Some(date)); + } + + #[cfg(feature = "tds73")] + #[test] + fn naive_time_round_trip() { + let time = NaiveTime::from_hms_opt(13, 37, 42).unwrap(); + let cd: ColumnData<'static> = time.into_sql(); + assert!(matches!(cd, ColumnData::Time(Some(_)))); + assert_eq!(NaiveTime::from_sql(&cd).unwrap(), Some(time)); + } + + #[cfg(feature = "tds73")] + #[test] + fn naive_datetime_round_trip() { + let dt = NaiveDateTime::new( + NaiveDate::from_ymd_opt(2000, 12, 31).unwrap(), + NaiveTime::from_hms_opt(23, 59, 58).unwrap(), + ); + let cd: ColumnData<'static> = dt.into_sql(); + assert!(matches!(cd, ColumnData::DateTime2(Some(_)))); + assert_eq!(NaiveDateTime::from_sql(&cd).unwrap(), Some(dt)); + } + + #[cfg(feature = "tds73")] + #[test] + fn datetime_utc_round_trip() { + let naive = NaiveDateTime::new( + NaiveDate::from_ymd_opt(2015, 3, 4).unwrap(), + NaiveTime::from_hms_opt(1, 2, 3).unwrap(), + ); + let dt = chrono::DateTime::::from_naive_utc_and_offset(naive, Utc); + let cd: ColumnData<'static> = dt.into_sql(); + assert!(matches!(cd, ColumnData::DateTime2(Some(_)))); + assert_eq!(chrono::DateTime::::from_sql(&cd).unwrap(), Some(dt)); + } + + #[cfg(feature = "tds73")] + #[test] + fn datetime_fixed_offset_round_trip() { + let offset = FixedOffset::east_opt(2 * 3600).unwrap(); + let naive = NaiveDateTime::new( + NaiveDate::from_ymd_opt(2015, 3, 4).unwrap(), + NaiveTime::from_hms_opt(1, 2, 3).unwrap(), + ); + let dt = chrono::DateTime::from_naive_utc_and_offset(naive, offset); + let cd: ColumnData<'static> = dt.into_sql(); + assert!(matches!(cd, ColumnData::DateTimeOffset(Some(_)))); + assert_eq!( + chrono::DateTime::::from_sql(&cd).unwrap(), + Some(dt) + ); + } + + // The tds73 `from_sql` NaiveDateTime path has a dedicated arm for the legacy + // `ColumnData::DateTime` wire type; exercise it directly (round-trips produce + // `DateTime2`, never `DateTime`, so this arm is otherwise unreachable). + #[cfg(feature = "tds73")] + #[test] + fn naive_datetime_from_legacy_datetime_column() { + let cd = ColumnData::DateTime(Some(crate::tds::time::DateTime::new(0, 0))); + let dt = NaiveDateTime::from_sql(&cd).unwrap().unwrap(); + assert_eq!( + dt, + NaiveDateTime::new( + NaiveDate::from_ymd_opt(1900, 1, 1).unwrap(), + NaiveTime::from_hms_opt(0, 0, 0).unwrap(), + ) + ); + } + + // The `DateTimeOffset -> DateTime` conversion arm (distinct from the + // `DateTime` arm exercised by the round-trip test). + #[cfg(feature = "tds73")] + #[test] + fn datetime_offset_reads_as_utc() { + let offset = FixedOffset::east_opt(2 * 3600).unwrap(); + let naive = NaiveDateTime::new( + NaiveDate::from_ymd_opt(2015, 3, 4).unwrap(), + NaiveTime::from_hms_opt(1, 2, 3).unwrap(), + ); + let dt: chrono::DateTime = + chrono::DateTime::from_naive_utc_and_offset(naive, offset); + let cd: ColumnData<'static> = dt.into_sql(); + assert!(matches!(cd, ColumnData::DateTimeOffset(Some(_)))); + + let utc = chrono::DateTime::::from_sql(&cd).unwrap(); + assert!(utc.is_some()); + } + + #[cfg(feature = "tds73")] + #[test] + fn out_of_range_naivedate_errors_at_encode_not_panic() { + use crate::tds::codec::Encode; + use bytes::BytesMut; + + // A `NaiveDate` legal for chrono but far outside SQL Server's 3-byte day + // range. The conversion (`into_sql`) must not panic; the out-of-range + // value must instead surface as a `Result::Err` at encode time. + let d = NaiveDate::from_ymd_opt(100_000, 1, 1).unwrap(); + let cd: ColumnData<'static> = d.into_sql(); + match cd { + ColumnData::Date(Some(date)) => { + let mut buf = BytesMut::new(); + assert!( + date.encode(&mut buf).is_err(), + "out-of-range date must error at encode, not silently truncate" + ); + } + other => panic!("expected ColumnData::Date, got {other:?}"), + } + } + + #[cfg(feature = "tds73")] + #[test] + fn out_of_range_naivedatetime_errors_at_encode_not_panic() { + use crate::tds::codec::Encode; + use bytes::BytesMut; + + let dt = NaiveDateTime::new( + NaiveDate::from_ymd_opt(100_000, 1, 1).unwrap(), + NaiveTime::from_hms_opt(0, 0, 0).unwrap(), + ); + let cd: ColumnData<'static> = dt.into_sql(); + match cd { + ColumnData::DateTime2(Some(dt2)) => { + let mut buf = BytesMut::new(); + assert!( + dt2.encode(&mut buf).is_err(), + "out-of-range datetime2 must error at encode, not silently truncate" + ); + } + other => panic!("expected ColumnData::DateTime2, got {other:?}"), + } + } + + #[cfg(feature = "tds73")] + #[test] + fn null_maps_to_none() { + assert_eq!(NaiveDate::from_sql(&ColumnData::Date(None)).unwrap(), None); + assert_eq!(NaiveTime::from_sql(&ColumnData::Time(None)).unwrap(), None); + } + + #[cfg(not(feature = "tds73"))] + #[test] + fn naive_datetime_round_trip_legacy() { + let dt = NaiveDateTime::new( + NaiveDate::from_ymd_opt(1990, 1, 1).unwrap(), + NaiveTime::from_hms_opt(12, 0, 0).unwrap(), + ); + let cd: ColumnData<'static> = dt.into_sql(); + assert!(matches!(cd, ColumnData::DateTime(Some(_)))); + assert_eq!(NaiveDateTime::from_sql(&cd).unwrap(), Some(dt)); + } +} diff --git a/src/tds/time/time.rs b/src/tds/time/time.rs index 5a2b1cfaa..ae165c891 100644 --- a/src/tds/time/time.rs +++ b/src/tds/time/time.rs @@ -10,15 +10,73 @@ pub use time::{Date, Month, OffsetDateTime, PrimitiveDateTime, Time, UtcOffset}; use crate::tds::codec::ColumnData; #[inline] -fn from_days(days: u64, start_year: i32) -> Date { - Date::from_calendar_date(start_year, Month::January, 1).unwrap() - + Duration::from_secs(60 * 60 * 24 * days) +fn from_days(days: i64, start_year: i32) -> crate::Result { + // Use the signed `time::Duration` so that negative day offsets (dates + // before `start_year`, e.g. `datetime` values prior to 1900) do not + // overflow. Casting a negative day count into an unsigned type and + // multiplying it out panics with "multiply with overflow". + // + // `days` ultimately comes from untrusted server bytes, so a malformed value + // can land outside the range `time::Date` can represent. Every valid SQL + // date (SQL's 0001-01-01..=9999-12-31 all fit within `time::Date`) succeeds; + // a genuinely out-of-range/malformed day offset is rejected as a protocol + // error rather than silently clamped to MIN/MAX (which would yield a wrong + // date). + let base = Date::from_calendar_date(start_year, Month::January, 1).unwrap(); + base.checked_add(time::Duration::days(days)).ok_or_else(|| { + crate::Error::Protocol( + format!("date day offset {days} is out of the representable range").into(), + ) + }) } +/// Validate and convert a server-supplied UTC offset (in whole minutes) into a +/// `UtcOffset`. SQL Server's `datetimeoffset` is only valid for the range +/// -14:00..=+14:00; a malformed offset outside that range is rejected as a +/// protocol error rather than silently falling back to UTC (which would shift +/// the represented instant). #[inline] #[cfg(feature = "tds73")] -fn from_secs(secs: u64) -> Time { - Time::from_hms(0, 0, 0).unwrap() + Duration::from_secs(secs) +fn offset_from_minutes(minutes: i16) -> crate::Result { + if !(-840..=840).contains(&minutes) { + return Err(crate::Error::Protocol( + format!( + "datetimeoffset offset {minutes} minutes is outside the valid -14:00..=+14:00 range" + ) + .into(), + )); + } + + UtcOffset::from_whole_seconds(minutes as i32 * 60) + .map_err(|_| crate::Error::Protocol("datetimeoffset offset is not representable".into())) +} + +/// Convert a server-supplied fractional-seconds `increments` at the given +/// `scale` into nanoseconds without panicking. `scale` and `increments` are +/// untrusted; a `scale > 9` would otherwise underflow `9 - scale`, and a large +/// `increments` would overflow the multiply. +#[inline] +#[cfg(feature = "tds73")] +fn nanos_from_increments(increments: u64, scale: u8) -> u64 { + let pow = 9u32.saturating_sub(scale as u32); + increments.saturating_mul(10u64.saturating_pow(pow)) +} + +#[inline] +#[cfg(feature = "tds73")] +fn from_secs(secs: u64) -> crate::Result1"); + data.set_schema(schema.clone()); + + let stored = data.schema().expect("schema present"); + assert_eq!(stored.db_name(), "db"); + assert_eq!(stored.owner(), "owner"); + assert_eq!(stored.collection(), "collection"); + } + + #[test] + fn encode_writes_plp_header_and_backpatches_length() { + let mut buf = BytesMut::new(); + XmlData::new("ab") + .encode(&mut buf) + .expect("encode succeeds"); + + // 8 (unknown-size marker) + 4 (length) + 2*2 (utf16 chars) + 4 (terminator) + assert_eq!(buf.len(), 8 + 4 + 4 + 4); + + // unknown size marker + assert_eq!(&buf[0..8], &0xfffffffffffffffe_u64.to_le_bytes()); + // backpatched length is number of chars * 2 bytes + assert_eq!(&buf[8..12], &(4u32).to_le_bytes()); + // 'a' then 'b' as UTF-16LE + assert_eq!(&buf[12..16], &[b'a', 0, b'b', 0]); + // PLP terminator + assert_eq!(&buf[16..20], &(0u32).to_le_bytes()); + } +} diff --git a/src/to_sql.rs b/src/to_sql.rs index cde353cd1..b7a6a2969 100644 --- a/src/to_sql.rs +++ b/src/to_sql.rs @@ -32,7 +32,8 @@ use uuid::Uuid; /// |[`XmlData`]|`xml`| /// |[`NaiveDate`] (with `chrono` feature, TDS 7.3 >)|`date`| /// |[`NaiveTime`] (with `chrono` feature, TDS 7.3 >)|`time`| -/// |[`DateTime`] (with `chrono` feature, TDS 7.3 >)|`datetimeoffset`| +/// |[`DateTime`]`` (with `chrono` feature, TDS 7.3 >)|`datetime2`| +/// |[`DateTime`]`` (with `chrono` feature, TDS 7.3 >)|`datetimeoffset`| /// |[`NaiveDateTime`] (with `chrono` feature, TDS 7.3 >)|`datetime2`| /// |[`NaiveDateTime`] (with `chrono` feature, TDS 7.2)|`datetime`| /// @@ -199,3 +200,208 @@ to_sql!(self_, XmlData: (ColumnData::Xml, Cow::Borrowed(self_)); Uuid: (ColumnData::Guid, *self_); ); + +#[cfg(test)] +mod tests { + use super::*; + use crate::tds::Numeric; + use crate::{IntoSql, ToSql}; + + #[test] + fn to_sql_scalars() { + assert_eq!(true.to_sql(), ColumnData::Bit(Some(true))); + assert_eq!(8u8.to_sql(), ColumnData::U8(Some(8))); + assert_eq!(16i16.to_sql(), ColumnData::I16(Some(16))); + assert_eq!(32i32.to_sql(), ColumnData::I32(Some(32))); + assert_eq!(64i64.to_sql(), ColumnData::I64(Some(64))); + assert_eq!(1.5f32.to_sql(), ColumnData::F32(Some(1.5))); + assert_eq!(2.5f64.to_sql(), ColumnData::F64(Some(2.5))); + } + + #[test] + // The `&Some(..)`/`&None` borrows are intentional: they exercise the + // `ToSql for &T` impls, not the by-value ones, so the borrow is not needless. + #[allow(clippy::needless_borrow)] + fn to_sql_option_some_and_none() { + assert_eq!(Some(1i32).to_sql(), ColumnData::I32(Some(1))); + assert_eq!(None::.to_sql(), ColumnData::I32(None)); + assert_eq!((&Some(1i32)).to_sql(), ColumnData::I32(Some(1))); + assert_eq!((&None::).to_sql(), ColumnData::I32(None)); + } + + #[test] + fn to_sql_strings_and_binary() { + assert_eq!("abc".to_sql(), ColumnData::String(Some(Cow::from("abc")))); + assert_eq!( + String::from("abc").to_sql(), + ColumnData::String(Some(Cow::from("abc"))) + ); + let v = vec![1u8, 2, 3]; + assert_eq!( + v.to_sql(), + ColumnData::Binary(Some(Cow::from(vec![1, 2, 3]))) + ); + assert_eq!( + [1u8, 2, 3].as_slice().to_sql(), + ColumnData::Binary(Some(Cow::from(vec![1, 2, 3]))) + ); + } + + #[test] + fn to_sql_numeric_and_uuid() { + let n = Numeric::new_with_scale(5, 1); + assert_eq!(n.to_sql(), ColumnData::Numeric(Some(n))); + + let uuid = Uuid::nil(); + assert_eq!(uuid.to_sql(), ColumnData::Guid(Some(uuid))); + } + + #[test] + fn into_sql_borrowed_and_owned() { + assert_eq!( + "abc".into_sql(), + ColumnData::String(Some(Cow::Borrowed("abc"))) + ); + assert_eq!( + Some("abc").into_sql(), + ColumnData::String(Some(Cow::Borrowed("abc"))) + ); + assert_eq!(None::<&str>.into_sql(), ColumnData::String(None)); + + let bytes = vec![9u8, 8, 7]; + assert_eq!( + bytes.as_slice().into_sql(), + ColumnData::Binary(Some(Cow::Borrowed(bytes.as_slice()))) + ); + assert_eq!( + (&bytes).into_sql(), + ColumnData::Binary(Some(Cow::from(&bytes))) + ); + + let uuid = Uuid::nil(); + assert_eq!((&uuid).into_sql(), ColumnData::Guid(Some(uuid))); + assert_eq!(Some(&uuid).into_sql(), ColumnData::Guid(Some(uuid))); + assert_eq!(None::<&Uuid>.into_sql(), ColumnData::Guid(None)); + } + + #[test] + fn into_sql_scalars() { + assert_eq!(true.into_sql(), ColumnData::Bit(Some(true))); + assert_eq!(5i32.into_sql(), ColumnData::I32(Some(5))); + assert_eq!(None::.into_sql(), ColumnData::I32(None)); + assert_eq!( + String::from("x").into_sql(), + ColumnData::String(Some(Cow::from("x"))) + ); + } + + #[test] + fn into_sql_owned_string_and_ref() { + let owned = String::from("abc"); + assert_eq!( + (&owned).into_sql(), + ColumnData::String(Some(Cow::from("abc"))) + ); + assert_eq!( + Some(&owned).into_sql(), + ColumnData::String(Some(Cow::from("abc"))) + ); + assert_eq!(None::<&String>.into_sql(), ColumnData::String(None)); + } + + #[test] + fn into_sql_binary_option_variants() { + assert_eq!(None::<&[u8]>.into_sql(), ColumnData::Binary(None)); + + let bytes = vec![1u8, 2, 3]; + assert_eq!( + Some(bytes.as_slice()).into_sql(), + ColumnData::Binary(Some(Cow::from(bytes.as_slice()))) + ); + assert_eq!(None::<&Vec>.into_sql(), ColumnData::Binary(None)); + assert_eq!( + bytes.into_sql(), + ColumnData::Binary(Some(Cow::from(vec![1, 2, 3]))) + ); + } + + #[test] + fn into_sql_cow_variants() { + let cow_str: Cow<'_, str> = Cow::Borrowed("hi"); + assert_eq!( + cow_str.into_sql(), + ColumnData::String(Some(Cow::from("hi"))) + ); + assert_eq!( + Some(Cow::Borrowed("hi")).into_sql(), + ColumnData::String(Some(Cow::from("hi"))) + ); + assert_eq!(None::>.into_sql(), ColumnData::String(None)); + + let cow_bin: Cow<'_, [u8]> = Cow::Borrowed(&[1u8, 2][..]); + assert_eq!( + cow_bin.into_sql(), + ColumnData::Binary(Some(Cow::from(vec![1u8, 2]))) + ); + assert_eq!( + Some(Cow::<[u8]>::Borrowed(&[1u8, 2][..])).into_sql(), + ColumnData::Binary(Some(Cow::from(vec![1u8, 2]))) + ); + assert_eq!(None::>.into_sql(), ColumnData::Binary(None)); + } + + #[test] + fn into_sql_xml_and_numeric() { + let xml = XmlData::new("".to_string()); + assert_eq!( + (&xml).into_sql(), + ColumnData::Xml(Some(Cow::Borrowed(&xml))) + ); + assert_eq!( + Some(&xml).into_sql(), + ColumnData::Xml(Some(Cow::Borrowed(&xml))) + ); + assert_eq!(None::<&XmlData>.into_sql(), ColumnData::Xml(None)); + + let xml_owned = XmlData::new("".to_string()); + assert_eq!( + xml_owned.clone().into_sql(), + ColumnData::Xml(Some(Cow::Owned(xml_owned))) + ); + + let n = Numeric::new_with_scale(42, 0); + assert_eq!(n.into_sql(), ColumnData::Numeric(Some(n))); + } + + #[test] + // The `&value` borrows are intentional: they exercise the `ToSql for &T` + // impls for the base scalar types, so the borrow is not needless. + #[allow(clippy::needless_borrow)] + fn to_sql_by_reference_scalars() { + // The macro-generated impls also cover `&T` for the base scalar types. + assert_eq!((&true).to_sql(), ColumnData::Bit(Some(true))); + assert_eq!((&8u8).to_sql(), ColumnData::U8(Some(8))); + assert_eq!((&16i16).to_sql(), ColumnData::I16(Some(16))); + assert_eq!((&64i64).to_sql(), ColumnData::I64(Some(64))); + assert_eq!((&1.5f32).to_sql(), ColumnData::F32(Some(1.5))); + assert_eq!((&2.5f64).to_sql(), ColumnData::F64(Some(2.5))); + } + + #[test] + fn to_sql_cow_variants() { + let cow_str: Cow<'_, str> = Cow::Borrowed("hi"); + assert_eq!(cow_str.to_sql(), ColumnData::String(Some(Cow::from("hi")))); + + let cow_bin: Cow<'_, [u8]> = Cow::Borrowed(&[1u8, 2][..]); + assert_eq!( + cow_bin.to_sql(), + ColumnData::Binary(Some(Cow::from(vec![1u8, 2]))) + ); + } + + #[test] + fn to_sql_xml() { + let xml = XmlData::new("".to_string()); + assert_eq!(xml.to_sql(), ColumnData::Xml(Some(Cow::Borrowed(&xml)))); + } +} diff --git a/tests/bulk.rs b/tests/bulk.rs index 33b90637a..e2655c88d 100644 --- a/tests/bulk.rs +++ b/tests/bulk.rs @@ -4,6 +4,7 @@ use once_cell::sync::Lazy; use std::cell::RefCell; use std::env; use std::sync::Once; +use tiberius::ColumnData; use tiberius::{IntoSql, Result, TokenRow}; #[cfg(all(feature = "tds73", feature = "chrono"))] @@ -25,7 +26,7 @@ static CONN_STR: Lazy = Lazy::new(|| { thread_local! { static NAMES: RefCell>> = - RefCell::new(None); + const { RefCell::new(None) }; } async fn random_table() -> String { @@ -148,6 +149,413 @@ test_bulk_type!(varchar_limited( vec!["aaaaaaaaaaaaaaaaaaaaaaa"; 1000].into_iter() )); +// Column types added by 97bbbfd (bulk support for #352/#358) that previously +// had no bulk coverage. `text`/`ntext` exercise the COLMETADATA TableName path +// (MS-TDS §2.2.7.4): without emitting TableName for these types the server +// rejects the bulk COLMETADATA, so these tests only pass with that fix in place. +test_bulk_type!(text( + "TEXT", + 1000, + vec!["some text value"; 1000].into_iter() +)); +test_bulk_type!(ntext( + "NTEXT", + 1000, + vec!["some ntext välue"; 1000].into_iter() +)); + +// `money`/`smallmoney` exercise the f64 money encoder. +test_bulk_type!(money("MONEY", 1000, vec![1234.5678f64; 1000].into_iter())); +test_bulk_type!(smallmoney( + "SMALLMONEY", + 1000, + vec![12.3456f64; 1000].into_iter() +)); + +// `numeric(p,s)` exercises the exact Numeric->wire path. +test_bulk_type!(numeric_28_4( + "NUMERIC(28,4)", + 1000, + vec![tiberius::numeric::Numeric::new_with_scale(12345, 4); 1000].into_iter() +)); + +// The `test_bulk_type!` cases above only assert the inserted row count. The +// following tests bulk-insert a known value and read it back, asserting the +// exact value survived the round-trip through our bulk encoders. (Requires a +// live SQL Server; compiles locally but only runs in CI.) + +#[test_on_runtimes] +async fn bulk_money_value_roundtrips(mut conn: tiberius::Client) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + let table = format!("##{}", random_table().await); + + conn.execute( + &format!("CREATE TABLE {} (content MONEY NOT NULL)", table), + &[], + ) + .await?; + + let mut req = conn.bulk_insert(&table).await?; + let mut row = TokenRow::new(); + row.push(1234.5678f64.into_sql()); + req.send(row).await?; + let res = req.finalize().await?; + assert_eq!(1, res.total()); + + let value: f64 = conn + .query(&format!("SELECT content FROM {}", table), &[]) + .await? + .into_row() + .await? + .unwrap() + .get(0) + .unwrap(); + + assert!((value - 1234.5678).abs() < 1e-6, "got {value}"); + + Ok(()) +} + +#[test_on_runtimes] +async fn bulk_numeric_value_roundtrips(mut conn: tiberius::Client) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + use tiberius::numeric::Numeric; + + let table = format!("##{}", random_table().await); + + conn.execute( + &format!("CREATE TABLE {} (content NUMERIC(28,4) NOT NULL)", table), + &[], + ) + .await?; + + // A magnitude whose scaled form exceeds 2^53, so an f64 detour would lose + // precision but the exact integer path must not. + let num = Numeric::new_with_scale(123_456_789_012_345_678, 4); + + let mut req = conn.bulk_insert(&table).await?; + let mut row = TokenRow::new(); + row.push(num.into_sql()); + req.send(row).await?; + let res = req.finalize().await?; + assert_eq!(1, res.total()); + + let value: Numeric = conn + .query(&format!("SELECT content FROM {}", table), &[]) + .await? + .into_row() + .await? + .unwrap() + .get(0) + .unwrap(); + + assert_eq!(value, num); + + Ok(()) +} + +#[test_on_runtimes] +async fn bulk_text_value_roundtrips(mut conn: tiberius::Client) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + let table = format!("##{}", random_table().await); + + conn.execute( + &format!("CREATE TABLE {} (content TEXT NOT NULL)", table), + &[], + ) + .await?; + + let expected = "hello bulk text"; + let mut req = conn.bulk_insert(&table).await?; + let mut row = TokenRow::new(); + row.push(expected.into_sql()); + req.send(row).await?; + let res = req.finalize().await?; + assert_eq!(1, res.total()); + + let row = conn + .query(&format!("SELECT content FROM {}", table), &[]) + .await? + .into_row() + .await? + .unwrap(); + let value: &str = row.get(0).unwrap(); + + assert_eq!(value, expected); + + Ok(()) +} + +#[test_on_runtimes] +async fn bulk_ntext_value_roundtrips(mut conn: tiberius::Client) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + let table = format!("##{}", random_table().await); + + conn.execute( + &format!("CREATE TABLE {} (content NTEXT NOT NULL)", table), + &[], + ) + .await?; + + let expected = "héllo bulk ñtext"; + let mut req = conn.bulk_insert(&table).await?; + let mut row = TokenRow::new(); + row.push(expected.into_sql()); + req.send(row).await?; + let res = req.finalize().await?; + assert_eq!(1, res.total()); + + let row = conn + .query(&format!("SELECT content FROM {}", table), &[]) + .await? + .into_row() + .await? + .unwrap(); + let value: &str = row.get(0).unwrap(); + + assert_eq!(value, expected); + + Ok(()) +} + +/// Bulk-insert "Привет" into Cyrillic_General_CI_AS (CP1251) columns of +/// `table` and check the stored bytes. The database default collation must +/// not use CP1251 for this to detect a missing `COLLATE` in `INSERT BULK`: +/// the server would then read the CP1251 bytes in the default code page. +async fn bulk_cyrillic_roundtrip(conn: &mut tiberius::Client, table: &str) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + let default_code_page = conn + .query( + "SELECT CONVERT(INT, COLLATIONPROPERTY(CONVERT(NVARCHAR(128), \ + DATABASEPROPERTYEX(DB_NAME(), 'Collation')), 'CodePage'))", + &[], + ) + .await? + .into_row() + .await? + .unwrap() + .get::(0); + assert_ne!(Some(1251), default_code_page); + + // A plain batch, not `execute`: a `#` temp table created inside the + // sp_executesql call `execute` sends is dropped when that call returns. + conn.simple_query(format!( + "CREATE TABLE {} (id INT NOT NULL, \ + v VARCHAR(20) COLLATE Cyrillic_General_CI_AS NOT NULL, \ + c CHAR(6) COLLATE Cyrillic_General_CI_AS NOT NULL, \ + t TEXT COLLATE Cyrillic_General_CI_AS NOT NULL, \ + n NVARCHAR(20) COLLATE Cyrillic_General_CI_AS NOT NULL)", + table + )) + .await? + .into_results() + .await?; + + let expected = "Привет"; + let mut req = conn.bulk_insert(table).await?; + let mut row = TokenRow::new(); + row.push(1i32.into_sql()); + row.push(expected.into_sql()); + row.push(expected.into_sql()); + row.push(expected.into_sql()); + row.push(expected.into_sql()); + req.send(row).await?; + assert_eq!(1, req.finalize().await?.total()); + + let row = conn + .query( + &format!( + "SELECT v, c, t, n, CONVERT(VARBINARY(20), v), CONVERT(VARBINARY(20), c), \ + CONVERT(VARBINARY(20), CONVERT(VARCHAR(20), t)) FROM {}", + table + ), + &[], + ) + .await? + .into_row() + .await? + .unwrap(); + + let cp1251: &[u8] = &[0xCF, 0xF0, 0xE8, 0xE2, 0xE5, 0xF2]; + assert_eq!(Some(expected), row.get::<&str, _>(0)); + assert_eq!(Some(expected), row.get::<&str, _>(1)); + assert_eq!(Some(expected), row.get::<&str, _>(2)); + assert_eq!(Some(expected), row.get::<&str, _>(3)); + assert_eq!(Some(cp1251), row.get::<&[u8], _>(4)); + assert_eq!(Some(cp1251), row.get::<&[u8], _>(5)); + assert_eq!(Some(cp1251), row.get::<&[u8], _>(6)); + + Ok(()) +} + +#[test_on_runtimes] +async fn bulk_text_keeps_a_non_default_column_collation( + mut conn: tiberius::Client, +) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + // A table in the current database, named with its schema. + let table = format!("dbo.bulk_collate_{}", random_table().await); + + let result = bulk_cyrillic_roundtrip(&mut conn, &table).await; + drop_table(&mut conn, &format!("N'{table}'"), &table).await?; + + result +} + +/// Drops `table`, whose `OBJECT_ID` name is `object_name`, in a plain batch. +async fn drop_table(conn: &mut tiberius::Client, object_name: &str, table: &str) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + conn.simple_query(format!( + "IF OBJECT_ID({object_name}) IS NOT NULL DROP TABLE {table}" + )) + .await? + .into_results() + .await?; + + Ok(()) +} + +#[test_on_runtimes] +async fn bulk_text_keeps_a_non_default_column_collation_in_a_temp_table( + mut conn: tiberius::Client, +) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + let table = format!("#{}", random_table().await); + + let result = bulk_cyrillic_roundtrip(&mut conn, &table).await; + drop_table(&mut conn, &format!("N'tempdb..{table}'"), &table).await?; + + result +} + +/// Bulk-insert every byte 0x80..=0xFF of the code page of `collation`, as +/// text decoded by the client, into char, varchar, varchar(max) and text +/// columns of the # temp table `table` and check the stored bytes. +async fn bulk_every_high_byte_roundtrip( + conn: &mut tiberius::Client, + table: &str, + collation: &str, +) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + let bytes: Vec = (0x80..=0xFF).collect(); + let hex: String = bytes.iter().map(|b| format!("{b:02X}")).collect(); + + conn.simple_query(format!( + "CREATE TABLE {table} (c CHAR(128) COLLATE {collation} NOT NULL, \ + v VARCHAR(128) COLLATE {collation} NOT NULL, \ + m VARCHAR(MAX) COLLATE {collation} NOT NULL, \ + t TEXT COLLATE {collation} NOT NULL)" + )) + .await? + .into_results() + .await?; + + // The text of the 128 bytes: a varchar value of the collation (a binary + // value converts to it byte for byte), decoded by the client, and the + // server's own decoding of it for comparison. + let row = conn + .simple_query(format!( + "DECLARE @t TABLE (v VARCHAR(128) COLLATE {collation}); \ + INSERT INTO @t VALUES (0x{hex}); \ + SELECT v, CONVERT(NVARCHAR(128), v), CONVERT(VARBINARY(128), v) FROM @t" + )) + .await? + .into_row() + .await? + .unwrap(); + assert_eq!( + Some(bytes.as_slice()), + row.get::<&[u8], _>(2), + "{collation}" + ); + let text = row.get::<&str, _>(0).unwrap().to_owned(); + assert_eq!(Some(text.as_str()), row.get::<&str, _>(1), "{collation}"); + assert_eq!(128, text.chars().count(), "{collation}"); + + let mut req = conn.bulk_insert(table).await?; + let mut row = TokenRow::new(); + for _ in 0..4 { + row.push(text.clone().into_sql()); + } + req.send(row).await?; + assert_eq!(1, req.finalize().await?.total()); + + let row = conn + .simple_query(format!( + "SELECT CONVERT(VARBINARY(MAX), c), CONVERT(VARBINARY(MAX), v), \ + CONVERT(VARBINARY(MAX), m), \ + CONVERT(VARBINARY(MAX), CONVERT(VARCHAR(MAX), t)) FROM {table}" + )) + .await? + .into_row() + .await? + .unwrap(); + + for (i, column) in ["c", "v", "m", "t"].into_iter().enumerate() { + assert_eq!( + Some(bytes.as_slice()), + row.get::<&[u8], _>(i), + "{collation} column {column}" + ); + } + + Ok(()) +} + +#[test_on_runtimes] +async fn bulk_legacy_code_pages_store_every_high_byte( + mut conn: tiberius::Client, +) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + // Detecting a missing `COLLATE` in `INSERT BULK` needs a database default + // code page other than the columns'. + let default_code_page = conn + .query( + "SELECT CONVERT(INT, COLLATIONPROPERTY(CONVERT(NVARCHAR(128), \ + DATABASEPROPERTYEX(DB_NAME(), 'Collation')), 'CodePage'))", + &[], + ) + .await? + .into_row() + .await? + .unwrap() + .get::(0); + + for (collation, code_page) in [ + ("SQL_Latin1_General_CP437_BIN", 437), + ("SQL_1xCompat_CP850_CI_AS", 850), + ] { + assert_ne!(Some(code_page), default_code_page); + + let table = format!("#{}", random_table().await); + let result = bulk_every_high_byte_roundtrip(&mut conn, &table, collation).await; + drop_table(&mut conn, &format!("N'tempdb..{table}'"), &table).await?; + result?; + } + + Ok(()) +} + #[cfg(all(feature = "tds73", feature = "chrono"))] test_bulk_type!(datetime2( "DATETIME2", @@ -218,3 +626,350 @@ test_bulk_type!(datetime2_7( 100, vec![DateTime::from_timestamp(1658524194, 123456789); 100].into_iter() )); + +#[test_on_runtimes] +async fn read_and_write_to_keyword_columns(mut conn: tiberius::Client) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + let table = format!("##{}", random_table().await); + + conn.simple_query(format!("CREATE TABLE {} ([End] INT)", table)) + .await?; + + let mut req = conn.bulk_insert(&table).await.unwrap(); + for num in [6, 7, 8] { + let mut row = TokenRow::new(); + row.push(ColumnData::I32(Some(num))); + req.send(row).await.unwrap(); + } + let result = req.finalize().await.unwrap(); + assert_eq!(result.rows_affected(), &[3]); + + let rows = conn + .query(format!("SELECT [End] FROM {}", table), &[]) + .await? + .into_first_result() + .await?; + + assert_eq!(rows.len(), 3); + assert_eq!(Some(6), rows[0].get(0)); + assert_eq!(Some(7), rows[1].get(0)); + assert_eq!(Some(8), rows[2].get(0)); + + Ok(()) +} + +macro_rules! test_bulk_columns { + ($name:ident($total_generated:literal $(, $sql_type:literal)+ $(, ($cols:expr, $generator:expr ))+ $(,)?)) => { + paste::item! { + #[test_on_runtimes] + async fn [< bulk_load_optional_ $name >](mut conn: tiberius::Client) -> Result<()> + where + S: AsyncRead + AsyncWrite + Unpin + Send, + { + use tiberius::IntoRow; + + let table = format!("##{}", random_table().await); + let column_defs = &[$($sql_type,)+]; + + conn.execute( + &format!( + "CREATE TABLE {} (id INT IDENTITY PRIMARY KEY, {})", + table, + column_defs.join(", "), + ), + &[], + ) + .await?; + + let mut count = 0; + + $( + let mut req = conn.bulk_insert_columns(&table, $cols).await?; + for i in $generator { + let row = i.into_row(); + req.send(row).await?; + } + + let res = req.finalize().await?; + count += res.total(); + )+ + assert_eq!($total_generated, count); + + Ok(()) + } + + #[test_on_runtimes] + async fn [< bulk_load_required_ $name >](mut conn: tiberius::Client) -> Result<()> + where + S: AsyncRead + AsyncWrite + Unpin + Send, + { + use tiberius::IntoRow; + let table = format!("##{}", random_table().await); + let column_defs = &[$(format!("{} NOT NULL", $sql_type),)+]; + + conn.execute( + &format!( + "CREATE TABLE {} (id INT IDENTITY PRIMARY KEY, {})", + table, + column_defs.join(", "), + ), + &[], + ) + .await?; + + let mut count = 0; + + $( + let mut req = conn.bulk_insert_columns(&table, $cols).await?; + for i in $generator { + let row = i.into_row(); + req.send(row).await?; + } + + let res = req.finalize().await?; + count += res.total(); + )+ + assert_eq!($total_generated, count); + + Ok(()) + } + + } + }; +} + +test_bulk_columns!(ab_ba_default_columns( + 200, + "a INT", + "b FLOAT", + "c INT DEFAULT 0", + (&["a", "b"], vec![(1i32, 1f64); 100]), + (&["b", "a"], vec![(2f64, 2i32); 100]), +)); + +test_bulk_columns!(ab_ba_override_default_columns( + 200, + "a INT", + "b FLOAT", + "c INT DEFAULT 0", + (&["a", "b", "c"], vec![(1i32, 1f64, 10i32); 100]), + (&["b", "c", "a"], vec![(2f64, 20i32, 2i32); 100]), +)); + +// Server-gated regression for the COLMETADATA `Flags` bit layout (MS-TDS +// §2.2.7.4): the bulk-insert column filter in `src/client.rs` selects only +// columns whose flags report them as `Updateable` and non-`Identity`. If the +// `Identity` or `Computed` bit is misdecoded, an identity/computed column would +// wrongly be treated as a bulk target (or a real target wrongly skipped) and +// the server would reject the row set. This drives a real server to prove the +// flags cause the identity (and computed) columns to be skipped while the +// normal column lands. Compiles locally; runs only with a live server in CI. +#[test_on_runtimes] +async fn bulk_insert_skips_identity_and_computed_columns( + mut conn: tiberius::Client, +) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + let table = format!("##{}", random_table().await); + + // `id` is IDENTITY (fIdentity set, not writeable), `c` is COMPUTED + // (fComputed set, not writeable); only `val` is a valid bulk target. + conn.execute( + &format!( + "CREATE TABLE {} (id INT IDENTITY PRIMARY KEY, val INT NOT NULL, c AS (val + 1))", + table + ), + &[], + ) + .await?; + + let mut req = conn.bulk_insert(&table).await?; + for v in 0..10i32 { + let mut row = TokenRow::new(); + row.push(v.into_sql()); + req.send(row).await?; + } + let res = req.finalize().await?; + assert_eq!(10, res.total()); + + // All rows landed, identity auto-populated, and the computed column + // reflects `val + 1` — confirming the identity/computed columns were + // correctly excluded from the bulk column list. + let count: i32 = conn + .query(&format!("SELECT COUNT(*) FROM {}", table), &[]) + .await? + .into_row() + .await? + .unwrap() + .get(0) + .unwrap(); + assert_eq!(10, count); + + let bad: i32 = conn + .query( + &format!("SELECT COUNT(*) FROM {} WHERE c <> val + 1", table), + &[], + ) + .await? + .into_row() + .await? + .unwrap() + .get(0) + .unwrap(); + assert_eq!(0, bad); + + Ok(()) +} + +// Server-gated: `KeepIdentity` must let the caller supply explicit identity +// values instead of the server auto-assigning them. It does so by keeping the +// identity column in the bulk column list (there is no `KEEP_IDENTITY` keyword +// in the `INSERT BULK` grammar); without the flag the identity column is +// filtered out and the `id`s would be reassigned. Compiles locally; runs only +// against a live server. +#[test_on_runtimes] +async fn bulk_insert_with_keep_identity_preserves_supplied_ids( + mut conn: tiberius::Client, +) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + use tiberius::{IntoRow, SqlBulkCopyOption}; + + let table = format!("##{}", random_table().await); + + conn.execute( + &format!( + "CREATE TABLE {} (id INT IDENTITY PRIMARY KEY, val INT NOT NULL)", + table + ), + &[], + ) + .await?; + + let mut req = conn + .bulk_insert_with_options( + &table, + &["id", "val"], + SqlBulkCopyOption::KeepIdentity | SqlBulkCopyOption::TableLock, + &[], + ) + .await?; + + for (id, val) in [(100i32, 1i32), (200, 2), (300, 3)] { + req.send((id, val).into_row()).await?; + } + let res = req.finalize().await?; + assert_eq!(3, res.total()); + + // The explicit ids survived because the identity column was kept in the + // bulk column list. + let kept: i32 = conn + .query( + &format!("SELECT COUNT(*) FROM {} WHERE id IN (100, 200, 300)", table), + &[], + ) + .await? + .into_row() + .await? + .unwrap() + .get(0) + .unwrap(); + assert_eq!(3, kept); + + Ok(()) +} + +// Server-gated: an `ORDER (...)` hint plus `TABLOCK` must produce a statement +// the server accepts, and all rows must land. Compiles locally; runs only +// against a live server. +#[test_on_runtimes] +async fn bulk_insert_with_order_hints_inserts_all_rows( + mut conn: tiberius::Client, +) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + use tiberius::{SortOrder, SqlBulkCopyOption}; + + let table = format!("##{}", random_table().await); + + conn.execute( + &format!( + "CREATE TABLE {} (id INT IDENTITY PRIMARY KEY, val INT NOT NULL)", + table + ), + &[], + ) + .await?; + + let mut req = conn + .bulk_insert_with_options( + &table, + &["val"], + SqlBulkCopyOption::TableLock.into(), + &[("val", SortOrder::Ascending)], + ) + .await?; + + for v in 0..10i32 { + let mut row = TokenRow::new(); + row.push(v.into_sql()); + req.send(row).await?; + } + let res = req.finalize().await?; + assert_eq!(10, res.total()); + + let count: i32 = conn + .query(&format!("SELECT COUNT(*) FROM {}", table), &[]) + .await? + .into_row() + .await? + .unwrap() + .get(0) + .unwrap(); + assert_eq!(10, count); + + Ok(()) +} + +// Server-gated: empty options + empty order hints via the new API must behave +// exactly like `bulk_insert_columns` (no `WITH` clause). Compiles locally; runs +// only against a live server. +#[test_on_runtimes] +async fn bulk_insert_with_options_empty_matches_plain_path( + mut conn: tiberius::Client, +) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + use tiberius::SqlBulkCopyOptions; + + let table = format!("##{}", random_table().await); + + conn.execute( + &format!( + "CREATE TABLE {} (id INT IDENTITY PRIMARY KEY, val INT NOT NULL)", + table + ), + &[], + ) + .await?; + + let mut req = conn + .bulk_insert_with_options(&table, &["val"], SqlBulkCopyOptions::empty(), &[]) + .await?; + + for v in 0..5i32 { + let mut row = TokenRow::new(); + row.push(v.into_sql()); + req.send(row).await?; + } + let res = req.finalize().await?; + assert_eq!(5, res.total()); + + Ok(()) +} diff --git a/tests/command.rs b/tests/command.rs new file mode 100644 index 000000000..2bf7101e9 --- /dev/null +++ b/tests/command.rs @@ -0,0 +1,314 @@ +use futures_util::io::{AsyncRead, AsyncWrite}; +use names::{Generator, Name}; +use once_cell::sync::Lazy; +use std::cell::RefCell; +use std::env; +use std::sync::Once; + +use tiberius::{numeric::Numeric, Command, Result, TableValueRow}; + +use runtimes_macro::test_on_runtimes; + +// Used by the test_on_runtimes macro. +#[allow(dead_code)] +static LOGGER_SETUP: Once = Once::new(); + +static CONN_STR: Lazy = Lazy::new(|| { + env::var("TIBERIUS_TEST_CONNECTION_STRING").unwrap_or_else(|_| { + "server=tcp:localhost,1433;user=SA;password=;IntegratedSecurity=true;TrustServerCertificate=true".to_owned() + }) +}); + +thread_local! { + static NAMES: RefCell>> = + const { RefCell::new(None) }; +} + +async fn random_table() -> String { + NAMES.with(|maybe_generator| { + maybe_generator + .borrow_mut() + .get_or_insert_with(|| Generator::with_naming(Name::Plain)) + .next() + .unwrap() + .replace('-', "") + }) +} + +#[test_on_runtimes] +async fn basic_proc_exec(mut conn: tiberius::Client) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + let table = random_table().await; + let proc = random_table().await; + + conn.simple_query(format!( + r#" + create table ##{} ( + id int identity(1,1), + other varchar(50), + ) + "#, + table + )) + .await?; + + conn.simple_query(format!( + r#" + create or alter procedure {} + @Param1 varchar(50) + as + insert into ##{} (other) + values (@Param1) + + return scope_identity() + "#, + proc, table, + )) + .await?; + + let mut ins_cmd = Command::new(&proc); + ins_cmd.bind_param("@Param1", "some text"); + + let result = ins_cmd.exec(&mut conn).await?.into_command_result().await?; + assert_eq!(1, result.return_code()); + + let mut ins_cmd = Command::new(&proc); + ins_cmd.bind_param("@Param1", "another text"); + + let result = ins_cmd.exec(&mut conn).await?.into_command_result().await?; + assert_eq!(2, result.return_code()); + + Ok(()) +} + +struct GeoTest { + id: i32, + lat: Numeric, + lon: Numeric, +} + +impl<'a> TableValueRow<'a> for GeoTest { + fn bind_fields(&self, data_row: &mut tiberius::SqlTableDataRow<'a>) { + data_row.add_field(self.id); + data_row.add_field(self.lat); + data_row.add_field(self.lon); + } + + fn get_db_type() -> &'static str { + "GeoTest" + } +} + +#[test_on_runtimes] +async fn tvp_proc_exec(mut conn: tiberius::Client) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + let table = random_table().await; + let proc_tvp = random_table().await; + let proc_ins = random_table().await; + let proc_get = random_table().await; + let db_type = random_table().await; + + conn.simple_query(format!( + r#" + create type dbo.[{}] as table + ( + [ID] int not null, + [lat] decimal(9,6), + [lon] decimal(9,6) + ) + "#, + db_type + )) + .await?; + + conn.simple_query(format!( + r#" + create table ##{} ( + id int not null, + lat decimal(9,6), + lon decimal(9,6), + n varchar(50) + ) + "#, + table + )) + .await?; + + conn.simple_query(format!( + r#" + create or alter procedure {} + @id int, + @geo dbo.[{}] readonly + as + update t set + [lat] = g.[lat], + [lon] = g.[lon] + from + ##{} t + inner join @geo g on g.id = t.id + + "#, + proc_tvp, db_type, table, + )) + .await?; + + conn.simple_query(format!( + r#" + create or alter procedure {} + @id int, + @name varchar(50) + as + insert into ##{} (id, n) + values (@id, @name) + + "#, + proc_ins, table, + )) + .await?; + + conn.simple_query(format!( + r#" + create or alter procedure {} + @id int, + @count int out + as + set @count = (select count(*) from ##{}) + select * from ##{} + "#, + proc_get, table, table + )) + .await?; + + let mut ins_cmd = Command::new(&proc_ins); + ins_cmd.bind_param("@id", 23); + ins_cmd.bind_param("@name", "the twenty three"); + + let result = ins_cmd.exec(&mut conn).await?.into_command_result().await?; + assert_eq!(0, result.return_code()); + + let g1 = GeoTest { + id: 23, + lon: Numeric::new_with_scale(141, 6), + lat: Numeric::new_with_scale(192, 6), + }; + let g2 = GeoTest { + id: 78, + lon: Numeric::new_with_scale(1141, 6), + lat: Numeric::new_with_scale(8192, 6), + }; + let tbl = vec![g1, g2]; + let mut tvp_cmd = Command::new(&proc_tvp); + tvp_cmd.bind_param("@id", 23); + tvp_cmd.bind_table_with_dbtype("@geo", &db_type, tbl); + + let result = tvp_cmd.exec(&mut conn).await?.into_command_result().await?; + assert_eq!(0, result.return_code()); + + let count = 0; + let mut get_cmd = Command::new(&proc_get); + get_cmd.bind_param("@id", 23); + get_cmd.bind_out_param("@count", count); + + let result = get_cmd.exec(&mut conn).await?.into_command_result().await?; + assert_eq!(0, result.return_code()); + let count: i32 = result.try_return_value("@count")?.unwrap(); + assert_eq!(1, count); + + let rows = result.to_query_result(0).unwrap(); + let lat: Numeric = rows[0].get("lat").unwrap(); + let lon: Numeric = rows[0].get("lon").unwrap(); + assert_eq!(Numeric::new_with_scale(141, 6), lon); + assert_eq!(Numeric::new_with_scale(192, 6), lat); + + Ok(()) +} + +/// Server-free test that exercises `#[derive(TableValueRow)]` end to end: +/// it declares structs with owned `String`/`Vec` and scalar fields (and a +/// borrowed `&str` field via a struct lifetime), constructs values, and asserts +/// the bound column data. This is the regression test for owned, non-`Copy` +/// fields, which previously failed to compile (E0507) in the generated +/// `bind_fields`. +#[test] +fn derive_table_value_row_binds_owned_and_borrowed_fields() { + use tiberius::{ColumnData, TableValue}; + + #[derive(TableValueRow)] + struct OwnedRow { + #[colname = "Id"] + id: i32, + #[colname = "Name"] + name: String, + #[colname = "Payload"] + payload: Vec, + #[colname = "Score"] + score: Numeric, + } + + // get_db_type derives from the struct name. + assert_eq!(OwnedRow::get_db_type(), "OwnedRow"); + + let rows = vec![ + OwnedRow { + id: 1, + name: "one".to_owned(), + payload: vec![0xDE, 0xAD], + score: Numeric::new_with_scale(10, 0), + }, + OwnedRow { + id: 2, + name: "two".to_owned(), + payload: vec![0xBE, 0xEF], + score: Numeric::new_with_scale(20, 0), + }, + ]; + + let data = TableValue::into_sql(rows); + assert_eq!(data.rows().len(), 2); + + let cols = data.rows()[0].columns(); + assert_eq!(cols.len(), 4); + assert!(matches!(cols[0], ColumnData::I32(Some(1)))); + match &cols[1] { + ColumnData::String(Some(s)) => assert_eq!(s.as_ref(), "one"), + other => panic!("expected String, got {other:?}"), + } + match &cols[2] { + ColumnData::Binary(Some(b)) => assert_eq!(b.as_ref(), [0xDE, 0xAD]), + other => panic!("expected Binary, got {other:?}"), + } + assert!(matches!(cols[3], ColumnData::Numeric(Some(_)))); + + // Second row values, to be sure per-row binding is correct. + match &data.rows()[1].columns()[1] { + ColumnData::String(Some(s)) => assert_eq!(s.as_ref(), "two"), + other => panic!("expected String, got {other:?}"), + } + + // Borrowed fields via a struct lifetime must also compile and bind, without + // cloning the underlying data. + #[derive(TableValueRow)] + struct BorrowedRow<'a> { + #[colname = "Id"] + id: i32, + #[colname = "Name"] + name: &'a str, + } + + let owned = String::from("borrowed"); + let brows = vec![BorrowedRow { + id: 7, + name: &owned, + }]; + let bdata = TableValue::into_sql(brows); + let bcols = bdata.rows()[0].columns(); + assert!(matches!(bcols[0], ColumnData::I32(Some(7)))); + match &bcols[1] { + ColumnData::String(Some(v)) => assert_eq!(v.as_ref(), "borrowed"), + other => panic!("expected String, got {other:?}"), + } +} diff --git a/tests/named-instance-async.rs b/tests/named-instance-async.rs deleted file mode 100644 index c3e48c657..000000000 --- a/tests/named-instance-async.rs +++ /dev/null @@ -1,44 +0,0 @@ -#![cfg(all(windows, feature = "sql-browser-async-std"))] - -use async_std::net::TcpStream; -use once_cell::sync::Lazy; -use std::env; -use std::sync::Once; -use tiberius::{Result, SqlBrowser}; - -// This is used in the testing macro :) -#[allow(dead_code)] -static LOGGER_SETUP: Once = Once::new(); - -static CONN_STR: Lazy = Lazy::new(|| { - env::var("TIBERIUS_TEST_CONNECTION_STRING").unwrap_or_else(|_| { - "server=tcp:localhost,1433;IntegratedSecurity=true;TrustServerCertificate=true".to_owned() - }) -}); - -static NAMED_INSTANCE_CONN_STR: Lazy = Lazy::new(|| { - let instance_name = env::var("TIBERIUS_TEST_INSTANCE").unwrap_or("MSSQLSERVER".to_owned()); - CONN_STR.replace(",1433", &format!("\\{}", instance_name)) -}); - -#[test] -fn connect_to_named_instance() -> Result<()> { - LOGGER_SETUP.call_once(|| { - env_logger::init(); - }); - async_std::task::block_on(async { - let config = tiberius::Config::from_ado_string(&NAMED_INSTANCE_CONN_STR)?; - let tcp = TcpStream::connect_named(&config).await?; - let mut client = tiberius::Client::connect(config, tcp).await?; - - let row = client - .query("SELECT @P1", &[&-4i32]) - .await? - .into_row() - .await? - .unwrap(); - - assert_eq!(Some(-4i32), row.get(0)); - Ok(()) - }) -} diff --git a/tests/query.rs b/tests/query.rs index 4cf3c62bd..c573e435d 100644 --- a/tests/query.rs +++ b/tests/query.rs @@ -24,7 +24,7 @@ static CONN_STR: Lazy = Lazy::new(|| { thread_local! { static NAMES: RefCell>> = - RefCell::new(None); + const { RefCell::new(None) }; } async fn random_table() -> String { @@ -40,11 +40,27 @@ async fn random_table() -> String { static DOT_CONN_STR: Lazy = Lazy::new(|| CONN_STR.replace("localhost", ".")); +static APP_NAME_CONN_STR: Lazy = + Lazy::new(|| format!("{};Application Name=meow", *CONN_STR)); + +// `encrypt=true` requires a TLS backend; without one it is a hard error +// (see #305), so this connection string and the test using it are only built +// when a TLS backend is compiled in. +#[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" +))] static ENCRYPTED_CONN_STR: Lazy = Lazy::new(|| format!("{};encrypt=true", *CONN_STR)); static PLAIN_TEXT_CONN_STR: Lazy = Lazy::new(|| format!("{};encrypt=DANGER_PLAINTEXT", *CONN_STR)); +#[cfg(any( + feature = "rustls", + feature = "native-tls", + feature = "vendored-openssl" +))] #[test_on_runtimes(connection_string = "ENCRYPTED_CONN_STR")] async fn connect_with_full_encryption(mut conn: tiberius::Client) -> Result<()> where @@ -397,6 +413,33 @@ where Ok(()) } +#[test_on_runtimes] +async fn read_and_write_to_keyword_columns(mut conn: tiberius::Client) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + let table = format!("##{}", random_table().await); + + conn.simple_query(format!("CREATE TABLE {} ([End] INT)", table)) + .await?; + + let res = conn + .execute(format!("INSERT INTO {} ([End]) VALUES (5)", table), &[]) + .await?; + + assert_eq!(1, res.total()); + + let rows = conn + .query(format!("SELECT [End] FROM {}", table), &[]) + .await? + .into_first_result() + .await?; + + assert_eq!(Some(5), rows[0].get(0)); + + Ok(()) +} + #[test_on_runtimes] async fn execute_insert_update_delete(mut conn: tiberius::Client) -> Result<()> where @@ -1749,7 +1792,7 @@ async fn numeric_type_u64_presentation(mut conn: tiberius::Client) -> Resu where S: AsyncRead + AsyncWrite + Unpin + Send, { - let num = Numeric::new_with_scale(std::i32::MAX as i128 + 10, 1); + let num = Numeric::new_with_scale(i32::MAX as i128 + 10, 1); let row = conn .query("SELECT @P1", &[&num]) @@ -1768,7 +1811,7 @@ async fn numeric_type_u96_presentation(mut conn: tiberius::Client) -> Resu where S: AsyncRead + AsyncWrite + Unpin + Send, { - let num = Numeric::new_with_scale(std::i64::MAX as i128, 19); + let num = Numeric::new_with_scale(i64::MAX as i128, 19); let row = conn .query("SELECT @P1", &[&num]) @@ -1787,7 +1830,7 @@ async fn numeric_type_u128_presentation(mut conn: tiberius::Client) -> Res where S: AsyncRead + AsyncWrite + Unpin + Send, { - let num = Numeric::new_with_scale(std::i64::MAX as i128, 37); + let num = Numeric::new_with_scale(i64::MAX as i128, 37); let row = conn .query("SELECT @P1", &[&num]) @@ -2685,94 +2728,196 @@ where Ok(()) } -#[test] -#[cfg(feature = "sql-browser-async-std")] -fn cyrillic_collations_should_work() -> Result<()> { - LOGGER_SETUP.call_once(|| { - env_logger::init(); - }); - - async_std::task::block_on(async { - let mut admin = { - let config = tiberius::Config::from_ado_string(&CONN_STR)?; +#[test_on_runtimes] +async fn cyrillic_collations_should_work(mut conn: tiberius::Client) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + conn.simple_query( + "CREATE TABLE #cyrillic_test ( + single CHAR(1) COLLATE Cyrillic_General_CI_AS, + multi VARCHAR(255) COLLATE Cyrillic_General_CI_AS, + huge TEXT COLLATE Cyrillic_General_CI_AS + )", + ) + .await?; - let tcp = async_std::net::TcpStream::connect(config.get_addr()).await?; - tcp.set_nodelay(true)?; + conn.execute( + "INSERT INTO #cyrillic_test (single, multi, huge) VALUES (@P1, @P2, @P3)", + &[ + &"Ж", + &"В Советском Союзе попытки борьбы с пьянством предпринимались не единожды. Первая антиалкогольная", + &"Первая антиалкогольная", + ], + ) + .await?; - tiberius::Client::connect(config, tcp).await? - }; + let row = conn + .query("SELECT single, multi, huge FROM #cyrillic_test", &[]) + .await? + .into_row() + .await? + .unwrap(); - admin - .simple_query("CREATE DATABASE ru_test COLLATE Cyrillic_General_CI_AS") - .await?; + assert_eq!(Some("Ж"), row.get(0)); + assert_eq!( + Some("В Советском Союзе попытки борьбы с пьянством предпринимались не единожды. Первая антиалкогольная"), + row.get(1) + ); + assert_eq!(Some("Первая антиалкогольная"), row.get(2)); - { - let mut client = { - let mut config = tiberius::Config::from_ado_string(&CONN_STR)?; - config.database("ru_test"); + Ok(()) +} - let tcp = async_std::net::TcpStream::connect(config.get_addr()).await?; - tcp.set_nodelay(true)?; +#[test_on_runtimes] +async fn legacy_codepages_query_round_trip(mut conn: tiberius::Client) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + for (collation, expected, expected_bytes) in [ + ( + "SQL_Latin1_General_CP437_BIN", + "Café α\u{a0}", + b"Caf\x82 \xe0\xff".as_slice(), + ), + ( + "SQL_1xCompat_CP850_CI_AS", + "Café ø\u{a0}", + b"Caf\x82 \x9b\xff".as_slice(), + ), + ] { + conn.simple_query(format!( + "CREATE TABLE #legacy_codepages ( + single CHAR(1) COLLATE {collation}, + multi VARCHAR(32) COLLATE {collation}, + huge VARCHAR(MAX) COLLATE {collation}, + legacy TEXT COLLATE {collation} + )" + )) + .await? + .into_results() + .await?; - tiberius::Client::connect(config, tcp).await? - }; + let long_value = expected.repeat(2000); + conn.execute( + "INSERT INTO #legacy_codepages VALUES (@P1, @P2, @P3, @P4)", + &[&"é", &expected, &long_value, &expected], + ) + .await?; - client - .simple_query( - "CREATE TABLE test (id INT IDENTITY PRIMARY KEY, single CHAR(1), multi VARCHAR(255), huge TEXT)", - ) - .await?; + let row = conn + .simple_query( + "SELECT single, multi, huge, legacy, CAST(multi AS SQL_VARIANT), + CONVERT(VARBINARY(32), multi) + FROM #legacy_codepages", + ) + .await? + .into_row() + .await? + .unwrap(); - client.execute( - "INSERT INTO test (single, multi, huge) VALUES (@P1, @P2, @P3)", - &[&"Ж", &"В Советском Союзе попытки борьбы с пьянством предпринимались не единожды. Первая антиалкогольная", &"Первая антиалкогольная"] - ).await?; + assert_eq!(row.get::<&str, _>(0), Some("é")); + assert_eq!(row.get::<&str, _>(1), Some(expected)); + assert_eq!(row.get::<&str, _>(2), Some(long_value.as_str())); + assert_eq!(row.get::<&str, _>(3), Some(expected)); + assert_eq!(row.get::<&str, _>(4), Some(expected)); + assert_eq!(row.get::<&[u8], _>(5), Some(expected_bytes)); - let row = client - .query("SELECT single, multi, huge FROM test", &[]) - .await? - .into_row() - .await? - .unwrap(); + conn.simple_query("DROP TABLE #legacy_codepages") + .await? + .into_results() + .await?; + } - assert_eq!(Some("Ж"), row.get(0)); - assert_eq!(Some("В Советском Союзе попытки борьбы с пьянством предпринимались не единожды. Первая антиалкогольная"), row.get(1)); - assert_eq!(Some("Первая антиалкогольная"), row.get(2)); - } + Ok(()) +} - admin.simple_query("DROP DATABASE ru_test").await?; +#[tokio::test] +async fn lossy_codepage_config_reaches_decoder() -> Result<()> { + use tokio_util::compat::TokioAsyncWriteCompatExt; - Ok(()) - }) -} + let mut config = tiberius::Config::from_ado_string(&CONN_STR)?; + config.database("master"); + let tcp = tokio::net::TcpStream::connect(config.get_addr()).await?; + tcp.set_nodelay(true)?; + let mut admin = tiberius::Client::connect(config, tcp.compat_write()).await?; + let database = format!("tiberius_lossy_{}", Uuid::new_v4().simple()); + admin + .simple_query(format!( + "CREATE DATABASE [{database}] COLLATE Chinese_PRC_CI_AS" + )) + .await? + .into_results() + .await?; -#[test] -#[cfg(feature = "sql-browser-async-std")] -fn application_name_should_be_set_correctly() -> Result<()> { - LOGGER_SETUP.call_once(|| { - env_logger::init(); - }); + let outcomes: Result<_> = async { + let mut outcomes = Vec::new(); + for lossy in [None, Some(false), Some(true)] { + let mut config = tiberius::Config::from_ado_string(&CONN_STR)?; + config.database(&database); + if let Some(lossy) = lossy { + config.lossy_codepage_decoding(lossy); + } + let tcp = tokio::net::TcpStream::connect(config.get_addr()).await?; + tcp.set_nodelay(true)?; + let mut client = tiberius::Client::connect(config, tcp.compat_write()).await?; + let result = client + .simple_query( + "SELECT CAST(0x61812062 AS VARCHAR(4)) AS malformed, + CAST(0xFF AS VARCHAR(1)) AS lone_byte; + SELECT CAST(0xD6D0CEC4 AS VARCHAR(4)) AS valid, + CAST('next' AS VARCHAR(4)) AS following", + ) + .await? + .into_results() + .await; + outcomes.push(result); + } + Ok(outcomes) + } + .await; - async_std::task::block_on(async { - let mut config = tiberius::Config::from_ado_string(&CONN_STR)?; - config.application_name("meow"); + admin + .simple_query(format!( + "ALTER DATABASE [{database}] SET SINGLE_USER WITH ROLLBACK IMMEDIATE; + DROP DATABASE [{database}]" + )) + .await? + .into_results() + .await?; - let tcp = async_std::net::TcpStream::connect(config.get_addr()).await?; - tcp.set_nodelay(true)?; + let mut outcomes = outcomes?.into_iter(); + for _ in 0..2 { + assert!(matches!( + outcomes.next().unwrap(), + Err(tiberius::error::Error::Encoding(_)) + )); + } + let results = outcomes.next().unwrap()?; + assert_eq!(results.len(), 2); + assert_eq!(results[0][0].get::<&str, _>(0), Some("a\u{fffd} b")); + assert_eq!(results[0][0].get::<&str, _>(1), Some("\u{fffd}")); + assert_eq!(results[1][0].get::<&str, _>(0), Some("中文")); + assert_eq!(results[1][0].get::<&str, _>(1), Some("next")); - let mut client = tiberius::Client::connect(config, tcp).await?; + Ok(()) +} - let row = client - .query("SELECT APP_NAME()", &[]) - .await? - .into_row() - .await? - .unwrap(); +#[test_on_runtimes(connection_string = "APP_NAME_CONN_STR")] +async fn application_name_should_be_set_correctly(mut conn: tiberius::Client) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + let row = conn + .query("SELECT APP_NAME()", &[]) + .await? + .into_row() + .await? + .unwrap(); - assert_eq!(Some("meow"), row.get(0)); + assert_eq!(Some("meow"), row.get(0)); - Ok(()) - }) + Ok(()) } #[test_on_runtimes] diff --git a/tests/serde.rs b/tests/serde.rs new file mode 100644 index 000000000..dd85d8297 --- /dev/null +++ b/tests/serde.rs @@ -0,0 +1,195 @@ +//! Tests that verify the optional serde Serialize/Deserialize impls +//! gated behind the `serde` feature. +//! +//! These tests do not require a live SQL Server. They round-trip the +//! exposed result-set types (Row pieces, ColumnData, Numeric, time +//! types, etc.) through `serde_json` and assert the values survive. + +#![cfg(feature = "serde")] + +use std::borrow::Cow; +use std::sync::Arc; + +use tiberius::numeric::Numeric; +use tiberius::time::DateTime; +use tiberius::xml::XmlData; +use tiberius::{Column, ColumnData, ColumnType, TokenRow}; +use uuid::Uuid; + +#[cfg(feature = "tds73")] +use tiberius::time::{Date, DateTime2, DateTimeOffset, Time}; + +fn json_round_trip(value: &T) -> T +where + T: serde::Serialize + serde::de::DeserializeOwned, +{ + let s = serde_json::to_string(value).expect("serialize"); + serde_json::from_str(&s).expect("deserialize") +} + +#[test] +fn column_type_round_trip() { + let value = ColumnType::Int4; + let back: ColumnType = json_round_trip(&value); + assert_eq!(value, back); +} + +#[test] +fn column_round_trip() { + let value = Column::new("hello".to_string(), ColumnType::NVarchar); + let back: Column = json_round_trip(&value); + assert_eq!(back.name(), "hello"); + assert_eq!(back.column_type(), ColumnType::NVarchar); +} + +#[test] +fn column_data_int_round_trip() { + let value = ColumnData::I32(Some(42)); + let back: ColumnData<'static> = json_round_trip(&value); + assert_eq!(value, back); +} + +#[test] +fn column_data_null_round_trip() { + let value: ColumnData<'static> = ColumnData::I64(None); + let back: ColumnData<'static> = json_round_trip(&value); + assert_eq!(value, back); +} + +#[test] +fn column_data_string_round_trip() { + let value = ColumnData::String(Some(Cow::Borrowed("héllo"))); + let back: ColumnData<'static> = json_round_trip(&value); + // Borrowed inputs deserialize as Cow::Owned but values match. + match back { + ColumnData::String(Some(s)) => assert_eq!(s.as_ref(), "héllo"), + other => panic!("unexpected: {:?}", other), + } +} + +#[test] +fn column_data_binary_round_trip() { + let bytes: &[u8] = &[0xde, 0xad, 0xbe, 0xef]; + let value = ColumnData::Binary(Some(Cow::Borrowed(bytes))); + let back: ColumnData<'static> = json_round_trip(&value); + match back { + ColumnData::Binary(Some(b)) => assert_eq!(b.as_ref(), bytes), + other => panic!("unexpected: {:?}", other), + } +} + +#[test] +fn column_data_guid_round_trip() { + let id = Uuid::from_u128(0xfeed_face_dead_beef_0000_1111_2222_3333u128); + let value = ColumnData::Guid(Some(id)); + let back: ColumnData<'static> = json_round_trip(&value); + assert_eq!(value, back); +} + +#[test] +fn column_data_bool_round_trip() { + let value = ColumnData::Bit(Some(true)); + let back: ColumnData<'static> = json_round_trip(&value); + assert_eq!(value, back); +} + +#[test] +fn column_data_float_round_trip() { + let value = ColumnData::F64(Some(std::f64::consts::PI)); + let back: ColumnData<'static> = json_round_trip(&value); + assert_eq!(value, back); +} + +#[test] +fn numeric_round_trip() { + let value = Numeric::new_with_scale(57705, 2); + let back: Numeric = json_round_trip(&value); + assert_eq!(value, back); + assert_eq!(back.value(), 57705); + assert_eq!(back.scale(), 2); +} + +#[test] +fn column_data_numeric_round_trip() { + let value = ColumnData::Numeric(Some(Numeric::new_with_scale(12345, 3))); + let back: ColumnData<'static> = json_round_trip(&value); + assert_eq!(value, back); +} + +#[test] +fn datetime_round_trip() { + let value = DateTime::new(200, 3000); + let back: DateTime = json_round_trip(&value); + assert_eq!(value, back); +} + +#[test] +fn column_data_datetime_round_trip() { + let value = ColumnData::DateTime(Some(DateTime::new(42, 84))); + let back: ColumnData<'static> = json_round_trip(&value); + assert_eq!(value, back); +} + +#[cfg(feature = "tds73")] +#[test] +fn time_types_round_trip() { + let date = Date::new(123); + let time = Time::new(7, 7); + let dt2 = DateTime2::new(date, time); + let dto = DateTimeOffset::new(dt2, -120); + + assert_eq!(date, json_round_trip(&date)); + assert_eq!(time, json_round_trip(&time)); + assert_eq!(dt2, json_round_trip(&dt2)); + assert_eq!(dto, json_round_trip(&dto)); +} + +#[test] +fn xml_data_round_trip() { + let value = XmlData::new("hi"); + let back: XmlData = json_round_trip(&value); + assert_eq!(value.as_ref(), back.as_ref()); +} + +#[test] +fn token_row_round_trip() { + let mut row: TokenRow<'static> = TokenRow::new(); + row.push(ColumnData::I32(Some(1))); + row.push(ColumnData::String(Some(Cow::Owned("hello".to_string())))); + row.push(ColumnData::Bit(Some(false))); + + let back: TokenRow<'static> = json_round_trip(&row); + assert_eq!(back.len(), 3); + assert_eq!(back.get(0), Some(&ColumnData::I32(Some(1)))); + match back.get(1).unwrap() { + ColumnData::String(Some(s)) => assert_eq!(s.as_ref(), "hello"), + other => panic!("unexpected: {:?}", other), + } + assert_eq!(back.get(2), Some(&ColumnData::Bit(Some(false)))); +} + +/// Mimic the "send query results across the network as JSON" flow from +/// the issue: build a row out of columns + data, then verify the whole +/// thing round-trips. +#[test] +fn row_shape_round_trip() { + // Build column metadata. + let columns = Arc::new(vec![ + Column::new("id".to_string(), ColumnType::Int4), + Column::new("name".to_string(), ColumnType::NVarchar), + ]); + let mut data: TokenRow<'static> = TokenRow::new(); + data.push(ColumnData::I32(Some(7))); + data.push(ColumnData::String(Some(Cow::Owned("ada".to_string())))); + + // Round-trip the column metadata. + let columns_back: Arc> = json_round_trip(&columns); + assert_eq!(columns_back.len(), 2); + assert_eq!(columns_back[0].name(), "id"); + assert_eq!(columns_back[1].column_type(), ColumnType::NVarchar); + + // Round-trip the row data. + let data_back: TokenRow<'static> = json_round_trip(&data); + assert_eq!(data_back.len(), 2); + assert_eq!(data_back.get(0), Some(&ColumnData::I32(Some(7)))); +} diff --git a/tests/special_char_password.rs b/tests/special_char_password.rs new file mode 100644 index 000000000..d2ec9fa53 --- /dev/null +++ b/tests/special_char_password.rs @@ -0,0 +1,99 @@ +//! End-to-end coverage for special-character passwords (issue #313). +//! +//! The unit tests in `src/client/config` prove the connection-string *parser*; +//! these prove the whole path — parse, LOGIN7, real SQL Server authentication — +//! for passwords that contain structural characters. They run only against a +//! live server (the CI integration lanes), like the other files in `tests/`. + +use once_cell::sync::Lazy; +use std::env; +use tiberius::{Client, Config}; +use tokio::net::TcpStream; +use tokio_util::compat::{Compat, TokioAsyncWriteCompatExt}; + +static ADMIN_CONN_STR: Lazy = Lazy::new(|| { + env::var("TIBERIUS_TEST_CONNECTION_STRING").unwrap_or_else(|_| { + "server=tcp:localhost,1433;user=SA;password=;TrustServerCertificate=true".to_owned() + }) +}); + +async fn connect(conn_str: &str) -> anyhow::Result>> { + let config = Config::from_ado_string(conn_str)?; + let tcp = TcpStream::connect(config.get_addr()).await?; + tcp.set_nodelay(true)?; + Ok(Client::connect(config, tcp.compat_write()).await?) +} + +/// Creates a SQL login + user with `password`, connects as it using +/// `ado_password` (the value exactly as written in the connection string), +/// runs a trivial query, then drops the login. Proves the special-character +/// password survives the whole login handshake, not just parsing. +async fn assert_login_roundtrip( + tag: &str, + password: &str, + ado_password: &str, +) -> anyhow::Result<()> { + // Unique per (tag, process, wall-clock nanos) so neither the three tests in + // this binary nor separate runs against a shared server collide on the login + // name. + let nanos = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .map(|d| d.as_nanos()) + .unwrap_or(0); + let login = format!("tib_pw_{tag}_{}_{nanos}", std::process::id()); + + let mut admin = connect(&ADMIN_CONN_STR).await?; + + // CHECK_POLICY=OFF so the test isn't at the mercy of the host's password + // policy; the passwords used here are still non-trivial. The password is a + // T-SQL string literal, so a `'` in it would need doubling — none is used. + let create = format!( + "IF SUSER_ID('{login}') IS NOT NULL DROP LOGIN [{login}]; \ + CREATE LOGIN [{login}] WITH PASSWORD = '{password}', CHECK_POLICY = OFF;" + ); + admin.simple_query(create).await?.into_results().await?; + + // Connect as the new login with the password written the way a user would. + // Built by appending overrides to the admin string: duplicate keys are + // last-wins, so this swaps the user/password and forces a plaintext login + // so the test runs identically in the TLS and no-TLS CI lanes (and also + // exercises last-wins end-to-end). + let conn_str = format!( + "{};user={login};password={ado_password};encrypt=DANGER_PLAINTEXT", + *ADMIN_CONN_STR + ); + let result: anyhow::Result<()> = async { + let mut client = connect(&conn_str).await?; + let row = client + .query("SELECT @@VERSION", &[]) + .await? + .into_row() + .await?; + anyhow::ensure!(row.is_some(), "login `{login}` should return a row"); + Ok(()) + } + .await; + + // Best-effort cleanup regardless of the connection result. + let _ = admin.simple_query(format!("DROP LOGIN [{login}]")).await; + + result +} + +#[tokio::test] +async fn base64_equals_password_authenticates() -> anyhow::Result<()> { + // The #313 case: `=` padding, written unquoted in the connection string. + assert_login_roundtrip("eqpad", "Ab1!Zm9vYmFy==", "Ab1!Zm9vYmFy==").await +} + +#[tokio::test] +async fn semicolon_password_authenticates_when_quoted() -> anyhow::Result<()> { + // A `;` in the password must be quoted in the connection string. + assert_login_roundtrip("semi", "Ab1!x;y=z", "'Ab1!x;y=z'").await +} + +#[tokio::test] +async fn braced_password_authenticates() -> anyhow::Result<()> { + // Brace quoting (the tiberius extension) around `;` and `=`. + assert_login_roundtrip("brace", "Ab1!p;q=r", "{Ab1!p;q=r}").await +} diff --git a/tests/transactions.rs b/tests/transactions.rs new file mode 100644 index 000000000..becd08201 --- /dev/null +++ b/tests/transactions.rs @@ -0,0 +1,166 @@ +//! Integration tests for the Transaction Manager requests (begin / commit / +//! rollback, and explicit isolation levels). These exercise the client-side +//! transaction API against a live SQL Server. + +use futures_util::io::{AsyncRead, AsyncWrite}; +use names::{Generator, Name}; +use once_cell::sync::Lazy; +use std::cell::RefCell; +use std::env; +use std::sync::Once; + +use runtimes_macro::test_on_runtimes; +use tiberius::{IsolationLevel, Result}; + +// This is used in the testing macro :) +#[allow(dead_code)] +static LOGGER_SETUP: Once = Once::new(); + +static CONN_STR: Lazy = Lazy::new(|| { + env::var("TIBERIUS_TEST_CONNECTION_STRING").unwrap_or_else(|_| { + "server=tcp:localhost,1433;user=SA;password=;IntegratedSecurity=true;TrustServerCertificate=true".to_owned() + }) +}); + +thread_local! { + static NAMES: RefCell>> = + const { RefCell::new(None) }; +} + +async fn random_table() -> String { + NAMES.with(|maybe_generator| { + maybe_generator + .borrow_mut() + .get_or_insert_with(|| Generator::with_naming(Name::Plain)) + .next() + .unwrap() + .replace('-', "") + }) +} + +async fn row_count(conn: &mut tiberius::Client, table: &str) -> Result +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + let row = conn + .query(format!("SELECT COUNT(*) FROM ##{}", table), &[]) + .await? + .into_row() + .await? + .expect("COUNT(*) always returns a row"); + + Ok(row.get::(0).expect("COUNT(*) is never NULL")) +} + +#[test_on_runtimes] +async fn transaction_commit_persists_rows(mut conn: tiberius::Client) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + let table = random_table().await; + // Create the table before BEGIN so it isn't part of the transaction under test. + conn.execute(format!("CREATE TABLE ##{} (id int)", table), &[]) + .await?; + + conn.begin_transaction().await?; + conn.execute( + format!("INSERT INTO ##{} (id) VALUES (@P1), (@P2)", table), + &[&1i32, &2i32], + ) + .await?; + conn.commit_transaction().await?; + + assert_eq!(2, row_count(&mut conn, &table).await?); + + Ok(()) +} + +#[test_on_runtimes] +async fn transaction_rollback_discards_rows(mut conn: tiberius::Client) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + let table = random_table().await; + conn.execute(format!("CREATE TABLE ##{} (id int)", table), &[]) + .await?; + conn.execute( + format!("INSERT INTO ##{} (id) VALUES (@P1)", table), + &[&1i32], + ) + .await?; + + conn.begin_transaction().await?; + conn.execute( + format!("INSERT INTO ##{} (id) VALUES (@P1), (@P2)", table), + &[&2i32, &3i32], + ) + .await?; + conn.rollback_transaction().await?; + + assert_eq!(1, row_count(&mut conn, &table).await?); + + Ok(()) +} + +#[test_on_runtimes] +async fn transaction_with_explicit_isolation_levels(mut conn: tiberius::Client) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + let table = random_table().await; + conn.execute(format!("CREATE TABLE ##{} (id int)", table), &[]) + .await?; + + // SNAPSHOT is intentionally omitted: it requires ALLOW_SNAPSHOT_ISOLATION to + // be enabled on the database, which the default test database is not. + for level in [ + IsolationLevel::ReadUncommitted, + IsolationLevel::ReadCommitted, + IsolationLevel::RepeatableRead, + IsolationLevel::Serializable, + ] { + conn.begin_transaction_with_isolation(level).await?; + conn.execute( + format!("INSERT INTO ##{} (id) VALUES (@P1)", table), + &[&1i32], + ) + .await?; + conn.commit_transaction().await?; + } + + assert_eq!(4, row_count(&mut conn, &table).await?); + + Ok(()) +} + +#[test_on_runtimes] +async fn rolled_back_transaction_can_be_followed_by_a_new_one( + mut conn: tiberius::Client, +) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send, +{ + let table = random_table().await; + conn.execute(format!("CREATE TABLE ##{} (id int)", table), &[]) + .await?; + + conn.begin_transaction().await?; + conn.execute( + format!("INSERT INTO ##{} (id) VALUES (@P1)", table), + &[&10i32], + ) + .await?; + conn.rollback_transaction().await?; + + conn.begin_transaction().await?; + conn.execute( + format!("INSERT INTO ##{} (id) VALUES (@P1)", table), + &[&20i32], + ) + .await?; + conn.commit_transaction().await?; + + assert_eq!(1, row_count(&mut conn, &table).await?); + + Ok(()) +} diff --git a/tiberius-macros/Cargo.toml b/tiberius-macros/Cargo.toml new file mode 100644 index 000000000..0a96b965f --- /dev/null +++ b/tiberius-macros/Cargo.toml @@ -0,0 +1,15 @@ +[package] +name = "tiberius-macros" +version = "0.1.0" +edition = "2021" +license = "MIT OR Apache-2.0" +description = "Derive macros for tiberius (e.g. TableValueRow for table-valued parameters)." +repository = "https://github.com/tiberius-rs/tiberius" + +[dependencies] +syn = "^2" +quote = "^1" +proc-macro2 = "^1" + +[lib] +proc-macro = true diff --git a/tiberius-macros/src/attr.rs b/tiberius-macros/src/attr.rs new file mode 100644 index 000000000..49d94abeb --- /dev/null +++ b/tiberius-macros/src/attr.rs @@ -0,0 +1,57 @@ +pub(crate) struct FieldAttr { + pub colname: Option, +} + +impl FieldAttr { + pub(crate) fn parse(attrs: &[syn::Attribute]) -> syn::Result> { + let mut result = None; + for attr in attrs.iter() { + match attr.style { + syn::AttrStyle::Outer => {} + _ => continue, + } + // A parsed attribute path always has at least one segment, so this + // is effectively infallible; keep an explicit guard just in case. + let Some(last_attr_path) = attr.path().segments.last() else { + continue; + }; + if last_attr_path.ident != "colname" { + continue; + } + let kv = match attr.meta { + syn::Meta::NameValue(ref kv) => kv, + _ if attr.path().is_ident("colname") => { + return Err(syn::Error::new_spanned( + attr, + "invalid `#[colname]` attribute on a `#[derive(TableValueRow)]` field: \ + expected the form `#[colname = \"SomeColName\"]`", + )); + } + _ => continue, + }; + if result.is_some() { + return Err(syn::Error::new_spanned( + attr, + "duplicate `#[colname]` attribute on a `#[derive(TableValueRow)]` field: \ + at most one `#[colname = \"...\"]` is allowed per field", + )); + } + if let syn::Expr::Lit(syn::ExprLit { + lit: syn::Lit::Str(ref s), + .. + }) = kv.value + { + result = Some(FieldAttr { + colname: Some(s.value()), + }); + } else { + return Err(syn::Error::new_spanned( + &kv.value, + "invalid `#[colname]` value on a `#[derive(TableValueRow)]` field: \ + expected a string literal, as in `#[colname = \"SomeColName\"]`", + )); + } + } + Ok(result) + } +} diff --git a/tiberius-macros/src/lib.rs b/tiberius-macros/src/lib.rs new file mode 100644 index 000000000..493c8eeb5 --- /dev/null +++ b/tiberius-macros/src/lib.rs @@ -0,0 +1,83 @@ +//! A utility proc-macro crate that generates trivial trait implementations used +//! in the Rust-to-SQL data exchange for tiberius table-valued parameters. +extern crate proc_macro; + +#[macro_use] +extern crate quote; +#[macro_use] +extern crate syn; + +use proc_macro::TokenStream; + +macro_rules! sp_quote { + ($($t:tt)*) => (quote_spanned!(proc_macro2::Span::call_site() => $($t)*)) +} + +mod attr; +mod table_value_param; + +/// Generates a trivial implementation of the `TableValueRow` trait. +/// +/// # Applications +/// Apply to structs that represent rows of a table-valued parameter. +/// +/// # Example +/// ```rust,ignore +/// # use tiberius::*; +/// #[derive(TableValueRow)] +/// pub struct SomeGeoList { +/// #[colname = "SomeID"] +/// pub id: i32, +/// #[colname = "Label"] +/// pub label: String, +/// #[colname = "LastSyncIPGeoLat"] +/// pub lat: Numeric, +/// #[colname = "LastSyncIPGeoLong"] +/// pub lon: Numeric, +/// } +/// ``` +/// +/// # Supported field types +/// +/// Each field is bound with `SqlTableDataRow::add_field(self..clone())`. +/// Because `bind_fields` receives `&self`, the field value cannot be moved out +/// and is instead cloned into the row, so every field type must implement both +/// `Clone` and `IntoSql`. That covers: +/// +/// - the `Copy` scalar/marker types (`i32`, `i64`, `bool`, `f64`, `u8`, `i16`, +/// `f32`, `Numeric`, `Uuid`), where `clone` is a plain copy; +/// - owned, non-`Copy` columns (`String`, `Vec`, `XmlData`), which are +/// cloned by value; +/// - borrowed strings/bytes (`&str`, `&[u8]`) via a struct lifetime, where +/// `clone` just copies the reference (no allocation); +/// - `Option` of any of the above. +/// +/// # Limitations +/// +/// The derive supports structs with named fields and at most one lifetime +/// parameter. Generic type parameters (`struct Row { .. }`), tuple/unit +/// structs, and enums/unions are not supported and produce a compile error. +#[proc_macro_derive(TableValueRow, attributes(colname))] +pub fn table_value_param(input: TokenStream) -> TokenStream { + let ast: syn::DeriveInput = match syn::parse(input) { + Ok(ast) => ast, + Err(e) => return e.to_compile_error().into(), + }; + + let result = match ast.data { + syn::Data::Struct(ref s) => table_value_param::for_struct(&ast, &s.fields), + syn::Data::Enum(_) => Err(syn::Error::new_spanned( + &ast.ident, + "TableValueRow can only be derived for structs, not enums", + )), + syn::Data::Union(_) => Err(syn::Error::new_spanned( + &ast.ident, + "TableValueRow can only be derived for structs, not unions", + )), + }; + + match result { + Ok(tokens) => tokens.into(), + Err(e) => e.to_compile_error().into(), + } +} diff --git a/tiberius-macros/src/table_value_param.rs b/tiberius-macros/src/table_value_param.rs new file mode 100644 index 000000000..9f3d36403 --- /dev/null +++ b/tiberius-macros/src/table_value_param.rs @@ -0,0 +1,188 @@ +use proc_macro2::TokenStream; +use syn::punctuated::Punctuated; + +use crate::attr::FieldAttr; + +pub(crate) fn for_struct(ast: &syn::DeriveInput, fields: &syn::Fields) -> syn::Result { + match *fields { + syn::Fields::Named(ref fields) => table_value_param_impl(ast, Some(&fields.named)), + _ => Err(syn::Error::new_spanned( + &ast.ident, + "TableValueRow can only be derived for structs with named fields \ + (tuple and unit structs are not supported)", + )), + } +} + +fn table_value_param_impl( + ast: &syn::DeriveInput, + fields: Option<&Punctuated>, +) -> syn::Result { + let name = &ast.ident; + let (lt_impl, lt_struct) = { + let mut lifetimes: Vec<&syn::Ident> = Vec::new(); + for gp in ast.generics.params.iter() { + if let syn::GenericParam::Lifetime(ltp) = gp { + lifetimes.push(<p.lifetime.ident); + } + } + if lifetimes.len() > 1 { + return Err(syn::Error::new_spanned( + &ast.generics, + format!( + "TableValueRow supports at most one lifetime parameter, found: {}", + lifetimes + .iter() + .map(|lt| lt.to_string()) + .collect::>() + .join(", ") + ), + )); + } + if lifetimes.is_empty() { + (sp_quote!(<'query>), sp_quote!()) + } else { + let lt = lifetimes[0]; + let ts: proc_macro2::TokenStream = format!("< '{} >", lt).parse().unwrap(); + (ts.clone(), ts) + } + }; + let empty = Default::default(); + let fields: Vec<_> = fields + .unwrap_or(&empty) + .iter() + .map(FieldExt::new) + .collect::>()?; + let col_names: Vec<_> = fields.iter().map(|f| f.get_col_name()).collect(); + // Column names are not needed for stored procedures, but will be required + // once TVPs are supported for ad-hoc queries. + let _col_names = sp_quote!( #(#col_names),* ); + let col_binds: Vec<_> = fields.iter().map(|f| f.as_bind()).collect(); + let col_binds = sp_quote!( #(#col_binds);*); + Ok(sp_quote! { + impl #lt_impl tiberius::TableValueRow #lt_impl for #name #lt_struct { + fn get_db_type() -> &'static str { + stringify!{ #name } + } + + fn bind_fields(&self, data_row: &mut tiberius::SqlTableDataRow #lt_impl) { + #col_binds; + } + } + }) +} + +struct FieldExt { + attr: Option, + ident: syn::Ident, +} + +impl FieldExt { + pub fn new(field: &syn::Field) -> syn::Result { + match field.ident.clone() { + Some(ident) => Ok(FieldExt { + attr: FieldAttr::parse(&field.attrs)?, + ident, + }), + None => Err(syn::Error::new_spanned( + field, + "TableValueRow fields must be named", + )), + } + } + pub(crate) fn get_col_name(&self) -> String { + if let Some(attr) = self.attr.as_ref() { + if let Some(colname) = attr.colname.as_ref() { + return colname.to_string(); + } + } + self.ident.to_string() + } + pub(crate) fn as_bind(&self) -> TokenStream { + let name = &self.ident; + // `bind_fields` receives `&self`, so we cannot move the field out (that + // would fail with E0507 for non-`Copy` types such as `String`/`Vec`). + // Borrowing (`&self.#name`) does not work either: the anonymous `&self` + // borrow is not guaranteed to outlive the `SqlTableDataRow<'a>` lifetime, + // so `&'_ String: IntoSql<'a>` fails to unify. Cloning produces an owned + // (or, for reference fields like `&str`, a copied reference) value that + // satisfies `IntoSql<'a>` for every documented field type. `.clone()` is + // a no-op copy for `Copy` scalars and reference fields. + sp_quote!(data_row.add_field(self.#name.clone())) + } +} + +#[cfg(test)] +mod tests { + use super::for_struct; + + #[test] + fn basic_nolifetime() { + let ast: syn::DeriveInput = syn::parse_str( + r#" + pub struct SomeGeoList { + #[colname = "SomeID"] + pub id: i32, + #[colname = "LastSyncIPGeoLat"] + pub lat: Numeric, // decimal(9,6) + #[colname = "LastSyncIPGeoLong"] + pub lon: Numeric, // decimal(9,6) + } + "#, + ) + .unwrap(); + let result = match ast.data { + syn::Data::Enum(_) => panic!("n/a for enums, makes sense for structs only"), + syn::Data::Struct(ref s) => for_struct(&ast, &s.fields).unwrap(), + syn::Data::Union(_) => panic!("doesn't work with unions"), + }; + let etalon = sp_quote!( + impl<'query> tiberius::TableValueRow<'query> for SomeGeoList { + fn get_db_type() -> &'static str { + stringify! { SomeGeoList } + } + fn bind_fields(&self, data_row: &mut tiberius::SqlTableDataRow<'query>) { + data_row.add_field(self.id.clone()); + data_row.add_field(self.lat.clone()); + data_row.add_field(self.lon.clone()); + } + } + ); + + assert_eq!(result.to_string(), etalon.to_string()); + } + + #[test] + fn basic_lifetime() { + let ast: syn::DeriveInput = syn::parse_str( + r#" + pub struct AnotherGeoList<'e> { + #[colname = "SomeID"] + pub id: i32, + #[colname = "SomeStr"] + pub s: &'e str, + } + "#, + ) + .unwrap(); + let result = match ast.data { + syn::Data::Enum(_) => panic!("n/a for enums, makes sense for structs only"), + syn::Data::Struct(ref s) => for_struct(&ast, &s.fields).unwrap(), + syn::Data::Union(_) => panic!("doesn't work with unions"), + }; + + let etalon = sp_quote!( + impl<'e> tiberius::TableValueRow<'e> for AnotherGeoList<'e> { + fn get_db_type() -> &'static str { + stringify! { AnotherGeoList } + } + fn bind_fields(&self, data_row: &mut tiberius::SqlTableDataRow<'e>) { + data_row.add_field(self.id.clone()); + data_row.add_field(self.s.clone()); + } + } + ); + + assert_eq!(result.to_string(), etalon.to_string()); + } +}