diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index e1755be..c7353fa 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -37,6 +37,11 @@ jobs: ruff format --check . - name: Type check run: mypy --strict src + - name: Check consumed wire types + run: python scripts/generate_wire.py --check + - name: Check breaking schema changes + if: matrix.python-version == '3.12' + run: python scripts/check_wire_drift.py - name: Test run: pytest -q --cov=shieldlabs --cov-report=term-missing @@ -70,6 +75,12 @@ jobs: python -m pip install build==1.6.1 twine==7.0.0 python -m build twine check --strict dist/* + - name: Test installed wheel + run: | + python -m venv /tmp/wheel-consumer + /tmp/wheel-consumer/bin/pip install dist/*.whl + cd /tmp + /tmp/wheel-consumer/bin/python "$GITHUB_WORKSPACE/scripts/smoke_wheel.py" generated: name: Generated API types @@ -78,7 +89,12 @@ jobs: - uses: actions/checkout@11d5960a326750d5838078e36cf38b85af677262 with: persist-credentials: false + - uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 + with: + python-version: "3.12" - name: Rebuild generated files - run: ./generate.sh + run: | + python3 -m pip install -e ".[dev]" + ./generate.sh - name: Fail if generated files drifted run: git diff --exit-code diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index 5e9ba13..736f506 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -39,11 +39,19 @@ jobs: ruff check . ruff format --check . mypy --strict src + python scripts/generate_wire.py --check + python scripts/check_wire_drift.py pytest -q - name: Build run: | python -m build twine check --strict dist/* + - name: Test installed wheel + run: | + python -m venv /tmp/wheel-consumer + /tmp/wheel-consumer/bin/pip install dist/*.whl + cd /tmp + /tmp/wheel-consumer/bin/python "$GITHUB_WORKSPACE/scripts/smoke_wheel.py" - uses: actions/upload-artifact@ea165f8d65b6e75b540449e92b4886f43607fa02 # v4.6.2 with: name: dist diff --git a/.gitignore b/.gitignore index dc5d337..23932d8 100644 --- a/.gitignore +++ b/.gitignore @@ -8,6 +8,8 @@ dist/ # Virtual environments .venv/ +.wheel-consumer/ +.wire-drift-*/ venv/ # Tooling caches and reports diff --git a/CHANGELOG.md b/CHANGELOG.md index d6427d3..9c4a13b 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -4,11 +4,14 @@ All notable changes to this project are documented in this file. The format foll [Keep a Changelog](https://keepachangelog.com/en/1.1.0/) and the project uses [Semantic Versioning](https://semver.org/spec/v2.0.0.html). -## [Unreleased] +## [1.0.1] - 2026-10-05 ### Added -- `sync.sh` downloads the OpenAPI description and `generate.sh` rebuilds `generated/` from it. The supported client is unchanged. +- `sync.sh` downloads the OpenAPI description. Schema-derived wire fields now drive History, + profile and webhook normalization, with generated request parameter types and CI checks + for stale output and incompatible schema changes. Public models and tolerant decoding stay + unchanged. The strict reference client in `generated/` remains separate. ## [1.0.0] - 2026-09-30 diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index c51a850..2fcaa3d 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -18,6 +18,8 @@ Run the same checks as CI: ```bash ruff check . && ruff format --check . && mypy --strict src +python scripts/generate_wire.py --check +python scripts/check_wire_drift.py pytest -q --cov=shieldlabs --cov-report=term-missing ``` @@ -30,6 +32,46 @@ pytest -q --cov=shieldlabs --cov-report=term-missing - The FastAPI example has its own smoke test: `pip install -r examples/requirements.txt && pytest tests/test_example_app.py`. +## Updating the HTTP contract + +`resources/shieldlabs-api.yaml` is the bundled OpenAPI input. Run +`python scripts/generate_wire.py` after updating it. The deterministic output +`src/shieldlabs/_generated_wire.py` is included in the wheel and is read by the real +History/profile/webhook normalizers and request builders. PyYAML and the pinned formatter +are development dependencies only; the installed SDK still depends only on httpx. + +The generated fields describe wire types, while the boundary helpers retain missing/null/ +malformed-value defaults, unknown strings and raw fields. They do not validate entire +responses or coerce UUIDs/dates. The strict models under `generated/` require extra runtime +dependencies and reject values the supported client accepts, so they remain reference code. +`./generate.sh` regenerates both layers (Docker is needed only for the reference client). + +The mutation check regenerates temporary schemas and type-checks copies of the actual SDK +source. Renamed fields, incompatible types, query parameters and headers must be rejected; +optional additive fields and parameters must compile. Unsupported new required parameters +on the consumed HTTP operations fail generation, including inherited path-level parameters. +The profile request path comes from the operation's OpenAPI route; a mutation test verifies +the changed route reaches an actual mocked HTTP request. Ping timestamp and version fields +are checked against the ping model separately from scored events. The checks never edit the +checked-in API description. + +History's route template also comes from OpenAPI, with lookup values escaped before template +substitution. The current HTTP operations require GET; changing their method fails generation. +The generated lookup enum must match the real validation list before generation proceeds. +Public and local IP objects have separate generated fields, so one can change without hiding +an incompatible change in the other. Both webhook discriminator definitions are checked +against the supported envelope field and event values before generating. + +To check the installed artifact without an editable checkout: + +```bash +python -m pip install build +python -m build +python -m venv .wheel-consumer +.wheel-consumer/bin/pip install dist/*.whl +.wheel-consumer/bin/python scripts/smoke_wheel.py +``` + ## Shared test fixtures `tests/data/` holds the test fixtures that every ShieldLabs server SDK passes: History API diff --git a/README.md b/README.md index 0be23b7..f249841 100644 --- a/README.md +++ b/README.md @@ -446,11 +446,16 @@ Keys and request bodies are never logged. Each request carries ## Development -Refresh the generated client when the API description changes. This does not replace the supported library in this repository. +The supported SDK consumes schema-derived wire fields for History, domain profiles and +webhooks, and generated request parameter types. Its public models, tolerant decoding, +retries and signature verification remain unchanged. The strict reference client under +`generated/` is separate and is not installed as part of the package. ```bash ./sync.sh # download the current OpenAPI description into resources/ -./generate.sh # rebuild generated/ from that file +python scripts/generate_wire.py # rebuild the wire types used by the supported SDK +python scripts/generate_wire.py --check # reject stale wire types +./generate.sh # rebuild both wire types and the reference client (requires Docker) ``` diff --git a/generate.sh b/generate.sh index 4293a5e..8fda005 100755 --- a/generate.sh +++ b/generate.sh @@ -3,6 +3,9 @@ set -euo pipefail cd "$(dirname "${BASH_SOURCE[0]}")" +# The supported package consumes these fields, rather than the strict reference transport. +python3 scripts/generate_wire.py + if ! docker info >/dev/null 2>&1; then echo "Docker is not running. Start Docker and run this script again." >&2 exit 1 diff --git a/pyproject.toml b/pyproject.toml index ceec1cf..6d894e9 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -45,7 +45,8 @@ dev = [ "pytest>=7.4,<10", "pytest-cov>=4.1,<8", "respx>=0.20.2,<0.24", - "ruff>=0.6,<0.17", + "ruff==0.16.10", + "PyYAML==6.0.3", ] [project.urls] @@ -62,7 +63,7 @@ path = "src/shieldlabs/_version.py" packages = ["src/shieldlabs"] [tool.hatch.build.targets.sdist] -include = ["src/shieldlabs", "tests", "examples", "README.md", "CHANGELOG.md", "LICENSE"] +include = ["src/shieldlabs", "tests", "examples", "scripts", "resources", "README.md", "CONTRIBUTING.md", "CHANGELOG.md", "LICENSE", "pyproject.toml"] [tool.pytest.ini_options] testpaths = ["tests"] @@ -98,6 +99,7 @@ keep-runtime-typing = true [tool.ruff.lint.per-file-ignores] "tests/**" = ["PT011"] +"src/shieldlabs/_generated_wire.py" = ["N815"] [tool.mypy] strict = true diff --git a/scripts/check_wire_drift.py b/scripts/check_wire_drift.py new file mode 100644 index 0000000..e6309fd --- /dev/null +++ b/scripts/check_wire_drift.py @@ -0,0 +1,265 @@ +"""Prove regenerated breaking schemas fail against the real client source.""" + +from __future__ import annotations + +import copy +import shutil +import subprocess +import sys +import tempfile +from pathlib import Path +from typing import Optional + +import yaml + +ROOT = Path(__file__).resolve().parents[1] + + +def run(command: list[str]) -> subprocess.CompletedProcess[str]: + return subprocess.run(command, cwd=ROOT, text=True, capture_output=True, timeout=90) + + +def main() -> None: + original = yaml.safe_load((ROOT / "resources/shieldlabs-api.yaml").read_text()) + # Mutate fields actually read by the supported client, never test-only typed literals. + cases = [ + ("history score rename", "HistoryRow", "score", "rename"), + ("history score type", "HistoryRow", "score", "string"), + ("history total type", "HistoryPage", "total", "string"), + ("profile Weight type", "DomainProfile", "Weight", "string"), + ("profile Domain rename", "DomainProfile", "Domain", "rename"), + ("webhook score type", "IdentificationScoredData", "risk_score", "string"), + ("webhook flag type", "DetectionFlags", "vpn", "string"), + ("webhook signal weight type", "Signal", "weight", "string"), + ("webhook timestamp type", "IdentificationScoredEvent", "created_at", "integer"), + ("ping timestamp type", "WebhookPingEvent", "created_at", "integer"), + ("ping timestamp rename", "WebhookPingEvent", "created_at", "rename"), + ("ping schema_version type", "WebhookPingEvent", "schema_version", "integer"), + ("ping schema_version rename", "WebhookPingEvent", "schema_version", "rename"), + ("query limit rename", "HistoryLimit", "name", "rename_parameter"), + ("query limit type", "HistoryLimit", "schema", "parameter_type"), + ("search type rename", "HistorySearchType", "name", "rename_parameter"), + ("profile header rename", "ShieldDomain", "name", "rename_parameter"), + ("query limit moved to header", "HistoryLimit", "in", "move_parameter"), + ("profile header moved to query", "ShieldDomain", "in", "move_parameter"), + ("local IP type", "IpInfo", "ip", "local_type"), + ("local country type", "IpInfo", "country", "local_type"), + ("ping discriminator rename", "WebhookPingEvent", "event_type", "rename"), + ("ping discriminator type", "WebhookPingEvent", "event_type", "integer"), + ] + with tempfile.TemporaryDirectory(prefix=".wire-drift-", dir=ROOT) as directory: + scratch = Path(directory) + shutil.copytree(ROOT / "src", scratch / "src", ignore=shutil.ignore_patterns("__pycache__")) + schema_path = scratch / "api.yaml" + output = scratch / "src/shieldlabs/_generated_wire.py" + + def check( + document: dict, + label: str, + should_pass: bool, + rejects_required: bool = False, + rejects_contract: Optional[str] = None, + ) -> None: + schema_path.write_text(yaml.safe_dump(document, sort_keys=False)) + generated = run( + [ + sys.executable, + "scripts/generate_wire.py", + "--schema", + str(schema_path), + "--output", + str(output), + ] + ) + if generated.returncode: + if rejects_required and "Unsupported required parameter" in generated.stderr: + print(f"PASS {label}: generator rejects unsupported required parameter") + return + if rejects_contract and rejects_contract in generated.stderr: + print(f"PASS {label}: generator rejects incompatible contract") + return + raise RuntimeError(f"{label}: generation failed unexpectedly\n{generated.stderr}") + if rejects_required: + raise RuntimeError(f"{label}: generator accepted unsupported required parameter") + if rejects_contract: + raise RuntimeError(f"{label}: generator accepted incompatible contract") + result = run( + [ + sys.executable, + "-m", + "mypy", + "--strict", + "--no-incremental", + "--cache-dir", + str(scratch / "cache"), + str(scratch / "src"), + ] + ) + if (result.returncode == 0) != should_pass: + raise RuntimeError( + f"{label}: unexpected compile result\n{result.stdout}{result.stderr}" + ) + if not should_pass and not any( + name in result.stdout + for name in ( + "_models.py", + "_client.py", + "_management.py", + "webhooks.py", + "_validation.py", + ) + ): + raise RuntimeError(f"{label}: did not fail in a real SDK consumer\n{result.stdout}") + print(f"PASS {label}: {'compiles' if should_pass else 'consumer rejects change'}") + + check(original, "baseline", True) + for label, model, field, change in cases: + document = copy.deepcopy(original) + if change in ("rename_parameter", "parameter_type", "move_parameter"): + parameter = document["components"]["parameters"][model] + if change == "rename_parameter": + parameter["name"] += "_changed" + elif change == "move_parameter": + parameter["in"] = "header" if parameter["in"] == "query" else "query" + else: + parameter["schema"] = {"type": "string"} + elif change == "local_type": + local = copy.deepcopy(document["components"]["schemas"]["IpInfo"]) + local["properties"][field] = {"type": "integer"} + document["components"]["schemas"]["IdentificationScoredData"]["properties"][ + "local_ip" + ] = local + else: + properties = document["components"]["schemas"][model]["properties"] + if change == "rename": + properties[field + "_changed"] = properties.pop(field) + else: + properties[field] = {"type": change} + rejects_required = change in ("rename_parameter", "move_parameter") and model in ( + "HistorySearchType", + "ShieldDomain", + ) + rejects_contract = ( + "Unsupported webhook discriminator" + if model == "WebhookPingEvent" and field == "event_type" + else None + ) + check( + document, + label, + False, + rejects_required=rejects_required, + rejects_contract=rejects_contract, + ) + document = copy.deepcopy(original) + document["components"]["parameters"]["HistorySearchType"]["schema"]["enum"].append( + "account_id" + ) + check(document, "lookup enum addition", False, rejects_contract="Unsupported lookup types") + for operation_id in ("searchHistory", "getDomainProfile"): + document = copy.deepcopy(original) + path = next( + value + for value in document["paths"].values() + if value.get("get", {}).get("operationId") == operation_id + ) + path["post"] = path.pop("get") + check( + document, + f"{operation_id} GET to POST", + False, + rejects_contract="Unsupported HTTP method", + ) + document = copy.deepcopy(original) + document["components"]["schemas"]["HistoryRow"]["properties"]["new_optional_field"] = { + "type": "string" + } + check(document, "optional additive field", True) + for operation_id in ("searchHistory", "getDomainProfile"): + for location in ("query", "header", "path"): + for inherited in (False, True): + for required in (False, True): + document = copy.deepcopy(original) + for path in document["paths"].values(): + operation = path.get("get", {}) + if operation.get("operationId") == operation_id: + target = path if inherited else operation + target.setdefault("parameters", []).append( + { + "name": "new_parameter", + "in": location, + "required": required, + "schema": {"type": "string"}, + } + ) + break + check( + document, + f"{operation_id} {location} " + f"{'inherited' if inherited else 'operation'} " + f"{'required' if required else 'optional'}", + not required, + rejects_required=required, + ) + document = copy.deepcopy(original) + document["paths"]["/v2/profile"] = document["paths"].pop("/v1/profile") + history_route = next( + path + for path, value in document["paths"].items() + if value.get("get", {}).get("operationId") == "searchHistory" + ) + document["paths"]["/api/v2/history/{search_type}/{value}"] = document["paths"].pop( + history_route + ) + for path in document["paths"].values(): + operation = path.get("get", {}) + if operation.get("operationId") in ("searchHistory", "getDomainProfile"): + operation.setdefault("parameters", []).extend( + [ + { + "name": "optional_query", + "in": "query", + "required": False, + "schema": {"type": "string", "default": "do-not-send"}, + }, + { + "name": "X-Optional-Header", + "in": "header", + "required": False, + "schema": {"type": "string", "default": "do-not-send"}, + }, + ] + ) + check(document, "profile and History route changes", True) + result = run( + [ + sys.executable, + "-B", + "-c", + "import sys; sys.path.insert(0, sys.argv[1]); import httpx; " + "from shieldlabs import ShieldLabsManagement, ShieldLabs; " + "requests=[]; transport=httpx.MockTransport(lambda r: " + "(requests.append(r), httpx.Response(200,json={}))[1]); " + "http=httpx.Client(transport=transport); " + "client=ShieldLabsManagement(secret_key='fixture'," + "domain='example.com',http_client=http); " + "client.get_profile(); " + "history=ShieldLabs(api_key='sec_aaaaaaaa-bbbbbbbb-cccccccc',http_client=http); " + "history.history.search('user_hid','a@b+c%next',limit=1); " + "assert [r.url.path for r in requests] == " + "['/v2/profile','/api/v2/history/user_hid/a@b+c%next']; " + "assert requests[1].url.raw_path.split(b'?')[0] == " + "b'/api/v2/history/user_hid/a@b+c%25next'; " + "assert not requests[0].url.params; " + "assert dict(requests[1].url.params) == {'limit':'1','offset':'0'}; " + "assert all('X-Optional-Header' not in r.headers for r in requests); http.close()", + str(scratch / "src"), + ] + ) + if result.returncode: + raise RuntimeError(f"Generated profile route not used at runtime: {result.stderr}") + print("PASS real HTTP: changed routes used; escaping preserved; optional defaults not sent") + + +if __name__ == "__main__": + main() diff --git a/scripts/generate_wire.py b/scripts/generate_wire.py new file mode 100644 index 0000000..b9764c2 --- /dev/null +++ b/scripts/generate_wire.py @@ -0,0 +1,289 @@ +"""Generate the dependency-free wire field layer consumed by the supported client.""" + +from __future__ import annotations + +import argparse +import ast +import keyword +import re +import subprocess +import sys +from pathlib import Path + +import yaml + +ROOT = Path(__file__).resolve().parents[1] +MODELS = ( + "HistoryRow", + "HistoryPage", + "DomainProfile", + "IdentificationScoredData", + "IdentificationScoredEvent", + "WebhookPingEvent", + "IpInfo", + "LocalIpInfo", + "TrafficSource", + "Signal", + "DetectionFlags", + "ScoreDetail", +) + + +def render(document: dict) -> str: + def resolve(schema: dict) -> dict: + if "$ref" in schema: + ref = schema["$ref"] + if not ref.startswith("#/"): + raise ValueError("Only bundled local references are supported") + value = document + for part in ref[2:].split("/"): + value = value[part] + return resolve(value) + return schema + + def wire_type(schema: dict) -> str: + schema = resolve(schema) + kind = schema.get("type") + if isinstance(kind, list): + return "Union[" + ", ".join(wire_type({**schema, "type": t}) for t in kind) + "]" + if kind == "object": + return "Mapping[str, object]" + if kind == "array": + return "list[" + wire_type(schema["items"]) + "]" + primitives = { + "string": "str", + "integer": "int", + "number": "float", + "boolean": "bool", + "null": "None", + } + if kind not in primitives: + raise ValueError(f"Unsupported wire schema: {schema}") + return primitives[kind] + + lines = [ + '"""Generated by scripts/generate_wire.py; do not edit manually.', + "", + "Types describe the wire contract, not runtime validation. Field.read preserves", + "missing, malformed and unknown data for the tolerant normalization boundary.", + '"""', + "from __future__ import annotations", + "", + "from collections.abc import Mapping", + "from typing import Generic, Literal, Optional, TypedDict, TypeVar, Union, cast", + "", + 'T = TypeVar("T", covariant=True)', + "", + "class Field(Generic[T]):", + " def __init__(self, name: str) -> None:", + " self.name = name", + "", + " def read(self, value: Mapping[str, object]) -> Optional[T]:", + " return cast(Optional[T], value.get(self.name))", + "", + ] + operations = {} + routes = {} + methods = {} + for surface in (document["paths"], document.get("webhooks", {})): + for route, path in surface.items(): + for method, operation in path.items(): + if isinstance(operation, dict) and "operationId" in operation: + operation_id = operation["operationId"] + operations[operation_id] = operation + routes[operation_id] = route + methods[operation_id] = method + # Operation parameters override inherited parameters with the same key. + merged = {} + for parameter in [ + *path.get("parameters", []), + *operation.get("parameters", []), + ]: + parameter = resolve(parameter) + merged[(parameter["in"], parameter["name"])] = parameter + operations[operation_id] = {**operation, "parameters": list(merged.values())} + supported = { + "searchHistory": { + ("path", "search_type"), + ("path", "value"), + ("query", "limit"), + ("query", "offset"), + }, + "getDomainProfile": {("header", "X-Shield-Domain")}, + } + for operation_id, allowed in supported.items(): + if methods[operation_id] != "get": + raise ValueError(f"Unsupported HTTP method on {operation_id}: {methods[operation_id]}") + for parameter in operations[operation_id]["parameters"]: + key = (parameter["in"], parameter["name"]) + if parameter.get("required", False) and key not in allowed: + raise ValueError(f"Unsupported required parameter on {operation_id}: {key}") + lookup_parameter = next( + parameter + for parameter in operations["searchHistory"]["parameters"] + if (parameter["in"], parameter["name"]) == ("path", "search_type") + ) + # The schema may expose a new lookup before its validation rules exist in this client. + validation = ast.parse((ROOT / "src/shieldlabs/_validation.py").read_text()) + lookup_types = next( + ast.literal_eval(statement.value) + for statement in validation.body + if isinstance(statement, ast.AnnAssign) + and isinstance(statement.target, ast.Name) + and statement.target.id == "LOOKUP_TYPES" + ) + if set(resolve(lookup_parameter["schema"])["enum"]) != set(lookup_types): + raise ValueError("Unsupported lookup types: update client validation before generation") + history_route = routes["searchHistory"] + if set(re.findall(r"\{([^{}]+)\}", history_route)) != {"search_type", "value"}: + raise ValueError("Unsupported History path placeholders") + lines.extend( + [ + f"PROFILE_PATH = {routes['getDomainProfile']!r}", + f"HISTORY_PATH = {history_route!r}", + "", + ] + ) + schemas = { + name: resolve(document["components"]["schemas"][name]) + for name in MODELS + if name != "LocalIpInfo" + } + for name, operation_id in ( + ("HistoryPage", "searchHistory"), + ("DomainProfile", "getDomainProfile"), + ): + response = resolve(operations[operation_id]["responses"]["200"]) + schemas[name] = resolve(response["content"]["application/json"]["schema"]) + for name, operation_id in ( + ("IdentificationScoredEvent", "identificationScored"), + ("WebhookPingEvent", "webhookPing"), + ): + body = resolve(operations[operation_id]["requestBody"]) + schemas[name] = resolve(body["content"]["application/json"]["schema"]) + for name, expected in ( + ("IdentificationScoredEvent", "identification.scored"), + ("WebhookPingEvent", "webhook.ping"), + ): + discriminator = schemas[name]["properties"].get("event_type") + if discriminator is None: + raise ValueError(f"Unsupported webhook discriminator on {name}: missing event_type") + discriminator = resolve(discriminator) + if wire_type(discriminator) != "str" or discriminator.get("const") != expected: + raise ValueError(f"Unsupported webhook discriminator on {name}") + schemas["HistoryRow"] = resolve(schemas["HistoryPage"]["properties"]["data"]["items"]) + schemas["IdentificationScoredData"] = resolve( + schemas["IdentificationScoredEvent"]["properties"]["data"] + ) + data = schemas["IdentificationScoredData"]["properties"] + for name, key in ( + ("IpInfo", "public_ip"), + ("LocalIpInfo", "local_ip"), + ("TrafficSource", "traffic_source"), + ("DetectionFlags", "detection_flags"), + ): + schemas[name] = resolve(data[key]) + schemas["Signal"] = resolve(data["signals"]["items"]) + for name in MODELS: + schema = schemas[name] + lines.append(f"class {name}:") + for key, prop in schema["properties"].items(): + if not key.isidentifier() or keyword.iskeyword(key): + raise ValueError(f"Unsupported Python field name: {key}") + lines.append(f" {key}: Field[{wire_type(prop)}] = Field({key!r})") + lines.append("") + + for operation_id, location, name in ( + ("searchHistory", "query", "HistoryQuery"), + ("searchHistory", "path", "HistoryPath"), + ("getDomainProfile", "header", "ProfileHeaders"), + ): + parameters = [resolve(p) for p in operations[operation_id]["parameters"]] + required_fields = {} + optional_fields = {} + for p in parameters: + if p["in"] != location: + continue + schema = resolve(p["schema"]) + annotation = wire_type(schema) + if "enum" in schema: + annotation = "Literal[" + ", ".join(repr(v) for v in schema["enum"]) + "]" + if name == "HistoryPath" and p["name"] == "search_type": + lines.append(f"LookupType = {annotation}") + fields = required_fields if p.get("required", False) else optional_fields + fields[p["name"]] = annotation + bases = [] + for suffix, fields, required in ( + ("Required", required_fields, True), + ("Optional", optional_fields, False), + ): + base = f"_{name}{suffix}" + bases.append(base) + if all(key.isidentifier() and not keyword.iskeyword(key) for key in fields): + lines.append(f"class {base}(TypedDict, total={required}):") + lines.extend(f" {key}: {annotation}" for key, annotation in fields.items()) + if not fields: + lines.append(" pass") + else: + entries = ", ".join(f"{key!r}: {annotation}" for key, annotation in fields.items()) + lines.append(f"{base} = TypedDict({base!r}, {{{entries}}}, total={required})") + lines.append("") + lines.extend([f"class {name}({', '.join(bases)}):", " pass"]) + lines.append("") + return "\n".join(lines) + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--schema", type=Path, default=ROOT / "resources/shieldlabs-api.yaml") + parser.add_argument("--output", type=Path, default=ROOT / "src/shieldlabs/_generated_wire.py") + parser.add_argument("--check", action="store_true") + args = parser.parse_args() + source = render(yaml.safe_load(args.schema.read_text())) + # Pin formatting to the same tool used by the repository, not host line endings. + formatted = subprocess.run( + [ + sys.executable, + "-m", + "ruff", + "format", + "--config", + str(ROOT / "pyproject.toml"), + "--stdin-filename", + "src/shieldlabs/_generated_wire.py", + "-", + ], + input=source, + text=True, + capture_output=True, + check=True, + ).stdout + formatted = subprocess.run( + [ + sys.executable, + "-m", + "ruff", + "check", + "--fix", + "--config", + str(ROOT / "pyproject.toml"), + "--stdin-filename", + "src/shieldlabs/_generated_wire.py", + "-", + ], + input=formatted, + text=True, + capture_output=True, + check=True, + ).stdout + if args.check: + if not args.output.exists() or args.output.read_text() != formatted: + print("Wire types are stale; run python scripts/generate_wire.py", file=sys.stderr) + return 1 + else: + args.output.write_text(formatted) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/scripts/smoke_wheel.py b/scripts/smoke_wheel.py new file mode 100644 index 0000000..2cd1e51 --- /dev/null +++ b/scripts/smoke_wheel.py @@ -0,0 +1,98 @@ +"""Exercise an installed wheel through mocked HTTP and signed synthetic webhooks.""" + +from __future__ import annotations + +import asyncio +import hashlib +import hmac +import json +from pathlib import Path + +import httpx + +import shieldlabs +from shieldlabs import ( + AsyncShieldLabs, + IdentificationScoredEvent, + ShieldLabs, + ShieldLabsManagement, + webhooks, +) + + +def main() -> None: + if "site-packages" not in Path(shieldlabs.__file__).parts: + raise RuntimeError("Run this smoke test with the installed wheel, not an editable checkout") + + def handle(request: httpx.Request) -> httpx.Response: + if request.url.path == "/v1/profile": + assert request.headers["X-Shield-Domain"] == "example.com" + return httpx.Response(200, json={"Domain": "example.com", "Weight": -9, "future": 1}) + assert request.url.path == "/api/v1/history/user_hid/anonymous" + assert request.url.params == httpx.QueryParams({"limit": 1, "offset": 0}) + return httpx.Response( + 200, + json={ + "total": 2, + "data": [ + { + "request_id": "02f1d973-84db-4156-a7f7-e799e6bf389b", + "score": 999, + "connection_type": "future_network", + "is_vpn": True, + "unknown": [1, 2], + } + ], + }, + ) + + with httpx.Client(transport=httpx.MockTransport(handle)) as http: + with ShieldLabs(api_key="sec_aaaaaaaa-bbbbbbbb-cccccccc", http_client=http) as client: + page = client.history.search("user_hid", "anonymous", limit=1) + assert page.total == 2 + assert page.data[0].risk_score == 999 + assert page.data[0].is_rate_limited + assert page.data[0].detection_flags.vpn + assert page.data[0].connection_type == "future_network" + assert page.data[0].raw["unknown"] == [1, 2] + with ShieldLabsManagement( + secret_key="fixture-secret", domain="example.com", http_client=http + ) as management: + profile = management.get_profile() + assert profile.remaining_identifications == -9 + assert profile.raw["future"] == 1 + + async def check_async() -> None: + http = httpx.AsyncClient(transport=httpx.MockTransport(handle)) + client = AsyncShieldLabs(api_key="sec_aaaaaaaa-bbbbbbbb-cccccccc", http_client=http) + async with http, client: + page = await client.history.search("user_hid", "anonymous", limit=1) + assert page.total == 2 + assert page.data[0].risk_score == 999 + + asyncio.run(check_async()) + payload = json.dumps( + { + "event_type": "identification.scored", + "schema_version": "2026-06-01", + "data": { + "risk_score": 999, + "detection_flags": {"vpn": True}, + "future": True, + "signals": [{"name": "future", "weight": -30}, {"name": "future", "weight": -30}], + }, + } + ).encode() + secret = "fixture-only" + signature = "sha256=" + hmac.new(secret.encode(), payload, hashlib.sha256).hexdigest() + event = webhooks.construct_event(payload, signature, secret) + assert isinstance(event, IdentificationScoredEvent) + assert event.data.risk_score == 999 + assert event.data.detection_flags.vpn + assert [s.weight for s in event.data.signals] == [-30, -30] + assert event.data.raw["future"] is True + print("Installed wheel: sync/async History, profile, signed webhook passed") + + +if __name__ == "__main__": + main() diff --git a/src/shieldlabs/_client.py b/src/shieldlabs/_client.py index fded10a..11b323d 100644 --- a/src/shieldlabs/_client.py +++ b/src/shieldlabs/_client.py @@ -5,7 +5,7 @@ import itertools from collections.abc import AsyncIterator, Iterator from types import TracebackType -from typing import Callable, Optional, Union +from typing import Callable, Optional, Union, cast from uuid import UUID import httpx @@ -17,6 +17,7 @@ ServerError, ShieldLabsError, ) +from ._generated_wire import HISTORY_PATH, HistoryPath, HistoryQuery from ._http import ( RATE_LIMIT_MIN_DELAY, RETRY_AFTER_CAP, @@ -61,7 +62,7 @@ # 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" +_HISTORY_PATH = HISTORY_PATH def poll_waits(initial: float) -> Iterator[float]: @@ -154,7 +155,9 @@ def _configure( 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}" + path = HistoryPath(search_type=cast(LookupType, checked_type), value=segment) + route = _HISTORY_PATH.format(search_type=path["search_type"], value=path["value"]) + return f"{self.base_url}{route}" def __repr__(self) -> str: return f"{type(self).__name__}(base_url={self.base_url!r})" @@ -209,7 +212,7 @@ def _search( response = self._transport.get( url, headers=self._headers, - params={"limit": validate_limit(limit), "offset": validate_offset(offset)}, + params=dict(HistoryQuery(limit=validate_limit(limit), offset=validate_offset(offset))), timeout=timeout, max_retries=max_retries, ) @@ -420,7 +423,7 @@ async def _search( response = await self._transport.get( url, headers=self._headers, - params={"limit": validate_limit(limit), "offset": validate_offset(offset)}, + params=dict(HistoryQuery(limit=validate_limit(limit), offset=validate_offset(offset))), timeout=timeout, max_retries=max_retries, ) diff --git a/src/shieldlabs/_generated_wire.py b/src/shieldlabs/_generated_wire.py new file mode 100644 index 0000000..aa003c0 --- /dev/null +++ b/src/shieldlabs/_generated_wire.py @@ -0,0 +1,226 @@ +"""Generated by scripts/generate_wire.py; do not edit manually. + +Types describe the wire contract, not runtime validation. Field.read preserves +missing, malformed and unknown data for the tolerant normalization boundary. +""" + +from __future__ import annotations + +from collections.abc import Mapping +from typing import Generic, Literal, Optional, TypedDict, TypeVar, Union, cast + +T = TypeVar("T", covariant=True) + + +class Field(Generic[T]): + def __init__(self, name: str) -> None: + self.name = name + + def read(self, value: Mapping[str, object]) -> Optional[T]: + return cast(Optional[T], value.get(self.name)) + + +PROFILE_PATH = "/v1/profile" +HISTORY_PATH = "/api/v1/history/{search_type}/{value}" + + +class HistoryRow: + request_id: Field[str] = Field("request_id") + session_id: Field[str] = Field("session_id") + cookie_id: Field[str] = Field("cookie_id") + domain: Field[str] = Field("domain") + site_domain: Field[str] = Field("site_domain") + user_hid: Field[str] = Field("user_hid") + device_id: Field[str] = Field("device_id") + visitor_id: Field[str] = Field("visitor_id") + ip: Field[str] = Field("ip") + os: Field[str] = Field("os") + browser: Field[str] = Field("browser") + device_type: Field[str] = Field("device_type") + country: Field[str] = Field("country") + connection_type: Field[str] = Field("connection_type") + score: Field[int] = Field("score") + score_details: Field[str] = Field("score_details") + created_at: Field[str] = Field("created_at") + ver: Field[int] = Field("ver") + web_rtc_ip: Field[str] = Field("web_rtc_ip") + web_rtc_country: Field[str] = Field("web_rtc_country") + web_rtc_connection_type: Field[str] = Field("web_rtc_connection_type") + webrtc_leak_ip: Field[str] = Field("webrtc_leak_ip") + webrtc_leak_country: Field[str] = Field("webrtc_leak_country") + webrtc_leak_connection_type: Field[str] = Field("webrtc_leak_connection_type") + webrtc_leak_source: Field[str] = Field("webrtc_leak_source") + is_vpn: Field[bool] = Field("is_vpn") + is_tor: Field[bool] = Field("is_tor") + is_proxy: Field[bool] = Field("is_proxy") + is_datacenter: Field[bool] = Field("is_datacenter") + is_abuser: Field[bool] = Field("is_abuser") + is_privacy_relay: Field[bool] = Field("is_privacy_relay") + is_stun_not_checked: Field[bool] = Field("is_stun_not_checked") + check_incomplete: Field[bool] = Field("check_incomplete") + is_antidetect: Field[bool] = Field("is_antidetect") + is_os_mismatch: Field[bool] = Field("is_os_mismatch") + is_os_not_detected: Field[bool] = Field("is_os_not_detected") + is_timezone_mismatch: Field[bool] = Field("is_timezone_mismatch") + is_js_disabled: Field[bool] = Field("is_js_disabled") + is_browser_automation: Field[bool] = Field("is_browser_automation") + is_incognito: Field[bool] = Field("is_incognito") + is_search_bot: Field[bool] = Field("is_search_bot") + is_suspicious_paid_click: Field[bool] = Field("is_suspicious_paid_click") + entry_url: Field[str] = Field("entry_url") + utm_source: Field[str] = Field("utm_source") + utm_medium: Field[str] = Field("utm_medium") + utm_campaign: Field[str] = Field("utm_campaign") + utm_content: Field[str] = Field("utm_content") + utm_term: Field[str] = Field("utm_term") + traffic_channel: Field[str] = Field("traffic_channel") + traffic_channel_group: Field[str] = Field("traffic_channel_group") + traffic_reason: Field[str] = Field("traffic_reason") + referrer_domain: Field[str] = Field("referrer_domain") + click_id_type: Field[str] = Field("click_id_type") + + +class HistoryPage: + data: Field[list[Mapping[str, object]]] = Field("data") + total: Field[int] = Field("total") + + +class DomainProfile: + Domain: Field[str] = Field("Domain") + Weight: Field[int] = Field("Weight") + Callback: Field[str] = Field("Callback") + PublicKey: Field[str] = Field("PublicKey") + Secret: Field[str] = Field("Secret") + CreatedAt: Field[str] = Field("CreatedAt") + + +class IdentificationScoredData: + request_id: Field[str] = Field("request_id") + visitor_id: Field[str] = Field("visitor_id") + device_id: Field[str] = Field("device_id") + session_id: Field[str] = Field("session_id") + cookie_id: Field[str] = Field("cookie_id") + user_hid: Field[Union[str, None]] = Field("user_hid") + domain: Field[str] = Field("domain") + public_ip: Field[Mapping[str, object]] = Field("public_ip") + local_ip: Field[Mapping[str, object]] = Field("local_ip") + connection_type: Field[str] = Field("connection_type") + os: Field[str] = Field("os") + browser: Field[str] = Field("browser") + device_type: Field[str] = Field("device_type") + traffic_source: Field[Mapping[str, object]] = Field("traffic_source") + risk_score: Field[int] = Field("risk_score") + signals: Field[list[Mapping[str, object]]] = Field("signals") + detection_flags: Field[Mapping[str, object]] = Field("detection_flags") + observed_at: Field[str] = Field("observed_at") + + +class IdentificationScoredEvent: + event_type: Field[str] = Field("event_type") + schema_version: Field[str] = Field("schema_version") + created_at: Field[str] = Field("created_at") + data: Field[Mapping[str, object]] = Field("data") + + +class WebhookPingEvent: + event_type: Field[str] = Field("event_type") + schema_version: Field[str] = Field("schema_version") + created_at: Field[str] = Field("created_at") + + +class IpInfo: + ip: Field[str] = Field("ip") + country: Field[str] = Field("country") + + +class LocalIpInfo: + ip: Field[str] = Field("ip") + country: Field[str] = Field("country") + + +class TrafficSource: + channel: Field[str] = Field("channel") + referrer_domain: Field[str] = Field("referrer_domain") + landing_url: Field[str] = Field("landing_url") + click_id_type: Field[str] = Field("click_id_type") + utm_source: Field[str] = Field("utm_source") + utm_medium: Field[str] = Field("utm_medium") + utm_campaign: Field[str] = Field("utm_campaign") + utm_content: Field[str] = Field("utm_content") + utm_term: Field[str] = Field("utm_term") + + +class Signal: + name: Field[str] = Field("name") + weight: Field[int] = Field("weight") + + +class DetectionFlags: + vpn: Field[bool] = Field("vpn") + privacy_relay: Field[bool] = Field("privacy_relay") + browser_vpn_proxy: Field[bool] = Field("browser_vpn_proxy") + tor: Field[bool] = Field("tor") + proxy: Field[bool] = Field("proxy") + datacenter_ip: Field[bool] = Field("datacenter_ip") + abuser: Field[bool] = Field("abuser") + os_mismatch: Field[bool] = Field("os_mismatch") + os_not_detected: Field[bool] = Field("os_not_detected") + timezone_mismatch: Field[bool] = Field("timezone_mismatch") + anti_detect_browser: Field[bool] = Field("anti_detect_browser") + browser_automation: Field[bool] = Field("browser_automation") + ip_mismatch: Field[bool] = Field("ip_mismatch") + incognito: Field[bool] = Field("incognito") + search_bot: Field[bool] = Field("search_bot") + suspicious_paid_click: Field[bool] = Field("suspicious_paid_click") + javascript_disabled: Field[bool] = Field("javascript_disabled") + stun_not_checked: Field[bool] = Field("stun_not_checked") + check_incomplete: Field[bool] = Field("check_incomplete") + + +class ScoreDetail: + Value: Field[int] = Field("Value") + Description: Field[str] = Field("Description") + + +class _HistoryQueryRequired(TypedDict, total=True): + pass + + +class _HistoryQueryOptional(TypedDict, total=False): + limit: int + offset: int + + +class HistoryQuery(_HistoryQueryRequired, _HistoryQueryOptional): + pass + + +LookupType = Literal[ + "request_id", "device_id", "user_hid", "visitor_id", "ip", "session_id", "cookie_id" +] + + +class _HistoryPathRequired(TypedDict, total=True): + search_type: Literal[ + "request_id", "device_id", "user_hid", "visitor_id", "ip", "session_id", "cookie_id" + ] + value: str + + +class _HistoryPathOptional(TypedDict, total=False): + pass + + +class HistoryPath(_HistoryPathRequired, _HistoryPathOptional): + pass + + +_ProfileHeadersRequired = TypedDict("_ProfileHeadersRequired", {"X-Shield-Domain": str}, total=True) + + +class _ProfileHeadersOptional(TypedDict, total=False): + pass + + +class ProfileHeaders(_ProfileHeadersRequired, _ProfileHeadersOptional): + pass diff --git a/src/shieldlabs/_management.py b/src/shieldlabs/_management.py index 37cb4f6..05cc3a7 100644 --- a/src/shieldlabs/_management.py +++ b/src/shieldlabs/_management.py @@ -7,6 +7,7 @@ import httpx +from ._generated_wire import PROFILE_PATH, ProfileHeaders from ._http import USER_AGENT, AsyncTransport, SyncTransport, json_object from ._models import DomainProfile from ._validation import ( @@ -19,7 +20,7 @@ __all__ = ["AsyncShieldLabsManagement", "ShieldLabsManagement"] -_PROFILE_PATH = "/v1/profile" +_PROFILE_PATH = PROFILE_PATH class _ManagementConfig: @@ -41,8 +42,9 @@ def _configure( 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, + profile_headers: ProfileHeaders = {"X-Shield-Domain": self.domain} + self._headers: dict[str, str] = { + "X-Shield-Domain": profile_headers["X-Shield-Domain"], "Authorization": f"Bearer {secret}", "Accept": "application/json", "User-Agent": USER_AGENT, diff --git a/src/shieldlabs/_models.py b/src/shieldlabs/_models.py index 74172d2..897834b 100644 --- a/src/shieldlabs/_models.py +++ b/src/shieldlabs/_models.py @@ -8,9 +8,10 @@ from datetime import datetime from typing import Any, Literal, Optional +from . import _generated_wire as wire +from . import _wire from ._normalize import ( FLAG_KEYS, - HISTORY_FLAG_MAP, NIL_UUID, TRAFFIC_KEYS, RiskBand, @@ -42,6 +43,60 @@ _IP_MISMATCH_DETAIL = "IP ≠ leakIP" +_HISTORY_FLAGS: dict[str, wire.Field[bool]] = { + "vpn": wire.HistoryRow.is_vpn, + "privacy_relay": wire.HistoryRow.is_privacy_relay, + "tor": wire.HistoryRow.is_tor, + "proxy": wire.HistoryRow.is_proxy, + "datacenter_ip": wire.HistoryRow.is_datacenter, + "abuser": wire.HistoryRow.is_abuser, + "os_mismatch": wire.HistoryRow.is_os_mismatch, + "os_not_detected": wire.HistoryRow.is_os_not_detected, + "timezone_mismatch": wire.HistoryRow.is_timezone_mismatch, + "anti_detect_browser": wire.HistoryRow.is_antidetect, + "browser_automation": wire.HistoryRow.is_browser_automation, + "incognito": wire.HistoryRow.is_incognito, + "search_bot": wire.HistoryRow.is_search_bot, + "suspicious_paid_click": wire.HistoryRow.is_suspicious_paid_click, + "javascript_disabled": wire.HistoryRow.is_js_disabled, + "stun_not_checked": wire.HistoryRow.is_stun_not_checked, + "check_incomplete": wire.HistoryRow.check_incomplete, +} + +_FLAG_FIELDS: dict[str, wire.Field[bool]] = { + "vpn": wire.DetectionFlags.vpn, + "privacy_relay": wire.DetectionFlags.privacy_relay, + "browser_vpn_proxy": wire.DetectionFlags.browser_vpn_proxy, + "tor": wire.DetectionFlags.tor, + "proxy": wire.DetectionFlags.proxy, + "datacenter_ip": wire.DetectionFlags.datacenter_ip, + "abuser": wire.DetectionFlags.abuser, + "os_mismatch": wire.DetectionFlags.os_mismatch, + "os_not_detected": wire.DetectionFlags.os_not_detected, + "timezone_mismatch": wire.DetectionFlags.timezone_mismatch, + "anti_detect_browser": wire.DetectionFlags.anti_detect_browser, + "browser_automation": wire.DetectionFlags.browser_automation, + "ip_mismatch": wire.DetectionFlags.ip_mismatch, + "incognito": wire.DetectionFlags.incognito, + "search_bot": wire.DetectionFlags.search_bot, + "suspicious_paid_click": wire.DetectionFlags.suspicious_paid_click, + "javascript_disabled": wire.DetectionFlags.javascript_disabled, + "stun_not_checked": wire.DetectionFlags.stun_not_checked, + "check_incomplete": wire.DetectionFlags.check_incomplete, +} + +_TRAFFIC_FIELDS: dict[str, wire.Field[str]] = { + "channel": wire.TrafficSource.channel, + "referrer_domain": wire.TrafficSource.referrer_domain, + "landing_url": wire.TrafficSource.landing_url, + "click_id_type": wire.TrafficSource.click_id_type, + "utm_source": wire.TrafficSource.utm_source, + "utm_medium": wire.TrafficSource.utm_medium, + "utm_campaign": wire.TrafficSource.utm_campaign, + "utm_content": wire.TrafficSource.utm_content, + "utm_term": wire.TrafficSource.utm_term, +} + def _mapping(value: object) -> Mapping[str, Any]: return value if isinstance(value, Mapping) else {} @@ -91,12 +146,23 @@ class IpInfo: 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"))) + return cls( + ip=clean_ip(_wire.text(wire.IpInfo.ip.read(data))), + country=_wire.text(wire.IpInfo.country.read(data)), + ) def to_dict(self) -> dict[str, Any]: return {"ip": self.ip, "country": self.country} +def _local_ip_info(value: Optional[Mapping[str, object]]) -> IpInfo: + data = _wire.mapping(value) + return IpInfo( + ip=clean_ip(_wire.text(wire.LocalIpInfo.ip.read(data))), + country=_wire.text(wire.LocalIpInfo.country.read(data)), + ) + + @dataclass(frozen=True) class TrafficSource: """Attribution of the visit. Every field is ``""`` when absent.""" @@ -115,7 +181,7 @@ class TrafficSource: 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}) + return cls(**{key: _wire.text(field.read(data)) for key, field in _TRAFFIC_FIELDS.items()}) def to_dict(self) -> dict[str, Any]: return {key: getattr(self, key) for key in TRAFFIC_KEYS} @@ -142,8 +208,8 @@ def from_dict(cls, data: Mapping[str, Any]) -> Signal: data = _mapping(data) description = data.get("description") return cls( - name=as_str(data.get("name")), - weight=as_int(data.get("weight")), + name=_wire.text(wire.Signal.name.read(data)), + weight=_wire.integer(wire.Signal.weight.read(data)), description=description if isinstance(description, str) else None, ) @@ -179,7 +245,7 @@ class DetectionFlags: 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}) + return cls(**{key: _wire.boolean(field.read(data)) for key, field in _FLAG_FIELDS.items()}) def to_dict(self) -> dict[str, bool]: return {key: getattr(self, key) for key in FLAG_KEYS} @@ -240,32 +306,41 @@ def has_device_signals(self) -> bool: 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") + flags = _wire.mapping(wire.IdentificationScoredData.detection_flags.read(data)) + raw_signals = wire.IdentificationScoredData.signals.read(data) 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 []) + Signal( + name=_wire.text(wire.Signal.name.read(item)), + weight=_wire.integer(wire.Signal.weight.read(item)), + ) + for item in _wire.records(raw_signals) 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")), + request_id=_wire.text(wire.IdentificationScoredData.request_id.read(data)), + visitor_id=_wire.text(wire.IdentificationScoredData.visitor_id.read(data)), + device_id=_wire.text(wire.IdentificationScoredData.device_id.read(data)), + session_id=_wire.text(wire.IdentificationScoredData.session_id.read(data)), + cookie_id=_wire.text(wire.IdentificationScoredData.cookie_id.read(data)), + user_hid=_optional_user_hid(wire.IdentificationScoredData.user_hid.read(data)), + domain=_wire.text(wire.IdentificationScoredData.domain.read(data)), + public_ip=IpInfo.from_dict( + _wire.mapping(wire.IdentificationScoredData.public_ip.read(data)) + ), + local_ip=_local_ip_info(wire.IdentificationScoredData.local_ip.read(data)), + connection_type=_wire.text(wire.IdentificationScoredData.connection_type.read(data)), + os=_wire.text(wire.IdentificationScoredData.os.read(data)), + browser=_wire.text(wire.IdentificationScoredData.browser.read(data)), + device_type=_wire.text(wire.IdentificationScoredData.device_type.read(data)), + traffic_source=TrafficSource.from_dict( + _wire.mapping(wire.IdentificationScoredData.traffic_source.read(data)) + ), + risk_score=_wire.integer(wire.IdentificationScoredData.risk_score.read(data)), signals=signals, detection_flags=DetectionFlags.from_dict(flags), - observed_at=parse_rfc3339(data.get("observed_at")), + observed_at=parse_rfc3339( + _wire.text(wire.IdentificationScoredData.observed_at.read(data)) + ), source="webhook", raw=dict(data), ) @@ -274,17 +349,17 @@ def from_webhook_data(cls, data: Mapping[str, Any]) -> Identification: 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() + leak_source = _wire.text(wire.HistoryRow.webrtc_leak_source.read(row)).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")) + local_ip = clean_ip(_wire.text(wire.HistoryRow.webrtc_leak_ip.read(row))) + local_country = _wire.text(wire.HistoryRow.webrtc_leak_country.read(row)) 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")) + local_ip = clean_ip(_wire.text(wire.HistoryRow.web_rtc_ip.read(row))) + local_country = _wire.text(wire.HistoryRow.web_rtc_country.read(row)) + public_ip = clean_ip(_wire.text(wire.HistoryRow.ip.read(row))) details: Any = [] - score_details = row.get("score_details") + score_details = wire.HistoryRow.score_details.read(row) if isinstance(score_details, str) and score_details: try: details = json.loads(score_details) @@ -298,56 +373,57 @@ def from_history_row(cls, row: Mapping[str, Any]) -> Identification: for detail in details: if not isinstance(detail, Mapping): continue - description = as_str(detail.get("Description")) + description = _wire.text(wire.ScoreDetail.Description.read(detail)) if description.startswith(_IP_MISMATCH_DETAIL): ip_leak_detail = True - value = detail.get("Value", 0) + value = wire.ScoreDetail.Value.read(detail) 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)) + search_bot = _wire.boolean(wire.HistoryRow.is_search_bot.read(row)) flags: dict[str, bool] = {} for key in FLAG_KEYS: if key == "browser_vpn_proxy": - flags[key] = row.get("connection_type") == "browser_vpn_proxy" + flags[key] = wire.HistoryRow.connection_type.read(row) == "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)) + flags[key] = _wire.boolean(_HISTORY_FLAGS[key].read(row)) 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"))), + request_id=_wire.text(wire.HistoryRow.request_id.read(row)), + visitor_id=_wire.text(wire.HistoryRow.visitor_id.read(row)), + device_id=_wire.text(wire.HistoryRow.device_id.read(row)), + session_id=_wire.text(wire.HistoryRow.session_id.read(row)), + cookie_id=_wire.text(wire.HistoryRow.cookie_id.read(row)), + user_hid=_optional_user_hid(wire.HistoryRow.user_hid.read(row)), + domain=_wire.text(wire.HistoryRow.site_domain.read(row)) + or _wire.text(wire.HistoryRow.domain.read(row)), + public_ip=IpInfo(ip=public_ip, country=_wire.text(wire.HistoryRow.country.read(row))), 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")), + connection_type=_wire.text(wire.HistoryRow.connection_type.read(row)), + os=_wire.text(wire.HistoryRow.os.read(row)), + browser=_wire.text(wire.HistoryRow.browser.read(row)), + device_type=_wire.text(wire.HistoryRow.device_type.read(row)), 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")), + channel=_wire.text(wire.HistoryRow.traffic_channel.read(row)), + referrer_domain=_wire.text(wire.HistoryRow.referrer_domain.read(row)), + landing_url=_wire.text(wire.HistoryRow.entry_url.read(row)), + click_id_type=_wire.text(wire.HistoryRow.click_id_type.read(row)), + utm_source=_wire.text(wire.HistoryRow.utm_source.read(row)), + utm_medium=_wire.text(wire.HistoryRow.utm_medium.read(row)), + utm_campaign=_wire.text(wire.HistoryRow.utm_campaign.read(row)), + utm_content=_wire.text(wire.HistoryRow.utm_content.read(row)), + utm_term=_wire.text(wire.HistoryRow.utm_term.read(row)), ), - risk_score=as_int(row.get("score")), + risk_score=_wire.integer(wire.HistoryRow.score.read(row)), signals=tuple(signals), detection_flags=DetectionFlags(**flags), - observed_at=parse_history_time(row.get("created_at")), + observed_at=parse_history_time(_wire.text(wire.HistoryRow.created_at.read(row))), source="history", raw=dict(row), ) @@ -432,13 +508,15 @@ class HistoryPage: 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") + rows = wire.HistoryPage.data.read(body) data = tuple( Identification.from_history_row(row) - for row in (rows if isinstance(rows, list) else []) + for row in _wire.records(rows) if isinstance(row, Mapping) ) - return cls(data=data, total=as_int(body.get("total"), default=len(data))) + return cls( + data=data, total=_wire.integer(wire.HistoryPage.total.read(body), default=len(data)) + ) @dataclass(frozen=True) @@ -467,11 +545,11 @@ 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")), + domain=_wire.text(wire.DomainProfile.Domain.read(body)), + remaining_identifications=_wire.integer(wire.DomainProfile.Weight.read(body)), + public_key_masked=_wire.text(wire.DomainProfile.PublicKey.read(body)), + secret_key_masked=_wire.text(wire.DomainProfile.Secret.read(body)), + created_at=parse_rfc3339(_wire.text(wire.DomainProfile.CreatedAt.read(body))), raw=dict(body), ) diff --git a/src/shieldlabs/_validation.py b/src/shieldlabs/_validation.py index d53a4a3..bca8932 100644 --- a/src/shieldlabs/_validation.py +++ b/src/shieldlabs/_validation.py @@ -7,11 +7,12 @@ import os import re import warnings -from typing import Literal, Optional, Union +from typing import Optional, Union from urllib.parse import quote, urlsplit from uuid import UUID from ._errors import ShieldLabsWarning, ValidationError +from ._generated_wire import LookupType as LookupType __all__ = [ "DEFAULT_HISTORY_BASE_URL", @@ -32,12 +33,7 @@ "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, ...] = ( +LOOKUP_TYPES: tuple[LookupType, ...] = ( "ip", "user_hid", "visitor_id", diff --git a/src/shieldlabs/_version.py b/src/shieldlabs/_version.py index 5becc17..5c4105c 100644 --- a/src/shieldlabs/_version.py +++ b/src/shieldlabs/_version.py @@ -1 +1 @@ -__version__ = "1.0.0" +__version__ = "1.0.1" diff --git a/src/shieldlabs/_wire.py b/src/shieldlabs/_wire.py new file mode 100644 index 0000000..ad856e3 --- /dev/null +++ b/src/shieldlabs/_wire.py @@ -0,0 +1,32 @@ +"""Typed, tolerant boundary between generated HTTP fields and public models. + +Annotations check the declared contract at build time. Runtime checks deliberately keep +the SDK's historical missing/null/malformed defaults, without validating whole responses. +""" + +from __future__ import annotations + +from collections.abc import Mapping +from typing import Any, Optional, Union + +from ._normalize import as_int, as_str + + +def text(value: Optional[str]) -> str: + return as_str(value) + + +def integer(value: Optional[Union[int, float]], default: int = 0) -> int: + return as_int(value, default) + + +def boolean(value: Optional[bool]) -> bool: + return bool(value) + + +def mapping(value: Optional[Mapping[str, object]]) -> Mapping[str, Any]: + return value if isinstance(value, Mapping) else {} + + +def records(value: Optional[list[Mapping[str, object]]]) -> list[Mapping[str, Any]]: + return [item for item in value if isinstance(item, Mapping)] if isinstance(value, list) else [] diff --git a/src/shieldlabs/webhooks.py b/src/shieldlabs/webhooks.py index cb6d2a8..4db901d 100644 --- a/src/shieldlabs/webhooks.py +++ b/src/shieldlabs/webhooks.py @@ -32,8 +32,11 @@ from typing import Any, Literal, Optional, Union from ._errors import ShieldLabsWarning, SignatureVerificationError, WebhookParseError +from ._generated_wire import IdentificationScoredEvent as WireEvent +from ._generated_wire import WebhookPingEvent as WirePing from ._models import Identification from ._normalize import parse_rfc3339 +from ._wire import text __all__ = [ "SCHEMA_VERSION", @@ -199,12 +202,11 @@ def _parse_event(body: bytes) -> WebhookEvent: 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") + event_type = text(WireEvent.event_type.read(envelope)) 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 = "" + envelope_fields = WirePing if event_type == "webhook.ping" else WireEvent + schema_version = text(envelope_fields.schema_version.read(envelope)) if schema_version != SCHEMA_VERSION: warnings.warn( f"Webhook schema_version {schema_version!r} is not {SCHEMA_VERSION!r}; " @@ -212,8 +214,8 @@ def _parse_event(body: bytes) -> WebhookEvent: ShieldLabsWarning, stacklevel=3, ) - created_at = parse_rfc3339(envelope.get("created_at")) - data = envelope.get("data") + created_at = parse_rfc3339(text(envelope_fields.created_at.read(envelope))) + data = WireEvent.data.read(envelope) if event_type == "identification.scored": if not isinstance(data, dict): raise WebhookParseError("identification.scored event has no data object") diff --git a/tests/test_history.py b/tests/test_history.py index b6d965b..38ad2eb 100644 --- a/tests/test_history.py +++ b/tests/test_history.py @@ -64,7 +64,7 @@ def test_search_sends_expected_request(client: ShieldLabs, mock: Any) -> None: 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 ") + assert USER_AGENT.startswith("shieldlabs-python/1.0.1 ") def test_search_defaults_and_empty_page(client: ShieldLabs, mock: Any) -> None: diff --git a/tests/test_package.py b/tests/test_package.py index 114e282..7b1e4a9 100644 --- a/tests/test_package.py +++ b/tests/test_package.py @@ -81,7 +81,7 @@ def test_error_hierarchy() -> None: def test_version_and_typing_marker() -> None: - assert shieldlabs.__version__ == "1.0.0" + assert shieldlabs.__version__ == "1.0.1" assert (Path(shieldlabs.__file__).parent / "py.typed").exists() diff --git a/tests/test_wire_contract.py b/tests/test_wire_contract.py new file mode 100644 index 0000000..510b1b4 --- /dev/null +++ b/tests/test_wire_contract.py @@ -0,0 +1,85 @@ +"""The generated boundary is used without making response parsing strict.""" + +from __future__ import annotations + +import json +import subprocess +import sys +from pathlib import Path +from typing import get_args + +import pytest + +from _support import sign +from shieldlabs import ( + DomainProfile, + HistoryPage, + Identification, + IdentificationScoredEvent, + webhooks, +) +from shieldlabs import _generated_wire as wire +from shieldlabs._validation import LOOKUP_TYPES, LookupType + + +def test_generated_output_is_current() -> None: + root = Path(__file__).resolve().parents[1] + result = subprocess.run( + [sys.executable, "scripts/generate_wire.py", "--check"], + cwd=root, + capture_output=True, + text=True, + timeout=30, + ) + assert result.returncode == 0, result.stdout + result.stderr + + +def test_lookup_validation_handles_every_declared_type() -> None: + assert set(LOOKUP_TYPES) == set(get_args(LookupType)) + + +def test_real_normalizers_read_generated_fields(monkeypatch: pytest.MonkeyPatch) -> None: + # A descriptor substitution changes actual outputs, proving it is not a sidecar. + monkeypatch.setattr(wire.HistoryRow, "score", wire.Field("new_score")) + monkeypatch.setattr(wire.HistoryPage, "total", wire.Field("new_total")) + monkeypatch.setattr(wire.DomainProfile, "Weight", wire.Field("new_weight")) + monkeypatch.setattr(wire.IdentificationScoredData, "risk_score", wire.Field("new_score")) + monkeypatch.setattr(wire.DetectionFlags.vpn, "name", "new_vpn") + monkeypatch.setattr(wire.LocalIpInfo, "ip", wire.Field("new_local_ip")) + page = HistoryPage.from_dict({"new_total": 8, "data": [{"new_score": 999}]}) + assert page.total == 8 + assert page.data[0].risk_score == 999 + assert DomainProfile.from_dict({"new_weight": -5}).remaining_identifications == -5 + payload = json.dumps( + { + "event_type": "identification.scored", + "schema_version": "2026-06-01", + "data": { + "new_score": 999, + "detection_flags": {"new_vpn": True}, + "local_ip": {"new_local_ip": "198.51.100.2", "country": "Germany"}, + }, + } + ).encode() + event = webhooks.construct_event(payload, sign("fixture", payload), "fixture") + assert isinstance(event, IdentificationScoredEvent) + assert event.data.risk_score == 999 + assert event.data.detection_flags.vpn + assert event.data.local_ip.ip == "198.51.100.2" + assert event.data.local_ip.country == "Germany" + + +@pytest.mark.parametrize("bad", [None, "bad", [], {}, True]) +def test_malformed_defaults_and_raw_are_preserved(bad: object) -> None: + row = {"score": bad, "connection_type": "future_network", "future": bad} + identification = Identification.from_history_row(row) + assert identification.risk_score == 0 + assert identification.connection_type == "future_network" + assert identification.raw == row + data = {"risk_score": bad, "signals": bad, "public_ip": bad, "future": bad} + result = Identification.from_webhook_data(data) + assert result.risk_score == 0 + assert not result.signals + assert result.public_ip.ip == "" + assert result.raw == data + assert DomainProfile.from_dict({"Weight": bad}).remaining_identifications == 0