From c15185eeff0ee2a7071ff729a53e3bef0fd9f9a7 Mon Sep 17 00:00:00 2001 From: EgorMajj <91486022+EgorMajj@users.noreply.github.com> Date: Sat, 3 Oct 2026 12:02:00 +0300 Subject: [PATCH] feat: connect supported client to generated OpenAPI wire types --- .github/workflows/ci.yml | 17 +- .github/workflows/release.yml | 10 + .gitignore | 1 + CHANGELOG.md | 4 +- CONTRIBUTING.md | 23 ++- README.md | 14 +- generate.sh | 2 + scripts/check-wire-drift.py | 84 ++++++++ scripts/consumer/Consumer.java | 68 +++++++ scripts/generate-wire.py | 157 +++++++++++++++ scripts/requirements.txt | 1 + scripts/test-wire-generation.py | 74 +++++++ scripts/verify-package.sh | 13 ++ .../java/ai/shieldlabs/DetectionFlag.java | 56 ++++-- .../java/ai/shieldlabs/DomainProfile.java | 11 +- .../java/ai/shieldlabs/HistoryService.java | 10 +- src/main/java/ai/shieldlabs/LookupType.java | 14 +- .../java/ai/shieldlabs/ManagementClient.java | 3 +- src/main/java/ai/shieldlabs/Normalizer.java | 141 +++++++------- src/main/java/ai/shieldlabs/Webhooks.java | 9 +- src/main/java/ai/shieldlabs/WireModels.java | 184 ++++++++++++++++++ src/main/java/ai/shieldlabs/WireValue.java | 41 ++++ .../java/ai/shieldlabs/WireModelsTest.java | 63 ++++++ 23 files changed, 887 insertions(+), 113 deletions(-) create mode 100644 scripts/check-wire-drift.py create mode 100644 scripts/consumer/Consumer.java create mode 100644 scripts/generate-wire.py create mode 100644 scripts/requirements.txt create mode 100644 scripts/test-wire-generation.py create mode 100644 scripts/verify-package.sh create mode 100644 src/main/java/ai/shieldlabs/WireModels.java create mode 100644 src/main/java/ai/shieldlabs/WireValue.java create mode 100644 src/test/java/ai/shieldlabs/WireModelsTest.java diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index e76c9e9..836f08c 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -32,12 +32,25 @@ jobs: java-version: ${{ matrix.java }} cache: maven + - name: Check supported-client wire generation + run: | + python3 -m pip install -r scripts/requirements.txt + python3 scripts/generate-wire.py --check + python3 scripts/test-wire-generation.py + + - name: Reject breaking contract mutations + if: matrix.java == '17' + run: python3 scripts/check-wire-drift.py + - name: Build, test, coverage and javadoc run: mvn -B -ntp install - name: Compile the example against the installed SDK run: mvn -B -ntp -f examples/httpserver/pom.xml verify + - name: Check a consumer of the packaged jar + run: bash scripts/verify-package.sh + generated: name: Generated API types runs-on: ubuntu-latest @@ -46,6 +59,8 @@ jobs: with: persist-credentials: false - name: Rebuild generated files - run: ./generate.sh + run: | + python3 -m pip install -r scripts/requirements.txt + ./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 ef97c8b..1adf403 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -43,9 +43,19 @@ jobs: exit 1 fi + - name: Check supported-client contract + run: | + python3 -m pip install -r scripts/requirements.txt + python3 scripts/generate-wire.py --check + python3 scripts/test-wire-generation.py + python3 scripts/check-wire-drift.py + - name: Build, test, coverage and javadoc run: mvn -B -ntp verify + - name: Check a consumer of the packaged jar + run: bash scripts/verify-package.sh + - uses: actions/upload-artifact@ea165f8d65b6e75b540449e92b4886f43607fa02 # v4.6.2 with: name: jars diff --git a/.gitignore b/.gitignore index 16b7b0d..9638605 100644 --- a/.gitignore +++ b/.gitignore @@ -1,6 +1,7 @@ # Build output target/ *.class +__pycache__/ # IDE and OS files .idea/ diff --git a/CHANGELOG.md b/CHANGELOG.md index b406c2e..6e35037 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -8,7 +8,9 @@ All notable changes to this project are documented in this file. The format foll ### 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 and `generate.sh` rebuilds the generated clients. +- Connect the supported client to OpenAPI-generated wire views, preserving tolerant normalization. +- Check generated freshness, breaking schema mutations and installed-jar consumption in CI. ## [1.0.0] - 2026-09-30 diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 17bf990..5a79c9e 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -24,7 +24,28 @@ mvn install -DskipTests mvn -f examples/httpserver/pom.xml verify ``` -## Guidelines +## OpenAPI contract updates + +Install the pinned generator dependency with `python3 -m pip install -r scripts/requirements.txt`. +After updating `resources/shieldlabs-api.yaml`, run `python3 scripts/generate-wire.py`, then +`python3 scripts/generate-wire.py --check` and `python3 scripts/check-wire-drift.py`. +The latter compiles the supported client against renamed and retyped fields and an optional additive +field. Use `--docker` if Maven is available only in the documented container. Run `mvn verify` and +`bash scripts/verify-package.sh` afterward. + +Run `python3 scripts/test-wire-generation.py` for operation-boundary regression checks. New required +parameters on either consumed operation and a changed Management profile route require an adapter +update. Optional new query/header parameters remain compatible. Ping and scored webhook envelope +fields must retain compatible names and declared kinds because the runtime shares envelope parsing. + +`WireModels.java` contains generated, package-private schema views. `WireValue` keeps the original +JSON value behind a declared-kind wrapper. Normalization requires the matching kind, so type drift +fails compilation without introducing strict deserialization of old or unexpected server values. +Enums remain strings on responses; unknown fields remain in `raw()`. New wire shapes or schema +constructs not supported by the narrow generator fail explicitly and need an adapter change. +The `generated/` directory is a separate reference client, not the supported runtime. + +## Code guidelines - Keep the public API small and the runtime dependencies to Jackson Databind only. Public methods need javadoc. diff --git a/README.md b/README.md index 126dd48..18e67a2 100644 --- a/README.md +++ b/README.md @@ -335,11 +335,21 @@ Retries apply to GET requests only (every SDK call is a GET): exponential backof ## Development -Refresh the generated client when the API description changes. This does not replace the supported library in this repository. +The supported client consumes schema-generated wire views for History, profiles, webhooks and +lookup types. The generator retains raw values so existing tolerant normalization and unknown-field +handling remain unchanged. Renaming a consumed field or changing its declared type requires the +client adapter to be updated: contract mutation checks exercise this in CI. + +The standalone client in `generated/` remains a reference implementation; its transport is not +used by the supported library. Retries, polling and webhook verification remain in this library. ```bash ./sync.sh # download the current OpenAPI description into resources/ -./generate.sh # rebuild generated/ from that file +python3 -m pip install -r scripts/requirements.txt +python3 scripts/generate-wire.py # supported client's wire views +python3 scripts/generate-wire.py --check # fail on stale views +python3 scripts/check-wire-drift.py # real compiler checks (requires Maven) +./generate.sh # also rebuild the standalone reference client (requires Docker) ``` diff --git a/generate.sh b/generate.sh index bf24bb9..baa2b22 100755 --- a/generate.sh +++ b/generate.sh @@ -3,6 +3,8 @@ set -euo pipefail cd "$(dirname "${BASH_SOURCE[0]}")" +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/scripts/check-wire-drift.py b/scripts/check-wire-drift.py new file mode 100644 index 0000000..75204f2 --- /dev/null +++ b/scripts/check-wire-drift.py @@ -0,0 +1,84 @@ +#!/usr/bin/env python3 +"""Compile schema mutations against the real supported client, in disposable directories.""" +import argparse +import copy +import importlib.util +from pathlib import Path +import shutil +import subprocess +import tempfile + +import yaml + +ROOT = Path(__file__).resolve().parents[1] +module = importlib.util.spec_from_file_location('wire_generator', ROOT / 'scripts/generate-wire.py') +generator = importlib.util.module_from_spec(module) +module.loader.exec_module(generator) + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument('--docker', action='store_true') + args = parser.parse_args() + original = yaml.safe_load((ROOT / 'resources/shieldlabs-api.yaml').read_text()) + cases = [('baseline', None, None, None)] + for model, field in [('HistoryRow', 'score'), ('HistoryRow', 'request_id'), + ('DomainProfile', 'Weight'), ('IdentificationScoredData', 'risk_score'), + ('IdentificationScoredEvent', 'event_type'), ('DetectionFlags', 'vpn'), + ('TrafficSource', 'channel'), ('HistoryPage', 'total')]: + cases.append((model + '.' + field + ' rename', model, field, 'rename')) + cases.append((model + '.' + field + ' type', model, field, 'type')) + cases.extend([('query rename', None, None, 'parameter'), + ('query type', None, None, 'parameter-type'), + ('search type enum removed', None, None, 'search-type'), + ('profile header renamed', None, None, 'header'), + ('profile header type', None, None, 'header-type'), + ('optional additive', 'HistoryRow', None, 'add')]) + for label, model, field, change in cases: + doc = copy.deepcopy(original) + if model: + props = doc['components']['schemas'][model]['properties'] + if change == 'rename': + props[field + '_changed'] = props.pop(field) + elif change == 'type': + props[field] = {'type': 'object'} + else: + props['future_optional'] = {'type': 'string'} + elif change == 'parameter': + doc['components']['parameters']['HistoryLimit']['name'] = 'changed_limit' + elif change == 'parameter-type': + doc['components']['parameters']['HistoryLimit']['schema'] = {'type': 'string'} + elif change == 'search-type': + doc['components']['parameters']['HistorySearchType']['schema']['enum'].remove('request_id') + elif change == 'header': + doc['components']['parameters']['ShieldDomain']['name'] = 'X-Changed-Domain' + elif change == 'header-type': + doc['components']['parameters']['ShieldDomain']['schema'] = {'type': 'integer'} + expected_success = change is None or change == 'add' + try: + output = generator.generate(doc) + except (KeyError, ValueError) as exc: + if expected_success: + raise + print('PASS rejected schema: ' + label + ' (' + str(exc) + ')', flush=True) + continue + with tempfile.TemporaryDirectory(prefix='shieldlabs-java-contract-') as tmp: + work = Path(tmp) + shutil.copy2(ROOT / 'pom.xml', work / 'pom.xml') + shutil.copytree(ROOT / 'src/main', work / 'src/main') + (work / generator.OUTPUT).write_text(output) + command = ['mvn', '-B', '-ntp', 'compile', '-DskipTests'] + if args.docker: + command = ['docker', 'run', '--rm', '-v', str(work) + ':/src', + '-v', 'shieldlabs-sdk-m2:/root/.m2', '-w', '/src', + 'maven:3.9-eclipse-temurin-17'] + command + run = subprocess.run(command, cwd=work, capture_output=True, text=True, timeout=180) + if (run.returncode == 0) != expected_success: + raise RuntimeError(label + '\n' + run.stdout + run.stderr) + if not expected_success and '[ERROR] COMPILATION ERROR' not in run.stdout: + raise RuntimeError('Not a compiler rejection: ' + label + '\n' + run.stdout + run.stderr) + print('PASS ' + ('compiled: ' if expected_success else 'compiler rejected: ') + label, flush=True) + + +if __name__ == '__main__': + main() diff --git a/scripts/consumer/Consumer.java b/scripts/consumer/Consumer.java new file mode 100644 index 0000000..02c9f67 --- /dev/null +++ b/scripts/consumer/Consumer.java @@ -0,0 +1,68 @@ +import ai.shieldlabs.DomainProfile; +import ai.shieldlabs.HistoryPage; +import ai.shieldlabs.Identification; +import ai.shieldlabs.IdentificationScoredEvent; +import ai.shieldlabs.LookupType; +import ai.shieldlabs.ManagementClient; +import ai.shieldlabs.ShieldLabsClient; +import ai.shieldlabs.Webhooks; +import com.sun.net.httpserver.HttpServer; +import java.net.InetSocketAddress; +import java.net.URI; +import java.nio.charset.StandardCharsets; +import javax.crypto.Mac; +import javax.crypto.spec.SecretKeySpec; + +/** Compiled and run using only the packaged jar and its published runtime dependencies. */ +public final class Consumer { + private Consumer() {} + + public static void main(String[] args) throws Exception { + HttpServer server = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0); + server.createContext("/", exchange -> { + String body = exchange.getRequestURI().getPath().startsWith("/v1/profile") + ? "{\"Domain\":\"example.com\",\"Weight\":123,\"PublicKey\":\"****abcd\",\"extra\":42}" + : "{\"data\":[{\"request_id\":\"00000000-0000-0000-0000-000000000001\"," + + "\"score\":999,\"connection_type\":\"future\",\"future_field\":42}],\"total\":1}"; + byte[] bytes = body.getBytes(StandardCharsets.UTF_8); + exchange.getResponseHeaders().add("Content-Type", "application/json"); + exchange.sendResponseHeaders(200, bytes.length); + try (java.io.OutputStream output = exchange.getResponseBody()) { + output.write(bytes); + } + }); + server.start(); + try { + URI origin = URI.create("http://127.0.0.1:" + server.getAddress().getPort()); + ShieldLabsClient client = ShieldLabsClient.builder().apiKey("sec_fixture_only").baseUrl(origin).build(); + HistoryPage page = client.history().search(LookupType.USER_HID, "anonymous"); + Identification row = page.getIdentifications().get(0); + require(row.getRiskScore() == 999 && row.getConnectionType().equals("future"), "History"); + require(row.raw().containsKey("future_field"), "unknown History field"); + DomainProfile profile = ManagementClient.builder().secretKey("fixture_only") + .domain("example.com").baseUrl(origin).build().getProfile(); + require(profile.getRemainingIdentifications() == 123 && profile.raw().containsKey("extra"), "profile"); + String body = "{\"event_type\":\"identification.scored\",\"schema_version\":\"2026-06-01\"," + + "\"data\":{\"risk_score\":30,\"connection_type\":\"future\",\"extra\":42}}"; + String secret = "whsec_fixture_only"; + Mac mac = Mac.getInstance("HmacSHA256"); + mac.init(new SecretKeySpec(secret.getBytes(StandardCharsets.UTF_8), "HmacSHA256")); + StringBuilder signature = new StringBuilder("sha256="); + for (byte b : mac.doFinal(body.getBytes(StandardCharsets.UTF_8))) { + signature.append(String.format("%02x", b & 255)); + } + IdentificationScoredEvent event = (IdentificationScoredEvent) Webhooks.constructEvent(body, signature.toString(), secret); + require(event.getIdentification().getRiskScore() == 30 + && event.getIdentification().raw().containsKey("extra"), "webhook"); + System.out.println("Packaged consumer: History, profile, verified webhook and unknown fields passed"); + } finally { + server.stop(0); + } + } + + private static void require(boolean condition, String label) { + if (!condition) { + throw new AssertionError(label); + } + } +} diff --git a/scripts/generate-wire.py b/scripts/generate-wire.py new file mode 100644 index 0000000..bc1db2f --- /dev/null +++ b/scripts/generate-wire.py @@ -0,0 +1,157 @@ +#!/usr/bin/env python3 +"""Generate tolerant, schema-typed views used by the supported client.""" +import argparse +import json +from pathlib import Path +import re +import sys + +import yaml + +ROOT = Path(__file__).resolve().parents[1] +OUTPUT = Path('src/main/java/ai/shieldlabs/WireModels.java') +MODELS = ('HistoryRow', 'HistoryPage', 'DomainProfile', 'ScoreDetail', + 'IdentificationScoredData', 'IdentificationScoredEvent', 'WebhookPingEvent', + 'TrafficSource', 'IpInfo', 'Signal', 'DetectionFlags') + + +def generate(doc): + def resolve(schema): + seen = set() + while '$ref' in schema: + ref = schema['$ref'] + if not ref.startswith('#/'): + raise ValueError('Only local references are supported') + if ref in seen: + raise ValueError('Circular alias: ' + ref) + seen.add(ref) + schema = doc + for key in ref[2:].split('/'): + schema = schema[key.replace('~1', '/').replace('~0', '~')] + return schema + + def kind(schema): + schema = resolve(schema) + typ = schema.get('type') + if isinstance(typ, list): + non_null = set(typ) - {'null'} + if len(non_null) != 1: + raise ValueError('Unsupported heterogeneous type list') + typ = non_null.pop() + if typ is None and ('anyOf' in schema or 'oneOf' in schema): + types = {kind(s) for s in schema.get('anyOf', schema.get('oneOf', [])) + if resolve(s).get('type') != 'null'} + if len(types) != 1: + raise ValueError('Unsupported heterogeneous union') + return types.pop() + return {'string': 'StringValue', 'integer': 'IntegerValue', 'number': 'NumberValue', + 'boolean': 'BooleanValue', 'array': 'ArrayValue', 'object': 'ObjectValue'}[typ] + + def ident(name): + if not re.fullmatch(r'[A-Za-z_][A-Za-z_0-9]*', name): + raise ValueError('Unsupported Java identifier: ' + name) + return name + + def parameters(item, op, supported): + # Operation-level definitions override path-item definitions with the same name/location. + params = [resolve(p) for p in item.get('parameters', []) + op.get('parameters', [])] + effective = {(p['in'], p['name']): p for p in params} + for key, parameter in effective.items(): + if key not in supported and parameter.get('required', False): + raise ValueError('Unsupported required parameter on ' + op['operationId'] + ': ' + str(key)) + return effective + + lines = ['// Generated by scripts/generate-wire.py. Do not edit.', + 'package ai.shieldlabs;', '', 'import java.util.Map;', '', + '/** Schema-typed views. Values remain uncoerced for compatibility with older responses. */', + 'final class WireModels {', ' private WireModels() {}'] + for model in MODELS: + schema = resolve(doc['components']['schemas'][model]) + if schema['type'] != 'object': + raise ValueError(model + ' must remain an object') + lines += [f' static final class {model} {{', ' private final Map raw;', + f' {model}(Map raw) {{ this.raw = raw; }}'] + for name, prop in schema['properties'].items(): + typ = kind(prop) + lines.append(f' WireValue.{typ} {ident(name)}() {{ return new WireValue.{typ}(raw.get({json.dumps(name)})); }}') + lines += [' }'] + # The actual operation, not only the reusable component, must still accept these parameters. + operations = {op['operationId']: (path, item, op) for path, item in doc['paths'].items() + for method, op in item.items() if method == 'get'} + path, item, op = operations['searchHistory'] + params = parameters(item, op, {('path', 'search_type'), ('path', 'value'), + ('query', 'limit'), ('query', 'offset')}) + for location, name, expected in [('path', 'search_type', 'StringValue'), ('path', 'value', 'StringValue'), + ('query', 'limit', 'IntegerValue'), ('query', 'offset', 'IntegerValue')]: + if kind(params[location, name]['schema']) != expected: + raise ValueError('Incompatible searchHistory parameter: ' + name) + response = resolve(op['responses']['200'])['content']['application/json']['schema'] + if response.get('$ref') != '#/components/schemas/HistoryPage': + raise ValueError('searchHistory response must map to HistoryPage') + if '{search_type}' not in path or '{value}' not in path: + raise ValueError('searchHistory path placeholders changed') + lines += [' enum SearchType {', ' ' + ', '.join(ident(v.upper()) for v in resolve(params['path', 'search_type']['schema'])['enum']) + ';', + ' String wire() { return name().toLowerCase(java.util.Locale.ROOT); }', ' }', + f' static final String HISTORY_PATH = {json.dumps(path)};'] + lines += [' static String historyQuery(int limit, long offset) {', + ' return "?' + params['query', 'limit']['name'] + '=" + limit + "&' + params['query', 'offset']['name'] + '=" + offset;', + ' }'] + profile_path, pitem, pop = operations['getDomainProfile'] + if profile_path != '/v1/profile': + raise ValueError('Unsupported profile route: update the Management client adapter') + pparams = list(parameters(pitem, pop, {('header', 'X-Shield-Domain')}).values()) + if not any(p['name'] == 'X-Shield-Domain' and p['in'] == 'header' and kind(p['schema']) == 'StringValue' for p in pparams): + raise ValueError('Management domain header changed') + domain_header = next(p for p in pparams if p['name'] == 'X-Shield-Domain') + lines += [' static String[] profileHeaders(String domain) {', + ' return new String[] {' + json.dumps(domain_header['name']) + ', domain};', ' }'] + if resolve(pop['responses']['200'])['content']['application/json']['schema'].get('$ref') != '#/components/schemas/DomainProfile': + raise ValueError('Profile response changed') + for model, field, target, array in [ + ('HistoryPage', 'data', 'HistoryRow', True), + ('IdentificationScoredData', 'signals', 'Signal', True), + ('IdentificationScoredData', 'public_ip', 'IpInfo', False), + ('IdentificationScoredData', 'local_ip', 'IpInfo', False), + ('IdentificationScoredData', 'traffic_source', 'TrafficSource', False), + ('IdentificationScoredData', 'detection_flags', 'DetectionFlags', False), + ('IdentificationScoredEvent', 'data', 'IdentificationScoredData', False), + ]: + prop = doc['components']['schemas'][model]['properties'][field] + if array: + prop = resolve(prop)['items'] + if prop.get('$ref') != '#/components/schemas/' + target: + raise ValueError('Unsupported nested shape: ' + model + '.' + field) + for name, target in [('identification.scored', 'IdentificationScoredEvent'), ('webhook.ping', 'WebhookPingEvent')]: + operation = doc['webhooks'][name]['post'] + body = resolve(operation['requestBody'])['content']['application/json']['schema'] + if body.get('$ref') != '#/components/schemas/' + target: + raise ValueError('Webhook body changed: ' + name) + # The runtime reads the common envelope through the scored view before dispatching on type. + # Ping must retain the same declared names/kinds; differing event_type constants are expected. + scored = doc['components']['schemas']['IdentificationScoredEvent']['properties'] + ping = doc['components']['schemas']['WebhookPingEvent']['properties'] + for field in ('event_type', 'schema_version', 'created_at'): + if field not in ping or field not in scored or kind(ping[field]) != kind(scored[field]): + raise ValueError('Incompatible shared webhook field: ' + field) + lines += ['}', ''] + return '\n'.join(lines) + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument('--check', action='store_true') + parser.add_argument('--schema', type=Path, default=ROOT / 'resources/shieldlabs-api.yaml') + parser.add_argument('--output-root', type=Path, default=ROOT) + args = parser.parse_args() + rendered = generate(yaml.safe_load(args.schema.read_text())) + output = args.output_root / OUTPUT + if args.check: + if not output.exists() or output.read_text() != rendered: + sys.exit('Wire models are stale. Run python3 scripts/generate-wire.py') + else: + output.parent.mkdir(parents=True, exist_ok=True) + output.write_text(rendered) + + +if __name__ == '__main__': + main() diff --git a/scripts/requirements.txt b/scripts/requirements.txt new file mode 100644 index 0000000..f62ce0c --- /dev/null +++ b/scripts/requirements.txt @@ -0,0 +1 @@ +PyYAML==6.0.3 diff --git a/scripts/test-wire-generation.py b/scripts/test-wire-generation.py new file mode 100644 index 0000000..fd9d1ea --- /dev/null +++ b/scripts/test-wire-generation.py @@ -0,0 +1,74 @@ +#!/usr/bin/env python3 +"""Regression tests for operation and shared webhook contract boundaries.""" +import copy +import importlib.util +from pathlib import Path +import unittest + +import yaml + +ROOT = Path(__file__).resolve().parents[1] +module = importlib.util.spec_from_file_location('wire_generator', ROOT / 'scripts/generate-wire.py') +generator = importlib.util.module_from_spec(module) +module.loader.exec_module(generator) + + +class GenerationBoundaryTests(unittest.TestCase): + def setUp(self): + self.original = yaml.safe_load((ROOT / 'resources/shieldlabs-api.yaml').read_text()) + + def operation(self, doc, name): + return next((item, op) for item in doc['paths'].values() + for method, op in item.items() + if method == 'get' and op['operationId'] == name) + + def test_unknown_required_parameters_are_rejected_on_both_operations(self): + for operation in ['searchHistory', 'getDomainProfile']: + for location in ['query', 'header', 'path']: + for level in ['path-item', 'operation']: + with self.subTest(operation=operation, location=location, level=level): + doc = copy.deepcopy(self.original) + item, op = self.operation(doc, operation) + target = item if level == 'path-item' else op + target.setdefault('parameters', []).append({ + 'name': 'tenant', 'in': location, 'required': True, + 'schema': {'type': 'string'}}) + with self.assertRaisesRegex(ValueError, 'Unsupported required parameter'): + generator.generate(doc) + + def test_optional_query_and_header_parameters_remain_compatible(self): + baseline = generator.generate(self.original) + for operation in ['searchHistory', 'getDomainProfile']: + for location in ['query', 'header']: + with self.subTest(operation=operation, location=location): + doc = copy.deepcopy(self.original) + _, op = self.operation(doc, operation) + op.setdefault('parameters', []).append({ + 'name': 'future', 'in': location, 'required': False, + 'schema': {'type': 'string'}}) + self.assertEqual(baseline, generator.generate(doc)) + + def test_profile_route_changes_require_adapter_update(self): + doc = copy.deepcopy(self.original) + item, _ = self.operation(doc, 'getDomainProfile') + old = next(path for path, value in doc['paths'].items() if value is item) + doc['paths']['/v2/profile'] = doc['paths'].pop(old) + with self.assertRaisesRegex(ValueError, 'Unsupported profile route'): + generator.generate(doc) + + def test_ping_common_field_rename_and_retype_require_adapter_update(self): + for field in ['event_type', 'schema_version', 'created_at']: + for change in ['rename', 'retype']: + with self.subTest(field=field, change=change): + doc = copy.deepcopy(self.original) + props = doc['components']['schemas']['WebhookPingEvent']['properties'] + if change == 'rename': + props[field + '_changed'] = props.pop(field) + else: + props[field] = {'type': 'integer'} + with self.assertRaisesRegex(ValueError, 'Incompatible shared webhook field'): + generator.generate(doc) + + +if __name__ == '__main__': + unittest.main() diff --git a/scripts/verify-package.sh b/scripts/verify-package.sh new file mode 100644 index 0000000..2ce39f7 --- /dev/null +++ b/scripts/verify-package.sh @@ -0,0 +1,13 @@ +#!/usr/bin/env bash +set -euo pipefail +cd "$(dirname "${BASH_SOURCE[0]}")/.." +work="$(mktemp -d)" +trap 'rm -rf "$work"' EXIT +version="$(mvn -B -ntp -q org.apache.maven.plugins:maven-help-plugin:3.5.2:evaluate -Dexpression=project.version -DforceStdout)" +test -f "target/shieldlabs-java-${version}.jar" +mvn -B -ntp org.apache.maven.plugins:maven-dependency-plugin:3.8.1:copy-dependencies \ + -DincludeScope=runtime -DoutputDirectory="$work/lib" +cp "target/shieldlabs-java-${version}.jar" "$work/lib/" +cp scripts/consumer/Consumer.java "$work/" +javac --release 11 -Xlint:all -Werror -cp "$work/lib/*" -d "$work" "$work/Consumer.java" +java -cp "$work:$work/lib/*" Consumer diff --git a/src/main/java/ai/shieldlabs/DetectionFlag.java b/src/main/java/ai/shieldlabs/DetectionFlag.java index 93fb626..79f1718 100644 --- a/src/main/java/ai/shieldlabs/DetectionFlag.java +++ b/src/main/java/ai/shieldlabs/DetectionFlag.java @@ -1,6 +1,7 @@ package ai.shieldlabs; import com.fasterxml.jackson.annotation.JsonValue; +import java.util.function.Function; /** * The 19 detection flags of an identification, in the order of the webhook contract. @@ -9,48 +10,55 @@ */ public enum DetectionFlag { /** A VPN was detected. */ - VPN("vpn", "is_vpn"), + VPN("vpn", "is_vpn", WireModels.DetectionFlags::vpn, WireModels.HistoryRow::is_vpn), /** A privacy relay (such as a platform relay service) was detected. */ - PRIVACY_RELAY("privacy_relay", "is_privacy_relay"), + PRIVACY_RELAY("privacy_relay", "is_privacy_relay", WireModels.DetectionFlags::privacy_relay, WireModels.HistoryRow::is_privacy_relay), /** A browser VPN or proxy extension was detected ({@code connection_type} {@code browser_vpn_proxy}). */ - BROWSER_VPN_PROXY("browser_vpn_proxy", null), + BROWSER_VPN_PROXY("browser_vpn_proxy", null, WireModels.DetectionFlags::browser_vpn_proxy, null), /** The request came through Tor. */ - TOR("tor", "is_tor"), + TOR("tor", "is_tor", WireModels.DetectionFlags::tor, WireModels.HistoryRow::is_tor), /** A proxy was detected. */ - PROXY("proxy", "is_proxy"), + PROXY("proxy", "is_proxy", WireModels.DetectionFlags::proxy, WireModels.HistoryRow::is_proxy), /** The public IP belongs to a hosting or datacenter network. */ - DATACENTER_IP("datacenter_ip", "is_datacenter"), + DATACENTER_IP("datacenter_ip", "is_datacenter", WireModels.DetectionFlags::datacenter_ip, WireModels.HistoryRow::is_datacenter), /** The public IP has a record of abuse. */ - ABUSER("abuser", "is_abuser"), + ABUSER("abuser", "is_abuser", WireModels.DetectionFlags::abuser, WireModels.HistoryRow::is_abuser), /** The operating system reported by the browser does not match the network evidence. */ - OS_MISMATCH("os_mismatch", "is_os_mismatch"), + OS_MISMATCH("os_mismatch", "is_os_mismatch", WireModels.DetectionFlags::os_mismatch, WireModels.HistoryRow::is_os_mismatch), /** The operating system could not be detected. */ - OS_NOT_DETECTED("os_not_detected", "is_os_not_detected"), + OS_NOT_DETECTED("os_not_detected", "is_os_not_detected", WireModels.DetectionFlags::os_not_detected, WireModels.HistoryRow::is_os_not_detected), /** The browser time zone does not match the IP location. */ - TIMEZONE_MISMATCH("timezone_mismatch", "is_timezone_mismatch"), + TIMEZONE_MISMATCH("timezone_mismatch", "is_timezone_mismatch", WireModels.DetectionFlags::timezone_mismatch, WireModels.HistoryRow::is_timezone_mismatch), /** An anti-detect browser was detected. */ - ANTI_DETECT_BROWSER("anti_detect_browser", "is_antidetect"), + ANTI_DETECT_BROWSER("anti_detect_browser", "is_antidetect", WireModels.DetectionFlags::anti_detect_browser, WireModels.HistoryRow::is_antidetect), /** Browser automation was detected. */ - BROWSER_AUTOMATION("browser_automation", "is_browser_automation"), + BROWSER_AUTOMATION("browser_automation", "is_browser_automation", WireModels.DetectionFlags::browser_automation, WireModels.HistoryRow::is_browser_automation), /** The public IP differs from the local network IP. Informational. */ - IP_MISMATCH("ip_mismatch", null), + IP_MISMATCH("ip_mismatch", null, WireModels.DetectionFlags::ip_mismatch, null), /** The browser runs in a private window. */ - INCOGNITO("incognito", "is_incognito"), + INCOGNITO("incognito", "is_incognito", WireModels.DetectionFlags::incognito, WireModels.HistoryRow::is_incognito), /** The visit is a search engine crawler (its Risk Score is forced to 0). */ - SEARCH_BOT("search_bot", "is_search_bot"), + SEARCH_BOT("search_bot", "is_search_bot", WireModels.DetectionFlags::search_bot, WireModels.HistoryRow::is_search_bot), /** A paid click with a high Risk Score. */ - SUSPICIOUS_PAID_CLICK("suspicious_paid_click", "is_suspicious_paid_click"), + SUSPICIOUS_PAID_CLICK("suspicious_paid_click", "is_suspicious_paid_click", WireModels.DetectionFlags::suspicious_paid_click, WireModels.HistoryRow::is_suspicious_paid_click), /** JavaScript was disabled. */ - JAVASCRIPT_DISABLED("javascript_disabled", "is_js_disabled"), + JAVASCRIPT_DISABLED("javascript_disabled", "is_js_disabled", WireModels.DetectionFlags::javascript_disabled, WireModels.HistoryRow::is_js_disabled), /** The network check did not complete. */ - STUN_NOT_CHECKED("stun_not_checked", "is_stun_not_checked"), + STUN_NOT_CHECKED("stun_not_checked", "is_stun_not_checked", WireModels.DetectionFlags::stun_not_checked, WireModels.HistoryRow::is_stun_not_checked), /** A browser check timed out. Informational. */ - CHECK_INCOMPLETE("check_incomplete", "check_incomplete"); + CHECK_INCOMPLETE("check_incomplete", "check_incomplete", WireModels.DetectionFlags::check_incomplete, WireModels.HistoryRow::check_incomplete); private final String value; private final String historyKey; - DetectionFlag(String value, String historyKey) { + private final Function webhookReader; + private final Function historyReader; + + DetectionFlag(String value, String historyKey, + Function webhookReader, + Function historyReader) { + this.webhookReader = webhookReader; + this.historyReader = historyReader; this.value = value; this.historyKey = historyKey; } @@ -65,6 +73,14 @@ public String getValue() { return value; } + Object historyValue(WireModels.HistoryRow row) { + return historyReader == null ? null : WireValue.bool(historyReader.apply(row)); + } + + Object webhookValue(WireModels.DetectionFlags data) { + return WireValue.bool(webhookReader.apply(data)); + } + /** The History API column for this flag, or {@code null} when the flag is derived. */ String historyKey() { return historyKey; diff --git a/src/main/java/ai/shieldlabs/DomainProfile.java b/src/main/java/ai/shieldlabs/DomainProfile.java index 5b9fb13..a221351 100644 --- a/src/main/java/ai/shieldlabs/DomainProfile.java +++ b/src/main/java/ai/shieldlabs/DomainProfile.java @@ -43,13 +43,14 @@ public final class DomainProfile { } static DomainProfile fromJson(Map body) { - Long weight = Json.longValue(body.get("Weight")); + WireModels.DomainProfile bodyWire = new WireModels.DomainProfile(body); + Long weight = Json.longValue(WireValue.integer(bodyWire.Weight())); return new DomainProfile( - Json.text(body.get("Domain")), + Json.text(WireValue.string(bodyWire.Domain())), weight == null ? 0L : weight, - Json.text(body.get("PublicKey")), - Json.text(body.get("Secret")), - Timestamps.parseRfc3339(body.get("CreatedAt")), + Json.text(WireValue.string(bodyWire.PublicKey())), + Json.text(WireValue.string(bodyWire.Secret())), + Timestamps.parseRfc3339(WireValue.string(bodyWire.CreatedAt())), Json.freezeObject(body)); } diff --git a/src/main/java/ai/shieldlabs/HistoryService.java b/src/main/java/ai/shieldlabs/HistoryService.java index ca6f4d9..fa21678 100644 --- a/src/main/java/ai/shieldlabs/HistoryService.java +++ b/src/main/java/ai/shieldlabs/HistoryService.java @@ -167,8 +167,9 @@ private Iterator iterator(LookupType type, String value, History URI uri(LookupType type, String value, int limit, long offset) { String checked = Validation.lookupValue(type, value); - return URI.create(origin + "/api/v1/history/" + type.getValue() + "/" + Urls.encodePathSegment(checked) - + "?limit=" + limit + "&offset=" + offset); + return URI.create(origin + WireModels.HISTORY_PATH.replace("{search_type}", type.getValue()) + .replace("{value}", Urls.encodePathSegment(checked)) + + WireModels.historyQuery(limit, offset)); } static HistoryPage parsePage(JsonResponse response) { @@ -176,7 +177,8 @@ static HistoryPage parsePage(JsonResponse response) { if (body == null) { throw response.unexpected("expected a JSON object with data and total"); } - Object data = body.get("data"); + WireModels.HistoryPage bodyWire = new WireModels.HistoryPage(body); + Object data = WireValue.array(bodyWire.data()); if (data != null && !(data instanceof List)) { throw response.unexpected("data is not an array"); } @@ -188,7 +190,7 @@ static HistoryPage parsePage(JsonResponse response) { items.add(Normalizer.fromHistoryRow(object)); } } - Long total = Json.longValue(body.get("total")); + Long total = Json.longValue(WireValue.integer(bodyWire.total())); return new HistoryPage(items, total == null ? rows.size() : total, rows.size()); } diff --git a/src/main/java/ai/shieldlabs/LookupType.java b/src/main/java/ai/shieldlabs/LookupType.java index efb1e26..ca5f4a8 100644 --- a/src/main/java/ai/shieldlabs/LookupType.java +++ b/src/main/java/ai/shieldlabs/LookupType.java @@ -11,22 +11,22 @@ */ public enum LookupType { /** Dotted IPv4 address (IPv6 addresses are not searchable). */ - IP("ip"), + IP(WireModels.SearchType.IP.wire()), /** * User HID exactly as passed to the agent, including {@code "anonymous"}. Values that contain * {@code /}, and the values {@code .} and {@code ..}, cannot be searched. */ - USER_HID("user_hid"), + USER_HID(WireModels.SearchType.USER_HID.wire()), /** Visitor ID (UUID). */ - VISITOR_ID("visitor_id"), + VISITOR_ID(WireModels.SearchType.VISITOR_ID.wire()), /** Request ID (UUID) of one identification. */ - REQUEST_ID("request_id"), + REQUEST_ID(WireModels.SearchType.REQUEST_ID.wire()), /** Device ID (UUID). */ - DEVICE_ID("device_id"), + DEVICE_ID(WireModels.SearchType.DEVICE_ID.wire()), /** Session ID (UUID). */ - SESSION_ID("session_id"), + SESSION_ID(WireModels.SearchType.SESSION_ID.wire()), /** Cookie ID (UUID). */ - COOKIE_ID("cookie_id"); + COOKIE_ID(WireModels.SearchType.COOKIE_ID.wire()); private final String value; diff --git a/src/main/java/ai/shieldlabs/ManagementClient.java b/src/main/java/ai/shieldlabs/ManagementClient.java index 9771048..86029f2 100644 --- a/src/main/java/ai/shieldlabs/ManagementClient.java +++ b/src/main/java/ai/shieldlabs/ManagementClient.java @@ -56,10 +56,11 @@ private ManagementClient(Builder builder) { builder.httpClient != null ? builder.httpClient : HttpClient.newBuilder().connectTimeout(builder.timeout).build(); + String[] domainHeader = WireModels.profileHeaders(normalized); String[] headers = { "Accept", "application/json", "User-Agent", Transport.userAgent(), - "X-Shield-Domain", normalized, + domainHeader[0], domainHeader[1], "Authorization", "Bearer " + secretKey, }; this.baseUrl = URI.create(origin); diff --git a/src/main/java/ai/shieldlabs/Normalizer.java b/src/main/java/ai/shieldlabs/Normalizer.java index 42e86bc..bb09aa3 100644 --- a/src/main/java/ai/shieldlabs/Normalizer.java +++ b/src/main/java/ai/shieldlabs/Normalizer.java @@ -121,148 +121,155 @@ static String ip(Object value) { } static Identification fromHistoryRow(Map row) { - Object leakSourceValue = row.get("webrtc_leak_source"); + WireModels.HistoryRow rowWire = new WireModels.HistoryRow(row); + Object leakSourceValue = WireValue.string(rowWire.webrtc_leak_source()); String leakSource = leakSourceValue instanceof String ? Text.strip((String) leakSourceValue) : ""; String localIp; String localCountry; if (!leakSource.isEmpty() && !"none".equals(leakSource)) { - localIp = ip(row.get("webrtc_leak_ip")); - localCountry = Json.orEmpty(row.get("webrtc_leak_country")); + localIp = ip(WireValue.string(rowWire.webrtc_leak_ip())); + localCountry = Json.orEmpty(WireValue.string(rowWire.webrtc_leak_country())); } else { - localIp = ip(row.get("web_rtc_ip")); - localCountry = Json.orEmpty(row.get("web_rtc_country")); + localIp = ip(WireValue.string(rowWire.web_rtc_ip())); + localCountry = Json.orEmpty(WireValue.string(rowWire.web_rtc_country())); } - String publicIp = ip(row.get("ip")); + String publicIp = ip(WireValue.string(rowWire.ip())); List signals = new ArrayList<>(); boolean ipLeakDetail = false; - for (Object entry : scoreDetails(row.get("score_details"))) { + for (Object entry : scoreDetails(WireValue.string(rowWire.score_details()))) { Map detail = Json.object(entry); if (detail == null) { continue; } - String description = Json.orEmpty(detail.get("Description")); + WireModels.ScoreDetail detailWire = new WireModels.ScoreDetail(detail); + String description = Json.orEmpty(WireValue.string(detailWire.Description())); if (description.startsWith(IP_LEAK_PREFIX)) { ipLeakDetail = true; } - Integer weight = Json.intValue(detail.containsKey("Value") ? detail.get("Value") : 0); + Integer weight = Json.intValue(detail.containsKey("Value") ? WireValue.integer(detailWire.Value()) : 0); if (weight == null || weight == 0) { continue; } signals.add(new Signal(signalSlug(description), weight, description)); } - boolean searchBot = Json.truthy(row.get("is_search_bot")); + boolean searchBot = Json.truthy(WireValue.bool(rowWire.is_search_bot())); EnumSet flags = EnumSet.noneOf(DetectionFlag.class); for (DetectionFlag flag : DetectionFlag.values()) { boolean value; if (flag == DetectionFlag.BROWSER_VPN_PROXY) { - value = ConnectionType.BROWSER_VPN_PROXY.equals(row.get("connection_type")); + value = ConnectionType.BROWSER_VPN_PROXY.equals(WireValue.string(rowWire.connection_type())); } else if (flag == DetectionFlag.IP_MISMATCH) { value = !searchBot && (ipLeakDetail || (!publicIp.isEmpty() && !localIp.isEmpty() && !publicIp.equals(localIp))); } else { - value = Json.truthy(row.get(flag.historyKey())); + value = Json.truthy(flag.historyValue(rowWire)); } if (value) { flags.add(flag); } } - String siteDomain = Json.orEmpty(row.get("site_domain")); + String siteDomain = Json.orEmpty(WireValue.string(rowWire.site_domain())); TrafficSource traffic = new TrafficSource( - Json.orEmpty(row.get("traffic_channel")), - Json.orEmpty(row.get("referrer_domain")), - Json.orEmpty(row.get("entry_url")), - Json.orEmpty(row.get("click_id_type")), - Json.orEmpty(row.get("utm_source")), - Json.orEmpty(row.get("utm_medium")), - Json.orEmpty(row.get("utm_campaign")), - Json.orEmpty(row.get("utm_content")), - Json.orEmpty(row.get("utm_term"))); + Json.orEmpty(WireValue.string(rowWire.traffic_channel())), + Json.orEmpty(WireValue.string(rowWire.referrer_domain())), + Json.orEmpty(WireValue.string(rowWire.entry_url())), + Json.orEmpty(WireValue.string(rowWire.click_id_type())), + Json.orEmpty(WireValue.string(rowWire.utm_source())), + Json.orEmpty(WireValue.string(rowWire.utm_medium())), + Json.orEmpty(WireValue.string(rowWire.utm_campaign())), + Json.orEmpty(WireValue.string(rowWire.utm_content())), + Json.orEmpty(WireValue.string(rowWire.utm_term()))); return new Identification( - Json.text(row.get("request_id")), - Json.text(row.get("visitor_id")), - Json.text(row.get("device_id")), - Json.text(row.get("session_id")), - Json.text(row.get("cookie_id")), - userHid(row.get("user_hid")), - siteDomain.isEmpty() ? Json.text(row.get("domain")) : siteDomain, - new IpInfo(publicIp, Json.orEmpty(row.get("country"))), + Json.text(WireValue.string(rowWire.request_id())), + Json.text(WireValue.string(rowWire.visitor_id())), + Json.text(WireValue.string(rowWire.device_id())), + Json.text(WireValue.string(rowWire.session_id())), + Json.text(WireValue.string(rowWire.cookie_id())), + userHid(WireValue.string(rowWire.user_hid())), + siteDomain.isEmpty() ? Json.text(WireValue.string(rowWire.domain())) : siteDomain, + new IpInfo(publicIp, Json.orEmpty(WireValue.string(rowWire.country()))), new IpInfo(localIp, localCountry), - Json.text(row.get("connection_type")), - Json.text(row.get("os")), - Json.text(row.get("browser")), - Json.text(row.get("device_type")), + Json.text(WireValue.string(rowWire.connection_type())), + Json.text(WireValue.string(rowWire.os())), + Json.text(WireValue.string(rowWire.browser())), + Json.text(WireValue.string(rowWire.device_type())), traffic, - score(row.get("score")), + score(WireValue.integer(rowWire.score())), signals, new DetectionFlags(flags), - Timestamps.parseHistoryTime(row.get("created_at")), + Timestamps.parseHistoryTime(WireValue.string(rowWire.created_at())), Identification.Source.HISTORY, Json.freezeObject(row)); } static Identification fromWebhookData(Map data) { - Map flagValues = Json.objectOrEmpty(data.get("detection_flags")); + WireModels.IdentificationScoredData dataWire = new WireModels.IdentificationScoredData(data); + Map flagValues = Json.objectOrEmpty(WireValue.object(dataWire.detection_flags())); + WireModels.DetectionFlags flagsWire = new WireModels.DetectionFlags(flagValues); EnumSet flags = EnumSet.noneOf(DetectionFlag.class); for (DetectionFlag flag : DetectionFlag.values()) { - if (Json.truthy(flagValues.get(flag.getValue()))) { + if (Json.truthy(flag.webhookValue(flagsWire))) { flags.add(flag); } } - Map ts = Json.objectOrEmpty(data.get("traffic_source")); + Map ts = Json.objectOrEmpty(WireValue.object(dataWire.traffic_source())); + WireModels.TrafficSource tsWire = new WireModels.TrafficSource(ts); TrafficSource traffic = new TrafficSource( - Json.orEmpty(ts.get("channel")), - Json.orEmpty(ts.get("referrer_domain")), - Json.orEmpty(ts.get("landing_url")), - Json.orEmpty(ts.get("click_id_type")), - Json.orEmpty(ts.get("utm_source")), - Json.orEmpty(ts.get("utm_medium")), - Json.orEmpty(ts.get("utm_campaign")), - Json.orEmpty(ts.get("utm_content")), - Json.orEmpty(ts.get("utm_term"))); + Json.orEmpty(WireValue.string(tsWire.channel())), + Json.orEmpty(WireValue.string(tsWire.referrer_domain())), + Json.orEmpty(WireValue.string(tsWire.landing_url())), + Json.orEmpty(WireValue.string(tsWire.click_id_type())), + Json.orEmpty(WireValue.string(tsWire.utm_source())), + Json.orEmpty(WireValue.string(tsWire.utm_medium())), + Json.orEmpty(WireValue.string(tsWire.utm_campaign())), + Json.orEmpty(WireValue.string(tsWire.utm_content())), + Json.orEmpty(WireValue.string(tsWire.utm_term()))); List signals = new ArrayList<>(); - Object signalValues = data.get("signals"); + Object signalValues = WireValue.array(dataWire.signals()); if (signalValues instanceof List) { for (Object entry : (List) signalValues) { Map signal = Json.object(entry); if (signal == null) { continue; } - Integer weight = Json.intValue(signal.get("weight")); - signals.add(new Signal(Json.text(signal.get("name")), weight == null ? 0 : weight, null)); + WireModels.Signal signalWire = new WireModels.Signal(signal); + Integer weight = Json.intValue(WireValue.integer(signalWire.weight())); + signals.add(new Signal(Json.text(WireValue.string(signalWire.name())), weight == null ? 0 : weight, null)); } } return new Identification( - Json.text(data.get("request_id")), - Json.text(data.get("visitor_id")), - Json.text(data.get("device_id")), - Json.text(data.get("session_id")), - Json.text(data.get("cookie_id")), - userHid(data.get("user_hid")), - Json.text(data.get("domain")), - ipInfo(data.get("public_ip")), - ipInfo(data.get("local_ip")), - Json.text(data.get("connection_type")), - Json.text(data.get("os")), - Json.text(data.get("browser")), - Json.text(data.get("device_type")), + Json.text(WireValue.string(dataWire.request_id())), + Json.text(WireValue.string(dataWire.visitor_id())), + Json.text(WireValue.string(dataWire.device_id())), + Json.text(WireValue.string(dataWire.session_id())), + Json.text(WireValue.string(dataWire.cookie_id())), + userHid(WireValue.string(dataWire.user_hid())), + Json.text(WireValue.string(dataWire.domain())), + ipInfo(WireValue.object(dataWire.public_ip())), + ipInfo(WireValue.object(dataWire.local_ip())), + Json.text(WireValue.string(dataWire.connection_type())), + Json.text(WireValue.string(dataWire.os())), + Json.text(WireValue.string(dataWire.browser())), + Json.text(WireValue.string(dataWire.device_type())), traffic, - score(data.get("risk_score")), + score(WireValue.integer(dataWire.risk_score())), signals, new DetectionFlags(flags), - Timestamps.parseRfc3339(data.get("observed_at")), + Timestamps.parseRfc3339(WireValue.string(dataWire.observed_at())), Identification.Source.WEBHOOK, Json.freezeObject(data)); } private static IpInfo ipInfo(Object value) { Map object = Json.objectOrEmpty(value); - return new IpInfo(ip(object.get("ip")), Json.orEmpty(object.get("country"))); + WireModels.IpInfo objectWire = new WireModels.IpInfo(object); + return new IpInfo(ip(WireValue.string(objectWire.ip())), Json.orEmpty(WireValue.string(objectWire.country()))); } /** {@code ""} becomes {@code null}; every other value, including placeholders, is kept. */ diff --git a/src/main/java/ai/shieldlabs/Webhooks.java b/src/main/java/ai/shieldlabs/Webhooks.java index 330f28a..34cb73f 100644 --- a/src/main/java/ai/shieldlabs/Webhooks.java +++ b/src/main/java/ai/shieldlabs/Webhooks.java @@ -243,18 +243,19 @@ private static WebhookEvent parse(byte[] payload) { if (envelope == null) { throw new WebhookParseException("The webhook body is not a JSON object"); } - Object type = envelope.get("event_type"); + WireModels.IdentificationScoredEvent envelopeWire = new WireModels.IdentificationScoredEvent(envelope); + Object type = WireValue.string(envelopeWire.event_type()); if (!(type instanceof String) || ((String) type).isEmpty()) { throw new WebhookParseException("The webhook body has no event_type"); } - Object version = envelope.get("schema_version"); + Object version = WireValue.string(envelopeWire.schema_version()); String schemaVersion = version instanceof String ? (String) version : null; warnOnUnknownSchemaVersion(schemaVersion); - Instant createdAt = Timestamps.parseRfc3339(envelope.get("created_at")); + Instant createdAt = Timestamps.parseRfc3339(WireValue.string(envelopeWire.created_at())); Map raw = Json.freezeObject(envelope); switch ((String) type) { case WebhookEvent.IDENTIFICATION_SCORED: - Map data = Json.object(envelope.get("data")); + Map data = Json.object(WireValue.object(envelopeWire.data())); if (data == null) { throw new WebhookParseException("The identification.scored event has no data object"); } diff --git a/src/main/java/ai/shieldlabs/WireModels.java b/src/main/java/ai/shieldlabs/WireModels.java new file mode 100644 index 0000000..b916071 --- /dev/null +++ b/src/main/java/ai/shieldlabs/WireModels.java @@ -0,0 +1,184 @@ +// Generated by scripts/generate-wire.py. Do not edit. +package ai.shieldlabs; + +import java.util.Map; + +/** Schema-typed views. Values remain uncoerced for compatibility with older responses. */ +final class WireModels { + private WireModels() {} + static final class HistoryRow { + private final Map raw; + HistoryRow(Map raw) { this.raw = raw; } + WireValue.StringValue request_id() { return new WireValue.StringValue(raw.get("request_id")); } + WireValue.StringValue session_id() { return new WireValue.StringValue(raw.get("session_id")); } + WireValue.StringValue cookie_id() { return new WireValue.StringValue(raw.get("cookie_id")); } + WireValue.StringValue domain() { return new WireValue.StringValue(raw.get("domain")); } + WireValue.StringValue site_domain() { return new WireValue.StringValue(raw.get("site_domain")); } + WireValue.StringValue user_hid() { return new WireValue.StringValue(raw.get("user_hid")); } + WireValue.StringValue device_id() { return new WireValue.StringValue(raw.get("device_id")); } + WireValue.StringValue visitor_id() { return new WireValue.StringValue(raw.get("visitor_id")); } + WireValue.StringValue ip() { return new WireValue.StringValue(raw.get("ip")); } + WireValue.StringValue os() { return new WireValue.StringValue(raw.get("os")); } + WireValue.StringValue browser() { return new WireValue.StringValue(raw.get("browser")); } + WireValue.StringValue device_type() { return new WireValue.StringValue(raw.get("device_type")); } + WireValue.StringValue country() { return new WireValue.StringValue(raw.get("country")); } + WireValue.StringValue connection_type() { return new WireValue.StringValue(raw.get("connection_type")); } + WireValue.IntegerValue score() { return new WireValue.IntegerValue(raw.get("score")); } + WireValue.StringValue score_details() { return new WireValue.StringValue(raw.get("score_details")); } + WireValue.StringValue created_at() { return new WireValue.StringValue(raw.get("created_at")); } + WireValue.IntegerValue ver() { return new WireValue.IntegerValue(raw.get("ver")); } + WireValue.StringValue web_rtc_ip() { return new WireValue.StringValue(raw.get("web_rtc_ip")); } + WireValue.StringValue web_rtc_country() { return new WireValue.StringValue(raw.get("web_rtc_country")); } + WireValue.StringValue web_rtc_connection_type() { return new WireValue.StringValue(raw.get("web_rtc_connection_type")); } + WireValue.StringValue webrtc_leak_ip() { return new WireValue.StringValue(raw.get("webrtc_leak_ip")); } + WireValue.StringValue webrtc_leak_country() { return new WireValue.StringValue(raw.get("webrtc_leak_country")); } + WireValue.StringValue webrtc_leak_connection_type() { return new WireValue.StringValue(raw.get("webrtc_leak_connection_type")); } + WireValue.StringValue webrtc_leak_source() { return new WireValue.StringValue(raw.get("webrtc_leak_source")); } + WireValue.BooleanValue is_vpn() { return new WireValue.BooleanValue(raw.get("is_vpn")); } + WireValue.BooleanValue is_tor() { return new WireValue.BooleanValue(raw.get("is_tor")); } + WireValue.BooleanValue is_proxy() { return new WireValue.BooleanValue(raw.get("is_proxy")); } + WireValue.BooleanValue is_datacenter() { return new WireValue.BooleanValue(raw.get("is_datacenter")); } + WireValue.BooleanValue is_abuser() { return new WireValue.BooleanValue(raw.get("is_abuser")); } + WireValue.BooleanValue is_privacy_relay() { return new WireValue.BooleanValue(raw.get("is_privacy_relay")); } + WireValue.BooleanValue is_stun_not_checked() { return new WireValue.BooleanValue(raw.get("is_stun_not_checked")); } + WireValue.BooleanValue check_incomplete() { return new WireValue.BooleanValue(raw.get("check_incomplete")); } + WireValue.BooleanValue is_antidetect() { return new WireValue.BooleanValue(raw.get("is_antidetect")); } + WireValue.BooleanValue is_os_mismatch() { return new WireValue.BooleanValue(raw.get("is_os_mismatch")); } + WireValue.BooleanValue is_os_not_detected() { return new WireValue.BooleanValue(raw.get("is_os_not_detected")); } + WireValue.BooleanValue is_timezone_mismatch() { return new WireValue.BooleanValue(raw.get("is_timezone_mismatch")); } + WireValue.BooleanValue is_js_disabled() { return new WireValue.BooleanValue(raw.get("is_js_disabled")); } + WireValue.BooleanValue is_browser_automation() { return new WireValue.BooleanValue(raw.get("is_browser_automation")); } + WireValue.BooleanValue is_incognito() { return new WireValue.BooleanValue(raw.get("is_incognito")); } + WireValue.BooleanValue is_search_bot() { return new WireValue.BooleanValue(raw.get("is_search_bot")); } + WireValue.BooleanValue is_suspicious_paid_click() { return new WireValue.BooleanValue(raw.get("is_suspicious_paid_click")); } + WireValue.StringValue entry_url() { return new WireValue.StringValue(raw.get("entry_url")); } + WireValue.StringValue utm_source() { return new WireValue.StringValue(raw.get("utm_source")); } + WireValue.StringValue utm_medium() { return new WireValue.StringValue(raw.get("utm_medium")); } + WireValue.StringValue utm_campaign() { return new WireValue.StringValue(raw.get("utm_campaign")); } + WireValue.StringValue utm_content() { return new WireValue.StringValue(raw.get("utm_content")); } + WireValue.StringValue utm_term() { return new WireValue.StringValue(raw.get("utm_term")); } + WireValue.StringValue traffic_channel() { return new WireValue.StringValue(raw.get("traffic_channel")); } + WireValue.StringValue traffic_channel_group() { return new WireValue.StringValue(raw.get("traffic_channel_group")); } + WireValue.StringValue traffic_reason() { return new WireValue.StringValue(raw.get("traffic_reason")); } + WireValue.StringValue referrer_domain() { return new WireValue.StringValue(raw.get("referrer_domain")); } + WireValue.StringValue click_id_type() { return new WireValue.StringValue(raw.get("click_id_type")); } + } + static final class HistoryPage { + private final Map raw; + HistoryPage(Map raw) { this.raw = raw; } + WireValue.ArrayValue data() { return new WireValue.ArrayValue(raw.get("data")); } + WireValue.IntegerValue total() { return new WireValue.IntegerValue(raw.get("total")); } + } + static final class DomainProfile { + private final Map raw; + DomainProfile(Map raw) { this.raw = raw; } + WireValue.StringValue Domain() { return new WireValue.StringValue(raw.get("Domain")); } + WireValue.IntegerValue Weight() { return new WireValue.IntegerValue(raw.get("Weight")); } + WireValue.StringValue Callback() { return new WireValue.StringValue(raw.get("Callback")); } + WireValue.StringValue PublicKey() { return new WireValue.StringValue(raw.get("PublicKey")); } + WireValue.StringValue Secret() { return new WireValue.StringValue(raw.get("Secret")); } + WireValue.StringValue CreatedAt() { return new WireValue.StringValue(raw.get("CreatedAt")); } + } + static final class ScoreDetail { + private final Map raw; + ScoreDetail(Map raw) { this.raw = raw; } + WireValue.IntegerValue Value() { return new WireValue.IntegerValue(raw.get("Value")); } + WireValue.StringValue Description() { return new WireValue.StringValue(raw.get("Description")); } + } + static final class IdentificationScoredData { + private final Map raw; + IdentificationScoredData(Map raw) { this.raw = raw; } + WireValue.StringValue request_id() { return new WireValue.StringValue(raw.get("request_id")); } + WireValue.StringValue visitor_id() { return new WireValue.StringValue(raw.get("visitor_id")); } + WireValue.StringValue device_id() { return new WireValue.StringValue(raw.get("device_id")); } + WireValue.StringValue session_id() { return new WireValue.StringValue(raw.get("session_id")); } + WireValue.StringValue cookie_id() { return new WireValue.StringValue(raw.get("cookie_id")); } + WireValue.StringValue user_hid() { return new WireValue.StringValue(raw.get("user_hid")); } + WireValue.StringValue domain() { return new WireValue.StringValue(raw.get("domain")); } + WireValue.ObjectValue public_ip() { return new WireValue.ObjectValue(raw.get("public_ip")); } + WireValue.ObjectValue local_ip() { return new WireValue.ObjectValue(raw.get("local_ip")); } + WireValue.StringValue connection_type() { return new WireValue.StringValue(raw.get("connection_type")); } + WireValue.StringValue os() { return new WireValue.StringValue(raw.get("os")); } + WireValue.StringValue browser() { return new WireValue.StringValue(raw.get("browser")); } + WireValue.StringValue device_type() { return new WireValue.StringValue(raw.get("device_type")); } + WireValue.ObjectValue traffic_source() { return new WireValue.ObjectValue(raw.get("traffic_source")); } + WireValue.IntegerValue risk_score() { return new WireValue.IntegerValue(raw.get("risk_score")); } + WireValue.ArrayValue signals() { return new WireValue.ArrayValue(raw.get("signals")); } + WireValue.ObjectValue detection_flags() { return new WireValue.ObjectValue(raw.get("detection_flags")); } + WireValue.StringValue observed_at() { return new WireValue.StringValue(raw.get("observed_at")); } + } + static final class IdentificationScoredEvent { + private final Map raw; + IdentificationScoredEvent(Map raw) { this.raw = raw; } + WireValue.StringValue event_type() { return new WireValue.StringValue(raw.get("event_type")); } + WireValue.StringValue schema_version() { return new WireValue.StringValue(raw.get("schema_version")); } + WireValue.StringValue created_at() { return new WireValue.StringValue(raw.get("created_at")); } + WireValue.ObjectValue data() { return new WireValue.ObjectValue(raw.get("data")); } + } + static final class WebhookPingEvent { + private final Map raw; + WebhookPingEvent(Map raw) { this.raw = raw; } + WireValue.StringValue event_type() { return new WireValue.StringValue(raw.get("event_type")); } + WireValue.StringValue schema_version() { return new WireValue.StringValue(raw.get("schema_version")); } + WireValue.StringValue created_at() { return new WireValue.StringValue(raw.get("created_at")); } + } + static final class TrafficSource { + private final Map raw; + TrafficSource(Map raw) { this.raw = raw; } + WireValue.StringValue channel() { return new WireValue.StringValue(raw.get("channel")); } + WireValue.StringValue referrer_domain() { return new WireValue.StringValue(raw.get("referrer_domain")); } + WireValue.StringValue landing_url() { return new WireValue.StringValue(raw.get("landing_url")); } + WireValue.StringValue click_id_type() { return new WireValue.StringValue(raw.get("click_id_type")); } + WireValue.StringValue utm_source() { return new WireValue.StringValue(raw.get("utm_source")); } + WireValue.StringValue utm_medium() { return new WireValue.StringValue(raw.get("utm_medium")); } + WireValue.StringValue utm_campaign() { return new WireValue.StringValue(raw.get("utm_campaign")); } + WireValue.StringValue utm_content() { return new WireValue.StringValue(raw.get("utm_content")); } + WireValue.StringValue utm_term() { return new WireValue.StringValue(raw.get("utm_term")); } + } + static final class IpInfo { + private final Map raw; + IpInfo(Map raw) { this.raw = raw; } + WireValue.StringValue ip() { return new WireValue.StringValue(raw.get("ip")); } + WireValue.StringValue country() { return new WireValue.StringValue(raw.get("country")); } + } + static final class Signal { + private final Map raw; + Signal(Map raw) { this.raw = raw; } + WireValue.StringValue name() { return new WireValue.StringValue(raw.get("name")); } + WireValue.IntegerValue weight() { return new WireValue.IntegerValue(raw.get("weight")); } + } + static final class DetectionFlags { + private final Map raw; + DetectionFlags(Map raw) { this.raw = raw; } + WireValue.BooleanValue vpn() { return new WireValue.BooleanValue(raw.get("vpn")); } + WireValue.BooleanValue privacy_relay() { return new WireValue.BooleanValue(raw.get("privacy_relay")); } + WireValue.BooleanValue browser_vpn_proxy() { return new WireValue.BooleanValue(raw.get("browser_vpn_proxy")); } + WireValue.BooleanValue tor() { return new WireValue.BooleanValue(raw.get("tor")); } + WireValue.BooleanValue proxy() { return new WireValue.BooleanValue(raw.get("proxy")); } + WireValue.BooleanValue datacenter_ip() { return new WireValue.BooleanValue(raw.get("datacenter_ip")); } + WireValue.BooleanValue abuser() { return new WireValue.BooleanValue(raw.get("abuser")); } + WireValue.BooleanValue os_mismatch() { return new WireValue.BooleanValue(raw.get("os_mismatch")); } + WireValue.BooleanValue os_not_detected() { return new WireValue.BooleanValue(raw.get("os_not_detected")); } + WireValue.BooleanValue timezone_mismatch() { return new WireValue.BooleanValue(raw.get("timezone_mismatch")); } + WireValue.BooleanValue anti_detect_browser() { return new WireValue.BooleanValue(raw.get("anti_detect_browser")); } + WireValue.BooleanValue browser_automation() { return new WireValue.BooleanValue(raw.get("browser_automation")); } + WireValue.BooleanValue ip_mismatch() { return new WireValue.BooleanValue(raw.get("ip_mismatch")); } + WireValue.BooleanValue incognito() { return new WireValue.BooleanValue(raw.get("incognito")); } + WireValue.BooleanValue search_bot() { return new WireValue.BooleanValue(raw.get("search_bot")); } + WireValue.BooleanValue suspicious_paid_click() { return new WireValue.BooleanValue(raw.get("suspicious_paid_click")); } + WireValue.BooleanValue javascript_disabled() { return new WireValue.BooleanValue(raw.get("javascript_disabled")); } + WireValue.BooleanValue stun_not_checked() { return new WireValue.BooleanValue(raw.get("stun_not_checked")); } + WireValue.BooleanValue check_incomplete() { return new WireValue.BooleanValue(raw.get("check_incomplete")); } + } + enum SearchType { + REQUEST_ID, DEVICE_ID, USER_HID, VISITOR_ID, IP, SESSION_ID, COOKIE_ID; + String wire() { return name().toLowerCase(java.util.Locale.ROOT); } + } + static final String HISTORY_PATH = "/api/v1/history/{search_type}/{value}"; + static String historyQuery(int limit, long offset) { + return "?limit=" + limit + "&offset=" + offset; + } + static String[] profileHeaders(String domain) { + return new String[] {"X-Shield-Domain", domain}; + } +} diff --git a/src/main/java/ai/shieldlabs/WireValue.java b/src/main/java/ai/shieldlabs/WireValue.java new file mode 100644 index 0000000..bc89fab --- /dev/null +++ b/src/main/java/ai/shieldlabs/WireValue.java @@ -0,0 +1,41 @@ +package ai.shieldlabs; + +/** + * Typed raw values separate the declared wire shape from tolerant normalization. Accessors accept + * only the schema's declared kind, while passing the original value to the existing coercion rules. + * This deliberately does not validate enums, formats, nulls or unknown response fields. + */ +final class WireValue { + private WireValue() {} + + static final class StringValue { + private final Object raw; + StringValue(Object raw) { this.raw = raw; } + } + static final class IntegerValue { + private final Object raw; + IntegerValue(Object raw) { this.raw = raw; } + } + static final class NumberValue { + private final Object raw; + NumberValue(Object raw) { this.raw = raw; } + } + static final class BooleanValue { + private final Object raw; + BooleanValue(Object raw) { this.raw = raw; } + } + static final class ArrayValue { + private final Object raw; + ArrayValue(Object raw) { this.raw = raw; } + } + static final class ObjectValue { + private final Object raw; + ObjectValue(Object raw) { this.raw = raw; } + } + static Object string(StringValue value) { return value.raw; } + static Object integer(IntegerValue value) { return value.raw; } + static Object number(NumberValue value) { return value.raw; } + static Object bool(BooleanValue value) { return value.raw; } + static Object array(ArrayValue value) { return value.raw; } + static Object object(ObjectValue value) { return value.raw; } +} diff --git a/src/test/java/ai/shieldlabs/WireModelsTest.java b/src/test/java/ai/shieldlabs/WireModelsTest.java new file mode 100644 index 0000000..e4481c1 --- /dev/null +++ b/src/test/java/ai/shieldlabs/WireModelsTest.java @@ -0,0 +1,63 @@ +package ai.shieldlabs; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertSame; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.lang.reflect.Method; +import java.util.HashMap; +import java.util.Map; +import org.junit.jupiter.api.Test; + +class WireModelsTest { + @Test + void everyGeneratedAccessorPreservesRawMissingNullAndUnexpectedValues() throws Exception { + for (Class model : WireModels.class.getDeclaredClasses()) { + if (model.isEnum()) { + continue; + } + Map raw = new HashMap<>(); + Object view = model.getDeclaredConstructor(Map.class).newInstance(raw); + for (Method getter : model.getDeclaredMethods()) { + if (getter.isSynthetic()) { + continue; + } + Method reader = null; + for (Method candidate : WireValue.class.getDeclaredMethods()) { + if (candidate.getParameterCount() == 1 && candidate.getParameterTypes()[0] == getter.getReturnType()) { + reader = candidate; + } + } + assertTrue(reader != null, getter.toString()); + assertNull(reader.invoke(null, getter.invoke(view))); + raw.put(getter.getName(), null); + assertNull(reader.invoke(null, getter.invoke(view))); + Object oddValue = new Object(); + raw.put(getter.getName(), oddValue); + assertSame(oddValue, reader.invoke(null, getter.invoke(view))); + raw.put(getter.getName(), "future-value"); + assertEquals("future-value", reader.invoke(null, getter.invoke(view))); + raw.clear(); + } + } + } + + @Test + void generatedContractDoesNotMakeNormalizationStrict() { + Map row = new HashMap<>(); + row.put("request_id", 123); + row.put("score", "999"); + row.put("connection_type", "future-connection"); + row.put("future_column", Map.of("answer", 42)); + Identification result = Identification.fromHistoryRow(row); + assertEquals("123", result.getRequestId()); + assertEquals(0, result.getRiskScore()); + assertEquals("future-connection", result.getConnectionType()); + assertEquals(row.get("future_column"), result.raw().get("future_column")); + DomainProfile profile = DomainProfile.fromJson(Map.of("Domain", true, "Weight", "12", "extra", 7)); + assertEquals("true", profile.getDomain()); + assertEquals(0, profile.getRemainingIdentifications()); + assertEquals(7, profile.raw().get("extra")); + } +}