diff --git a/uraniborg/docs/automate_observation.md b/uraniborg/docs/automate_observation.md index fcd98ae..228ee5f 100644 --- a/uraniborg/docs/automate_observation.md +++ b/uraniborg/docs/automate_observation.md @@ -100,6 +100,7 @@ Optional fields are left out rather than set to `null`. | :--- | :--- | :--- | | `run_started` | `argv` (arguments, without the script name), `pid` | First event of every run | | `step` | `step`, `state` (`started` / `finished` / `failed`), `device`\*, `duration_ms`\*\*, `message`\* | Around each phase (see below) | +| `step_progress` | `step`, `device`\*, `done`, `total` | Between a step's `started` and its `finished` / `failed`, for long steps (see [Step progress](#step-progress)) | | `devices` | `devices` (list of `{serial, unauthorized, model?, product?, device?}`), `selected` (serials to observe), `missing` (requested via `--serial` but not connected) | Once, after listing devices | | `device_started` | `device` | Before processing each selected device | | `prompt` | `device`, `kind`, `message`, `expects_input` | The script is waiting for a person (see [Manual intervention](#manual-intervention)) | @@ -123,6 +124,35 @@ pre-fetching is attempted) and `inclusion_proof_check`. A failed `inclusion_proof_prefetch` is not fatal: verification falls back to fetching entries on demand. +#### Step progress + +Some steps take minutes. While they run, `step_progress` events report how far +they have got, e.g. to show "120 of 314 checked". Currently only +`inclusion_proof_check` reports progress: `done` is the number of APK splits +verified so far (found in the log or not) and `total` the number of splits to +verify (splits skipped for missing fields are not counted). + +- Each event carries the same `step` and `device` as the step it belongs to, + and appears only between that step's `started` and `finished` / `failed`. +- The first event has `done` = `0`, once verification starts. For + `inclusion_proof_check`, that is after the package list has been read; if it + cannot be read, the step fails without any `step_progress`. +- `done` only increases, and `total` stays the same within a step. +- Events are throttled: at most one every 2 seconds, plus the first one and the + one where `done` reaches `total`, which is always sent. A step with nothing + to verify sends a single event with `done` = `total` = `0`. +- A step that is stopped early (Ctrl-C, `SIGTERM`, an error) ends with + `failed` without reaching `done` = `total`. + +```json +{"v": 1, "ts": "2026-01-02T03:05:00.000Z", "type": "step", "step": "inclusion_proof_check", "state": "started", "device": "ABCDEF012345"} +{"v": 1, "ts": "2026-01-02T03:05:00.010Z", "type": "step_progress", "step": "inclusion_proof_check", "device": "ABCDEF012345", "done": 0, "total": 314} +{"v": 1, "ts": "2026-01-02T03:05:02.020Z", "type": "step_progress", "step": "inclusion_proof_check", "device": "ABCDEF012345", "done": 11, "total": 314} +... +{"v": 1, "ts": "2026-01-02T03:05:55.400Z", "type": "step_progress", "step": "inclusion_proof_check", "device": "ABCDEF012345", "done": 314, "total": 314} +{"v": 1, "ts": "2026-01-02T03:05:55.410Z", "type": "step", "step": "inclusion_proof_check", "state": "finished", "device": "ABCDEF012345", "duration_ms": 55410} +``` + #### Manual intervention Two situations need a person. Each is reported as a `prompt` event, and is @@ -294,6 +324,30 @@ flags: `preinstalled_packages.txt` instead of `packages.txt`, writing results to `preinstalled_packages_with_inclusion_proof_signal.txt`. +### Batch Verification +All APK splits are verified by a single `verifier` run in batch mode +(`--payloads_path`, see the [verifier README](../../verifier_tools/verify/README.md)). +It fetches each log's checkpoint, and searches the log, once for all splits. +This keeps a cold cache (e.g. with `--no_prefetch`, or after pre-fetching +failed) from being filled by many verifiers downloading the same files at +once. Results are reported as the verifier finds them. + +With a `verifier` built before batch mode existed, or for any splits a batch +run gives no result for (e.g. if it crashes), each split is verified by a +`verifier` run of its own instead, one at a time. That is much slower, so a +warning suggests rebuilding the `verifier`. The results file is the same +either way, including the order of packages and splits. + +Log lines about individual splits (shown with `-D`) start with the package and +split they are about, e.g. `com.android.chrome [config.en]: ...`. + +Ctrl-C stops the running verifier and starts no new ones. `SIGTERM` does the +same when `--events` is given (and always for `inclusion_proof_check.py`). +Without `--events`, `automate_observation.py` keeps the default `SIGTERM` +behaviour: the script exits at once, and a verifier that was already running +stops on its own: a per-split verifier after its split (about a second or +two), a batch verifier when it next tries to write a result. + ### Pulling Pre-installed APKs Only By default, `--pull-all-apks` downloads all packages listed in `packages.txt`. You can pass `--pull-preinstalled-apks-only` to download only pre-installed diff --git a/uraniborg/scripts/python/automate_observation.py b/uraniborg/scripts/python/automate_observation.py index 2de00bb..7dc9fda 100644 --- a/uraniborg/scripts/python/automate_observation.py +++ b/uraniborg/scripts/python/automate_observation.py @@ -39,6 +39,7 @@ import inclusion_proof_check import syscall_wrapper +import termination AdbWrapper = syscall_wrapper.AdbWrapper SyscallWrapper = syscall_wrapper.SyscallWrapper @@ -227,6 +228,10 @@ def set_up_logging(args: argparse.Namespace) -> logging.Logger: EVENTS_SCHEMA_VERSION = 1 +# step_progress events are written at most this often (plus the first and +# the final one of each step). +STEP_PROGRESS_INTERVAL_SECONDS = 2.0 + # Per-device outcomes. These are the same four outcomes the final log summary # distinguishes. STATUS_SUCCESS = "success" @@ -273,30 +278,11 @@ def set_up_logging(args: argparse.Namespace) -> logging.Logger: PROMPT_OUTCOME_STDIN_CLOSED = REASON_STDIN_CLOSED -class Terminated(BaseException): - """Raised by the SIGTERM handler that main() installs while --events is on. - - Derives from BaseException, like KeyboardInterrupt, so that the generic - `except Exception` handlers do not swallow it. SyscallWrapper only catches - KeyboardInterrupt, so this also propagates out of adb calls. - """ - - def __init__(self): - super().__init__("received SIGTERM") - - -def _raise_terminated(signum, frame): # pylint: disable=unused-argument - # Ignore repeated SIGTERMs while cleaning up so that run_finished is still - # written; _die_by_sigterm() restores the default action afterwards. - signal.signal(signal.SIGTERM, signal.SIG_IGN) - raise Terminated() - - def _interruption_error(e: BaseException) -> dict: """Describes an exception that is not an Exception, for device errors.""" if isinstance(e, KeyboardInterrupt): return _error(REASON_INTERRUPTED, "Interrupted.") - if isinstance(e, Terminated): + if isinstance(e, termination.Terminated): return _error(REASON_TERMINATED, "Terminated.") return _error(REASON_UNEXPECTED_ERROR, "{}: {}".format(type(e).__name__, e)) @@ -359,11 +345,13 @@ class EventEmitter: def __init__(self, stream: Optional[TextIO] = None, logger: Optional[logging.Logger] = None, - clock=None): + clock=None, monotonic=None): self._stream = stream self._logger = logger self._clock = clock or ( lambda: datetime.datetime.now(datetime.timezone.utc)) + # Used only to throttle step_progress events. + self._monotonic = monotonic or time.monotonic self._run_finished = False # serial -> {status, results_dir?}, in device_finished order. self._finished_devices = {} @@ -402,6 +390,31 @@ def device_finished(self, device: str, status: str, self.emit("device_finished", device=device, status=status, results_dir=results_dir, error=error) + def progress_reporter( + self, step: str, device: Optional[str] = None, + interval: float = STEP_PROGRESS_INTERVAL_SECONDS): + """Returns a progress(done, total) callback emitting step_progress. + + Meant to be called often (e.g. once per item); events are throttled: + the first call is always reported, and so is the one where done reaches + total, but in between at most one event is written per `interval` + seconds. A call repeating the last reported `done` is never reported. + """ + last = {"done": None, "at": None} + + def report(done: int, total: int) -> None: + if self._stream is None or done == last["done"]: + return + now = self._monotonic() + if (last["at"] is not None and done != total and + now - last["at"] < interval): + return + last["done"], last["at"] = done, now + self.emit("step_progress", step=step, device=device, done=done, + total=total) + + return report + @contextlib.contextmanager def step(self, name: str, device: Optional[str] = None): """Brackets a phase with step started/finished/failed events. @@ -1559,7 +1572,8 @@ def _observe_device(target_device, args: argparse.Namespace, concurrency=args.cache_prefetch_concurrency, timeout=args.cache_prefetch_timeout, prefetch=False, - preinstalled_only=args.check_preinstalled_only): + preinstalled_only=args.check_preinstalled_only, + progress=events.progress_reporter("inclusion_proof_check", serial)): progress.check_incomplete = True # False means the check could not complete (bad input or unwritable # output), not that some splits are absent from the log; per-split @@ -1748,45 +1762,32 @@ def main(): # being written, turn it into an exception so the run can report itself, # then die by SIGTERM anyway so the parent sees the usual exit status. # Without --events, SIGTERM handling is left untouched. - previous_sigterm_handler = None - if events.enabled: - previous_sigterm_handler = signal.signal(signal.SIGTERM, _raise_terminated) - terminated = False - try: - exit_code = run(args, logger, events) - except Terminated: - events.finish_run(128 + signal.SIGTERM, error=_error( - REASON_TERMINATED, "Terminated by SIGTERM.")) - terminated = True - except KeyboardInterrupt: - events.finish_run(130, error=_error(REASON_INTERRUPTED, - "Interrupted by user.")) - raise - except BaseException as e: - events.finish_run(1, error=_error( - REASON_UNEXPECTED_ERROR, "{}: {}".format(type(e).__name__, e))) - raise - finally: - events.close() - if previous_sigterm_handler is not None: - signal.signal(signal.SIGTERM, previous_sigterm_handler) + with (termination.sigterm_raises() if events.enabled + else contextlib.nullcontext()): + try: + exit_code = run(args, logger, events) + except termination.Terminated: + events.finish_run(128 + signal.SIGTERM, error=_error( + REASON_TERMINATED, "Terminated by SIGTERM.")) + terminated = True + except KeyboardInterrupt: + events.finish_run(130, error=_error(REASON_INTERRUPTED, + "Interrupted by user.")) + raise + except BaseException as e: + events.finish_run(1, error=_error( + REASON_UNEXPECTED_ERROR, "{}: {}".format(type(e).__name__, e))) + raise + finally: + events.close() if terminated: - _die_by_sigterm(logger) + termination.die_by_sigterm(logger) return if exit_code: sys.exit(exit_code) -def _die_by_sigterm(logger): - """Terminates this process with SIGTERM's default action.""" - logger.error("Terminated by SIGTERM.") - signal.signal(signal.SIGTERM, signal.SIG_DFL) - os.kill(os.getpid(), signal.SIGTERM) - # Not reached unless SIGTERM is blocked; fall back to the conventional code. - sys.exit(128 + signal.SIGTERM) - - if __name__ == "__main__": main() diff --git a/uraniborg/scripts/python/inclusion_proof_check.py b/uraniborg/scripts/python/inclusion_proof_check.py index 63795d3..990deb7 100644 --- a/uraniborg/scripts/python/inclusion_proof_check.py +++ b/uraniborg/scripts/python/inclusion_proof_check.py @@ -17,20 +17,46 @@ """Performs inclusion proof check against packages in packages.txt or preinstalled_packages.txt.""" import argparse +import contextlib import json import logging import os import subprocess import sys import tempfile -from typing import Optional +from typing import Callable, NamedTuple, Optional + +import termination OUTPUT_FILENAME = 'packages_with_inclusion_proof_signal.txt' PREINSTALLED_OUTPUT_FILENAME = ( 'preinstalled_packages_with_inclusion_proof_signal.txt' ) +# progress(done, total); see perform_inclusion_proof_check(). +ProgressCallback = Callable[[int, int], None] DEFAULT_PREFETCH_CONCURRENCY = 16 DEFAULT_PREFETCH_TIMEOUT = 600 +# What a verifier without batch mode prints (Go's flag package, exit code 2) +# when given --payloads_path. Exit code 2 alone is not enough: a Go panic +# exits with 2 too. +_BATCH_UNSUPPORTED_EXIT_CODE = 2 +_BATCH_UNSUPPORTED_MESSAGE = "flag provided but not defined: -payloads_path" + + +class _SplitJob(NamedTuple): + """One APK split to verify.""" + split: dict # The split's entry in the package list; receives the result. + payload: str + label: str # Names the package and split in log lines. + + +def _split_label(package_name: str, split: dict, index: int, + split_count: int) -> str: + """Describes a split for log lines, e.g. 'com.android.chrome [chrome]'.""" + split_name = split.get("name") + if not split_name: + split_name = "split {}".format(index) if split_count > 1 else "base" + return "{} [{}]".format(package_name, split_name) def prefetch_log_entries(verifier_executable: str, @@ -89,8 +115,26 @@ def prefetch_log_entries(verifier_executable: str, def run_verifier(verifier_executable: str, payload_path: str, logger: logging.Logger, - cache_dir: Optional[str] = None) -> bool: - """Runs verifier tool and returns True if inclusion proof is successful.""" + cache_dir: Optional[str] = None, + label: Optional[str] = None) -> bool: + """Runs verifier tool and returns True if inclusion proof is successful. + + If interrupted (KeyboardInterrupt, or an exception raised by a signal + handler), the verifier is killed and the exception propagates. + + Args: + verifier_executable: path to verifier tool. + payload_path: path to the payload file to verify. + logger: logger instance. + cache_dir: optional custom root directory for local cache. + label: names the package and split in log lines. Defaults to + payload_path. + + Returns: + True if the inclusion proof succeeded; False if it failed or the verifier + could not be run. + """ + what = label or payload_path try: cmd = [verifier_executable, f"--payload_path={payload_path}", "--log_type=google_1p_apk"] @@ -98,26 +142,215 @@ def run_verifier(verifier_executable: str, payload_path: str, cmd.append(f"--cache_dir={cache_dir}") with open(payload_path, "r") as f_in: payload = f_in.read() - logger.debug("payload content: %s", payload) - logger.debug("Running verifier: %s", " ".join(cmd)) - result = subprocess.run(cmd, capture_output=True, text=True, check=False) - logger.debug("Verifier stdout: %s", result.stdout) - logger.debug("Verifier stderr: %s", result.stderr) - if ("OK. inclusion check success!" in result.stdout or - "OK. inclusion check success!" in result.stderr): - logger.debug("Verifier check passed.") + logger.debug("%s: payload content: %s", what, payload) + logger.debug("%s: Running verifier: %s", what, " ".join(cmd)) + with subprocess.Popen(cmd, stdout=subprocess.PIPE, stderr=subprocess.PIPE, + text=True) as process: + try: + stdout, stderr = process.communicate() + except BaseException: + _kill(process) + raise + logger.debug("%s: Verifier stdout: %s", what, stdout) + logger.debug("%s: Verifier stderr: %s", what, stderr) + if ("OK. inclusion check success!" in stdout or + "OK. inclusion check success!" in stderr): + logger.debug("%s: Verifier check passed.", what) return True else: - logger.debug("Verifier check failed.") + logger.debug("%s: Verifier check failed.", what) return False except FileNotFoundError: - logger.error("`%s` command not found.", verifier_executable) + logger.error("%s: `%s` command not found.", what, verifier_executable) return False except Exception as e: - logger.error("Error running verifier: %s", e) + logger.error("%s: Error running verifier: %s", what, e) return False +def _kill(process: subprocess.Popen) -> None: + """Kills and reaps a verifier. + + Popen's context manager does not wait for the process after a + KeyboardInterrupt, so wait here to not leave a zombie behind. + """ + process.kill() + process.wait() + + +def _verify_split(job: _SplitJob, verifier_executable: str, + logger: logging.Logger, cache_dir: Optional[str]) -> bool: + """Verifies one split with a verifier run of its own.""" + temp_payload_path = "" + try: + with tempfile.NamedTemporaryFile(mode="w", delete=False, + suffix=".txt") as fp: + fp.write(job.payload) + temp_payload_path = fp.name + + return run_verifier( + verifier_executable, + temp_payload_path, + logger, + cache_dir=cache_dir, + label=job.label) + finally: + if temp_payload_path and os.path.exists(temp_payload_path): + os.remove(temp_payload_path) + + +def _write_payloads_file(jobs: list) -> str: + """Writes the jobs' payloads to a new JSON Lines file; returns its path.""" + fd, path = tempfile.mkstemp(suffix=".jsonl") + try: + with os.fdopen(fd, "w") as fp: + for job in jobs: + fp.write(json.dumps({"payload": job.payload}) + "\n") + except BaseException: + os.remove(path) + raise + return path + + +def _verify_splits_batch(jobs: list, verifier_executable: str, + logger: logging.Logger, cache_dir: Optional[str], + report: Callable[[_SplitJob, bool], None]) -> list: + """Verifies jobs with a single verifier run in batch mode (--payloads_path). + + One run fetches each log's checkpoint and searches its entries once for + all splits, instead of once per split. Calls report(job, verified) as each + result arrives. If interrupted, the verifier is killed and the exception + propagates. + + Returns: + The jobs that got no result, in their original order: all of them if the + verifier does not support batch mode or could not be run, and the rest + if it stopped early or printed malformed results. Empty if the verifier + executable is missing or not executable: then every job is reported as + not verified, since running it per split would fail the same way. + """ + reported = set() + returncode = None + with contextlib.ExitStack() as cleanup: + try: + payloads_path = _write_payloads_file(jobs) + cleanup.callback(os.remove, payloads_path) + # stderr (the verifier's log) goes to a file, so that it cannot fill a + # pipe and block the verifier while stdout is read here. + stderr_file = cleanup.enter_context(tempfile.TemporaryFile(mode="w+")) + except OSError as e: + logger.warning("Could not prepare batch verification: %s", e) + return list(jobs) + + cmd = [verifier_executable, f"--payloads_path={payloads_path}", + "--log_type=google_1p_apk"] + if cache_dir: + cmd.append(f"--cache_dir={cache_dir}") + logger.debug("Running verifier in batch mode: %s", " ".join(cmd)) + try: + process = subprocess.Popen(cmd, stdout=subprocess.PIPE, + stderr=stderr_file, text=True) + except (FileNotFoundError, PermissionError) as e: + # E.g. a wrong --verifier_path. Running it once per split would fail + # the same way for each, so give every split its result now. + logger.error("Cannot run verifier `%s`: %s. Marking %d split(s) as not " + "verified.", verifier_executable, e, len(jobs)) + for job in jobs: + report(job, False) + return [] + except OSError as e: + logger.warning("Could not run verifier in batch mode: %s", e) + return list(jobs) + + with process: + try: + # Results are read as they arrive. In the main thread, signals + # interrupt this blocking read. + for line in process.stdout: + if not line.strip(): + continue + try: + result = json.loads(line) + index, verified = result["index"], result["verified"] + except (ValueError, KeyError, TypeError): + logger.warning("Ignoring malformed verifier result: %r", line) + continue + if (type(index) is not int or not 0 <= index < len(jobs) or + index in reported or type(verified) is not bool): + logger.warning("Ignoring unexpected verifier result: %r", line) + continue + reported.add(index) + job = jobs[index] + if result.get("error"): + logger.warning("%s: Inclusion proof failed: %s", job.label, + result["error"]) + logger.debug("%s: Verifier check %s.", job.label, + "passed" if verified else "failed") + report(job, verified) + returncode = process.wait() + except BaseException: + # E.g. Ctrl-C: do not leave the verifier behind. + _kill(process) + raise + stderr_file.seek(0) + stderr = stderr_file.read() + + logger.debug("Batch verifier stderr: %s", stderr) + remaining = [job for i, job in enumerate(jobs) if i not in reported] + if (returncode == _BATCH_UNSUPPORTED_EXIT_CODE and not reported and + _BATCH_UNSUPPORTED_MESSAGE in stderr): + # An older verifier, which rejects the unknown --payloads_path flag. + logger.warning("Verifier does not support batch mode; verifying %d " + "split(s) one at a time, which is slow. Rebuild the " + "verifier for faster checks.", len(jobs)) + elif remaining and returncode is not None: + logger.warning("Batch verifier exited with code %d after %d of %d " + "results; verifying the remaining %d split(s) one at a " + "time.", returncode, len(reported), len(jobs), + len(remaining)) + return remaining + + +def _verify_splits(jobs: list, verifier_executable: str, + logger: logging.Logger, cache_dir: Optional[str], + progress: Optional[ProgressCallback] = None) -> None: + """Verifies jobs, preferably with one batch-mode verifier run. + + Splits that the batch run gives no result for (e.g. with a verifier that + does not support batch mode) are then verified one at a time, with a + verifier run each. + + Stores each result as job.split["inclusion_proof_verified"]. + + If interrupted (KeyboardInterrupt, or an exception raised by a SIGTERM + handler), the running verifier is killed, no more are started, and the + exception propagates. + """ + total = len(jobs) + done_count = 0 + + def report(job: _SplitJob, verified: bool) -> None: + nonlocal done_count + job.split["inclusion_proof_verified"] = verified + done_count += 1 + if progress is not None: + progress(done_count, total) + + try: + if progress is not None: + progress(0, total) + remaining = jobs + if jobs: + remaining = _verify_splits_batch(jobs, verifier_executable, logger, + cache_dir, report) + for job in remaining: + report(job, _verify_split(job, verifier_executable, logger, cache_dir)) + except (KeyboardInterrupt, termination.Terminated): + logger.warning("Inclusion proof check interrupted after %d of %d " + "split(s).", done_count, total) + raise + + def perform_inclusion_proof_check( verifier_executable: str, packages_file_path: str, @@ -126,7 +359,8 @@ def perform_inclusion_proof_check( concurrency: int = DEFAULT_PREFETCH_CONCURRENCY, timeout: int = DEFAULT_PREFETCH_TIMEOUT, prefetch: bool = True, - preinstalled_only: Optional[bool] = None) -> bool: + preinstalled_only: Optional[bool] = None, + progress: Optional[ProgressCallback] = None) -> bool: """Reads packages.txt or preinstalled_packages.txt and performs inclusion proof check. By default, pre-fetches and locally caches transparency log entries before @@ -144,6 +378,13 @@ def perform_inclusion_proof_check( preinstalled_only: whether the input file is preinstalled_packages.txt (expects 'preinstalledPackages' key). If None, inferred from packages_file_path basename. + progress: optional callback, called in the calling thread as + progress(done, total): once with done=0 when verification + starts (after the input was read and pre-fetching, if any, + is over), then after each split is verified. total is the + number of splits to verify; skipped splits are not counted. + Not called if the input file cannot be read. Calls are not + throttled. Returns: True if the input file was valid and inclusion proof results were @@ -183,6 +424,10 @@ def perform_inclusion_proof_check( concurrency=concurrency, timeout=timeout) + # Collect the splits to verify first, then verify them (in one batch run + # if possible). Each job refers to its split's own dict, so the output + # keeps the input's package and split order. + jobs = [] for package in packages_list: if "name" not in package or "versionCode" not in package: logger.warning("Skipping package due to missing fields: %s", @@ -201,7 +446,8 @@ def perform_inclusion_proof_check( package_name) continue - for split in package["splits"]: + splits = package["splits"] + for index, split in enumerate(splits): split_hash = "" if "hash" in split: split_hash = split["hash"] @@ -216,22 +462,16 @@ def perform_inclusion_proof_check( split["hash"] = split_hash payload = f"{split_hash}\nSHA256(APK)\n{package_name}\n{version_code}\n" - temp_payload_path = "" - try: - with tempfile.NamedTemporaryFile(mode="w", delete=False, - suffix=".txt") as fp: - fp.write(payload) - temp_payload_path = fp.name - - verified = run_verifier( - verifier_executable, - temp_payload_path, - logger, - cache_dir=cache_dir) - split["inclusion_proof_verified"] = verified - finally: - if temp_payload_path and os.path.exists(temp_payload_path): - os.remove(temp_payload_path) + jobs.append(_SplitJob( + split=split, + payload=payload, + label=_split_label(package_name, split, index, len(splits)))) + + if jobs: + logger.info("Verifying %d APK split(s) in one batch verifier run...", + len(jobs)) + _verify_splits(jobs, verifier_executable, logger, cache_dir, + progress=progress) filtered_packages = [] for p in packages_list: @@ -316,14 +556,24 @@ def main(): s_handler.setFormatter(s_format) logger.addHandler(s_handler) - if not perform_inclusion_proof_check( - args.verifier_path, - args.packages_file, - logger, - cache_dir=args.cache_dir, - concurrency=args.cache_prefetch_concurrency, - timeout=args.cache_prefetch_timeout, - prefetch=not args.no_prefetch): + # Turn SIGTERM into an exception so that running verifiers are killed + # instead of being left behind, then die by SIGTERM as before. + terminated = False + with termination.sigterm_raises(): + try: + ok = perform_inclusion_proof_check( + args.verifier_path, + args.packages_file, + logger, + cache_dir=args.cache_dir, + concurrency=args.cache_prefetch_concurrency, + timeout=args.cache_prefetch_timeout, + prefetch=not args.no_prefetch) + except termination.Terminated: + terminated = True + if terminated: + termination.die_by_sigterm(logger) + if not ok: sys.exit(1) diff --git a/uraniborg/scripts/python/termination.py b/uraniborg/scripts/python/termination.py new file mode 100644 index 0000000000..ec97800 --- /dev/null +++ b/uraniborg/scripts/python/termination.py @@ -0,0 +1,83 @@ +#!/usr/bin/python3 +# Copyright 2026 Uraniborg authors. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + +"""Turns SIGTERM into an exception, so that scripts can clean up first. + +SIGTERM is how CI timeouts and cancellations usually stop a process. By +default it kills the script at once, leaving e.g. child verifiers behind or +an event stream without its last event. Usage: + + terminated = False + with termination.sigterm_raises(): + try: + ... + except termination.Terminated: + terminated = True # Clean up here. + if terminated: + termination.die_by_sigterm(logger) +""" + +import contextlib +import logging +import os +import signal +import sys + + +class Terminated(BaseException): + """Raised by the SIGTERM handler that sigterm_raises() installs. + + Derives from BaseException, like KeyboardInterrupt, so that the generic + `except Exception` handlers do not swallow it. SyscallWrapper only catches + KeyboardInterrupt, so this also propagates out of adb calls. + """ + + def __init__(self): + super().__init__("received SIGTERM") + + +def raise_terminated(signum, frame): # pylint: disable=unused-argument + """SIGTERM handler that raises Terminated.""" + # Ignore repeated SIGTERMs while cleaning up; die_by_sigterm() restores the + # default action afterwards. + signal.signal(signal.SIGTERM, signal.SIG_IGN) + raise Terminated() + + +@contextlib.contextmanager +def sigterm_raises(): + """Makes SIGTERM raise Terminated within the block. + + The previous handler is restored on leaving the block, however it is left. + """ + previous = signal.signal(signal.SIGTERM, raise_terminated) + try: + yield + finally: + signal.signal(signal.SIGTERM, previous) + + +def die_by_sigterm(logger: logging.Logger) -> None: + """Terminates this process with SIGTERM's default action. + + So the parent process sees the same exit status as without + sigterm_raises(). + """ + logger.error("Terminated by SIGTERM.") + signal.signal(signal.SIGTERM, signal.SIG_DFL) + os.kill(os.getpid(), signal.SIGTERM) + # Not reached unless SIGTERM is blocked; fall back to the conventional code. + sys.exit(128 + signal.SIGTERM) diff --git a/uraniborg/scripts/python/tests/test_automate_observation_events.py b/uraniborg/scripts/python/tests/test_automate_observation_events.py index 9eaf5f3..202b659 100644 --- a/uraniborg/scripts/python/tests/test_automate_observation_events.py +++ b/uraniborg/scripts/python/tests/test_automate_observation_events.py @@ -36,6 +36,7 @@ sys.path.insert(0, SCRIPT_DIR) import automate_observation +import termination # Shared fixture and helpers for driving main() with every collaborator mocked. from test_automate_observation import _make_mock_device # pylint: disable=g-importing-member from test_automate_observation import _set_argv # pylint: disable=g-importing-member @@ -153,6 +154,69 @@ def test_step_reports_finished_when_body_continues_a_loop(): ("loop", "started"), ("loop", "failed")] +class _FakeMonotonic: + + def __init__(self): + self.now = 100.0 + + def __call__(self): + return self.now + + +def _progress(events: list[dict]) -> list[tuple[int, int]]: + return [(e["done"], e["total"]) for e in _of_type(events, "step_progress")] + + +def test_progress_reporter_event_shape(): + stream = io.StringIO() + emitter = EventEmitter(stream, clock=lambda: _FIXED_TIME) + emitter.progress_reporter("inclusion_proof_check", "D1")(0, 5) + emitter.progress_reporter("other_step")(0, 2) + assert _parse(stream) == [ + {"v": 1, "ts": "2026-01-02T03:04:05.678Z", "type": "step_progress", + "step": "inclusion_proof_check", "device": "D1", "done": 0, + "total": 5}, + {"v": 1, "ts": "2026-01-02T03:04:05.678Z", "type": "step_progress", + "step": "other_step", "done": 0, "total": 2}, + ] + + +def test_progress_reporter_throttles_but_reports_first_and_last(): + stream = io.StringIO() + clock = _FakeMonotonic() + emitter = EventEmitter(stream, clock=lambda: _FIXED_TIME, monotonic=clock) + report = emitter.progress_reporter("s", "D1", interval=2.0) + + report(0, 10) # first: reported + report(1, 10) # too soon + clock.now += 1.9 + report(2, 10) # still too soon + clock.now += 0.1 + report(3, 10) # 2 s after the last event: reported + report(4, 10) # too soon + clock.now += 5 + report(5, 10) # reported + report(5, 10) # same done again: never reported + report(10, 10) # last: always reported, even right away + clock.now += 5 + report(10, 10) # repeat of the last: not reported + + assert _progress(_parse(stream)) == [(0, 10), (3, 10), (5, 10), (10, 10)] + + +def test_progress_reporter_empty_step_reports_once(): + stream = io.StringIO() + report = _emitter(stream).progress_reporter("s") + report(0, 0) + assert _progress(_parse(stream)) == [(0, 0)] + + +def test_progress_reporter_without_stream_is_a_noop(): + report = EventEmitter().progress_reporter("s", "D1") + report(0, 3) + report(3, 3) # must not raise + + def test_finish_run_is_emitted_once_with_ok_semantics(): stream = io.StringIO() emitter = _emitter(stream) @@ -370,6 +434,52 @@ def test_main_events_inclusion_proof_outcomes( "message": "Inclusion proof check could not complete."} +def test_main_events_inclusion_proof_progress_inside_step( + serial_main_mocks, monkeypatch: pytest.MonkeyPatch, tmp_path: Path, +): + m = serial_main_mocks + m["AdbWrapper"].devices.return_value = [_make_mock_device("DEV1")] + events_path = tmp_path / "events.jsonl" + _set_argv(monkeypatch, "--events", str(events_path), + "--perform_inclusion_proof_check", "--verifier_path=/v", + "--no_prefetch") + clock = _FakeMonotonic() + real_emitter = automate_observation.EventEmitter + + def emitter_with_fake_clock(*args, **kwargs): + return real_emitter(*args, monotonic=clock, **kwargs) + + def fake_check(*args, progress, **kwargs): + for done in range(5): + progress(done, 4) + clock.now += 1.5 # 1.5 s per split: every other one is reported + return True + + with mock.patch("automate_observation.os.path.isfile", return_value=True), \ + mock.patch("automate_observation.EventEmitter", + side_effect=emitter_with_fake_clock), \ + mock.patch("inclusion_proof_check.perform_inclusion_proof_check", + side_effect=fake_check): + automate_observation.main() + + events = _read_events(events_path) + check = [i for i, e in enumerate(events) + if e.get("step") == "inclusion_proof_check"] + started, finished = check[0], check[-1] + assert events[started]["type"] == "step" + assert events[started]["state"] == "started" + progress = events[started + 1:finished] + assert [(e["type"], e["step"], e["device"], e["done"], e["total"]) + for e in progress] == [ + ("step_progress", "inclusion_proof_check", "DEV1", 0, 4), + ("step_progress", "inclusion_proof_check", "DEV1", 2, 4), + ("step_progress", "inclusion_proof_check", "DEV1", 4, 4), + ] + assert events[finished]["state"] == "finished" + # No progress events outside the step. + assert len(_of_type(events, "step_progress")) == 3 + + def test_main_events_check_preinstalled_only_missing_file( serial_main_mocks, monkeypatch: pytest.MonkeyPatch, tmp_path: Path, ): @@ -675,7 +785,7 @@ def _wait(*unused_args): _set_argv(monkeypatch, "--events", str(tmp_path / "events.jsonl")) automate_observation.main() - assert seen == [before, automate_observation._raise_terminated] + assert seen == [before, termination.raise_terminated] assert signal.getsignal(signal.SIGTERM) == before @@ -684,12 +794,12 @@ def test_main_events_terminated_in_process( ): """Terminated -> run_finished{terminated}, then die by SIGTERM (patched).""" m = serial_main_mocks - _two_devices_interrupted_on_second(m, automate_observation.Terminated()) + _two_devices_interrupted_on_second(m, termination.Terminated()) events_path = tmp_path / "events.jsonl" _set_argv(monkeypatch, "--events", str(events_path)) before = signal.getsignal(signal.SIGTERM) - with mock.patch("automate_observation._die_by_sigterm") as die: + with mock.patch("termination.die_by_sigterm") as die: automate_observation.main() die.assert_called_once_with(mock.ANY) assert signal.getsignal(signal.SIGTERM) == before @@ -902,7 +1012,7 @@ def test_prompt_outcome_set_by_body(body, expected_outcome): "exc, expected_outcome", [ (KeyboardInterrupt(), "interrupted"), - (automate_observation.Terminated(), "terminated"), + (termination.Terminated(), "terminated"), (RuntimeError("boom"), "unexpected_error"), ], ids=["ctrl_c", "sigterm", "exception"], diff --git a/uraniborg/scripts/python/tests/test_inclusion_proof_check.py b/uraniborg/scripts/python/tests/test_inclusion_proof_check.py index 75dc5da..dc54619 100644 --- a/uraniborg/scripts/python/tests/test_inclusion_proof_check.py +++ b/uraniborg/scripts/python/tests/test_inclusion_proof_check.py @@ -20,8 +20,12 @@ import logging import os from pathlib import Path +import signal import subprocess import sys +import textwrap +import threading +import time from unittest import mock import pytest @@ -30,6 +34,11 @@ import inclusion_proof_check +_SCRIPT = os.path.abspath(os.path.join( + os.path.dirname(__file__), "..", "inclusion_proof_check.py")) +_OK = (0, "OK. inclusion check success!", "") +_NOT_FOUND = (1, "", "inclusion check failed") + @pytest.fixture def logger() -> logging.Logger: @@ -38,6 +47,123 @@ def logger() -> logging.Logger: return log +_UNSUPPORTED_BATCH = (2, [], "flag provided but not defined: -payloads_path\n") + + +def _is_batch(cmd) -> bool: + return any(arg.startswith("--payloads_path=") for arg in cmd) + + +class _FakeProcess: + """subprocess.Popen stand-in for one verifier run.""" + + def __init__(self, cmd, result_for, batch_for, **kwargs): + self.cmd = cmd + self.returncode = None + self.killed = False + self._result_for = result_for + if _is_batch(cmd): + self._batch_returncode, lines, stderr = batch_for(cmd) + self.stdout = iter(lines) + kwargs["stderr"].write(stderr) + + def __enter__(self): + return self + + def __exit__(self, *exc_info): + return False + + def communicate(self, input=None, timeout=None): # pylint: disable=redefined-builtin + self.returncode, stdout, stderr = self._result_for(self.cmd) + return stdout, stderr + + def poll(self): + return self.returncode + + def wait(self): + if self.returncode is None: + self.returncode = self._batch_returncode + return self.returncode + + def kill(self): + self.killed = True + self.returncode = -signal.SIGKILL + + +def _fake_popen(result_for, batch_for=lambda cmd: _UNSUPPORTED_BATCH): + """Returns a subprocess.Popen stand-in for verifier runs. + + Args: + result_for: for a single-payload run, called with the command when the + process is waited for (in communicate(), like a real + process's work); returns (returncode, stdout, stderr). + batch_for: for a batch run (--payloads_path), called with the command when + the process starts; returns (returncode, stdout lines, stderr). + By default, the verifier does not support batch mode. + """ + def factory(cmd, **kwargs): + return _FakeProcess(cmd, result_for, batch_for, **kwargs) + return mock.Mock(side_effect=factory) + + +def _split_cmds(popen) -> list: + """The single-payload verifier commands a fake Popen was called with.""" + return [c.args[0] for c in popen.call_args_list if not _is_batch(c.args[0])] + + +def _batch_cmds(popen) -> list: + return [c.args[0] for c in popen.call_args_list if _is_batch(c.args[0])] + + +def _batch_payloads(cmd) -> list: + """Reads the payloads of a batch verifier command.""" + prefix = "--payloads_path=" + path = next(arg[len(prefix):] for arg in cmd if arg.startswith(prefix)) + return [json.loads(line)["payload"] + for line in Path(path).read_text().splitlines()] + + +def _batch_result(index, verified, **extra) -> str: + return json.dumps(dict(index=index, verified=verified, **extra)) + "\n" + + +def _payload_of(cmd) -> str: + """Reads the payload file a verifier command refers to.""" + prefix = "--payload_path=" + path = next(arg[len(prefix):] for arg in cmd if arg.startswith(prefix)) + return Path(path).read_text() + + +def _write_packages(path: Path, packages, key="packages") -> Path: + path.write_text(json.dumps({key: packages})) + return path + + +def _write_fake_verifier(tmp_path: Path, body: str) -> Path: + """Writes an executable Python script that stands in for the verifier.""" + path = tmp_path / "fake_verifier" + path.write_text("#!{}\n".format(sys.executable) + textwrap.dedent(body)) + path.chmod(0o755) + return path + + +def _wait_for(predicate, timeout=10.0): + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + if predicate(): + return True + time.sleep(0.05) + return False + + +def _process_exists(pid: int) -> bool: + try: + os.kill(pid, 0) + except ProcessLookupError: + return False + return True + + @mock.patch("subprocess.run") def test_prefetch_log_entries_default( mock_run: mock.MagicMock, logger: logging.Logger @@ -124,83 +250,94 @@ def test_prefetch_log_entries_file_not_found( mock_error.assert_called_once() -@mock.patch("subprocess.run") -def test_run_verifier_with_cache_dir( - mock_run: mock.MagicMock, tmp_path: Path, logger: logging.Logger -): - mock_run.return_value = mock.Mock( - returncode=0, - stdout="OK. inclusion check success!", - stderr="", - ) +def test_run_verifier_with_cache_dir(tmp_path: Path, logger: logging.Logger): payload_file = tmp_path / "payload.txt" payload_file.write_text("hash\nSHA256(APK)\ncom.example\n1\n") - verified = inclusion_proof_check.run_verifier( - "/path/to/verifier", - str(payload_file), - logger, - cache_dir="/custom/cache", - ) + with mock.patch("subprocess.Popen", _fake_popen(lambda cmd: _OK)) as popen: + verified = inclusion_proof_check.run_verifier( + "/path/to/verifier", + str(payload_file), + logger, + cache_dir="/custom/cache", + ) assert verified is True - mock_run.assert_called_once_with( + popen.assert_called_once_with( [ "/path/to/verifier", f"--payload_path={payload_file}", "--log_type=google_1p_apk", "--cache_dir=/custom/cache", ], - capture_output=True, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, text=True, - check=False, ) +def test_run_verifier_names_split_in_log_lines( + tmp_path: Path, logger: logging.Logger, caplog: pytest.LogCaptureFixture +): + payload_file = tmp_path / "payload.txt" + payload_file.write_text("hash\nSHA256(APK)\ncom.example\n1\n") + + with caplog.at_level(logging.DEBUG, logger=logger.name): + with mock.patch("subprocess.Popen", _fake_popen(lambda cmd: _NOT_FOUND)): + verified = inclusion_proof_check.run_verifier( + "/path/to/verifier", str(payload_file), logger, + label="com.example [config.en]") + + assert verified is False + assert caplog.records + for record in caplog.records: + assert record.getMessage().startswith("com.example [config.en]: ") + + +def test_run_verifier_file_not_found_is_failure( + tmp_path: Path, logger: logging.Logger +): + payload_file = tmp_path / "payload.txt" + payload_file.write_text("payload") + + with mock.patch("subprocess.Popen", side_effect=FileNotFoundError()): + assert inclusion_proof_check.run_verifier( + "/bad/verifier", str(payload_file), logger) is False + + @mock.patch("subprocess.run") def test_perform_inclusion_proof_check_with_prefetch( mock_run: mock.MagicMock, tmp_path: Path, logger: logging.Logger ): - def side_effect(cmd, **kwargs): - if "--fetch_entries" in cmd: - return mock.Mock(returncode=0, stdout="prefetched", stderr="") - return mock.Mock( - returncode=0, - stdout="OK. inclusion check success!", - stderr="", + mock_run.return_value = mock.Mock(returncode=0, stdout="prefetched", stderr="") + + packages_file = _write_packages(tmp_path / "packages.txt", [{ + "name": "com.google.android.gm", + "versionCode": 123, + "splits": [{"hash": "abc123hash"}], + }]) + + with mock.patch("subprocess.Popen", _fake_popen(lambda cmd: _OK)) as popen: + success = inclusion_proof_check.perform_inclusion_proof_check( + "/path/to/verifier", + str(packages_file), + logger, + cache_dir="/tmp/cache", + concurrency=32, + timeout=45, + prefetch=True, ) - mock_run.side_effect = side_effect - - packages_file = tmp_path / "packages.txt" - packages_file.write_text(json.dumps({ - "packages": [{ - "name": "com.google.android.gm", - "versionCode": 123, - "splits": [{"hash": "abc123hash"}], - }] - })) - - success = inclusion_proof_check.perform_inclusion_proof_check( - "/path/to/verifier", - str(packages_file), - logger, - cache_dir="/tmp/cache", - concurrency=32, - timeout=45, - prefetch=True, - ) - assert success is True - assert mock_run.call_count == 2 - + mock_run.assert_called_once() prefetch_cmd = mock_run.call_args_list[0][0][0] assert "--fetch_entries" in prefetch_cmd assert "--concurrency=32" in prefetch_cmd assert "--cache_dir=/tmp/cache" in prefetch_cmd assert mock_run.call_args_list[0].kwargs["timeout"] == 45 - verify_cmd = mock_run.call_args_list[1][0][0] + assert len(_split_cmds(popen)) == 1 + verify_cmd = _split_cmds(popen)[0] assert "--log_type=google_1p_apk" in verify_cmd assert "--cache_dir=/tmp/cache" in verify_cmd @@ -214,32 +351,25 @@ def side_effect(cmd, **kwargs): def test_perform_inclusion_proof_check_prefetch_disabled( mock_run: mock.MagicMock, tmp_path: Path, logger: logging.Logger ): - mock_run.return_value = mock.Mock( - returncode=0, - stdout="OK. inclusion check success!", - stderr="", - ) - - packages_file = tmp_path / "packages.txt" - packages_file.write_text(json.dumps({ - "packages": [{ - "name": "com.google.android.gm", - "versionCode": 123, - "splits": [{"hash": "abc123hash"}], - }] - })) - - success = inclusion_proof_check.perform_inclusion_proof_check( - "/path/to/verifier", - str(packages_file), - logger, - cache_dir="/tmp/cache", - prefetch=False, - ) + packages_file = _write_packages(tmp_path / "packages.txt", [{ + "name": "com.google.android.gm", + "versionCode": 123, + "splits": [{"hash": "abc123hash"}], + }]) + + with mock.patch("subprocess.Popen", _fake_popen(lambda cmd: _OK)) as popen: + success = inclusion_proof_check.perform_inclusion_proof_check( + "/path/to/verifier", + str(packages_file), + logger, + cache_dir="/tmp/cache", + prefetch=False, + ) assert success is True - assert mock_run.call_count == 1 - verify_cmd = mock_run.call_args_list[0][0][0] + mock_run.assert_not_called() + assert len(_split_cmds(popen)) == 1 + verify_cmd = _split_cmds(popen)[0] assert "--fetch_entries" not in verify_cmd assert "--log_type=google_1p_apk" in verify_cmd @@ -253,35 +383,26 @@ def test_perform_inclusion_proof_check_prefetch_disabled( def test_perform_inclusion_proof_check_fail_open_on_prefetch_failure( mock_run: mock.MagicMock, tmp_path: Path, logger: logging.Logger ): - def side_effect(cmd, **kwargs): - if "--fetch_entries" in cmd: - return mock.Mock(returncode=1, stdout="", stderr="prefetch network timeout") - return mock.Mock( - returncode=0, - stdout="OK. inclusion check success!", - stderr="", + mock_run.return_value = mock.Mock( + returncode=1, stdout="", stderr="prefetch network timeout") + + packages_file = _write_packages(tmp_path / "packages.txt", [{ + "name": "com.google.android.gm", + "versionCode": 123, + "splits": [{"hash": "abc123hash"}], + }]) + + with mock.patch("subprocess.Popen", _fake_popen(lambda cmd: _OK)) as popen: + success = inclusion_proof_check.perform_inclusion_proof_check( + "/path/to/verifier", + str(packages_file), + logger, + prefetch=True, ) - mock_run.side_effect = side_effect - - packages_file = tmp_path / "packages.txt" - packages_file.write_text(json.dumps({ - "packages": [{ - "name": "com.google.android.gm", - "versionCode": 123, - "splits": [{"hash": "abc123hash"}], - }] - })) - - success = inclusion_proof_check.perform_inclusion_proof_check( - "/path/to/verifier", - str(packages_file), - logger, - prefetch=True, - ) - assert success is True - assert mock_run.call_count == 2 + mock_run.assert_called_once() + assert len(_split_cmds(popen)) == 1 output_file = tmp_path / inclusion_proof_check.OUTPUT_FILENAME assert output_file.exists() @@ -338,16 +459,7 @@ def test_perform_inclusion_proof_check_empty_packages_list_skips_prefetch_and_su def test_perform_inclusion_proof_check_with_preinstalled_packages_and_metadata( mock_run: mock.MagicMock, tmp_path: Path, logger: logging.Logger ): - def side_effect(cmd, **kwargs): - if "--fetch_entries" in cmd: - return mock.Mock(returncode=0, stdout="prefetched", stderr="") - return mock.Mock( - returncode=0, - stdout="OK. inclusion check success!", - stderr="", - ) - - mock_run.side_effect = side_effect + mock_run.return_value = mock.Mock(returncode=0, stdout="prefetched", stderr="") # Also create an existing full-run output file to verify it is NOT overwritten existing_full_output = tmp_path / inclusion_proof_check.OUTPUT_FILENAME @@ -381,13 +493,14 @@ def side_effect(cmd, **kwargs): ], })) - success = inclusion_proof_check.perform_inclusion_proof_check( - "/path/to/verifier", - str(preinstalled_file), - logger, - prefetch=True, - preinstalled_only=True, - ) + with mock.patch("subprocess.Popen", _fake_popen(lambda cmd: _OK)): + success = inclusion_proof_check.perform_inclusion_proof_check( + "/path/to/verifier", + str(preinstalled_file), + logger, + prefetch=True, + preinstalled_only=True, + ) assert success is True # Full-run output artifact remains intact @@ -414,6 +527,529 @@ def side_effect(cmd, **kwargs): assert pkg1["splits"][0]["inclusion_proof_verified"] is True +def _many_packages(): + """Packages with several splits, in a deliberately non-sorted order.""" + return [ + {"name": "com.z.last", "versionCode": 3, "splits": [ + {"name": "base", "hash": "z0"}, + {"name": "config.en", "hash": "z1-ok"}, + {"name": "config.xxhdpi", "hash": "z2"}, + ]}, + {"name": "com.a.first", "versionCode": 1, "hash": "a-ok"}, + {"name": "com.m.middle", "versionCode": 2, "splits": [ + {"hash": "m0-ok"}, + {"hash": "m1"}, + ]}, + {"name": "com.no.hash", "versionCode": 4}, + ] + + +def test_perform_inclusion_proof_check_one_at_a_time_results_and_order( + tmp_path: Path, logger: logging.Logger +): + packages_file = _write_packages(tmp_path / "packages.txt", _many_packages()) + + def result_for(cmd): + split_hash = _payload_of(cmd).split("\n")[0] + # Finish in a scrambled order, so results complete out of order. + time.sleep((hash(split_hash) % 5) / 500) + return _OK if split_hash.endswith("-ok") else _NOT_FOUND + + with mock.patch("subprocess.Popen", _fake_popen(result_for)) as popen: + success = inclusion_proof_check.perform_inclusion_proof_check( + "/path/to/verifier", str(packages_file), logger, prefetch=False) + + assert success is True + assert len(_split_cmds(popen)) == 6 + result_json = json.loads( + (tmp_path / inclusion_proof_check.OUTPUT_FILENAME).read_text()) + assert [ + (p["name"], [(s["hash"], s["inclusion_proof_verified"]) + for s in p["splits"]]) + for p in result_json["packages"] + ] == [ + ("com.z.last", [("z0", False), ("z1-ok", True), ("z2", False)]), + ("com.a.first", [("a-ok", True)]), + ("com.m.middle", [("m0-ok", True), ("m1", False)]), + ] + + +def test_perform_inclusion_proof_check_reports_progress( + tmp_path: Path, logger: logging.Logger +): + packages_file = _write_packages(tmp_path / "packages.txt", _many_packages()) + calls = [] + threads = set() + + def progress(done, total): + calls.append((done, total)) + threads.add(threading.get_ident()) + + with mock.patch("subprocess.Popen", _fake_popen(lambda cmd: _OK)): + assert inclusion_proof_check.perform_inclusion_proof_check( + "/path/to/verifier", str(packages_file), logger, prefetch=False, progress=progress) + + # Skipped packages (com.no.hash) are not counted. + assert calls == [(done, 6) for done in range(7)] + assert threads == {threading.get_ident()} # Always the calling thread. + + +@mock.patch("subprocess.run") +def test_perform_inclusion_proof_check_progress_starts_after_prefetch( + mock_run: mock.MagicMock, tmp_path: Path, logger: logging.Logger +): + order = [] + mock_run.side_effect = lambda *a, **k: order.append("prefetch") or mock.Mock( + returncode=0) + packages_file = _write_packages(tmp_path / "packages.txt", _many_packages()) + + with mock.patch("subprocess.Popen", _fake_popen(lambda cmd: _OK)): + inclusion_proof_check.perform_inclusion_proof_check( + "/path/to/verifier", str(packages_file), logger, prefetch=True, + progress=lambda done, total: order.append((done, total))) + + assert order[:2] == ["prefetch", (0, 6)] + + +def test_perform_inclusion_proof_check_progress_empty_and_invalid_input( + tmp_path: Path, logger: logging.Logger +): + calls = [] + empty = _write_packages(tmp_path / "packages.txt", []) + assert inclusion_proof_check.perform_inclusion_proof_check( + "/v", str(empty), logger, prefetch=False, + progress=lambda *a: calls.append(a)) + assert calls == [(0, 0)] + + calls.clear() + invalid = tmp_path / "bad" / "packages.txt" + invalid.parent.mkdir() + invalid.write_text("{not json") + assert not inclusion_proof_check.perform_inclusion_proof_check( + "/v", str(invalid), logger, prefetch=False, + progress=lambda *a: calls.append(a)) + assert not inclusion_proof_check.perform_inclusion_proof_check( + "/v", str(tmp_path / "missing.txt"), logger, prefetch=False, + progress=lambda *a: calls.append(a)) + assert calls == [] + + +def test_perform_inclusion_proof_check_gives_each_split_its_own_payload( + tmp_path: Path, logger: logging.Logger +): + packages_file = _write_packages(tmp_path / "packages.txt", _many_packages()) + seen = {} + lock = threading.Lock() + + def result_for(cmd): + path = next(a for a in cmd if a.startswith("--payload_path=")) + payload = _payload_of(cmd) + time.sleep(0.02) + # The file still holds this split's payload after other workers ran. + assert _payload_of(cmd) == payload + with lock: + seen[path] = payload + return _OK + + with mock.patch("subprocess.Popen", _fake_popen(result_for)): + assert inclusion_proof_check.perform_inclusion_proof_check( + "/path/to/verifier", str(packages_file), logger, prefetch=False) + + assert len(seen) == 6 # One temp file per split... + assert sorted(seen.values()) == sorted([ + "z0\nSHA256(APK)\ncom.z.last\n3\n", + "z1-ok\nSHA256(APK)\ncom.z.last\n3\n", + "z2\nSHA256(APK)\ncom.z.last\n3\n", + "a-ok\nSHA256(APK)\ncom.a.first\n1\n", + "m0-ok\nSHA256(APK)\ncom.m.middle\n2\n", + "m1\nSHA256(APK)\ncom.m.middle\n2\n", + ]) + for path in seen: # ...removed afterwards. + assert not os.path.exists(path[len("--payload_path="):]) + + +def test_perform_inclusion_proof_check_labels_name_package_and_split( + tmp_path: Path, logger: logging.Logger, caplog: pytest.LogCaptureFixture +): + packages_file = _write_packages(tmp_path / "packages.txt", _many_packages()) + + with caplog.at_level(logging.DEBUG, logger=logger.name): + with mock.patch("subprocess.Popen", _fake_popen(lambda cmd: _OK)): + inclusion_proof_check.perform_inclusion_proof_check( + "/path/to/verifier", str(packages_file), logger, prefetch=False) + + passed = sorted( + r.getMessage()[:-len(": Verifier check passed.")] + for r in caplog.records + if r.getMessage().endswith(": Verifier check passed.")) + assert passed == [ + "com.a.first [base]", + "com.m.middle [split 0]", + "com.m.middle [split 1]", + "com.z.last [base]", + "com.z.last [config.en]", + "com.z.last [config.xxhdpi]", + ] + + +_HANGING_VERIFIER = """ + import os, sys, time + if any(a.startswith("--payloads_path=") for a in sys.argv): + # No batch mode: what Go's flag package does with an unknown flag. + print("flag provided but not defined: -payloads_path", file=sys.stderr) + sys.exit(2) + pid_dir = os.environ["FAKE_VERIFIER_PID_DIR"] + open(os.path.join(pid_dir, str(os.getpid())), "w").close() + time.sleep(60) +""" + +# Supports batch mode: reports the first payload, then hangs. +_HANGING_BATCH_VERIFIER = """ + import os, sys, time + pid_dir = os.environ["FAKE_VERIFIER_PID_DIR"] + print('{"index": 0, "verified": true}', flush=True) + open(os.path.join(pid_dir, str(os.getpid())), "w").close() + time.sleep(60) +""" + + +def test_interrupt_kills_running_verifier_and_starts_no_more( + tmp_path: Path, logger: logging.Logger, monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture +): + pid_dir = tmp_path / "pids" + pid_dir.mkdir() + monkeypatch.setenv("FAKE_VERIFIER_PID_DIR", str(pid_dir)) + verifier = _write_fake_verifier(tmp_path, _HANGING_VERIFIER) + packages_file = _write_packages(tmp_path / "packages.txt", [ + {"name": "com.example.{}".format(i), "versionCode": 1, + "hash": "h{}".format(i)} + for i in range(6) + ]) + + def interrupt_when_started(): + if _wait_for(lambda: os.listdir(pid_dir)): + os.kill(os.getpid(), signal.SIGINT) + + threading.Thread(target=interrupt_when_started, daemon=True).start() + started = time.monotonic() + with pytest.raises(KeyboardInterrupt): + inclusion_proof_check.perform_inclusion_proof_check( + str(verifier), str(packages_file), logger, prefetch=False) + + assert time.monotonic() - started < 20 + pids = [int(name) for name in os.listdir(pid_dir)] + # Pending splits were never started. + assert len(pids) == 1 + for pid in pids: + assert _wait_for(lambda: not _process_exists(pid), timeout=5) + assert not (tmp_path / inclusion_proof_check.OUTPUT_FILENAME).exists() + assert "Inclusion proof check interrupted after 0 of 6 split(s)." in [ + r.getMessage() for r in caplog.records] + + +@pytest.mark.parametrize("verifier_body", [ + _HANGING_VERIFIER, _HANGING_BATCH_VERIFIER], + ids=["one_at_a_time", "batch"]) +def test_main_sigterm_kills_running_verifier( + tmp_path: Path, verifier_body: str +): + pid_dir = tmp_path / "pids" + pid_dir.mkdir() + verifier = _write_fake_verifier(tmp_path, verifier_body) + packages_file = _write_packages(tmp_path / "packages.txt", [ + {"name": "com.example.{}".format(i), "versionCode": 1, + "hash": "h{}".format(i)} + for i in range(4) + ]) + + process = subprocess.Popen( + [sys.executable, _SCRIPT, "--packages_file", str(packages_file), + "--verifier_path", str(verifier), "--no_prefetch"], + env=dict(os.environ, FAKE_VERIFIER_PID_DIR=str(pid_dir)), + stdout=subprocess.DEVNULL, stderr=subprocess.PIPE, text=True) + try: + assert _wait_for(lambda: os.listdir(pid_dir)) + time.sleep(0.2) # Let the batch result be read, if any. + process.send_signal(signal.SIGTERM) + _, stderr = process.communicate(timeout=20) + finally: + if process.poll() is None: + process.kill() + + assert process.returncode == -signal.SIGTERM + assert "Terminated by SIGTERM." in stderr + pids = [int(name) for name in os.listdir(pid_dir)] + assert len(pids) == 1 # No second verifier was started. + for pid in pids: + assert _wait_for(lambda: not _process_exists(pid), timeout=5) + assert not (tmp_path / inclusion_proof_check.OUTPUT_FILENAME).exists() + + +def test_interrupt_kills_batch_verifier( + tmp_path: Path, logger: logging.Logger, monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture +): + pid_dir = tmp_path / "pids" + pid_dir.mkdir() + monkeypatch.setenv("FAKE_VERIFIER_PID_DIR", str(pid_dir)) + verifier = _write_fake_verifier(tmp_path, _HANGING_BATCH_VERIFIER) + packages_file = _write_packages(tmp_path / "packages.txt", [ + {"name": "com.example.{}".format(i), "versionCode": 1, + "hash": "h{}".format(i)} + for i in range(3) + ]) + calls = [] + + def interrupt_when_started(): + if _wait_for(lambda: os.listdir(pid_dir) and len(calls) >= 2): + os.kill(os.getpid(), signal.SIGINT) + + threading.Thread(target=interrupt_when_started, daemon=True).start() + started = time.monotonic() + with pytest.raises(KeyboardInterrupt): + inclusion_proof_check.perform_inclusion_proof_check( + str(verifier), str(packages_file), logger, prefetch=False, + progress=lambda *a: calls.append(a)) + + assert time.monotonic() - started < 20 + # The result printed before the interrupt was reported as it arrived. + assert calls == [(0, 3), (1, 3)] + pids = [int(name) for name in os.listdir(pid_dir)] + assert len(pids) == 1 # No per-split verifiers were started. + assert _wait_for(lambda: not _process_exists(pids[0]), timeout=5) + assert "Inclusion proof check interrupted after 1 of 3 split(s)." in [ + r.getMessage() for r in caplog.records] + + +def _results_by_hash(tmp_path: Path) -> dict: + result_json = json.loads( + (tmp_path / inclusion_proof_check.OUTPUT_FILENAME).read_text()) + return {s["hash"]: s["inclusion_proof_verified"] + for p in result_json["packages"] for s in p["splits"]} + + +_MANY_PACKAGES_RESULTS = { + "z0": False, "z1-ok": True, "z2": False, "a-ok": True, "m0-ok": True, + "m1": False, +} + + +def test_perform_inclusion_proof_check_batch_mode( + tmp_path: Path, logger: logging.Logger +): + packages_file = _write_packages(tmp_path / "packages.txt", _many_packages()) + seen = {} + calls = [] + + def batch_for(cmd): + payloads = _batch_payloads(cmd) + seen["payloads"] = payloads + # Results arrive in any order. + lines = [_batch_result(i, p.split("\n")[0].endswith("-ok")) + for i, p in reversed(list(enumerate(payloads)))] + return 0, lines, "INFO Verified payloads\n" + + with mock.patch("subprocess.Popen", + _fake_popen(lambda cmd: pytest.fail("per-split run"), + batch_for)) as popen: + assert inclusion_proof_check.perform_inclusion_proof_check( + "/path/to/verifier", str(packages_file), logger, prefetch=False, + cache_dir="/tmp/cache", progress=lambda *a: calls.append(a)) + + (cmd,) = _batch_cmds(popen) + assert _split_cmds(popen) == [] + assert "--log_type=google_1p_apk" in cmd + assert "--cache_dir=/tmp/cache" in cmd + # Payloads in input order, same format as per-split payload files. + assert seen["payloads"] == [ + "z0\nSHA256(APK)\ncom.z.last\n3\n", + "z1-ok\nSHA256(APK)\ncom.z.last\n3\n", + "z2\nSHA256(APK)\ncom.z.last\n3\n", + "a-ok\nSHA256(APK)\ncom.a.first\n1\n", + "m0-ok\nSHA256(APK)\ncom.m.middle\n2\n", + "m1\nSHA256(APK)\ncom.m.middle\n2\n", + ] + payloads_path = next(a for a in cmd if a.startswith("--payloads_path=")) + assert not os.path.exists(payloads_path[len("--payloads_path="):]) + assert _results_by_hash(tmp_path) == _MANY_PACKAGES_RESULTS + assert calls == [(done, 6) for done in range(7)] + + +def test_perform_inclusion_proof_check_falls_back_without_batch_mode( + tmp_path: Path, logger: logging.Logger, caplog: pytest.LogCaptureFixture +): + packages_file = _write_packages(tmp_path / "packages.txt", _many_packages()) + + def result_for(cmd): + return _OK if _payload_of(cmd).split("\n")[0].endswith("-ok") else _NOT_FOUND + + with caplog.at_level(logging.INFO, logger=logger.name): + with mock.patch("subprocess.Popen", _fake_popen(result_for)) as popen: + assert inclusion_proof_check.perform_inclusion_proof_check( + "/path/to/verifier", str(packages_file), logger, prefetch=False) + + assert len(_batch_cmds(popen)) == 1 + assert len(_split_cmds(popen)) == 6 + assert _results_by_hash(tmp_path) == _MANY_PACKAGES_RESULTS + messages = [r.getMessage() for r in caplog.records] + # Slow, and fixed by rebuilding the verifier: worth a warning. + assert ("Verifier does not support batch mode; verifying 6 split(s) one at " + "a time, which is slow. Rebuild the verifier for faster checks." + in [r.getMessage() for r in caplog.records + if r.levelno == logging.WARNING]) + + +@pytest.mark.parametrize("returncode", [0, 1]) +def test_perform_inclusion_proof_check_batch_leftovers_verified_per_split( + tmp_path: Path, logger: logging.Logger, caplog: pytest.LogCaptureFixture, + returncode: int +): + packages_file = _write_packages(tmp_path / "packages.txt", _many_packages()) + calls = [] + + def batch_for(cmd): + return returncode, [ + _batch_result(1, True), # z1-ok + "not json\n", + _batch_result(1, False), # Duplicate: ignored. + _batch_result(99, True), # Out of range: ignored. + _batch_result(True, True), # Not an index: ignored. + json.dumps({"index": 2}) + "\n", # No verdict: ignored. + _batch_result(5, False, error="inclusion check error"), # m1 + "\n", + ], "verifier log\n" + + per_split = [] + lock = threading.Lock() + + def result_for(cmd): + split_hash = _payload_of(cmd).split("\n")[0] + with lock: + per_split.append(split_hash) + return _OK if split_hash.endswith("-ok") else _NOT_FOUND + + with caplog.at_level(logging.DEBUG, logger=logger.name): + with mock.patch("subprocess.Popen", _fake_popen(result_for, batch_for)): + assert inclusion_proof_check.perform_inclusion_proof_check( + "/path/to/verifier", str(packages_file), logger, prefetch=False, + progress=lambda *a: calls.append(a)) + + # Only splits without a batch result were verified one by one. + assert sorted(per_split) == ["a-ok", "m0-ok", "z0", "z2"] + assert _results_by_hash(tmp_path) == _MANY_PACKAGES_RESULTS + assert [c[0] for c in calls] == list(range(7)) + assert {c[1] for c in calls} == {6} + + warnings = [r.getMessage() for r in caplog.records + if r.levelno == logging.WARNING] + assert "com.m.middle [split 1]: Inclusion proof failed: inclusion check error" in warnings + assert len([w for w in warnings if w.startswith("Ignoring ")]) == 5 + assert ("Batch verifier exited with code {} after 2 of 6 results; verifying " + "the remaining 4 split(s) one at a time.".format(returncode) + in warnings) + assert "Batch verifier stderr: verifier log\n" in [ + r.getMessage() for r in caplog.records] + + +def test_perform_inclusion_proof_check_batch_crash_is_not_missing_batch_mode( + tmp_path: Path, logger: logging.Logger, caplog: pytest.LogCaptureFixture +): + """A Go panic also exits 2; only the flag error means an old verifier.""" + packages_file = _write_packages(tmp_path / "packages.txt", _many_packages()) + + def batch_for(cmd): + return 2, [], "panic: runtime error: index out of range\n\ngoroutine 1\n" + + def result_for(cmd): + return _OK if _payload_of(cmd).split("\n")[0].endswith("-ok") else _NOT_FOUND + + with mock.patch("subprocess.Popen", + _fake_popen(result_for, batch_for)) as popen: + assert inclusion_proof_check.perform_inclusion_proof_check( + "/path/to/verifier", str(packages_file), logger, prefetch=False) + + # Every split is still verified, one per verifier run... + assert len(_split_cmds(popen)) == 6 + assert _results_by_hash(tmp_path) == _MANY_PACKAGES_RESULTS + # ...but the crash is reported as such. + messages = [(r.levelno, r.getMessage()) for r in caplog.records] + assert not any("does not support batch mode" in m for _, m in messages) + assert (logging.WARNING, + "Batch verifier exited with code 2 after 0 of 6 results; verifying " + "the remaining 6 split(s) one at a time.") in messages + + +@pytest.mark.parametrize("executable", [False, True], + ids=["missing", "not_executable"]) +def test_perform_inclusion_proof_check_batch_verifier_cannot_run( + tmp_path: Path, logger: logging.Logger, caplog: pytest.LogCaptureFixture, + executable: bool +): + packages_file = _write_packages(tmp_path / "packages.txt", _many_packages()) + verifier = tmp_path / "verifier" + if executable: + verifier.write_text("not a program") + verifier.chmod(0o644) + + with mock.patch("subprocess.Popen", wraps=subprocess.Popen) as popen: + assert inclusion_proof_check.perform_inclusion_proof_check( + str(verifier), str(packages_file), logger, prefetch=False) + + assert set(_results_by_hash(tmp_path).values()) == {False} + assert len(_results_by_hash(tmp_path)) == 6 + # Tried once, in batch mode, not once more per split. + assert popen.call_count == 1 + errors = [r.getMessage() for r in caplog.records + if r.levelno == logging.ERROR] + assert len(errors) == 1 + assert errors[0].startswith("Cannot run verifier `{}`".format(verifier)) + assert errors[0].endswith("Marking 6 split(s) as not verified.") + + +@pytest.mark.parametrize("error", [FileNotFoundError, PermissionError]) +def test_perform_inclusion_proof_check_batch_temp_file_error_falls_back( + tmp_path: Path, logger: logging.Logger, caplog: pytest.LogCaptureFixture, + error: type +): + """E.g. a bad TMPDIR is not blamed on the verifier.""" + packages_file = _write_packages(tmp_path / "packages.txt", _many_packages()) + + def result_for(cmd): + return _OK if _payload_of(cmd).split("\n")[0].endswith("-ok") else _NOT_FOUND + + with mock.patch.object(inclusion_proof_check, "_write_payloads_file", + side_effect=error("no temp dir")), \ + mock.patch("subprocess.Popen", _fake_popen(result_for)) as popen: + assert inclusion_proof_check.perform_inclusion_proof_check( + "/path/to/verifier", str(packages_file), logger, prefetch=False) + + assert not _batch_cmds(popen) + assert len(_split_cmds(popen)) == 6 + assert _results_by_hash(tmp_path) == _MANY_PACKAGES_RESULTS + messages = [r.getMessage() for r in caplog.records] + assert "Could not prepare batch verification: no temp dir" in messages + assert not any(m.startswith("Cannot run verifier") for m in messages) + + +def test_error_during_check_is_not_logged_as_interrupted( + tmp_path: Path, logger: logging.Logger, caplog: pytest.LogCaptureFixture +): + packages_file = _write_packages(tmp_path / "packages.txt", _many_packages()) + + def progress(done, total): + if done: + raise RuntimeError("event stream broke") + + with mock.patch("subprocess.Popen", _fake_popen(lambda cmd: _OK)): + with pytest.raises(RuntimeError): + inclusion_proof_check.perform_inclusion_proof_check( + "/path/to/verifier", str(packages_file), logger, prefetch=False, + progress=progress) + + assert not any("interrupted" in r.getMessage() for r in caplog.records) + + def test_main_exits_nonzero_on_failure( tmp_path: Path, monkeypatch: pytest.MonkeyPatch ): @@ -432,5 +1068,23 @@ def test_main_exits_nonzero_on_failure( assert exc_info.value.code == 1 +@pytest.mark.parametrize("error", [KeyboardInterrupt, RuntimeError]) +def test_main_restores_sigterm_handler_on_exception( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, error: type +): + monkeypatch.setattr(sys, "argv", [ + "inclusion_proof_check.py", + f"--packages_file={tmp_path / 'packages.txt'}", + "--verifier_path=/path/to/verifier", + ]) + before = signal.getsignal(signal.SIGTERM) + with mock.patch.object(inclusion_proof_check, + "perform_inclusion_proof_check", + side_effect=error): + with pytest.raises(error): + inclusion_proof_check.main() + assert signal.getsignal(signal.SIGTERM) == before + + if __name__ == "__main__": sys.exit(pytest.main([__file__])) diff --git a/verifier_tools/verify/README.md b/verifier_tools/verify/README.md index ebc40e4..77b19ff 100644 --- a/verifier_tools/verify/README.md +++ b/verifier_tools/verify/README.md @@ -51,6 +51,26 @@ where `log_type` is one of the following: * `google_1p_apk` (for Google Product Applications) * `mainline_module` (for Android Mainline Modules) +### Batch Verification Mode + +To verify many candidate binaries in one run: +``` +$ ./verifier --payloads_path=${PAYLOADS_PATH} --log_type= [--cache_dir=] +``` +The input is a [JSON Lines](https://jsonlines.org/) file with one payload per line: +``` +{"payload": "\n\n\n\n"} +``` +Each payload gets the same result as a separate `--payload_path` run, but each log's checkpoint is fetched, and its entries searched, only once for all of them. This is much faster than one run per payload, especially on a cold cache. + +One JSON result per payload is written to stdout, as soon as it is known (so not necessarily in input order). `index` is the payload's 0-based record number in the input: the first `{"payload": ...}` record is `0`, the next `1`, and so on (blank lines do not count): +``` +{"index":0,"verified":true,"log":"google_1p_apk (2026/02 Tessera)"} +{"index":2,"verified":false,"error":"inclusion check error in tlog.CheckRecord ..."} +{"index":1,"verified":false} +``` +`error` is set if a payload was found in a log but its inclusion proof failed; a payload that is in no log has neither `log` nor `error`. The exit code is 0 once every result has been written, whether or not the payloads verified, and 1 if the input cannot be read or the run is interrupted (in which case some payloads have no result). + ### Pre-fetching & Offline Cache Mode To pre-fetch and locally cache all entry tiles or legacy info files up to the current checkpoint (without requiring a payload or running an inclusion proof): @@ -69,6 +89,7 @@ This enables: | --- | --- | --- | | `--log_type`, `--log-type` | Target transparency log (`pixel`, `google_1p_code`, `google_1p_apk`, `mainline_module`). Required. | `""` | | `--payload_path`, `--payload-path` | Path to the payload file describing the candidate binary. Required for verification mode. | `""` | +| `--payloads_path`, `--payloads-path` | Path to a JSON Lines file of payloads. Required for batch verification mode; cannot be combined with `--payload_path`. | `""` | | `--fetch_entries`, `--fetch-entries` | Pre-fetch and cache all log entries locally up to the latest checkpoint. | `false` | | `--concurrency` | Number of concurrent workers for fetching Tessera entry tiles. | `16` | | `--cache_dir`, `--cache-dir` | Custom root directory for local cache. If unspecified, defaults to system cache. | OS user cache dir | diff --git a/verifier_tools/verify/cmd/verifier/verifier.go b/verifier_tools/verify/cmd/verifier/verifier.go index cd188b6..bcd57cc 100644 --- a/verifier_tools/verify/cmd/verifier/verifier.go +++ b/verifier_tools/verify/cmd/verifier/verifier.go @@ -22,11 +22,15 @@ package main import ( "bytes" "context" + "encoding/json" "flag" "fmt" + "io" "log/slog" + "maps" "os" "os/signal" + "slices" "syscall" "github.com/android/android-binary-transparency/verifier_tools/verify/internal/checkpoint" @@ -83,6 +87,7 @@ var mainlineModuleLogPubKey []byte var ( payloadPath = flag.String("payload_path", "", "Path to the payload describing the binary of interest.") + payloadsPath = flag.String("payloads_path", "", "Path to a JSON Lines file of payloads to verify in one run, one {\"payload\": \"...\"} object per line. Writes one JSON result per payload to stdout.") logType = flag.String("log_type", "", "Which log: 'pixel' or 'google_1p_code' or 'google_1p_apk' or 'mainline_module'.") fetchEntries = flag.Bool("fetch_entries", false, "Pre-fetch and cache all entries/tiles locally for the specified --log_type, without performing an inclusion proof.") concurrency = flag.Int("concurrency", tiles.DefaultTesseraFetchConcurrency, "Number of concurrent workers for fetching Tessera entry tiles.") @@ -92,6 +97,7 @@ var ( func init() { flag.StringVar(logType, "log-type", "", "Alias for --log_type.") flag.StringVar(payloadPath, "payload-path", "", "Alias for --payload_path.") + flag.StringVar(payloadsPath, "payloads-path", "", "Alias for --payloads_path.") flag.StringVar(cacheDir, "cache-dir", "", "Alias for --cache_dir.") flag.BoolVar(fetchEntries, "fetch-entries", false, "Alias for --fetch_entries.") @@ -102,14 +108,17 @@ Modes: 1. Verify binary inclusion in transparency log: %s --log_type= --payload_path= [--cache_dir=] - 2. Pre-fetch and cache entries locally for offline verification: + 2. Verify many binaries in one run (JSON Lines in, JSON Lines out): + %s --log_type= --payloads_path= [--cache_dir=] + + 3. Pre-fetch and cache entries locally for offline verification: %s --log_type= --fetch_entries [--concurrency=16] [--cache_dir=] Supported log types: pixel, google_1p_code, google_1p_apk, mainline_module Flags: -`, os.Args[0], os.Args[0], os.Args[0]) +`, os.Args[0], os.Args[0], os.Args[0], os.Args[0]) flag.PrintDefaults() } } @@ -291,8 +300,17 @@ func main() { return } + if *payloadsPath != "" { + if *payloadPath != "" { + slog.Error("specify only one of '--payload_path' and '--payloads_path'") + flag.Usage() + os.Exit(1) + } + os.Exit(runBatch(ctx, targets, *payloadsPath, os.Stdout)) + } + if *payloadPath == "" { - slog.Error("must specify either '--payload_path' to verify a binary, or '--fetch_entries' to pre-fetch log entries") + slog.Error("must specify '--payload_path' or '--payloads_path' to verify binaries, or '--fetch_entries' to pre-fetch log entries") flag.Usage() os.Exit(1) } @@ -301,16 +319,137 @@ func main() { slog.Error("Unable to open file", "path", *payloadPath, "error", err) os.Exit(1) } - // Payload should not contain excessive leading or trailing whitespace. - payloadBytes := bytes.TrimSpace(b) - payloadBytes = append(payloadBytes, '\n') + payloadBytes := normalizePayload(b) if string(b) != string(payloadBytes) { slog.Info("Reformatted payload content", "from", b, "to", payloadBytes) } - var verified bool + var result payloadResult + if err := verifyPayloads(ctx, targets, [][]byte{payloadBytes}, func(r payloadResult) { result = r }); err != nil { + slog.Error("FAILURE: verification interrupted", "error", err) + os.Exit(1) + } + switch { + case result.Verified: + slog.Info("OK. inclusion check success!", "log", result.Log) + case result.Error != "": + slog.Error("FAILURE: " + result.Error) + os.Exit(1) + default: + slog.Error("FAILURE: payload not verified in any log") + os.Exit(1) + } +} + +// normalizePayload trims excessive leading or trailing whitespace from a +// payload and terminates it with a single newline, as it appears in the logs. +func normalizePayload(b []byte) []byte { + p := append([]byte(nil), bytes.TrimSpace(b)...) + return append(p, '\n') +} + +// payloadResult is the outcome for one payload; in --payloads_path mode it is +// written to stdout as one JSON line. +type payloadResult struct { + // Index is the payload's 0-based record number in the input. + Index int `json:"index"` + // Verified is true if the payload's inclusion proof succeeded. + Verified bool `json:"verified"` + // Log names the log the payload was verified in, if Verified. + Log string `json:"log,omitempty"` + // Error is set if the payload was found in a log but could not be proven + // included. A payload that is simply not in any log has no error. + Error string `json:"error,omitempty"` +} + +// readPayloads reads a JSON Lines file of {"payload": "..."} objects. Blank +// lines are ignored, as are unknown fields. +func readPayloads(r io.Reader) ([][]byte, error) { + var payloads [][]byte + dec := json.NewDecoder(r) + for { + var rec struct { + Payload *string `json:"payload"` + } + if err := dec.Decode(&rec); err == io.EOF { + return payloads, nil + } else if err != nil { + return nil, fmt.Errorf("record %d: %w", len(payloads), err) + } + if rec.Payload == nil { + return nil, fmt.Errorf("record %d: missing \"payload\"", len(payloads)) + } + payloads = append(payloads, normalizePayload([]byte(*rec.Payload))) + } +} + +// runBatch verifies every payload in the JSON Lines file at path and writes +// one payloadResult per payload to out. Results are written as soon as they +// are known, so they are not necessarily in input order. +// +// Returns the process exit code: 0 once every result has been written, +// whether or not the payloads were verified; 1 if the input cannot be read, +// the results cannot be written, or ctx is cancelled first (in which case +// some payloads have no result). +func runBatch(ctx context.Context, targets []logTarget, path string, out io.Writer) int { + f, err := os.Open(path) + if err != nil { + slog.Error("Unable to open file", "path", path, "error", err) + return 1 + } + payloads, err := readPayloads(f) + f.Close() + if err != nil { + slog.Error("Malformed payloads file", "path", path, "error", err) + return 1 + } + + enc := json.NewEncoder(out) + var writeErr error + verified := 0 + err = verifyPayloads(ctx, targets, payloads, func(r payloadResult) { + if r.Verified { + verified++ + } + if writeErr == nil { + writeErr = enc.Encode(r) + } + }) + if err != nil { + slog.Error("FAILURE: verification interrupted", "error", err) + return 1 + } + if writeErr != nil { + slog.Error("FAILURE: unable to write results", "error", writeErr) + return 1 + } + slog.Info("Verified payloads", "total", len(payloads), "verified", verified) + return 0 +} + +// verifyPayloads checks each (normalized) payload against targets, in order, +// and calls emit exactly once per payload. +// +// A payload's outcome is the same as verifying it alone: the first log that +// contains it decides, and a failed proof there is final. But each checkpoint +// is fetched, and each log searched, once for all payloads. +// +// If ctx is cancelled, verifyPayloads returns ctx.Err() without calling emit +// for payloads it has not finished, rather than reporting them as not found. +func verifyPayloads(ctx context.Context, targets []logTarget, payloads [][]byte, emit func(payloadResult)) error { + pending := make(map[int]bool, len(payloads)) + for i := range payloads { + pending[i] = true + } + for _, target := range targets { - slog.Info("Checking log", "log", target.name, "url", target.baseURL) + if len(pending) == 0 { + break + } + if err := ctx.Err(); err != nil { + return err + } + slog.Info("Checking log", "log", target.name, "url", target.baseURL, "payloads", len(pending)) root, err := checkpoint.FromURLWithPathContext(ctx, target.baseURL, target.checkpointPath, target.verifier) if err != nil { slog.Warn("Failed to read checkpoint", "log", target.name, "error", err) @@ -318,71 +457,109 @@ func main() { } logSize := int64(root.Size) - var binaryInfoIndex int64 - var found bool - - if target.isTessera { - idx, ok, err := tiles.TesseraFindPayloadIndex(target.baseURL, logSize, payloadBytes) - if err != nil { - slog.Warn("Failed to search Tessera entry tiles", "log", target.name, "error", err) - continue - } - binaryInfoIndex = idx - found = ok - } else { - for _, filename := range target.binaryInfoFilenames { - m, err := tiles.BinaryInfosIndex(target.baseURL, filename, logSize) - if err != nil { - slog.Warn("Failed to load binary info map", "log", target.name, "file", filename, "error", err) - continue - } - if idx, ok := m[string(payloadBytes)]; ok { - binaryInfoIndex = idx - found = true - break - } - } - } - - if !found { - slog.Info("Payload not found in log", "log", target.name) - continue + found := findPayloads(target, logSize, payloads, slices.Sorted(maps.Keys(pending))) + if missing := len(pending) - len(found); missing > 0 { + slog.Info("Payload not found in log", "log", target.name, "count", missing) } var th tlog.Hash copy(th[:], root.Hash) - r := tiles.HashReader{ URL: target.baseURL, TileHeight: target.tileHeight, TreeSize: logSize, IsTessera: target.isTessera, + TileCache: make(map[string][]byte), } - slog.Debug("tlog.ProveRecord", "log", target.name, "logSize", logSize, "binaryInfoIndex", binaryInfoIndex) - rp, err := tlog.ProveRecord(logSize, binaryInfoIndex, r) - if err != nil { - slog.Error("error in tlog.ProveRecord", "log", target.name, "error", err) - os.Exit(1) + for _, i := range slices.Sorted(maps.Keys(found)) { + if err := ctx.Err(); err != nil { + return err + } + delete(pending, i) + emit(proveInclusion(target, r, logSize, th, i, payloads[i], found[i])) } + } + + // A cancelled checkpoint fetch above is only logged, so check again before + // declaring the remaining payloads not found. + if err := ctx.Err(); err != nil { + return err + } + for _, i := range slices.Sorted(maps.Keys(pending)) { + emit(payloadResult{Index: i}) + } + return nil +} - leafHash, err := tiles.PayloadHash(payloadBytes) +// findPayloads returns the leaf index in target of each payloads[i], for i in +// indices, that the log contains. +func findPayloads(target logTarget, logSize int64, payloads [][]byte, indices []int) map[int]int64 { + found := make(map[int]int64) + if target.isTessera { + wanted := make([][]byte, 0, len(indices)) + for _, i := range indices { + wanted = append(wanted, payloads[i]) + } + m, err := tiles.TesseraFindPayloadIndices(target.baseURL, logSize, wanted) if err != nil { - slog.Error("error hashing payload", "error", err) - os.Exit(1) + // m still holds the exact matches found before the error; the + // other payloads move on to the next log, as in a single-payload + // run that hits the same error. + slog.Warn("Failed to search Tessera entry tiles", "log", target.name, "found", len(m), "error", err) } - - if err := tlog.CheckRecord(rp, logSize, th, binaryInfoIndex, leafHash); err != nil { - slog.Error("FAILURE: inclusion check error in tlog.CheckRecord", "log", target.name, "error", err) - os.Exit(1) + for _, i := range indices { + if idx, ok := m[string(bytes.TrimSpace(payloads[i]))]; ok { + found[i] = idx + } } + return found + } - slog.Info("OK. inclusion check success!", "log", target.name) - verified = true - break + // Load each info file only if some payloads are still not found, like a + // single-payload run does. + remaining := indices + for _, filename := range target.binaryInfoFilenames { + if len(remaining) == 0 { + break + } + m, err := tiles.BinaryInfosIndex(target.baseURL, filename, logSize) + if err != nil { + slog.Warn("Failed to load binary info map", "log", target.name, "file", filename, "error", err) + continue + } + var next []int + for _, i := range remaining { + if idx, ok := m[string(payloads[i])]; ok { + found[i] = idx + } else { + next = append(next, i) + } + } + remaining = next } + return found +} - if !verified { - slog.Error("FAILURE: payload not verified in any log") - os.Exit(1) +// proveInclusion proves that payload is the leaf at leafIndex in target's tree +// of size logSize and root hash rootHash. +func proveInclusion(target logTarget, r tiles.HashReader, logSize int64, rootHash tlog.Hash, index int, payload []byte, leafIndex int64) payloadResult { + result := payloadResult{Index: index} + slog.Debug("tlog.ProveRecord", "log", target.name, "logSize", logSize, "binaryInfoIndex", leafIndex) + rp, err := tlog.ProveRecord(logSize, leafIndex, r) + if err != nil { + result.Error = fmt.Sprintf("error in tlog.ProveRecord for log %s: %v", target.name, err) + return result + } + leafHash, err := tiles.PayloadHash(payload) + if err != nil { + result.Error = fmt.Sprintf("error hashing payload: %v", err) + return result + } + if err := tlog.CheckRecord(rp, logSize, rootHash, leafIndex, leafHash); err != nil { + result.Error = fmt.Sprintf("inclusion check error in tlog.CheckRecord for log %s: %v", target.name, err) + return result } + result.Verified = true + result.Log = target.name + return result } diff --git a/verifier_tools/verify/cmd/verifier/verifier_batch_test.go b/verifier_tools/verify/cmd/verifier/verifier_batch_test.go new file mode 100644 index 0000000000..e7641e7 --- /dev/null +++ b/verifier_tools/verify/cmd/verifier/verifier_batch_test.go @@ -0,0 +1,431 @@ +package main + +import ( + "bytes" + "context" + "crypto/rand" + "encoding/base64" + "encoding/binary" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "sync/atomic" + "testing" + + "github.com/android/android-binary-transparency/verifier_tools/verify/internal/tiles" + "github.com/google/go-cmp/cmp" + "golang.org/x/mod/sumdb/note" + "golang.org/x/mod/sumdb/tlog" +) + +// fakeLog serves a small signed transparency log over HTTP: a checkpoint, +// hash tiles and either Tessera entry tiles or a legacy package_info.txt. +type fakeLog struct { + t *testing.T + origin string + tessera bool + tileHeight int + entries [][]byte + // corruptHashTiles serves zeroed hash tiles, so inclusion proofs fail. + corruptHashTiles bool + // failFirstEntryTile makes the oldest Tessera entry tile return HTTP 500. + failFirstEntryTile bool + + server *httptest.Server + verifier note.Verifier + checkpointRequests atomic.Int64 + infoFileRequests atomic.Int64 + entryTileRequests atomic.Int64 + hashes []tlog.Hash + signedCheckpointTxt []byte +} + +func (l *fakeLog) readHashes(indices []int64) ([]tlog.Hash, error) { + out := make([]tlog.Hash, len(indices)) + for i, idx := range indices { + if idx >= int64(len(l.hashes)) { + return nil, fmt.Errorf("no stored hash %d", idx) + } + out[i] = l.hashes[idx] + } + return out, nil +} + +func (l *fakeLog) start() *fakeLog { + l.t.Helper() + for i, e := range l.entries { + hs, err := tlog.StoredHashes(int64(i), e, tlog.HashReaderFunc(l.readHashes)) + if err != nil { + l.t.Fatalf("StoredHashes: %v", err) + } + l.hashes = append(l.hashes, hs...) + } + root, err := tlog.TreeHash(int64(len(l.entries)), tlog.HashReaderFunc(l.readHashes)) + if err != nil { + l.t.Fatalf("TreeHash: %v", err) + } + + name := strings.TrimSuffix(l.origin, "\n") + skey, vkey, err := note.GenerateKey(rand.Reader, name) + if err != nil { + l.t.Fatalf("GenerateKey: %v", err) + } + signer, err := note.NewSigner(skey) + if err != nil { + l.t.Fatalf("NewSigner: %v", err) + } + if l.verifier, err = note.NewVerifier(vkey); err != nil { + l.t.Fatalf("NewVerifier: %v", err) + } + text := fmt.Sprintf("%s%d\n%s\n", l.origin, len(l.entries), base64.StdEncoding.EncodeToString(root[:])) + if l.signedCheckpointTxt, err = note.Sign(¬e.Note{Text: text}, signer); err != nil { + l.t.Fatalf("Sign: %v", err) + } + + l.server = httptest.NewServer(http.HandlerFunc(l.serve)) + l.t.Cleanup(l.server.Close) + return l +} + +func (l *fakeLog) serve(w http.ResponseWriter, r *http.Request) { + p := strings.TrimPrefix(r.URL.Path, "/") + switch { + case p == "checkpoint": + l.checkpointRequests.Add(1) + w.Write(l.signedCheckpointTxt) + case p == PackageInfoFilename && !l.tessera: + l.infoFileRequests.Add(1) + var records []string + for i, e := range l.entries { + records = append(records, fmt.Sprintf("%d\n%s", i, bytes.TrimSpace(e))) + } + w.Write([]byte(strings.Join(records, "\n\n"))) + case strings.HasPrefix(p, "tile/entries/") && l.tessera: + l.entryTileRequests.Add(1) + var n, width int + rest := strings.TrimPrefix(p, "tile/entries/") + if _, err := fmt.Sscanf(rest, "%03d.p/%d", &n, &width); err != nil { + if _, err := fmt.Sscanf(rest, "%03d", &n); err != nil { + http.NotFound(w, r) + return + } + width = 256 + } + if n == 0 && l.failFirstEntryTile { + http.Error(w, "boom", http.StatusInternalServerError) + return + } + var buf bytes.Buffer + for _, e := range l.entries[n*256 : n*256+width] { + binary.Write(&buf, binary.BigEndian, uint16(len(e))) + buf.Write(e) + } + w.Write(buf.Bytes()) + case strings.HasPrefix(p, "tile/"): + if l.tessera { + p = fmt.Sprintf("tile/%d/%s", l.tileHeight, strings.TrimPrefix(p, "tile/")) + } + tile, err := tlog.ParseTilePath(p) + if err != nil || tile.H != l.tileHeight { + http.NotFound(w, r) + return + } + data, err := tlog.ReadTileData(tile, tlog.HashReaderFunc(l.readHashes)) + if err != nil { + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + if l.corruptHashTiles { + data = make([]byte, len(data)) + } + w.Write(data) + default: + http.NotFound(w, r) + } +} + +func (l *fakeLog) target(name string) logTarget { + target := logTarget{ + name: name, + baseURL: l.server.URL, + checkpointPath: "checkpoint", + verifier: l.verifier, + tileHeight: l.tileHeight, + isTessera: l.tessera, + } + if !l.tessera { + target.binaryInfoFilenames = []string{PackageInfoFilename} + } + return target +} + +func testPayload(name string) []byte { + return []byte(fmt.Sprintf("hash_%s\nSHA256(Signed APK)\n%s\n1\n", name, name)) +} + +// testLogs returns a Tessera log followed by a legacy log. "both" is in each; +// "new" only in the Tessera log (in its last, partial entry tile); "old" only +// in the legacy log. +func testLogs(t *testing.T, corruptTessera bool) (*fakeLog, *fakeLog) { + t.Helper() + tiles.SetCacheDir(t.TempDir()) + t.Cleanup(func() { tiles.SetCacheDir("") }) + + var tesseraEntries [][]byte + for i := 0; i < 300; i++ { + switch i { + case 10: + tesseraEntries = append(tesseraEntries, testPayload("both")) + case 290: + tesseraEntries = append(tesseraEntries, testPayload("new")) + default: + tesseraEntries = append(tesseraEntries, testPayload(fmt.Sprintf("t%d", i))) + } + } + tessera := (&fakeLog{ + t: t, + origin: "android.transparency.goog/google1p/apk/2026/1\n", + tessera: true, + tileHeight: 8, + entries: tesseraEntries, + corruptHashTiles: corruptTessera, + }).start() + legacy := (&fakeLog{ + t: t, + origin: "gstatic.com/android/binary_transparency/google1p/apk/2026/0\n", + tileHeight: 2, + entries: [][]byte{testPayload("l0"), testPayload("both"), testPayload("l2"), testPayload("old"), testPayload("l4")}, + }).start() + return tessera, legacy +} + +func collect(t *testing.T, ctx context.Context, targets []logTarget, payloads [][]byte) (map[int]payloadResult, error) { + t.Helper() + got := make(map[int]payloadResult) + err := verifyPayloads(ctx, targets, payloads, func(r payloadResult) { + if _, dup := got[r.Index]; dup { + t.Errorf("result for payload %d emitted twice", r.Index) + } + got[r.Index] = r + }) + return got, err +} + +func TestVerifyPayloads(t *testing.T) { + tessera, legacy := testLogs(t, false) + targets := []logTarget{tessera.target("tessera"), legacy.target("legacy")} + + payloads := [][]byte{ + normalizePayload(testPayload("both")), + normalizePayload(testPayload("missing")), + normalizePayload(testPayload("old")), + normalizePayload(append([]byte("\n "), testPayload("new")...)), + } + got, err := collect(t, context.Background(), targets, payloads) + if err != nil { + t.Fatalf("verifyPayloads: %v", err) + } + want := map[int]payloadResult{ + // The first log that contains a payload decides. + 0: {Index: 0, Verified: true, Log: "tessera"}, + 1: {Index: 1}, + 2: {Index: 2, Verified: true, Log: "legacy"}, + 3: {Index: 3, Verified: true, Log: "tessera"}, + } + if diff := cmp.Diff(want, got); diff != "" { + t.Errorf("verifyPayloads mismatch (-want +got):\n%s", diff) + } + + // Each log is consulted once for all payloads. + if n := tessera.checkpointRequests.Load(); n != 1 { + t.Errorf("Tessera checkpoint fetched %d times, want 1", n) + } + if n := legacy.checkpointRequests.Load(); n != 1 { + t.Errorf("legacy checkpoint fetched %d times, want 1", n) + } + if n := legacy.infoFileRequests.Load(); n != 1 { + t.Errorf("legacy info file fetched %d times, want 1", n) + } + if n := tessera.entryTileRequests.Load(); n != 2 { + t.Errorf("Tessera entry tiles fetched %d times, want 2", n) + } +} + +func TestVerifyPayloadsSkipsLaterLogsWhenAllFound(t *testing.T) { + tessera, legacy := testLogs(t, false) + targets := []logTarget{tessera.target("tessera"), legacy.target("legacy")} + + got, err := collect(t, context.Background(), targets, [][]byte{normalizePayload(testPayload("new"))}) + if err != nil { + t.Fatalf("verifyPayloads: %v", err) + } + if want := (payloadResult{Index: 0, Verified: true, Log: "tessera"}); got[0] != want { + t.Errorf("got %+v, want %+v", got[0], want) + } + if n := legacy.checkpointRequests.Load(); n != 0 { + t.Errorf("legacy checkpoint fetched %d times, want 0", n) + } + // "new" is in the last entry tile, so the first tile is never read. + if n := tessera.entryTileRequests.Load(); n != 1 { + t.Errorf("Tessera entry tiles fetched %d times, want 1", n) + } +} + +func TestVerifyPayloadsEntryTileErrorMatchesSinglePayloadRuns(t *testing.T) { + tessera, legacy := testLogs(t, false) + tessera.failFirstEntryTile = true // Tile 0 holds "both"; tile 1 holds "new". + targets := []logTarget{tessera.target("tessera"), legacy.target("legacy")} + + names := []string{"new", "both", "old", "missing"} + var payloads [][]byte + for _, name := range names { + payloads = append(payloads, normalizePayload(testPayload(name))) + } + got, err := collect(t, context.Background(), targets, payloads) + if err != nil { + t.Fatalf("verifyPayloads: %v", err) + } + want := map[int]payloadResult{ + // Found in tile 1, before the failing tile: kept. + 0: {Index: 0, Verified: true, Log: "tessera"}, + // Its Tessera search fails, so it is found in the next log. + 1: {Index: 1, Verified: true, Log: "legacy"}, + 2: {Index: 2, Verified: true, Log: "legacy"}, + 3: {Index: 3}, + } + if diff := cmp.Diff(want, got); diff != "" { + t.Errorf("verifyPayloads mismatch (-want +got):\n%s", diff) + } + + // Each result is what verifying that payload alone gives. + for i, p := range payloads { + single, err := collect(t, context.Background(), targets, [][]byte{p}) + if err != nil { + t.Fatalf("verifyPayloads(%s): %v", names[i], err) + } + r := single[0] + r.Index = i + if r != got[i] { + t.Errorf("%s: batch result %+v, single-payload result %+v", names[i], got[i], r) + } + } +} + +func TestVerifyPayloadsFailedProofIsFinal(t *testing.T) { + tessera, legacy := testLogs(t, true) + targets := []logTarget{tessera.target("tessera"), legacy.target("legacy")} + + got, err := collect(t, context.Background(), targets, [][]byte{ + normalizePayload(testPayload("both")), + normalizePayload(testPayload("old")), + }) + if err != nil { + t.Fatalf("verifyPayloads: %v", err) + } + // "both" is found in the Tessera log, whose proof fails; like a + // single-payload run, the legacy log is not tried for it. + if r := got[0]; r.Verified || r.Error == "" || r.Log != "" { + t.Errorf("payload 0 = %+v, want unverified with an error", r) + } + if want := (payloadResult{Index: 1, Verified: true, Log: "legacy"}); got[1] != want { + t.Errorf("payload 1 = %+v, want %+v", got[1], want) + } +} + +func TestVerifyPayloadsCancelled(t *testing.T) { + tessera, legacy := testLogs(t, false) + targets := []logTarget{tessera.target("tessera"), legacy.target("legacy")} + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + got, err := collect(t, ctx, targets, [][]byte{normalizePayload(testPayload("missing"))}) + if err == nil { + t.Errorf("verifyPayloads with cancelled context returned nil error") + } + if len(got) != 0 { + t.Errorf("verifyPayloads with cancelled context emitted %v, want nothing", got) + } +} + +func TestReadPayloads(t *testing.T) { + for _, tc := range []struct { + desc string + input string + want [][]byte + wantErr bool + }{ + {desc: "empty", input: ""}, + { + desc: "normalized, blank lines and unknown fields ignored", + input: "{\"payload\": \"a\\nb\"}\n\n{\"payload\": \" c\\n\\n\", \"label\": \"x\"}\n", + want: [][]byte{[]byte("a\nb\n"), []byte("c\n")}, + }, + {desc: "missing payload", input: "{\"payload\": \"a\"}\n{\"label\": \"x\"}\n", wantErr: true}, + {desc: "malformed", input: "{\"payload\": \"a\"}\nnot json\n", wantErr: true}, + {desc: "wrong type", input: "{\"payload\": 1}\n", wantErr: true}, + } { + t.Run(tc.desc, func(t *testing.T) { + got, err := readPayloads(strings.NewReader(tc.input)) + if (err != nil) != tc.wantErr { + t.Fatalf("readPayloads error = %v, wantErr %v", err, tc.wantErr) + } + if diff := cmp.Diff(tc.want, got); diff != "" { + t.Errorf("readPayloads mismatch (-want +got):\n%s", diff) + } + }) + } +} + +func TestRunBatch(t *testing.T) { + tessera, legacy := testLogs(t, false) + targets := []logTarget{tessera.target("tessera"), legacy.target("legacy")} + + dir := t.TempDir() + path := filepath.Join(dir, "payloads.jsonl") + var in bytes.Buffer + enc := json.NewEncoder(&in) + for _, name := range []string{"missing", "old", "new"} { + enc.Encode(map[string]string{"payload": string(testPayload(name))}) + } + if err := os.WriteFile(path, in.Bytes(), 0o600); err != nil { + t.Fatal(err) + } + + var out bytes.Buffer + if code := runBatch(context.Background(), targets, path, &out); code != 0 { + t.Fatalf("runBatch exit code = %d, want 0", code) + } + got := make(map[int]payloadResult) + for _, line := range strings.Split(strings.TrimSpace(out.String()), "\n") { + var r payloadResult + if err := json.Unmarshal([]byte(line), &r); err != nil { + t.Fatalf("output line %q: %v", line, err) + } + got[r.Index] = r + } + want := map[int]payloadResult{ + 0: {Index: 0}, + 1: {Index: 1, Verified: true, Log: "legacy"}, + 2: {Index: 2, Verified: true, Log: "tessera"}, + } + if diff := cmp.Diff(want, got); diff != "" { + t.Errorf("runBatch results mismatch (-want +got):\n%s", diff) + } + // Optional fields are omitted when empty. + if !strings.Contains(out.String(), `{"index":0,"verified":false}`) { + t.Errorf("runBatch output %q lacks compact not-found result", out.String()) + } + + malformed := filepath.Join(dir, "malformed.jsonl") + os.WriteFile(malformed, []byte("nope\n"), 0o600) + for _, p := range []string{malformed, filepath.Join(dir, "absent.jsonl")} { + out.Reset() + if code := runBatch(context.Background(), targets, p, &out); code != 1 || out.Len() != 0 { + t.Errorf("runBatch(%s) = %d with output %q, want 1 and no output", filepath.Base(p), code, out.String()) + } + } +} diff --git a/verifier_tools/verify/internal/tiles/reader.go b/verifier_tools/verify/internal/tiles/reader.go index adeb597..1818f0e 100644 --- a/verifier_tools/verify/internal/tiles/reader.go +++ b/verifier_tools/verify/internal/tiles/reader.go @@ -31,6 +31,10 @@ type HashReader struct { TileHeight int TreeSize int64 IsTessera bool + // TileCache, if non-nil, keeps downloaded hash tiles (path -> content) + // across ReadHashes calls, e.g. for several proofs against the same tree. + // Not safe for concurrent use. If nil, tiles are kept for one call only. + TileCache map[string][]byte } // Domain separation prefix for Merkle tree hashing with second preimage @@ -42,7 +46,10 @@ const ( // ReadHashes implements tlog.HashReader's ReadHashes. // See: https://pkg.go.dev/golang.org/x/mod/sumdb/tlog#HashReader. func (h HashReader) ReadHashes(indices []int64) ([]tlog.Hash, error) { - tiles := make(map[string][]byte) // cache tile path -> content + tiles := h.TileCache // cache tile path -> content + if tiles == nil { + tiles = make(map[string][]byte) + } hashes := make([]tlog.Hash, 0, len(indices)) for _, index := range indices { // A tlog index is a pointer to a hash at a given level in the tree. @@ -595,16 +602,48 @@ func FetchAllTesseraEntries(ctx context.Context, logBaseURL string, treeSize int // in the log. It returns the highest index if the payload appears more than once. // Returns (index, true, nil) if found, (-1, false, nil) if not found. func TesseraFindPayloadIndex(logBaseURL string, treeSize int64, targetPayload []byte) (int64, bool, error) { - if treeSize <= 0 { + found, err := TesseraFindPayloadIndices(logBaseURL, treeSize, [][]byte{targetPayload}) + if err != nil { + return -1, false, err + } + idx, ok := found[string(bytes.TrimSpace(targetPayload))] + if !ok { return -1, false, nil } + return idx, true, nil +} + +// TesseraFindPayloadIndices searches the entry tiles once for several +// payloads. It returns a map from each found payload, with surrounding +// whitespace trimmed, to its highest index in the log; payloads that are not +// found are absent. +// +// Tiles are read from latest to oldest, as in TesseraFindPayloadIndex, and +// the search stops as soon as every payload has been found. A payload that is +// not in the log therefore makes the search read every tile, but only once for +// all payloads. +// +// If a tile cannot be read, the search stops and returns the payloads found +// so far together with the error. Those matches are still exact: they were +// found in newer tiles, so a search for one payload alone would have stopped +// there too, without reaching the failing tile. Only the payloads not yet +// found are affected, as they would be when searched for alone. +func TesseraFindPayloadIndices(logBaseURL string, treeSize int64, targetPayloads [][]byte) (map[string]int64, error) { + found := make(map[string]int64) + wanted := make(map[string]bool, len(targetPayloads)) + for _, p := range targetPayloads { + wanted[string(bytes.TrimSpace(p))] = true + } + if treeSize <= 0 || len(wanted) == 0 { + return found, nil + } numTiles := (treeSize + 255) / 256 - target := bytes.TrimSpace(targetPayload) // Search in reverse (latest to oldest) to find recent packages with typically - // fewer HTTP requests on cold caches. - for tileN := numTiles - 1; tileN >= 0; tileN-- { + // fewer HTTP requests on cold caches. The first match seen for a payload is + // therefore its highest index. + for tileN := numTiles - 1; tileN >= 0 && len(found) < len(wanted); tileN-- { w := 256 if (tileN+1)*256 > treeSize { w = int(treeSize - tileN*256) @@ -612,20 +651,24 @@ func TesseraFindPayloadIndex(logBaseURL string, treeSize int64, targetPayload [] b, err := readCachedEntryTile(logBaseURL, tileN, w) if err != nil { - return -1, false, fmt.Errorf("failed to fetch entry tile %d (width %d): %w", tileN, w, err) + return found, fmt.Errorf("failed to fetch entry tile %d (width %d): %w", tileN, w, err) } entries, err := ParseEntryBundle(b) if err != nil { - return -1, false, fmt.Errorf("failed to parse entry tile %d: %w", tileN, err) + return found, fmt.Errorf("failed to parse entry tile %d: %w", tileN, err) } for idx := len(entries) - 1; idx >= 0; idx-- { - if bytes.Equal(bytes.TrimSpace(entries[idx]), target) { - return tileN*256 + int64(idx), true, nil + key := string(bytes.TrimSpace(entries[idx])) + if !wanted[key] { + continue + } + if _, seen := found[key]; !seen { + found[key] = tileN*256 + int64(idx) } } } - return -1, false, nil + return found, nil } diff --git a/verifier_tools/verify/internal/tiles/reader_test.go b/verifier_tools/verify/internal/tiles/reader_test.go index 386aa40..8458c9e 100644 --- a/verifier_tools/verify/internal/tiles/reader_test.go +++ b/verifier_tools/verify/internal/tiles/reader_test.go @@ -416,6 +416,121 @@ func TestTesseraFindPayloadIndex(t *testing.T) { } } +func TestTesseraFindPayloadIndices(t *testing.T) { + tempDir := t.TempDir() + t.Setenv("HOME", tempDir) + t.Setenv("XDG_CACHE_HOME", tempDir) + + entry := func(i int) []byte { return []byte(fmt.Sprintf("hash_%d\nhash_desc\npackage_%d\n%d\n", i, i, i)) } + dup := []byte("hash_dup\nhash_desc\npackage_dup\n1\n") + + // Tile 0 holds indices 0..255 (dup at 5 and 250); tile 1 holds 256..258 + // (dup at 257). + var tile0Entries [][]byte + for i := 0; i < 256; i++ { + if i == 5 || i == 250 { + tile0Entries = append(tile0Entries, dup) + } else { + tile0Entries = append(tile0Entries, entry(i)) + } + } + tile0Data := createTestEntryBundle(tile0Entries) + tile1Data := createTestEntryBundle([][]byte{entry(256), dup, entry(258)}) + + var tile0Requests, tile1Requests atomic.Int64 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/tile/entries/000": + tile0Requests.Add(1) + w.Write(tile0Data) + case "/tile/entries/001.p/3": + tile1Requests.Add(1) + w.Write(tile1Data) + default: + http.NotFound(w, r) + } + })) + defer server.Close() + + // All payloads in the latest tile: tile 0 is never read. Surrounding + // whitespace is ignored, as in TesseraFindPayloadIndex. + got, err := TesseraFindPayloadIndices(server.URL, 259, [][]byte{entry(258), append([]byte(" "), dup...)}) + if err != nil { + t.Fatalf("TesseraFindPayloadIndices error: %v", err) + } + want := map[string]int64{string(bytes.TrimSpace(entry(258))): 258, string(bytes.TrimSpace(dup)): 257} + if diff := cmp.Diff(want, got); diff != "" { + t.Errorf("TesseraFindPayloadIndices mismatch (-want +got):\n%s", diff) + } + if tile0Requests.Load() != 0 || tile1Requests.Load() != 1 { + t.Errorf("requests (tile0, tile1) = (%d, %d), want (0, 1)", tile0Requests.Load(), tile1Requests.Load()) + } + + // Payloads across both tiles plus one that is absent: every tile is read, + // each only once, and the absent payload is omitted. + got, err = TesseraFindPayloadIndices(server.URL, 259, [][]byte{entry(3), entry(256), []byte("absent"), entry(3)}) + if err != nil { + t.Fatalf("TesseraFindPayloadIndices error: %v", err) + } + want = map[string]int64{string(bytes.TrimSpace(entry(3))): 3, string(bytes.TrimSpace(entry(256))): 256} + if diff := cmp.Diff(want, got); diff != "" { + t.Errorf("TesseraFindPayloadIndices mismatch (-want +got):\n%s", diff) + } + if tile0Requests.Load() != 1 || tile1Requests.Load() != 1 { + t.Errorf("requests (tile0, tile1) = (%d, %d), want (1, 1)", tile0Requests.Load(), tile1Requests.Load()) + } + + // No payloads, or an empty tree: nothing to search. + for _, tc := range []struct { + size int64 + payloads [][]byte + }{{259, nil}, {0, [][]byte{entry(3)}}} { + got, err := TesseraFindPayloadIndices(server.URL, tc.size, tc.payloads) + if err != nil || len(got) != 0 { + t.Errorf("TesseraFindPayloadIndices(size=%d, %d payloads) = (%v, %v), want empty map", tc.size, len(tc.payloads), got, err) + } + } +} + +func TestTesseraFindPayloadIndicesKeepsMatchesOnTileError(t *testing.T) { + tempDir := t.TempDir() + t.Setenv("HOME", tempDir) + t.Setenv("XDG_CACHE_HOME", tempDir) + + entry := func(i int) []byte { return []byte(fmt.Sprintf("hash_%d\nhash_desc\npackage_%d\n%d\n", i, i, i)) } + tile1Data := createTestEntryBundle([][]byte{entry(256), entry(257)}) + + // Tile 0 (the oldest) fails; tile 1 is served. + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/tile/entries/001.p/2": + w.Write(tile1Data) + case "/tile/entries/000": + http.Error(w, "boom", http.StatusInternalServerError) + default: + http.NotFound(w, r) + } + })) + defer server.Close() + + // entry(3) would be in tile 0, so the search reaches it and fails. The + // match from tile 1 is exact and must be kept alongside the error. + got, err := TesseraFindPayloadIndices(server.URL, 258, [][]byte{entry(257), entry(3)}) + if err == nil { + t.Fatalf("TesseraFindPayloadIndices error = nil, want tile 0 failure") + } + want := map[string]int64{string(bytes.TrimSpace(entry(257))): 257} + if diff := cmp.Diff(want, got); diff != "" { + t.Errorf("TesseraFindPayloadIndices partial result mismatch (-want +got):\n%s", diff) + } + + // A single-payload search that is satisfied by tile 1 never reaches the + // failing tile. + if idx, found, err := TesseraFindPayloadIndex(server.URL, 258, entry(257)); err != nil || !found || idx != 257 { + t.Errorf("TesseraFindPayloadIndex = (%d, %v, %v), want (257, true, nil)", idx, found, err) + } +} + func createTestEntryBundle(entries [][]byte) []byte { var buf bytes.Buffer for _, entry := range entries {