diff --git a/.editorconfig b/.editorconfig new file mode 100644 index 0000000..54baea0 --- /dev/null +++ b/.editorconfig @@ -0,0 +1,22 @@ +root = true + +[*] +charset = utf-8 +end_of_line = lf +indent_style = space +indent_size = 4 +insert_final_newline = true +trim_trailing_whitespace = true + +[*.{yml,yaml,toml,json}] +indent_size = 2 + +[*.md] +indent_size = 2 +trim_trailing_whitespace = false + +# Shared test fixtures are compared byte for byte. +[tests/data/**] +insert_final_newline = unset +trim_trailing_whitespace = unset +indent_size = unset diff --git a/.gitattributes b/.gitattributes new file mode 100644 index 0000000..f858f1b --- /dev/null +++ b/.gitattributes @@ -0,0 +1 @@ +tests/data/** -text diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 9b41bad..9976f0f 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -5,16 +5,68 @@ on: branches: [main] pull_request: +permissions: + contents: read + +concurrency: + group: ci-${{ github.ref }} + cancel-in-progress: true + jobs: test: + name: Python ${{ matrix.python-version }} runs-on: ubuntu-latest strategy: + fail-fast: false matrix: - python-version: ["3.8", "3.9", "3.10", "3.11", "3.12"] + python-version: ["3.9", "3.10", "3.11", "3.12", "3.13"] steps: - - uses: actions/checkout@v4 - - uses: actions/setup-python@v5 + - uses: actions/checkout@11d5960a326750d5838078e36cf38b85af677262 # v4.4.0 + with: + persist-credentials: false + - uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 with: python-version: ${{ matrix.python-version }} - - run: pip install -e ".[dev]" - - run: pytest -q + - name: Install + run: | + python -m pip install --upgrade pip + python -m pip install -e ".[dev]" + - name: Lint + run: | + ruff check . + ruff format --check . + - name: Type check + run: mypy --strict src + - name: Test + run: pytest -q --cov=shieldlabs --cov-report=term-missing + + example: + name: FastAPI example + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@11d5960a326750d5838078e36cf38b85af677262 # v4.4.0 + with: + persist-credentials: false + - uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 + with: + python-version: "3.12" + - name: Install + run: python -m pip install -e ".[dev]" -r examples/requirements.txt + - name: Smoke test + run: pytest -q tests/test_example_app.py + + build: + name: Build distribution + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@11d5960a326750d5838078e36cf38b85af677262 # v4.4.0 + with: + persist-credentials: false + - uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 + with: + python-version: "3.12" + - name: Build + run: | + python -m pip install build==1.6.1 twine==7.0.0 + python -m build + twine check --strict dist/* diff --git a/.github/workflows/publish.yml b/.github/workflows/publish.yml deleted file mode 100644 index b06d1c3..0000000 --- a/.github/workflows/publish.yml +++ /dev/null @@ -1,25 +0,0 @@ -name: Publish - -on: - push: - tags: - - "v*" - -permissions: - contents: read - id-token: write # PyPI trusted publishing - -jobs: - publish: - runs-on: ubuntu-latest - steps: - - uses: actions/checkout@v4 - - uses: actions/setup-python@v5 - with: - python-version: "3.12" - - run: pip install -e ".[dev]" build - - run: pytest -q - - run: python -m build - # Configure trusted publisher on PyPI for project "shieldlabs" - # (owner ShieldLabs-ai, repo shieldlabs-python, workflow publish.yml). - - uses: pypa/gh-action-pypi-publish@release/v1 diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml new file mode 100644 index 0000000..5e9ba13 --- /dev/null +++ b/.github/workflows/release.yml @@ -0,0 +1,70 @@ +name: Release + +# Publishes the package to PyPI when a version tag (v1.0.0, v1.0.1, ...) is pushed. +# The build job has read-only access. The publish job only receives the built files and uses +# PyPI trusted publishing (OpenID Connect), so no API token is stored in the repository. +# One-time setup: on PyPI, add this repository, the workflow file release.yml and the +# environment "pypi" as a trusted publisher of the "shieldlabs" project. +# Re-running the workflow for the same tag is safe: files already on PyPI are skipped. + +on: + push: + tags: ["v*"] + +permissions: + contents: read + +jobs: + build: + name: Check and build + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@11d5960a326750d5838078e36cf38b85af677262 # v4.4.0 + with: + persist-credentials: false + - uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 + with: + python-version: "3.12" + - name: Install + run: python -m pip install -e ".[dev]" build==1.6.1 twine==7.0.0 + - name: Check that the tag matches the package version + run: | + version="$(python -c 'import shieldlabs; print(shieldlabs.__version__)')" + if [ "v${version}" != "${GITHUB_REF_NAME}" ]; then + echo "::error::Tag ${GITHUB_REF_NAME} does not match the package version ${version}." + exit 1 + fi + - name: Verify + run: | + ruff check . + ruff format --check . + mypy --strict src + pytest -q + - name: Build + run: | + python -m build + twine check --strict dist/* + - uses: actions/upload-artifact@ea165f8d65b6e75b540449e92b4886f43607fa02 # v4.6.2 + with: + name: dist + path: dist/ + if-no-files-found: error + + publish: + name: Publish to PyPI + needs: build + runs-on: ubuntu-latest + environment: + name: pypi + url: https://pypi.org/project/shieldlabs/ + permissions: + id-token: write # trusted publishing and attestations + steps: + - uses: actions/download-artifact@d3f86a106a0bac45b974a628896c90dbdf5c8093 # v4.3.0 + with: + name: dist + path: dist/ + - uses: pypa/gh-action-pypi-publish@dc37677b2e1c63e2034f94d8a5b11f265b73ba33 # v1.14.2 + with: + skip-existing: true + attestations: true diff --git a/.gitignore b/.gitignore index 878d7a5..dc5d337 100644 --- a/.gitignore +++ b/.gitignore @@ -1,16 +1,30 @@ +# Python __pycache__/ *.py[cod] -dist/ -build/ *.egg-info/ +.eggs/ +build/ +dist/ + +# Virtual environments .venv/ venv/ -.env -.env.* -.DS_Store + +# Tooling caches and reports .pytest_cache/ .mypy_cache/ .ruff_cache/ +.coverage +.coverage.* +coverage.xml +htmlcov/ + +# Local configuration +.env +.env.* -src/shieldlabs.egg-info -.history +# Editors and OS +.DS_Store +.history/ +.idea/ +.vscode/ diff --git a/CHANGELOG.md b/CHANGELOG.md index ca512e6..9c6e331 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,65 @@ # Changelog -## 2026-09-06 +All notable changes to this project are documented in this file. The format follows +[Keep a Changelog](https://keepachangelog.com/en/1.1.0/) and the project uses +[Semantic Versioning](https://semver.org/spec/v2.0.0.html). -- Minor improvements and bug fixes +## [Unreleased] + +## [1.0.0] - 2026-09-30 + +First stable release. It replaces the 0.1.0 preview package completely. + +### Added + +- `ShieldLabs` and `AsyncShieldLabs` History API clients: `history.search`, `history.iter` + (de-duplicated on `request_id`) and `identifications.get`, which waits for the verdict: + - `timeout` (10 s by default) is the total budget of the call. Polls run immediately, then + after waits of `poll_interval` times 1, 2, 4, 6 and 8, then 8 again, each capped at 2 s, or + at `poll_interval` when that is longer (0.25 s, 0.5 s, 1 s, 1.5 s and every 2 s by default; + every 3 s for `poll_interval=3`), and a last time at the deadline. + - Each poll is one HTTP attempt with a timeout of the client timeout cut to the time left, + but at least 1 s. + - A `429`, a 5xx response, a connection error or a timeout keeps it polling. The error of the + last poll is raised at the deadline; `None` means the last poll found no row. + - After a `429` the next wait is the longest of the ladder step, 1 s and `Retry-After` + capped at 10 s (`Retry-After: 0` or a past date counts as 0), cut to the deadline. A capped + `Retry-After` longer than the time left is raised at once. + - `400`, `401`, `403` and `404` end the wait at once. +- `ShieldLabsManagement` and `AsyncShieldLabsManagement` with `get_profile()`. The domain is + normalized before it is sent, and a `429` is never retried. +- `webhooks.verify_signature` and `webhooks.construct_event`. Both accept one signing secret or + a list of secrets for rotation. Events are typed: `IdentificationScoredEvent`, + `WebhookPingEvent` and `UnknownWebhookEvent`. +- One `Identification` model for webhook data and History rows: 19 detection flags, risk + signals (with descriptions on History rows), timezone-aware `observed_at` and the original + payload in `raw`. Also `DomainProfile`, `HistoryPage` and `SignalName`. +- Helpers: `risk_band`, `is_rate_limited`, `evaluate_identification` and `user_hid`. +- Error hierarchy including `QuotaExceededError`; retries with jittered exponential backoff and + `Retry-After` as sent, up to 10 s (at least 1 s after a `429` without it); a + `User-Agent: shieldlabs-python/` header. +- Client-side validation of every History lookup. User HIDs are percent-encoded in the form the + History API matches, and values that cannot be matched in the request path (`.`, `..` and + anything that contains `/`) raise `ValidationError` instead of returning an empty page. +- Base URLs must use https; plain http is accepted only for `localhost`, `127.0.0.1` and `::1`. +- A FastAPI example, the shared test fixtures and CI on Python 3.9 to 3.13. + +### Changed + +- History requests go to `https://account.shieldlabs.ai/api/v1/history/...`. A custom base URL + that ends in `/api` is accepted and the suffix is removed, so requests never reach + `/api/api/...` (the preview built that URL and received a 404). +- httpx is the only runtime dependency. Python 3.9 is the minimum version. + +### Removed + +- `verify_webhook` and `ShieldLabsClient` from the preview. Use `webhooks.verify_signature` and + `ShieldLabs().history.search` instead. + +## [0.1.0] - 2026-09-06 + +Preview package with `verify_webhook` and a minimal History API client. + +[Unreleased]: https://github.com/ShieldLabs-ai/shieldlabs-python/compare/v1.0.0...HEAD +[1.0.0]: https://github.com/ShieldLabs-ai/shieldlabs-python/releases/tag/v1.0.0 +[0.1.0]: https://github.com/ShieldLabs-ai/shieldlabs-python/commits/main diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md new file mode 100644 index 0000000..c51a850 --- /dev/null +++ b/CONTRIBUTING.md @@ -0,0 +1,64 @@ +# Contributing + +Thank you for helping improve the ShieldLabs Python SDK. + +## Set up + +Check out the repository, then: + +```bash +python3 -m venv .venv +source .venv/bin/activate +pip install -e ".[dev]" +``` + +## Before you open a pull request + +Run the same checks as CI: + +```bash +ruff check . && ruff format --check . && mypy --strict src +pytest -q --cov=shieldlabs --cov-report=term-missing +``` + +- Every change comes with tests. Line coverage stays at 90% or above (CI enforces it). +- Code must run on Python 3.9: use `typing.Optional` and `typing.Union` in annotations. +- Keep httpx the only runtime dependency. +- Changes to the public API need an entry under `Unreleased` in `CHANGELOG.md`. +- The development tools have upper version bounds in `pyproject.toml`, so a new tool release + cannot change CI results on its own. Raise a bound in its own pull request. +- The FastAPI example has its own smoke test: + `pip install -r examples/requirements.txt && pytest tests/test_example_app.py`. + +## Shared test fixtures + +`tests/data/` holds the test fixtures that every ShieldLabs server SDK passes: History API +bodies, webhook bodies, signature vectors, error responses and the expected normalized results. +They are identical in every SDK and compared byte for byte (the `.raw.txt` files are exact +webhook bodies without a trailing newline), so do not edit them in a pull request. If a fixture +looks wrong, open an issue. + +## Writing style + +Docs, docstrings and comments use plain technical English and the terms used in the README. + +## Commits + +Use conventional commit messages: `feat:`, `fix:`, `docs:`, `test:`, `ci:`, `chore:`. + +## Releasing + +Maintainers bump `src/shieldlabs/_version.py`, move the `Unreleased` notes under a new version +heading in `CHANGELOG.md`, and push a `vX.Y.Z` tag. The release workflow checks that the tag +matches the version, runs the checks, builds the package once and publishes that build to PyPI +with trusted publishing and attestations, from a separate job that runs in the `pypi` +environment. + +One-time setup: on PyPI, add this repository, the workflow file `release.yml` and the +environment `pypi` as a trusted publisher of the `shieldlabs` project; on GitHub, protect the +`pypi` environment with required reviewers. No API token is stored in the repository. +Re-running the workflow for a tag is safe: files that are already on PyPI are skipped. + +## Security + +Please report vulnerabilities privately to contact@shieldlabs.ai instead of opening an issue. diff --git a/LICENSE b/LICENSE index 39bb7e4..597e96b 100644 --- a/LICENSE +++ b/LICENSE @@ -1,6 +1,6 @@ MIT License -Copyright (c) 2026 ShieldLabs +Copyright (c) 2026 ShieldLabs Inc. Permission is hereby granted, free of charge, to any person obtaining a copy of this software and associated documentation files (the "Software"), to deal diff --git a/PUBLISHING.md b/PUBLISHING.md deleted file mode 100644 index 7f4b863..0000000 --- a/PUBLISHING.md +++ /dev/null @@ -1,21 +0,0 @@ -# Publishing - -Release packages by pushing a semver tag (`v0.1.0`, `v1.0.0`, …). Workflows live in `.github/workflows/publish.yml`. - -## Prerequisites (founder / org admin) - -| Registry | Package | Setup | -|---|---|---| -| npm | `@shieldlabs/node`, `@shieldlabs/js`, … | Create npm org `@shieldlabs`. Add GitHub Actions secret `NPM_TOKEN` **or** configure [trusted publishing / OIDC](https://docs.npmjs.com/trusted-publishers) for each package. Provenance is enabled (`id-token: write`). | -| PyPI | `shieldlabs` | Claim project name. Add a [trusted publisher](https://docs.pypi.org/trusted-publishers/) for `ShieldLabs-ai/shieldlabs-python` → workflow `publish.yml`. | -| Packagist | `shieldlabs/shieldlabs` | Submit https://github.com/ShieldLabs-ai/shieldlabs-php. Optional secrets `PACKAGIST_USERNAME` + `PACKAGIST_TOKEN` for update hooks. | -| Go | `github.com/ShieldLabs-ai/shieldlabs-go` | Consumers `go get` by git tag; publish workflow creates a GitHub Release. | - -## Publish - -```bash -git tag v0.1.0 -git push origin v0.1.0 -``` - -Do not publish until namespaces are reserved and secrets/OIDC are configured — the workflows will fail otherwise. diff --git a/README.md b/README.md index dc156e0..470b8c8 100644 --- a/README.md +++ b/README.md @@ -1,6 +1,33 @@ -# shieldlabs (Python) +# ShieldLabs Python SDK -ShieldLabs server SDK: webhook verification and History API client. +Read identification verdicts, verify webhooks and turn risk scores into decisions on your Python +backend. + +[![CI](https://github.com/ShieldLabs-ai/shieldlabs-python/actions/workflows/ci.yml/badge.svg)](https://github.com/ShieldLabs-ai/shieldlabs-python/actions/workflows/ci.yml) +[![License: MIT](https://img.shields.io/badge/license-MIT-blue.svg)](LICENSE) +[![PyPI](https://img.shields.io/pypi/v/shieldlabs.svg)](https://pypi.org/project/shieldlabs/) + +ShieldLabs identifies visitors and scores risk with device intelligence. New to ShieldLabs? +[Start free](https://app.shieldlabs.ai) and read the docs at +[docs.shieldlabs.ai](https://docs.shieldlabs.ai). + +## How it fits + +``` + 1. Browser 2. Your backend 3. Decision + ShieldLabs agent ---> receives requestId with the signup, ---> allow, step up, + returns requestId login or checkout; reads the verdict review or refuse + with this SDK (History API) or gets it + as a signed identification.scored webhook +``` + +1. **Browser.** The ShieldLabs agent runs an identification and hands your page a request ID. + The browser never sees a Risk Score, a visitor ID or a device ID. +2. **Your backend.** It receives the request ID together with the protected action and reads the + verdict for it from the History API with this SDK, or receives the verdict by a signed + webhook. +3. **Decision.** Your backend acts on `risk_score`, the three risk bands, `detection_flags` and + the identifiers (for example, how many accounts share one `device_id`). ## Install @@ -8,16 +35,430 @@ ShieldLabs server SDK: webhook verification and History API client. pip install shieldlabs ``` -## Verify webhooks +Python 3.9 or newer. The only dependency is [httpx](https://www.python-httpx.org/). + +## Quick start + +```python +import os +from typing import Optional + +from shieldlabs import ShieldLabs, evaluate_identification, webhooks + +client = ShieldLabs(api_key=os.environ["SHIELDLABS_API_KEY"]) # Private API Key, sec_... +used_request_ids: set[str] = set() # use your database or cache in production + + +def allow_signup(request_id: str) -> bool: + # 1. Read the identification for the request ID the browser sent with the form. + # Scoring is asynchronous, so this waits (up to 10 s by default) for the verdict. + identification = client.identifications.get(request_id) + + # 2. Evaluate it: missing, reused, stale, rate-limited, automated or dangerous is refused. + verdict = evaluate_identification(identification, is_replay=lambda rid: rid in used_request_ids) + if identification is not None: + used_request_ids.add(identification.request_id) + return verdict.ok + + +def on_webhook(raw_body: bytes, signature_header: Optional[str]) -> None: + # 3. Verify and parse a delivery: the raw body bytes and the X-Shield-Signature header. + event = webhooks.construct_event( + raw_body, + signature_header, + os.environ["SHIELDLABS_WEBHOOK_SECRET"], # whsec_... + ) + print(event.event_type) +``` + +A runnable FastAPI app that does all of this is in [`examples/`](examples/). + +## Guide + +### Wait for the verdict + +Scoring is asynchronous. The History row for an identification appears about 1-3 seconds after +the browser call and can be refined for up to about 10 seconds while follow-up checks finish. +Start the identification in the browser when the user begins the action (for example when they +start filling in the signup form) rather than on submit, so the verdict is usually ready when +your backend asks for it. An identification older than `max_age` (5 minutes by default) counts +as stale, so start a new one when the user comes back later. + +`identifications.get` polls the History API by `request_id` until the row appears and returns +that first version: ```python -from shieldlabs import verify_webhook +identification = client.identifications.get(request_id) # wait up to 10 s +identification = client.identifications.get(request_id, timeout=5) # shorter budget +identification = client.identifications.get(request_id, wait=False) # one lookup only + +if identification is None: + ... # not scored in time: treat it as unverified, never as clean +``` + +How the wait works: + +- **Total budget.** `timeout` (10 seconds by default) is the time budget of the whole call, not + of one request. +- **Schedule.** The first poll runs immediately, then after waits of 0.25 s, 0.5 s, 1 s and + 1.5 s, then every 2 s. The last poll runs at the deadline. `poll_interval` (p, 0.25 s by + default) sets the ladder: waits of p, 2p, 4p, 6p and 8p, then 8p again, each capped at + 2 seconds, or at p when p is longer. For example, `poll_interval=0.1` waits 0.1, 0.2, 0.4 and + 0.6 s, then every 0.8 s; `poll_interval=1` waits 1 s, then every 2 s; and `poll_interval=3` + polls every 3 s. +- **One attempt per poll.** Each poll is a single HTTP request, never retried inside the poll. + Its timeout is the client `timeout`, cut to the time left before the deadline but never + shorter than 1 second. +- **Transient errors keep polling.** A `429`, a 5xx response, a connection error or a timeout + does not end the wait: the next poll follows the schedule. If the last poll fails, its error + is raised; if it answers without a row, the result is `None`. +- **After a `429`.** The next wait is the longest of the ladder step, 1 second (the History API + limit is counted per second) and `Retry-After` capped at 10 seconds. A `Retry-After` of `0` or + a date in the past counts as 0, so the 1-second minimum still applies. A wait that would pass + the deadline is cut, and the last poll runs at the deadline. When the capped `Retry-After` is + longer than the time left, the `RateLimitError` is raised at once. +- **Errors that stop at once.** A `400`, `401`, `403` or `404` ends the wait immediately and + raises `BadRequestError`, `AuthenticationError` or `NotFoundError`. +- **`wait=False`** makes one lookup with the client's regular retries and returns `None` when + there is no row yet. +- `request_id` must be a UUID. An invalid value raises `ValidationError` before any request. +- For the refined state (for example in a later review job), read the row again with + `client.history.search("request_id", request_id, limit=1)`. + +### Decide with `evaluate_identification` + +`evaluate_identification` applies the checks our tutorials use before a protected action, in +this order, and reports the first one that fails: + +| `reason` | Refused when | +|---|---| +| `missing` | there is no identification (unverified, never clean) | +| `replayed` | `is_replay(request_id)` returns `True` (one identification authorizes one action) | +| `stale` | `observed_at` is older than `max_age` (default 300 seconds) | +| `rate_limited` | the Risk Score is the rate-limit marker (above 100, in practice 999) | +| `no_device_signals` | the device ID is the all-zero UUID `00000000-0000-0000-0000-000000000000` | +| `blocked_flag` | a flag in `block_flags` is set (default `browser_automation`, `javascript_disabled`) | +| `blocked_band` | the risk band is in `block_bands` (default `dangerous`) | + +```python +from datetime import timedelta + +from shieldlabs import evaluate_identification + +verdict = evaluate_identification( + identification, + max_age=timedelta(minutes=5), + block_bands=["suspicious", "dangerous"], + block_flags=["browser_automation", "javascript_disabled", "anti_detect_browser"], + is_replay=replay_store.seen_before, +) +# Evaluation(ok=False, reason='blocked_flag', band='dangerous', flag='anti_detect_browser') +``` + +The defaults are a starting point: tune the bands, flags and freshness window for each action. +The SDK stores nothing, so keep used request IDs in your own store (for example a Redis +`SET key NX EX 600`) and pass a lookup as `is_replay`. + +A visitor IP that sends too many identifications is blocked for 10 minutes. The block can show +up once as a separate identification with the marker 999 (`rate_limited`) and its own request +ID. Request IDs that the browser receives during the block get no History row at all, so they +end up as `missing`. + +Risk bands are computed on the client from the score: + +| Band | Risk Score | +|---|---| +| `trusted` | 0-29 | +| `suspicious` | 30-59 | +| `dangerous` | 60-100 | +| `rate_limited` | above 100: the rate-limit marker, not a score | + +```python +from shieldlabs import is_rate_limited, risk_band + +risk_band(45) # 'suspicious' +is_rate_limited(999) # True +identification.risk_band # same helpers as properties +``` + +Branch on `risk_score` and `detection_flags`. Risk signal names (`identification.signals`) are +for display and logging: the set is open (`SignalName` lists known values), names can repeat, +and weights can be negative, so never add weights up yourself. + +### The `Identification` model + +Webhook deliveries and History API rows are normalized into one frozen dataclass with the webhook +field names: + +| Field | Type | Notes | +|---|---|---| +| `request_id`, `visitor_id`, `device_id`, `session_id`, `cookie_id` | `str` | UUIDs; the all-zero UUID is possible | +| `user_hid` | `str` or `None` | your User HID; `"anonymous"` for anonymous checks | +| `domain` | `str` | the registered domain | +| `public_ip`, `local_ip` | `IpInfo(ip, country)` | IPv4 or `""`; `country` is an English country name such as `"Germany"`, or `""` | +| `connection_type` | `str` | `direct`, `mobile`, `vpn`, `proxy`, `tor`, `privacy_relay`, `browser_vpn_proxy`, `unknown` (unknown values are kept) | +| `os`, `browser`, `device_type` | `str` | | +| `traffic_source` | `TrafficSource` | `channel`, `referrer_domain`, `landing_url`, `click_id_type`, `utm_*`; `""` when absent | +| `risk_score` | `int` | 0-100, or 999 for the rate-limit marker | +| `signals` | `tuple[Signal, ...]` | `Signal(name, weight, description)`; `description` is set on History rows | +| `detection_flags` | `DetectionFlags` | 19 booleans; `.active()` lists the set ones | +| `observed_at` | `datetime` or `None` | timezone-aware UTC | +| `source` | `"webhook"` or `"history"` | | +| `raw` | `Mapping` | the original webhook `data` object or History row | + +`identification.to_dict()` returns a JSON-ready dict and `Identification.from_dict()` reads it +back. + +### Search history for account-abuse checks + +`history.search` reads one page of identifications that share one identifier, newest first. +`history.iter` walks every page for you, removes rows repeated between pages (new +identifications can shift offsets) and stops at `total`, at an empty page or after `max_items`. + +```python +page = client.history.search("user_hid", account_hid, limit=50) +print(page.total, len(page.data)) + +# How many accounts signed in from this device? +# User HID values that name no account (anonymous checks and sentinel values): +NOT_ACCOUNTS = (None, "anonymous", "fail", "-1", "unknown") + +if identification.has_device_signals: # the all-zero device ID matches unrelated rows + accounts = { + item.user_hid + for item in client.history.iter("device_id", identification.device_id, max_items=500) + if item.user_hid not in NOT_ACCOUNTS + } + if len(accounts) > 2: + send_to_review(identification.request_id, accounts) +``` + +| `type` | `value` | +|---|---| +| `request_id`, `device_id`, `visitor_id`, `session_id`, `cookie_id` | a UUID (sent lowercase) | +| `user_hid` | a non-empty string, matched exactly and case-sensitively | +| `ip` | a dotted IPv4 address | + +Arguments are validated before any request (`ValidationError`): unknown types, malformed UUIDs, +IPv6 addresses, `limit` outside 1-100, a negative `offset`, and a `user_hid` that is empty, +contains `/` or is `.` or `..`. Those User HIDs cannot be matched in the request path, so the SDK +refuses them instead of returning an empty page. Every other `user_hid` is percent-encoded in the +form the History API matches. + +### User HID + +Pass a stable, pseudonymous account ID to the browser agent instead of an email address or raw +account ID. `user_hid` derives one on your server: + +```python +from shieldlabs import user_hid + +hid = user_hid(str(account.id), user_hid_secret) # 64 lowercase hex characters +``` + +It is HMAC-SHA256 keyed with a secret of your own (any server-side secret, separate from your +ShieldLabs keys). Keep it private and stable: changing it changes every User HID. The hex output +is always searchable with `history.search("user_hid", ...)`; if you build User HIDs another way, +avoid `/` (standard base64 contains it, base64url does not). + +### Webhooks + +ShieldLabs sends `identification.scored` to every enabled endpoint (register them in the +analytics dashboard under **Integration > Webhooks**). Each delivery carries +`X-Shield-Signature: sha256=`, keyed with the endpoint signing +secret including its `whsec_` prefix. + +```python +from typing import Optional + +from shieldlabs import ( + IdentificationScoredEvent, + SignatureVerificationError, + WebhookParseError, + WebhookPingEvent, + webhooks, +) + + +def handle_delivery(raw_body: bytes, signature: Optional[str]) -> int: + try: + event = webhooks.construct_event(raw_body, signature, [current_secret, previous_secret]) + except SignatureVerificationError: + return 401 + except WebhookParseError: + return 400 + if isinstance(event, IdentificationScoredEvent): + if already_processed(event.data.request_id): + return 200 + enqueue(event.data) # do slow work after responding + elif isinstance(event, WebhookPingEvent): + pass # sent by the Verify button + return 200 # unknown event types: acknowledge and ignore +``` + +- Verify the raw bytes exactly as received, before parsing. Re-serialized JSON does not match. +- `secret` can be a list: a delivery is valid when any secret matches, so you can rotate an + endpoint secret without downtime. +- ShieldLabs sends one delivery per identification and endpoint, with a 1-second timeout and no + retries. Respond with a 2xx within 1 second and do slow work afterwards. +- Make handlers idempotent on `data.request_id`: a future release retries deliveries, and a + retry resends identical bytes. +- Use the History API for guaranteed reads and for the latest state: a delivery that fails is + not sent again, and a History row can be refined after its webhook was sent. +- `construct_event` returns `IdentificationScoredEvent`, `WebhookPingEvent` or + `UnknownWebhookEvent`, and never raises for an unknown event type. The **Test** delivery sent + from the analytics dashboard parses like production traffic. + +`webhooks.verify_signature(payload, signature_header, secret)` returns a bool when you only need +the check. + +### Management API: domain profile + +```python +from shieldlabs import ShieldLabsManagement + +management = ShieldLabsManagement( + secret_key=os.environ["SHIELDLABS_SECRET_KEY"], + # Normalized before use: "https://www.Example.com/" becomes "example.com". + domain=os.environ["SHIELDLABS_DOMAIN"], +) +profile = management.get_profile() +profile.remaining_identifications # negative when the account is over its included volume +profile.public_key_masked # "****************************a3f8" +``` + +The Management API allows about 15 requests per minute per caller IP and then blocks that IP +for 10 minutes. The client never retries a `429`, so call it sparingly and cache the profile. + +### Rate limits + +| API | Limit | What the SDK does | +|---|---|---| +| History API | about 15 requests per second per domain, shared by all your callers | `identifications.get` spaces its polls, and inside its wait a `429` waits at least 1 second (and at least `Retry-After`, up to 10 seconds). Ordinary calls (`history.search`, `history.iter`, `identifications.get` with `wait=False`) follow `Retry-After` as sent, up to 10 seconds, and wait at least 1 second after a `429` without it | +| Management API | about 15 requests per minute per IP, then a 10-minute block | raises `RateLimitError` without retrying | + +### Async + +`AsyncShieldLabs` and `AsyncShieldLabsManagement` mirror the sync clients: + +```python +from shieldlabs import AsyncShieldLabs + +async with AsyncShieldLabs() as client: # reads SHIELDLABS_API_KEY + identification = await client.identifications.get(request_id) + async for item in client.history.iter("user_hid", account_hid, max_items=200): + ... +``` + +### Configuration + +| Option | Default | Notes | +|---|---|---| +| `api_key` | `SHIELDLABS_API_KEY` | Private API Key `sec_...`; a key of another shape triggers a `ShieldLabsWarning` | +| `base_url` | `SHIELDLABS_API_BASE_URL`, else `https://account.shieldlabs.ai` | the origin; a trailing `/api` is removed | +| `secret_key`, `domain` | `SHIELDLABS_SECRET_KEY`, `SHIELDLABS_DOMAIN` | Management client | +| `base_url` (Management) | `SHIELDLABS_MANAGEMENT_BASE_URL`, else `https://api.shieldlabs.ai` | | +| `timeout` | `10.0` | seconds per HTTP attempt | +| `max_retries` | `2` | retries for connection errors, timeouts, `429` (History only) and 5xx | +| `http_client` | a new `httpx.Client` / `httpx.AsyncClient` | pass your own for proxies or custom transports; it is not closed for you | + +Base URLs must use https. Plain `http://` is accepted only for `localhost`, `127.0.0.1` and +`[::1]` (local test servers), because every request carries a key. + +Clients are safe to share across threads (sync) or tasks (async): create one per process and +reuse it. Use them as context managers or call `close()` / `aclose()`. + +For development and staging, register a separate domain (for example `dev.example.com`) and use +its keys with the default hosts. + +## Reference + +| Call | Returns | +|---|---| +| `ShieldLabs(api_key=None, base_url=None, timeout=10.0, max_retries=2, http_client=None)` | History API client | +| `client.identifications.get(request_id, wait=True, timeout=10.0, poll_interval=0.25)` | `Identification` or `None`; `timeout` is the total wait in seconds; `poll_interval` p sets the waits p, 2p, 4p, 6p, 8p, then 8p again, each at most 2 seconds, or p when p is longer | +| `client.history.search(type, value, limit=20, offset=0)` | `HistoryPage(data, total)` | +| `client.history.iter(type, value, page_size=100, max_items=None)` | iterator of `Identification`; `max_items=0` yields nothing | +| `AsyncShieldLabs(...)` | same methods as coroutines; `history.iter` is an async iterator | +| `ShieldLabsManagement(secret_key=None, domain=None, base_url=None, timeout=10.0, max_retries=2, http_client=None)` | Management API client | +| `management.get_profile()` | `DomainProfile` | +| `AsyncShieldLabsManagement(...)` | same method as a coroutine | +| `webhooks.verify_signature(payload, signature_header, secret)` | `bool` | +| `webhooks.construct_event(payload, signature_header, secret)` | `IdentificationScoredEvent`, `WebhookPingEvent` or `UnknownWebhookEvent` | +| `evaluate_identification(identification, *, max_age=300.0, now=None, block_bands=("dangerous",), block_flags=("browser_automation", "javascript_disabled"), is_replay=None)` | `Evaluation(ok, reason, band, flag)` | +| `risk_band(score)` | `"trusted"`, `"suspicious"`, `"dangerous"` or `"rate_limited"` | +| `is_rate_limited(score)` | `bool` | +| `user_hid(user_id, secret)` | 64-character lowercase hex `str` | + +Models: `Identification`, `IpInfo`, `TrafficSource`, `Signal`, `SignalName`, `DetectionFlags`, +`HistoryPage`, `DomainProfile`, `Evaluation`. Type aliases: `LookupType`, `RiskBand`, +`EvaluationReason`, `WebhookEvent`. + +## Errors and retries + +| Exception | When | +|---|---| +| `ShieldLabsError` | base class of everything below | +| `ApiError` | any non-2xx response; has `status`, `message`, `body` (parsed JSON or text) and `headers` | +| `BadRequestError` | 400 | +| `AuthenticationError` | 401 or 403: wrong, rotated or disabled key | +| `QuotaExceededError` | 402. Neither the History API nor the Management API returns it today: an account over its included volume shows a negative `remaining_identifications` | +| `NotFoundError` | 404: usually a wrong base URL or path prefix | +| `RateLimitError` | 429; `retry_after` holds seconds when the server sent `Retry-After` | +| `ServerError` | 5xx | +| `APIConnectionError` | DNS, TCP, TLS or protocol failure | +| `APITimeoutError` | an attempt exceeded `timeout` | +| `SignatureVerificationError` | a webhook signature is missing or wrong | +| `WebhookParseError` | a verified webhook body is not an event envelope | +| `ValidationError` | an invalid argument, raised before any request (also a `ValueError`) | + +Only GET requests are sent, and they are retried on connection errors, timeouts, `429` and 5xx +with exponential backoff and jitter (0.5 s base, doubling, capped at 8 s), up to `max_retries` +times. `Retry-After` is followed as sent, up to 10 seconds (`0` retries at once), and a `429` +without it waits at least 1 second (the History API limit is counted per second). `400`, `401`, +`402`, `403` and `404` are never retried, and the Management client never retries `429`. Inside +`identifications.get` a poll is never retried: the wait polls again on its schedule instead (see +[Wait for the verdict](#wait-for-the-verdict)). + +```python +from shieldlabs import ApiError, AuthenticationError, RateLimitError + +try: + page = client.history.search("device_id", device_id) +except AuthenticationError: + ... # check SHIELDLABS_API_KEY +except RateLimitError as exc: + ... # exc.retry_after +except ApiError as exc: + print(exc.status, exc.message) +``` + +Keys and request bodies are never logged. Each request carries +`User-Agent: shieldlabs-python/` with the Python and httpx versions. + +## Compatibility + +- Python 3.9, 3.10, 3.11, 3.12 and 3.13 (CPython), tested in CI. +- httpx 0.25 or newer, below 1.0. The async client runs on asyncio and trio. +- Webhook `schema_version` `2026-06-01`. Other versions are parsed with a `ShieldLabsWarning`. +- Unknown fields and enum values in responses are tolerated and kept in `raw`. +- The package follows semantic versioning and ships type information (`py.typed`). + +## Development + +```bash +python3 -m venv .venv +source .venv/bin/activate +pip install -e ".[dev]" -ok = verify_webhook(raw_body, request.headers["X-Shield-Signature"], secret) +ruff check . && ruff format --check . && mypy --strict src +pytest -q --cov=shieldlabs --cov-report=term-missing ``` -Signature: `X-Shield-Signature: sha256=` + hex(HMAC-SHA256(secret, raw_body)). Schema `2026-06-01`. +`tests/data/` holds the shared test fixtures (History rows, webhook bodies, signature vectors, +error responses) that every ShieldLabs server SDK passes. See +[CONTRIBUTING.md](CONTRIBUTING.md). Questions and security reports: contact@shieldlabs.ai. ## License -[MIT](./LICENSE) +[MIT](LICENSE) diff --git a/examples/README.md b/examples/README.md new file mode 100644 index 0000000..134305c --- /dev/null +++ b/examples/README.md @@ -0,0 +1,44 @@ +# Examples + +## `fastapi_app.py` + +A FastAPI app that does both halves of a server integration: + +- `POST /signup` reads the `requestId` sent by the browser, waits for the verdict with + `identifications.get`, and refuses the signup when the identification is missing, reused, + stale, carries the rate-limit marker, has no device signals, shows browser automation or + disabled JavaScript, or falls in the dangerous band (`evaluate_identification` defaults). +- `POST /webhooks/shieldlabs` verifies `X-Shield-Signature` over the raw body, parses the event, + handles each `request_id` once (so a retried delivery is harmless), and logs it. + +Run it from the repository root: + +```bash +python -m venv .venv && source .venv/bin/activate +pip install -r examples/requirements.txt # before the package is on PyPI: pip install -e . first +export SHIELDLABS_API_KEY=sec_your_private_key +export SHIELDLABS_WEBHOOK_SECRET=whsec_your_signing_secret +uvicorn examples.fastapi_app:app --port 8000 +``` + +Try the signup route with a request ID from your page: + +```bash +curl -X POST http://localhost:8000/signup \ + -H 'content-type: application/json' \ + -d '{"email": "user@example.com", "requestId": "8f14e45f-ceea-4c1e-a3b2-1d2c3b4a5f60"}' +``` + +To receive webhooks locally, expose port 8000 with an HTTPS tunnel and register +`https:///webhooks/shieldlabs` in the analytics dashboard under +**Integration > Webhooks**. Press **Verify** to send a `webhook.ping`. + +The in-memory sets keep the example short. In production, store used request IDs and processed +webhook request IDs in your database or cache. + +The smoke test for this app is `tests/test_example_app.py`; it runs when FastAPI is installed: + +```bash +pip install -e ".[dev]" -r examples/requirements.txt +pytest tests/test_example_app.py +``` diff --git a/examples/fastapi_app.py b/examples/fastapi_app.py new file mode 100644 index 0000000..3ba6cc2 --- /dev/null +++ b/examples/fastapi_app.py @@ -0,0 +1,141 @@ +"""Signup protection and a webhook receiver with FastAPI and the ShieldLabs Python SDK. + +The browser runs an identification with the ShieldLabs agent and sends the resulting +``requestId`` with the signup form. This app reads the verdict for that request ID from the +History API, applies a policy, and separately receives signed ``identification.scored`` +webhooks. + +Environment: + SHIELDLABS_API_KEY Private API Key (sec_...), reads verdicts from the History API. + SHIELDLABS_WEBHOOK_SECRET Endpoint signing secret (whsec_...), verifies webhook deliveries. + Several secrets can be given, separated by commas, while you + rotate one. + +Run from the repository root: + pip install -r examples/requirements.txt + uvicorn examples.fastapi_app:app --port 8000 +""" + +from __future__ import annotations + +import logging +import os +from collections.abc import AsyncIterator +from contextlib import asynccontextmanager + +from fastapi import FastAPI, Request, Response +from fastapi.responses import JSONResponse +from pydantic import BaseModel, Field + +from shieldlabs import ( + AsyncShieldLabs, + IdentificationScoredEvent, + ShieldLabsError, + SignatureVerificationError, + ValidationError, + WebhookParseError, + WebhookPingEvent, + evaluate_identification, + webhooks, +) + +logging.basicConfig(level=logging.INFO) +logger = logging.getLogger("shieldlabs.example") + +# In-memory stores keep the example self-contained. Use your database or cache in production: +# one request ID authorizes one protected action, and the webhook handler stays idempotent on +# the request ID (one delivery per identification today; a future release retries deliveries +# with identical bytes). +used_request_ids: set[str] = set() +processed_request_ids: set[str] = set() + + +def _webhook_secrets() -> list[str]: + # Several secrets, separated by commas, while you rotate one: "whsec_new, whsec_old". + raw = os.environ.get("SHIELDLABS_WEBHOOK_SECRET", "") + return [secret.strip() for secret in raw.split(",") if secret.strip()] + + +@asynccontextmanager +async def lifespan(app: FastAPI) -> AsyncIterator[None]: + # One client for the whole process; it reads SHIELDLABS_API_KEY and is safe to share. + app.state.shieldlabs = AsyncShieldLabs() + try: + yield + finally: + await app.state.shieldlabs.aclose() + + +app = FastAPI(title="ShieldLabs signup example", lifespan=lifespan) + + +class SignupForm(BaseModel): + email: str + request_id: str = Field(alias="requestId") + + +@app.post("/signup") +async def signup(form: SignupForm, request: Request) -> JSONResponse: + client: AsyncShieldLabs = request.app.state.shieldlabs + try: + # Waits for the verdict (up to 10 s): the History row appears about 1-3 s after the + # browser call, so the page starts the identification when the user starts filling in + # the form rather than on submit. + identification = await client.identifications.get(form.request_id) + except ValidationError: + return JSONResponse({"error": "invalid_request_id"}, status_code=400) + except ShieldLabsError: + # Unverified is never clean: refuse and let the user retry. + logger.exception("ShieldLabs lookup failed") + return JSONResponse({"error": "verification_unavailable"}, status_code=503) + + # Default policy: refuse a missing, reused or stale identification, the rate-limit marker, + # an identification without device signals, browser automation or disabled JavaScript, + # and the dangerous band. Tune it for your product. + verdict = evaluate_identification( + identification, + is_replay=lambda request_id: request_id in used_request_ids, + ) + if identification is not None: + used_request_ids.add(identification.request_id) + if not verdict.ok: + logger.info( + "signup refused: reason=%s band=%s flag=%s", verdict.reason, verdict.band, verdict.flag + ) + return JSONResponse({"error": "signup_refused", "reason": verdict.reason}, status_code=403) + + # Create the account here. + return JSONResponse({"status": "created"}, status_code=201) + + +@app.post("/webhooks/shieldlabs") +async def shieldlabs_webhook(request: Request) -> Response: + payload = await request.body() # the raw bytes: verify them before parsing anything + try: + event = webhooks.construct_event( + payload, request.headers.get(webhooks.SIGNATURE_HEADER), _webhook_secrets() + ) + except SignatureVerificationError: + return Response(status_code=401) + except WebhookParseError: + return Response(status_code=400) + + if isinstance(event, IdentificationScoredEvent): + data = event.data + if data.request_id in processed_request_ids: + return Response(status_code=200) # already handled: acknowledge and stop + processed_request_ids.add(data.request_id) + # Answer fast (under 1 s): a slow or failed delivery is not sent again, so queue slow + # work instead of doing it here, and read the History API when you must not miss one. + logger.info( + "identification.scored request_id=%s risk_score=%s band=%s flags=%s", + data.request_id, + data.risk_score, + data.risk_band, + ",".join(data.detection_flags.active()) or "none", + ) + elif isinstance(event, WebhookPingEvent): + logger.info("webhook.ping received") + else: + logger.info("ignored event_type=%s", event.event_type) + return Response(status_code=200) diff --git a/examples/requirements.txt b/examples/requirements.txt new file mode 100644 index 0000000..3a9d236 --- /dev/null +++ b/examples/requirements.txt @@ -0,0 +1,3 @@ +shieldlabs>=1.0.0,<2 +fastapi>=0.110,<0.143 +uvicorn>=0.29,<0.55 diff --git a/pyproject.toml b/pyproject.toml index 2738967..3866c69 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,33 +1,103 @@ [build-system] -requires = ["hatchling"] +requires = ["hatchling>=1.27"] build-backend = "hatchling.build" [project] name = "shieldlabs" -version = "0.1.0" -description = "ShieldLabs server SDK for Python: API client, webhook verification, and types." +dynamic = ["version"] +description = "ShieldLabs server SDK for Python: History API and Management API clients, webhook verification, typed events and risk helpers." readme = "README.md" -requires-python = ">=3.8" -license = { text = "MIT" } -authors = [{ name = "ShieldLabs", email = "contact@shieldlabs.ai" }] -keywords = ["shieldlabs", "server-sdk", "webhook-verification", "risk-score", "fraud-prevention"] +requires-python = ">=3.9" +license = "MIT" +license-files = ["LICENSE"] +authors = [{ name = "ShieldLabs Inc.", email = "contact@shieldlabs.ai" }] +keywords = [ + "shieldlabs", + "device-intelligence", + "fraud-detection", + "risk-score", + "webhooks", + "account-abuse", +] classifiers = [ - "Development Status :: 4 - Beta", + "Development Status :: 5 - Production/Stable", + "Intended Audience :: Developers", + "Operating System :: OS Independent", "Programming Language :: Python :: 3", - "License :: OSI Approved :: MIT License", + "Programming Language :: Python :: 3 :: Only", + "Programming Language :: Python :: 3.9", + "Programming Language :: Python :: 3.10", + "Programming Language :: Python :: 3.11", + "Programming Language :: Python :: 3.12", + "Programming Language :: Python :: 3.13", + "Topic :: Internet :: WWW/HTTP", + "Topic :: Security", + "Topic :: Software Development :: Libraries :: Python Modules", + "Typing :: Typed", ] -dependencies = [] +dependencies = ["httpx>=0.25,<1"] [project.optional-dependencies] -dev = ["pytest>=7.0"] +# Upper bounds keep CI results stable when a tool publishes a new release; raise them on purpose. +dev = [ + "anyio>=3.7,<5", + "mypy>=1.8,<2.4", + "pytest>=7.4,<10", + "pytest-cov>=4.1,<8", + "respx>=0.20.2,<0.24", + "ruff>=0.6,<0.17", +] [project.urls] Homepage = "https://shieldlabs.ai" +Documentation = "https://docs.shieldlabs.ai" Repository = "https://github.com/ShieldLabs-ai/shieldlabs-python" Issues = "https://github.com/ShieldLabs-ai/shieldlabs-python/issues" +Changelog = "https://github.com/ShieldLabs-ai/shieldlabs-python/blob/main/CHANGELOG.md" + +[tool.hatch.version] +path = "src/shieldlabs/_version.py" [tool.hatch.build.targets.wheel] packages = ["src/shieldlabs"] +[tool.hatch.build.targets.sdist] +include = ["src/shieldlabs", "tests", "examples", "README.md", "CHANGELOG.md", "LICENSE"] + [tool.pytest.ini_options] testpaths = ["tests"] +addopts = "-ra" +filterwarnings = ["error::DeprecationWarning:shieldlabs.*"] + +[tool.coverage.run] +source = ["shieldlabs"] +branch = true + +[tool.coverage.report] +fail_under = 90 +show_missing = true +skip_covered = false +exclude_also = [ + "if TYPE_CHECKING:", + "raise NotImplementedError", + "@overload", +] + +[tool.ruff] +line-length = 100 +target-version = "py39" +src = ["src", "tests"] + +[tool.ruff.lint] +select = ["E", "W", "F", "I", "B", "UP", "SIM", "C4", "RUF", "N", "PT", "RET", "PIE"] + +[tool.ruff.lint.pyupgrade] +# Keep typing.Optional / typing.Union so annotations stay evaluable on Python 3.9. +keep-runtime-typing = true + +[tool.ruff.lint.per-file-ignores] +"tests/**" = ["PT011"] + +[tool.mypy] +strict = true +warn_unreachable = true diff --git a/src/shieldlabs/__init__.py b/src/shieldlabs/__init__.py index 489ba24..bf7876f 100644 --- a/src/shieldlabs/__init__.py +++ b/src/shieldlabs/__init__.py @@ -1,81 +1,99 @@ """ShieldLabs server SDK for Python. -Talks to the ShieldLabs API and verifies inbound webhooks. Your code decides -what to do with the score: you set the rules. This SDK never makes the -decision for you. -""" - -from __future__ import annotations - -import hashlib -import hmac -from typing import Any, Dict, Mapping, Optional, Union -from urllib.parse import urlencode -from urllib.request import Request, urlopen - -__version__ = "0.1.0" - -WEBHOOK_SCHEMA_VERSION = "2026-06-01" - -_DEFAULT_BASE_URL = "https://account.shieldlabs.ai/api" +Read identification verdicts from the History API, read the domain profile from the Management +API, verify and parse webhook deliveries, and apply risk helpers. +Quick start:: -def verify_webhook( - payload: Union[bytes, bytearray, memoryview], - signature: str, - secret: str, -) -> bool: - """Verify X-Shield-Signature against HMAC-SHA256(secret, raw body). + from shieldlabs import ShieldLabs, evaluate_identification - ``payload`` must be the raw request body bytes. Re-serializing parsed JSON - will fail verification. Comparison is constant-time. - """ - if not signature or not secret: - return False - body = bytes(payload) - expected = "sha256=" + hmac.new( - secret.encode("utf-8"), body, hashlib.sha256 - ).hexdigest() - return hmac.compare_digest(expected, signature) - - -class ShieldLabsClient: - """Thin client for ShieldLabs History API (account.shieldlabs.ai).""" - - def __init__(self, api_key: str, base_url: Optional[str] = None) -> None: - self.api_key = api_key - self.base_url = (base_url or _DEFAULT_BASE_URL).rstrip("/") - - def get_history( - self, - search_type: str, - value: str, - *, - limit: Optional[int] = None, - offset: Optional[int] = None, - ) -> Dict[str, Any]: - """GET /api/v1/history/{search_type}/{value} → {data, total}.""" - path = f"{self.base_url}/api/v1/history/{search_type}/{value}" - query: Dict[str, str] = {} - if limit is not None: - query["limit"] = str(limit) - if offset is not None: - query["offset"] = str(offset) - if query: - path = f"{path}?{urlencode(query)}" - req = Request( - path, - headers={"Authorization": f"Bearer {self.api_key}"}, - ) - with urlopen(req) as resp: - import json - - return json.loads(resp.read().decode("utf-8")) + client = ShieldLabs(api_key="sec_your_private_key") + identification = client.identifications.get(request_id) + verdict = evaluate_identification(identification) +""" +from . import webhooks +from ._client import AsyncShieldLabs, ShieldLabs +from ._errors import ( + APIConnectionError, + ApiError, + APITimeoutError, + AuthenticationError, + BadRequestError, + NotFoundError, + QuotaExceededError, + RateLimitError, + ServerError, + ShieldLabsError, + ShieldLabsWarning, + SignatureVerificationError, + ValidationError, + WebhookParseError, +) +from ._helpers import Evaluation, EvaluationReason, evaluate_identification, user_hid +from ._management import AsyncShieldLabsManagement, ShieldLabsManagement +from ._models import ( + DetectionFlags, + DomainProfile, + HistoryPage, + Identification, + IdentificationSource, + IpInfo, + Signal, + SignalName, + TrafficSource, +) +from ._normalize import NIL_UUID, RiskBand, is_rate_limited, risk_band +from ._validation import LookupType +from ._version import __version__ +from .webhooks import ( + IdentificationScoredEvent, + UnknownWebhookEvent, + WebhookEvent, + WebhookPingEvent, +) __all__ = [ - "ShieldLabsClient", - "WEBHOOK_SCHEMA_VERSION", + "NIL_UUID", + "APIConnectionError", + "APITimeoutError", + "ApiError", + "AsyncShieldLabs", + "AsyncShieldLabsManagement", + "AuthenticationError", + "BadRequestError", + "DetectionFlags", + "DomainProfile", + "Evaluation", + "EvaluationReason", + "HistoryPage", + "Identification", + "IdentificationScoredEvent", + "IdentificationSource", + "IpInfo", + "LookupType", + "NotFoundError", + "QuotaExceededError", + "RateLimitError", + "RiskBand", + "ServerError", + "ShieldLabs", + "ShieldLabsError", + "ShieldLabsManagement", + "ShieldLabsWarning", + "Signal", + "SignalName", + "SignatureVerificationError", + "TrafficSource", + "UnknownWebhookEvent", + "ValidationError", + "WebhookEvent", + "WebhookParseError", + "WebhookPingEvent", "__version__", - "verify_webhook", + "evaluate_identification", + "is_rate_limited", + "risk_band", + "user_hid", + "webhooks", ] diff --git a/src/shieldlabs/_client.py b/src/shieldlabs/_client.py new file mode 100644 index 0000000..fded10a --- /dev/null +++ b/src/shieldlabs/_client.py @@ -0,0 +1,542 @@ +"""History API clients (sync and async).""" + +from __future__ import annotations + +import itertools +from collections.abc import AsyncIterator, Iterator +from types import TracebackType +from typing import Callable, Optional, Union +from uuid import UUID + +import httpx + +from ._errors import ( + APIConnectionError, + APITimeoutError, + RateLimitError, + ServerError, + ShieldLabsError, +) +from ._http import ( + RATE_LIMIT_MIN_DELAY, + RETRY_AFTER_CAP, + USER_AGENT, + AsyncTransport, + SyncTransport, + json_object, +) +from ._models import HistoryPage, Identification +from ._validation import ( + LookupType, + history_origin, + resolve_api_key, + validate_count, + validate_limit, + validate_lookup, + validate_max_retries, + validate_offset, + validate_seconds, + validate_uuid, +) + +__all__ = [ + "AsyncHistory", + "AsyncIdentifications", + "AsyncShieldLabs", + "History", + "Identifications", + "ShieldLabs", +] + +POLL_WAIT_CAP = 2.0 +"""Cap of every wait between two polls of ``identifications.get``, in seconds, unless +``poll_interval`` is longer: then ``poll_interval`` is the cap.""" + +POLL_MIN_ATTEMPT_TIMEOUT = 1.0 +"""Shortest timeout of one poll, in seconds, even when less time is left before the deadline.""" + +POLL_STEPS = (1, 2, 4, 6, 8) +"""Waits between polls as multiples of ``poll_interval``; the last multiple repeats.""" + +# Errors that do not end a wait: the next poll can still succeed. Everything else (400, 401, +# 403, 404 and any other API error) is raised at once. +_TRANSIENT_ERRORS = (RateLimitError, ServerError, APIConnectionError, APITimeoutError) +_HISTORY_PATH = "/api/v1/history" + + +def poll_waits(initial: float) -> Iterator[float]: + """Waits between polls: ``initial`` times 1, 2, 4, 6 and 8, then 8 again. + + Each wait is capped at ``max(2 s, initial)``. With the default initial wait of 0.25 s: 0.25, + 0.5, 1, 1.5, 2, 2, ... With 0.1 s: 0.1, 0.2, 0.4, 0.6, 0.8, 0.8, ... With 1 s: 1, 2, 2, ... + An initial wait of 2 s or more is used for every wait: 3 s gives 3, 3, 3, ... + """ + cap = max(POLL_WAIT_CAP, initial) + for factor in POLL_STEPS: + yield min(cap, initial * factor) + yield from itertools.repeat(min(cap, initial * POLL_STEPS[-1])) + + +class _PollPlan: + """Timing of one ``identifications.get`` wait, shared by the sync and async clients. + + ``budget`` is the total time the wait may take. Polls follow ``poll_waits`` and the last + one runs at the deadline. + """ + + def __init__( + self, + clock: Callable[[], float], + budget: float, + initial_wait: float, + client_timeout: float, + ) -> None: + self._clock = clock + self._deadline = clock() + budget + self._waits = poll_waits(initial_wait) + self._client_timeout = client_timeout + self._final = False + + def attempt_timeout(self) -> float: + """Timeout of the next poll: the client timeout, cut to the time left, at least 1 s.""" + remaining = self._deadline - self._clock() + return min(self._client_timeout, max(remaining, POLL_MIN_ATTEMPT_TIMEOUT)) + + def next_wait(self, error: Optional[ShieldLabsError]) -> Optional[float]: + """Seconds to wait before the next poll, or ``None`` when the wait ends now. + + ``error`` is the transient error of the poll that just ran, or ``None`` when it answered + without a row. Every wait takes the next ladder step. After a 429 the wait is the + longest of that step, 1 s and ``Retry-After`` capped at 10 s (a missing header, ``0`` or + a date in the past counts as 0). A capped ``Retry-After`` longer than the time left ends + the wait at once; any other wait is cut so that the last poll runs at the deadline. + """ + remaining = self._deadline - self._clock() + if self._final or remaining <= 0: + return None + wait = next(self._waits) + if isinstance(error, RateLimitError): + retry_after = min(error.retry_after or 0.0, RETRY_AFTER_CAP) + if retry_after > remaining: + return None + # The History API limit is counted per second: an earlier poll is refused again. + wait = max(wait, RATE_LIMIT_MIN_DELAY, retry_after) + if wait >= remaining: + wait = remaining + self._final = True + return wait + + +class _HistoryConfig: + """Configuration shared by the sync and async History clients.""" + + base_url: str + """History API origin, for example ``https://account.shieldlabs.ai``.""" + + def _configure( + self, + api_key: Optional[str], + base_url: Optional[str], + timeout: float, + max_retries: int, + ) -> tuple[float, int]: + key = resolve_api_key(api_key) + self.base_url = history_origin(base_url) + self._headers = { + "Authorization": f"Bearer {key}", + "Accept": "application/json", + "User-Agent": USER_AGENT, + } + return ( + validate_seconds(timeout, "timeout", allow_zero=False), + validate_max_retries(max_retries), + ) + + def _history_url(self, lookup_type: str, value: Union[str, UUID]) -> str: + checked_type, segment = validate_lookup(lookup_type, value) + return f"{self.base_url}{_HISTORY_PATH}/{checked_type}/{segment}" + + def __repr__(self) -> str: + return f"{type(self).__name__}(base_url={self.base_url!r})" + + +class ShieldLabs(_HistoryConfig): + """Client for the ShieldLabs History API (read identifications with a Private API Key). + + Args: + api_key: Private API Key (``sec_...``). Defaults to ``SHIELDLABS_API_KEY``. + base_url: History API origin. Defaults to ``SHIELDLABS_API_BASE_URL`` or + ``https://account.shieldlabs.ai``. A trailing ``/api`` is removed. Must use https; + plain http is accepted only for localhost, 127.0.0.1 and ::1 (local test servers). + timeout: Timeout of one HTTP attempt, in seconds. + max_retries: Retries for connection errors, timeouts, 429 and 5xx responses. + http_client: Your own ``httpx.Client`` (proxies, transports). It is not closed by + ``close()``. + + The client is safe to share across threads. Use it as a context manager or call + ``close()`` when you are done. + """ + + def __init__( + self, + api_key: Optional[str] = None, + base_url: Optional[str] = None, + timeout: float = 10.0, + max_retries: int = 2, + http_client: Optional[httpx.Client] = None, + ) -> None: + checked_timeout, checked_retries = self._configure(api_key, base_url, timeout, max_retries) + self._transport = SyncTransport( + client=http_client, + timeout=checked_timeout, + max_retries=checked_retries, + retry_rate_limited=True, + ) + self.history = History(self) + self.identifications = Identifications(self) + + def _search( + self, + lookup_type: str, + value: Union[str, UUID], + limit: int, + offset: int, + *, + timeout: Optional[float] = None, + max_retries: Optional[int] = None, + ) -> HistoryPage: + url = self._history_url(lookup_type, value) + response = self._transport.get( + url, + headers=self._headers, + params={"limit": validate_limit(limit), "offset": validate_offset(offset)}, + timeout=timeout, + max_retries=max_retries, + ) + return HistoryPage.from_dict(json_object(response)) + + def close(self) -> None: + """Close the underlying HTTP client (unless it was passed in).""" + self._transport.close() + + def __enter__(self) -> ShieldLabs: + return self + + def __exit__( + self, + exc_type: Optional[type[BaseException]], + exc: Optional[BaseException], + tb: Optional[TracebackType], + ) -> None: + self.close() + + +class History: + """``client.history``: search identifications by one identifier.""" + + def __init__(self, client: ShieldLabs) -> None: + self._client = client + + def search( + self, + type: LookupType, + value: Union[str, UUID], + limit: int = 20, + offset: int = 0, + ) -> HistoryPage: + """Read one page of identifications that match ``type`` = ``value``, newest first. + + Args: + type: ``ip``, ``user_hid``, ``visitor_id``, ``request_id``, ``device_id``, + ``session_id`` or ``cookie_id``. + value: A UUID for the ID types, a dotted IPv4 address for ``ip``, a non-empty + string for ``user_hid`` (matched exactly; ``.``, ``..`` and values that contain + ``/`` cannot be matched in the request path and raise ``ValidationError``). + limit: Page size, 1 to 100. + offset: Rows to skip. + + Raises: + ValidationError: Invalid arguments (nothing is sent). + ApiError: The API answered with an error status. + """ + return self._client._search(type, value, limit, offset) + + def iter( + self, + type: LookupType, + value: Union[str, UUID], + page_size: int = 100, + max_items: Optional[int] = None, + ) -> Iterator[Identification]: + """Iterate over every identification that matches, newest first, page by page. + + Rows are de-duplicated on ``request_id`` (offset paging can repeat a row while new + identifications arrive). Iteration stops at ``total``, at an empty page, or after + ``max_items`` identifications. + """ + validate_lookup(type, value) + size = validate_limit(page_size, "page_size") + if max_items is not None: + max_items = validate_count(max_items, "max_items") + return self._iterate(type, value, size, max_items) + + def _iterate( + self, + lookup_type: str, + value: Union[str, UUID], + page_size: int, + max_items: Optional[int], + ) -> Iterator[Identification]: + seen: set[str] = set() + offset = 0 + count = 0 + while max_items is None or count < max_items: + page = self._client._search(lookup_type, value, page_size, offset) + if not page.data: + return + for identification in page.data: + if identification.request_id in seen: + continue + seen.add(identification.request_id) + yield identification + count += 1 + if max_items is not None and count >= max_items: + return + offset += len(page.data) + if offset >= page.total: + return + + +class Identifications: + """``client.identifications``: read the verdict for one request ID.""" + + def __init__(self, client: ShieldLabs) -> None: + self._client = client + + def get( + self, + request_id: Union[str, UUID], + wait: bool = True, + timeout: float = 10.0, + poll_interval: float = 0.25, + ) -> Optional[Identification]: + """Return the identification for ``request_id``, waiting for it to be scored. + + Scoring is asynchronous: the History row appears about 1-3 s after the browser call + and can be refined for up to about 10 s while follow-up checks finish. This method + returns the first version it finds; read the row again with ``history.search`` when you + need the refined state. + + With ``wait=True``, ``timeout`` is the total time budget of the call. The History API is + polled immediately, then after waits of ``poll_interval`` times 1, 2, 4, 6 and 8, then + 8 again, each wait capped at ``max(2 s, poll_interval)`` (by default 0.25 s, 0.5 s, 1 s, + 1.5 s and then every 2 s; every 3 s for ``poll_interval=3``), and a last time at the + deadline. Each poll is one HTTP attempt, never retried inside the poll, with a timeout + of ``min(client timeout, max(time left, 1 s))``. A 429, a 5xx response, a connection + error or a timeout does not end the wait: polling continues. After a 429 the next wait + is the longest of the ladder step, 1 s and ``Retry-After`` capped at 10 s + (``Retry-After: 0`` or a date in the past counts as 0), cut to the deadline; a capped + ``Retry-After`` longer than the time left raises the 429 at once. A 400, 401, 403 or 404 + is raised at once. + + With ``wait=False`` one lookup is made, with the client's regular retries. + + Returns: + The ``Identification``, or ``None`` when the last poll answered without a row (treat + that as unverified, never as clean). + + Raises: + ValidationError: ``request_id`` is not a UUID (nothing is sent). + BadRequestError, AuthenticationError, NotFoundError: A 400, a 401 or 403, or a 404 + answer; polling stops at once. + RateLimitError, ServerError, APIConnectionError, APITimeoutError: The last poll + failed with this error, or a 429 asked for a pause longer than the time left. + """ + rid = validate_uuid(request_id, "request_id") + total_timeout = validate_seconds(timeout, "timeout", allow_zero=True) + initial_wait = validate_seconds(poll_interval, "poll_interval", allow_zero=False) + client = self._client + if not wait: + page = client._search("request_id", rid, 1, 0) + return page.data[0] if page.data else None + + transport = client._transport + plan = _PollPlan(transport._clock, total_timeout, initial_wait, transport.timeout) + while True: + error: Optional[ShieldLabsError] = None + try: + page = client._search( + "request_id", rid, 1, 0, timeout=plan.attempt_timeout(), max_retries=0 + ) + except _TRANSIENT_ERRORS as exc: + error = exc + else: + if page.data: + return page.data[0] + delay = plan.next_wait(error) + if delay is None: + if error is not None: + raise error + return None + transport._sleep(delay) + + +class AsyncShieldLabs(_HistoryConfig): + """Asynchronous client for the ShieldLabs History API. + + Same arguments as ``ShieldLabs``; ``http_client`` is an ``httpx.AsyncClient``. Safe to share + across tasks. Use ``async with`` or call ``aclose()`` when you are done. + """ + + def __init__( + self, + api_key: Optional[str] = None, + base_url: Optional[str] = None, + timeout: float = 10.0, + max_retries: int = 2, + http_client: Optional[httpx.AsyncClient] = None, + ) -> None: + checked_timeout, checked_retries = self._configure(api_key, base_url, timeout, max_retries) + self._transport = AsyncTransport( + client=http_client, + timeout=checked_timeout, + max_retries=checked_retries, + retry_rate_limited=True, + ) + self.history = AsyncHistory(self) + self.identifications = AsyncIdentifications(self) + + async def _search( + self, + lookup_type: str, + value: Union[str, UUID], + limit: int, + offset: int, + *, + timeout: Optional[float] = None, + max_retries: Optional[int] = None, + ) -> HistoryPage: + url = self._history_url(lookup_type, value) + response = await self._transport.get( + url, + headers=self._headers, + params={"limit": validate_limit(limit), "offset": validate_offset(offset)}, + timeout=timeout, + max_retries=max_retries, + ) + return HistoryPage.from_dict(json_object(response)) + + async def aclose(self) -> None: + """Close the underlying HTTP client (unless it was passed in).""" + await self._transport.aclose() + + async def __aenter__(self) -> AsyncShieldLabs: + return self + + async def __aexit__( + self, + exc_type: Optional[type[BaseException]], + exc: Optional[BaseException], + tb: Optional[TracebackType], + ) -> None: + await self.aclose() + + +class AsyncHistory: + """``client.history`` on ``AsyncShieldLabs``.""" + + def __init__(self, client: AsyncShieldLabs) -> None: + self._client = client + + async def search( + self, + type: LookupType, + value: Union[str, UUID], + limit: int = 20, + offset: int = 0, + ) -> HistoryPage: + """Async ``ShieldLabs.history.search``.""" + return await self._client._search(type, value, limit, offset) + + def iter( + self, + type: LookupType, + value: Union[str, UUID], + page_size: int = 100, + max_items: Optional[int] = None, + ) -> AsyncIterator[Identification]: + """Async ``ShieldLabs.history.iter``: use it with ``async for``.""" + validate_lookup(type, value) + size = validate_limit(page_size, "page_size") + if max_items is not None: + max_items = validate_count(max_items, "max_items") + return self._iterate(type, value, size, max_items) + + async def _iterate( + self, + lookup_type: str, + value: Union[str, UUID], + page_size: int, + max_items: Optional[int], + ) -> AsyncIterator[Identification]: + seen: set[str] = set() + offset = 0 + count = 0 + while max_items is None or count < max_items: + page = await self._client._search(lookup_type, value, page_size, offset) + if not page.data: + return + for identification in page.data: + if identification.request_id in seen: + continue + seen.add(identification.request_id) + yield identification + count += 1 + if max_items is not None and count >= max_items: + return + offset += len(page.data) + if offset >= page.total: + return + + +class AsyncIdentifications: + """``client.identifications`` on ``AsyncShieldLabs``.""" + + def __init__(self, client: AsyncShieldLabs) -> None: + self._client = client + + async def get( + self, + request_id: Union[str, UUID], + wait: bool = True, + timeout: float = 10.0, + poll_interval: float = 0.25, + ) -> Optional[Identification]: + """Async ``ShieldLabs.identifications.get``: same waiting rules, errors and result.""" + rid = validate_uuid(request_id, "request_id") + total_timeout = validate_seconds(timeout, "timeout", allow_zero=True) + initial_wait = validate_seconds(poll_interval, "poll_interval", allow_zero=False) + client = self._client + if not wait: + page = await client._search("request_id", rid, 1, 0) + return page.data[0] if page.data else None + + transport = client._transport + plan = _PollPlan(transport._clock, total_timeout, initial_wait, transport.timeout) + while True: + error: Optional[ShieldLabsError] = None + try: + page = await client._search( + "request_id", rid, 1, 0, timeout=plan.attempt_timeout(), max_retries=0 + ) + except _TRANSIENT_ERRORS as exc: + error = exc + else: + if page.data: + return page.data[0] + delay = plan.next_wait(error) + if delay is None: + if error is not None: + raise error + return None + await transport._sleep(delay) diff --git a/src/shieldlabs/_errors.py b/src/shieldlabs/_errors.py new file mode 100644 index 0000000..c8e81d0 --- /dev/null +++ b/src/shieldlabs/_errors.py @@ -0,0 +1,133 @@ +"""Exceptions and warnings raised by the ShieldLabs SDK.""" + +from __future__ import annotations + +from collections.abc import Mapping +from typing import Any, Optional + +__all__ = [ + "APIConnectionError", + "APITimeoutError", + "ApiError", + "AuthenticationError", + "BadRequestError", + "NotFoundError", + "QuotaExceededError", + "RateLimitError", + "ServerError", + "ShieldLabsError", + "ShieldLabsWarning", + "SignatureVerificationError", + "ValidationError", + "WebhookParseError", +] + + +class ShieldLabsError(Exception): + """Base class of every exception raised by this package.""" + + message: str + + def __init__(self, message: str) -> None: + super().__init__(message) + self.message = message + + +class ApiError(ShieldLabsError): + """A ShieldLabs API answered with a status outside 2xx. + + Attributes: + status: HTTP status code. + body: The response body parsed as JSON when possible, else the text, else ``None``. + headers: The response headers (case-insensitive when they come from a response). + """ + + status: int + body: Any + headers: Mapping[str, str] + + def __init__( + self, + message: str, + *, + status: int = 0, + body: Any = None, + headers: Optional[Mapping[str, str]] = None, + ) -> None: + super().__init__(message) + self.status = status + self.body = body + self.headers = headers if headers is not None else {} + + def __str__(self) -> str: + return f"{self.message} (HTTP {self.status})" + + +class BadRequestError(ApiError): + """HTTP 400: the server rejected the request.""" + + +class AuthenticationError(ApiError): + """HTTP 401 or 403: the key is missing, wrong, rotated, or the domain is disabled.""" + + +class QuotaExceededError(ApiError): + """HTTP 402. Neither the History API nor the Management API returns it today. + + An account over its included volume shows a negative ``remaining_identifications`` in the + domain profile instead. + """ + + +class NotFoundError(ApiError): + """HTTP 404: the path does not exist (check the base URL).""" + + +class RateLimitError(ApiError): + """HTTP 429: too many requests. + + Attributes: + retry_after: Seconds to wait before the next request, when the server sent ``Retry-After``. + """ + + retry_after: Optional[float] + + def __init__( + self, + message: str, + *, + status: int = 429, + body: Any = None, + headers: Optional[Mapping[str, str]] = None, + retry_after: Optional[float] = None, + ) -> None: + super().__init__(message, status=status, body=body, headers=headers) + self.retry_after = retry_after + + +class ServerError(ApiError): + """HTTP 5xx: a server or edge proxy error.""" + + +class APIConnectionError(ShieldLabsError): + """The request never produced a response (DNS, TCP, TLS or protocol failure).""" + + +class APITimeoutError(ShieldLabsError): + """The request did not complete within the configured timeout.""" + + +class SignatureVerificationError(ShieldLabsError): + """A webhook delivery did not carry a valid ``X-Shield-Signature`` for the given secret.""" + + +class WebhookParseError(ShieldLabsError): + """A verified webhook body could not be parsed into an event.""" + + +class ValidationError(ShieldLabsError, ValueError): + """An argument is invalid. Raised before any HTTP request is sent.""" + + +class ShieldLabsWarning(UserWarning): + """Non-fatal configuration or data issue detected by the SDK.""" diff --git a/src/shieldlabs/_helpers.py b/src/shieldlabs/_helpers.py new file mode 100644 index 0000000..c5e9bc8 --- /dev/null +++ b/src/shieldlabs/_helpers.py @@ -0,0 +1,152 @@ +"""Policy and identity helpers.""" + +from __future__ import annotations + +import hashlib +import hmac +from collections.abc import Iterable +from dataclasses import dataclass +from datetime import datetime, timedelta, timezone +from typing import Callable, Literal, Optional, Union, get_args + +from ._errors import ValidationError +from ._models import Identification +from ._normalize import FLAG_KEYS, RiskBand, is_rate_limited, risk_band +from ._validation import validate_seconds + +__all__ = ["Evaluation", "EvaluationReason", "evaluate_identification", "user_hid"] + +EvaluationReason = Literal[ + "missing", + "replayed", + "stale", + "rate_limited", + "no_device_signals", + "blocked_flag", + "blocked_band", +] +"""Why ``evaluate_identification`` refused an identification.""" + +_BANDS = frozenset(get_args(RiskBand)) +_FLAGS = frozenset(FLAG_KEYS) + + +@dataclass(frozen=True) +class Evaluation: + """Result of ``evaluate_identification``. + + Attributes: + ok: ``True`` when no check failed. + reason: The first check that failed, or ``None`` when ``ok``. + band: Risk band of the identification (``None`` when it is missing). + flag: The detection flag that blocked it, when ``reason`` is ``"blocked_flag"``. + """ + + ok: bool + reason: Optional[EvaluationReason] + band: Optional[RiskBand] + flag: Optional[str] = None + + +def _names( + values: Union[str, Iterable[str]], allowed: frozenset[str], what: str +) -> tuple[str, ...]: + items = (values,) if isinstance(values, str) else tuple(values) + for item in items: + if item not in allowed: + known = ", ".join(sorted(allowed)) + raise ValidationError(f"unknown {what} {item!r}; expected one of {known}") + return items + + +def _max_age_seconds(max_age: Union[float, timedelta]) -> float: + # A NaN, infinite or negative window would silently disable or invert the freshness check. + if isinstance(max_age, timedelta): + return validate_seconds(max_age.total_seconds(), "max_age", allow_zero=True) + if isinstance(max_age, bool) or not isinstance(max_age, (int, float)): + raise ValidationError("max_age must be a number of seconds or a timedelta") + return validate_seconds(max_age, "max_age", allow_zero=True) + + +def evaluate_identification( + identification: Optional[Identification], + *, + max_age: Union[float, timedelta] = 300.0, + now: Optional[datetime] = None, + block_bands: Iterable[str] = ("dangerous",), + block_flags: Iterable[str] = ("browser_automation", "javascript_disabled"), + is_replay: Optional[Callable[[str], bool]] = None, +) -> Evaluation: + """Apply a starting-point policy to one identification before a protected action. + + Checks run in this order and the first failure wins: + + 1. ``missing``: no identification (unverified, never clean). + 2. ``replayed``: ``is_replay(request_id)`` returned ``True``. The SDK stores nothing; keep + used request IDs in your own store. + 3. ``stale``: ``observed_at`` is older than ``max_age`` (a finite number of seconds, 0 or + greater, or a timedelta). + 4. ``rate_limited``: the Risk Score is the rate-limit marker (above 100). + 5. ``no_device_signals``: the device ID is the all-zero UUID. + 6. ``blocked_flag``: a flag in ``block_flags`` is set (reported in ``flag``). + 7. ``blocked_band``: the risk band is in ``block_bands``. + + The defaults are a starting point: tune ``block_bands``, ``block_flags`` and ``max_age`` + to your product. + + Raises: + ValidationError: An unknown band or flag name, or an invalid ``max_age``. + """ + bands = _names(block_bands, _BANDS, "band") + flags = _names(block_flags, _FLAGS, "detection flag") + max_age_seconds = _max_age_seconds(max_age) + + if identification is None: + return Evaluation(ok=False, reason="missing", band=None) + band = risk_band(identification.risk_score) + if is_replay is not None and is_replay(identification.request_id): + return Evaluation(ok=False, reason="replayed", band=band) + + current = now if now is not None else datetime.now(timezone.utc) + if current.tzinfo is None: + current = current.replace(tzinfo=timezone.utc) + observed = identification.observed_at + if observed is not None and observed.tzinfo is None: + observed = observed.replace(tzinfo=timezone.utc) + if observed is None or (current - observed).total_seconds() > max_age_seconds: + return Evaluation(ok=False, reason="stale", band=band) + + if is_rate_limited(identification.risk_score): + return Evaluation(ok=False, reason="rate_limited", band=band) + if not identification.has_device_signals: + return Evaluation(ok=False, reason="no_device_signals", band=band) + for name in flags: + if getattr(identification.detection_flags, name): + return Evaluation(ok=False, reason="blocked_flag", band=band, flag=name) + if band in bands: + return Evaluation(ok=False, reason="blocked_band", band=band) + return Evaluation(ok=True, reason=None, band=band) + + +def user_hid(user_id: str, secret: Union[str, bytes]) -> str: + """Derive a stable, irreversible User HID for one of your accounts. + + Returns HMAC-SHA256(key = ``secret``, message = ``user_id``) as 64 lowercase hex + characters. Compute it on your server and pass it to the browser agent instead of a raw + email address or account ID. Keep ``secret`` private and stable: changing it changes + every User HID. + + Raises: + ValidationError: ``user_id`` or ``secret`` is empty or not a string. + """ + if not isinstance(user_id, str) or not user_id: + raise ValidationError("user_id must be a non-empty string") + if isinstance(secret, str): + key = secret.encode("utf-8") + elif isinstance(secret, (bytes, bytearray)): + key = bytes(secret) + else: + raise ValidationError("secret must be a string or bytes") + if not key: + raise ValidationError("secret must not be empty") + return hmac.new(key, user_id.encode("utf-8"), hashlib.sha256).hexdigest() diff --git a/src/shieldlabs/_http.py b/src/shieldlabs/_http.py new file mode 100644 index 0000000..32b2c88 --- /dev/null +++ b/src/shieldlabs/_http.py @@ -0,0 +1,341 @@ +"""HTTP transport: headers, retries with backoff, and error mapping.""" + +from __future__ import annotations + +import email.utils +import http +import json +import math +import platform +import random +import re +import time +from collections.abc import Awaitable, Mapping +from datetime import datetime, timezone +from typing import Any, Callable, Optional + +import httpx + +from ._errors import ( + APIConnectionError, + ApiError, + APITimeoutError, + AuthenticationError, + BadRequestError, + NotFoundError, + QuotaExceededError, + RateLimitError, + ServerError, +) +from ._version import __version__ + +try: # anyio ships with httpx; it lets the async client run under asyncio and trio. + from anyio import sleep as _async_sleep +except ImportError: # pragma: no cover + from asyncio import sleep as _async_sleep + +__all__ = [ + "USER_AGENT", + "AsyncTransport", + "SyncTransport", + "backoff_delay", + "error_from_response", + "json_object", + "parse_retry_after", +] + +USER_AGENT = ( + f"shieldlabs-python/{__version__} " + f"({platform.python_implementation()} {platform.python_version()}; httpx {httpx.__version__})" +) + +RETRY_BASE_DELAY = 0.5 +RETRY_MAX_DELAY = 8.0 +RETRY_AFTER_CAP = 10.0 +RATE_LIMIT_MIN_DELAY = 1.0 +"""Shortest wait after a 429, in seconds (the History API limit is a 1-second window): after any +429 inside the wait of ``identifications.get``, and before retrying an ordinary request whose 429 +had no ``Retry-After``. With ``Retry-After`` an ordinary request waits as long as it says, capped +at ``RETRY_AFTER_CAP``.""" + +_STATUS_ERRORS: dict[int, type[ApiError]] = { + 400: BadRequestError, + 401: AuthenticationError, + 402: QuotaExceededError, + 403: AuthenticationError, + 404: NotFoundError, + 429: RateLimitError, +} + +Params = Mapping[str, Any] + + +def backoff_delay(attempt: int, rand: Callable[[], float]) -> float: + """Exponential backoff with jitter: base 0.5 s, factor 2, cap 8 s, times a random 0.5 to 1.""" + delay = min(RETRY_MAX_DELAY, RETRY_BASE_DELAY * (2**attempt)) + return float(delay * (0.5 + 0.5 * rand())) + + +_RETRY_AFTER_SECONDS = re.compile(r"[0-9]+(?:\.[0-9]+)?") + + +def parse_retry_after(value: Optional[str]) -> Optional[float]: + """Parse ``Retry-After`` (seconds or an HTTP date) into seconds, or ``None``. + + Seconds are digits with an optional fraction. Anything else that is not an HTTP date, such as + a sign, an exponent or an underscore, counts as no ``Retry-After``. + """ + if value is None: + return None + text = value.strip() + if not text: + return None + if _RETRY_AFTER_SECONDS.fullmatch(text): + seconds = float(text) + else: + try: + moment = email.utils.parsedate_to_datetime(text) + except (TypeError, ValueError, IndexError): + return None + if moment.tzinfo is None: + moment = moment.replace(tzinfo=timezone.utc) + seconds = (moment - datetime.now(timezone.utc)).total_seconds() + if math.isnan(seconds) or math.isinf(seconds): + return None + return max(0.0, seconds) + + +def _describe_body(content: bytes) -> tuple[Any, Optional[str]]: + """Return ``(parsed body, message)`` without ever raising.""" + try: + text = content.decode("utf-8", errors="replace").strip() + if not text: + return None, None + try: + parsed = json.loads(text) + except (ValueError, RecursionError): + if text.startswith("<"): + return text, None + return text, text.splitlines()[0][:200] + if isinstance(parsed, dict): + for key in ("error", "message"): + detail = parsed.get(key) + if isinstance(detail, str) and detail.strip(): + return parsed, detail.strip()[:500] + return parsed, None + if isinstance(parsed, str) and parsed.strip(): + return parsed, parsed.strip()[:500] + return parsed, None + except Exception: # pragma: no cover - defensive: building an error must never fail + return None, None + + +def _reason(status: int) -> str: + try: + return http.HTTPStatus(status).phrase + except ValueError: + return "Unexpected response" + + +def error_from_response(response: httpx.Response) -> ApiError: + """Map a non-2xx response to the matching ``ApiError`` subclass.""" + status = response.status_code + body, detail = _describe_body(response.content) + message = detail or _reason(status) + if status == 429: + return RateLimitError( + message, + status=status, + body=body, + headers=response.headers, + retry_after=parse_retry_after(response.headers.get("retry-after")), + ) + cls = _STATUS_ERRORS.get(status) or (ServerError if 500 <= status <= 599 else ApiError) + return cls(message, status=status, body=body, headers=response.headers) + + +def json_object(response: httpx.Response) -> Mapping[str, Any]: + """Decode a 2xx response body that must be a JSON object.""" + try: + parsed = json.loads(response.content) + except (ValueError, RecursionError): + raise ApiError( + "Response body is not valid JSON", + status=response.status_code, + body=response.content.decode("utf-8", errors="replace"), + headers=response.headers, + ) from None + if not isinstance(parsed, dict): + raise ApiError( + "Response body is not a JSON object", + status=response.status_code, + body=parsed, + headers=response.headers, + ) + return parsed + + +class _TransportBase: + """Retry policy shared by the sync and async transports.""" + + def __init__( + self, + *, + timeout: float, + max_retries: int, + retry_rate_limited: bool, + ) -> None: + self.timeout = timeout + self.max_retries = max_retries + self.retry_rate_limited = retry_rate_limited + self._random: Callable[[], float] = random.random + self._clock: Callable[[], float] = time.monotonic + + def _retryable(self, status: int) -> bool: + if status == 429: + return self.retry_rate_limited + return 500 <= status <= 599 + + def _delay_for(self, attempt: int, error: ApiError) -> float: + retry_after = parse_retry_after(error.headers.get("retry-after")) + if retry_after is not None: + return min(retry_after, RETRY_AFTER_CAP) + delay = backoff_delay(attempt, self._random) + if error.status == 429: + # A retry inside the same 1-second window would only be refused again. + delay = max(delay, RATE_LIMIT_MIN_DELAY) + return delay + + +def _timeout_error(exc: httpx.TimeoutException, timeout: float) -> APITimeoutError: + kind = type(exc).__name__ + return APITimeoutError(f"Request to ShieldLabs timed out after {timeout:g} s ({kind})") + + +def _connection_error(exc: httpx.RequestError) -> APIConnectionError: + detail = str(exc) or type(exc).__name__ + return APIConnectionError(f"Could not reach ShieldLabs: {detail}") + + +class SyncTransport(_TransportBase): + """Blocking transport over ``httpx.Client``. Safe to share across threads.""" + + def __init__( + self, + *, + client: Optional[httpx.Client], + timeout: float, + max_retries: int, + retry_rate_limited: bool, + ) -> None: + super().__init__( + timeout=timeout, max_retries=max_retries, retry_rate_limited=retry_rate_limited + ) + self.client = client if client is not None else httpx.Client() + self.owns_client = client is None + self._sleep: Callable[[float], None] = time.sleep + + def get( + self, + url: str, + *, + headers: Mapping[str, str], + params: Optional[Params] = None, + timeout: Optional[float] = None, + max_retries: Optional[int] = None, + ) -> httpx.Response: + """GET with retries. Returns a 2xx response or raises a ``ShieldLabsError``.""" + attempt_timeout = self.timeout if timeout is None else timeout + retries = self.max_retries if max_retries is None else max_retries + attempt = 0 + while True: + try: + response = self.client.get( + url, headers=dict(headers), params=params, timeout=attempt_timeout + ) + except httpx.TimeoutException as exc: + if attempt < retries: + self._sleep(backoff_delay(attempt, self._random)) + attempt += 1 + continue + raise _timeout_error(exc, attempt_timeout) from exc + except httpx.RequestError as exc: + if attempt < retries: + self._sleep(backoff_delay(attempt, self._random)) + attempt += 1 + continue + raise _connection_error(exc) from exc + if response.is_success: + return response + error = error_from_response(response) + if attempt < retries and self._retryable(response.status_code): + self._sleep(self._delay_for(attempt, error)) + attempt += 1 + continue + raise error + + def close(self) -> None: + if self.owns_client: + self.client.close() + + +class AsyncTransport(_TransportBase): + """Asynchronous transport over ``httpx.AsyncClient``. Safe for concurrent tasks.""" + + def __init__( + self, + *, + client: Optional[httpx.AsyncClient], + timeout: float, + max_retries: int, + retry_rate_limited: bool, + ) -> None: + super().__init__( + timeout=timeout, max_retries=max_retries, retry_rate_limited=retry_rate_limited + ) + self.client = client if client is not None else httpx.AsyncClient() + self.owns_client = client is None + self._sleep: Callable[[float], Awaitable[None]] = _async_sleep + + async def get( + self, + url: str, + *, + headers: Mapping[str, str], + params: Optional[Params] = None, + timeout: Optional[float] = None, + max_retries: Optional[int] = None, + ) -> httpx.Response: + """GET with retries. Returns a 2xx response or raises a ``ShieldLabsError``.""" + attempt_timeout = self.timeout if timeout is None else timeout + retries = self.max_retries if max_retries is None else max_retries + attempt = 0 + while True: + try: + response = await self.client.get( + url, headers=dict(headers), params=params, timeout=attempt_timeout + ) + except httpx.TimeoutException as exc: + if attempt < retries: + await self._sleep(backoff_delay(attempt, self._random)) + attempt += 1 + continue + raise _timeout_error(exc, attempt_timeout) from exc + except httpx.RequestError as exc: + if attempt < retries: + await self._sleep(backoff_delay(attempt, self._random)) + attempt += 1 + continue + raise _connection_error(exc) from exc + if response.is_success: + return response + error = error_from_response(response) + if attempt < retries and self._retryable(response.status_code): + await self._sleep(self._delay_for(attempt, error)) + attempt += 1 + continue + raise error + + async def aclose(self) -> None: + if self.owns_client: + await self.client.aclose() diff --git a/src/shieldlabs/_management.py b/src/shieldlabs/_management.py new file mode 100644 index 0000000..37cb4f6 --- /dev/null +++ b/src/shieldlabs/_management.py @@ -0,0 +1,169 @@ +"""Management API clients (sync and async).""" + +from __future__ import annotations + +from types import TracebackType +from typing import Optional + +import httpx + +from ._http import USER_AGENT, AsyncTransport, SyncTransport, json_object +from ._models import DomainProfile +from ._validation import ( + management_origin, + normalize_domain, + require_secret, + validate_max_retries, + validate_seconds, +) + +__all__ = ["AsyncShieldLabsManagement", "ShieldLabsManagement"] + +_PROFILE_PATH = "/v1/profile" + + +class _ManagementConfig: + """Configuration shared by the sync and async Management clients.""" + + base_url: str + """Management API origin, for example ``https://api.shieldlabs.ai``.""" + domain: str + """The registered domain, normalized (sent as ``X-Shield-Domain``).""" + + def _configure( + self, + secret_key: Optional[str], + domain: Optional[str], + base_url: Optional[str], + timeout: float, + max_retries: int, + ) -> tuple[float, int]: + secret = require_secret(secret_key, "SHIELDLABS_SECRET_KEY", "secret_key") + self.domain = normalize_domain(domain) + self.base_url = management_origin(base_url) + self._headers = { + "X-Shield-Domain": self.domain, + "Authorization": f"Bearer {secret}", + "Accept": "application/json", + "User-Agent": USER_AGENT, + } + return ( + validate_seconds(timeout, "timeout", allow_zero=False), + validate_max_retries(max_retries), + ) + + def __repr__(self) -> str: + return f"{type(self).__name__}(domain={self.domain!r}, base_url={self.base_url!r})" + + +class ShieldLabsManagement(_ManagementConfig): + """Client for the ShieldLabs Management API (domain profile, Secret Key + domain). + + The Management API allows about 15 requests per minute per caller IP and then blocks the + IP for 10 minutes, so this client never retries a 429 (``RateLimitError`` is raised at + once). Call it sparingly and cache the profile. + + Args: + secret_key: Secret Key of the domain. Defaults to ``SHIELDLABS_SECRET_KEY``. + domain: Registered domain. Defaults to ``SHIELDLABS_DOMAIN``. Normalized before use + (``https://www.Example.com/`` becomes ``example.com``). + base_url: Management API origin. Defaults to ``SHIELDLABS_MANAGEMENT_BASE_URL`` or + ``https://api.shieldlabs.ai``. Must use https; plain http is accepted only for + localhost, 127.0.0.1 and ::1 (local test servers). + timeout: Timeout of one HTTP attempt, in seconds. + max_retries: Retries for connection errors, timeouts and 5xx responses (never 429). + http_client: Your own ``httpx.Client``. It is not closed by ``close()``. + """ + + def __init__( + self, + secret_key: Optional[str] = None, + domain: Optional[str] = None, + base_url: Optional[str] = None, + timeout: float = 10.0, + max_retries: int = 2, + http_client: Optional[httpx.Client] = None, + ) -> None: + checked_timeout, checked_retries = self._configure( + secret_key, domain, base_url, timeout, max_retries + ) + self._transport = SyncTransport( + client=http_client, + timeout=checked_timeout, + max_retries=checked_retries, + retry_rate_limited=False, + ) + + def get_profile(self) -> DomainProfile: + """Read the domain profile (``GET /v1/profile``). + + Raises: + AuthenticationError: Wrong Secret Key, unknown or disabled domain. + RateLimitError: Per-IP limit reached; the IP stays blocked for 10 minutes. + """ + response = self._transport.get(f"{self.base_url}{_PROFILE_PATH}", headers=self._headers) + return DomainProfile.from_dict(json_object(response)) + + def close(self) -> None: + """Close the underlying HTTP client (unless it was passed in).""" + self._transport.close() + + def __enter__(self) -> ShieldLabsManagement: + return self + + def __exit__( + self, + exc_type: Optional[type[BaseException]], + exc: Optional[BaseException], + tb: Optional[TracebackType], + ) -> None: + self.close() + + +class AsyncShieldLabsManagement(_ManagementConfig): + """Asynchronous client for the ShieldLabs Management API. + + Same arguments and rate-limit rules as ``ShieldLabsManagement``; ``http_client`` is an + ``httpx.AsyncClient``. + """ + + def __init__( + self, + secret_key: Optional[str] = None, + domain: Optional[str] = None, + base_url: Optional[str] = None, + timeout: float = 10.0, + max_retries: int = 2, + http_client: Optional[httpx.AsyncClient] = None, + ) -> None: + checked_timeout, checked_retries = self._configure( + secret_key, domain, base_url, timeout, max_retries + ) + self._transport = AsyncTransport( + client=http_client, + timeout=checked_timeout, + max_retries=checked_retries, + retry_rate_limited=False, + ) + + async def get_profile(self) -> DomainProfile: + """Async ``ShieldLabsManagement.get_profile``.""" + response = await self._transport.get( + f"{self.base_url}{_PROFILE_PATH}", headers=self._headers + ) + return DomainProfile.from_dict(json_object(response)) + + async def aclose(self) -> None: + """Close the underlying HTTP client (unless it was passed in).""" + await self._transport.aclose() + + async def __aenter__(self) -> AsyncShieldLabsManagement: + return self + + async def __aexit__( + self, + exc_type: Optional[type[BaseException]], + exc: Optional[BaseException], + tb: Optional[TracebackType], + ) -> None: + await self.aclose() diff --git a/src/shieldlabs/_models.py b/src/shieldlabs/_models.py new file mode 100644 index 0000000..74172d2 --- /dev/null +++ b/src/shieldlabs/_models.py @@ -0,0 +1,486 @@ +"""Typed models returned by the SDK.""" + +from __future__ import annotations + +import json +from collections.abc import Mapping +from dataclasses import dataclass, field +from datetime import datetime +from typing import Any, Literal, Optional + +from ._normalize import ( + FLAG_KEYS, + HISTORY_FLAG_MAP, + NIL_UUID, + TRAFFIC_KEYS, + RiskBand, + as_int, + as_str, + clean_ip, + format_timestamp, + is_rate_limited, + parse_history_time, + parse_rfc3339, + risk_band, + signal_slug, +) + +__all__ = [ + "DetectionFlags", + "DomainProfile", + "HistoryPage", + "Identification", + "IdentificationSource", + "IpInfo", + "Signal", + "SignalName", + "TrafficSource", +] + +IdentificationSource = Literal["webhook", "history"] +"""Where an ``Identification`` came from.""" + +_IP_MISMATCH_DETAIL = "IP ≠ leakIP" + + +def _mapping(value: object) -> Mapping[str, Any]: + return value if isinstance(value, Mapping) else {} + + +def _optional_user_hid(value: object) -> Optional[str]: + # "" becomes None; sentinels such as "anonymous", "fail", "-1" and "unknown" are kept. + return value if isinstance(value, str) and value != "" else None + + +class SignalName: + """Known risk signal names (``Signal.name``). + + The set of names is open: new names can appear at any time, so compare against these + constants but never reject a name that is not listed. Branch on ``risk_score`` and + ``detection_flags``; use signal names for display and logging. + """ + + TOR = "tor" + JAVASCRIPT_DISABLED = "javascript_disabled" + OS_MISMATCH = "os_mismatch" + ANTIDETECT_BROWSER = "antidetect_browser" + PROXY_ROUTED_ANTIDETECT = "proxy_routed_antidetect" + PORT_SCAN_ROUTED_VIA_PROXY = "port_scan_routed_via_proxy" + BROWSER_AUTOMATION = "browser_automation" + STUN_NOT_CHECKED = "stun_not_checked" + STUN_LATE_CORRECTION = "stun_late_correction" + OS_NOT_DETECTED = "os_not_detected" + BROWSER_VPN_PROXY = "browser_vpn_proxy" + VPN = "vpn" + PRIVACY_RELAY = "privacy_relay" + PROXY = "proxy" + DATACENTER_IP = "datacenter_ip" + ABUSER = "abuser" + TIMEZONE_MISMATCH = "timezone_mismatch" + RATE_LIMITED = "rate_limited" + + +@dataclass(frozen=True) +class IpInfo: + """An IP address with the English name of its country. Both are ``""`` when unknown.""" + + ip: str = "" + country: str = "" + + @classmethod + def from_dict(cls, data: Mapping[str, Any]) -> IpInfo: + """Build from ``{"ip": ..., "country": ...}``. ``0.0.0.0`` becomes ``""``.""" + data = _mapping(data) + return cls(ip=clean_ip(data.get("ip")), country=as_str(data.get("country"))) + + def to_dict(self) -> dict[str, Any]: + return {"ip": self.ip, "country": self.country} + + +@dataclass(frozen=True) +class TrafficSource: + """Attribution of the visit. Every field is ``""`` when absent.""" + + channel: str = "" + referrer_domain: str = "" + landing_url: str = "" + click_id_type: str = "" + utm_source: str = "" + utm_medium: str = "" + utm_campaign: str = "" + utm_content: str = "" + utm_term: str = "" + + @classmethod + def from_dict(cls, data: Mapping[str, Any]) -> TrafficSource: + """Build from a webhook ``traffic_source`` object. Missing keys become ``""``.""" + data = _mapping(data) + return cls(**{key: as_str(data.get(key)) for key in TRAFFIC_KEYS}) + + def to_dict(self) -> dict[str, Any]: + return {key: getattr(self, key) for key in TRAFFIC_KEYS} + + +@dataclass(frozen=True) +class Signal: + """One weighted risk signal behind the score. + + Attributes: + name: Signal name, an open set (see ``SignalName`` for known values). Names can repeat. + weight: Weight in points. Can be negative (corrections) or informational. + description: Human-readable detail. Present on identifications read from the History + API, ``None`` on webhook deliveries. + """ + + name: str + weight: int + description: Optional[str] = None + + @classmethod + def from_dict(cls, data: Mapping[str, Any]) -> Signal: + """Build from ``{"name": ..., "weight": ..., "description"?: ...}``.""" + data = _mapping(data) + description = data.get("description") + return cls( + name=as_str(data.get("name")), + weight=as_int(data.get("weight")), + description=description if isinstance(description, str) else None, + ) + + def to_dict(self) -> dict[str, Any]: + return {"name": self.name, "weight": self.weight, "description": self.description} + + +@dataclass(frozen=True) +class DetectionFlags: + """The 19 detection flags. Branch on these and on ``risk_score``.""" + + vpn: bool = False + privacy_relay: bool = False + browser_vpn_proxy: bool = False + tor: bool = False + proxy: bool = False + datacenter_ip: bool = False + abuser: bool = False + os_mismatch: bool = False + os_not_detected: bool = False + timezone_mismatch: bool = False + anti_detect_browser: bool = False + browser_automation: bool = False + ip_mismatch: bool = False + incognito: bool = False + search_bot: bool = False + suspicious_paid_click: bool = False + javascript_disabled: bool = False + stun_not_checked: bool = False + check_incomplete: bool = False + + @classmethod + def from_dict(cls, data: Mapping[str, Any]) -> DetectionFlags: + """Build from a webhook ``detection_flags`` object. A missing key is ``False``.""" + data = _mapping(data) + return cls(**{key: bool(data.get(key, False)) for key in FLAG_KEYS}) + + def to_dict(self) -> dict[str, bool]: + return {key: getattr(self, key) for key in FLAG_KEYS} + + def active(self) -> tuple[str, ...]: + """Names of the flags that are set, in wire order.""" + return tuple(key for key in FLAG_KEYS if getattr(self, key)) + + +@dataclass(frozen=True) +class Identification: + """One identification (one run of the browser agent), normalized. + + Webhook deliveries and History API rows produce the same model. Field names follow the + webhook contract. + """ + + request_id: str + visitor_id: str + device_id: str + session_id: str + cookie_id: str + user_hid: Optional[str] + domain: str + public_ip: IpInfo + local_ip: IpInfo + connection_type: str + os: str + browser: str + device_type: str + traffic_source: TrafficSource + risk_score: int + signals: tuple[Signal, ...] + detection_flags: DetectionFlags + observed_at: Optional[datetime] + """When the identification was observed, as an aware UTC datetime. ``None`` only when the + server value could not be parsed.""" + source: IdentificationSource + raw: Mapping[str, Any] = field(default_factory=dict, compare=False, repr=False) + """The original webhook ``data`` object or History row, including fields the model omits.""" + + @property + def risk_band(self) -> RiskBand: + """Band of ``risk_score``: trusted, suspicious, dangerous, or rate_limited (marker).""" + return risk_band(self.risk_score) + + @property + def is_rate_limited(self) -> bool: + """``True`` when ``risk_score`` is the rate-limit marker (above 100).""" + return is_rate_limited(self.risk_score) + + @property + def has_device_signals(self) -> bool: + """``False`` when the device ID is the all-zero UUID (no usable device signals).""" + return self.device_id not in ("", NIL_UUID) + + @classmethod + def from_webhook_data(cls, data: Mapping[str, Any]) -> Identification: + """Normalize the ``data`` object of an ``identification.scored`` webhook.""" + data = _mapping(data) + flags = _mapping(data.get("detection_flags")) + raw_signals = data.get("signals") + signals = tuple( + Signal(name=as_str(item.get("name")), weight=as_int(item.get("weight"))) + for item in (raw_signals if isinstance(raw_signals, list) else []) + if isinstance(item, Mapping) + ) + return cls( + request_id=as_str(data.get("request_id")), + visitor_id=as_str(data.get("visitor_id")), + device_id=as_str(data.get("device_id")), + session_id=as_str(data.get("session_id")), + cookie_id=as_str(data.get("cookie_id")), + user_hid=_optional_user_hid(data.get("user_hid")), + domain=as_str(data.get("domain")), + public_ip=IpInfo.from_dict(_mapping(data.get("public_ip"))), + local_ip=IpInfo.from_dict(_mapping(data.get("local_ip"))), + connection_type=as_str(data.get("connection_type")), + os=as_str(data.get("os")), + browser=as_str(data.get("browser")), + device_type=as_str(data.get("device_type")), + traffic_source=TrafficSource.from_dict(_mapping(data.get("traffic_source"))), + risk_score=as_int(data.get("risk_score")), + signals=signals, + detection_flags=DetectionFlags.from_dict(flags), + observed_at=parse_rfc3339(data.get("observed_at")), + source="webhook", + raw=dict(data), + ) + + @classmethod + def from_history_row(cls, row: Mapping[str, Any]) -> Identification: + """Normalize one row of a History API response.""" + row = _mapping(row) + leak_source = as_str(row.get("webrtc_leak_source")).strip() + if leak_source and leak_source != "none": + local_ip = clean_ip(row.get("webrtc_leak_ip")) + local_country = as_str(row.get("webrtc_leak_country")) + else: + local_ip = clean_ip(row.get("web_rtc_ip")) + local_country = as_str(row.get("web_rtc_country")) + public_ip = clean_ip(row.get("ip")) + + details: Any = [] + score_details = row.get("score_details") + if isinstance(score_details, str) and score_details: + try: + details = json.loads(score_details) + except (ValueError, RecursionError): + details = [] + if not isinstance(details, list): + details = [] + + signals: list[Signal] = [] + ip_leak_detail = False + for detail in details: + if not isinstance(detail, Mapping): + continue + description = as_str(detail.get("Description")) + if description.startswith(_IP_MISMATCH_DETAIL): + ip_leak_detail = True + value = detail.get("Value", 0) + if isinstance(value, bool) or not isinstance(value, int) or value == 0: + continue + signals.append( + Signal(name=signal_slug(description), weight=value, description=description) + ) + + search_bot = bool(row.get("is_search_bot", False)) + flags: dict[str, bool] = {} + for key in FLAG_KEYS: + if key == "browser_vpn_proxy": + flags[key] = row.get("connection_type") == "browser_vpn_proxy" + elif key == "ip_mismatch": + differs = public_ip != "" and local_ip != "" and public_ip != local_ip + flags[key] = (not search_bot) and (ip_leak_detail or differs) + else: + flags[key] = bool(row.get(HISTORY_FLAG_MAP[key], False)) + + return cls( + request_id=as_str(row.get("request_id")), + visitor_id=as_str(row.get("visitor_id")), + device_id=as_str(row.get("device_id")), + session_id=as_str(row.get("session_id")), + cookie_id=as_str(row.get("cookie_id")), + user_hid=_optional_user_hid(row.get("user_hid")), + domain=as_str(row.get("site_domain")) or as_str(row.get("domain")), + public_ip=IpInfo(ip=public_ip, country=as_str(row.get("country"))), + local_ip=IpInfo(ip=local_ip, country=local_country), + connection_type=as_str(row.get("connection_type")), + os=as_str(row.get("os")), + browser=as_str(row.get("browser")), + device_type=as_str(row.get("device_type")), + traffic_source=TrafficSource( + channel=as_str(row.get("traffic_channel")), + referrer_domain=as_str(row.get("referrer_domain")), + landing_url=as_str(row.get("entry_url")), + click_id_type=as_str(row.get("click_id_type")), + utm_source=as_str(row.get("utm_source")), + utm_medium=as_str(row.get("utm_medium")), + utm_campaign=as_str(row.get("utm_campaign")), + utm_content=as_str(row.get("utm_content")), + utm_term=as_str(row.get("utm_term")), + ), + risk_score=as_int(row.get("score")), + signals=tuple(signals), + detection_flags=DetectionFlags(**flags), + observed_at=parse_history_time(row.get("created_at")), + source="history", + raw=dict(row), + ) + + @classmethod + def from_dict(cls, data: Mapping[str, Any]) -> Identification: + """Rebuild an identification from the output of ``to_dict()``. + + Webhook ``data`` objects are accepted as well (use ``from_webhook_data`` for those). + """ + data = _mapping(data) + raw_signals = data.get("signals") + observed_at = data.get("observed_at") + source = data.get("source") + return cls( + request_id=as_str(data.get("request_id")), + visitor_id=as_str(data.get("visitor_id")), + device_id=as_str(data.get("device_id")), + session_id=as_str(data.get("session_id")), + cookie_id=as_str(data.get("cookie_id")), + user_hid=_optional_user_hid(data.get("user_hid")), + domain=as_str(data.get("domain")), + public_ip=IpInfo.from_dict(_mapping(data.get("public_ip"))), + local_ip=IpInfo.from_dict(_mapping(data.get("local_ip"))), + connection_type=as_str(data.get("connection_type")), + os=as_str(data.get("os")), + browser=as_str(data.get("browser")), + device_type=as_str(data.get("device_type")), + traffic_source=TrafficSource.from_dict(_mapping(data.get("traffic_source"))), + risk_score=as_int(data.get("risk_score")), + signals=tuple( + Signal.from_dict(item) + for item in (raw_signals if isinstance(raw_signals, list) else []) + if isinstance(item, Mapping) + ), + detection_flags=DetectionFlags.from_dict(_mapping(data.get("detection_flags"))), + observed_at=observed_at + if isinstance(observed_at, datetime) + else parse_rfc3339(observed_at), + source="history" if source == "history" else "webhook", + raw=dict(data), + ) + + def to_dict(self) -> dict[str, Any]: + """Plain JSON-ready dict (``observed_at`` as RFC 3339 UTC with milliseconds; no ``raw``).""" + return { + "request_id": self.request_id, + "visitor_id": self.visitor_id, + "device_id": self.device_id, + "session_id": self.session_id, + "cookie_id": self.cookie_id, + "user_hid": self.user_hid, + "domain": self.domain, + "public_ip": self.public_ip.to_dict(), + "local_ip": self.local_ip.to_dict(), + "connection_type": self.connection_type, + "os": self.os, + "browser": self.browser, + "device_type": self.device_type, + "traffic_source": self.traffic_source.to_dict(), + "risk_score": self.risk_score, + "signals": [signal.to_dict() for signal in self.signals], + "detection_flags": self.detection_flags.to_dict(), + "observed_at": format_timestamp(self.observed_at), + "source": self.source, + } + + +@dataclass(frozen=True) +class HistoryPage: + """One page of History API results, newest first. + + Attributes: + data: The identifications on this page. + total: Number of identifications that match the lookup, across all pages. + """ + + data: tuple[Identification, ...] + total: int + + @classmethod + def from_dict(cls, body: Mapping[str, Any]) -> HistoryPage: + """Build from a History API response body ``{"data": [...], "total": N}``.""" + body = _mapping(body) + rows = body.get("data") + data = tuple( + Identification.from_history_row(row) + for row in (rows if isinstance(rows, list) else []) + if isinstance(row, Mapping) + ) + return cls(data=data, total=as_int(body.get("total"), default=len(data))) + + +@dataclass(frozen=True) +class DomainProfile: + """Management API profile of one registered domain. + + Attributes: + domain: The registered domain. + remaining_identifications: Included identifications left on the account. Can be + negative when the account is over its included volume. + public_key_masked: Public Key with every character except the last 4 replaced by ``*``. + secret_key_masked: Secret Key masked the same way. + created_at: When the domain was added (aware UTC datetime), ``None`` if unparsable. + raw: The original response body. + """ + + domain: str + remaining_identifications: int + public_key_masked: str + secret_key_masked: str + created_at: Optional[datetime] + raw: Mapping[str, Any] = field(default_factory=dict, compare=False, repr=False) + + @classmethod + def from_dict(cls, body: Mapping[str, Any]) -> DomainProfile: + """Build from a ``GET /v1/profile`` response body.""" + body = _mapping(body) + return cls( + domain=as_str(body.get("Domain")), + remaining_identifications=as_int(body.get("Weight")), + public_key_masked=as_str(body.get("PublicKey")), + secret_key_masked=as_str(body.get("Secret")), + created_at=parse_rfc3339(body.get("CreatedAt")), + raw=dict(body), + ) + + def to_dict(self) -> dict[str, Any]: + """Plain JSON-ready dict (``created_at`` as RFC 3339 UTC with milliseconds; no ``raw``).""" + return { + "domain": self.domain, + "remaining_identifications": self.remaining_identifications, + "public_key_masked": self.public_key_masked, + "secret_key_masked": self.secret_key_masked, + "created_at": format_timestamp(self.created_at), + } diff --git a/src/shieldlabs/_normalize.py b/src/shieldlabs/_normalize.py new file mode 100644 index 0000000..a77d68d --- /dev/null +++ b/src/shieldlabs/_normalize.py @@ -0,0 +1,282 @@ +"""Shared normalization rules. + +History API rows and webhook ``data`` objects describe the same identification with different +field names. The helpers here implement the normalization rules that every ShieldLabs server SDK +follows (signal names, timestamps, IP sentinels, risk bands), so that both sources produce the +same ``Identification``. +""" + +from __future__ import annotations + +import re +import unicodedata +from datetime import datetime, timedelta, timezone +from typing import Literal, Optional + +__all__ = [ + "FLAG_KEYS", + "HISTORY_FLAG_MAP", + "NIL_UUID", + "TRAFFIC_KEYS", + "RiskBand", + "as_int", + "as_str", + "clean_ip", + "fallback_slug", + "format_timestamp", + "is_rate_limited", + "parse_history_time", + "parse_rfc3339", + "risk_band", + "signal_slug", +] + +RiskBand = Literal["trusted", "suspicious", "dangerous", "rate_limited"] +"""Client-side label of a Risk Score. ``rate_limited`` marks the 999 marker, not a band.""" + +NIL_UUID = "00000000-0000-0000-0000-000000000000" +"""The all-zero UUID. As a device ID it means "no usable device signals".""" + +FLAG_KEYS: tuple[str, ...] = ( + "vpn", + "privacy_relay", + "browser_vpn_proxy", + "tor", + "proxy", + "datacenter_ip", + "abuser", + "os_mismatch", + "os_not_detected", + "timezone_mismatch", + "anti_detect_browser", + "browser_automation", + "ip_mismatch", + "incognito", + "search_bot", + "suspicious_paid_click", + "javascript_disabled", + "stun_not_checked", + "check_incomplete", +) +"""The 19 detection flags, in wire order.""" + +HISTORY_FLAG_MAP: dict[str, str] = { + "vpn": "is_vpn", + "privacy_relay": "is_privacy_relay", + "tor": "is_tor", + "proxy": "is_proxy", + "datacenter_ip": "is_datacenter", + "abuser": "is_abuser", + "os_mismatch": "is_os_mismatch", + "os_not_detected": "is_os_not_detected", + "timezone_mismatch": "is_timezone_mismatch", + "anti_detect_browser": "is_antidetect", + "browser_automation": "is_browser_automation", + "incognito": "is_incognito", + "search_bot": "is_search_bot", + "suspicious_paid_click": "is_suspicious_paid_click", + "javascript_disabled": "is_js_disabled", + "stun_not_checked": "is_stun_not_checked", + "check_incomplete": "check_incomplete", +} +"""Detection flag to History row column (``browser_vpn_proxy`` and ``ip_mismatch`` are derived).""" + +TRAFFIC_KEYS: tuple[str, ...] = ( + "channel", + "referrer_domain", + "landing_url", + "click_id_type", + "utm_source", + "utm_medium", + "utm_campaign", + "utm_content", + "utm_term", +) + +_EXACT_SLUGS: dict[str, str] = { + "Is tor": "tor", + "Is VPN": "vpn", + "Is privacy relay": "privacy_relay", + "Is proxy": "proxy", + "Is datacenter": "datacenter_ip", + "Is abuser": "abuser", + "Stun is not checked": "stun_not_checked", + "Stun passed (late arrival, corrected)": "stun_late_correction", + "UA OS is not detected": "os_not_detected", + "Network OS is not detected": "os_not_detected", + "Browser timezone ≠ IP-timezone": "timezone_mismatch", + "Browser VPN/Proxy": "browser_vpn_proxy", + "Browser Automation": "browser_automation", + "Port scan routed via proxy (antidetect browser pattern)": "proxy_routed_antidetect", + "User has been banned 1H, to many requests": "rate_limited", +} + +_PREFIX_SLUGS: tuple[tuple[str, str], ...] = ( + ("Antidetect browser", "antidetect_browser"), + ("Os_mismatch", "os_mismatch"), + ("OS mismatch2", "os_mismatch2"), + ("TCP handshake", "tcp_handshake_v2"), + ("Latency test", "ws_tcp_latency"), + ("JavaScript disabled", "javascript_disabled"), +) + +_STICKY_PREFIX = "Sticky verdict: " +_NOT_EQUAL = "≠" +_SEPARATORS = (" ", "-", "/") + + +def _is_letter_or_digit(ch: str) -> bool: + category = unicodedata.category(ch) + return category.startswith("L") or category == "Nd" + + +def fallback_slug(description: str) -> str: + """Slugify a free-text description: the text before the first ``(``, lowercased.""" + text = description.strip() + paren = text.find("(") + if paren >= 0: + text = text[:paren].strip() + out: list[str] = [] + prev_sep = False + for ch in text: + if ch in _SEPARATORS: + if not prev_sep and out: + out.append("_") + prev_sep = True + elif ch == _NOT_EQUAL: + out.append("_neq_") + prev_sep = False + elif _is_letter_or_digit(ch): + out.append(ch.lower()) + prev_sep = False + slug = "".join(out).strip("_") + return slug or "unknown" + + +def signal_slug(description: str) -> str: + """Map a History ``score_details`` description to the risk signal name used on webhooks.""" + exact = _EXACT_SLUGS.get(description) + if exact is not None: + return exact + for prefix, slug in _PREFIX_SLUGS: + if description.startswith(prefix): + return slug + if description.startswith(_STICKY_PREFIX): + colon = description.find(":") + rest = description[colon + 1 :].strip() + return fallback_slug(rest) + return fallback_slug(description) + + +_HISTORY_TIME = re.compile( + r"^([0-9]{4})-([0-9]{2})-([0-9]{2})[ T]([0-9]{2}):([0-9]{2}):([0-9]{2})" + r"(?:\.([0-9]{1,9}))?(Z|[+-][0-9]{2}:?[0-9]{2})?$" +) +_RFC3339_TIME = re.compile( + r"^([0-9]{4})-([0-9]{2})-([0-9]{2})T([0-9]{2}):([0-9]{2}):([0-9]{2})" + r"(?:\.([0-9]{1,9}))?(Z|[+-][0-9]{2}:[0-9]{2})$" +) + + +def _build_utc(match: re.Match[str]) -> Optional[datetime]: + year, month, day, hour, minute, second = (int(match.group(i)) for i in range(1, 7)) + fraction = match.group(7) or "" + microsecond = int(fraction[:6].ljust(6, "0")) + try: + return datetime(year, month, day, hour, minute, second, microsecond, tzinfo=timezone.utc) + except ValueError: + return None + + +def parse_history_time(value: object) -> Optional[datetime]: + """Parse History ``created_at`` (``YYYY-MM-DD HH:MM:SS[.fff]``, UTC) into an aware datetime. + + The value carries no zone; it is always UTC. A trailing designator, if any, is ignored. + Returns ``None`` when the value is empty or not in that format. + """ + text = value.strip() if isinstance(value, str) else "" + if not text: + return None + match = _HISTORY_TIME.match(text) + if match is None: + return None + return _build_utc(match) + + +def parse_rfc3339(value: object) -> Optional[datetime]: + """Parse an RFC 3339 timestamp (up to 9 fractional digits) into an aware UTC datetime. + + Sub-microsecond digits are truncated. Returns ``None`` when the value is not RFC 3339. + """ + if not isinstance(value, str): + return None + match = _RFC3339_TIME.match(value) + if match is None: + return None + moment = _build_utc(match) + if moment is None: + return None + designator = match.group(8) + if designator == "Z": + return moment + sign = 1 if designator[0] == "+" else -1 + offset = timedelta(hours=int(designator[1:3]), minutes=int(designator[4:6])) + try: + return moment - sign * offset + except OverflowError: + return None + + +def format_timestamp(moment: Optional[datetime]) -> Optional[str]: + """Format a datetime as RFC 3339 UTC with millisecond precision (truncated) and ``Z``.""" + if moment is None: + return None + if moment.tzinfo is not None: + moment = moment.astimezone(timezone.utc) + return ( + f"{moment.year:04d}-{moment.month:02d}-{moment.day:02d}T" + f"{moment.hour:02d}:{moment.minute:02d}:{moment.second:02d}." + f"{moment.microsecond // 1000:03d}Z" + ) + + +def as_str(value: object) -> str: + """Return ``value`` when it is a string, else ``""``.""" + return value if isinstance(value, str) else "" + + +def as_int(value: object, default: int = 0) -> int: + """Return ``value`` as an integer when it is one (booleans excluded), else ``default``.""" + if isinstance(value, bool): + return default + if isinstance(value, int): + return value + if isinstance(value, float) and value.is_integer(): + return int(value) + return default + + +def clean_ip(value: object) -> str: + """Normalize an IP string: surrounding whitespace removed, ``0.0.0.0`` becomes ``""``.""" + text = as_str(value).strip() + return "" if text in ("", "0.0.0.0") else text + + +def risk_band(score: int) -> RiskBand: + """Return the risk band of a Risk Score. + + ``trusted`` for 0-29, ``suspicious`` for 30-59, ``dangerous`` for 60-100, and + ``rate_limited`` for values above 100 (the 999 rate-limit marker is not a score). + """ + if score > 100: + return "rate_limited" + if score >= 60: + return "dangerous" + if score >= 30: + return "suspicious" + return "trusted" + + +def is_rate_limited(score: int) -> bool: + """Return ``True`` for the rate-limit marker (any value above 100, in practice 999).""" + return score > 100 diff --git a/src/shieldlabs/_validation.py b/src/shieldlabs/_validation.py new file mode 100644 index 0000000..d53a4a3 --- /dev/null +++ b/src/shieldlabs/_validation.py @@ -0,0 +1,284 @@ +"""Argument validation shared by the clients. Every check runs before any HTTP request.""" + +from __future__ import annotations + +import ipaddress +import math +import os +import re +import warnings +from typing import Literal, Optional, Union +from urllib.parse import quote, urlsplit +from uuid import UUID + +from ._errors import ShieldLabsWarning, ValidationError + +__all__ = [ + "DEFAULT_HISTORY_BASE_URL", + "DEFAULT_MANAGEMENT_BASE_URL", + "LOOKUP_TYPES", + "LookupType", + "history_origin", + "management_origin", + "normalize_domain", + "require_secret", + "resolve_api_key", + "validate_count", + "validate_limit", + "validate_lookup", + "validate_max_retries", + "validate_offset", + "validate_seconds", + "validate_uuid", +] + +LookupType = Literal[ + "ip", "user_hid", "visitor_id", "request_id", "device_id", "session_id", "cookie_id" +] +"""The identifier a History API lookup searches by.""" + +LOOKUP_TYPES: tuple[str, ...] = ( + "ip", + "user_hid", + "visitor_id", + "request_id", + "device_id", + "session_id", + "cookie_id", +) + +DEFAULT_HISTORY_BASE_URL = "https://account.shieldlabs.ai" +DEFAULT_MANAGEMENT_BASE_URL = "https://api.shieldlabs.ai" + +_UUID_TYPES = frozenset({"visitor_id", "request_id", "device_id", "session_id", "cookie_id"}) +_UUID_PATTERN = re.compile( + r"[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}" +) +_OCTET = r"(?:25[0-5]|2[0-4][0-9]|1[0-9]{2}|[1-9]?[0-9])" +_IPV4_PATTERN = re.compile(rf"{_OCTET}(?:\.{_OCTET}){{3}}") +_API_KEY_PATTERN = re.compile(r"sec_[a-z0-9]{8}-[a-z0-9]{8}-[a-z0-9]{8}") +_URL_PATTERN = re.compile(r"https?://[^/\s?#]+(?:/[^\s?#]*)?", re.IGNORECASE) + +# The History API matches a User HID only when its path segment uses this canonical escaping: +# letters, digits, "-._~" and the characters below stay as they are, everything else is +# percent-encoded as UTF-8 with uppercase hex. Any other spelling (for example "%40" for "@") +# is compared literally and matches nothing. +_USER_HID_SAFE = "$&+,:;=@" + + +def _text(value: Union[str, UUID], name: str) -> str: + if isinstance(value, UUID): + return str(value) + if not isinstance(value, str): + raise ValidationError(f"{name} must be a string, got {type(value).__name__}") + return value + + +def validate_uuid(value: Union[str, UUID], name: str = "request_id") -> str: + """Return the UUID in lowercase, or raise ``ValidationError``. Any version, nil allowed.""" + text = _text(value, name) + if not _UUID_PATTERN.fullmatch(text): + raise ValidationError(f"{name} must be a UUID like 8f14e45f-ceea-4c1e-a3b2-1d2c3b4a5f60") + return text.lower() + + +def validate_lookup(lookup_type: str, value: Union[str, UUID]) -> tuple[str, str]: + """Validate a History lookup and return ``(type, value as an encoded path segment)``.""" + if lookup_type not in LOOKUP_TYPES: + allowed = ", ".join(LOOKUP_TYPES) + raise ValidationError(f"type must be one of {allowed}; got {lookup_type!r}") + if lookup_type in _UUID_TYPES: + return lookup_type, validate_uuid(value, lookup_type) + text = _text(value, lookup_type) + if lookup_type == "ip": + if ":" in text: + raise ValidationError( + "ip must be a dotted IPv4 address; IPv6 addresses are not searchable" + ) + if not _IPV4_PATTERN.fullmatch(text): + raise ValidationError( + f"ip must be a dotted IPv4 address such as 203.0.113.7; got {text!r}" + ) + return lookup_type, text + return lookup_type, _user_hid_segment(text) + + +def _user_hid_segment(text: str) -> str: + """Encode a User HID as one path segment the History API can match, or raise.""" + if text == "": + raise ValidationError("user_hid must be a non-empty string") + if text in (".", ".."): + raise ValidationError( + f"user_hid {text!r} cannot be searched: URL handling removes '.' and '..' " + "path segments, so the lookup would request a different path" + ) + if "/" in text: + raise ValidationError( + "user_hid values that contain '/' cannot be searched in the History API; " + "use User HIDs without '/', such as the hex output of user_hid()" + ) + try: + return quote(text, safe=_USER_HID_SAFE) + except UnicodeEncodeError: + raise ValidationError( + "user_hid must be valid Unicode text (it contains an unpaired surrogate)" + ) from None + + +def _integer(value: object, name: str) -> int: + if isinstance(value, bool) or not isinstance(value, int): + raise ValidationError(f"{name} must be an integer, got {value!r}") + return value + + +def validate_limit(value: object, name: str = "limit") -> int: + """An integer from 1 to 100 (the server silently replaces other values with 20).""" + number = _integer(value, name) + if not 1 <= number <= 100: + raise ValidationError(f"{name} must be between 1 and 100, got {number}") + return number + + +def validate_count(value: object, name: str) -> int: + """An integer that is 0 or greater.""" + number = _integer(value, name) + if number < 0: + raise ValidationError(f"{name} must be 0 or greater, got {number}") + return number + + +def validate_offset(value: object) -> int: + return validate_count(value, "offset") + + +def validate_max_retries(value: object) -> int: + return validate_count(value, "max_retries") + + +def validate_seconds(value: object, name: str, *, allow_zero: bool) -> float: + """A finite number of seconds; positive, or non-negative when ``allow_zero``.""" + if isinstance(value, bool) or not isinstance(value, (int, float)): + raise ValidationError(f"{name} must be a number of seconds, got {value!r}") + try: + seconds = float(value) + except OverflowError: # an int too large for a float + raise ValidationError(f"{name} must be a finite number of seconds") from None + if not math.isfinite(seconds): + raise ValidationError(f"{name} must be a finite number of seconds, got {value!r}") + if seconds < 0 or (seconds == 0 and not allow_zero): + qualifier = "0 or greater" if allow_zero else "greater than 0" + raise ValidationError(f"{name} must be {qualifier}, got {value!r}") + return seconds + + +def _base_url(value: Optional[str], env_name: str, default: str) -> str: + if value is None: + value = os.environ.get(env_name) or default + if not isinstance(value, str): + raise ValidationError("base_url must be a string") + url = value.strip().rstrip("/") + host = _url_host(url) if _URL_PATTERN.fullmatch(url) else None + if not host: + raise ValidationError(f"base_url must be an http(s) URL such as {default}; got {value!r}") + if url[:5].lower() == "http:" and not _is_loopback(host): + # Every request carries a credential, so plain http is limited to local test servers. + raise ValidationError( + f"base_url must use https; plain http is accepted only for localhost, 127.0.0.1 " + f"and ::1 (got a plain http URL for {host})" + ) + return url + + +def _url_host(url: str) -> Optional[str]: + try: + return urlsplit(url).hostname + except ValueError: # for example an unclosed IPv6 bracket + return None + + +def _is_loopback(host: str) -> bool: + if host == "localhost": + return True + try: + return ipaddress.ip_address(host).is_loopback + except ValueError: + return False + + +def history_origin(base_url: Optional[str]) -> str: + """History API origin. A trailing ``/api`` is removed: request paths already start with it.""" + url = _base_url(base_url, "SHIELDLABS_API_BASE_URL", DEFAULT_HISTORY_BASE_URL) + if url.endswith("/api"): + url = url[: -len("/api")] + return url + + +def management_origin(base_url: Optional[str]) -> str: + """Management API origin.""" + return _base_url(base_url, "SHIELDLABS_MANAGEMENT_BASE_URL", DEFAULT_MANAGEMENT_BASE_URL) + + +def _require_header_safe(text: str, name: str) -> None: + # Checked up front so that an unusable value never reaches the HTTP layer, whose errors + # could echo it. The value itself is never included in the message. + if not all("!" <= ch <= "~" for ch in text): + raise ValidationError( + f"{name} contains characters that cannot be sent in an HTTP header " + "(only visible ASCII characters are allowed)" + ) + + +def require_secret(value: Optional[str], env_name: str, name: str) -> str: + """Return a credential from the argument or the environment, stripped; never empty.""" + if value is None: + value = os.environ.get(env_name) + if value is not None and not isinstance(value, str): + raise ValidationError(f"{name} must be a string") + text = (value or "").strip() + if not text: + raise ValidationError(f"{name} is required: pass it or set {env_name}") + _require_header_safe(text, name) + return text + + +def resolve_api_key(value: Optional[str]) -> str: + """Private API Key from the argument or ``SHIELDLABS_API_KEY``. Warns on an unusual shape.""" + key = require_secret(value, "SHIELDLABS_API_KEY", "api_key") + if not _API_KEY_PATTERN.fullmatch(key): + warnings.warn( + "api_key does not look like a ShieldLabs Private API Key " + "(sec_xxxxxxxx-xxxxxxxx-xxxxxxxx). The History API expects the Private API Key, " + "not the Public Key or the Secret Key.", + ShieldLabsWarning, + stacklevel=4, + ) + return key + + +def normalize_domain(value: Optional[str]) -> str: + """Normalize a registered domain the way the server stores it. + + Trims and lowercases, then removes a scheme, any path, query or trailing slash, and a + leading ``www.``. ``https://www.Example.com/`` becomes ``example.com``. + """ + if value is None: + value = os.environ.get("SHIELDLABS_DOMAIN") + if value is not None and not isinstance(value, str): + raise ValidationError("domain must be a string") + text = (value or "").strip().lower() + if "://" in text: + text = text.split("://", 1)[1] + elif text.startswith("//"): + text = text[2:] + for separator in ("/", "?", "#"): + text = text.split(separator, 1)[0] + if text.startswith("www."): + text = text[len("www.") :] + if not text: + raise ValidationError( + "domain is required: pass the registered domain or set SHIELDLABS_DOMAIN" + ) + if not text.isascii(): + raise ValidationError("domain must be ASCII: pass the punycode form (xn--...)") + _require_header_safe(text, "domain") + return text diff --git a/src/shieldlabs/_version.py b/src/shieldlabs/_version.py new file mode 100644 index 0000000..5becc17 --- /dev/null +++ b/src/shieldlabs/_version.py @@ -0,0 +1 @@ +__version__ = "1.0.0" diff --git a/src/shieldlabs/py.typed b/src/shieldlabs/py.typed new file mode 100644 index 0000000..e69de29 diff --git a/src/shieldlabs/webhooks.py b/src/shieldlabs/webhooks.py new file mode 100644 index 0000000..cb6d2a8 --- /dev/null +++ b/src/shieldlabs/webhooks.py @@ -0,0 +1,240 @@ +"""Verify and parse ShieldLabs webhook deliveries. + +Every delivery is a ``POST`` with a JSON body and the header +``X-Shield-Signature: sha256=``. The HMAC key is the endpoint signing secret +string exactly as shown in the analytics dashboard, ``whsec_`` prefix included, and the message +is the raw request body. Always verify the raw bytes you received, before parsing them. + +ShieldLabs sends one delivery per identification and endpoint, with a 1-second timeout and no +retries. Respond with a 2xx within 1 second and do slow work afterwards. Keep handlers +idempotent on ``data.request_id``: a future release retries deliveries, and a retry resends +identical bytes. Use the History API for guaranteed reads and for the latest state. + +Example:: + + from shieldlabs import webhooks, IdentificationScoredEvent + + event = webhooks.construct_event(raw_body, headers.get("X-Shield-Signature"), secret) + if isinstance(event, IdentificationScoredEvent): + handle(event.data) +""" + +from __future__ import annotations + +import hashlib +import hmac +import json +import re +import warnings +from collections.abc import Mapping, Sequence +from dataclasses import dataclass, field +from datetime import datetime +from typing import Any, Literal, Optional, Union + +from ._errors import ShieldLabsWarning, SignatureVerificationError, WebhookParseError +from ._models import Identification +from ._normalize import parse_rfc3339 + +__all__ = [ + "SCHEMA_VERSION", + "SIGNATURE_HEADER", + "IdentificationScoredEvent", + "UnknownWebhookEvent", + "WebhookEvent", + "WebhookPingEvent", + "construct_event", + "verify_signature", +] + +SIGNATURE_HEADER = "X-Shield-Signature" +"""Name of the header that carries the signature.""" + +SCHEMA_VERSION = "2026-06-01" +"""Webhook ``schema_version`` this SDK was built for. Other values are parsed with a warning.""" + +_SIGNATURE_PREFIX = "sha256=" +_HEX_DIGEST = re.compile(r"[0-9a-f]{64}") + +Payload = Union[bytes, bytearray, memoryview, str] +Secrets = Union[str, Sequence[str]] + + +@dataclass(frozen=True) +class IdentificationScoredEvent: + """``identification.scored``: one identification was scored. + + Attributes: + event_type: Always ``"identification.scored"``. + schema_version: Envelope schema version, for example ``"2026-06-01"``. + created_at: When the event was created (aware UTC datetime), ``None`` if unparsable. + data: The normalized identification. + raw: The whole parsed envelope. + """ + + event_type: Literal["identification.scored"] + schema_version: str + created_at: Optional[datetime] + data: Identification + raw: Mapping[str, Any] = field(default_factory=dict, compare=False, repr=False) + + +@dataclass(frozen=True) +class WebhookPingEvent: + """``webhook.ping``: sent when you press Verify on an endpoint. It carries no data.""" + + event_type: Literal["webhook.ping"] + schema_version: str + created_at: Optional[datetime] + raw: Mapping[str, Any] = field(default_factory=dict, compare=False, repr=False) + + +@dataclass(frozen=True) +class UnknownWebhookEvent: + """An event type this SDK version does not know. Acknowledge it and ignore it or log it. + + Attributes: + data: The ``data`` object of the envelope when it has one, else ``None``. + """ + + event_type: str + schema_version: str + created_at: Optional[datetime] + data: Optional[Mapping[str, Any]] = None + raw: Mapping[str, Any] = field(default_factory=dict, compare=False, repr=False) + + +WebhookEvent = Union[IdentificationScoredEvent, WebhookPingEvent, UnknownWebhookEvent] +"""Every event ``construct_event`` can return.""" + + +def _payload_bytes(payload: Payload) -> bytes: + if isinstance(payload, str): + return payload.encode("utf-8", errors="surrogateescape") + if isinstance(payload, (bytes, bytearray, memoryview)): + return bytes(payload) + raise TypeError( + "payload must be the raw request body as bytes or str, " + f"not {type(payload).__name__}: never re-serialize parsed JSON" + ) + + +def _secret_list(secret: Optional[Secrets]) -> list[str]: + if secret is None: + return [] + if isinstance(secret, str): + return [secret] + if isinstance(secret, (bytes, bytearray, memoryview)) or not isinstance(secret, Sequence): + raise TypeError("secret must be a string or a sequence of strings") + items = list(secret) + for item in items: + if not isinstance(item, str): + raise TypeError("secret must be a string or a sequence of strings") + return items + + +def verify_signature( + payload: Payload, + signature_header: Optional[str], + secret: Secrets, +) -> bool: + """Check the ``X-Shield-Signature`` header of a delivery. + + Args: + payload: The raw request body, as bytes or str. Never re-serialized JSON. + signature_header: The ``X-Shield-Signature`` header value (``None`` when absent). + secret: The endpoint signing secret (``whsec_...``), or a list of secrets while you + rotate one: the delivery is valid when any of them matches. + + Returns: + ``True`` when the signature is valid. ``False`` for a missing or malformed header, + an empty secret or a wrong signature. + """ + try: + body = _payload_bytes(payload) + except UnicodeEncodeError: + return False + secrets = [item for item in _secret_list(secret) if item] + if not secrets or not isinstance(signature_header, str): + return False + header = signature_header.strip() + if not header.startswith(_SIGNATURE_PREFIX): + return False + received = header[len(_SIGNATURE_PREFIX) :].lower() + if not _HEX_DIGEST.fullmatch(received): + return False + valid = False + for item in secrets: + expected = hmac.new(item.encode("utf-8"), body, hashlib.sha256).hexdigest() + if hmac.compare_digest(expected, received): + valid = True + return valid + + +def construct_event( + payload: Payload, + signature_header: Optional[str], + secret: Secrets, +) -> WebhookEvent: + """Verify a delivery, then parse it into a typed event. + + Returns ``IdentificationScoredEvent``, ``WebhookPingEvent`` or ``UnknownWebhookEvent`` + (never raises for an unknown ``event_type``). + + Raises: + SignatureVerificationError: The signature is missing or does not match. + WebhookParseError: The verified body is not a webhook envelope. + """ + if not verify_signature(payload, signature_header, secret): + raise SignatureVerificationError( + "Webhook signature verification failed: check the endpoint signing secret and " + "pass the raw request body" + ) + return _parse_event(_payload_bytes(payload)) + + +def _parse_event(body: bytes) -> WebhookEvent: + try: + envelope = json.loads(body) + except (ValueError, RecursionError) as exc: + raise WebhookParseError("Webhook body is not valid JSON") from exc + if not isinstance(envelope, dict): + raise WebhookParseError("Webhook body is not a JSON object") + event_type = envelope.get("event_type") + if not isinstance(event_type, str) or not event_type: + raise WebhookParseError("Webhook body has no event_type") + schema_version = envelope.get("schema_version") + if not isinstance(schema_version, str): + schema_version = "" + if schema_version != SCHEMA_VERSION: + warnings.warn( + f"Webhook schema_version {schema_version!r} is not {SCHEMA_VERSION!r}; " + "parsing it anyway. Upgrade the shieldlabs package to get the latest fields.", + ShieldLabsWarning, + stacklevel=3, + ) + created_at = parse_rfc3339(envelope.get("created_at")) + data = envelope.get("data") + if event_type == "identification.scored": + if not isinstance(data, dict): + raise WebhookParseError("identification.scored event has no data object") + return IdentificationScoredEvent( + event_type="identification.scored", + schema_version=schema_version, + created_at=created_at, + data=Identification.from_webhook_data(data), + raw=envelope, + ) + if event_type == "webhook.ping": + return WebhookPingEvent( + event_type="webhook.ping", + schema_version=schema_version, + created_at=created_at, + raw=envelope, + ) + return UnknownWebhookEvent( + event_type=event_type, + schema_version=schema_version, + created_at=created_at, + data=data if isinstance(data, dict) else None, + raw=envelope, + ) diff --git a/tests/_support.py b/tests/_support.py new file mode 100644 index 0000000..dd66bd9 --- /dev/null +++ b/tests/_support.py @@ -0,0 +1,85 @@ +"""Shared test helpers.""" + +from __future__ import annotations + +import hashlib +import hmac +import json +from pathlib import Path +from typing import Any, Union + +import httpx + +from shieldlabs._http import AsyncTransport, SyncTransport + +DATA_DIR = Path(__file__).parent / "data" + +API_KEY = "sec_test0001-test0002-test0003" +SECRET_KEY = "0123456789abcdef0123456789abcdef" +DOMAIN = "example.com" +HISTORY_HOST = "account.shieldlabs.ai" +MANAGEMENT_HOST = "api.shieldlabs.ai" +REQUEST_ID = "02f1d973-84db-4156-a7f7-e799e6bf389b" +REQUEST_PATH = f"/api/v1/history/request_id/{REQUEST_ID}" + + +def load_json(name: str) -> Any: + return json.loads((DATA_DIR / name).read_text(encoding="utf-8")) + + +def load_bytes(name: str) -> bytes: + return (DATA_DIR / name).read_bytes() + + +def sign(secret: str, body: bytes) -> str: + return "sha256=" + hmac.new(secret.encode("utf-8"), body, hashlib.sha256).hexdigest() + + +def history_body(*rows: dict[str, Any], total: Union[int, None] = None) -> dict[str, Any]: + return {"data": list(rows), "total": len(rows) if total is None else total} + + +def row(request_id: str, **extra: Any) -> dict[str, Any]: + """A minimal History row for paging tests.""" + base: dict[str, Any] = { + "request_id": request_id, + "device_id": "ac7c303d-971b-41d1-8e25-cd5b46b46aed", + "score": 10, + "created_at": "2026-09-30 12:00:00.000", + } + base.update(extra) + return base + + +def uuid_for(index: int) -> str: + return f"00000000-0000-4000-8000-{index:012d}" + + +def empty_page() -> httpx.Response: + return httpx.Response(200, json={"data": [], "total": 0}) + + +class FakeTime: + """Deterministic clock, sleep and jitter for transports.""" + + def __init__(self) -> None: + self.now = 0.0 + self.sleeps: list[float] = [] + + def clock(self) -> float: + return self.now + + def sleep(self, seconds: float) -> None: + self.sleeps.append(round(seconds, 6)) + self.now += seconds + + async def async_sleep(self, seconds: float) -> None: + self.sleep(seconds) + + def install(self, transport: Union[SyncTransport, AsyncTransport]) -> None: + transport._clock = self.clock + transport._random = lambda: 1.0 + if isinstance(transport, AsyncTransport): + transport._sleep = self.async_sleep + else: + transport._sleep = self.sleep diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..113682b --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,30 @@ +from __future__ import annotations + +import pytest + +from _support import FakeTime + +_ENV_VARS = ( + "SHIELDLABS_API_KEY", + "SHIELDLABS_API_BASE_URL", + "SHIELDLABS_SECRET_KEY", + "SHIELDLABS_DOMAIN", + "SHIELDLABS_MANAGEMENT_BASE_URL", + "SHIELDLABS_WEBHOOK_SECRET", +) + + +@pytest.fixture(autouse=True) +def _isolated_env(monkeypatch: pytest.MonkeyPatch) -> None: + for name in _ENV_VARS: + monkeypatch.delenv(name, raising=False) + + +@pytest.fixture +def fake_time() -> FakeTime: + return FakeTime() + + +@pytest.fixture +def anyio_backend() -> str: + return "asyncio" diff --git a/tests/data/error-responses.json b/tests/data/error-responses.json new file mode 100644 index 0000000..2c96d9b --- /dev/null +++ b/tests/data/error-responses.json @@ -0,0 +1,109 @@ +{ + "description": "Error responses exactly as the ShieldLabs servers send them. Bodies are not uniform: parse defensively and branch on status.", + "cases": [ + { + "surface": "history", + "status": 401, + "content_type": "text/plain; charset=utf-8", + "body": "{\"error\":\"missing or invalid authorization header\"}\n", + "expected_error": "AuthenticationError", + "retry": false + }, + { + "surface": "history", + "status": 401, + "content_type": "text/plain; charset=utf-8", + "body": "{\"error\":\"invalid api key\"}\n", + "expected_error": "AuthenticationError", + "retry": false + }, + { + "surface": "history", + "status": 429, + "content_type": "application/json", + "body": "{\"error\":\"too many requests\"}\n", + "expected_error": "RateLimitError", + "retry": true + }, + { + "surface": "history", + "status": 500, + "content_type": "application/json", + "body": "{\"error\":\"code: 53, message: Cannot convert string 'abc' to type UUID\"}\n", + "expected_error": "ServerError", + "retry": true + }, + { + "surface": "history", + "status": 500, + "content_type": "text/plain; charset=utf-8", + "body": "{\"error\":\"internal error\"}\n", + "expected_error": "ServerError", + "retry": true + }, + { + "surface": "history", + "status": 404, + "content_type": "text/plain; charset=utf-8", + "body": "404 page not found", + "expected_error": "NotFoundError", + "retry": false + }, + { + "surface": "history", + "status": 502, + "content_type": "text/html", + "body": "

502 Bad Gateway

", + "expected_error": "ServerError", + "retry": true + }, + { + "surface": "management", + "status": 401, + "content_type": null, + "body": "", + "expected_error": "AuthenticationError", + "retry": false + }, + { + "surface": "management", + "status": 429, + "content_type": "application/json; charset=utf-8", + "body": "{\"error\":\"too many requests\"}", + "expected_error": "RateLimitError", + "retry": false + }, + { + "surface": "management", + "status": 503, + "content_type": "application/json; charset=utf-8", + "body": "{\"error\":\"server is busy\"}", + "expected_error": "ServerError", + "retry": true + }, + { + "surface": "management", + "status": 400, + "content_type": "application/json; charset=utf-8", + "body": "null", + "expected_error": "BadRequestError", + "retry": false + }, + { + "surface": "management", + "status": 402, + "content_type": null, + "body": "", + "expected_error": "QuotaExceededError", + "retry": false + }, + { + "surface": "management", + "status": 404, + "content_type": "text/plain; charset=utf-8", + "body": "404 page not found", + "expected_error": "NotFoundError", + "retry": false + } + ] +} diff --git a/tests/data/history-empty.json b/tests/data/history-empty.json new file mode 100644 index 0000000..e88e669 --- /dev/null +++ b/tests/data/history-empty.json @@ -0,0 +1,4 @@ +{ + "data": [], + "total": 0 +} diff --git a/tests/data/history-page.json b/tests/data/history-page.json new file mode 100644 index 0000000..047866b --- /dev/null +++ b/tests/data/history-page.json @@ -0,0 +1,278 @@ +{ + "data": [ + { + "request_id": "02f1d973-84db-4156-a7f7-e799e6bf389b", + "session_id": "bde78778-efd2-4c49-952f-1f11b9c05f35", + "cookie_id": "4449bb58-590c-444c-ae1f-d1ddc768dbdd", + "domain": "shop.example.com", + "site_domain": "example.com", + "user_hid": "9f86d081884c7d659a2feaa0c55ad015", + "device_id": "ac7c303d-971b-41d1-8e25-cd5b46b46aed", + "visitor_id": "bde0e249-20d8-4544-838c-ed9a0b6d7a36", + "ip": "203.0.113.24", + "os": "Windows", + "browser": "Chrome", + "device_type": "desktop", + "country": "Netherlands", + "connection_type": "proxy", + "score": 80, + "score_details": "[{\"Value\":10,\"Description\":\"Is proxy\"},{\"Value\":10,\"Description\":\"Is datacenter\"},{\"Value\":60,\"Description\":\"Antidetect browser (turn_block)\"},{\"Value\":0,\"Description\":\"Check Incomplete\"}]", + "created_at": "2026-09-30 12:34:56.123", + "ver": 1790771696123, + "web_rtc_ip": "198.51.100.23", + "web_rtc_country": "Germany", + "web_rtc_connection_type": "direct", + "scanner_web_rtc_ip": "0.0.0.0", + "scanner_web_rtc_country": "", + "scanner_web_rtc_connection_type": "", + "webrtc_leak_ip": "0.0.0.0", + "webrtc_leak_country": "", + "webrtc_leak_connection_type": "", + "webrtc_leak_source": "none", + "tcp_mss": 1460, + "mtu_value": 1500, + "mtu_hint": "direct", + "is_vpn": false, + "is_tor": false, + "is_proxy": true, + "is_datacenter": true, + "is_abuser": false, + "is_privacy_relay": false, + "is_stun_not_checked": false, + "check_incomplete": false, + "is_antidetect": true, + "is_os_mismatch": false, + "is_os_not_detected": false, + "is_timezone_mismatch": false, + "is_js_disabled": false, + "is_browser_automation": false, + "is_incognito": false, + "is_search_bot": false, + "stun_request_seen": true, + "is_scanner_stun_passed": false, + "stun_flow_status": "ok", + "entry_url": "https://shop.example.com/signup?utm_source=google&utm_medium=cpc&gclid=abc123", + "utm_source": "google", + "utm_medium": "cpc", + "traffic_channel": "Google Ads", + "traffic_channel_group": "Paid Search", + "traffic_reason": "gclid_present", + "click_id_type": "gclid", + "is_suspicious_paid_click": true + }, + { + "request_id": "7c1e2f4a-3b6d-4e8f-9a0b-1c2d3e4f5a6b", + "session_id": "a1b2c3d4-e5f6-4a7b-8c9d-0e1f2a3b4c5d", + "cookie_id": "f0e1d2c3-b4a5-4968-8776-655443322110", + "domain": "example.com", + "user_hid": "anonymous", + "device_id": "5d9a1f3e-7b2c-5e4d-8f6a-9b0c1d2e3f4a", + "visitor_id": "3c4d5e6f-7a8b-5c9d-8e0f-1a2b3c4d5e6f", + "ip": "192.0.2.44", + "os": "Mac OS X", + "browser": "Safari", + "device_type": "desktop", + "country": "United States", + "connection_type": "direct", + "score": 10, + "score_details": "[{\"Value\":10,\"Description\":\"Browser timezone ≠ IP-timezone\"}]", + "created_at": "2026-09-30 12:40:01.007", + "ver": 1790772001007, + "web_rtc_ip": "192.0.2.44", + "web_rtc_country": "United States", + "web_rtc_connection_type": "direct", + "scanner_web_rtc_ip": "0.0.0.0", + "scanner_web_rtc_country": "", + "scanner_web_rtc_connection_type": "", + "webrtc_leak_ip": "0.0.0.0", + "webrtc_leak_country": "", + "webrtc_leak_connection_type": "", + "webrtc_leak_source": "", + "tcp_mss": 1460, + "mtu_value": 1500, + "mtu_hint": "direct", + "is_vpn": false, + "is_tor": false, + "is_proxy": false, + "is_datacenter": false, + "is_abuser": false, + "is_privacy_relay": false, + "is_stun_not_checked": false, + "check_incomplete": false, + "is_antidetect": false, + "is_os_mismatch": false, + "is_os_not_detected": false, + "is_timezone_mismatch": true, + "is_js_disabled": false, + "is_browser_automation": false, + "is_incognito": true, + "is_search_bot": false, + "stun_request_seen": true, + "is_scanner_stun_passed": false, + "stun_flow_status": "ok" + }, + { + "request_id": "9e8d7c6b-5a49-4382-9716-05f4e3d2c1b0", + "session_id": "b0c1d2e3-f4a5-4b6c-9d7e-8f9a0b1c2d3e", + "cookie_id": "c3d4e5f6-a7b8-4c9d-8e0f-a1b2c3d4e5f6", + "domain": "example.com", + "user_hid": "", + "device_id": "e1f2a3b4-c5d6-5e7f-8a9b-0c1d2e3f4a5b", + "visitor_id": "d2e3f4a5-b6c7-5d8e-9f0a-1b2c3d4e5f6a", + "ip": "198.51.100.7", + "os": "Android", + "browser": "Chrome", + "device_type": "mobile", + "country": "France", + "connection_type": "vpn", + "score": 45, + "score_details": "[{\"Value\":15,\"Description\":\"Is VPN\"},{\"Value\":30,\"Description\":\"Stun is not checked\"},{\"Value\":-30,\"Description\":\"Stun passed (late arrival, corrected)\"},{\"Value\":30,\"Description\":\"Sticky verdict: Stun is not checked (request 11111111-2222-4333-8444-555555555555)\"},{\"Value\":0,\"Description\":\"IP ≠ leakIP (198.51.100.7 ≠ 203.0.113.9, source=scanner)\"}]", + "created_at": "2026-09-30 13:05:12", + "ver": 1790773512000, + "web_rtc_ip": "0.0.0.0", + "web_rtc_country": "", + "web_rtc_connection_type": "", + "scanner_web_rtc_ip": "203.0.113.9", + "scanner_web_rtc_country": "Spain", + "scanner_web_rtc_connection_type": "direct", + "webrtc_leak_ip": "203.0.113.9", + "webrtc_leak_country": "Spain", + "webrtc_leak_connection_type": "direct", + "webrtc_leak_source": "scanner", + "tcp_mss": 1380, + "mtu_value": 1420, + "mtu_hint": "vpn_likely", + "is_vpn": true, + "is_tor": false, + "is_proxy": false, + "is_datacenter": false, + "is_abuser": false, + "is_privacy_relay": false, + "is_stun_not_checked": true, + "check_incomplete": true, + "is_antidetect": false, + "is_os_mismatch": false, + "is_os_not_detected": false, + "is_timezone_mismatch": false, + "is_js_disabled": false, + "is_browser_automation": false, + "is_incognito": false, + "is_search_bot": false, + "stun_request_seen": false, + "is_scanner_stun_passed": true, + "stun_flow_status": "reply_without_request", + "entry_url": "https://example.com/pricing", + "referrer_domain": "news.example.org", + "traffic_channel": "Referral", + "traffic_channel_group": "Referral", + "traffic_reason": "external_referrer" + }, + { + "request_id": "1a2b3c4d-5e6f-4a7b-8c9d-0e1f2a3b4c5d", + "session_id": "00000000-0000-0000-0000-000000000000", + "cookie_id": "00000000-0000-0000-0000-000000000000", + "domain": "example.com", + "user_hid": "-1", + "device_id": "00000000-0000-0000-0000-000000000000", + "visitor_id": "00000000-0000-0000-0000-000000000000", + "ip": "203.0.113.200", + "os": "Unknown", + "browser": "Unknown", + "device_type": "desktop", + "country": "", + "connection_type": "unknown", + "score": 999, + "score_details": "[{\"Value\":999,\"Description\":\"User has been banned 1H, to many requests\"}]", + "created_at": "2026-09-30 13:10:00.500", + "ver": 1790773800500, + "web_rtc_ip": "0.0.0.0", + "web_rtc_country": "", + "web_rtc_connection_type": "", + "scanner_web_rtc_ip": "0.0.0.0", + "scanner_web_rtc_country": "", + "scanner_web_rtc_connection_type": "", + "webrtc_leak_ip": "0.0.0.0", + "webrtc_leak_country": "", + "webrtc_leak_connection_type": "", + "webrtc_leak_source": "", + "tcp_mss": 0, + "mtu_value": 0, + "mtu_hint": "", + "is_vpn": false, + "is_tor": false, + "is_proxy": false, + "is_datacenter": false, + "is_abuser": false, + "is_privacy_relay": false, + "is_stun_not_checked": false, + "check_incomplete": false, + "is_antidetect": false, + "is_os_mismatch": false, + "is_os_not_detected": false, + "is_timezone_mismatch": false, + "is_js_disabled": false, + "is_browser_automation": false, + "is_incognito": false, + "is_search_bot": false, + "stun_request_seen": false, + "is_scanner_stun_passed": false, + "stun_flow_status": "" + }, + { + "request_id": "4f5e6d7c-8b9a-4c1d-9e2f-3a4b5c6d7e8f", + "session_id": "5a6b7c8d-9e0f-4a1b-8c2d-3e4f5a6b7c8d", + "cookie_id": "6b7c8d9e-0f1a-4b2c-9d3e-4f5a6b7c8d9e", + "domain": "example.com", + "user_hid": "anonymous", + "device_id": "7c8d9e0f-1a2b-5c3d-8e4f-5a6b7c8d9e0f", + "visitor_id": "8d9e0f1a-2b3c-5d4e-9f5a-6b7c8d9e0f1a", + "ip": "198.51.100.66", + "os": "Linux", + "browser": "Chrome", + "device_type": "desktop", + "country": "United States", + "connection_type": "proxy", + "score": 0, + "score_details": "", + "created_at": "2026-09-30 13:20:30.250", + "ver": 1790774430250, + "web_rtc_ip": "203.0.113.77", + "web_rtc_country": "United States", + "web_rtc_connection_type": "direct", + "scanner_web_rtc_ip": "0.0.0.0", + "scanner_web_rtc_country": "", + "scanner_web_rtc_connection_type": "", + "webrtc_leak_ip": "0.0.0.0", + "webrtc_leak_country": "", + "webrtc_leak_connection_type": "", + "webrtc_leak_source": "none", + "tcp_mss": 1460, + "mtu_value": 1500, + "mtu_hint": "direct", + "is_vpn": false, + "is_tor": false, + "is_proxy": false, + "is_datacenter": false, + "is_abuser": false, + "is_privacy_relay": false, + "is_stun_not_checked": false, + "check_incomplete": false, + "is_antidetect": false, + "is_os_mismatch": false, + "is_os_not_detected": false, + "is_timezone_mismatch": false, + "is_js_disabled": false, + "is_browser_automation": false, + "is_incognito": false, + "is_search_bot": true, + "stun_request_seen": false, + "is_scanner_stun_passed": false, + "stun_flow_status": "", + "referrer_domain": "GoogleBot", + "traffic_channel": "Search bot", + "traffic_channel_group": "Bot", + "traffic_reason": "ip_crawler_detected" + } + ], + "total": 37 +} diff --git a/tests/data/management-profile-expected.json b/tests/data/management-profile-expected.json new file mode 100644 index 0000000..99c5f3f --- /dev/null +++ b/tests/data/management-profile-expected.json @@ -0,0 +1,7 @@ +{ + "domain": "example.com", + "remaining_identifications": 148230, + "public_key_masked": "****************************a3f8", + "secret_key_masked": "****************************9c2d", + "created_at": "2026-01-15T09:00:00.000Z" +} diff --git a/tests/data/management-profile.json b/tests/data/management-profile.json new file mode 100644 index 0000000..b0d1c74 --- /dev/null +++ b/tests/data/management-profile.json @@ -0,0 +1,8 @@ +{ + "Domain": "example.com", + "Weight": 148230, + "Callback": "", + "PublicKey": "****************************a3f8", + "Secret": "****************************9c2d", + "CreatedAt": "2026-01-15T09:00:00Z" +} diff --git a/tests/data/normalization-cases.json b/tests/data/normalization-cases.json new file mode 100644 index 0000000..0ccaff3 --- /dev/null +++ b/tests/data/normalization-cases.json @@ -0,0 +1,1050 @@ +{ + "description": "History API rows and webhook data objects with the Identification every SDK must produce. Compare observed_at at millisecond precision (truncate, never round). Keys not listed in expected (for example raw) are free.", + "cases": [ + { + "name": "history_02f1d973", + "source": "history", + "input": { + "request_id": "02f1d973-84db-4156-a7f7-e799e6bf389b", + "session_id": "bde78778-efd2-4c49-952f-1f11b9c05f35", + "cookie_id": "4449bb58-590c-444c-ae1f-d1ddc768dbdd", + "domain": "shop.example.com", + "site_domain": "example.com", + "user_hid": "9f86d081884c7d659a2feaa0c55ad015", + "device_id": "ac7c303d-971b-41d1-8e25-cd5b46b46aed", + "visitor_id": "bde0e249-20d8-4544-838c-ed9a0b6d7a36", + "ip": "203.0.113.24", + "os": "Windows", + "browser": "Chrome", + "device_type": "desktop", + "country": "Netherlands", + "connection_type": "proxy", + "score": 80, + "score_details": "[{\"Value\":10,\"Description\":\"Is proxy\"},{\"Value\":10,\"Description\":\"Is datacenter\"},{\"Value\":60,\"Description\":\"Antidetect browser (turn_block)\"},{\"Value\":0,\"Description\":\"Check Incomplete\"}]", + "created_at": "2026-09-30 12:34:56.123", + "ver": 1790771696123, + "web_rtc_ip": "198.51.100.23", + "web_rtc_country": "Germany", + "web_rtc_connection_type": "direct", + "scanner_web_rtc_ip": "0.0.0.0", + "scanner_web_rtc_country": "", + "scanner_web_rtc_connection_type": "", + "webrtc_leak_ip": "0.0.0.0", + "webrtc_leak_country": "", + "webrtc_leak_connection_type": "", + "webrtc_leak_source": "none", + "tcp_mss": 1460, + "mtu_value": 1500, + "mtu_hint": "direct", + "is_vpn": false, + "is_tor": false, + "is_proxy": true, + "is_datacenter": true, + "is_abuser": false, + "is_privacy_relay": false, + "is_stun_not_checked": false, + "check_incomplete": false, + "is_antidetect": true, + "is_os_mismatch": false, + "is_os_not_detected": false, + "is_timezone_mismatch": false, + "is_js_disabled": false, + "is_browser_automation": false, + "is_incognito": false, + "is_search_bot": false, + "stun_request_seen": true, + "is_scanner_stun_passed": false, + "stun_flow_status": "ok", + "entry_url": "https://shop.example.com/signup?utm_source=google&utm_medium=cpc&gclid=abc123", + "utm_source": "google", + "utm_medium": "cpc", + "traffic_channel": "Google Ads", + "traffic_channel_group": "Paid Search", + "traffic_reason": "gclid_present", + "click_id_type": "gclid", + "is_suspicious_paid_click": true + }, + "expected": { + "request_id": "02f1d973-84db-4156-a7f7-e799e6bf389b", + "visitor_id": "bde0e249-20d8-4544-838c-ed9a0b6d7a36", + "device_id": "ac7c303d-971b-41d1-8e25-cd5b46b46aed", + "session_id": "bde78778-efd2-4c49-952f-1f11b9c05f35", + "cookie_id": "4449bb58-590c-444c-ae1f-d1ddc768dbdd", + "user_hid": "9f86d081884c7d659a2feaa0c55ad015", + "domain": "example.com", + "public_ip": { + "ip": "203.0.113.24", + "country": "Netherlands" + }, + "local_ip": { + "ip": "198.51.100.23", + "country": "Germany" + }, + "connection_type": "proxy", + "os": "Windows", + "browser": "Chrome", + "device_type": "desktop", + "traffic_source": { + "channel": "Google Ads", + "referrer_domain": "", + "landing_url": "https://shop.example.com/signup?utm_source=google&utm_medium=cpc&gclid=abc123", + "click_id_type": "gclid", + "utm_source": "google", + "utm_medium": "cpc", + "utm_campaign": "", + "utm_content": "", + "utm_term": "" + }, + "risk_score": 80, + "signals": [ + { + "name": "proxy", + "weight": 10, + "description": "Is proxy" + }, + { + "name": "datacenter_ip", + "weight": 10, + "description": "Is datacenter" + }, + { + "name": "antidetect_browser", + "weight": 60, + "description": "Antidetect browser (turn_block)" + } + ], + "detection_flags": { + "vpn": false, + "privacy_relay": false, + "browser_vpn_proxy": false, + "tor": false, + "proxy": true, + "datacenter_ip": true, + "abuser": false, + "os_mismatch": false, + "os_not_detected": false, + "timezone_mismatch": false, + "anti_detect_browser": true, + "browser_automation": false, + "ip_mismatch": true, + "incognito": false, + "search_bot": false, + "suspicious_paid_click": true, + "javascript_disabled": false, + "stun_not_checked": false, + "check_incomplete": false + }, + "observed_at": "2026-09-30T12:34:56.123Z", + "source": "history" + } + }, + { + "name": "history_7c1e2f4a", + "source": "history", + "input": { + "request_id": "7c1e2f4a-3b6d-4e8f-9a0b-1c2d3e4f5a6b", + "session_id": "a1b2c3d4-e5f6-4a7b-8c9d-0e1f2a3b4c5d", + "cookie_id": "f0e1d2c3-b4a5-4968-8776-655443322110", + "domain": "example.com", + "user_hid": "anonymous", + "device_id": "5d9a1f3e-7b2c-5e4d-8f6a-9b0c1d2e3f4a", + "visitor_id": "3c4d5e6f-7a8b-5c9d-8e0f-1a2b3c4d5e6f", + "ip": "192.0.2.44", + "os": "Mac OS X", + "browser": "Safari", + "device_type": "desktop", + "country": "United States", + "connection_type": "direct", + "score": 10, + "score_details": "[{\"Value\":10,\"Description\":\"Browser timezone ≠ IP-timezone\"}]", + "created_at": "2026-09-30 12:40:01.007", + "ver": 1790772001007, + "web_rtc_ip": "192.0.2.44", + "web_rtc_country": "United States", + "web_rtc_connection_type": "direct", + "scanner_web_rtc_ip": "0.0.0.0", + "scanner_web_rtc_country": "", + "scanner_web_rtc_connection_type": "", + "webrtc_leak_ip": "0.0.0.0", + "webrtc_leak_country": "", + "webrtc_leak_connection_type": "", + "webrtc_leak_source": "", + "tcp_mss": 1460, + "mtu_value": 1500, + "mtu_hint": "direct", + "is_vpn": false, + "is_tor": false, + "is_proxy": false, + "is_datacenter": false, + "is_abuser": false, + "is_privacy_relay": false, + "is_stun_not_checked": false, + "check_incomplete": false, + "is_antidetect": false, + "is_os_mismatch": false, + "is_os_not_detected": false, + "is_timezone_mismatch": true, + "is_js_disabled": false, + "is_browser_automation": false, + "is_incognito": true, + "is_search_bot": false, + "stun_request_seen": true, + "is_scanner_stun_passed": false, + "stun_flow_status": "ok" + }, + "expected": { + "request_id": "7c1e2f4a-3b6d-4e8f-9a0b-1c2d3e4f5a6b", + "visitor_id": "3c4d5e6f-7a8b-5c9d-8e0f-1a2b3c4d5e6f", + "device_id": "5d9a1f3e-7b2c-5e4d-8f6a-9b0c1d2e3f4a", + "session_id": "a1b2c3d4-e5f6-4a7b-8c9d-0e1f2a3b4c5d", + "cookie_id": "f0e1d2c3-b4a5-4968-8776-655443322110", + "user_hid": "anonymous", + "domain": "example.com", + "public_ip": { + "ip": "192.0.2.44", + "country": "United States" + }, + "local_ip": { + "ip": "192.0.2.44", + "country": "United States" + }, + "connection_type": "direct", + "os": "Mac OS X", + "browser": "Safari", + "device_type": "desktop", + "traffic_source": { + "channel": "", + "referrer_domain": "", + "landing_url": "", + "click_id_type": "", + "utm_source": "", + "utm_medium": "", + "utm_campaign": "", + "utm_content": "", + "utm_term": "" + }, + "risk_score": 10, + "signals": [ + { + "name": "timezone_mismatch", + "weight": 10, + "description": "Browser timezone ≠ IP-timezone" + } + ], + "detection_flags": { + "vpn": false, + "privacy_relay": false, + "browser_vpn_proxy": false, + "tor": false, + "proxy": false, + "datacenter_ip": false, + "abuser": false, + "os_mismatch": false, + "os_not_detected": false, + "timezone_mismatch": true, + "anti_detect_browser": false, + "browser_automation": false, + "ip_mismatch": false, + "incognito": true, + "search_bot": false, + "suspicious_paid_click": false, + "javascript_disabled": false, + "stun_not_checked": false, + "check_incomplete": false + }, + "observed_at": "2026-09-30T12:40:01.007Z", + "source": "history" + } + }, + { + "name": "history_9e8d7c6b", + "source": "history", + "input": { + "request_id": "9e8d7c6b-5a49-4382-9716-05f4e3d2c1b0", + "session_id": "b0c1d2e3-f4a5-4b6c-9d7e-8f9a0b1c2d3e", + "cookie_id": "c3d4e5f6-a7b8-4c9d-8e0f-a1b2c3d4e5f6", + "domain": "example.com", + "user_hid": "", + "device_id": "e1f2a3b4-c5d6-5e7f-8a9b-0c1d2e3f4a5b", + "visitor_id": "d2e3f4a5-b6c7-5d8e-9f0a-1b2c3d4e5f6a", + "ip": "198.51.100.7", + "os": "Android", + "browser": "Chrome", + "device_type": "mobile", + "country": "France", + "connection_type": "vpn", + "score": 45, + "score_details": "[{\"Value\":15,\"Description\":\"Is VPN\"},{\"Value\":30,\"Description\":\"Stun is not checked\"},{\"Value\":-30,\"Description\":\"Stun passed (late arrival, corrected)\"},{\"Value\":30,\"Description\":\"Sticky verdict: Stun is not checked (request 11111111-2222-4333-8444-555555555555)\"},{\"Value\":0,\"Description\":\"IP ≠ leakIP (198.51.100.7 ≠ 203.0.113.9, source=scanner)\"}]", + "created_at": "2026-09-30 13:05:12", + "ver": 1790773512000, + "web_rtc_ip": "0.0.0.0", + "web_rtc_country": "", + "web_rtc_connection_type": "", + "scanner_web_rtc_ip": "203.0.113.9", + "scanner_web_rtc_country": "Spain", + "scanner_web_rtc_connection_type": "direct", + "webrtc_leak_ip": "203.0.113.9", + "webrtc_leak_country": "Spain", + "webrtc_leak_connection_type": "direct", + "webrtc_leak_source": "scanner", + "tcp_mss": 1380, + "mtu_value": 1420, + "mtu_hint": "vpn_likely", + "is_vpn": true, + "is_tor": false, + "is_proxy": false, + "is_datacenter": false, + "is_abuser": false, + "is_privacy_relay": false, + "is_stun_not_checked": true, + "check_incomplete": true, + "is_antidetect": false, + "is_os_mismatch": false, + "is_os_not_detected": false, + "is_timezone_mismatch": false, + "is_js_disabled": false, + "is_browser_automation": false, + "is_incognito": false, + "is_search_bot": false, + "stun_request_seen": false, + "is_scanner_stun_passed": true, + "stun_flow_status": "reply_without_request", + "entry_url": "https://example.com/pricing", + "referrer_domain": "news.example.org", + "traffic_channel": "Referral", + "traffic_channel_group": "Referral", + "traffic_reason": "external_referrer" + }, + "expected": { + "request_id": "9e8d7c6b-5a49-4382-9716-05f4e3d2c1b0", + "visitor_id": "d2e3f4a5-b6c7-5d8e-9f0a-1b2c3d4e5f6a", + "device_id": "e1f2a3b4-c5d6-5e7f-8a9b-0c1d2e3f4a5b", + "session_id": "b0c1d2e3-f4a5-4b6c-9d7e-8f9a0b1c2d3e", + "cookie_id": "c3d4e5f6-a7b8-4c9d-8e0f-a1b2c3d4e5f6", + "user_hid": null, + "domain": "example.com", + "public_ip": { + "ip": "198.51.100.7", + "country": "France" + }, + "local_ip": { + "ip": "203.0.113.9", + "country": "Spain" + }, + "connection_type": "vpn", + "os": "Android", + "browser": "Chrome", + "device_type": "mobile", + "traffic_source": { + "channel": "Referral", + "referrer_domain": "news.example.org", + "landing_url": "https://example.com/pricing", + "click_id_type": "", + "utm_source": "", + "utm_medium": "", + "utm_campaign": "", + "utm_content": "", + "utm_term": "" + }, + "risk_score": 45, + "signals": [ + { + "name": "vpn", + "weight": 15, + "description": "Is VPN" + }, + { + "name": "stun_not_checked", + "weight": 30, + "description": "Stun is not checked" + }, + { + "name": "stun_late_correction", + "weight": -30, + "description": "Stun passed (late arrival, corrected)" + }, + { + "name": "stun_is_not_checked", + "weight": 30, + "description": "Sticky verdict: Stun is not checked (request 11111111-2222-4333-8444-555555555555)" + } + ], + "detection_flags": { + "vpn": true, + "privacy_relay": false, + "browser_vpn_proxy": false, + "tor": false, + "proxy": false, + "datacenter_ip": false, + "abuser": false, + "os_mismatch": false, + "os_not_detected": false, + "timezone_mismatch": false, + "anti_detect_browser": false, + "browser_automation": false, + "ip_mismatch": true, + "incognito": false, + "search_bot": false, + "suspicious_paid_click": false, + "javascript_disabled": false, + "stun_not_checked": true, + "check_incomplete": true + }, + "observed_at": "2026-09-30T13:05:12.000Z", + "source": "history" + } + }, + { + "name": "history_1a2b3c4d", + "source": "history", + "input": { + "request_id": "1a2b3c4d-5e6f-4a7b-8c9d-0e1f2a3b4c5d", + "session_id": "00000000-0000-0000-0000-000000000000", + "cookie_id": "00000000-0000-0000-0000-000000000000", + "domain": "example.com", + "user_hid": "-1", + "device_id": "00000000-0000-0000-0000-000000000000", + "visitor_id": "00000000-0000-0000-0000-000000000000", + "ip": "203.0.113.200", + "os": "Unknown", + "browser": "Unknown", + "device_type": "desktop", + "country": "", + "connection_type": "unknown", + "score": 999, + "score_details": "[{\"Value\":999,\"Description\":\"User has been banned 1H, to many requests\"}]", + "created_at": "2026-09-30 13:10:00.500", + "ver": 1790773800500, + "web_rtc_ip": "0.0.0.0", + "web_rtc_country": "", + "web_rtc_connection_type": "", + "scanner_web_rtc_ip": "0.0.0.0", + "scanner_web_rtc_country": "", + "scanner_web_rtc_connection_type": "", + "webrtc_leak_ip": "0.0.0.0", + "webrtc_leak_country": "", + "webrtc_leak_connection_type": "", + "webrtc_leak_source": "", + "tcp_mss": 0, + "mtu_value": 0, + "mtu_hint": "", + "is_vpn": false, + "is_tor": false, + "is_proxy": false, + "is_datacenter": false, + "is_abuser": false, + "is_privacy_relay": false, + "is_stun_not_checked": false, + "check_incomplete": false, + "is_antidetect": false, + "is_os_mismatch": false, + "is_os_not_detected": false, + "is_timezone_mismatch": false, + "is_js_disabled": false, + "is_browser_automation": false, + "is_incognito": false, + "is_search_bot": false, + "stun_request_seen": false, + "is_scanner_stun_passed": false, + "stun_flow_status": "" + }, + "expected": { + "request_id": "1a2b3c4d-5e6f-4a7b-8c9d-0e1f2a3b4c5d", + "visitor_id": "00000000-0000-0000-0000-000000000000", + "device_id": "00000000-0000-0000-0000-000000000000", + "session_id": "00000000-0000-0000-0000-000000000000", + "cookie_id": "00000000-0000-0000-0000-000000000000", + "user_hid": "-1", + "domain": "example.com", + "public_ip": { + "ip": "203.0.113.200", + "country": "" + }, + "local_ip": { + "ip": "", + "country": "" + }, + "connection_type": "unknown", + "os": "Unknown", + "browser": "Unknown", + "device_type": "desktop", + "traffic_source": { + "channel": "", + "referrer_domain": "", + "landing_url": "", + "click_id_type": "", + "utm_source": "", + "utm_medium": "", + "utm_campaign": "", + "utm_content": "", + "utm_term": "" + }, + "risk_score": 999, + "signals": [ + { + "name": "rate_limited", + "weight": 999, + "description": "User has been banned 1H, to many requests" + } + ], + "detection_flags": { + "vpn": false, + "privacy_relay": false, + "browser_vpn_proxy": false, + "tor": false, + "proxy": false, + "datacenter_ip": false, + "abuser": false, + "os_mismatch": false, + "os_not_detected": false, + "timezone_mismatch": false, + "anti_detect_browser": false, + "browser_automation": false, + "ip_mismatch": false, + "incognito": false, + "search_bot": false, + "suspicious_paid_click": false, + "javascript_disabled": false, + "stun_not_checked": false, + "check_incomplete": false + }, + "observed_at": "2026-09-30T13:10:00.500Z", + "source": "history" + } + }, + { + "name": "history_4f5e6d7c", + "source": "history", + "input": { + "request_id": "4f5e6d7c-8b9a-4c1d-9e2f-3a4b5c6d7e8f", + "session_id": "5a6b7c8d-9e0f-4a1b-8c2d-3e4f5a6b7c8d", + "cookie_id": "6b7c8d9e-0f1a-4b2c-9d3e-4f5a6b7c8d9e", + "domain": "example.com", + "user_hid": "anonymous", + "device_id": "7c8d9e0f-1a2b-5c3d-8e4f-5a6b7c8d9e0f", + "visitor_id": "8d9e0f1a-2b3c-5d4e-9f5a-6b7c8d9e0f1a", + "ip": "198.51.100.66", + "os": "Linux", + "browser": "Chrome", + "device_type": "desktop", + "country": "United States", + "connection_type": "proxy", + "score": 0, + "score_details": "", + "created_at": "2026-09-30 13:20:30.250", + "ver": 1790774430250, + "web_rtc_ip": "203.0.113.77", + "web_rtc_country": "United States", + "web_rtc_connection_type": "direct", + "scanner_web_rtc_ip": "0.0.0.0", + "scanner_web_rtc_country": "", + "scanner_web_rtc_connection_type": "", + "webrtc_leak_ip": "0.0.0.0", + "webrtc_leak_country": "", + "webrtc_leak_connection_type": "", + "webrtc_leak_source": "none", + "tcp_mss": 1460, + "mtu_value": 1500, + "mtu_hint": "direct", + "is_vpn": false, + "is_tor": false, + "is_proxy": false, + "is_datacenter": false, + "is_abuser": false, + "is_privacy_relay": false, + "is_stun_not_checked": false, + "check_incomplete": false, + "is_antidetect": false, + "is_os_mismatch": false, + "is_os_not_detected": false, + "is_timezone_mismatch": false, + "is_js_disabled": false, + "is_browser_automation": false, + "is_incognito": false, + "is_search_bot": true, + "stun_request_seen": false, + "is_scanner_stun_passed": false, + "stun_flow_status": "", + "referrer_domain": "GoogleBot", + "traffic_channel": "Search bot", + "traffic_channel_group": "Bot", + "traffic_reason": "ip_crawler_detected" + }, + "expected": { + "request_id": "4f5e6d7c-8b9a-4c1d-9e2f-3a4b5c6d7e8f", + "visitor_id": "8d9e0f1a-2b3c-5d4e-9f5a-6b7c8d9e0f1a", + "device_id": "7c8d9e0f-1a2b-5c3d-8e4f-5a6b7c8d9e0f", + "session_id": "5a6b7c8d-9e0f-4a1b-8c2d-3e4f5a6b7c8d", + "cookie_id": "6b7c8d9e-0f1a-4b2c-9d3e-4f5a6b7c8d9e", + "user_hid": "anonymous", + "domain": "example.com", + "public_ip": { + "ip": "198.51.100.66", + "country": "United States" + }, + "local_ip": { + "ip": "203.0.113.77", + "country": "United States" + }, + "connection_type": "proxy", + "os": "Linux", + "browser": "Chrome", + "device_type": "desktop", + "traffic_source": { + "channel": "Search bot", + "referrer_domain": "GoogleBot", + "landing_url": "", + "click_id_type": "", + "utm_source": "", + "utm_medium": "", + "utm_campaign": "", + "utm_content": "", + "utm_term": "" + }, + "risk_score": 0, + "signals": [], + "detection_flags": { + "vpn": false, + "privacy_relay": false, + "browser_vpn_proxy": false, + "tor": false, + "proxy": false, + "datacenter_ip": false, + "abuser": false, + "os_mismatch": false, + "os_not_detected": false, + "timezone_mismatch": false, + "anti_detect_browser": false, + "browser_automation": false, + "ip_mismatch": false, + "incognito": false, + "search_bot": true, + "suspicious_paid_click": false, + "javascript_disabled": false, + "stun_not_checked": false, + "check_incomplete": false + }, + "observed_at": "2026-09-30T13:20:30.250Z", + "source": "history" + } + }, + { + "name": "webhook_scored", + "source": "webhook", + "input": { + "request_id": "02f1d973-84db-4156-a7f7-e799e6bf389b", + "visitor_id": "bde0e249-20d8-4544-838c-ed9a0b6d7a36", + "device_id": "ac7c303d-971b-41d1-8e25-cd5b46b46aed", + "session_id": "bde78778-efd2-4c49-952f-1f11b9c05f35", + "cookie_id": "4449bb58-590c-444c-ae1f-d1ddc768dbdd", + "user_hid": "9f86d081884c7d659a2feaa0c55ad015", + "domain": "example.com", + "public_ip": { + "ip": "203.0.113.24", + "country": "Netherlands" + }, + "local_ip": { + "ip": "198.51.100.23", + "country": "Germany" + }, + "connection_type": "proxy", + "os": "Windows", + "browser": "Chrome", + "device_type": "desktop", + "traffic_source": { + "channel": "Google Ads", + "referrer_domain": "google.com", + "landing_url": "https://shop.example.com/signup?utm_source=google&utm_medium=cpc&gclid=abc123", + "click_id_type": "gclid", + "utm_source": "google", + "utm_medium": "cpc", + "utm_campaign": "", + "utm_content": "", + "utm_term": "" + }, + "risk_score": 80, + "signals": [ + { + "name": "proxy", + "weight": 10 + }, + { + "name": "datacenter_ip", + "weight": 10 + }, + { + "name": "antidetect_browser", + "weight": 60 + } + ], + "detection_flags": { + "vpn": false, + "privacy_relay": false, + "browser_vpn_proxy": false, + "tor": false, + "proxy": true, + "datacenter_ip": true, + "abuser": false, + "os_mismatch": false, + "os_not_detected": false, + "timezone_mismatch": false, + "anti_detect_browser": true, + "browser_automation": false, + "ip_mismatch": true, + "incognito": false, + "search_bot": false, + "suspicious_paid_click": true, + "javascript_disabled": false, + "stun_not_checked": false, + "check_incomplete": false + }, + "observed_at": "2026-09-30T12:34:57.482913041Z" + }, + "expected": { + "request_id": "02f1d973-84db-4156-a7f7-e799e6bf389b", + "visitor_id": "bde0e249-20d8-4544-838c-ed9a0b6d7a36", + "device_id": "ac7c303d-971b-41d1-8e25-cd5b46b46aed", + "session_id": "bde78778-efd2-4c49-952f-1f11b9c05f35", + "cookie_id": "4449bb58-590c-444c-ae1f-d1ddc768dbdd", + "user_hid": "9f86d081884c7d659a2feaa0c55ad015", + "domain": "example.com", + "public_ip": { + "ip": "203.0.113.24", + "country": "Netherlands" + }, + "local_ip": { + "ip": "198.51.100.23", + "country": "Germany" + }, + "connection_type": "proxy", + "os": "Windows", + "browser": "Chrome", + "device_type": "desktop", + "traffic_source": { + "channel": "Google Ads", + "referrer_domain": "google.com", + "landing_url": "https://shop.example.com/signup?utm_source=google&utm_medium=cpc&gclid=abc123", + "click_id_type": "gclid", + "utm_source": "google", + "utm_medium": "cpc", + "utm_campaign": "", + "utm_content": "", + "utm_term": "" + }, + "risk_score": 80, + "signals": [ + { + "name": "proxy", + "weight": 10, + "description": null + }, + { + "name": "datacenter_ip", + "weight": 10, + "description": null + }, + { + "name": "antidetect_browser", + "weight": 60, + "description": null + } + ], + "detection_flags": { + "vpn": false, + "privacy_relay": false, + "browser_vpn_proxy": false, + "tor": false, + "proxy": true, + "datacenter_ip": true, + "abuser": false, + "os_mismatch": false, + "os_not_detected": false, + "timezone_mismatch": false, + "anti_detect_browser": true, + "browser_automation": false, + "ip_mismatch": true, + "incognito": false, + "search_bot": false, + "suspicious_paid_click": true, + "javascript_disabled": false, + "stun_not_checked": false, + "check_incomplete": false + }, + "observed_at": "2026-09-30T12:34:57.482Z", + "source": "webhook" + } + }, + { + "name": "webhook_rate_limited", + "source": "webhook", + "input": { + "request_id": "1a2b3c4d-5e6f-4a7b-8c9d-0e1f2a3b4c5d", + "visitor_id": "00000000-0000-0000-0000-000000000000", + "device_id": "00000000-0000-0000-0000-000000000000", + "session_id": "00000000-0000-0000-0000-000000000000", + "cookie_id": "00000000-0000-0000-0000-000000000000", + "user_hid": "-1", + "domain": "example.com", + "public_ip": { + "ip": "203.0.113.200", + "country": "" + }, + "local_ip": { + "ip": "", + "country": "" + }, + "connection_type": "unknown", + "os": "Unknown", + "browser": "Unknown", + "device_type": "desktop", + "traffic_source": { + "channel": "", + "referrer_domain": "", + "landing_url": "", + "click_id_type": "", + "utm_source": "", + "utm_medium": "", + "utm_campaign": "", + "utm_content": "", + "utm_term": "" + }, + "risk_score": 999, + "signals": [ + { + "name": "rate_limited", + "weight": 999 + } + ], + "detection_flags": { + "vpn": false, + "privacy_relay": false, + "browser_vpn_proxy": false, + "tor": false, + "proxy": false, + "datacenter_ip": false, + "abuser": false, + "os_mismatch": false, + "os_not_detected": false, + "timezone_mismatch": false, + "anti_detect_browser": false, + "browser_automation": false, + "ip_mismatch": false, + "incognito": false, + "search_bot": false, + "suspicious_paid_click": false, + "javascript_disabled": false, + "stun_not_checked": false, + "check_incomplete": false + }, + "observed_at": "2026-09-30T13:10:00.5Z" + }, + "expected": { + "request_id": "1a2b3c4d-5e6f-4a7b-8c9d-0e1f2a3b4c5d", + "visitor_id": "00000000-0000-0000-0000-000000000000", + "device_id": "00000000-0000-0000-0000-000000000000", + "session_id": "00000000-0000-0000-0000-000000000000", + "cookie_id": "00000000-0000-0000-0000-000000000000", + "user_hid": "-1", + "domain": "example.com", + "public_ip": { + "ip": "203.0.113.200", + "country": "" + }, + "local_ip": { + "ip": "", + "country": "" + }, + "connection_type": "unknown", + "os": "Unknown", + "browser": "Unknown", + "device_type": "desktop", + "traffic_source": { + "channel": "", + "referrer_domain": "", + "landing_url": "", + "click_id_type": "", + "utm_source": "", + "utm_medium": "", + "utm_campaign": "", + "utm_content": "", + "utm_term": "" + }, + "risk_score": 999, + "signals": [ + { + "name": "rate_limited", + "weight": 999, + "description": null + } + ], + "detection_flags": { + "vpn": false, + "privacy_relay": false, + "browser_vpn_proxy": false, + "tor": false, + "proxy": false, + "datacenter_ip": false, + "abuser": false, + "os_mismatch": false, + "os_not_detected": false, + "timezone_mismatch": false, + "anti_detect_browser": false, + "browser_automation": false, + "ip_mismatch": false, + "incognito": false, + "search_bot": false, + "suspicious_paid_click": false, + "javascript_disabled": false, + "stun_not_checked": false, + "check_incomplete": false + }, + "observed_at": "2026-09-30T13:10:00.500Z", + "source": "webhook" + } + }, + { + "name": "webhook_test_delivery", + "source": "webhook", + "input": { + "browser": "Chrome", + "connection_type": "proxy", + "cookie_id": "2c9d1e8f-4b7a-4c3e-9d2f-1a8b7c6d5e4f", + "detection_flags": { + "abuser": true, + "anti_detect_browser": false, + "check_incomplete": false, + "datacenter_ip": true, + "incognito": false, + "ip_mismatch": false, + "javascript_disabled": false, + "os_mismatch": false, + "os_not_detected": false, + "privacy_relay": false, + "proxy": true, + "stun_not_checked": false, + "suspicious_paid_click": false, + "timezone_mismatch": false, + "tor": false, + "vpn": false, + "browser_vpn_proxy": false + }, + "device_id": "6f1e2d3c-4b5a-5968-8776-655443322110", + "device_type": "desktop", + "domain": "example.com", + "local_ip": { + "country": "BY", + "ip": "198.51.100.10" + }, + "observed_at": "2026-09-30T12:34:56Z", + "os": "Windows", + "public_ip": { + "country": "BY", + "ip": "203.0.113.10" + }, + "request_id": "13f84f05-7c2a-4e9b-9f1d-2a6b8c0e4d11", + "risk_score": 30, + "session_id": "3a2b1c0d-9e8f-4a7b-8c6d-5e4f3a2b1c0d", + "signals": [ + { + "name": "proxy", + "weight": 10 + }, + { + "name": "datacenter_ip", + "weight": 10 + }, + { + "name": "abuser", + "weight": 10 + } + ], + "traffic_source": { + "channel": "Direct", + "click_id_type": "", + "landing_url": "https://example.com/", + "referrer_domain": "", + "utm_campaign": "", + "utm_content": "", + "utm_medium": "", + "utm_source": "", + "utm_term": "" + }, + "user_hid": null, + "visitor_id": "7a6b5c4d-3e2f-5a1b-9c8d-7e6f5a4b3c2d" + }, + "expected": { + "request_id": "13f84f05-7c2a-4e9b-9f1d-2a6b8c0e4d11", + "visitor_id": "7a6b5c4d-3e2f-5a1b-9c8d-7e6f5a4b3c2d", + "device_id": "6f1e2d3c-4b5a-5968-8776-655443322110", + "session_id": "3a2b1c0d-9e8f-4a7b-8c6d-5e4f3a2b1c0d", + "cookie_id": "2c9d1e8f-4b7a-4c3e-9d2f-1a8b7c6d5e4f", + "user_hid": null, + "domain": "example.com", + "public_ip": { + "ip": "203.0.113.10", + "country": "BY" + }, + "local_ip": { + "ip": "198.51.100.10", + "country": "BY" + }, + "connection_type": "proxy", + "os": "Windows", + "browser": "Chrome", + "device_type": "desktop", + "traffic_source": { + "channel": "Direct", + "referrer_domain": "", + "landing_url": "https://example.com/", + "click_id_type": "", + "utm_source": "", + "utm_medium": "", + "utm_campaign": "", + "utm_content": "", + "utm_term": "" + }, + "risk_score": 30, + "signals": [ + { + "name": "proxy", + "weight": 10, + "description": null + }, + { + "name": "datacenter_ip", + "weight": 10, + "description": null + }, + { + "name": "abuser", + "weight": 10, + "description": null + } + ], + "detection_flags": { + "vpn": false, + "privacy_relay": false, + "browser_vpn_proxy": false, + "tor": false, + "proxy": true, + "datacenter_ip": true, + "abuser": true, + "os_mismatch": false, + "os_not_detected": false, + "timezone_mismatch": false, + "anti_detect_browser": false, + "browser_automation": false, + "ip_mismatch": false, + "incognito": false, + "search_bot": false, + "suspicious_paid_click": false, + "javascript_disabled": false, + "stun_not_checked": false, + "check_incomplete": false + }, + "observed_at": "2026-09-30T12:34:56.000Z", + "source": "webhook" + } + } + ] +} diff --git a/tests/data/risk-band-cases.json b/tests/data/risk-band-cases.json new file mode 100644 index 0000000..98c7d1c --- /dev/null +++ b/tests/data/risk-band-cases.json @@ -0,0 +1,44 @@ +{ + "cases": [ + { + "score": 0, + "band": "trusted" + }, + { + "score": 5, + "band": "trusted" + }, + { + "score": 29, + "band": "trusted" + }, + { + "score": 30, + "band": "suspicious" + }, + { + "score": 45, + "band": "suspicious" + }, + { + "score": 59, + "band": "suspicious" + }, + { + "score": 60, + "band": "dangerous" + }, + { + "score": 85, + "band": "dangerous" + }, + { + "score": 100, + "band": "dangerous" + }, + { + "score": 999, + "band": "rate_limited" + } + ] +} diff --git a/tests/data/signal-slug-cases.json b/tests/data/signal-slug-cases.json new file mode 100644 index 0000000..2a73957 --- /dev/null +++ b/tests/data/signal-slug-cases.json @@ -0,0 +1,129 @@ +{ + "description": "Description (History score_details) to signal slug (webhook signals[].name). The slug function: exact map, then known prefixes, then 'Sticky verdict: ' prefix, then the fallback slugger.", + "cases": [ + { + "description": "Is tor", + "slug": "tor" + }, + { + "description": "Is VPN", + "slug": "vpn" + }, + { + "description": "Is privacy relay", + "slug": "privacy_relay" + }, + { + "description": "Is proxy", + "slug": "proxy" + }, + { + "description": "Is datacenter", + "slug": "datacenter_ip" + }, + { + "description": "Is abuser", + "slug": "abuser" + }, + { + "description": "Stun is not checked", + "slug": "stun_not_checked" + }, + { + "description": "Stun passed (late arrival, corrected)", + "slug": "stun_late_correction" + }, + { + "description": "UA OS is not detected", + "slug": "os_not_detected" + }, + { + "description": "Network OS is not detected", + "slug": "os_not_detected" + }, + { + "description": "Browser timezone ≠ IP-timezone", + "slug": "timezone_mismatch" + }, + { + "description": "Browser VPN/Proxy", + "slug": "browser_vpn_proxy" + }, + { + "description": "Browser Automation", + "slug": "browser_automation" + }, + { + "description": "Port scan routed via proxy (antidetect browser pattern)", + "slug": "proxy_routed_antidetect" + }, + { + "description": "User has been banned 1H, to many requests", + "slug": "rate_limited" + }, + { + "description": "Antidetect browser (turn_block)", + "slug": "antidetect_browser" + }, + { + "description": "Antidetect browser (port: 51234)", + "slug": "antidetect_browser" + }, + { + "description": "Os_mismatch (Fail by windows detect)", + "slug": "os_mismatch" + }, + { + "description": "OS mismatch2 (TCP behaviour: ttl)", + "slug": "os_mismatch2" + }, + { + "description": "TCP handshake (window)", + "slug": "tcp_handshake_v2" + }, + { + "description": "Latency test (WS vs TCP: 40ms)", + "slug": "ws_tcp_latency" + }, + { + "description": "JavaScript disabled (WebRTC, WebGL)", + "slug": "javascript_disabled" + }, + { + "description": "Sticky verdict: Antidetect browser (turn_block) (request 11111111-2222-4333-8444-555555555555)", + "slug": "antidetect_browser" + }, + { + "description": "Sticky verdict: Port scan routed via proxy (antidetect browser pattern) (request 11111111-2222-4333-8444-555555555555)", + "slug": "port_scan_routed_via_proxy" + }, + { + "description": "IP ≠ leakIP (198.51.100.7 ≠ 203.0.113.9, source=scanner)", + "slug": "ip__neq__leakip" + }, + { + "description": "Device spoofing", + "slug": "device_spoofing" + }, + { + "description": " Weird -- Name / Here ", + "slug": "weird_name_here" + }, + { + "description": "État Spécial (x)", + "slug": "état_spécial" + }, + { + "description": "(only parens)", + "slug": "unknown" + }, + { + "description": "!!!", + "slug": "unknown" + }, + { + "description": "", + "slug": "unknown" + } + ] +} diff --git a/tests/data/webhook-identification-scored.json b/tests/data/webhook-identification-scored.json new file mode 100644 index 0000000..f2de2d2 --- /dev/null +++ b/tests/data/webhook-identification-scored.json @@ -0,0 +1,74 @@ +{ + "event_type": "identification.scored", + "schema_version": "2026-06-01", + "created_at": "2026-09-30T12:34:57.482913041Z", + "data": { + "request_id": "02f1d973-84db-4156-a7f7-e799e6bf389b", + "visitor_id": "bde0e249-20d8-4544-838c-ed9a0b6d7a36", + "device_id": "ac7c303d-971b-41d1-8e25-cd5b46b46aed", + "session_id": "bde78778-efd2-4c49-952f-1f11b9c05f35", + "cookie_id": "4449bb58-590c-444c-ae1f-d1ddc768dbdd", + "user_hid": "9f86d081884c7d659a2feaa0c55ad015", + "domain": "example.com", + "public_ip": { + "ip": "203.0.113.24", + "country": "Netherlands" + }, + "local_ip": { + "ip": "198.51.100.23", + "country": "Germany" + }, + "connection_type": "proxy", + "os": "Windows", + "browser": "Chrome", + "device_type": "desktop", + "traffic_source": { + "channel": "Google Ads", + "referrer_domain": "google.com", + "landing_url": "https://shop.example.com/signup?utm_source=google&utm_medium=cpc&gclid=abc123", + "click_id_type": "gclid", + "utm_source": "google", + "utm_medium": "cpc", + "utm_campaign": "", + "utm_content": "", + "utm_term": "" + }, + "risk_score": 80, + "signals": [ + { + "name": "proxy", + "weight": 10 + }, + { + "name": "datacenter_ip", + "weight": 10 + }, + { + "name": "antidetect_browser", + "weight": 60 + } + ], + "detection_flags": { + "vpn": false, + "privacy_relay": false, + "browser_vpn_proxy": false, + "tor": false, + "proxy": true, + "datacenter_ip": true, + "abuser": false, + "os_mismatch": false, + "os_not_detected": false, + "timezone_mismatch": false, + "anti_detect_browser": true, + "browser_automation": false, + "ip_mismatch": true, + "incognito": false, + "search_bot": false, + "suspicious_paid_click": true, + "javascript_disabled": false, + "stun_not_checked": false, + "check_incomplete": false + }, + "observed_at": "2026-09-30T12:34:57.482913041Z" + } +} diff --git a/tests/data/webhook-identification-scored.raw.txt b/tests/data/webhook-identification-scored.raw.txt new file mode 100644 index 0000000..a4da089 --- /dev/null +++ b/tests/data/webhook-identification-scored.raw.txt @@ -0,0 +1 @@ +{"event_type":"identification.scored","schema_version":"2026-06-01","created_at":"2026-09-30T12:34:57.482913041Z","data":{"request_id":"02f1d973-84db-4156-a7f7-e799e6bf389b","visitor_id":"bde0e249-20d8-4544-838c-ed9a0b6d7a36","device_id":"ac7c303d-971b-41d1-8e25-cd5b46b46aed","session_id":"bde78778-efd2-4c49-952f-1f11b9c05f35","cookie_id":"4449bb58-590c-444c-ae1f-d1ddc768dbdd","user_hid":"9f86d081884c7d659a2feaa0c55ad015","domain":"example.com","public_ip":{"ip":"203.0.113.24","country":"Netherlands"},"local_ip":{"ip":"198.51.100.23","country":"Germany"},"connection_type":"proxy","os":"Windows","browser":"Chrome","device_type":"desktop","traffic_source":{"channel":"Google Ads","referrer_domain":"google.com","landing_url":"https://shop.example.com/signup?utm_source=google\u0026utm_medium=cpc\u0026gclid=abc123","click_id_type":"gclid","utm_source":"google","utm_medium":"cpc","utm_campaign":"","utm_content":"","utm_term":""},"risk_score":80,"signals":[{"name":"proxy","weight":10},{"name":"datacenter_ip","weight":10},{"name":"antidetect_browser","weight":60}],"detection_flags":{"vpn":false,"privacy_relay":false,"browser_vpn_proxy":false,"tor":false,"proxy":true,"datacenter_ip":true,"abuser":false,"os_mismatch":false,"os_not_detected":false,"timezone_mismatch":false,"anti_detect_browser":true,"browser_automation":false,"ip_mismatch":true,"incognito":false,"search_bot":false,"suspicious_paid_click":true,"javascript_disabled":false,"stun_not_checked":false,"check_incomplete":false},"observed_at":"2026-09-30T12:34:57.482913041Z"}} \ No newline at end of file diff --git a/tests/data/webhook-ping.json b/tests/data/webhook-ping.json new file mode 100644 index 0000000..94568cc --- /dev/null +++ b/tests/data/webhook-ping.json @@ -0,0 +1,5 @@ +{ + "created_at": "2026-09-30T12:34:56Z", + "event_type": "webhook.ping", + "schema_version": "2026-06-01" +} diff --git a/tests/data/webhook-ping.raw.txt b/tests/data/webhook-ping.raw.txt new file mode 100644 index 0000000..0f9f73b --- /dev/null +++ b/tests/data/webhook-ping.raw.txt @@ -0,0 +1 @@ +{"created_at":"2026-09-30T12:34:56Z","event_type":"webhook.ping","schema_version":"2026-06-01"} \ No newline at end of file diff --git a/tests/data/webhook-rate-limited.json b/tests/data/webhook-rate-limited.json new file mode 100644 index 0000000..b9edf47 --- /dev/null +++ b/tests/data/webhook-rate-limited.json @@ -0,0 +1,66 @@ +{ + "event_type": "identification.scored", + "schema_version": "2026-06-01", + "created_at": "2026-09-30T13:10:00.5Z", + "data": { + "request_id": "1a2b3c4d-5e6f-4a7b-8c9d-0e1f2a3b4c5d", + "visitor_id": "00000000-0000-0000-0000-000000000000", + "device_id": "00000000-0000-0000-0000-000000000000", + "session_id": "00000000-0000-0000-0000-000000000000", + "cookie_id": "00000000-0000-0000-0000-000000000000", + "user_hid": "-1", + "domain": "example.com", + "public_ip": { + "ip": "203.0.113.200", + "country": "" + }, + "local_ip": { + "ip": "", + "country": "" + }, + "connection_type": "unknown", + "os": "Unknown", + "browser": "Unknown", + "device_type": "desktop", + "traffic_source": { + "channel": "", + "referrer_domain": "", + "landing_url": "", + "click_id_type": "", + "utm_source": "", + "utm_medium": "", + "utm_campaign": "", + "utm_content": "", + "utm_term": "" + }, + "risk_score": 999, + "signals": [ + { + "name": "rate_limited", + "weight": 999 + } + ], + "detection_flags": { + "vpn": false, + "privacy_relay": false, + "browser_vpn_proxy": false, + "tor": false, + "proxy": false, + "datacenter_ip": false, + "abuser": false, + "os_mismatch": false, + "os_not_detected": false, + "timezone_mismatch": false, + "anti_detect_browser": false, + "browser_automation": false, + "ip_mismatch": false, + "incognito": false, + "search_bot": false, + "suspicious_paid_click": false, + "javascript_disabled": false, + "stun_not_checked": false, + "check_incomplete": false + }, + "observed_at": "2026-09-30T13:10:00.5Z" + } +} diff --git a/tests/data/webhook-signature-vectors.json b/tests/data/webhook-signature-vectors.json new file mode 100644 index 0000000..7529cf0 --- /dev/null +++ b/tests/data/webhook-signature-vectors.json @@ -0,0 +1,201 @@ +{ + "description": "Shared webhook signature test vectors for every ShieldLabs server SDK. Header format: X-Shield-Signature: sha256=. body_base64 holds the exact raw bytes; body is the same bytes decoded as UTF-8 for convenience. A vector has either 'secret' (string) or 'secrets' (list: valid when any matches).", + "header_name": "X-Shield-Signature", + "vectors": [ + { + "name": "valid_lowercase_hex", + "secret": "whsec_test_4f1a9c2e7b5d3a8f6e0c1b2d9a7e5f3c", + "body": "{\"event_type\":\"identification.scored\",\"schema_version\":\"2026-06-01\",\"created_at\":\"2026-09-30T10:15:00Z\",\"data\":{\"request_id\":\"3f2b8c1e-9d4a-4e6b-8a7c-2d1e0f9b6a53\",\"risk_score\":35,\"user_hid\":null}}", + "body_base64": "eyJldmVudF90eXBlIjoiaWRlbnRpZmljYXRpb24uc2NvcmVkIiwic2NoZW1hX3ZlcnNpb24iOiIyMDI2LTA2LTAxIiwiY3JlYXRlZF9hdCI6IjIwMjYtMDktMzBUMTA6MTU6MDBaIiwiZGF0YSI6eyJyZXF1ZXN0X2lkIjoiM2YyYjhjMWUtOWQ0YS00ZTZiLThhN2MtMmQxZTBmOWI2YTUzIiwicmlza19zY29yZSI6MzUsInVzZXJfaGlkIjpudWxsfX0=", + "signature_header": "sha256=1c2d0251d83eba5ed62c98d3758b805d227022850cdb8ddb1aa3754524345453", + "valid": true, + "note": "Canonical header as sent by ShieldLabs" + }, + { + "name": "valid_uppercase_hex", + "secret": "whsec_test_4f1a9c2e7b5d3a8f6e0c1b2d9a7e5f3c", + "body": "{\"event_type\":\"identification.scored\",\"schema_version\":\"2026-06-01\",\"created_at\":\"2026-09-30T10:15:00Z\",\"data\":{\"request_id\":\"3f2b8c1e-9d4a-4e6b-8a7c-2d1e0f9b6a53\",\"risk_score\":35,\"user_hid\":null}}", + "body_base64": "eyJldmVudF90eXBlIjoiaWRlbnRpZmljYXRpb24uc2NvcmVkIiwic2NoZW1hX3ZlcnNpb24iOiIyMDI2LTA2LTAxIiwiY3JlYXRlZF9hdCI6IjIwMjYtMDktMzBUMTA6MTU6MDBaIiwiZGF0YSI6eyJyZXF1ZXN0X2lkIjoiM2YyYjhjMWUtOWQ0YS00ZTZiLThhN2MtMmQxZTBmOWI2YTUzIiwicmlza19zY29yZSI6MzUsInVzZXJfaGlkIjpudWxsfX0=", + "signature_header": "sha256=1C2D0251D83EBA5ED62C98D3758B805D227022850CDB8DDB1AA3754524345453", + "valid": true, + "note": "Hex digest compared case-insensitively" + }, + { + "name": "valid_surrounding_whitespace", + "secret": "whsec_test_4f1a9c2e7b5d3a8f6e0c1b2d9a7e5f3c", + "body": "{\"event_type\":\"identification.scored\",\"schema_version\":\"2026-06-01\",\"created_at\":\"2026-09-30T10:15:00Z\",\"data\":{\"request_id\":\"3f2b8c1e-9d4a-4e6b-8a7c-2d1e0f9b6a53\",\"risk_score\":35,\"user_hid\":null}}", + "body_base64": "eyJldmVudF90eXBlIjoiaWRlbnRpZmljYXRpb24uc2NvcmVkIiwic2NoZW1hX3ZlcnNpb24iOiIyMDI2LTA2LTAxIiwiY3JlYXRlZF9hdCI6IjIwMjYtMDktMzBUMTA6MTU6MDBaIiwiZGF0YSI6eyJyZXF1ZXN0X2lkIjoiM2YyYjhjMWUtOWQ0YS00ZTZiLThhN2MtMmQxZTBmOWI2YTUzIiwicmlza19zY29yZSI6MzUsInVzZXJfaGlkIjpudWxsfX0=", + "signature_header": " sha256=1c2d0251d83eba5ed62c98d3758b805d227022850cdb8ddb1aa3754524345453 ", + "valid": true, + "note": "Leading/trailing whitespace is trimmed from the header value" + }, + { + "name": "valid_utf8_multibyte_body", + "secret": "whsec_test_4f1a9c2e7b5d3a8f6e0c1b2d9a7e5f3c", + "body": "{\"event_type\":\"webhook.ping\",\"schema_version\":\"2026-06-01\",\"created_at\":\"2026-09-30T10:15:00Z\",\"note\":\"проверка ✓ 日本\"}", + "body_base64": "eyJldmVudF90eXBlIjoid2ViaG9vay5waW5nIiwic2NoZW1hX3ZlcnNpb24iOiIyMDI2LTA2LTAxIiwiY3JlYXRlZF9hdCI6IjIwMjYtMDktMzBUMTA6MTU6MDBaIiwibm90ZSI6ItC/0YDQvtCy0LXRgNC60LAg4pyTIOaXpeacrCJ9", + "signature_header": "sha256=64a1212632b11b58dde55df7689bad609ae82f84681608ce3394da4d399b4362", + "valid": true, + "note": "HMAC is computed over the raw UTF-8 bytes" + }, + { + "name": "valid_pretty_printed_body", + "secret": "whsec_test_4f1a9c2e7b5d3a8f6e0c1b2d9a7e5f3c", + "body": "{\n \"event_type\": \"identification.scored\",\n \"schema_version\": \"2026-06-01\",\n \"created_at\": \"2026-09-30T10:15:00Z\",\n \"data\": {\n \"request_id\": \"3f2b8c1e-9d4a-4e6b-8a7c-2d1e0f9b6a53\",\n \"risk_score\": 35,\n \"user_hid\": null\n }\n}", + "body_base64": "ewogICJldmVudF90eXBlIjogImlkZW50aWZpY2F0aW9uLnNjb3JlZCIsCiAgInNjaGVtYV92ZXJzaW9uIjogIjIwMjYtMDYtMDEiLAogICJjcmVhdGVkX2F0IjogIjIwMjYtMDktMzBUMTA6MTU6MDBaIiwKICAiZGF0YSI6IHsKICAgICJyZXF1ZXN0X2lkIjogIjNmMmI4YzFlLTlkNGEtNGU2Yi04YTdjLTJkMWUwZjliNmE1MyIsCiAgICAicmlza19zY29yZSI6IDM1LAogICAgInVzZXJfaGlkIjogbnVsbAogIH0KfQ==", + "signature_header": "sha256=263308293bcccb571464f60e0f8ef6076f9d23559b1658b6dc6f74ac3fdffd35", + "valid": true, + "note": "Whatever bytes arrive are what was signed" + }, + { + "name": "invalid_reserialized_body", + "secret": "whsec_test_4f1a9c2e7b5d3a8f6e0c1b2d9a7e5f3c", + "body": "{\n \"event_type\": \"identification.scored\",\n \"schema_version\": \"2026-06-01\",\n \"created_at\": \"2026-09-30T10:15:00Z\",\n \"data\": {\n \"request_id\": \"3f2b8c1e-9d4a-4e6b-8a7c-2d1e0f9b6a53\",\n \"risk_score\": 35,\n \"user_hid\": null\n }\n}", + "body_base64": "ewogICJldmVudF90eXBlIjogImlkZW50aWZpY2F0aW9uLnNjb3JlZCIsCiAgInNjaGVtYV92ZXJzaW9uIjogIjIwMjYtMDYtMDEiLAogICJjcmVhdGVkX2F0IjogIjIwMjYtMDktMzBUMTA6MTU6MDBaIiwKICAiZGF0YSI6IHsKICAgICJyZXF1ZXN0X2lkIjogIjNmMmI4YzFlLTlkNGEtNGU2Yi04YTdjLTJkMWUwZjliNmE1MyIsCiAgICAicmlza19zY29yZSI6IDM1LAogICAgInVzZXJfaGlkIjogbnVsbAogIH0KfQ==", + "signature_header": "sha256=1c2d0251d83eba5ed62c98d3758b805d227022850cdb8ddb1aa3754524345453", + "valid": false, + "note": "Signature of the compact body does not match a re-serialized body: always verify raw bytes" + }, + { + "name": "invalid_tampered_body", + "secret": "whsec_test_4f1a9c2e7b5d3a8f6e0c1b2d9a7e5f3c", + "body": "{\"event_type\":\"identification.scored\",\"schema_version\":\"2026-06-01\",\"created_at\":\"2026-09-30T10:15:00Z\",\"data\":{\"request_id\":\"3f2b8c1e-9d4a-4e6b-8a7c-2d1e0f9b6a53\",\"risk_score\":5,\"user_hid\":null}}", + "body_base64": "eyJldmVudF90eXBlIjoiaWRlbnRpZmljYXRpb24uc2NvcmVkIiwic2NoZW1hX3ZlcnNpb24iOiIyMDI2LTA2LTAxIiwiY3JlYXRlZF9hdCI6IjIwMjYtMDktMzBUMTA6MTU6MDBaIiwiZGF0YSI6eyJyZXF1ZXN0X2lkIjoiM2YyYjhjMWUtOWQ0YS00ZTZiLThhN2MtMmQxZTBmOWI2YTUzIiwicmlza19zY29yZSI6NSwidXNlcl9oaWQiOm51bGx9fQ==", + "signature_header": "sha256=1c2d0251d83eba5ed62c98d3758b805d227022850cdb8ddb1aa3754524345453", + "valid": false, + "note": "Body changed after signing" + }, + { + "name": "invalid_wrong_secret", + "secret": "whsec_test_0000000000000000000000000000000", + "body": "{\"event_type\":\"identification.scored\",\"schema_version\":\"2026-06-01\",\"created_at\":\"2026-09-30T10:15:00Z\",\"data\":{\"request_id\":\"3f2b8c1e-9d4a-4e6b-8a7c-2d1e0f9b6a53\",\"risk_score\":35,\"user_hid\":null}}", + "body_base64": "eyJldmVudF90eXBlIjoiaWRlbnRpZmljYXRpb24uc2NvcmVkIiwic2NoZW1hX3ZlcnNpb24iOiIyMDI2LTA2LTAxIiwiY3JlYXRlZF9hdCI6IjIwMjYtMDktMzBUMTA6MTU6MDBaIiwiZGF0YSI6eyJyZXF1ZXN0X2lkIjoiM2YyYjhjMWUtOWQ0YS00ZTZiLThhN2MtMmQxZTBmOWI2YTUzIiwicmlza19zY29yZSI6MzUsInVzZXJfaGlkIjpudWxsfX0=", + "signature_header": "sha256=1c2d0251d83eba5ed62c98d3758b805d227022850cdb8ddb1aa3754524345453", + "valid": false, + "note": "Different endpoint secret" + }, + { + "name": "invalid_missing_prefix", + "secret": "whsec_test_4f1a9c2e7b5d3a8f6e0c1b2d9a7e5f3c", + "body": "{\"event_type\":\"identification.scored\",\"schema_version\":\"2026-06-01\",\"created_at\":\"2026-09-30T10:15:00Z\",\"data\":{\"request_id\":\"3f2b8c1e-9d4a-4e6b-8a7c-2d1e0f9b6a53\",\"risk_score\":35,\"user_hid\":null}}", + "body_base64": "eyJldmVudF90eXBlIjoiaWRlbnRpZmljYXRpb24uc2NvcmVkIiwic2NoZW1hX3ZlcnNpb24iOiIyMDI2LTA2LTAxIiwiY3JlYXRlZF9hdCI6IjIwMjYtMDktMzBUMTA6MTU6MDBaIiwiZGF0YSI6eyJyZXF1ZXN0X2lkIjoiM2YyYjhjMWUtOWQ0YS00ZTZiLThhN2MtMmQxZTBmOWI2YTUzIiwicmlza19zY29yZSI6MzUsInVzZXJfaGlkIjpudWxsfX0=", + "signature_header": "1c2d0251d83eba5ed62c98d3758b805d227022850cdb8ddb1aa3754524345453", + "valid": false, + "note": "The sha256= prefix is required" + }, + { + "name": "invalid_wrong_algorithm_prefix", + "secret": "whsec_test_4f1a9c2e7b5d3a8f6e0c1b2d9a7e5f3c", + "body": "{\"event_type\":\"identification.scored\",\"schema_version\":\"2026-06-01\",\"created_at\":\"2026-09-30T10:15:00Z\",\"data\":{\"request_id\":\"3f2b8c1e-9d4a-4e6b-8a7c-2d1e0f9b6a53\",\"risk_score\":35,\"user_hid\":null}}", + "body_base64": "eyJldmVudF90eXBlIjoiaWRlbnRpZmljYXRpb24uc2NvcmVkIiwic2NoZW1hX3ZlcnNpb24iOiIyMDI2LTA2LTAxIiwiY3JlYXRlZF9hdCI6IjIwMjYtMDktMzBUMTA6MTU6MDBaIiwiZGF0YSI6eyJyZXF1ZXN0X2lkIjoiM2YyYjhjMWUtOWQ0YS00ZTZiLThhN2MtMmQxZTBmOWI2YTUzIiwicmlza19zY29yZSI6MzUsInVzZXJfaGlkIjpudWxsfX0=", + "signature_header": "sha1=1c2d0251d83eba5ed62c98d3758b805d227022850cdb8ddb1aa3754524345453", + "valid": false, + "note": "Only sha256 is accepted" + }, + { + "name": "invalid_truncated_digest", + "secret": "whsec_test_4f1a9c2e7b5d3a8f6e0c1b2d9a7e5f3c", + "body": "{\"event_type\":\"identification.scored\",\"schema_version\":\"2026-06-01\",\"created_at\":\"2026-09-30T10:15:00Z\",\"data\":{\"request_id\":\"3f2b8c1e-9d4a-4e6b-8a7c-2d1e0f9b6a53\",\"risk_score\":35,\"user_hid\":null}}", + "body_base64": "eyJldmVudF90eXBlIjoiaWRlbnRpZmljYXRpb24uc2NvcmVkIiwic2NoZW1hX3ZlcnNpb24iOiIyMDI2LTA2LTAxIiwiY3JlYXRlZF9hdCI6IjIwMjYtMDktMzBUMTA6MTU6MDBaIiwiZGF0YSI6eyJyZXF1ZXN0X2lkIjoiM2YyYjhjMWUtOWQ0YS00ZTZiLThhN2MtMmQxZTBmOWI2YTUzIiwicmlza19zY29yZSI6MzUsInVzZXJfaGlkIjpudWxsfX0=", + "signature_header": "sha256=1c2d0251d83eba5ed62c98d3758b805d227022850cdb8ddb1aa37545243454", + "valid": false, + "note": "Digest must be 64 hex characters" + }, + { + "name": "invalid_non_hex_digest", + "secret": "whsec_test_4f1a9c2e7b5d3a8f6e0c1b2d9a7e5f3c", + "body": "{\"event_type\":\"identification.scored\",\"schema_version\":\"2026-06-01\",\"created_at\":\"2026-09-30T10:15:00Z\",\"data\":{\"request_id\":\"3f2b8c1e-9d4a-4e6b-8a7c-2d1e0f9b6a53\",\"risk_score\":35,\"user_hid\":null}}", + "body_base64": "eyJldmVudF90eXBlIjoiaWRlbnRpZmljYXRpb24uc2NvcmVkIiwic2NoZW1hX3ZlcnNpb24iOiIyMDI2LTA2LTAxIiwiY3JlYXRlZF9hdCI6IjIwMjYtMDktMzBUMTA6MTU6MDBaIiwiZGF0YSI6eyJyZXF1ZXN0X2lkIjoiM2YyYjhjMWUtOWQ0YS00ZTZiLThhN2MtMmQxZTBmOWI2YTUzIiwicmlza19zY29yZSI6MzUsInVzZXJfaGlkIjpudWxsfX0=", + "signature_header": "sha256=zzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzz", + "valid": false, + "note": "Digest must be hexadecimal" + }, + { + "name": "invalid_empty_header", + "secret": "whsec_test_4f1a9c2e7b5d3a8f6e0c1b2d9a7e5f3c", + "body": "{\"event_type\":\"identification.scored\",\"schema_version\":\"2026-06-01\",\"created_at\":\"2026-09-30T10:15:00Z\",\"data\":{\"request_id\":\"3f2b8c1e-9d4a-4e6b-8a7c-2d1e0f9b6a53\",\"risk_score\":35,\"user_hid\":null}}", + "body_base64": "eyJldmVudF90eXBlIjoiaWRlbnRpZmljYXRpb24uc2NvcmVkIiwic2NoZW1hX3ZlcnNpb24iOiIyMDI2LTA2LTAxIiwiY3JlYXRlZF9hdCI6IjIwMjYtMDktMzBUMTA6MTU6MDBaIiwiZGF0YSI6eyJyZXF1ZXN0X2lkIjoiM2YyYjhjMWUtOWQ0YS00ZTZiLThhN2MtMmQxZTBmOWI2YTUzIiwicmlza19zY29yZSI6MzUsInVzZXJfaGlkIjpudWxsfX0=", + "signature_header": "", + "valid": false, + "note": "Missing signature header" + }, + { + "name": "invalid_prefix_only", + "secret": "whsec_test_4f1a9c2e7b5d3a8f6e0c1b2d9a7e5f3c", + "body": "{\"event_type\":\"identification.scored\",\"schema_version\":\"2026-06-01\",\"created_at\":\"2026-09-30T10:15:00Z\",\"data\":{\"request_id\":\"3f2b8c1e-9d4a-4e6b-8a7c-2d1e0f9b6a53\",\"risk_score\":35,\"user_hid\":null}}", + "body_base64": "eyJldmVudF90eXBlIjoiaWRlbnRpZmljYXRpb24uc2NvcmVkIiwic2NoZW1hX3ZlcnNpb24iOiIyMDI2LTA2LTAxIiwiY3JlYXRlZF9hdCI6IjIwMjYtMDktMzBUMTA6MTU6MDBaIiwiZGF0YSI6eyJyZXF1ZXN0X2lkIjoiM2YyYjhjMWUtOWQ0YS00ZTZiLThhN2MtMmQxZTBmOWI2YTUzIiwicmlza19zY29yZSI6MzUsInVzZXJfaGlkIjpudWxsfX0=", + "signature_header": "sha256=", + "valid": false, + "note": "Prefix without digest" + }, + { + "name": "invalid_empty_secret", + "secret": "", + "body": "{\"event_type\":\"identification.scored\",\"schema_version\":\"2026-06-01\",\"created_at\":\"2026-09-30T10:15:00Z\",\"data\":{\"request_id\":\"3f2b8c1e-9d4a-4e6b-8a7c-2d1e0f9b6a53\",\"risk_score\":35,\"user_hid\":null}}", + "body_base64": "eyJldmVudF90eXBlIjoiaWRlbnRpZmljYXRpb24uc2NvcmVkIiwic2NoZW1hX3ZlcnNpb24iOiIyMDI2LTA2LTAxIiwiY3JlYXRlZF9hdCI6IjIwMjYtMDktMzBUMTA6MTU6MDBaIiwiZGF0YSI6eyJyZXF1ZXN0X2lkIjoiM2YyYjhjMWUtOWQ0YS00ZTZiLThhN2MtMmQxZTBmOWI2YTUzIiwicmlza19zY29yZSI6MzUsInVzZXJfaGlkIjpudWxsfX0=", + "signature_header": "sha256=68ae6c4becfcbc22bb017c96fec281973bc554f96afe7ba875854cd65b3368dc", + "valid": false, + "note": "An empty secret is a configuration error and must never verify" + }, + { + "name": "valid_real_ping_delivery", + "secret": "whsec_00112233445566778899aabbccddeeff", + "body": "{\"created_at\":\"2026-09-30T12:34:56Z\",\"event_type\":\"webhook.ping\",\"schema_version\":\"2026-06-01\"}", + "signature_header": "sha256=ea2685733d254f7028fb031c4214583b0650de01e6c8c93131236024edd9fdd8", + "valid": true, + "note": "Exact webhook.ping body and signature as produced by the ShieldLabs servers", + "body_base64": "eyJjcmVhdGVkX2F0IjoiMjAyNi0wOS0zMFQxMjozNDo1NloiLCJldmVudF90eXBlIjoid2ViaG9vay5waW5nIiwic2NoZW1hX3ZlcnNpb24iOiIyMDI2LTA2LTAxIn0=" + }, + { + "name": "valid_scored_with_escaped_ampersand", + "secret": "whsec_00112233445566778899aabbccddeeff", + "body": "{\"event_type\":\"identification.scored\",\"schema_version\":\"2026-06-01\",\"created_at\":\"2026-09-30T12:34:57.482913041Z\",\"data\":{\"request_id\":\"02f1d973-84db-4156-a7f7-e799e6bf389b\",\"visitor_id\":\"bde0e249-20d8-4544-838c-ed9a0b6d7a36\",\"device_id\":\"ac7c303d-971b-41d1-8e25-cd5b46b46aed\",\"session_id\":\"bde78778-efd2-4c49-952f-1f11b9c05f35\",\"cookie_id\":\"4449bb58-590c-444c-ae1f-d1ddc768dbdd\",\"user_hid\":\"9f86d081884c7d659a2feaa0c55ad015\",\"domain\":\"example.com\",\"public_ip\":{\"ip\":\"203.0.113.24\",\"country\":\"Netherlands\"},\"local_ip\":{\"ip\":\"198.51.100.23\",\"country\":\"Germany\"},\"connection_type\":\"proxy\",\"os\":\"Windows\",\"browser\":\"Chrome\",\"device_type\":\"desktop\",\"traffic_source\":{\"channel\":\"Google Ads\",\"referrer_domain\":\"google.com\",\"landing_url\":\"https://shop.example.com/signup?utm_source=google\\u0026utm_medium=cpc\\u0026gclid=abc123\",\"click_id_type\":\"gclid\",\"utm_source\":\"google\",\"utm_medium\":\"cpc\",\"utm_campaign\":\"\",\"utm_content\":\"\",\"utm_term\":\"\"},\"risk_score\":80,\"signals\":[{\"name\":\"proxy\",\"weight\":10},{\"name\":\"datacenter_ip\",\"weight\":10},{\"name\":\"antidetect_browser\",\"weight\":60}],\"detection_flags\":{\"vpn\":false,\"privacy_relay\":false,\"browser_vpn_proxy\":false,\"tor\":false,\"proxy\":true,\"datacenter_ip\":true,\"abuser\":false,\"os_mismatch\":false,\"os_not_detected\":false,\"timezone_mismatch\":false,\"anti_detect_browser\":true,\"browser_automation\":false,\"ip_mismatch\":true,\"incognito\":false,\"search_bot\":false,\"suspicious_paid_click\":true,\"javascript_disabled\":false,\"stun_not_checked\":false,\"check_incomplete\":false},\"observed_at\":\"2026-09-30T12:34:57.482913041Z\"}}", + "signature_header": "sha256=c4d44b7873625bdfda98cdb7a02d460f8492ca6a68fad8d147a5f30ad16f92a9", + "valid": true, + "note": "Server bodies escape & as \\u0026; verify the bytes as received", + "body_base64": "eyJldmVudF90eXBlIjoiaWRlbnRpZmljYXRpb24uc2NvcmVkIiwic2NoZW1hX3ZlcnNpb24iOiIyMDI2LTA2LTAxIiwiY3JlYXRlZF9hdCI6IjIwMjYtMDktMzBUMTI6MzQ6NTcuNDgyOTEzMDQxWiIsImRhdGEiOnsicmVxdWVzdF9pZCI6IjAyZjFkOTczLTg0ZGItNDE1Ni1hN2Y3LWU3OTllNmJmMzg5YiIsInZpc2l0b3JfaWQiOiJiZGUwZTI0OS0yMGQ4LTQ1NDQtODM4Yy1lZDlhMGI2ZDdhMzYiLCJkZXZpY2VfaWQiOiJhYzdjMzAzZC05NzFiLTQxZDEtOGUyNS1jZDViNDZiNDZhZWQiLCJzZXNzaW9uX2lkIjoiYmRlNzg3NzgtZWZkMi00YzQ5LTk1MmYtMWYxMWI5YzA1ZjM1IiwiY29va2llX2lkIjoiNDQ0OWJiNTgtNTkwYy00NDRjLWFlMWYtZDFkZGM3NjhkYmRkIiwidXNlcl9oaWQiOiI5Zjg2ZDA4MTg4NGM3ZDY1OWEyZmVhYTBjNTVhZDAxNSIsImRvbWFpbiI6ImV4YW1wbGUuY29tIiwicHVibGljX2lwIjp7ImlwIjoiMjAzLjAuMTEzLjI0IiwiY291bnRyeSI6Ik5ldGhlcmxhbmRzIn0sImxvY2FsX2lwIjp7ImlwIjoiMTk4LjUxLjEwMC4yMyIsImNvdW50cnkiOiJHZXJtYW55In0sImNvbm5lY3Rpb25fdHlwZSI6InByb3h5Iiwib3MiOiJXaW5kb3dzIiwiYnJvd3NlciI6IkNocm9tZSIsImRldmljZV90eXBlIjoiZGVza3RvcCIsInRyYWZmaWNfc291cmNlIjp7ImNoYW5uZWwiOiJHb29nbGUgQWRzIiwicmVmZXJyZXJfZG9tYWluIjoiZ29vZ2xlLmNvbSIsImxhbmRpbmdfdXJsIjoiaHR0cHM6Ly9zaG9wLmV4YW1wbGUuY29tL3NpZ251cD91dG1fc291cmNlPWdvb2dsZVx1MDAyNnV0bV9tZWRpdW09Y3BjXHUwMDI2Z2NsaWQ9YWJjMTIzIiwiY2xpY2tfaWRfdHlwZSI6ImdjbGlkIiwidXRtX3NvdXJjZSI6Imdvb2dsZSIsInV0bV9tZWRpdW0iOiJjcGMiLCJ1dG1fY2FtcGFpZ24iOiIiLCJ1dG1fY29udGVudCI6IiIsInV0bV90ZXJtIjoiIn0sInJpc2tfc2NvcmUiOjgwLCJzaWduYWxzIjpbeyJuYW1lIjoicHJveHkiLCJ3ZWlnaHQiOjEwfSx7Im5hbWUiOiJkYXRhY2VudGVyX2lwIiwid2VpZ2h0IjoxMH0seyJuYW1lIjoiYW50aWRldGVjdF9icm93c2VyIiwid2VpZ2h0Ijo2MH1dLCJkZXRlY3Rpb25fZmxhZ3MiOnsidnBuIjpmYWxzZSwicHJpdmFjeV9yZWxheSI6ZmFsc2UsImJyb3dzZXJfdnBuX3Byb3h5IjpmYWxzZSwidG9yIjpmYWxzZSwicHJveHkiOnRydWUsImRhdGFjZW50ZXJfaXAiOnRydWUsImFidXNlciI6ZmFsc2UsIm9zX21pc21hdGNoIjpmYWxzZSwib3Nfbm90X2RldGVjdGVkIjpmYWxzZSwidGltZXpvbmVfbWlzbWF0Y2giOmZhbHNlLCJhbnRpX2RldGVjdF9icm93c2VyIjp0cnVlLCJicm93c2VyX2F1dG9tYXRpb24iOmZhbHNlLCJpcF9taXNtYXRjaCI6dHJ1ZSwiaW5jb2duaXRvIjpmYWxzZSwic2VhcmNoX2JvdCI6ZmFsc2UsInN1c3BpY2lvdXNfcGFpZF9jbGljayI6dHJ1ZSwiamF2YXNjcmlwdF9kaXNhYmxlZCI6ZmFsc2UsInN0dW5fbm90X2NoZWNrZWQiOmZhbHNlLCJjaGVja19pbmNvbXBsZXRlIjpmYWxzZX0sIm9ic2VydmVkX2F0IjoiMjAyNi0wOS0zMFQxMjozNDo1Ny40ODI5MTMwNDFaIn19" + }, + { + "name": "invalid_scored_unescaped_reserialization", + "secret": "whsec_00112233445566778899aabbccddeeff", + "body": "{\"event_type\":\"identification.scored\",\"schema_version\":\"2026-06-01\",\"created_at\":\"2026-09-30T12:34:57.482913041Z\",\"data\":{\"request_id\":\"02f1d973-84db-4156-a7f7-e799e6bf389b\",\"visitor_id\":\"bde0e249-20d8-4544-838c-ed9a0b6d7a36\",\"device_id\":\"ac7c303d-971b-41d1-8e25-cd5b46b46aed\",\"session_id\":\"bde78778-efd2-4c49-952f-1f11b9c05f35\",\"cookie_id\":\"4449bb58-590c-444c-ae1f-d1ddc768dbdd\",\"user_hid\":\"9f86d081884c7d659a2feaa0c55ad015\",\"domain\":\"example.com\",\"public_ip\":{\"ip\":\"203.0.113.24\",\"country\":\"Netherlands\"},\"local_ip\":{\"ip\":\"198.51.100.23\",\"country\":\"Germany\"},\"connection_type\":\"proxy\",\"os\":\"Windows\",\"browser\":\"Chrome\",\"device_type\":\"desktop\",\"traffic_source\":{\"channel\":\"Google Ads\",\"referrer_domain\":\"google.com\",\"landing_url\":\"https://shop.example.com/signup?utm_source=google&utm_medium=cpc&gclid=abc123\",\"click_id_type\":\"gclid\",\"utm_source\":\"google\",\"utm_medium\":\"cpc\",\"utm_campaign\":\"\",\"utm_content\":\"\",\"utm_term\":\"\"},\"risk_score\":80,\"signals\":[{\"name\":\"proxy\",\"weight\":10},{\"name\":\"datacenter_ip\",\"weight\":10},{\"name\":\"antidetect_browser\",\"weight\":60}],\"detection_flags\":{\"vpn\":false,\"privacy_relay\":false,\"browser_vpn_proxy\":false,\"tor\":false,\"proxy\":true,\"datacenter_ip\":true,\"abuser\":false,\"os_mismatch\":false,\"os_not_detected\":false,\"timezone_mismatch\":false,\"anti_detect_browser\":true,\"browser_automation\":false,\"ip_mismatch\":true,\"incognito\":false,\"search_bot\":false,\"suspicious_paid_click\":true,\"javascript_disabled\":false,\"stun_not_checked\":false,\"check_incomplete\":false},\"observed_at\":\"2026-09-30T12:34:57.482913041Z\"}}", + "signature_header": "sha256=c4d44b7873625bdfda98cdb7a02d460f8492ca6a68fad8d147a5f30ad16f92a9", + "valid": false, + "note": "Re-serializing (unescaping \\u0026) changes the bytes", + "body_base64": "eyJldmVudF90eXBlIjoiaWRlbnRpZmljYXRpb24uc2NvcmVkIiwic2NoZW1hX3ZlcnNpb24iOiIyMDI2LTA2LTAxIiwiY3JlYXRlZF9hdCI6IjIwMjYtMDktMzBUMTI6MzQ6NTcuNDgyOTEzMDQxWiIsImRhdGEiOnsicmVxdWVzdF9pZCI6IjAyZjFkOTczLTg0ZGItNDE1Ni1hN2Y3LWU3OTllNmJmMzg5YiIsInZpc2l0b3JfaWQiOiJiZGUwZTI0OS0yMGQ4LTQ1NDQtODM4Yy1lZDlhMGI2ZDdhMzYiLCJkZXZpY2VfaWQiOiJhYzdjMzAzZC05NzFiLTQxZDEtOGUyNS1jZDViNDZiNDZhZWQiLCJzZXNzaW9uX2lkIjoiYmRlNzg3NzgtZWZkMi00YzQ5LTk1MmYtMWYxMWI5YzA1ZjM1IiwiY29va2llX2lkIjoiNDQ0OWJiNTgtNTkwYy00NDRjLWFlMWYtZDFkZGM3NjhkYmRkIiwidXNlcl9oaWQiOiI5Zjg2ZDA4MTg4NGM3ZDY1OWEyZmVhYTBjNTVhZDAxNSIsImRvbWFpbiI6ImV4YW1wbGUuY29tIiwicHVibGljX2lwIjp7ImlwIjoiMjAzLjAuMTEzLjI0IiwiY291bnRyeSI6Ik5ldGhlcmxhbmRzIn0sImxvY2FsX2lwIjp7ImlwIjoiMTk4LjUxLjEwMC4yMyIsImNvdW50cnkiOiJHZXJtYW55In0sImNvbm5lY3Rpb25fdHlwZSI6InByb3h5Iiwib3MiOiJXaW5kb3dzIiwiYnJvd3NlciI6IkNocm9tZSIsImRldmljZV90eXBlIjoiZGVza3RvcCIsInRyYWZmaWNfc291cmNlIjp7ImNoYW5uZWwiOiJHb29nbGUgQWRzIiwicmVmZXJyZXJfZG9tYWluIjoiZ29vZ2xlLmNvbSIsImxhbmRpbmdfdXJsIjoiaHR0cHM6Ly9zaG9wLmV4YW1wbGUuY29tL3NpZ251cD91dG1fc291cmNlPWdvb2dsZSZ1dG1fbWVkaXVtPWNwYyZnY2xpZD1hYmMxMjMiLCJjbGlja19pZF90eXBlIjoiZ2NsaWQiLCJ1dG1fc291cmNlIjoiZ29vZ2xlIiwidXRtX21lZGl1bSI6ImNwYyIsInV0bV9jYW1wYWlnbiI6IiIsInV0bV9jb250ZW50IjoiIiwidXRtX3Rlcm0iOiIifSwicmlza19zY29yZSI6ODAsInNpZ25hbHMiOlt7Im5hbWUiOiJwcm94eSIsIndlaWdodCI6MTB9LHsibmFtZSI6ImRhdGFjZW50ZXJfaXAiLCJ3ZWlnaHQiOjEwfSx7Im5hbWUiOiJhbnRpZGV0ZWN0X2Jyb3dzZXIiLCJ3ZWlnaHQiOjYwfV0sImRldGVjdGlvbl9mbGFncyI6eyJ2cG4iOmZhbHNlLCJwcml2YWN5X3JlbGF5IjpmYWxzZSwiYnJvd3Nlcl92cG5fcHJveHkiOmZhbHNlLCJ0b3IiOmZhbHNlLCJwcm94eSI6dHJ1ZSwiZGF0YWNlbnRlcl9pcCI6dHJ1ZSwiYWJ1c2VyIjpmYWxzZSwib3NfbWlzbWF0Y2giOmZhbHNlLCJvc19ub3RfZGV0ZWN0ZWQiOmZhbHNlLCJ0aW1lem9uZV9taXNtYXRjaCI6ZmFsc2UsImFudGlfZGV0ZWN0X2Jyb3dzZXIiOnRydWUsImJyb3dzZXJfYXV0b21hdGlvbiI6ZmFsc2UsImlwX21pc21hdGNoIjp0cnVlLCJpbmNvZ25pdG8iOmZhbHNlLCJzZWFyY2hfYm90IjpmYWxzZSwic3VzcGljaW91c19wYWlkX2NsaWNrIjp0cnVlLCJqYXZhc2NyaXB0X2Rpc2FibGVkIjpmYWxzZSwic3R1bl9ub3RfY2hlY2tlZCI6ZmFsc2UsImNoZWNrX2luY29tcGxldGUiOmZhbHNlfSwib2JzZXJ2ZWRfYXQiOiIyMDI2LTA5LTMwVDEyOjM0OjU3LjQ4MjkxMzA0MVoifX0=" + }, + { + "name": "valid_rotation_second_secret_matches", + "secrets": [ + "whsec_old_secret_value_0000000000", + "whsec_00112233445566778899aabbccddeeff" + ], + "body": "{\"created_at\":\"2026-09-30T12:34:56Z\",\"event_type\":\"webhook.ping\",\"schema_version\":\"2026-06-01\"}", + "signature_header": "sha256=ea2685733d254f7028fb031c4214583b0650de01e6c8c93131236024edd9fdd8", + "valid": true, + "note": "verify accepts a list of secrets during rotation; valid when any matches", + "body_base64": "eyJjcmVhdGVkX2F0IjoiMjAyNi0wOS0zMFQxMjozNDo1NloiLCJldmVudF90eXBlIjoid2ViaG9vay5waW5nIiwic2NoZW1hX3ZlcnNpb24iOiIyMDI2LTA2LTAxIn0=" + }, + { + "name": "invalid_rotation_no_secret_matches", + "secrets": [ + "whsec_old_secret_value_0000000000", + "whsec_other_secret_value_11111111" + ], + "body": "{\"created_at\":\"2026-09-30T12:34:56Z\",\"event_type\":\"webhook.ping\",\"schema_version\":\"2026-06-01\"}", + "signature_header": "sha256=ea2685733d254f7028fb031c4214583b0650de01e6c8c93131236024edd9fdd8", + "valid": false, + "note": "No secret in the list matches", + "body_base64": "eyJjcmVhdGVkX2F0IjoiMjAyNi0wOS0zMFQxMjozNDo1NloiLCJldmVudF90eXBlIjoid2ViaG9vay5waW5nIiwic2NoZW1hX3ZlcnNpb24iOiIyMDI2LTA2LTAxIn0=" + }, + { + "name": "invalid_secret_without_prefix", + "secret": "00112233445566778899aabbccddeeff", + "body": "{\"created_at\":\"2026-09-30T12:34:56Z\",\"event_type\":\"webhook.ping\",\"schema_version\":\"2026-06-01\"}", + "signature_header": "sha256=ea2685733d254f7028fb031c4214583b0650de01e6c8c93131236024edd9fdd8", + "valid": false, + "note": "The key is the full secret string including whsec_; stripping the prefix breaks verification", + "body_base64": "eyJjcmVhdGVkX2F0IjoiMjAyNi0wOS0zMFQxMjozNDo1NloiLCJldmVudF90eXBlIjoid2ViaG9vay5waW5nIiwic2NoZW1hX3ZlcnNpb24iOiIyMDI2LTA2LTAxIn0=" + } + ] +} diff --git a/tests/data/webhook-test-delivery.json b/tests/data/webhook-test-delivery.json new file mode 100644 index 0000000..8e214ce --- /dev/null +++ b/tests/data/webhook-test-delivery.json @@ -0,0 +1,72 @@ +{ + "created_at": "2026-09-30T12:34:56Z", + "data": { + "browser": "Chrome", + "connection_type": "proxy", + "cookie_id": "2c9d1e8f-4b7a-4c3e-9d2f-1a8b7c6d5e4f", + "detection_flags": { + "abuser": true, + "anti_detect_browser": false, + "browser_vpn_proxy": false, + "check_incomplete": false, + "datacenter_ip": true, + "incognito": false, + "ip_mismatch": false, + "javascript_disabled": false, + "os_mismatch": false, + "os_not_detected": false, + "privacy_relay": false, + "proxy": true, + "stun_not_checked": false, + "suspicious_paid_click": false, + "timezone_mismatch": false, + "tor": false, + "vpn": false + }, + "device_id": "6f1e2d3c-4b5a-5968-8776-655443322110", + "device_type": "desktop", + "domain": "example.com", + "local_ip": { + "country": "BY", + "ip": "198.51.100.10" + }, + "observed_at": "2026-09-30T12:34:56Z", + "os": "Windows", + "public_ip": { + "country": "BY", + "ip": "203.0.113.10" + }, + "request_id": "13f84f05-7c2a-4e9b-9f1d-2a6b8c0e4d11", + "risk_score": 30, + "session_id": "3a2b1c0d-9e8f-4a7b-8c6d-5e4f3a2b1c0d", + "signals": [ + { + "name": "proxy", + "weight": 10 + }, + { + "name": "datacenter_ip", + "weight": 10 + }, + { + "name": "abuser", + "weight": 10 + } + ], + "traffic_source": { + "channel": "Direct", + "click_id_type": "", + "landing_url": "https://example.com/", + "referrer_domain": "", + "utm_campaign": "", + "utm_content": "", + "utm_medium": "", + "utm_source": "", + "utm_term": "" + }, + "user_hid": null, + "visitor_id": "7a6b5c4d-3e2f-5a1b-9c8d-7e6f5a4b3c2d" + }, + "event_type": "identification.scored", + "schema_version": "2026-06-01" +} diff --git a/tests/test_async.py b/tests/test_async.py new file mode 100644 index 0000000..75f703f --- /dev/null +++ b/tests/test_async.py @@ -0,0 +1,199 @@ +"""AsyncShieldLabs and AsyncShieldLabsManagement mirror the sync clients.""" + +from __future__ import annotations + +from typing import Any + +import httpx +import pytest +import respx + +from _support import ( + API_KEY, + DOMAIN, + HISTORY_HOST, + MANAGEMENT_HOST, + REQUEST_ID, + REQUEST_PATH, + SECRET_KEY, + FakeTime, + empty_page, + history_body, + load_json, + row, + uuid_for, +) +from shieldlabs import ( + APIConnectionError, + APITimeoutError, + AsyncShieldLabs, + AsyncShieldLabsManagement, + HistoryPage, + RateLimitError, + ServerError, + ValidationError, +) + +pytestmark = pytest.mark.anyio + + +def _found() -> httpx.Response: + page = load_json("history-page.json") + return httpx.Response(200, json={"data": page["data"][:1], "total": 1}) + + +@pytest.fixture +def mock() -> Any: + with respx.mock(assert_all_called=False) as router: + yield router + + +@pytest.fixture +def client(fake_time: FakeTime) -> AsyncShieldLabs: + instance = AsyncShieldLabs(api_key=API_KEY) + fake_time.install(instance._transport) + return instance + + +async def test_search(client: AsyncShieldLabs, mock: Any) -> None: + route = mock.get(host=HISTORY_HOST, path=REQUEST_PATH).respond( + 200, json=load_json("history-page.json") + ) + page = await client.history.search("request_id", REQUEST_ID, limit=3) + assert isinstance(page, HistoryPage) + assert len(page.data) == 5 + assert route.calls.last.request.url.params["limit"] == "3" + assert route.calls.last.request.headers["authorization"] == f"Bearer {API_KEY}" + + +async def test_search_validation_sends_nothing(client: AsyncShieldLabs, mock: Any) -> None: + route = mock.get(host=HISTORY_HOST).respond(200, json=history_body()) + with pytest.raises(ValidationError): + await client.history.search("user_hid", "") + with pytest.raises(ValidationError): + await client.history.search("request_id", REQUEST_ID, limit=101) + with pytest.raises(ValidationError): + client.history.iter("ip", "2001:db8::1") + with pytest.raises(ValidationError): + await client.identifications.get("nope") + assert route.call_count == 0 + + +async def test_iter_dedupes_and_stops(client: AsyncShieldLabs, mock: Any) -> None: + route = mock.get(host=HISTORY_HOST).mock( + side_effect=[ + httpx.Response(200, json=history_body(row(uuid_for(1)), row(uuid_for(2)), total=4)), + httpx.Response(200, json=history_body(row(uuid_for(2)), row(uuid_for(3)), total=4)), + ] + ) + items = [item async for item in client.history.iter("user_hid", "acct-1", page_size=2)] + assert [item.request_id for item in items] == [uuid_for(1), uuid_for(2), uuid_for(3)] + assert route.call_count == 2 + + +async def test_iter_empty_page_and_max_items(client: AsyncShieldLabs, mock: Any) -> None: + route = mock.get(host=HISTORY_HOST) + route.side_effect = [httpx.Response(200, json=history_body(total=9))] + assert [item async for item in client.history.iter("user_hid", "acct-1")] == [] + + route.side_effect = [ + httpx.Response(200, json=history_body(*(row(uuid_for(i)) for i in range(5)), total=9)) + ] + items = [item async for item in client.history.iter("user_hid", "acct-1", max_items=2)] + assert len(items) == 2 + assert [item async for item in client.history.iter("user_hid", "a", max_items=0)] == [] + + +# identifications.get waiting rules run against both clients in test_polling.py. + + +async def test_user_hid_escaping_and_unsearchable_values( + client: AsyncShieldLabs, mock: Any +) -> None: + route = mock.get(host=HISTORY_HOST).respond(200, json=history_body()) + await client.history.search("user_hid", "a@b c!") + raw_path = route.calls.last.request.url.raw_path.decode("ascii") + assert raw_path == "/api/v1/history/user_hid/a@b%20c%21?limit=20&offset=0" + await client.history.search("user_hid", "-._~$&+,:;=@ %?#'()*é") + raw_path = route.calls.last.request.url.raw_path.decode("ascii") + assert raw_path == ( + "/api/v1/history/user_hid/-._~$&+,:;=@%20%25%3F%23%27%28%29%2A%C3%A9?limit=20&offset=0" + ) + for value in (".", "..", "a/b", "/"): + with pytest.raises(ValidationError): + await client.history.search("user_hid", value) + with pytest.raises(ValidationError): + client.history.iter("user_hid", value) + assert route.call_count == 2 + + +async def test_retries_and_errors(client: AsyncShieldLabs, mock: Any, fake_time: FakeTime) -> None: + route = mock.get(host=HISTORY_HOST) + route.side_effect = [httpx.Response(500), httpx.Response(200, json=history_body())] + assert (await client.history.search("request_id", REQUEST_ID)).total == 0 + assert fake_time.sleeps == [0.5] + + route.side_effect = httpx.ConnectError("refused") + with pytest.raises(APIConnectionError): + await client.history.search("request_id", REQUEST_ID) + + route.side_effect = httpx.ConnectTimeout("slow") + with pytest.raises(APITimeoutError): + await client.history.search("request_id", REQUEST_ID) + + route.side_effect = lambda request: httpx.Response(502, text="") + with pytest.raises(ServerError): + await client.history.search("request_id", REQUEST_ID) + + +async def test_context_manager_and_injected_client(mock: Any) -> None: + async with AsyncShieldLabs(api_key=API_KEY) as client: + inner = client._transport.client + assert repr(client) == "AsyncShieldLabs(base_url='https://account.shieldlabs.ai')" + assert inner.is_closed + + route = mock.get(host=HISTORY_HOST).respond(200, json=history_body()) + http_client = httpx.AsyncClient() + async with AsyncShieldLabs(api_key=API_KEY, http_client=http_client) as injected: + await injected.history.search("request_id", REQUEST_ID) + assert not http_client.is_closed + assert route.call_count == 1 + await http_client.aclose() + + +async def test_async_management(mock: Any, fake_time: FakeTime) -> None: + route = mock.get(host=MANAGEMENT_HOST, path="/v1/profile").respond( + 200, json=load_json("management-profile.json") + ) + async with AsyncShieldLabsManagement(secret_key=SECRET_KEY, domain="www.example.com") as client: + fake_time.install(client._transport) + profile = await client.get_profile() + assert profile.to_dict() == load_json("management-profile-expected.json") + assert route.calls.last.request.headers["x-shield-domain"] == DOMAIN + + +async def test_async_management_never_retries_429(mock: Any, fake_time: FakeTime) -> None: + route = mock.get(host=MANAGEMENT_HOST).respond(429, json={"error": "too many requests"}) + client = AsyncShieldLabsManagement(secret_key=SECRET_KEY, domain=DOMAIN, max_retries=5) + fake_time.install(client._transport) + with pytest.raises(RateLimitError): + await client.get_profile() + assert route.call_count == 1 + await client.aclose() + + +async def test_async_management_injected_client(mock: Any) -> None: + http_client = httpx.AsyncClient() + client = AsyncShieldLabsManagement( + secret_key=SECRET_KEY, domain=DOMAIN, http_client=http_client + ) + await client.aclose() + assert not http_client.is_closed + await http_client.aclose() + + +async def test_real_async_sleep_is_used_by_default(mock: Any) -> None: + mock.get(host=HISTORY_HOST).mock(side_effect=[empty_page(), _found()]) + async with AsyncShieldLabs(api_key=API_KEY) as client: + found = await client.identifications.get(REQUEST_ID, poll_interval=0.01) + assert found is not None diff --git a/tests/test_errors.py b/tests/test_errors.py new file mode 100644 index 0000000..6d902b0 --- /dev/null +++ b/tests/test_errors.py @@ -0,0 +1,259 @@ +"""Error mapping and retries, driven by the shared error-responses fixture.""" + +from __future__ import annotations + +from datetime import datetime, timedelta, timezone +from email.utils import format_datetime +from typing import Any + +import httpx +import pytest +import respx + +import shieldlabs +from _support import ( + API_KEY, + DOMAIN, + HISTORY_HOST, + MANAGEMENT_HOST, + REQUEST_ID, + SECRET_KEY, + FakeTime, + load_json, +) +from shieldlabs import ( + APIConnectionError, + ApiError, + APITimeoutError, + RateLimitError, + ServerError, + ShieldLabs, + ShieldLabsError, + ShieldLabsManagement, +) +from shieldlabs._http import backoff_delay, error_from_response, parse_retry_after + +CASES = load_json("error-responses.json")["cases"] + + +def _response(case: dict[str, Any]) -> httpx.Response: + headers = {"content-type": case["content_type"]} if case["content_type"] else {} + return httpx.Response(case["status"], content=case["body"].encode("utf-8"), headers=headers) + + +@pytest.fixture +def mock() -> Any: + with respx.mock(assert_all_called=False) as router: + yield router + + +def _call(surface: str, fake_time: FakeTime, max_retries: int = 1) -> None: + if surface == "history": + client = ShieldLabs(api_key=API_KEY, max_retries=max_retries) + fake_time.install(client._transport) + client.history.search("request_id", REQUEST_ID) + else: + management = ShieldLabsManagement( + secret_key=SECRET_KEY, domain=DOMAIN, max_retries=max_retries + ) + fake_time.install(management._transport) + management.get_profile() + + +@pytest.mark.parametrize( + "case", + CASES, + ids=[f"{c['surface']}-{c['status']}-{i}" for i, c in enumerate(CASES)], +) +def test_error_responses_fixture(case: dict[str, Any], mock: Any, fake_time: FakeTime) -> None: + host = HISTORY_HOST if case["surface"] == "history" else MANAGEMENT_HOST + route = mock.get(host=host).mock(side_effect=lambda request: _response(case)) + expected_class = getattr(shieldlabs, case["expected_error"]) + + with pytest.raises(expected_class) as caught: + _call(case["surface"], fake_time) + + error = caught.value + assert type(error) is expected_class + assert isinstance(error, ApiError) + assert isinstance(error, ShieldLabsError) + assert error.status == case["status"] + assert f"(HTTP {case['status']})" in str(error) + assert route.call_count == (2 if case["retry"] else 1) + if case["content_type"]: + assert error.headers["content-type"] == case["content_type"] + + +def test_error_messages_are_parsed_defensively() -> None: + def build(status: int, body: bytes, content_type: str = "text/plain") -> ApiError: + response = httpx.Response(status, content=body, headers={"content-type": content_type}) + return error_from_response(response) + + assert build(401, b'{"error":"invalid api key"}\n').message == "invalid api key" + assert build(401, b'{"error":"invalid api key"}\n').body == {"error": "invalid api key"} + assert build(400, b"null").message == "Bad Request" + assert build(400, b"null").body is None + assert build(400, b'"fail parse uuid"').message == "fail parse uuid" + assert build(401, b"").message == "Unauthorized" + assert build(401, b"").body is None + assert build(404, b"404 page not found").message == "404 page not found" + assert build(502, b"

502

", "text/html").message == "Bad Gateway" + assert build(502, b"

502

").body.startswith("") + assert build(503, b'{"message":"busy"}').message == "busy" + assert build(500, b'{"other":1}').message == "Internal Server Error" + assert build(500, b"[1,2]").body == [1, 2] + assert build(500, b"\xff\xfe garbage").message is not None + assert build(599, b"").message == "Unexpected response" + unknown = build(418, b"teapot") + assert type(unknown) is ApiError + assert unknown.message == "teapot" + + +def test_rate_limit_error_carries_retry_after() -> None: + response = httpx.Response( + 429, json={"error": "too many requests"}, headers={"Retry-After": "3"} + ) + error = error_from_response(response) + assert isinstance(error, RateLimitError) + assert error.retry_after == 3.0 + assert error_from_response(httpx.Response(429)).retry_after is None + + +def test_parse_retry_after_formats() -> None: + assert parse_retry_after(None) is None + assert parse_retry_after("") is None + assert parse_retry_after(" 2.5 ") == 2.5 + for malformed in ("-5", "-0", "+5", "1e3", "5_0", "0x10", "\u0665"): + assert parse_retry_after(malformed) is None, malformed + assert parse_retry_after("nan") is None + assert parse_retry_after("soon") is None + future = datetime.now(timezone.utc) + timedelta(seconds=30) + seconds = parse_retry_after(format_datetime(future, usegmt=True)) + assert seconds is not None + assert 25 <= seconds <= 31 + naive = format_datetime(future.replace(tzinfo=None)) + assert parse_retry_after(naive) is not None + past = datetime.now(timezone.utc) - timedelta(hours=1) + assert parse_retry_after(format_datetime(past, usegmt=True)) == 0.0 + + +def test_backoff_delay_has_jitter_and_cap() -> None: + assert backoff_delay(0, lambda: 0.0) == 0.25 + assert backoff_delay(0, lambda: 1.0) == 0.5 + assert backoff_delay(1, lambda: 1.0) == 1.0 + assert backoff_delay(3, lambda: 1.0) == 4.0 + assert backoff_delay(10, lambda: 1.0) == 8.0 + assert backoff_delay(10, lambda: 0.0) == 4.0 + + +def test_retries_use_backoff_then_raise(mock: Any, fake_time: FakeTime) -> None: + route = mock.get(host=HISTORY_HOST).mock(side_effect=lambda request: httpx.Response(503)) + client = ShieldLabs(api_key=API_KEY, max_retries=3) + fake_time.install(client._transport) + with pytest.raises(ServerError): + client.history.search("request_id", REQUEST_ID) + assert route.call_count == 4 + assert fake_time.sleeps == [0.5, 1.0, 2.0] + + +def test_rate_limit_without_retry_after_waits_at_least_one_second( + mock: Any, fake_time: FakeTime +) -> None: + route = mock.get(host=HISTORY_HOST).mock(side_effect=lambda request: httpx.Response(429)) + client = ShieldLabs(api_key=API_KEY, max_retries=3) + fake_time.install(client._transport) + with pytest.raises(RateLimitError): + client.history.search("request_id", REQUEST_ID) + assert route.call_count == 4 + assert fake_time.sleeps == [1.0, 1.0, 2.0] # backoff 0.5, 1, 2 with a 1 s floor + + +def test_retry_after_is_honoured_and_capped(mock: Any, fake_time: FakeTime) -> None: + route = mock.get(host=HISTORY_HOST).mock( + side_effect=[ + httpx.Response(429, headers={"Retry-After": "3"}), + httpx.Response(503, headers={"Retry-After": "120"}), + httpx.Response(200, json={"data": [], "total": 0}), + ] + ) + client = ShieldLabs(api_key=API_KEY, max_retries=2) + fake_time.install(client._transport) + page = client.history.search("request_id", REQUEST_ID) + assert page.total == 0 + assert route.call_count == 3 + assert fake_time.sleeps == [3.0, 10.0] + + +def test_retry_after_is_followed_as_sent_even_below_one_second( + mock: Any, fake_time: FakeTime +) -> None: + # Only the wait of identifications.get raises the delay after a 429 to at least 1 s. + route = mock.get(host=HISTORY_HOST).mock( + side_effect=[ + httpx.Response(429, headers={"Retry-After": "0"}), + httpx.Response(429, headers={"Retry-After": "0.5"}), + httpx.Response(200, json={"data": [], "total": 0}), + ] + ) + client = ShieldLabs(api_key=API_KEY, max_retries=2) + fake_time.install(client._transport) + assert client.history.search("request_id", REQUEST_ID).total == 0 + assert route.call_count == 3 + assert fake_time.sleeps == [0.0, 0.5] + + +@pytest.mark.parametrize("status", [400, 401, 402, 403, 404]) +def test_client_errors_are_never_retried(mock: Any, fake_time: FakeTime, status: int) -> None: + route = mock.get(host=HISTORY_HOST).mock(side_effect=lambda request: httpx.Response(status)) + client = ShieldLabs(api_key=API_KEY, max_retries=5) + fake_time.install(client._transport) + with pytest.raises(ApiError): + client.history.search("request_id", REQUEST_ID) + assert route.call_count == 1 + assert fake_time.sleeps == [] + + +def test_connection_errors_are_retried(mock: Any, fake_time: FakeTime) -> None: + route = mock.get(host=HISTORY_HOST).mock( + side_effect=[httpx.ConnectError("refused"), httpx.Response(200, json={"data": []})] + ) + client = ShieldLabs(api_key=API_KEY) + fake_time.install(client._transport) + assert client.history.search("request_id", REQUEST_ID).total == 0 + assert route.call_count == 2 + assert fake_time.sleeps == [0.5] + + +def test_connection_error_after_retries(mock: Any, fake_time: FakeTime) -> None: + mock.get(host=HISTORY_HOST).mock(side_effect=httpx.ConnectError("refused")) + client = ShieldLabs(api_key=API_KEY, max_retries=1) + fake_time.install(client._transport) + with pytest.raises(APIConnectionError, match="refused") as caught: + client.history.search("request_id", REQUEST_ID) + assert isinstance(caught.value.__cause__, httpx.ConnectError) + assert API_KEY not in str(caught.value) + + +def test_timeouts_are_retried_then_raised(mock: Any, fake_time: FakeTime) -> None: + route = mock.get(host=HISTORY_HOST).mock(side_effect=httpx.ReadTimeout("slow")) + client = ShieldLabs(api_key=API_KEY, timeout=2.5, max_retries=2) + fake_time.install(client._transport) + with pytest.raises(APITimeoutError, match=r"2\.5 s"): + client.history.search("request_id", REQUEST_ID) + assert route.call_count == 3 + assert fake_time.sleeps == [0.5, 1.0] + + +def test_max_retries_zero(mock: Any, fake_time: FakeTime) -> None: + route = mock.get(host=HISTORY_HOST).mock(side_effect=lambda request: httpx.Response(500)) + client = ShieldLabs(api_key=API_KEY, max_retries=0) + fake_time.install(client._transport) + with pytest.raises(ServerError): + client.history.search("request_id", REQUEST_ID) + assert route.call_count == 1 + + +def test_timeout_is_passed_to_each_attempt(mock: Any) -> None: + route = mock.get(host=HISTORY_HOST).respond(200, json={"data": [], "total": 0}) + ShieldLabs(api_key=API_KEY, timeout=4).history.search("request_id", REQUEST_ID) + assert route.calls.last.request.extensions["timeout"]["read"] == 4.0 diff --git a/tests/test_example_app.py b/tests/test_example_app.py new file mode 100644 index 0000000..0f8cbba --- /dev/null +++ b/tests/test_example_app.py @@ -0,0 +1,129 @@ +"""Smoke test of examples/fastapi_app.py. Skipped unless FastAPI is installed. + +Run it with: + pip install -e ".[dev]" -r examples/requirements.txt + pytest tests/test_example_app.py +""" + +from __future__ import annotations + +import importlib.util +import json +from collections.abc import Iterator +from datetime import datetime, timezone +from pathlib import Path +from types import ModuleType +from typing import Any + +import httpx +import pytest +import respx + +from _support import API_KEY, HISTORY_HOST, FakeTime, load_json, sign + +pytest.importorskip("fastapi") +from fastapi.testclient import TestClient + +EXAMPLE = Path(__file__).resolve().parent.parent / "examples" / "fastapi_app.py" +SECRET = "whsec_your_signing_secret" +REQUEST_ID = "3f2b8c1e-9d4a-4e6b-8a7c-2d1e0f9b6a53" + + +def _now_history_time() -> str: + return datetime.now(timezone.utc).strftime("%Y-%m-%d %H:%M:%S.000") + + +def _row(**changes: Any) -> dict[str, Any]: + row = dict(load_json("history-page.json")["data"][1]) # trusted, anonymous + row.update(request_id=REQUEST_ID, created_at=_now_history_time()) + row.update(changes) + return row + + +def _load_example() -> ModuleType: + spec = importlib.util.spec_from_file_location("shieldlabs_fastapi_example", EXAMPLE) + assert spec is not None + assert spec.loader is not None + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +@pytest.fixture +def app_client(monkeypatch: pytest.MonkeyPatch) -> Iterator[tuple[TestClient, Any]]: + monkeypatch.setenv("SHIELDLABS_API_KEY", API_KEY) + # Two secrets while one is rotated; the space after the comma must not break verification. + monkeypatch.setenv("SHIELDLABS_WEBHOOK_SECRET", f"whsec_previous, {SECRET} ,") + module = _load_example() + with respx.mock(assert_all_called=False) as router, TestClient(module.app) as client: + FakeTime().install(module.app.state.shieldlabs._transport) + yield client, router + + +def _signup(client: TestClient, request_id: str = REQUEST_ID) -> httpx.Response: + return client.post("/signup", json={"email": "user@example.com", "requestId": request_id}) + + +def test_trusted_signup_is_created_once(app_client: tuple[TestClient, Any]) -> None: + client, router = app_client + router.get(host=HISTORY_HOST).respond(200, json={"data": [_row()], "total": 1}) + assert _signup(client).status_code == 201 + replay = _signup(client) + assert replay.status_code == 403 + assert replay.json()["reason"] == "replayed" + + +@pytest.mark.parametrize( + ("changes", "reason"), + [ + ({"score": 80}, "blocked_band"), + ({"is_browser_automation": True}, "blocked_flag"), + ({"score": 999}, "rate_limited"), + ({"created_at": "2026-01-01 00:00:00.000"}, "stale"), + ({"device_id": "00000000-0000-0000-0000-000000000000"}, "no_device_signals"), + ], +) +def test_policy_refusals( + app_client: tuple[TestClient, Any], changes: dict[str, Any], reason: str +) -> None: + client, router = app_client + router.get(host=HISTORY_HOST).respond(200, json={"data": [_row(**changes)], "total": 1}) + response = _signup(client) + assert response.status_code == 403 + assert response.json() == {"error": "signup_refused", "reason": reason} + + +def test_missing_identification_is_refused(app_client: tuple[TestClient, Any]) -> None: + client, router = app_client + router.get(host=HISTORY_HOST).respond(200, json={"data": [], "total": 0}) + response = _signup(client) + assert response.status_code == 403 + assert response.json()["reason"] == "missing" + + +def test_invalid_request_id_and_unavailable_api(app_client: tuple[TestClient, Any]) -> None: + client, router = app_client + assert _signup(client, "not-a-uuid").status_code == 400 + router.get(host=HISTORY_HOST).respond(401, text='{"error":"invalid api key"}\n') + assert _signup(client).status_code == 503 + + +def _deliver(client: TestClient, body: bytes, secret: Any = SECRET) -> int: + headers = {"X-Shield-Signature": sign(secret, body)} if secret else {} + return client.post("/webhooks/shieldlabs", content=body, headers=headers).status_code + + +def test_webhook_receiver(app_client: tuple[TestClient, Any]) -> None: + client, _ = app_client + data = Path(__file__).parent / "data" + scored = (data / "webhook-identification-scored.raw.txt").read_bytes() + ping = (data / "webhook-ping.raw.txt").read_bytes() + other = json.dumps({"event_type": "other", "schema_version": "2026-06-01"}).encode() + + assert _deliver(client, scored) == 200 + assert _deliver(client, scored) == 200 # a repeated delivery is acknowledged once more + assert _deliver(client, scored, secret="whsec_wrong") == 401 + assert _deliver(client, scored, secret=None) == 401 + assert _deliver(client, ping) == 200 + assert _deliver(client, other) == 200 + assert _deliver(client, b"not json") == 400 diff --git a/tests/test_helpers.py b/tests/test_helpers.py new file mode 100644 index 0000000..798d929 --- /dev/null +++ b/tests/test_helpers.py @@ -0,0 +1,214 @@ +"""evaluate_identification and user_hid.""" + +from __future__ import annotations + +import hashlib +import hmac +from dataclasses import replace +from datetime import datetime, timedelta, timezone +from typing import Any + +import pytest + +from _support import load_json +from shieldlabs import ( + NIL_UUID, + DetectionFlags, + Evaluation, + Identification, + ValidationError, + evaluate_identification, + user_hid, +) + +CASES = {c["name"]: c for c in load_json("normalization-cases.json")["cases"]} +OBSERVED = datetime(2026, 9, 30, 12, 34, 57, tzinfo=timezone.utc) +NOW = OBSERVED + timedelta(seconds=30) + + +def _identification(**changes: Any) -> Identification: + base = Identification.from_webhook_data(CASES["webhook_test_delivery"]["input"]) + base = replace(base, observed_at=OBSERVED, risk_score=10) + return replace(base, **changes) + + +def test_missing_identification() -> None: + assert evaluate_identification(None) == Evaluation(ok=False, reason="missing", band=None) + + +def test_clean_identification_is_ok() -> None: + verdict = evaluate_identification(_identification(), now=NOW) + assert verdict == Evaluation(ok=True, reason=None, band="trusted", flag=None) + + +def test_replay_callback_receives_request_id() -> None: + seen: list[str] = [] + + def is_replay(request_id: str) -> bool: + seen.append(request_id) + return True + + verdict = evaluate_identification(_identification(), now=NOW, is_replay=is_replay) + assert verdict.reason == "replayed" + assert verdict.band == "trusted" + assert seen == ["13f84f05-7c2a-4e9b-9f1d-2a6b8c0e4d11"] + assert evaluate_identification(_identification(), now=NOW, is_replay=lambda _: False).ok + + +def test_stale_identification() -> None: + late = OBSERVED + timedelta(seconds=301) + assert evaluate_identification(_identification(), now=late).reason == "stale" + assert evaluate_identification(_identification(), now=OBSERVED + timedelta(seconds=300)).ok + assert evaluate_identification(_identification(), now=late, max_age=600).ok + assert evaluate_identification(_identification(), now=late, max_age=timedelta(minutes=10)).ok + assert evaluate_identification(_identification(observed_at=None), now=NOW).reason == "stale" + + +def test_now_defaults_to_current_time() -> None: + fresh = _identification(observed_at=datetime.now(timezone.utc) - timedelta(seconds=5)) + assert evaluate_identification(fresh).ok + old = _identification(observed_at=datetime(2020, 1, 1, tzinfo=timezone.utc)) + assert evaluate_identification(old).reason == "stale" + + +def test_naive_datetimes_are_treated_as_utc() -> None: + naive_now = NOW.replace(tzinfo=None) + assert evaluate_identification(_identification(), now=naive_now).ok + naive_observed = _identification(observed_at=OBSERVED.replace(tzinfo=None)) + assert evaluate_identification(naive_observed, now=NOW).ok + + +def test_future_observed_at_is_not_stale() -> None: + ahead = _identification(observed_at=NOW + timedelta(seconds=5)) + assert evaluate_identification(ahead, now=NOW).ok + + +def test_rate_limit_marker() -> None: + marker = Identification.from_webhook_data(CASES["webhook_rate_limited"]["input"]) + marker = replace(marker, observed_at=OBSERVED) + verdict = evaluate_identification(marker, now=NOW) + assert verdict == Evaluation(ok=False, reason="rate_limited", band="rate_limited") + + +def test_nil_device_id() -> None: + verdict = evaluate_identification(_identification(device_id=NIL_UUID), now=NOW) + assert verdict.reason == "no_device_signals" + assert evaluate_identification(_identification(device_id=""), now=NOW).reason == ( + "no_device_signals" + ) + + +def test_blocked_flags_in_order() -> None: + flags = DetectionFlags(javascript_disabled=True, browser_automation=True) + verdict = evaluate_identification(_identification(detection_flags=flags), now=NOW) + assert verdict == Evaluation( + ok=False, reason="blocked_flag", band="trusted", flag="browser_automation" + ) + only_js = DetectionFlags(javascript_disabled=True) + assert ( + evaluate_identification(_identification(detection_flags=only_js), now=NOW).flag + == "javascript_disabled" + ) + + +def test_custom_block_flags() -> None: + flags = DetectionFlags(vpn=True, browser_automation=True) + identification = _identification(detection_flags=flags) + assert evaluate_identification(identification, now=NOW, block_flags=["vpn"]).flag == "vpn" + assert evaluate_identification(identification, now=NOW, block_flags="vpn").flag == "vpn" + assert evaluate_identification(identification, now=NOW, block_flags=()).ok + + +def test_blocked_bands() -> None: + dangerous = _identification(risk_score=80) + assert evaluate_identification(dangerous, now=NOW) == Evaluation( + ok=False, reason="blocked_band", band="dangerous" + ) + suspicious = _identification(risk_score=45) + assert evaluate_identification(suspicious, now=NOW).ok + both = ["suspicious", "dangerous"] + assert evaluate_identification(suspicious, now=NOW, block_bands=both).reason == "blocked_band" + assert evaluate_identification(dangerous, now=NOW, block_bands="suspicious").ok + assert evaluate_identification(dangerous, now=NOW, block_bands=[]).ok + + +def test_check_order_replay_before_stale_before_marker() -> None: + marker = _identification(risk_score=999, device_id=NIL_UUID) + late = NOW + timedelta(hours=1) + assert evaluate_identification(marker, now=late, is_replay=lambda _: True).reason == "replayed" + assert evaluate_identification(marker, now=late).reason == "stale" + assert evaluate_identification(marker, now=NOW).reason == "rate_limited" + + +@pytest.mark.parametrize( + "kwargs", + [ + {"block_bands": ["critical"]}, + {"block_flags": ["automation"]}, + {"block_flags": "anti_detect"}, + {"max_age": "300"}, + {"max_age": True}, + {"max_age": None}, + ], +) +def test_invalid_policy_arguments(kwargs: dict[str, Any]) -> None: + with pytest.raises(ValidationError): + evaluate_identification(_identification(), now=NOW, **kwargs) + + +@pytest.mark.parametrize( + "max_age", + [ + float("nan"), + float("inf"), + float("-inf"), + -1, + -0.5, + timedelta(seconds=-1), + 10**400, + ], +) +def test_freshness_window_must_be_finite_and_not_negative(max_age: Any) -> None: + # A window that can never be exceeded would silently accept stale identifications. + old = _identification(observed_at=datetime(2020, 1, 1, tzinfo=timezone.utc)) + with pytest.raises(ValidationError, match="max_age"): + evaluate_identification(old, now=NOW, max_age=max_age) + with pytest.raises(ValidationError, match="max_age"): + evaluate_identification(None, max_age=max_age) + + +def test_zero_freshness_window() -> None: + assert evaluate_identification(_identification(), now=OBSERVED, max_age=0).ok + assert evaluate_identification(_identification(), now=NOW, max_age=timedelta()).reason == ( + "stale" + ) + + +def test_user_hid_is_hmac_sha256_hex() -> None: + expected = hmac.new(b"server-secret", b"user-42", hashlib.sha256).hexdigest() + assert user_hid("user-42", "server-secret") == expected + assert user_hid("user-42", b"server-secret") == expected + assert len(expected) == 64 + assert user_hid("user-42", "server-secret") == user_hid("user-42", "server-secret") + assert user_hid("user-43", "server-secret") != expected + + +def test_user_hid_known_vector() -> None: + # HMAC-SHA256(key="key", message="The quick brown fox jumps over the lazy dog") + assert user_hid("The quick brown fox jumps over the lazy dog", "key") == ( + "f7bc83f430538424b13298e6aa6fb143ef4d59a14946175997479dbc2d1a3cd8" + ) + + +def test_user_hid_unicode_is_utf8() -> None: + expected = hmac.new(b"k", "élève".encode(), hashlib.sha256).hexdigest() + assert user_hid("élève", "k") == expected + + +@pytest.mark.parametrize( + ("user_id", "secret"), + [("", "secret"), ("user", ""), ("user", b""), (42, "secret"), ("user", None)], +) +def test_user_hid_rejects_empty_input(user_id: Any, secret: Any) -> None: + with pytest.raises(ValidationError): + user_hid(user_id, secret) diff --git a/tests/test_history.py b/tests/test_history.py new file mode 100644 index 0000000..b6d965b --- /dev/null +++ b/tests/test_history.py @@ -0,0 +1,499 @@ +"""History API client: requests, validation, iteration and configuration.""" + +from __future__ import annotations + +import copy +import string +import warnings +from typing import Any +from uuid import UUID + +import httpx +import pytest +import respx + +from _support import ( + API_KEY, + HISTORY_HOST, + REQUEST_ID, + REQUEST_PATH, + history_body, + load_json, + row, + uuid_for, +) +from shieldlabs import ( + ApiError, + HistoryPage, + Identification, + ShieldLabs, + ShieldLabsWarning, + ValidationError, +) +from shieldlabs._http import USER_AGENT + +DEVICE_ID = "AC7C303D-971B-41D1-8E25-CD5B46B46AED" + + +@pytest.fixture +def client() -> ShieldLabs: + return ShieldLabs(api_key=API_KEY) + + +@pytest.fixture +def mock() -> Any: + with respx.mock(assert_all_called=False) as router: + yield router + + +def test_search_sends_expected_request(client: ShieldLabs, mock: Any) -> None: + route = mock.get(host=HISTORY_HOST, path=REQUEST_PATH).respond( + 200, json=load_json("history-page.json") + ) + page = client.history.search("request_id", REQUEST_ID, limit=5, offset=10) + + assert isinstance(page, HistoryPage) + assert page.total == 37 + assert [item.request_id for item in page.data][:2] == [ + REQUEST_ID, + "7c1e2f4a-3b6d-4e8f-9a0b-1c2d3e4f5a6b", + ] + request = route.calls.last.request + assert request.method == "GET" + assert str(request.url) == f"https://{HISTORY_HOST}{REQUEST_PATH}?limit=5&offset=10" + assert request.headers["authorization"] == f"Bearer {API_KEY}" + assert request.headers["accept"] == "application/json" + assert request.headers["user-agent"] == USER_AGENT + assert USER_AGENT.startswith("shieldlabs-python/1.0.0 ") + + +def test_search_defaults_and_empty_page(client: ShieldLabs, mock: Any) -> None: + route = mock.get(host=HISTORY_HOST).respond(200, json=load_json("history-empty.json")) + page = client.history.search("user_hid", "anonymous") + assert page == HistoryPage(data=(), total=0) + assert route.calls.last.request.url.params["limit"] == "20" + assert route.calls.last.request.url.params["offset"] == "0" + + +def test_uuid_values_are_sent_lowercase(client: ShieldLabs, mock: Any) -> None: + route = mock.get(host=HISTORY_HOST).respond(200, json=history_body()) + client.history.search("device_id", DEVICE_ID) + client.history.search("visitor_id", UUID("bde0e249-20d8-4544-838c-ed9a0b6d7a36")) + paths = [call.request.url.path for call in route.calls] + assert paths == [ + f"/api/v1/history/device_id/{DEVICE_ID.lower()}", + "/api/v1/history/visitor_id/bde0e249-20d8-4544-838c-ed9a0b6d7a36", + ] + + +@pytest.mark.parametrize( + ("user_hid", "segment"), + [ + # The History API matches only this canonical escaping: letters, digits, "-._~" and + # "$&+,:;=@" as they are, everything else percent-encoded with uppercase hex. + ("anonymous", "anonymous"), + ("a@b", "a@b"), + ("a+b", "a+b"), + ("a:b=c", "a:b=c"), + ("x$y&z", "x$y&z"), + ("a,b", "a,b"), + ("a;b", "a;b"), + ("-._~", "-._~"), + ("...", "..."), + (".hidden", ".hidden"), + ("a b", "a%20b"), + ("ü", "%C3%BC"), + ("a%2Fb", "a%252Fb"), + ("a!b c", "a%21b%20c"), + ("it's (1)*", "it%27s%20%281%29%2A"), + ("Team A?c#d%e é", "Team%20A%3Fc%23d%25e%20%C3%A9"), + ('"<>[]^`{|}\\', "%22%3C%3E%5B%5D%5E%60%7B%7C%7D%5C"), + ("line\nbreak", "line%0Abreak"), + ("\u00a0", "%C2%A0"), + ("日本", "%E6%97%A5%E6%9C%AC"), + ("\U0001f600", "%F0%9F%98%80"), + ], +) +def test_user_hid_uses_the_escaping_the_history_api_matches( + client: ShieldLabs, mock: Any, user_hid: str, segment: str +) -> None: + route = mock.get(host=HISTORY_HOST).respond(200, json=history_body()) + client.history.search("user_hid", user_hid) + raw_path = route.calls.last.request.url.raw_path.decode("ascii") + assert raw_path == f"/api/v1/history/user_hid/{segment}?limit=20&offset=0" + + +_KEPT_IN_PATH = frozenset(string.ascii_letters + string.digits + "-._~" + "$&+,:;=@") + + +def test_every_ascii_character_uses_the_canonical_path_escaping( + client: ShieldLabs, mock: Any +) -> None: + # Canonical form: letters, digits, "-._~" and "$&+,:;=@" as they are, every other byte as + # %XX with uppercase hex. "/" is refused (see below), so it is not part of this check. + route = mock.get(host=HISTORY_HOST).respond(200, json=history_body()) + mismatches = {} + for code in range(128): + char = chr(code) + if char == "/": + continue + client.history.search("user_hid", f"a{char}z") + raw_path = route.calls.last.request.url.raw_path.decode("ascii") + segment = raw_path.split("?", 1)[0].rsplit("/", 1)[1] + expected = "a" + (char if char in _KEPT_IN_PATH else f"%{code:02X}") + "z" + if segment != expected: + mismatches[char] = segment + assert mismatches == {} + assert route.call_count == 127 + + +@pytest.mark.parametrize( + ("user_hid", "message"), + [ + (".", "cannot be searched"), + ("..", "cannot be searched"), + ("a/b", "contain '/'"), + ("/", "contain '/'"), + ("acct/", "contain '/'"), + ("../anonymous", "contain '/'"), + ("\ud800", "unpaired surrogate"), + ], +) +def test_user_hids_the_history_api_cannot_search_send_nothing( + client: ShieldLabs, mock: Any, user_hid: str, message: str +) -> None: + route = mock.get(host=HISTORY_HOST).respond(200, json=history_body()) + with pytest.raises(ValidationError, match=message): + client.history.search("user_hid", user_hid) + with pytest.raises(ValidationError, match=message): + client.history.iter("user_hid", user_hid) + assert route.call_count == 0 + + +def test_ip_lookup(client: ShieldLabs, mock: Any) -> None: + route = mock.get(host=HISTORY_HOST).respond(200, json=history_body()) + client.history.search("ip", "203.0.113.24") + assert route.calls.last.request.url.path == "/api/v1/history/ip/203.0.113.24" + + +@pytest.mark.parametrize( + ("lookup_type", "value", "message"), + [ + ("email", "user@example.com", "type must be one of"), + ("auto", "anonymous", "type must be one of"), + ("REQUEST_ID", REQUEST_ID, "type must be one of"), + ("request_id", "not-a-uuid", "must be a UUID"), + ("request_id", REQUEST_ID + " ", "must be a UUID"), + ("device_id", "ac7c303d971b41d18e25cd5b46b46aed", "must be a UUID"), + ("session_id", "{bde78778-efd2-4c49-952f-1f11b9c05f35}", "must be a UUID"), + ("cookie_id", "", "must be a UUID"), + ("ip", "2001:db8::1", "IPv6"), + ("ip", "203.0.113", "dotted IPv4"), + ("ip", "203.0.113.256", "dotted IPv4"), + ("ip", "203.0.113.07", "dotted IPv4"), + ("ip", "203.0.113.7\n", "dotted IPv4"), + ("user_hid", "", "non-empty"), + ("user_hid", 12345, "must be a string"), + ("request_id", None, "must be a string"), + ], +) +def test_validation_errors_send_nothing( + client: ShieldLabs, mock: Any, lookup_type: Any, value: Any, message: str +) -> None: + route = mock.get(host=HISTORY_HOST).respond(200, json=history_body()) + with pytest.raises(ValidationError, match=message): + client.history.search(lookup_type, value) + assert route.call_count == 0 + + +@pytest.mark.parametrize( + ("kwargs", "message"), + [ + ({"limit": 0}, "between 1 and 100"), + ({"limit": 101}, "between 1 and 100"), + ({"limit": 20.0}, "must be an integer"), + ({"limit": True}, "must be an integer"), + ({"offset": -1}, "0 or greater"), + ({"offset": "10"}, "must be an integer"), + ], +) +def test_limit_and_offset_validation( + client: ShieldLabs, mock: Any, kwargs: dict[str, Any], message: str +) -> None: + route = mock.get(host=HISTORY_HOST).respond(200, json=history_body()) + with pytest.raises(ValidationError, match=message): + client.history.search("request_id", REQUEST_ID, **kwargs) + assert route.call_count == 0 + + +def test_limit_boundaries_are_accepted(client: ShieldLabs, mock: Any) -> None: + route = mock.get(host=HISTORY_HOST).respond(200, json=history_body()) + client.history.search("request_id", REQUEST_ID, limit=1) + client.history.search("request_id", REQUEST_ID, limit=100, offset=5000) + assert [c.request.url.params["limit"] for c in route.calls] == ["1", "100"] + + +def test_validation_error_is_also_value_error(client: ShieldLabs) -> None: + with pytest.raises(ValueError): + client.history.search("request_id", "nope") + + +# --- iteration --------------------------------------------------------------------------- + + +def test_iter_dedupes_on_request_id_and_stops_at_total(client: ShieldLabs, mock: Any) -> None: + route = mock.get(host=HISTORY_HOST).mock( + side_effect=[ + httpx.Response(200, json=history_body(row(uuid_for(1)), row(uuid_for(2)), total=5)), + httpx.Response(200, json=history_body(row(uuid_for(2)), row(uuid_for(3)), total=5)), + httpx.Response(200, json=history_body(row(uuid_for(4)), total=5)), + ] + ) + items = list(client.history.iter("user_hid", "acct-1", page_size=2)) + assert [item.request_id for item in items] == [uuid_for(i) for i in (1, 2, 3, 4)] + offsets = [call.request.url.params["offset"] for call in route.calls] + limits = {call.request.url.params["limit"] for call in route.calls} + assert offsets == ["0", "2", "4"] + assert limits == {"2"} + + +def test_iter_stops_at_empty_page(client: ShieldLabs, mock: Any) -> None: + route = mock.get(host=HISTORY_HOST).mock( + side_effect=[ + httpx.Response(200, json=history_body(row(uuid_for(1)), total=100)), + httpx.Response(200, json=history_body(total=100)), + ] + ) + items = list(client.history.iter("device_id", DEVICE_ID, page_size=1)) + assert len(items) == 1 + assert route.call_count == 2 + + +def test_iter_respects_max_items(client: ShieldLabs, mock: Any) -> None: + route = mock.get(host=HISTORY_HOST).mock( + side_effect=[ + httpx.Response(200, json=history_body(*(row(uuid_for(i)) for i in range(3)), total=10)), + httpx.Response( + 200, json=history_body(*(row(uuid_for(i)) for i in range(3, 6)), total=10) + ), + ] + ) + items = list(client.history.iter("ip", "203.0.113.24", page_size=3, max_items=4)) + assert [item.request_id for item in items] == [uuid_for(i) for i in range(4)] + assert route.call_count == 2 + + +def test_iter_max_items_zero_sends_nothing(client: ShieldLabs, mock: Any) -> None: + route = mock.get(host=HISTORY_HOST).respond(200, json=history_body()) + assert list(client.history.iter("ip", "203.0.113.24", max_items=0)) == [] + assert route.call_count == 0 + + +def test_iter_default_page_size_and_single_page(client: ShieldLabs, mock: Any) -> None: + route = mock.get(host=HISTORY_HOST).respond(200, json=load_json("history-page.json")) + items = list(client.history.iter("user_hid", "acct-1", max_items=5)) + assert len(items) == 5 + assert route.calls.last.request.url.params["limit"] == "100" + assert all(isinstance(item, Identification) for item in items) + + +@pytest.mark.parametrize( + ("args", "kwargs"), + [ + (("user_hid", ""), {}), + (("email", "x"), {}), + (("ip", "2001:db8::1"), {}), + (("user_hid", "a"), {"page_size": 0}), + (("user_hid", "a"), {"page_size": 101}), + (("user_hid", "a"), {"max_items": -1}), + ], +) +def test_iter_validates_eagerly( + client: ShieldLabs, mock: Any, args: tuple[Any, ...], kwargs: dict[str, Any] +) -> None: + route = mock.get(host=HISTORY_HOST).respond(200, json=history_body()) + with pytest.raises(ValidationError): + client.history.iter(*args, **kwargs) + assert route.call_count == 0 + + +# --- configuration ------------------------------------------------------------------------ + + +@pytest.mark.parametrize( + "base_url", + [ + "https://account.shieldlabs.ai/api", + "https://account.shieldlabs.ai/api/", + "https://account.shieldlabs.ai/", + " https://account.shieldlabs.ai ", + ], +) +def test_api_suffix_is_stripped(mock: Any, base_url: str) -> None: + route = mock.get(host=HISTORY_HOST).respond(200, json=history_body()) + client = ShieldLabs(api_key=API_KEY, base_url=base_url) + assert client.base_url == "https://account.shieldlabs.ai" + client.history.search("request_id", REQUEST_ID) + assert route.calls.last.request.url.path == REQUEST_PATH + + +def test_custom_base_url_with_prefix_and_dev_host(mock: Any) -> None: + route = mock.get(host="dev.account.shieldlabs.ai").respond(200, json=history_body()) + client = ShieldLabs(api_key=API_KEY, base_url="https://dev.account.shieldlabs.ai") + client.history.search("request_id", REQUEST_ID) + assert route.call_count == 1 + proxied = ShieldLabs(api_key=API_KEY, base_url="http://localhost:8080/shield/api") + assert proxied.base_url == "http://localhost:8080/shield" + + +@pytest.mark.parametrize( + "base_url", + ["account.shieldlabs.ai", "ftp://x", "", "https://", "http://:8080", "https://[::1"], +) +def test_invalid_base_url(base_url: str) -> None: + with pytest.raises(ValidationError, match="must be an http"): + ShieldLabs(api_key=API_KEY, base_url=base_url) + + +@pytest.mark.parametrize( + "base_url", + [ + "http://account.shieldlabs.ai", + "HTTP://account.shieldlabs.ai/api", + "http://10.0.0.5:8080", + "http://localhost.example.com", + "http://0.0.0.0:8080", + "http://user:secret@shieldlabs.example", + ], +) +def test_plain_http_is_refused_for_remote_hosts(base_url: str) -> None: + with pytest.raises(ValidationError, match="must use https") as caught: + ShieldLabs(api_key=API_KEY, base_url=base_url) + assert "secret" not in str(caught.value) + + +@pytest.mark.parametrize( + ("base_url", "origin"), + [ + ("http://localhost:8080", "http://localhost:8080"), + ("http://LOCALHOST:3000/api", "http://LOCALHOST:3000"), + ("http://127.0.0.1:8080", "http://127.0.0.1:8080"), + ("http://127.0.0.2", "http://127.0.0.2"), + ("http://[::1]:8080/", "http://[::1]:8080"), + ("https://203.0.113.10:8443", "https://203.0.113.10:8443"), + ], +) +def test_plain_http_is_accepted_for_loopback_hosts(base_url: str, origin: str) -> None: + assert ShieldLabs(api_key=API_KEY, base_url=base_url).base_url == origin + + +def test_plain_http_from_the_environment_is_refused(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("SHIELDLABS_API_BASE_URL", "http://account.shieldlabs.ai") + with pytest.raises(ValidationError, match="must use https"): + ShieldLabs(api_key=API_KEY) + + +def test_env_fallbacks(monkeypatch: pytest.MonkeyPatch, mock: Any) -> None: + monkeypatch.setenv("SHIELDLABS_API_KEY", API_KEY) + monkeypatch.setenv("SHIELDLABS_API_BASE_URL", "https://dev.account.shieldlabs.ai/api") + route = mock.get(host="dev.account.shieldlabs.ai").respond(200, json=history_body()) + client = ShieldLabs() + assert client.base_url == "https://dev.account.shieldlabs.ai" + client.history.search("request_id", REQUEST_ID) + assert route.calls.last.request.headers["authorization"] == f"Bearer {API_KEY}" + + +@pytest.mark.parametrize("api_key", [None, "", " "]) +def test_missing_api_key(api_key: Any) -> None: + with pytest.raises(ValidationError, match="SHIELDLABS_API_KEY"): + ShieldLabs(api_key=api_key) + + +def test_non_string_api_key() -> None: + with pytest.raises(ValidationError, match="api_key must be a string"): + ShieldLabs(api_key=12345) # type: ignore[arg-type] + + +def test_unusual_api_key_warns_without_leaking_it() -> None: + with pytest.warns(ShieldLabsWarning, match="Private API Key") as record: + ShieldLabs(api_key="0123456789abcdef0123456789abcdef") + assert "0123456789abcdef" not in str(record[0].message) + assert record[0].filename == __file__ + + +def test_valid_api_key_does_not_warn() -> None: + with warnings.catch_warnings(): + warnings.simplefilter("error") + ShieldLabs(api_key=f" {API_KEY}\n") + + +def test_repr_never_contains_the_key(client: ShieldLabs) -> None: + assert API_KEY not in repr(client) + assert repr(client) == "ShieldLabs(base_url='https://account.shieldlabs.ai')" + + +@pytest.mark.parametrize( + "kwargs", + [{"timeout": 0}, {"timeout": -1}, {"timeout": "10"}, {"max_retries": -1}, {"max_retries": 1.5}], +) +def test_invalid_options(kwargs: dict[str, Any]) -> None: + with pytest.raises(ValidationError): + ShieldLabs(api_key=API_KEY, **kwargs) + + +def test_context_manager_closes_owned_client() -> None: + with ShieldLabs(api_key=API_KEY) as client: + inner = client._transport.client + assert not inner.is_closed + assert inner.is_closed + + +def test_injected_http_client_is_used_and_left_open(mock: Any) -> None: + route = mock.get(host=HISTORY_HOST).respond(200, json=history_body()) + http_client = httpx.Client(headers={"X-Trace": "1"}) + with ShieldLabs(api_key=API_KEY, http_client=http_client) as client: + client.history.search("request_id", REQUEST_ID) + assert not http_client.is_closed + assert route.calls.last.request.headers["x-trace"] == "1" + assert route.calls.last.request.headers["authorization"] == f"Bearer {API_KEY}" + http_client.close() + + +def test_malformed_success_bodies_raise_api_error(client: ShieldLabs, mock: Any) -> None: + mock.get(host=HISTORY_HOST).mock( + side_effect=[ + httpx.Response(200, text="maintenance"), + httpx.Response(200, json=[1, 2, 3]), + ] + ) + with pytest.raises(ApiError, match="not valid JSON") as first: + client.history.search("request_id", REQUEST_ID) + assert first.value.status == 200 + with pytest.raises(ApiError, match="not a JSON object"): + client.history.search("request_id", REQUEST_ID) + + +def test_errors_survive_deepcopy(client: ShieldLabs, mock: Any) -> None: + mock.get(host=HISTORY_HOST).respond( + 401, text='{"error":"invalid api key"}\n', headers={"content-type": "text/plain"} + ) + with pytest.raises(ApiError) as caught: + client.history.search("request_id", REQUEST_ID) + restored = copy.deepcopy(caught.value) + assert type(restored) is type(caught.value) + assert restored.status == 401 + assert restored.message == "invalid api key" + assert restored.headers["content-type"] == "text/plain" + assert str(restored) == "invalid api key (HTTP 401)" + + +@pytest.mark.parametrize("api_key", ["sec_abc def", "sec_abc\x00def", "sec_été"]) +def test_api_key_must_be_header_safe_and_is_never_echoed(api_key: str) -> None: + with pytest.raises(ValidationError, match="HTTP header") as caught: + ShieldLabs(api_key=api_key) + assert api_key not in str(caught.value) + + +def test_non_string_base_url() -> None: + with pytest.raises(ValidationError, match="base_url must be a string"): + ShieldLabs(api_key=API_KEY, base_url=8080) # type: ignore[arg-type] diff --git a/tests/test_management.py b/tests/test_management.py new file mode 100644 index 0000000..656b7b4 --- /dev/null +++ b/tests/test_management.py @@ -0,0 +1,226 @@ +"""Management API client: profile, headers, domain normalization and the no-retry-on-429 rule.""" + +from __future__ import annotations + +from datetime import datetime, timezone +from typing import Any + +import httpx +import pytest +import respx + +from _support import DOMAIN, MANAGEMENT_HOST, SECRET_KEY, FakeTime, load_json +from shieldlabs import ( + DomainProfile, + QuotaExceededError, + RateLimitError, + ServerError, + ShieldLabsManagement, + ValidationError, +) +from shieldlabs._http import USER_AGENT +from shieldlabs._validation import normalize_domain + + +@pytest.fixture +def mock() -> Any: + with respx.mock(assert_all_called=False) as router: + yield router + + +@pytest.fixture +def management(fake_time: FakeTime) -> ShieldLabsManagement: + client = ShieldLabsManagement(secret_key=SECRET_KEY, domain="https://www.Example.com/") + fake_time.install(client._transport) + return client + + +def test_get_profile_matches_expected_fixture(management: ShieldLabsManagement, mock: Any) -> None: + route = mock.get(host=MANAGEMENT_HOST, path="/v1/profile").respond( + 200, json=load_json("management-profile.json") + ) + profile = management.get_profile() + + assert isinstance(profile, DomainProfile) + assert profile.to_dict() == load_json("management-profile-expected.json") + assert profile.created_at == datetime(2026, 1, 15, 9, 0, tzinfo=timezone.utc) + assert profile.raw["Callback"] == "" + assert not hasattr(profile, "callback") + + request = route.calls.last.request + assert str(request.url) == "https://api.shieldlabs.ai/v1/profile" + assert request.headers["x-shield-domain"] == DOMAIN + assert request.headers["authorization"] == f"Bearer {SECRET_KEY}" + assert request.headers["accept"] == "application/json" + assert request.headers["user-agent"] == USER_AGENT + + +def test_negative_remaining_identifications(management: ShieldLabsManagement, mock: Any) -> None: + body = dict(load_json("management-profile.json"), Weight=-120, CreatedAt="0001-01-01T00:00:00Z") + mock.get(host=MANAGEMENT_HOST).respond(200, json=body) + profile = management.get_profile() + assert profile.remaining_identifications == -120 + assert profile.created_at == datetime(1, 1, 1, tzinfo=timezone.utc) + assert profile.to_dict()["created_at"] == "0001-01-01T00:00:00.000Z" + + +def test_profile_tolerates_missing_fields() -> None: + profile = DomainProfile.from_dict({"Domain": "example.com"}) + assert profile.remaining_identifications == 0 + assert profile.public_key_masked == "" + assert profile.created_at is None + + +def test_rate_limit_is_never_retried( + management: ShieldLabsManagement, mock: Any, fake_time: FakeTime +) -> None: + route = mock.get(host=MANAGEMENT_HOST).respond( + 429, json={"error": "too many requests"}, headers={"Retry-After": "1"} + ) + with pytest.raises(RateLimitError, match="too many requests"): + management.get_profile() + assert route.call_count == 1 + assert fake_time.sleeps == [] + + +def test_rate_limit_not_retried_even_with_many_retries(mock: Any, fake_time: FakeTime) -> None: + client = ShieldLabsManagement(secret_key=SECRET_KEY, domain=DOMAIN, max_retries=10) + fake_time.install(client._transport) + route = mock.get(host=MANAGEMENT_HOST).respond(429, json={"error": "too many requests"}) + with pytest.raises(RateLimitError): + client.get_profile() + assert route.call_count == 1 + + +def test_server_busy_is_retried( + management: ShieldLabsManagement, mock: Any, fake_time: FakeTime +) -> None: + route = mock.get(host=MANAGEMENT_HOST).mock( + side_effect=[ + httpx.Response(503, json={"error": "server is busy"}), + httpx.Response(200, json=load_json("management-profile.json")), + ] + ) + assert management.get_profile().domain == "example.com" + assert route.call_count == 2 + assert fake_time.sleeps == [0.5] + + +def test_server_busy_exhausts_retries( + management: ShieldLabsManagement, mock: Any, fake_time: FakeTime +) -> None: + route = mock.get(host=MANAGEMENT_HOST).respond(503, json={"error": "server is busy"}) + with pytest.raises(ServerError, match="server is busy"): + management.get_profile() + assert route.call_count == 3 + + +def test_quota_exceeded_class(management: ShieldLabsManagement, mock: Any) -> None: + mock.get(host=MANAGEMENT_HOST).respond(402) + with pytest.raises(QuotaExceededError): + management.get_profile() + + +@pytest.mark.parametrize( + ("raw", "expected"), + [ + ("example.com", "example.com"), + (" Example.COM ", "example.com"), + ("https://www.example.com/", "example.com"), + ("http://www.example.com/signup?x=1#top", "example.com"), + ("www.example.com", "example.com"), + ("//example.com/path", "example.com"), + ("shop.example.com", "shop.example.com"), + ("example.com:8443", "example.com:8443"), + ("example.com?x", "example.com"), + ("wwwexample.com", "wwwexample.com"), + ], +) +def test_domain_normalization(raw: str, expected: str) -> None: + assert normalize_domain(raw) == expected + + +@pytest.mark.parametrize("raw", ["", " ", "https://", "www.", "/path"]) +def test_empty_domain_is_rejected(raw: str) -> None: + with pytest.raises(ValidationError, match="domain"): + ShieldLabsManagement(secret_key=SECRET_KEY, domain=raw) + + +def test_non_string_domain() -> None: + with pytest.raises(ValidationError, match="domain must be a string"): + ShieldLabsManagement(secret_key=SECRET_KEY, domain=42) # type: ignore[arg-type] + + +@pytest.mark.parametrize("secret", [None, "", " "]) +def test_secret_key_is_required(secret: Any) -> None: + with pytest.raises(ValidationError, match="SHIELDLABS_SECRET_KEY"): + ShieldLabsManagement(secret_key=secret, domain=DOMAIN) + + +def test_env_fallbacks(monkeypatch: pytest.MonkeyPatch, mock: Any) -> None: + monkeypatch.setenv("SHIELDLABS_SECRET_KEY", SECRET_KEY) + monkeypatch.setenv("SHIELDLABS_DOMAIN", "WWW.Example.com") + monkeypatch.setenv("SHIELDLABS_MANAGEMENT_BASE_URL", "https://dev.api.shieldlabs.ai/") + route = mock.get(host="dev.api.shieldlabs.ai", path="/v1/profile").respond( + 200, json=load_json("management-profile.json") + ) + with ShieldLabsManagement() as client: + assert client.domain == "example.com" + assert client.base_url == "https://dev.api.shieldlabs.ai" + client.get_profile() + assert route.calls.last.request.headers["x-shield-domain"] == "example.com" + + +def test_repr_hides_the_secret(management: ShieldLabsManagement) -> None: + assert SECRET_KEY not in repr(management) + assert "example.com" in repr(management) + + +def test_context_manager_closes_owned_client() -> None: + with ShieldLabsManagement(secret_key=SECRET_KEY, domain=DOMAIN) as client: + inner = client._transport.client + assert inner.is_closed + + +def test_injected_client_is_not_closed() -> None: + http_client = httpx.Client() + client = ShieldLabsManagement(secret_key=SECRET_KEY, domain=DOMAIN, http_client=http_client) + client.close() + assert not http_client.is_closed + http_client.close() + + +def test_invalid_management_options() -> None: + with pytest.raises(ValidationError): + ShieldLabsManagement(secret_key=SECRET_KEY, domain=DOMAIN, timeout=0) + with pytest.raises(ValidationError): + ShieldLabsManagement(secret_key=SECRET_KEY, domain=DOMAIN, base_url="api.shieldlabs.ai") + + +def test_plain_http_base_url_only_for_loopback(monkeypatch: pytest.MonkeyPatch) -> None: + with pytest.raises(ValidationError, match="must use https"): + ShieldLabsManagement( + secret_key=SECRET_KEY, domain=DOMAIN, base_url="http://api.shieldlabs.ai" + ) + monkeypatch.setenv("SHIELDLABS_MANAGEMENT_BASE_URL", "http://dev.api.shieldlabs.ai") + with pytest.raises(ValidationError, match="must use https"): + ShieldLabsManagement(secret_key=SECRET_KEY, domain=DOMAIN) + monkeypatch.setenv("SHIELDLABS_MANAGEMENT_BASE_URL", "http://127.0.0.1:9000/") + local = ShieldLabsManagement(secret_key=SECRET_KEY, domain=DOMAIN) + assert local.base_url == "http://127.0.0.1:9000" + + +def test_secret_key_must_be_header_safe() -> None: + with pytest.raises(ValidationError, match="HTTP header") as caught: + ShieldLabsManagement(secret_key="abc\ndef", domain=DOMAIN) + assert "abc" not in str(caught.value) + + +@pytest.mark.parametrize( + ("raw", "message"), + [("пример.рф", "punycode"), ("exa mple.com", "HTTP header")], +) +def test_domain_must_be_header_safe(raw: str, message: str) -> None: + with pytest.raises(ValidationError, match=message): + ShieldLabsManagement(secret_key=SECRET_KEY, domain=raw) + assert normalize_domain("xn--e1afmkfd.xn--p1ai") == "xn--e1afmkfd.xn--p1ai" diff --git a/tests/test_normalization.py b/tests/test_normalization.py new file mode 100644 index 0000000..caec542 --- /dev/null +++ b/tests/test_normalization.py @@ -0,0 +1,348 @@ +"""Shared fixture tests: History rows and webhook data normalize into the same Identification.""" + +from __future__ import annotations + +import json +from datetime import datetime, timedelta, timezone +from typing import Any + +import pytest + +from _support import load_json +from shieldlabs import ( + NIL_UUID, + DetectionFlags, + HistoryPage, + Identification, + IpInfo, + Signal, + SignalName, + TrafficSource, + is_rate_limited, + risk_band, +) +from shieldlabs._normalize import ( + FLAG_KEYS, + fallback_slug, + format_timestamp, + parse_history_time, + parse_rfc3339, + signal_slug, +) + +NORMALIZATION = load_json("normalization-cases.json")["cases"] +SLUGS = load_json("signal-slug-cases.json")["cases"] +BANDS = load_json("risk-band-cases.json")["cases"] + + +def _normalize(case: dict[str, Any]) -> Identification: + if case["source"] == "history": + return Identification.from_history_row(case["input"]) + return Identification.from_webhook_data(case["input"]) + + +def _truncate_ms(value: str) -> str: + return value[:23] + + +def test_fixture_covers_both_sources() -> None: + sources = {case["source"] for case in NORMALIZATION} + assert sources == {"history", "webhook"} + assert len(NORMALIZATION) == 8 + + +@pytest.mark.parametrize("case", NORMALIZATION, ids=[c["name"] for c in NORMALIZATION]) +def test_normalization_cases(case: dict[str, Any]) -> None: + identification = _normalize(case) + expected = case["expected"] + actual = identification.to_dict() + + # observed_at is compared at millisecond precision (truncated). + assert _truncate_ms(actual.pop("observed_at")) == _truncate_ms(expected["observed_at"]) + expected_rest = {key: value for key, value in expected.items() if key != "observed_at"} + assert actual == expected_rest + + assert identification.raw == case["input"] + assert identification.source == case["source"] + assert identification.observed_at is not None + assert identification.observed_at.tzinfo is not None + assert identification.observed_at.utcoffset() == timedelta(0) + assert isinstance(identification.signals, tuple) + assert all(isinstance(signal.weight, int) for signal in identification.signals) + assert len(identification.detection_flags.to_dict()) == 19 + + +@pytest.mark.parametrize("case", NORMALIZATION, ids=[c["name"] for c in NORMALIZATION]) +def test_to_dict_round_trip(case: dict[str, Any]) -> None: + identification = _normalize(case) + rebuilt = Identification.from_dict(identification.to_dict()) + assert rebuilt.to_dict() == identification.to_dict() + assert rebuilt.source == identification.source + json.dumps(identification.to_dict()) + + +def test_webhook_keeps_nanosecond_input_as_microseconds() -> None: + case = next(c for c in NORMALIZATION if c["name"] == "webhook_scored") + identification = _normalize(case) + assert identification.observed_at == datetime( + 2026, 9, 30, 12, 34, 57, 482913, tzinfo=timezone.utc + ) + + +@pytest.mark.parametrize("case", SLUGS, ids=[c["slug"] + ":" + c["description"] for c in SLUGS]) +def test_signal_slug_cases(case: dict[str, str]) -> None: + assert signal_slug(case["description"]) == case["slug"] + + +@pytest.mark.parametrize("case", BANDS, ids=[str(c["score"]) for c in BANDS]) +def test_risk_band_cases(case: dict[str, Any]) -> None: + assert risk_band(case["score"]) == case["band"] + assert is_rate_limited(case["score"]) is (case["band"] == "rate_limited") + + +def test_identification_band_properties() -> None: + marker = next(c for c in NORMALIZATION if c["name"] == "webhook_rate_limited") + identification = _normalize(marker) + assert identification.risk_score == 999 + assert identification.risk_band == "rate_limited" + assert identification.is_rate_limited + assert not identification.has_device_signals + assert identification.device_id == NIL_UUID + + scored = _normalize(next(c for c in NORMALIZATION if c["name"] == "webhook_scored")) + assert scored.risk_band == "dangerous" + assert not scored.is_rate_limited + assert scored.has_device_signals + assert scored.detection_flags.active() == ( + "proxy", + "datacenter_ip", + "anti_detect_browser", + "ip_mismatch", + "suspicious_paid_click", + ) + + +def test_history_row_tolerates_missing_and_malformed_fields() -> None: + identification = Identification.from_history_row( + { + "request_id": 42, + "score": "high", + "score_details": "not json", + "created_at": "yesterday", + "user_hid": "", + "ip": " 0.0.0.0 ", + } + ) + assert identification.request_id == "" + assert identification.risk_score == 0 + assert identification.signals == () + assert identification.observed_at is None + assert identification.user_hid is None + assert identification.public_ip == IpInfo("", "") + assert identification.detection_flags == DetectionFlags() + + +@pytest.mark.parametrize( + "score_details", + ["", "{}", "[1, 2]", '"text"', "[" * 100_000, None, 7], +) +def test_history_row_score_details_edge_cases(score_details: Any) -> None: + identification = Identification.from_history_row({"score_details": score_details}) + assert identification.signals == () + + +def test_history_row_keeps_negative_and_repeated_signals_and_skips_non_integers() -> None: + details = [ + {"Value": 30, "Description": "Stun is not checked"}, + {"Value": -30, "Description": "Stun passed (late arrival, corrected)"}, + {"Value": 30, "Description": "Stun is not checked"}, + {"Value": 0, "Description": "Check Incomplete"}, + {"Value": 10.5, "Description": "Is proxy"}, + {"Value": True, "Description": "Is VPN"}, + "garbage", + {"Description": "no value"}, + ] + identification = Identification.from_history_row({"score_details": json.dumps(details)}) + assert [(s.name, s.weight) for s in identification.signals] == [ + ("stun_not_checked", 30), + ("stun_late_correction", -30), + ("stun_not_checked", 30), + ] + + +def test_history_ip_mismatch_rules() -> None: + base = {"ip": "203.0.113.1", "web_rtc_ip": "198.51.100.2"} + assert Identification.from_history_row(base).detection_flags.ip_mismatch + same = {"ip": "203.0.113.1", "web_rtc_ip": "203.0.113.1"} + assert not Identification.from_history_row(same).detection_flags.ip_mismatch + bot = dict(base, is_search_bot=True) + assert not Identification.from_history_row(bot).detection_flags.ip_mismatch + detail = {"score_details": json.dumps([{"Value": 0, "Description": "IP ≠ leakIP (a ≠ b)"}])} + assert Identification.from_history_row(detail).detection_flags.ip_mismatch + + +def test_history_browser_vpn_proxy_is_derived_from_connection_type() -> None: + row = {"connection_type": "browser_vpn_proxy"} + assert Identification.from_history_row(row).detection_flags.browser_vpn_proxy + assert Identification.from_history_row(row).connection_type == "browser_vpn_proxy" + + +def test_history_domain_prefers_site_domain() -> None: + row = {"domain": "shop.example.com", "site_domain": "example.com"} + assert Identification.from_history_row(row).domain == "example.com" + assert Identification.from_history_row({"domain": "shop.example.com"}).domain == ( + "shop.example.com" + ) + + +def test_history_local_ip_uses_leak_source_only_when_set() -> None: + row = { + "web_rtc_ip": "198.51.100.2", + "web_rtc_country": "Germany", + "webrtc_leak_ip": "203.0.113.9", + "webrtc_leak_country": "Spain", + "webrtc_leak_source": "none", + } + assert Identification.from_history_row(row).local_ip == IpInfo("198.51.100.2", "Germany") + row["webrtc_leak_source"] = "shield" + assert Identification.from_history_row(row).local_ip == IpInfo("203.0.113.9", "Spain") + + +def test_webhook_data_tolerates_partial_objects() -> None: + identification = Identification.from_webhook_data( + { + "request_id": "3f2b8c1e-9d4a-4e6b-8a7c-2d1e0f9b6a53", + "risk_score": 35, + "user_hid": None, + "signals": [{"name": "vpn", "weight": 15}, "bad", {"weight": 5}], + "detection_flags": ["not", "a", "map"], + "public_ip": "not a map", + } + ) + assert identification.risk_score == 35 + assert identification.user_hid is None + assert identification.signals == (Signal("vpn", 15), Signal("", 5)) + assert identification.detection_flags == DetectionFlags() + assert identification.public_ip == IpInfo() + assert identification.traffic_source == TrafficSource() + assert identification.observed_at is None + assert identification.source == "webhook" + + +def test_user_hid_sentinels_are_kept() -> None: + for sentinel in ("anonymous", "fail", "-1", "unknown"): + assert Identification.from_history_row({"user_hid": sentinel}).user_hid == sentinel + assert Identification.from_webhook_data({"user_hid": sentinel}).user_hid == sentinel + + +def test_from_dict_accepts_datetime_and_signal_descriptions() -> None: + moment = datetime(2026, 9, 30, 12, 0, tzinfo=timezone.utc) + identification = Identification.from_dict( + { + "request_id": "3f2b8c1e-9d4a-4e6b-8a7c-2d1e0f9b6a53", + "signals": [{"name": "proxy", "weight": 10, "description": "Is proxy"}], + "observed_at": moment, + "source": "history", + } + ) + assert identification.observed_at == moment + assert identification.signals == (Signal("proxy", 10, "Is proxy"),) + assert identification.source == "history" + assert Identification.from_dict({}).source == "webhook" + + +def test_history_page_from_dict() -> None: + page = HistoryPage.from_dict(load_json("history-page.json")) + assert page.total == 37 + assert len(page.data) == 5 + assert all(item.source == "history" for item in page.data) + empty = HistoryPage.from_dict(load_json("history-empty.json")) + assert empty == HistoryPage(data=(), total=0) + odd = HistoryPage.from_dict({"data": [{"request_id": "x"}, "junk"], "total": None}) + assert odd.total == 1 + assert HistoryPage.from_dict({"data": "nope"}).data == () + + +def test_history_page_rows_match_normalization_cases() -> None: + page = HistoryPage.from_dict(load_json("history-page.json")) + expected = {c["name"]: c["expected"] for c in NORMALIZATION if c["source"] == "history"} + by_request = {v["request_id"]: v for v in expected.values()} + for identification in page.data: + assert identification.to_dict() == by_request[identification.request_id] + + +def test_known_signal_names_are_plain_strings() -> None: + assert SignalName.ANTIDETECT_BROWSER == "antidetect_browser" + assert SignalName.STUN_LATE_CORRECTION == "stun_late_correction" + assert SignalName.RATE_LIMITED == "rate_limited" + + +def test_models_are_frozen() -> None: + signal = Signal("vpn", 15) + with pytest.raises(AttributeError): + signal.weight = 99 # type: ignore[misc] + + +def test_flag_keys_order_matches_detection_flags() -> None: + assert tuple(DetectionFlags().to_dict()) == FLAG_KEYS + + +@pytest.mark.parametrize( + ("value", "expected"), + [ + ("2026-09-30 12:34:56.123", "2026-09-30T12:34:56.123Z"), + ("2026-09-30 12:34:56", "2026-09-30T12:34:56.000Z"), + ("2026-09-30T12:34:56.1", "2026-09-30T12:34:56.100Z"), + ("2026-09-30 12:34:56.123456789", "2026-09-30T12:34:56.123Z"), + (" 2026-09-30 12:34:56.999 ", "2026-09-30T12:34:56.999Z"), + ("2026-09-30 12:34:56Z", "2026-09-30T12:34:56.000Z"), + ("2026-09-30 12:34:56+02:00", "2026-09-30T12:34:56.000Z"), + ("2026-02-30 12:34:56", None), + ("2026-09-30", None), + ("", None), + (None, None), + (1790771696123, None), + ], +) +def test_parse_history_time(value: Any, expected: Any) -> None: + assert format_timestamp(parse_history_time(value)) == expected + + +@pytest.mark.parametrize( + ("value", "expected"), + [ + ("2026-09-30T12:34:57.482913041Z", "2026-09-30T12:34:57.482Z"), + ("2026-09-30T13:10:00.5Z", "2026-09-30T13:10:00.500Z"), + ("2026-09-30T12:34:56Z", "2026-09-30T12:34:56.000Z"), + ("2026-09-30T14:34:56+02:00", "2026-09-30T12:34:56.000Z"), + ("2026-09-30T10:04:56-02:30", "2026-09-30T12:34:56.000Z"), + ("0001-01-01T00:00:00Z", "0001-01-01T00:00:00.000Z"), + ("0001-01-01T00:00:00+01:00", None), + ("2026-09-30 12:34:56Z", None), + ("2026-09-30T12:34:56", None), + ("2026-13-30T12:34:56Z", None), + (None, None), + ], +) +def test_parse_rfc3339(value: Any, expected: Any) -> None: + assert format_timestamp(parse_rfc3339(value)) == expected + + +def test_format_timestamp_converts_offsets_and_handles_naive() -> None: + aware = datetime(2026, 9, 30, 14, 0, 0, 999999, tzinfo=timezone(timedelta(hours=2))) + assert format_timestamp(aware) == "2026-09-30T12:00:00.999Z" + assert format_timestamp(datetime(2026, 9, 30, 12, 0)) == "2026-09-30T12:00:00.000Z" + assert format_timestamp(None) is None + + +def test_fallback_slug_examples() -> None: + assert fallback_slug("A - . - B") == "a_b" + assert fallback_slug("--x--") == "x" + assert fallback_slug("Latency (x)") == "latency" + + +def test_numeric_coercion_keeps_integers() -> None: + assert Identification.from_webhook_data({"risk_score": 45.0}).risk_score == 45 + assert Identification.from_webhook_data({"risk_score": 45.5}).risk_score == 0 + assert Identification.from_webhook_data({"risk_score": True}).risk_score == 0 + assert Identification.from_webhook_data({"risk_score": "45"}).risk_score == 0 diff --git a/tests/test_package.py b/tests/test_package.py new file mode 100644 index 0000000..114e282 --- /dev/null +++ b/tests/test_package.py @@ -0,0 +1,91 @@ +"""Public surface of the package.""" + +from __future__ import annotations + +import builtins +from pathlib import Path + +import shieldlabs + +EXPECTED_EXPORTS = { + "ShieldLabs", + "AsyncShieldLabs", + "ShieldLabsManagement", + "AsyncShieldLabsManagement", + "webhooks", + "risk_band", + "is_rate_limited", + "evaluate_identification", + "user_hid", + "Identification", + "Signal", + "DetectionFlags", + "TrafficSource", + "IpInfo", + "HistoryPage", + "DomainProfile", + "LookupType", + "RiskBand", + "Evaluation", + "IdentificationScoredEvent", + "WebhookPingEvent", + "UnknownWebhookEvent", + "ShieldLabsError", + "ApiError", + "BadRequestError", + "AuthenticationError", + "QuotaExceededError", + "NotFoundError", + "RateLimitError", + "ServerError", + "APIConnectionError", + "APITimeoutError", + "SignatureVerificationError", + "WebhookParseError", + "ValidationError", +} + + +def test_every_documented_name_is_exported() -> None: + missing = EXPECTED_EXPORTS - set(shieldlabs.__all__) + assert not missing + for name in shieldlabs.__all__: + assert hasattr(shieldlabs, name), name + + +def test_no_export_shadows_a_builtin() -> None: + assert not set(shieldlabs.__all__) & set(dir(builtins)) + + +def test_error_hierarchy() -> None: + api_errors = [ + shieldlabs.BadRequestError, + shieldlabs.AuthenticationError, + shieldlabs.QuotaExceededError, + shieldlabs.NotFoundError, + shieldlabs.RateLimitError, + shieldlabs.ServerError, + ] + for cls in api_errors: + assert issubclass(cls, shieldlabs.ApiError) + for cls in [ + shieldlabs.ApiError, + shieldlabs.APIConnectionError, + shieldlabs.APITimeoutError, + shieldlabs.SignatureVerificationError, + shieldlabs.WebhookParseError, + shieldlabs.ValidationError, + ]: + assert issubclass(cls, shieldlabs.ShieldLabsError) + assert not issubclass(shieldlabs.APITimeoutError, shieldlabs.APIConnectionError) + + +def test_version_and_typing_marker() -> None: + assert shieldlabs.__version__ == "1.0.0" + assert (Path(shieldlabs.__file__).parent / "py.typed").exists() + + +def test_webhooks_module_surface() -> None: + assert shieldlabs.webhooks.SIGNATURE_HEADER == "X-Shield-Signature" + assert shieldlabs.webhooks.SCHEMA_VERSION == "2026-06-01" + assert shieldlabs.webhooks.IdentificationScoredEvent is shieldlabs.IdentificationScoredEvent diff --git a/tests/test_polling.py b/tests/test_polling.py new file mode 100644 index 0000000..8155c68 --- /dev/null +++ b/tests/test_polling.py @@ -0,0 +1,785 @@ +"""identifications.get: the wait-for-verdict rules, for the sync and the async client. + +- ``timeout`` is the total budget of the call. +- The first poll runs at once, then after waits of ``poll_interval`` times 1, 2, 4, 6 and 8, then + 8 again, each capped at ``max(2 s, poll_interval)`` (0.25 s, 0.5 s, 1 s, 1.5 s and every 2 s by + default); the last poll runs at the deadline. +- Each poll is one HTTP attempt with a timeout of ``min(client timeout, max(time left, 1 s))``. +- 429, 5xx, connection errors and timeouts keep the wait going. At the deadline the error of the + last poll is raised; when the last poll answered without a row, the result is ``None``. +- After a 429 the next wait is ``max(ladder step, 1 s, min(Retry-After or 0, 10 s))``, cut to the + deadline (``Retry-After: 0`` and past dates count as 0); a capped ``Retry-After`` longer than + the time left raises the 429 at once. +- 400, 401, 403 and 404 stop the wait at once. +""" + +from __future__ import annotations + +from collections.abc import Iterator +from typing import Any, Callable, Optional, Union + +import httpx +import pytest +import respx + +from _support import ( + API_KEY, + HISTORY_HOST, + REQUEST_ID, + REQUEST_PATH, + FakeTime, + empty_page, + load_json, +) +from shieldlabs import ( + APIConnectionError, + ApiError, + APITimeoutError, + AsyncShieldLabs, + AuthenticationError, + BadRequestError, + Identification, + NotFoundError, + QuotaExceededError, + RateLimitError, + ServerError, + ShieldLabs, + ValidationError, +) +from shieldlabs._client import poll_waits + +pytestmark = pytest.mark.anyio + +Client = Union[ShieldLabs, AsyncShieldLabs] +Answer = Callable[[], Union[httpx.Response, Exception]] + +LADDER_TIMES = [0.0, 0.25, 0.75, 1.75, 3.25, 5.25, 7.25, 9.25, 10.0] +"""When the polls of a default 10-second wait run.""" + + +def _found() -> httpx.Response: + page = load_json("history-page.json") + return httpx.Response(200, json={"data": page["data"][:1], "total": 1}) + + +def _status(code: int, retry_after: Any = None) -> Answer: + def answer() -> httpx.Response: + headers = {"content-type": "application/json"} + if retry_after is not None: + headers["retry-after"] = str(retry_after) + return httpx.Response(code, text='{"error":"request failed"}\n', headers=headers) + + return answer + + +class Server: + """Answers each poll with the next scripted answer (the last one repeats). + + An answer is a response or an exception to raise. The server records when each poll ran on + the fake clock and the timeout of its HTTP attempt, and can let every answer take + ``latency`` seconds. + """ + + def __init__(self, fake_time: FakeTime, *answers: Answer, latency: float = 0.0) -> None: + self._fake_time = fake_time + self._answers = answers + self._latency = latency + self.times: list[float] = [] + self.timeouts: list[float] = [] + + def __call__(self, request: httpx.Request) -> httpx.Response: + self.times.append(round(self._fake_time.now, 6)) + self.timeouts.append(request.extensions["timeout"]["read"]) + self._fake_time.now += self._latency + answer = self._answers[min(len(self.times), len(self._answers)) - 1]() + if isinstance(answer, Exception): + raise answer + return answer + + +@pytest.fixture(params=["sync", "async"]) +def make_client(request: pytest.FixtureRequest, fake_time: FakeTime) -> Callable[..., Client]: + def make(**options: Any) -> Client: + cls = ShieldLabs if request.param == "sync" else AsyncShieldLabs + instance: Client = cls(api_key=API_KEY, **options) + fake_time.install(instance._transport) + return instance + + return make + + +@pytest.fixture +def client(make_client: Callable[..., Client]) -> Client: + return make_client() + + +@pytest.fixture +def mock() -> Iterator[respx.MockRouter]: + with respx.mock(assert_all_called=False) as router: + yield router + + +def serve(mock: respx.MockRouter, fake_time: FakeTime, *answers: Answer, **kw: Any) -> Server: + server = Server(fake_time, *answers, **kw) + mock.get(host=HISTORY_HOST, path=REQUEST_PATH).mock(side_effect=server) + return server + + +async def get(client: Client, *args: Any, **kwargs: Any) -> Optional[Identification]: + if isinstance(client, AsyncShieldLabs): + return await client.identifications.get(*args, **kwargs) + return client.identifications.get(*args, **kwargs) + + +def _take(iterator: Iterator[float], count: int) -> list[float]: + return [next(iterator) for _ in range(count)] + + +# The ladder + + +def test_poll_waits_default_schedule() -> None: + assert _take(poll_waits(0.25), 8) == [0.25, 0.5, 1.0, 1.5, 2.0, 2.0, 2.0, 2.0] + + +@pytest.mark.parametrize( + ("initial", "expected"), + [ + # poll_interval times 1, 2, 4, 6 and 8, then 8 again, each wait capped at + # max(2 s, poll_interval). + (0.01, [0.01, 0.02, 0.04, 0.06, 0.08, 0.08, 0.08, 0.08]), # no lower limit + (0.1, [0.1, 0.2, 0.4, 0.6, 0.8, 0.8, 0.8, 0.8]), + (0.2, [0.2, 0.4, 0.8, 1.2, 1.6, 1.6, 1.6, 1.6]), + (0.3, [0.3, 0.6, 1.2, 1.8, 2.0, 2.0, 2.0, 2.0]), + (0.5, [0.5, 1.0, 2.0, 2.0, 2.0, 2.0, 2.0, 2.0]), + (1.0, [1.0, 2.0, 2.0, 2.0, 2.0, 2.0, 2.0, 2.0]), + (2.0, [2.0] * 8), + (3.0, [3.0] * 8), + (60.0, [60.0] * 8), + ], +) +def test_poll_waits_scale_with_the_poll_interval(initial: float, expected: list[float]) -> None: + assert _take(poll_waits(initial), 8) == pytest.approx(expected) + + +async def test_first_poll_is_immediate_and_found( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + server = serve(mock, fake_time, _found) + identification = await get(client, REQUEST_ID) + assert isinstance(identification, Identification) + assert identification.request_id == REQUEST_ID + assert server.times == [0.0] + assert fake_time.sleeps == [] + + +async def test_polls_on_the_ladder_until_the_row_appears( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + route = mock.get(host=HISTORY_HOST).mock( + side_effect=[empty_page(), empty_page(), empty_page(), _found()] + ) + assert await get(client, REQUEST_ID.upper()) is not None + assert fake_time.sleeps == [0.25, 0.5, 1.0] + assert route.call_count == 4 + assert route.calls.last.request.url.path == REQUEST_PATH + params = route.calls.last.request.url.params + assert (params["limit"], params["offset"]) == ("1", "0") + + +async def test_custom_poll_interval_sets_the_first_wait( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + serve(mock, fake_time, empty_page, empty_page, _found) + assert await get(client, REQUEST_ID, poll_interval=0.5) is not None + assert fake_time.sleeps == [0.5, 1.0] + + +async def test_small_poll_interval_keeps_its_own_ladder_until_the_deadline( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + # 0.1 s gives waits of 0.1, 0.2, 0.4, 0.6 and 0.8 s, then 0.8 s again; the last one is cut + # so that the last poll runs at the deadline. + server = serve(mock, fake_time, empty_page) + assert await get(client, REQUEST_ID, timeout=5, poll_interval=0.1) is None + assert fake_time.sleeps == pytest.approx([0.1, 0.2, 0.4, 0.6, 0.8, 0.8, 0.8, 0.8, 0.5]) + assert server.times == pytest.approx([0.0, 0.1, 0.3, 0.7, 1.3, 2.1, 2.9, 3.7, 4.5, 5.0]) + + +@pytest.mark.parametrize( + ("poll_interval", "sleeps", "times"), + [ + # Waits of p, 2p, 4p, 6p and 8p, then 8p again, each at most max(2 s, p). The last wait + # is cut so that the last poll runs at the 10 s deadline. + (0.25, [0.25, 0.5, 1.0, 1.5, 2.0, 2.0, 2.0, 0.75], LADDER_TIMES), + (1.0, [1.0, 2.0, 2.0, 2.0, 2.0, 1.0], [0.0, 1.0, 3.0, 5.0, 7.0, 9.0, 10.0]), + (3.0, [3.0, 3.0, 3.0, 1.0], [0.0, 3.0, 6.0, 9.0, 10.0]), + ], +) +async def test_the_ladder_of_a_poll_interval_until_the_deadline( + client: Client, + mock: respx.MockRouter, + fake_time: FakeTime, + poll_interval: float, + sleeps: list[float], + times: list[float], +) -> None: + server = serve(mock, fake_time, empty_page) + assert await get(client, REQUEST_ID, poll_interval=poll_interval) is None + assert fake_time.sleeps == sleeps + assert server.times == times + + +# The deadline: total budget, last poll at the deadline, None when nothing was found + + +async def test_last_poll_runs_at_the_deadline_and_none_is_returned( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + server = serve(mock, fake_time, empty_page) + assert await get(client, REQUEST_ID) is None + assert server.times == LADDER_TIMES + assert fake_time.sleeps == [0.25, 0.5, 1.0, 1.5, 2.0, 2.0, 2.0, 0.75] + assert fake_time.now == 10.0 + + +async def test_short_budget_cuts_the_last_wait_to_the_deadline( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + server = serve(mock, fake_time, empty_page) + assert await get(client, REQUEST_ID, timeout=3) is None + assert server.times == [0.0, 0.25, 0.75, 1.75, 3.0] + assert fake_time.sleeps == [0.25, 0.5, 1.0, 1.25] + + +async def test_budget_counts_the_time_polls_take( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + # Every answer takes 0.3 s: the waits keep their length and the last poll still starts at + # the deadline (3 s), so the call ends when that answer arrives. + server = serve(mock, fake_time, empty_page, latency=0.3) + assert await get(client, REQUEST_ID, timeout=3) is None + assert server.times == [0.0, 0.55, 1.35, 2.65, 3.0] + assert fake_time.sleeps == [0.25, 0.5, 1.0, 0.05] + assert fake_time.now == pytest.approx(3.3) + + +async def test_poll_after_the_last_wait_is_final_even_when_a_sleep_ends_early( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + # Event loop timers can fire slightly early. The poll that follows the wait cut to the + # deadline is still the last one, even though a sliver of time seems to be left. + def early_sleep(seconds: float) -> None: + assert len(fake_time.sleeps) < 20, "polling did not stop" + fake_time.sleeps.append(round(seconds, 6)) + fake_time.now += max(0.0, seconds - 0.001) + + async def early_async_sleep(seconds: float) -> None: + early_sleep(seconds) + + if isinstance(client, AsyncShieldLabs): + client._transport._sleep = early_async_sleep + else: + client._transport._sleep = early_sleep + server = serve(mock, fake_time, empty_page) + assert await get(client, REQUEST_ID, timeout=3) is None + assert server.times == [0.0, 0.249, 0.748, 1.747, 2.999] + + +async def test_a_poll_that_ends_after_the_deadline_is_the_last_one( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + server = serve(mock, fake_time, empty_page, latency=4.0) + assert await get(client, REQUEST_ID, timeout=3) is None + assert server.times == [0.0] + assert fake_time.sleeps == [] + + +async def test_a_failed_poll_that_ends_after_the_deadline_raises( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + server = serve(mock, fake_time, _status(504), latency=4.0) + with pytest.raises(ServerError): + await get(client, REQUEST_ID, timeout=3) + assert server.times == [0.0] + + +async def test_timeout_zero_polls_once( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + server = serve(mock, fake_time, empty_page) + assert await get(client, REQUEST_ID, timeout=0) is None + assert server.times == [0.0] + assert fake_time.sleeps == [] + + +async def test_none_when_the_last_poll_finds_nothing_after_errors( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + server = serve( + mock, + fake_time, + _status(503), + lambda: httpx.ConnectError("connection refused"), + empty_page, + ) + assert await get(client, REQUEST_ID, timeout=1) is None + assert server.times == [0.0, 0.25, 0.75, 1.0] + + +# One HTTP attempt per poll and its timeout + + +@pytest.mark.parametrize( + ("client_timeout", "budget", "expected"), + [ + # min(client timeout, max(time left, 1 s)) at the start of each poll. + (10.0, 10.0, [10.0, 9.75, 9.25, 8.25, 6.75, 4.75, 2.75, 1.0, 1.0]), + (2.0, 5.0, [2.0, 2.0, 2.0, 2.0, 1.75, 1.0]), + (0.5, 1.0, [0.5, 0.5, 0.5, 0.5]), + ], +) +async def test_attempt_timeout_is_the_client_timeout_cut_to_the_time_left( + make_client: Callable[..., Client], + mock: respx.MockRouter, + fake_time: FakeTime, + client_timeout: float, + budget: float, + expected: list[float], +) -> None: + server = serve(mock, fake_time, empty_page) + client = make_client(timeout=client_timeout) + assert await get(client, REQUEST_ID, timeout=budget) is None + assert server.timeouts == pytest.approx(expected) + + +async def test_attempt_timeout_counts_the_time_polls_take( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + server = serve(mock, fake_time, empty_page, latency=0.3) + await get(client, REQUEST_ID, timeout=3) + assert server.times == [0.0, 0.55, 1.35, 2.65, 3.0] + assert server.timeouts == pytest.approx([3.0, 2.45, 1.65, 1.0, 1.0]) + + +async def test_each_poll_is_one_http_attempt( + make_client: Callable[..., Client], mock: respx.MockRouter, fake_time: FakeTime +) -> None: + # max_retries applies to single requests; inside the wait a failed poll is not retried. + server = serve(mock, fake_time, _status(500), _status(500), _found) + client = make_client(max_retries=5) + assert await get(client, REQUEST_ID) is not None + assert server.times == [0.0, 0.25, 0.75] + assert fake_time.sleeps == [0.25, 0.5] + + +# Transient errors keep the wait going + + +_TRANSIENT: list[tuple[str, Answer, type[Exception]]] = [ + ("500", _status(500), ServerError), + ("502", lambda: httpx.Response(502, text="bad gateway"), ServerError), + ("503", _status(503), ServerError), + ("connect", lambda: httpx.ConnectError("connection refused"), APIConnectionError), + ("read-timeout", lambda: httpx.ReadTimeout("slow"), APITimeoutError), + ("connect-timeout", lambda: httpx.ConnectTimeout("slow"), APITimeoutError), +] + + +@pytest.mark.parametrize( + ("answer", "error_type"), + [pytest.param(answer, error_type, id=name) for name, answer, error_type in _TRANSIENT], +) +async def test_transient_error_keeps_polling_and_is_raised_at_the_deadline( + client: Client, + mock: respx.MockRouter, + fake_time: FakeTime, + answer: Answer, + error_type: type[Exception], +) -> None: + server = serve(mock, fake_time, answer) + with pytest.raises(error_type): + await get(client, REQUEST_ID, timeout=4) + assert server.times == [0.0, 0.25, 0.75, 1.75, 3.25, 4.0] + assert fake_time.sleeps == [0.25, 0.5, 1.0, 1.5, 0.75] + + +@pytest.mark.parametrize( + "answer", [pytest.param(answer, id=name) for name, answer, _ in _TRANSIENT] +) +async def test_transient_error_then_row_returns_the_row( + client: Client, mock: respx.MockRouter, fake_time: FakeTime, answer: Answer +) -> None: + server = serve(mock, fake_time, answer, answer, _found) + identification = await get(client, REQUEST_ID) + assert identification is not None + assert identification.request_id == REQUEST_ID + assert server.times == [0.0, 0.25, 0.75] + + +async def test_mixed_transient_errors_keep_polling( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + server = serve( + mock, + fake_time, + lambda: httpx.Response(502, text="bad gateway"), + lambda: httpx.ConnectError("connection refused"), + lambda: httpx.ReadTimeout("slow"), + _found, + ) + assert await get(client, REQUEST_ID) is not None + assert server.times == [0.0, 0.25, 0.75, 1.75] + + +async def test_error_of_the_last_poll_is_raised( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + server = serve( + mock, + fake_time, + _status(500), + _status(429), + lambda: httpx.ConnectError("refused at the deadline"), + ) + with pytest.raises(APIConnectionError, match="refused at the deadline"): + await get(client, REQUEST_ID, timeout=0.5) + assert server.times == [0.0, 0.25, 0.5] + + +async def test_last_poll_failing_after_empty_answers_raises( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + server = serve(mock, fake_time, empty_page, _status(500)) + with pytest.raises(ServerError): + await get(client, REQUEST_ID, timeout=0.25) + assert server.times == [0.0, 0.25] + + +# 429 inside the wait + + +async def test_rate_limit_without_retry_after_waits_at_least_one_second( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + server = serve(mock, fake_time, _status(429), _status(429), empty_page, _status(429), _found) + assert await get(client, REQUEST_ID) is not None + # The 0.25 s and 0.5 s steps become 1 s after a 429; the 1.5 s step is already long enough. + assert fake_time.sleeps == [1.0, 1.0, 1.0, 1.5] + assert server.times == [0.0, 1.0, 2.0, 3.0, 4.5] + + +async def test_rate_limit_without_retry_after_until_the_deadline( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + server = serve(mock, fake_time, _status(429)) + with pytest.raises(RateLimitError) as caught: + await get(client, REQUEST_ID, timeout=4) + assert caught.value.status == 429 + assert caught.value.retry_after is None + assert server.times == [0.0, 1.0, 2.0, 3.0, 4.0] + + +async def test_one_second_floor_never_passes_the_deadline( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + server = serve(mock, fake_time, _status(429)) + with pytest.raises(RateLimitError): + await get(client, REQUEST_ID, timeout=0.2) + assert server.times == [0.0, 0.2] + assert fake_time.sleeps == [0.2] + + +async def test_retry_after_is_honoured_then_the_ladder_continues( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + serve(mock, fake_time, _status(429, retry_after=1.5), empty_page, _found) + assert await get(client, REQUEST_ID) is not None + assert fake_time.sleeps == [1.5, 0.5] + + +async def test_retry_after_is_capped_at_ten_seconds( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + server = serve(mock, fake_time, _status(429, retry_after=60), _found) + assert await get(client, REQUEST_ID, timeout=30) is not None + assert server.times == [0.0, 10.0] + + +async def test_retry_after_shorter_than_the_ladder_step_keeps_the_step( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + serve( + mock, + fake_time, + empty_page, + empty_page, + empty_page, + empty_page, + _status(429, retry_after=1), + _status(429, retry_after=0), + _found, + ) + assert await get(client, REQUEST_ID, timeout=30) is not None + assert fake_time.sleeps == [0.25, 0.5, 1.0, 1.5, 2.0, 2.0] + + +@pytest.mark.parametrize( + "retry_after", + [ + pytest.param("0", id="zero"), + pytest.param("Wed, 21 Oct 2015 07:28:00 GMT", id="past-date"), + ], +) +async def test_zero_or_past_retry_after_still_waits_one_second( + client: Client, mock: respx.MockRouter, fake_time: FakeTime, retry_after: str +) -> None: + # Both count as 0: the wait is max(ladder step, 1 s, 0), so the 0.25 s step becomes 1 s. + server = serve(mock, fake_time, _status(429, retry_after=retry_after), empty_page, _found) + assert await get(client, REQUEST_ID) is not None + assert fake_time.sleeps == [1.0, 0.5] + assert server.times == [0.0, 1.0, 1.5] + + +@pytest.mark.parametrize( + "retry_after", + [ + pytest.param("0", id="zero"), + pytest.param("Wed, 21 Oct 2015 07:28:00 GMT", id="past-date"), + ], +) +async def test_zero_retry_after_near_the_deadline_polls_at_the_deadline( + client: Client, mock: respx.MockRouter, fake_time: FakeTime, retry_after: str +) -> None: + server = serve(mock, fake_time, _status(429, retry_after=retry_after)) + with pytest.raises(RateLimitError) as caught: + await get(client, REQUEST_ID, timeout=0.5) + assert caught.value.retry_after == 0.0 + assert server.times == [0.0, 0.5] + assert fake_time.sleeps == [0.5] + + +@pytest.mark.parametrize( + ("retry_after", "sleeps"), + [ + # The 3 s step is longer than the 1 s floor and a shorter Retry-After. A longer + # Retry-After wins, and the ladder then goes on with its next step. + pytest.param(None, [3.0, 3.0], id="no-retry-after"), + pytest.param(2, [3.0, 3.0], id="shorter-retry-after"), + pytest.param(5, [5.0, 3.0], id="longer-retry-after"), + ], +) +async def test_rate_limit_keeps_the_ladder_of_a_long_poll_interval( + client: Client, + mock: respx.MockRouter, + fake_time: FakeTime, + retry_after: Optional[float], + sleeps: list[float], +) -> None: + serve(mock, fake_time, _status(429, retry_after=retry_after), empty_page, _found) + assert await get(client, REQUEST_ID, poll_interval=3) is not None + assert fake_time.sleeps == sleeps + + +async def test_rate_limit_floor_applies_to_a_custom_ladder( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + serve(mock, fake_time, _status(429), empty_page, empty_page, _found) + assert await get(client, REQUEST_ID, poll_interval=0.1) is not None + # The 0.1 s step becomes 1 s after the 429; the ladder then goes on with 0.2 s and 0.4 s. + assert fake_time.sleeps == pytest.approx([1.0, 0.2, 0.4]) + + +async def test_retry_after_between_floor_and_step_is_cut_to_the_deadline( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + # At 3.25 s the step is 2 s, longer than Retry-After (1.5 s) and the 1 s floor, and only + # 1.75 s is left: the wait is cut and the last poll runs at the deadline. + server = serve( + mock, + fake_time, + empty_page, + empty_page, + empty_page, + empty_page, + _status(429, retry_after=1.5), + _found, + ) + assert await get(client, REQUEST_ID, timeout=5) is not None + assert fake_time.sleeps == [0.25, 0.5, 1.0, 1.5, 1.75] + assert server.times == [0.0, 0.25, 0.75, 1.75, 3.25, 5.0] + + +@pytest.mark.parametrize( + ("retry_after", "budget"), + [ + (5, 4), + (60, 4), + (60, 9.5), # capped at 10 s, which is still longer than the time left + ], +) +async def test_retry_after_longer_than_the_time_left_raises_at_once( + client: Client, + mock: respx.MockRouter, + fake_time: FakeTime, + retry_after: float, + budget: float, +) -> None: + server = serve(mock, fake_time, _status(429, retry_after=retry_after), _found) + with pytest.raises(RateLimitError) as caught: + await get(client, REQUEST_ID, timeout=budget) + assert caught.value.retry_after == retry_after + assert server.times == [0.0] + assert fake_time.sleeps == [] + + +async def test_retry_after_later_in_the_wait_raises_when_it_does_not_fit( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + server = serve(mock, fake_time, empty_page, empty_page, empty_page, _status(429, retry_after=9)) + with pytest.raises(RateLimitError): + await get(client, REQUEST_ID) + assert server.times == [0.0, 0.25, 0.75, 1.75] + assert fake_time.now == 1.75 + + +@pytest.mark.parametrize( + ("retry_after", "budget"), + [ + (2, 2), + (60, 10), # capped at 10 s, which fits exactly + ], +) +async def test_retry_after_equal_to_the_time_left_polls_at_the_deadline( + client: Client, + mock: respx.MockRouter, + fake_time: FakeTime, + retry_after: float, + budget: float, +) -> None: + server = serve(mock, fake_time, _status(429, retry_after=retry_after), _found) + assert await get(client, REQUEST_ID, timeout=budget) is not None + assert server.times == [0.0, budget] + + +# Errors that stop the wait at once + +STOP_STATUSES = [ + (400, BadRequestError), + (401, AuthenticationError), + (403, AuthenticationError), + (404, NotFoundError), +] + + +@pytest.mark.parametrize(("status", "error_type"), STOP_STATUSES) +async def test_client_errors_stop_polling_at_once( + client: Client, + mock: respx.MockRouter, + fake_time: FakeTime, + status: int, + error_type: type[ApiError], +) -> None: + server = serve(mock, fake_time, _status(status)) + with pytest.raises(error_type) as caught: + await get(client, REQUEST_ID) + assert caught.value.status == status + assert server.times == [0.0] + assert fake_time.sleeps == [] + + +@pytest.mark.parametrize(("status", "error_type"), STOP_STATUSES) +async def test_client_error_after_a_transient_error_stops_at_once( + client: Client, + mock: respx.MockRouter, + fake_time: FakeTime, + status: int, + error_type: type[ApiError], +) -> None: + server = serve(mock, fake_time, _status(503), _status(status), _found) + with pytest.raises(error_type): + await get(client, REQUEST_ID) + assert server.times == [0.0, 0.25] + + +async def test_authentication_error_message_is_kept( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + mock.get(host=HISTORY_HOST).respond( + 401, text='{"error":"invalid api key"}\n', headers={"content-type": "text/plain"} + ) + with pytest.raises(AuthenticationError, match="invalid api key"): + await get(client, REQUEST_ID) + + +async def test_not_found_page_stops_polling( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + route = mock.get(host=HISTORY_HOST).respond(404, text="404 page not found") + with pytest.raises(NotFoundError, match="404 page not found"): + await get(client, REQUEST_ID) + assert route.call_count == 1 + + +@pytest.mark.parametrize(("status", "error_type"), [(402, QuotaExceededError), (409, ApiError)]) +async def test_other_error_statuses_stop_polling_too( + client: Client, + mock: respx.MockRouter, + fake_time: FakeTime, + status: int, + error_type: type[ApiError], +) -> None: + server = serve(mock, fake_time, _status(status)) + with pytest.raises(error_type): + await get(client, REQUEST_ID) + assert server.times == [0.0] + + +# wait=False and validation + + +async def test_wait_false_sends_one_request( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + route = mock.get(host=HISTORY_HOST).mock(side_effect=[empty_page(), _found()]) + assert await get(client, REQUEST_ID, wait=False) is None + assert await get(client, REQUEST_ID, wait=False) is not None + assert route.call_count == 2 + assert fake_time.sleeps == [] + + +async def test_wait_false_uses_regular_retries( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + route = mock.get(host=HISTORY_HOST).mock(side_effect=lambda request: _status(429)()) + with pytest.raises(RateLimitError): + await get(client, REQUEST_ID, wait=False) + assert route.call_count == 3 + assert fake_time.sleeps == [1.0, 1.0] # a 429 without Retry-After waits at least 1 s + + +async def test_connection_error_without_wait_is_raised( + client: Client, mock: respx.MockRouter, fake_time: FakeTime +) -> None: + mock.get(host=HISTORY_HOST).mock(side_effect=httpx.ConnectError("down")) + with pytest.raises(APIConnectionError, match="down"): + await get(client, REQUEST_ID, wait=False) + + +@pytest.mark.parametrize( + ("args", "kwargs"), + [ + (("not-a-uuid",), {}), + ((REQUEST_ID,), {"timeout": -1}), + ((REQUEST_ID,), {"timeout": float("inf")}), + ((REQUEST_ID,), {"poll_interval": 0}), + ((REQUEST_ID,), {"poll_interval": float("nan")}), + ], +) +async def test_get_validates_before_sending( + client: Client, mock: respx.MockRouter, args: tuple[Any, ...], kwargs: dict[str, Any] +) -> None: + route = mock.get(host=HISTORY_HOST).mock(return_value=empty_page()) + with pytest.raises(ValidationError): + await get(client, *args, **kwargs) + assert route.call_count == 0 diff --git a/tests/test_readme.py b/tests/test_readme.py new file mode 100644 index 0000000..ae737b4 --- /dev/null +++ b/tests/test_readme.py @@ -0,0 +1,78 @@ +"""README code samples: every Python block compiles, and the Quick start runs as pasted.""" + +from __future__ import annotations + +import ast +import importlib.util +import re +from datetime import datetime, timezone +from pathlib import Path +from types import ModuleType + +import pytest +import respx + +from _support import API_KEY, HISTORY_HOST, load_bytes, load_json, sign + +README = Path(__file__).resolve().parent.parent / "README.md" +SECRET = "whsec_your_signing_secret" +_PYTHON_BLOCK = re.compile(r"```python\n(.*?)```", re.DOTALL) + + +def _readme() -> str: + return README.read_text(encoding="utf-8") + + +def _python_blocks() -> list[str]: + return _PYTHON_BLOCK.findall(_readme()) + + +def _quick_start() -> str: + section = _readme().split("\n## Quick start\n", 1)[1].split("\n## ", 1)[0] + blocks = _PYTHON_BLOCK.findall(section) + assert len(blocks) == 1 + return blocks[0] + + +def _import_file(path: Path) -> ModuleType: + spec = importlib.util.spec_from_file_location(path.stem, path) + assert spec is not None + assert spec.loader is not None + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +@pytest.mark.parametrize("index", range(len(_python_blocks()))) +def test_readme_python_blocks_compile(index: int) -> None: + # Top-level await is allowed because the async section shows statements, not a module. + compile( + _python_blocks()[index], + f"README.md python block {index}", + "exec", + flags=ast.PyCF_ALLOW_TOP_LEVEL_AWAIT, + ) + + +def test_quick_start_runs_as_pasted( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, capsys: pytest.CaptureFixture[str] +) -> None: + monkeypatch.setenv("SHIELDLABS_API_KEY", API_KEY) + monkeypatch.setenv("SHIELDLABS_WEBHOOK_SECRET", SECRET) + script = tmp_path / "readme_quick_start.py" + script.write_text(_quick_start(), encoding="utf-8") + row = dict(load_json("history-page.json")["data"][1]) # trusted, with device signals + row["created_at"] = datetime.now(timezone.utc).strftime("%Y-%m-%d %H:%M:%S.000") + + with respx.mock(assert_all_called=False) as router: + route = router.get(host=HISTORY_HOST).respond(200, json={"data": [row], "total": 1}) + quick_start = _import_file(script) + assert quick_start.allow_signup(row["request_id"]) is True + # One identification authorizes one action: the same request ID is refused next time. + assert quick_start.allow_signup(row["request_id"]) is False + quick_start.client.close() + assert route.call_count == 2 + + body = load_bytes("webhook-ping.raw.txt") + quick_start.on_webhook(body, sign(SECRET, body)) + assert capsys.readouterr().out == "webhook.ping\n" diff --git a/tests/test_verify_webhook.py b/tests/test_verify_webhook.py deleted file mode 100644 index 2d3fe2b..0000000 --- a/tests/test_verify_webhook.py +++ /dev/null @@ -1,36 +0,0 @@ -import hashlib -import hmac - -from shieldlabs import verify_webhook - -SECRET = "whsec_test_secret" -BODY = b'{"event_type":"webhook.ping","schema_version":"2026-06-01","created_at":"2026-06-26T14:20:42Z"}' - - -def _sign(secret: str, body: bytes) -> str: - return "sha256=" + hmac.new(secret.encode(), body, hashlib.sha256).hexdigest() - - -def test_accepts_valid_signature(): - assert verify_webhook(BODY, _sign(SECRET, BODY), SECRET) is True - - -def test_accepts_bytearray(): - buf = bytearray(BODY) - assert verify_webhook(buf, _sign(SECRET, BODY), SECRET) is True - - -def test_rejects_wrong_secret(): - assert verify_webhook(BODY, _sign(SECRET, BODY), "other") is False - - -def test_rejects_tampered_body(): - assert verify_webhook(BODY + b" ", _sign(SECRET, BODY), SECRET) is False - - -def test_rejects_missing_header(): - assert verify_webhook(BODY, "", SECRET) is False - - -def test_rejects_truncated_signature(): - assert verify_webhook(BODY, "sha256=ab", SECRET) is False diff --git a/tests/test_webhooks.py b/tests/test_webhooks.py new file mode 100644 index 0000000..1406e85 --- /dev/null +++ b/tests/test_webhooks.py @@ -0,0 +1,221 @@ +"""Webhook signature vectors and typed event parsing.""" + +from __future__ import annotations + +import base64 +import json +import warnings +from datetime import datetime, timezone +from typing import Any + +import pytest + +from _support import load_bytes, load_json, sign +from shieldlabs import ( + IdentificationScoredEvent, + ShieldLabsWarning, + SignatureVerificationError, + UnknownWebhookEvent, + WebhookParseError, + WebhookPingEvent, + webhooks, +) + +VECTORS = load_json("webhook-signature-vectors.json") +SECRET = "whsec_00112233445566778899aabbccddeeff" +NORMALIZATION = {c["name"]: c for c in load_json("normalization-cases.json")["cases"]} + + +def _secret(vector: dict[str, Any]) -> Any: + return vector["secret"] if "secret" in vector else vector["secrets"] + + +def test_all_signature_vectors_present() -> None: + assert VECTORS["header_name"] == webhooks.SIGNATURE_HEADER + assert len(VECTORS["vectors"]) == 21 + assert any("secrets" in vector for vector in VECTORS["vectors"]) + + +@pytest.mark.parametrize("vector", VECTORS["vectors"], ids=[v["name"] for v in VECTORS["vectors"]]) +def test_signature_vector_bytes_and_str(vector: dict[str, Any]) -> None: + body = base64.b64decode(vector["body_base64"]) + assert body.decode("utf-8") == vector["body"] + header = vector["signature_header"] + expected = vector["valid"] + assert webhooks.verify_signature(body, header, _secret(vector)) is expected + assert webhooks.verify_signature(vector["body"], header, _secret(vector)) is expected + assert webhooks.verify_signature(bytearray(body), header, _secret(vector)) is expected + assert webhooks.verify_signature(memoryview(body), header, _secret(vector)) is expected + + +@pytest.mark.parametrize("vector", VECTORS["vectors"], ids=[v["name"] for v in VECTORS["vectors"]]) +def test_signature_vector_construct_event(vector: dict[str, Any]) -> None: + body = base64.b64decode(vector["body_base64"]) + if vector["valid"]: + event = webhooks.construct_event(body, vector["signature_header"], _secret(vector)) + assert event.event_type == json.loads(body)["event_type"] + else: + with pytest.raises(SignatureVerificationError): + webhooks.construct_event(body, vector["signature_header"], _secret(vector)) + + +def test_verify_rejects_missing_header_and_empty_secrets() -> None: + body = load_bytes("webhook-ping.raw.txt") + header = sign(SECRET, body) + assert webhooks.verify_signature(body, header, SECRET) + assert not webhooks.verify_signature(body, None, SECRET) + assert not webhooks.verify_signature(body, header, []) + assert not webhooks.verify_signature(body, header, ["", ""]) + assert not webhooks.verify_signature(body, header, None) # type: ignore[arg-type] + assert not webhooks.verify_signature(body, "sha256 =" + header[7:], SECRET) + assert not webhooks.verify_signature(body, "SHA256=" + header[7:], SECRET) + assert webhooks.verify_signature(body, header, ("whsec_other", SECRET)) + + +def test_verify_rejects_parsed_json_and_bad_secret_types() -> None: + body = load_bytes("webhook-ping.raw.txt") + header = sign(SECRET, body) + with pytest.raises(TypeError, match="raw request body"): + webhooks.verify_signature(json.loads(body), header, SECRET) # type: ignore[arg-type] + with pytest.raises(TypeError): + webhooks.verify_signature(body, header, SECRET.encode()) # type: ignore[arg-type] + with pytest.raises(TypeError): + webhooks.verify_signature(body, header, [SECRET.encode()]) # type: ignore[list-item] + with pytest.raises(TypeError): + webhooks.verify_signature(body, header, 42) # type: ignore[arg-type] + + +def test_verify_returns_false_for_unencodable_str_payload() -> None: + assert not webhooks.verify_signature("\ud800", "sha256=" + "0" * 64, SECRET) + with pytest.raises(SignatureVerificationError): + webhooks.construct_event("\ud800", "sha256=" + "0" * 64, SECRET) + + +def test_scored_event_from_raw_bytes_matches_normalization() -> None: + body = load_bytes("webhook-identification-scored.raw.txt") + event = webhooks.construct_event(body, sign(SECRET, body), SECRET) + assert isinstance(event, IdentificationScoredEvent) + assert event.event_type == "identification.scored" + assert event.schema_version == webhooks.SCHEMA_VERSION + assert event.created_at == datetime(2026, 9, 30, 12, 34, 57, 482913, tzinfo=timezone.utc) + assert event.data.to_dict() == NORMALIZATION["webhook_scored"]["expected"] + assert event.data.traffic_source.landing_url.endswith("utm_medium=cpc&gclid=abc123") + assert event.raw == load_json("webhook-identification-scored.json") + + +def test_scored_event_accepts_str_body() -> None: + body = load_bytes("webhook-identification-scored.raw.txt") + event = webhooks.construct_event(body.decode("utf-8"), sign(SECRET, body), SECRET) + assert isinstance(event, IdentificationScoredEvent) + + +def test_ping_event_from_raw_bytes() -> None: + body = load_bytes("webhook-ping.raw.txt") + event = webhooks.construct_event(body, sign(SECRET, body), [SECRET]) + assert isinstance(event, WebhookPingEvent) + assert event.event_type == "webhook.ping" + assert event.created_at == datetime(2026, 9, 30, 12, 34, 56, tzinfo=timezone.utc) + assert event.raw == load_json("webhook-ping.json") + + +def test_rate_limited_event() -> None: + body = json.dumps(load_json("webhook-rate-limited.json")).encode() + event = webhooks.construct_event(body, sign(SECRET, body), SECRET) + assert isinstance(event, IdentificationScoredEvent) + assert event.data.risk_score == 999 + assert event.data.is_rate_limited + assert event.data.risk_band == "rate_limited" + assert [(s.name, s.weight) for s in event.data.signals] == [("rate_limited", 999)] + assert event.data.to_dict() == NORMALIZATION["webhook_rate_limited"]["expected"] + + +def test_dashboard_test_delivery_with_17_flags() -> None: + delivery = load_json("webhook-test-delivery.json") + assert len(delivery["data"]["detection_flags"]) == 17 + body = json.dumps(delivery, separators=(",", ":"), sort_keys=True).encode() + event = webhooks.construct_event(body, sign(SECRET, body), SECRET) + assert isinstance(event, IdentificationScoredEvent) + flags = event.data.detection_flags + assert flags.browser_automation is False + assert flags.search_bot is False + assert flags.proxy + assert flags.datacenter_ip + assert flags.abuser + assert event.data.user_hid is None + assert event.data.to_dict() == NORMALIZATION["webhook_test_delivery"]["expected"] + + +def test_unknown_event_type_is_not_an_error() -> None: + envelope = { + "event_type": "identification.refined", + "schema_version": "2026-06-01", + "created_at": "2026-09-30T12:00:00Z", + "data": {"request_id": "02f1d973-84db-4156-a7f7-e799e6bf389b"}, + "extra": True, + } + body = json.dumps(envelope).encode() + event = webhooks.construct_event(body, sign(SECRET, body), SECRET) + assert isinstance(event, UnknownWebhookEvent) + assert event.event_type == "identification.refined" + assert event.data == envelope["data"] + assert event.raw == envelope + + no_data = json.dumps({"event_type": "x", "schema_version": "2026-06-01"}).encode() + parsed = webhooks.construct_event(no_data, sign(SECRET, no_data), SECRET) + assert isinstance(parsed, UnknownWebhookEvent) + assert parsed.data is None + assert parsed.created_at is None + + +def test_unknown_schema_version_warns_but_parses() -> None: + body = ( + b'{"event_type":"webhook.ping","schema_version":"2027-01-01",' + b'"created_at":"2027-01-01T00:00:00Z"}' + ) + with pytest.warns(ShieldLabsWarning, match="schema_version"): + event = webhooks.construct_event(body, sign(SECRET, body), SECRET) + assert isinstance(event, WebhookPingEvent) + assert event.schema_version == "2027-01-01" + + missing = b'{"event_type":"webhook.ping"}' + with pytest.warns(ShieldLabsWarning): + event = webhooks.construct_event(missing, sign(SECRET, missing), SECRET) + assert event.schema_version == "" + + +def test_known_schema_version_does_not_warn() -> None: + body = load_bytes("webhook-ping.raw.txt") + with warnings.catch_warnings(): + warnings.simplefilter("error") + webhooks.construct_event(body, sign(SECRET, body), SECRET) + + +@pytest.mark.parametrize( + ("body", "message"), + [ + (b"not json", "not valid JSON"), + (b"\xff\xfe\x00", "not valid JSON"), + (b"[1, 2]", "not a JSON object"), + (b'{"schema_version":"2026-06-01"}', "no event_type"), + (b'{"event_type":7,"schema_version":"2026-06-01"}', "no event_type"), + ( + b'{"event_type":"identification.scored","schema_version":"2026-06-01"}', + "no data object", + ), + ( + b'{"event_type":"identification.scored","schema_version":"2026-06-01","data":[]}', + "no data object", + ), + ], +) +def test_parse_errors_after_valid_signature(body: bytes, message: str) -> None: + with pytest.raises(WebhookParseError, match=message): + webhooks.construct_event(body, sign(SECRET, body), SECRET) + + +def test_scored_event_is_hashable_free_of_raw_in_equality() -> None: + body = load_bytes("webhook-ping.raw.txt") + first = webhooks.construct_event(body, sign(SECRET, body), SECRET) + second = webhooks.construct_event(body, sign(SECRET, body), SECRET) + assert first == second + assert "raw" not in repr(first)