diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 041345b8..5de1e780 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -58,16 +58,38 @@ jobs: reports/hypothesis/seed.txt if-no-files-found: ignore - mutation: - name: Mutation Gate + mutation-plan: + name: Mutation inventory + runs-on: ubuntu-latest + timeout-minutes: 5 + permissions: + contents: read + outputs: + modules: ${{ steps.plan.outputs.modules }} + steps: + - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + with: + persist-credentials: false + - name: Inventory every handwritten runtime module + id: plan + run: | + modules=$(bash scripts/mutation.sh --matrix) + echo "modules=$modules" >> "$GITHUB_OUTPUT" + + mutation-module: + name: Mutation module (${{ matrix.module }}) + needs: mutation-plan runs-on: ubuntu-latest timeout-minutes: 40 permissions: contents: read + strategy: + fail-fast: false + matrix: + module: ${{ fromJSON(needs.mutation-plan.outputs.modules) }} steps: - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 with: - fetch-depth: 0 persist-credentials: false - uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5 with: @@ -76,18 +98,30 @@ jobs: with: version: "0.12.17" - run: uv sync --locked - - name: Mutate changed and critical runtime modules + - name: Mutate assigned runtime modules env: - MUTATION_BASE_SHA: ${{ github.event.pull_request.base.sha || github.event.merge_group.base_sha || github.event.before }} + MUTATION_MODULE: ${{ matrix.module }} run: uv run --locked poe mutation - name: Preserve mutation outcomes if: ${{ always() }} uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1 with: - name: mutation-outcomes + name: mutation-outcomes-${{ matrix.module }} path: reports/mutation.json if-no-files-found: error + mutation: + name: Mutation Gate + if: ${{ always() }} + needs: [mutation-plan, mutation-module] + runs-on: ubuntu-latest + timeout-minutes: 5 + steps: + - name: Require every mutation module to succeed + env: + RESULTS: ${{ toJSON(needs.*.result) }} + run: jq -e 'length == 2 and all(. == "success")' <<< "$RESULTS" + quality: name: Quality Gate if: ${{ always() }} diff --git a/.github/workflows/mutation-audit.yml b/.github/workflows/mutation-audit.yml index 92b9b1a2..3c504c07 100644 --- a/.github/workflows/mutation-audit.yml +++ b/.github/workflows/mutation-audit.yml @@ -25,7 +25,7 @@ jobs: with: version: "0.12.17" - run: uv sync --locked - - run: uv run --locked poe mutation-full + - run: uv run --locked poe mutation - name: Preserve mutation outcomes if: ${{ always() }} uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1 diff --git a/maintainers/mutation-testing.md b/maintainers/mutation-testing.md index db16609c..1f386acc 100644 --- a/maintainers/mutation-testing.md +++ b/maintainers/mutation-testing.md @@ -1,20 +1,25 @@ # Mutation testing -`uv run --locked poe quality` runs all native checks and mutates every changed -handwritten runtime module plus the lock acquisition, guard, renewal, and worker -modules. CI uses the same `checks` and `mutation` tasks in separate jobs, then -requires both through `Quality Gate`. The weekly `poe mutation-full` task audits -every handwritten runtime module without a debt baseline. +`uv run --locked poe quality` runs all native checks and mutates every +handwritten runtime module. CI runs the same `checks` and `mutation` tasks, +assigning one handwritten module to each independent matrix job. Git's current +handwritten module inventory determines the jobs, including newly added modules. +`Mutation Gate` requires every module job, and `Quality Gate` requires it and the Python test matrix. The +weekly audit runs the same full mutation task without a debt baseline. Mutmut's [native configuration](https://mutmut.readthedocs.io/en/latest/) lives in `pyproject.toml`; it excludes only the generated OpenAPI client. Mutmut can -select modules by name but has no Git-changed-module option and returns success -when mutants survive. `scripts/mutation.sh` selects module names from Git and -`scripts/mutation_results.py` reads only those modules' native metadata. The +select modules by name but returns success when mutants survive. +`scripts/mutation.sh` selects the full handwritten inventory from Git and +`scripts/mutation_results.py` reads each selected module's native metadata. The report at `reports/mutation.json` distinguishes killed, statically invalid, surviving, uncovered, timed-out, crashed, interrupted, and missing results. +Mutmut creates mutants inside functions. Export-only modules remain in the +inventory; the runner verifies that they define no functions and records them +as unmutatable. Coverage and installed-package checks still include them. A pytest internal error is a harness crash, not a killed mutant. -The pinned Pyrefly check rejects type-invalid realtime and auth mutants before pytest; +The pinned Pyrefly check covers the handwritten runtime and rejects +type-invalid mutants before pytest; the report counts these as `type_checked`, separately from test-killed mutants. Surviving, uncovered, timed-out, crashed, and incomplete mutants still fail. Mutmut passes pytest `-x` so a selected test's first assertion failure kills the @@ -29,12 +34,13 @@ Mutmut uses its native forkserver isolation because forking from a process that has already run asyncio tests can crash a worker before its tests report a result. The mutation runner sets `NO_PROXY=*` for its hermetic transport tests because macOS system-proxy discovery can abort after a fork with active threads. -The pinned pytest-order plugin runs bounded callback and presence assertions -first when mutmut's unordered test selection could otherwise reach a blocked test. +The pinned pytest-order plugin runs bounded callback, presence, and fetch-worker +assertions first when mutmut's unordered test selection could otherwise reach a +blocked test. -The full audit runs mutmut once across the entire source tree. It reports any -surviving mutants without treating them as an approved baseline. Equivalent -mutants require a reviewed, exact exception before a gate can accept them. +The local and weekly full runs execute the inventory in one Mutmut process. +Equivalent mutants require a reviewed, exact exception before a gate can +accept them. Ruff, Mypy, Basedpyright, pytest, and Tox tasks pass `pyproject.toml` explicitly. Their documented config-file precedence can otherwise select a diff --git a/maintainers/quality-policy.lock.json b/maintainers/quality-policy.lock.json index 54d94fc7..b795e3f4 100644 --- a/maintainers/quality-policy.lock.json +++ b/maintainers/quality-policy.lock.json @@ -299,16 +299,9 @@ "pyrefly", "check", "--output-format=json", - "src/volcano_sdk/realtime.py", - "src/volcano_sdk/auth.py", - "src/volcano_sdk/storage.py", - "src/volcano_sdk/functions.py", - "src/volcano_sdk/durable_authoring.py", - "src/volcano_sdk/database.py", - "src/volcano_sdk/_function_resolution.py", - "src/volcano_sdk/_transport.py", - "src/volcano_sdk/_session.py", - "src/volcano_sdk/_session_operations.py" + "--project-excludes", + "src/volcano_sdk/_generated/**", + "src/volcano_sdk" ] }, "mypy": { @@ -400,7 +393,6 @@ "generated": "python -m scripts.check_openapi", "lint": "ruff check --config pyproject.toml .", "mutation": "bash scripts/mutation.sh", - "mutation-full": "MUTATION_FULL=1 bash scripts/mutation.sh", "mypy": "mypy --config-file pyproject.toml", "package-check": { "interpreter": "bash", diff --git a/pyproject.toml b/pyproject.toml index d012d250..d7ef67f6 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -458,16 +458,8 @@ partial_branches = [] source_paths = ["src/volcano_sdk"] type_check_command = [ "pyrefly", "check", "--output-format=json", - "src/volcano_sdk/realtime.py", - "src/volcano_sdk/auth.py", - "src/volcano_sdk/storage.py", - "src/volcano_sdk/functions.py", - "src/volcano_sdk/durable_authoring.py", - "src/volcano_sdk/database.py", - "src/volcano_sdk/_function_resolution.py", - "src/volcano_sdk/_transport.py", - "src/volcano_sdk/_session.py", - "src/volcano_sdk/_session_operations.py", + "--project-excludes", "src/volcano_sdk/_generated/**", + "src/volcano_sdk", ] cache_invalidation_files = ["tests/**/*.py"] pytest_add_cli_args_test_selection = ["tests/unit"] @@ -498,7 +490,6 @@ types = ["mypy", "basedpyright"] mypy = "mypy --config-file pyproject.toml" basedpyright = "basedpyright --project pyproject.toml" mutation = "bash scripts/mutation.sh" -mutation-full = "MUTATION_FULL=1 bash scripts/mutation.sh" [tool.poe.tasks.test] interpreter = "bash" diff --git a/scripts/check_quality_policy.py b/scripts/check_quality_policy.py index 784e2d51..9b401358 100644 --- a/scripts/check_quality_policy.py +++ b/scripts/check_quality_policy.py @@ -18,7 +18,7 @@ from collections.abc import Iterable GENERATED = "src/volcano_sdk/_generated" -LOCK_SHA256 = "489b1a41632f67d528dbb10267bb8921c38c78726c430d990b4ee99ebdd9236a" +LOCK_SHA256 = "3eea92085dce48f7454bc9e2b82854cc99a8787083da471609ae41da57d75560" TYPE_FIXTURES = { "tests/typing/contract_steps.py", "tests/typing/durable_callbacks.py", diff --git a/scripts/mutation.sh b/scripts/mutation.sh index d05e3ae7..1f40d14c 100644 --- a/scripts/mutation.sh +++ b/scripts/mutation.sh @@ -4,71 +4,87 @@ set -euo pipefail # Forked macOS workers must not query SystemConfiguration through urllib/httpx. export NO_PROXY='*' no_proxy='*' -# Refresh mutmut's test-to-mutant map so newly added tests are selected. -rm -rf -- mutants +paths=() +modules=() +while IFS= read -r -d '' path; do + if [[ $path == src/volcano_sdk/*.py && $path != src/volcano_sdk/_generated/* && -f $path ]]; then + module=${path#src/} + module=${module%.py} + if [[ $module == */__init__ ]]; then + module=${module%/__init__} + fi + module=${module//\//.} + if [[ ! $module =~ ^[a-zA-Z_][a-zA-Z_0-9]*(\.[a-zA-Z_][a-zA-Z_0-9]*)*$ ]]; then + echo "Invalid Python runtime module: $path" >&2 + exit 1 + fi + for existing in "${modules[@]-}"; do + if [[ $module == "$existing" ]]; then + echo "Duplicate Python runtime module: $module" >&2 + exit 1 + fi + done + paths+=("$path") + modules+=("$module") + fi +done < <(git ls-files --cached --others --exclude-standard -z -- src/volcano_sdk) + +if (( ${#paths[@]} == 0 )); then + echo 'No handwritten SDK runtime modules found' >&2 + exit 1 +fi +if [[ ${1:-} == --matrix && $# == 1 ]]; then + printf '[' + for index in "${!modules[@]}"; do + ((index == 0)) || printf ',' + printf '"%s"' "${modules[$index]}" + done + printf ']\n' + exit +fi +if (( $# != 0 )); then + echo 'Usage: scripts/mutation.sh [--matrix]' >&2 + exit 2 +fi mkdir -p reports targets=reports/mutation-targets.bin failed=reports/mutation-failed.bin : > "$targets" : > "$failed" +patterns=() +for index in "${!paths[@]}"; do + if [[ -n ${MUTATION_MODULE:-} ]] && [[ ${modules[$index]} != "$MUTATION_MODULE" ]]; then + continue + fi + path=${paths[$index]} + selected_path=$path + printf '%s\0' "$path" >> "$targets" + patterns+=("${modules[$index]}.x*") +done -modules=() -add_module() { - for known in "${modules[@]-}"; do - [[ $known == "$1" ]] && return - done - modules+=("$1") -} -if [[ ${MUTATION_FULL:-0} == 1 ]]; then - git ls-files -z 'src/volcano_sdk/*.py' > reports/mutation-source.bin - while IFS= read -r -d '' path; do - if [[ $path != src/volcano_sdk/_generated/* && -f $path ]]; then - add_module "$path" - fi - done < reports/mutation-source.bin -else - for path in \ - src/volcano_sdk/locks.py \ - src/volcano_sdk/_lock_guard.py \ - src/volcano_sdk/_lock_renewer.py \ - src/volcano_sdk/_lock_worker.py; do - add_module "$path" - done - - base=${MUTATION_BASE_SHA:-origin/main} - ancestor=$(git merge-base "$base" HEAD) - changed() { - git diff --name-only --diff-filter=ACMRT -z "$ancestor" HEAD - git diff --name-only --diff-filter=ACMRT -z HEAD - git ls-files -z --others --exclude-standard - } - changed > reports/mutation-changed.bin - while IFS= read -r -d '' path; do - if [[ $path == src/volcano_sdk/*.py && $path != src/volcano_sdk/_generated/* && -f $path ]]; then - add_module "$path" - fi - done < reports/mutation-changed.bin +if (( ${#patterns[@]} == 0 )); then + echo "Unknown mutation module: ${MUTATION_MODULE:-}" >&2 + exit 2 fi -for path in "${modules[@]}"; do - printf '%s\0' "$path" >> "$targets" -done +# A fresh run must not inherit stale test-to-mutant mappings or verdicts. +rm -rf -- mutants -if [[ ${MUTATION_FULL:-0} == 1 ]]; then - if ! mutmut run --max-children 1; then - printf '%s\0' "full mutation run" >> "$failed" - fi -else - patterns=() - for path in "${modules[@]}"; do - module=${path#src/} - module=${module%.py} - patterns+=("${module//\//.}.x*") - done - if ! mutmut run --max-children 1 "${patterns[@]}"; then - printf '%s\0' "scoped mutation run" >> "$failed" - fi +# Mutmut rejects an exact wildcard for a module with no functions. Record that +# module explicitly instead of treating a native no-match assertion as a kill. +if [[ -n ${MUTATION_MODULE:-} ]] && ! python -c ' +import sys +from pathlib import Path +from scripts.mutation_results import has_functions +raise SystemExit(0 if has_functions(Path(sys.argv[1])) else 1) +' "$selected_path"; then + python -m scripts.mutation_results "$targets" "$failed" + exit +fi + +if ! mutmut run --max-children 1 "${patterns[@]}"; then + printf '%s\0' 'mutation run' >> "$failed" fi python -m scripts.mutation_results "$targets" "$failed" diff --git a/scripts/mutation_results.py b/scripts/mutation_results.py index 3c8c774d..d5892e42 100644 --- a/scripts/mutation_results.py +++ b/scripts/mutation_results.py @@ -134,7 +134,7 @@ def main(targets_path: Path, failed_path: Path) -> int: counts.update(module_counts) if empty: unmutatable.append(name) - if not counts: + if not counts and (not targets or len(unmutatable) != len(targets)): failures.append("No mutants were tested") failures.extend( f"{name}: {count} mutant(s)" diff --git a/src/volcano_sdk/_realtime_fetch_worker.py b/src/volcano_sdk/_realtime_fetch_worker.py index 421e73f5..e0f0b835 100644 --- a/src/volcano_sdk/_realtime_fetch_worker.py +++ b/src/volcano_sdk/_realtime_fetch_worker.py @@ -293,9 +293,11 @@ async def _fetch_and_deliver( if isinstance(result, BaseException): await self._deliver_failure(jobs, result) return - if len(result) != len(jobs): - raise RuntimeError(_INVALID_RESULT_COUNT) - for job, record in zip(jobs, result, strict=True): + try: + pairs = tuple(zip(jobs, result, strict=True)) + except ValueError: + raise RuntimeError(_INVALID_RESULT_COUNT) from None + for job, record in pairs: await self._deliver(PostgresFetchOutcome(job=job, record=record)) async def _deliver_failure( diff --git a/src/volcano_sdk/connection_string.py b/src/volcano_sdk/connection_string.py index b69ad544..55bef179 100644 --- a/src/volcano_sdk/connection_string.py +++ b/src/volcano_sdk/connection_string.py @@ -1,7 +1,7 @@ """Postgres connection helpers for Volcano functions.""" import re -from urllib.parse import quote, unquote +from urllib.parse import quote, unquote, urlencode _FULL_ACCESS_APP_NAME = "volcano_full_access" _USER_ACCESS_APP_NAME = "volcano_user_access" @@ -40,8 +40,11 @@ def database_connection_string( target, query = _connection_parts(base_connection_string) parameters = _query_parameters(query) - application_name = quote(_database_application_name(user_id), safe="") - parameters.append(f"application_name={application_name}") + parameters.append( + urlencode( + {"application_name": _database_application_name(user_id)}, quote_via=quote + ) + ) return f"{target}?{'&'.join(parameters)}" @@ -50,32 +53,24 @@ def _connection_parts(value: str) -> tuple[str, str]: if prefix is None or _INVALID_PERCENT_ENCODING.search(value): raise ValueError(_INVALID_ERROR) - authority_end = value.find("/", prefix.end()) - possible_userinfo_end = value.find("@", prefix.end()) - userinfo_end = ( - possible_userinfo_end - if possible_userinfo_end != -1 - and (authority_end == -1 or possible_userinfo_end < authority_end) - else -1 - ) - query_start = value.find("?", max(prefix.end(), userinfo_end + 1)) + before_at, at_separator, _ = value[prefix.end() :].partition("@") + query_search_start = prefix.end() + if at_separator and "/" not in before_at: + query_search_start += len(before_at) + 1 + query_start = value.find("?", query_search_start) if query_start == -1: return value, "" return value[:query_start], value[query_start + 1 :] def _query_parameters(query: str) -> list[str]: - if not query: - return [] parameters = [ parameter for parameter in query.split("&") if unquote(parameter.partition("=")[0]) != "application_name" ] - for index in range(len(parameters) - 1, -1, -1): - if parameters[index]: - return parameters[: index + 1] - return [] + serialized = "&".join(parameters).rstrip("&") + return serialized.split("&") if serialized else [] def _database_application_name(user_id: str | None) -> str: diff --git a/tests/unit/test_connection_string.py b/tests/unit/test_connection_string.py index 3ad56efc..f92c9006 100644 --- a/tests/unit/test_connection_string.py +++ b/tests/unit/test_connection_string.py @@ -89,6 +89,19 @@ def test_database_connection_string_drops_only_empty_parameters() -> None: ) +@pytest.mark.parametrize( + "query", ["mode=X", "mode=with space", "mode=XX&XX", "mode=x&&other=y"] +) +def test_database_connection_string_preserves_unrelated_query_fields( + query: str, +) -> None: + base = f"postgresql://host/db?{query}" + + assert database_connection_string(base) == ( + f"{base}&application_name=volcano_full_access" + ) + + def test_database_connection_string_ignores_at_sign_in_query_value() -> None: assert database_connection_string("postgresql://host/db?options=foo@bar") == ( "postgresql://host/db?options=foo@bar&application_name=volcano_full_access" @@ -109,6 +122,12 @@ def test_database_connection_string_keeps_credential_question_mark() -> None: ) +def test_database_connection_string_finds_query_immediately_after_userinfo() -> None: + assert database_connection_string("postgres://u@?sslmode=require") == ( + "postgres://u@?sslmode=require&application_name=volcano_full_access" + ) + + def test_database_connection_string_ignores_later_at_sign_in_query() -> None: base = "postgres://u@host?options=a@b" diff --git a/tests/unit/test_database_refresh.py b/tests/unit/test_database_refresh.py index d409bad0..43a31064 100644 --- a/tests/unit/test_database_refresh.py +++ b/tests/unit/test_database_refresh.py @@ -362,6 +362,9 @@ def test_refresh_listener_can_wait_for_another_refresh_thread() -> None: def on_refresh(event: str, _session: Session | None) -> None: if event != "TOKEN_REFRESHED": return + if workers: + completed.append(False) + return subscription.unsubscribe() worker = Thread(target=client.auth.refresh_session) workers.append(worker) diff --git a/tests/unit/test_durable_authoring.py b/tests/unit/test_durable_authoring.py index 64178228..0c873662 100644 --- a/tests/unit/test_durable_authoring.py +++ b/tests/unit/test_durable_authoring.py @@ -131,6 +131,22 @@ def test_retry_false_produces_an_immediate_no_retry_decision() -> None: assert decision.delay.to_seconds() == 0 +@pytest.mark.order(0) +def test_step_forwards_a_disabled_retry_before_scheduling() -> None: + runtime = RecordingContext() + context = DurableContext(runtime, durable_authoring._Engine()) + + with pytest.raises(AssertionError, match="unexpected runtime operation"): + _ = context.step("once", lambda _scope: "done", retry=False) + + assert isinstance(runtime.config, StepConfig) + retry = runtime.config.retry_strategy + assert retry is not None + decision = retry(RuntimeError("failed"), 1) + assert decision.should_retry is False + assert decision.delay.to_seconds() == 0 + + def test_custom_retry_receives_the_original_error() -> None: engine = durable_authoring._Engine() failure = RuntimeError("failed") @@ -185,6 +201,18 @@ def record( assert configured.initial_state is False +def test_durable_runtime_adapter_is_loaded_once( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(durable_authoring._Engine, "_loaded", None) + + first = durable_authoring._Engine.load() + second = durable_authoring._Engine.load() + + assert isinstance(first, durable_authoring._Engine) + assert first is second + + def test_wait_options_name_invalid_timing_fields() -> None: engine = durable_authoring._Engine() @@ -989,17 +1017,16 @@ def handler(_event: object, ctx: DurableContext) -> object: def test_wait_until_refuses_a_timeout() -> None: - @durable - def handler(_event: object, ctx: DurableContext) -> object: - return ctx.wait_until( + context = DurableContext(RecordingContext(), durable_authoring._Engine()) + + # Validate before handing the condition to the runtime, which may wait + # indefinitely when the unsupported timeout is silently ignored. + with pytest.raises(TypeError, match="has no `timeout`"): + _ = context.wait_until( lambda state, _scope: state, WaitUntilOptions(until=bool, initial_state=False, timeout="1h"), ) - # A condition is bounded by checks, not by a deadline: the platform holds - # the wait between them and has no clock to compare against on resume. - assert "has no `timeout`" in failing_handler(handler) - def test_wait_until_requires_an_initial_state() -> None: @durable diff --git a/tests/unit/test_mutation_results.py b/tests/unit/test_mutation_results.py index 626a5c08..e4df23b6 100644 --- a/tests/unit/test_mutation_results.py +++ b/tests/unit/test_mutation_results.py @@ -5,48 +5,74 @@ import json import os import subprocess +import sys from pathlib import Path from typing import cast import pytest +from mutmut.utils.format_utils import get_mutant_name from scripts.mutation_results import main PROJECT = Path(__file__).parents[2] -def test_scoped_mutation_excludes_prefix_sibling_module( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch +def mutation_harness( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, + modules: list[str], + *, + real_python: bool = False, ) -> None: - source = tmp_path / "src/volcano_sdk/durable.py" - source.parent.mkdir(parents=True) - _ = source.write_text("def durable() -> bool: return True\n", encoding="utf-8") - sibling = source.with_name("durable_authoring.py") - _ = sibling.write_text("def authoring() -> bool: return True\n", encoding="utf-8") + """Install shell stubs for testing the native Mutmut selector.""" + for module in modules: + source = tmp_path / module + source.parent.mkdir(parents=True, exist_ok=True) + _ = source.write_text("def probe() -> bool: return True\n", encoding="utf-8") + _ = (tmp_path / "git-paths.bin").write_bytes( + b"\0".join(module.encode() for module in modules) + b"\0" + ) scripts = tmp_path / "scripts" scripts.mkdir() _ = (scripts / "mutation.sh").write_bytes( (PROJECT / "scripts/mutation.sh").read_bytes() ) + if real_python: + _ = (scripts / "mutation_results.py").write_bytes( + (PROJECT / "scripts/mutation_results.py").read_bytes() + ) bin_dir = tmp_path / "bin" bin_dir.mkdir() stubs = { - "git": """#!/bin/sh -case "$1" in - merge-base) printf 'base\\n' ;; - diff) printf 'src/volcano_sdk/durable.py\\0' ;; -esac -""", + "git": ( + "#!/bin/sh\n" + '[ "$*" = \'ls-files --cached --others --exclude-standard ' + "-z -- src/volcano_sdk' ] || exit 1\n" + "cat git-paths.bin\n" + ), "mutmut": "#!/bin/sh\nprintf '%s\\n' \"$@\" > mutation-args.txt\n", - "python": "#!/bin/sh\nexit 0\n", } for name, content in stubs.items(): stub = bin_dir / name _ = stub.write_text(content, encoding="utf-8") stub.chmod(0o755) + python = bin_dir / "python" + if real_python: + python.symlink_to(sys.executable) + else: + _ = python.write_text("#!/bin/sh\nexit 0\n", encoding="utf-8") + python.chmod(0o755) monkeypatch.setenv("PATH", f"{bin_dir}:{os.environ['PATH']}") - result = subprocess.run( + +def run_mutation_script(tmp_path: Path) -> subprocess.CompletedProcess[str]: + """Run the mutation orchestrator against the shell stubs. + + Returns: + Completed script process. + + """ + return subprocess.run( ["/bin/bash", "scripts/mutation.sh"], cwd=tmp_path, capture_output=True, @@ -54,10 +80,159 @@ def test_scoped_mutation_excludes_prefix_sibling_module( check=False, ) - assert result.returncode == 0, result.stdout + result.stderr - arguments = (tmp_path / "mutation-args.txt").read_text(encoding="utf-8") - assert "volcano_sdk.durable.x*" in arguments.splitlines() - assert "volcano_sdk.durable*" not in arguments.splitlines() + +def run_mutation_matrix(tmp_path: Path) -> subprocess.CompletedProcess[str]: + """Read the complete CI module plan. + + Returns: + Completed script process. + + """ + return subprocess.run( + ["/bin/bash", "scripts/mutation.sh", "--matrix"], + cwd=tmp_path, + capture_output=True, + text=True, + check=False, + ) + + +def mutant_pattern(module: str) -> str: + """Mirror Mutmut's package-init name normalization. + + Returns: + Exact function-prefix selector for the module. + + """ + name = module.removeprefix("src/").removesuffix(".py") + return f"{name.removesuffix('/__init__').replace('/', '.')}.x*" + + +def module_name(path: str) -> str: + """Return the Python module selected by the native Mutmut wildcard. + + Returns: + Import name for a handwritten Python source file. + + """ + return mutant_pattern(path).removesuffix(".x*") + + +@pytest.mark.parametrize( + ("path", "expected"), + [ + ("src/volcano_sdk/__init__.py", "volcano_sdk.x_probe__mutmut_1"), + ("src/volcano_sdk/nested/__init__.py", "volcano_sdk.nested.x_probe__mutmut_1"), + ], +) +def test_package_init_selector_matches_native_mutmut_name( + path: str, expected: str +) -> None: + assert get_mutant_name(Path(path), "x_probe__mutmut_1") == expected + assert expected.startswith(mutant_pattern(path).removesuffix("*")) + + +def test_mutation_matrix_covers_all_handwritten_modules( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + modules = [f"src/volcano_sdk/module_{index}.py" for index in range(8)] + modules.extend( + [ + "src/volcano_sdk/__init__.py", + "src/volcano_sdk/durable.py", + "src/volcano_sdk/durable_authoring.py", + "src/volcano_sdk/nested/cache.py", + "src/volcano_sdk/nested/__init__.py", + ] + ) + generated = "src/volcano_sdk/_generated/wire.py" + mutation_harness(tmp_path, monkeypatch, [*modules, generated]) + + plan = run_mutation_matrix(tmp_path) + assert plan.returncode == 0, plan.stderr + planned: object = cast("object", json.loads(plan.stdout)) + assert isinstance(planned, list) + values = cast("list[object]", planned) + matrix_modules: list[str] = [] + for name in values: + assert isinstance(name, str) + matrix_modules.append(name) + assert matrix_modules == [module_name(path) for path in modules] + + selected: list[str] = [] + selected_patterns: list[str] = [] + for module in matrix_modules: + monkeypatch.setenv("MUTATION_MODULE", module) + result = run_mutation_script(tmp_path) + assert result.returncode == 0, result.stderr + paths = (tmp_path / "reports/mutation-targets.bin").read_bytes().split(b"\0") + selected_modules = [os.fsdecode(path) for path in paths if path] + assert len(selected_modules) == 1 + selected.extend(selected_modules) + patterns = [mutant_pattern(path) for path in selected_modules] + selected_patterns.extend(patterns) + arguments = (tmp_path / "mutation-args.txt").read_text(encoding="utf-8") + assert arguments.splitlines() == ["run", "--max-children", "1", *patterns] + + assert sorted(selected) == sorted(modules) + assert len(selected) == len(set(selected)) + assert "volcano_sdk.durable.x*" in selected_patterns + assert "volcano_sdk.durable*" not in selected_patterns + assert "volcano_sdk.nested.x*" in selected_patterns + assert "volcano_sdk.x*" in selected_patterns + + +def test_empty_mutation_inventory_fails( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + mutation_harness(tmp_path, monkeypatch, []) + result = run_mutation_matrix(tmp_path) + assert result.returncode == 1 + assert "No handwritten SDK runtime modules found" in result.stderr + + +def test_duplicate_import_name_fails_inventory( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + mutation_harness( + tmp_path, + monkeypatch, + ["src/volcano_sdk/nested.py", "src/volcano_sdk/nested/__init__.py"], + ) + result = run_mutation_matrix(tmp_path) + assert result.returncode == 1 + assert "Duplicate Python runtime module" in result.stderr + + +def test_functionless_package_init_skips_native_no_match( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + source = "src/volcano_sdk/__init__.py" + mutation_harness(tmp_path, monkeypatch, [source], real_python=True) + _ = (tmp_path / source).write_text('"""Public exports."""\n', encoding="utf-8") + stale_metadata = tmp_path / "mutants" / f"{source}.meta" + stale_metadata.parent.mkdir(parents=True) + _ = stale_metadata.write_text('{"exit_code_by_key": {"stale": 1}}\n') + monkeypatch.setenv("MUTATION_MODULE", "volcano_sdk") + result = run_mutation_script(tmp_path) + assert result.returncode == 0, result.stderr + assert not (tmp_path / "mutation-args.txt").exists() + report = cast( + "object", json.loads((tmp_path / "reports/mutation.json").read_text()) + ) + assert isinstance(report, dict) + assert report["unmutatable_modules"] == [source] + + +@pytest.mark.parametrize("module", ["volcano_sdk.missing", "-1", "not-a-module"]) +def test_invalid_mutation_module_fails( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, module: str +) -> None: + mutation_harness(tmp_path, monkeypatch, ["src/volcano_sdk/probe.py"]) + monkeypatch.setenv("MUTATION_MODULE", module) + result = run_mutation_script(tmp_path) + assert result.returncode == 2 + assert "Unknown mutation module" in result.stderr def fixture_report(tmp_path: Path, code: int | None) -> tuple[Path, Path]: @@ -115,6 +290,25 @@ def test_killed_mutant_passes(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) - assert main(targets, failed) == 0 +def test_functionless_package_init_is_reported( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.chdir(tmp_path) + source = Path("src/volcano_sdk/__init__.py") + source.parent.mkdir(parents=True) + _ = source.write_text('"""Public package exports."""\n', encoding="utf-8") + targets = tmp_path / "targets.bin" + _ = targets.write_bytes(f"{source}\0".encode()) + failed = tmp_path / "failed.bin" + _ = failed.write_bytes(b"") + assert main(targets, failed) == 0 + report = cast( + "object", json.loads(Path("reports/mutation.json").read_text(encoding="utf-8")) + ) + assert isinstance(report, dict) + assert report["unmutatable_modules"] == [str(source)] + + def test_statically_invalid_mutant_is_reported_separately( tmp_path: Path, monkeypatch: pytest.MonkeyPatch ) -> None: diff --git a/tests/unit/test_realtime_fetch_cleanup.py b/tests/unit/test_realtime_fetch_cleanup.py new file mode 100644 index 00000000..60c3287c --- /dev/null +++ b/tests/unit/test_realtime_fetch_cleanup.py @@ -0,0 +1,198 @@ +from __future__ import annotations + +import asyncio + +import pytest +from test_realtime_fetch_lifecycle import cancel_operation +from test_realtime_fetch_worker import ( + BlockingRowFetch, + OutcomeRecorder, + RecordingBatchFetch, + fetch_job, +) +from typing_extensions import override + +from volcano_sdk._realtime_fetch_worker import ( + PostgresFetchJob, + PostgresFetchWorker, + _StopWorker, + _wait_for_close, +) + + +class DelayedCancellation: + def __init__(self) -> None: + self.started: asyncio.Event = asyncio.Event() + self.cancelled: asyncio.Event = asyncio.Event() + self.release_cleanup: asyncio.Event = asyncio.Event() + + async def wait(self) -> None: + self.started.set() + try: + _ = await asyncio.Event().wait() + except asyncio.CancelledError: + self.cancelled.set() + _ = await self.release_cleanup.wait() + raise + + +class CleanupQueue(asyncio.Queue[PostgresFetchJob[str] | _StopWorker]): + def __init__(self) -> None: + super().__init__(maxsize=1) + self.started: asyncio.Event = asyncio.Event() + self.cancelled: asyncio.Event = asyncio.Event() + self.release_cleanup: asyncio.Event = asyncio.Event() + + @override + async def put(self, item: PostgresFetchJob[str] | _StopWorker) -> None: + self.started.set() + try: + await super().put(item) + except asyncio.CancelledError: + self.cancelled.set() + _ = await self.release_cleanup.wait() + raise + + +async def fail_worker() -> None: + message = "delivery failed" + raise RuntimeError(message) + + +async def test_close_waits_for_cancelled_stop_cleanup() -> None: + stop = DelayedCancellation() + stop_task = asyncio.create_task(stop.wait()) + task = asyncio.create_task(fail_worker()) + closing = asyncio.create_task(_wait_for_close(task, stop_task)) + try: + _ = await asyncio.wait_for(stop.cancelled.wait(), timeout=1) + await asyncio.sleep(0) + assert not closing.done() + stop.release_cleanup.set() + with pytest.raises(RuntimeError, match="delivery failed"): + await asyncio.wait_for(closing, timeout=1) + assert stop_task.cancelled() + finally: + stop.release_cleanup.set() + await cancel_operation(closing) + await cancel_operation(stop_task) + await cancel_operation(task) + + +async def test_cancelled_enqueue_waits_for_its_queue_put_cleanup() -> None: + queue = CleanupQueue() + queue.put_nowait(fetch_job(1)) + worker = PostgresFetchWorker( + RecordingBatchFetch(), OutcomeRecorder(), queue_limit=1 + ) + worker._queue = queue + waiting = DelayedCancellation() + task = asyncio.create_task(waiting.wait()) + enqueueing = asyncio.create_task(worker._put_while_running(fetch_job(2), task)) + try: + _ = await asyncio.wait_for(queue.started.wait(), timeout=1) + _ = enqueueing.cancel() + _ = await asyncio.wait_for(queue.cancelled.wait(), timeout=1) + await asyncio.sleep(0) + assert not enqueueing.done() + queue.release_cleanup.set() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(enqueueing, timeout=1) + assert queue.get_nowait() == fetch_job(1) + assert queue.empty() + finally: + queue.release_cleanup.set() + waiting.release_cleanup.set() + if not queue.empty(): + _ = queue.get_nowait() + await cancel_operation(enqueueing) + await cancel_operation(task) + + +async def test_abort_cancels_an_outstanding_stop_request() -> None: + fetch = BlockingRowFetch() + worker = PostgresFetchWorker(fetch, OutcomeRecorder(), queue_limit=1) + closing = None + try: + await asyncio.wait_for(worker.enqueue(fetch_job(1)), timeout=1) + _ = await asyncio.wait_for(fetch.started.wait(), timeout=1) + await asyncio.wait_for(worker.enqueue(fetch_job(2)), timeout=1) + closing = asyncio.create_task(worker.close()) + await asyncio.sleep(0) + stop_task = worker._stop_task + assert stop_task is not None + assert not stop_task.done() + await cancel_operation(closing) + await asyncio.wait_for(worker.abort(), timeout=1) + assert stop_task.cancelled() + assert fetch.cancelled.is_set() + finally: + fetch.release.set() + await cancel_operation(closing) + await asyncio.wait_for(worker.abort(), timeout=1) + + +async def test_closed_worker_rejects_jobs_before_abort() -> None: + worker = PostgresFetchWorker( + RecordingBatchFetch(), OutcomeRecorder(), queue_limit=1 + ) + try: + await asyncio.wait_for(worker.close(), timeout=1) + with pytest.raises(RuntimeError, match="fetch worker is closed"): + await asyncio.wait_for(worker.enqueue(fetch_job(1)), timeout=1) + finally: + await asyncio.wait_for(worker.abort(), timeout=1) + + +async def test_repeated_close_completes_all_queued_tasks() -> None: + worker = PostgresFetchWorker( + RecordingBatchFetch(), + OutcomeRecorder(), + queue_limit=2, + batch_window_seconds=1, + max_batch_size=2, + ) + try: + await asyncio.wait_for(worker.enqueue(fetch_job(1)), timeout=1) + await asyncio.wait_for(worker.enqueue(fetch_job(2)), timeout=1) + await asyncio.wait_for(worker.close(), timeout=1) + await asyncio.wait_for(worker.close(), timeout=1) + await asyncio.wait_for(worker._queue.join(), timeout=1) + finally: + await asyncio.wait_for(worker.abort(), timeout=1) + + +async def test_batch_capacity_flushes_without_an_extra_row() -> None: + fetch = RecordingBatchFetch() + delivered = asyncio.Event() + + async def deliver(_outcome: object) -> None: + delivered.set() + + worker = PostgresFetchWorker[str]( + fetch, deliver, queue_limit=3, max_batch_size=2, batch_window_seconds=60 + ) + try: + await asyncio.wait_for(worker.enqueue(fetch_job(1)), timeout=1) + await asyncio.wait_for(worker.enqueue(fetch_job(2)), timeout=1) + _ = await asyncio.wait_for(delivered.wait(), timeout=1) + assert fetch.calls == [(1, 2)] + await asyncio.wait_for(worker.enqueue(fetch_job(3)), timeout=1) + await asyncio.wait_for(worker.close(), timeout=1) + assert fetch.calls == [(1, 2), (3,)] + finally: + await asyncio.wait_for(worker.abort(), timeout=1) + + +async def test_expired_batch_deadline_preserves_queued_row( + monkeypatch: pytest.MonkeyPatch, +) -> None: + worker = PostgresFetchWorker( + RecordingBatchFetch(), OutcomeRecorder(), queue_limit=1 + ) + worker._queue.put_nowait(fetch_job(1)) + loop = asyncio.get_running_loop() + with monkeypatch.context() as scoped: + scoped.setattr(loop, "time", lambda: 42.0) + assert await worker._next_before(42.0) is None + assert worker._queue.get_nowait() == fetch_job(1) diff --git a/tests/unit/test_realtime_fetch_lifecycle.py b/tests/unit/test_realtime_fetch_lifecycle.py index 9f06716f..669ac63d 100644 --- a/tests/unit/test_realtime_fetch_lifecycle.py +++ b/tests/unit/test_realtime_fetch_lifecycle.py @@ -51,11 +51,11 @@ async def test_unused_worker_can_close_and_abort_repeatedly() -> None: await worker.close() await worker.close() - await worker.abort() - await worker.abort() + await asyncio.wait_for(worker.abort(), timeout=1) + await asyncio.wait_for(worker.abort(), timeout=1) with pytest.raises(RuntimeError, match="fetch worker is closed"): - await worker.enqueue(fetch_job(1)) + await asyncio.wait_for(worker.enqueue(fetch_job(1)), timeout=1) assert fetch.calls == [] assert deliver.items == [] @@ -70,9 +70,9 @@ async def fail_delivery(_outcome: PostgresFetchOutcome[str]) -> None: worker = PostgresFetchWorker(fetch, fail_delivery, queue_limit=1) closing = None try: - await worker.enqueue(fetch_job(1)) + await asyncio.wait_for(worker.enqueue(fetch_job(1)), timeout=1) _ = await asyncio.wait_for(fetch.started.wait(), timeout=1) - await worker.enqueue(fetch_job(2)) + await asyncio.wait_for(worker.enqueue(fetch_job(2)), timeout=1) closing = asyncio.create_task(worker.close()) await asyncio.sleep(0) stop_task = worker._stop_task @@ -87,7 +87,7 @@ async def fail_delivery(_outcome: PostgresFetchOutcome[str]) -> None: assert stop_task.cancelled() assert fetch.row_ids == [1] finally: - await worker.abort() + await asyncio.wait_for(worker.abort(), timeout=1) await cancel_operation(closing) @@ -97,9 +97,9 @@ async def test_abort_unblocks_close_with_a_full_queue() -> None: worker = PostgresFetchWorker(fetch, deliver, queue_limit=1) closing = None try: - await worker.enqueue(fetch_job(1)) + await asyncio.wait_for(worker.enqueue(fetch_job(1)), timeout=1) _ = await asyncio.wait_for(fetch.started.wait(), timeout=1) - await worker.enqueue(fetch_job(2)) + await asyncio.wait_for(worker.enqueue(fetch_job(2)), timeout=1) closing = asyncio.create_task(worker.close()) await asyncio.sleep(0) stop_task = worker._stop_task @@ -115,7 +115,7 @@ async def test_abort_unblocks_close_with_a_full_queue() -> None: assert fetch.cancelled.is_set() assert deliver.items == [] finally: - await worker.abort() + await asyncio.wait_for(worker.abort(), timeout=1) await cancel_operation(closing) @@ -125,9 +125,9 @@ async def test_abort_rejects_an_enqueue_waiting_for_capacity() -> None: worker = PostgresFetchWorker(fetch, deliver, queue_limit=1) enqueueing = None try: - await worker.enqueue(fetch_job(1)) + await asyncio.wait_for(worker.enqueue(fetch_job(1)), timeout=1) _ = await asyncio.wait_for(fetch.started.wait(), timeout=1) - await worker.enqueue(fetch_job(2)) + await asyncio.wait_for(worker.enqueue(fetch_job(2)), timeout=1) enqueueing = asyncio.create_task(worker.enqueue(fetch_job(3))) await asyncio.sleep(0) assert not enqueueing.done() @@ -139,29 +139,32 @@ async def test_abort_rejects_an_enqueue_waiting_for_capacity() -> None: assert fetch.cancelled.is_set() assert deliver.items == [] finally: - await worker.abort() + await asyncio.wait_for(worker.abort(), timeout=1) await cancel_operation(enqueueing) -async def test_wrong_batch_result_count_fails_before_delivery() -> None: +@pytest.mark.parametrize("records", [(), ({"id": 1}, {"id": 2})]) +async def test_wrong_batch_result_count_fails_before_delivery( + records: tuple[dict[str, int], ...], +) -> None: release = asyncio.Event() - async def no_results( + async def malformed_results( _requests: tuple[_PostgresFetchRequest, ...], ) -> tuple[dict[str, int], ...]: _ = await release.wait() - return () + return records deliver = OutcomeRecorder() - worker = PostgresFetchWorker(no_results, deliver, queue_limit=1) + worker = PostgresFetchWorker(malformed_results, deliver, queue_limit=1) try: - await worker.enqueue(fetch_job(1)) + await asyncio.wait_for(worker.enqueue(fetch_job(1)), timeout=1) release.set() with pytest.raises(RuntimeError, match="unexpected result count"): await asyncio.wait_for(worker.close(), timeout=1) assert deliver.items == [] finally: - await worker.abort() + await asyncio.wait_for(worker.abort(), timeout=1) async def test_fetch_cancellation_propagates_without_delivering_a_fallback() -> None: @@ -178,7 +181,7 @@ async def cancelled_fetch( deliver = OutcomeRecorder() worker = PostgresFetchWorker(cancelled_fetch, deliver, queue_limit=1) try: - await worker.enqueue(fetch_job(1)) + await asyncio.wait_for(worker.enqueue(fetch_job(1)), timeout=1) _ = await asyncio.wait_for(started.wait(), timeout=1) release.set() @@ -187,7 +190,7 @@ async def cancelled_fetch( assert deliver.items == [] finally: - await worker.abort() + await asyncio.wait_for(worker.abort(), timeout=1) async def test_batch_window_flushes_without_waiting_for_another_row_or_close() -> None: @@ -202,13 +205,13 @@ async def deliver(outcome: PostgresFetchOutcome[str]) -> None: fetch, deliver, queue_limit=2, max_batch_size=2, batch_window_seconds=0.1 ) try: - await worker.enqueue(fetch_job(1)) + await asyncio.wait_for(worker.enqueue(fetch_job(1)), timeout=1) _ = await asyncio.wait_for(delivered.wait(), timeout=1) await asyncio.wait_for(worker.close(), timeout=1) assert fetch.calls == [(1,)] finally: - await worker.abort() + await asyncio.wait_for(worker.abort(), timeout=1) @pytest.mark.parametrize(("window", "expected"), [(0, [(1,), (2,)]), (60, [(1, 2)])]) @@ -221,12 +224,44 @@ async def test_batch_capacity_and_zero_window_preserve_delivery_order( fetch, deliver, queue_limit=2, max_batch_size=2, batch_window_seconds=window ) try: - await worker.enqueue(fetch_job(1)) - await worker.enqueue(fetch_job(2)) + await asyncio.wait_for(worker.enqueue(fetch_job(1)), timeout=1) + await asyncio.wait_for(worker.enqueue(fetch_job(2)), timeout=1) await asyncio.wait_for(worker.close(), timeout=1) - await worker.abort() + await asyncio.wait_for(worker.abort(), timeout=1) assert fetch.calls == expected assert [item.record for item in deliver.items] == [{"id": 1}, {"id": 2}] finally: - await worker.abort() + await asyncio.wait_for(worker.abort(), timeout=1) + + +async def test_zero_batch_window_keeps_queued_followup_fetches_separate() -> None: + first_fetch_started = asyncio.Event() + release_first_fetch = asyncio.Event() + recording_fetch = RecordingBatchFetch() + + async def fetch( + requests: tuple[_PostgresFetchRequest, ...], + ) -> tuple[dict[str, int], ...]: + if requests[0].row_id == 1: + first_fetch_started.set() + _ = await release_first_fetch.wait() + return await recording_fetch(requests) + + worker = PostgresFetchWorker( + fetch, + OutcomeRecorder(), + queue_limit=3, + max_batch_size=3, + ) + try: + await asyncio.wait_for(worker.enqueue(fetch_job(1)), timeout=1) + _ = await asyncio.wait_for(first_fetch_started.wait(), timeout=0.2) + await asyncio.wait_for(worker.enqueue(fetch_job(2)), timeout=1) + await asyncio.wait_for(worker.enqueue(fetch_job(3)), timeout=1) + release_first_fetch.set() + await asyncio.wait_for(worker.close(), timeout=1) + + assert recording_fetch.calls == [(1,), (2,), (3,)] + finally: + await asyncio.wait_for(worker.abort(), timeout=1) diff --git a/tests/unit/test_realtime_fetch_worker.py b/tests/unit/test_realtime_fetch_worker.py index a7effb92..006bf7da 100644 --- a/tests/unit/test_realtime_fetch_worker.py +++ b/tests/unit/test_realtime_fetch_worker.py @@ -96,6 +96,7 @@ def passthrough_job(name: str) -> PostgresFetchJob[str]: return PostgresFetchJob(request=None, fallback=name) +@pytest.mark.order(0) def test_postgres_fetch_worker_bounds_and_orders_fetches() -> None: async def scenario() -> None: fetch = BlockingRowFetch() @@ -103,9 +104,9 @@ async def scenario() -> None: outcomes = deliver.items worker = PostgresFetchWorker(fetch, deliver, queue_limit=1) - await worker.enqueue(fetch_job(1)) + await asyncio.wait_for(worker.enqueue(fetch_job(1)), timeout=1) _ = await asyncio.wait_for(fetch.started.wait(), timeout=0.2) - await worker.enqueue(fetch_job(2)) + await asyncio.wait_for(worker.enqueue(fetch_job(2)), timeout=1) third_enqueue = asyncio.create_task(worker.enqueue(fetch_job(3))) await asyncio.sleep(0) @@ -143,9 +144,9 @@ async def scenario() -> None: batch_window_seconds=0.05, max_batch_size=3, ) - await worker.enqueue(fetch_job(1)) - await worker.enqueue(fetch_job(2)) - await worker.enqueue(fetch_job(3)) + await asyncio.wait_for(worker.enqueue(fetch_job(1)), timeout=1) + await asyncio.wait_for(worker.enqueue(fetch_job(2)), timeout=1) + await asyncio.wait_for(worker.enqueue(fetch_job(3)), timeout=1) await asyncio.wait_for(worker.close(), timeout=0.2) assert fetch.calls == [(1, 2, 3)] @@ -176,8 +177,8 @@ async def scenario() -> None: batch_window_seconds=0.01, max_batch_size=2, ) - await worker.enqueue(fetch_job(1)) - await worker.enqueue(fetch_job(2, table="archive")) + await asyncio.wait_for(worker.enqueue(fetch_job(1)), timeout=1) + await asyncio.wait_for(worker.enqueue(fetch_job(2, table="archive")), timeout=1) await asyncio.wait_for(worker.close(), timeout=0.2) assert fetch.calls == [(1,), (2,)] @@ -202,8 +203,8 @@ async def scenario() -> None: batch_window_seconds=0.01, max_batch_size=2, ) - await worker.enqueue(fetch_job(1)) - await worker.enqueue(fetch_job(1)) + await asyncio.wait_for(worker.enqueue(fetch_job(1)), timeout=1) + await asyncio.wait_for(worker.enqueue(fetch_job(1)), timeout=1) await asyncio.wait_for(worker.close(), timeout=0.2) assert fetch.calls == [(1,), (1,)] @@ -219,9 +220,11 @@ async def scenario() -> None: outcomes = deliver.items worker = PostgresFetchWorker(fetch, deliver, queue_limit=1) - await worker.enqueue(fetch_job(1)) + await asyncio.wait_for(worker.enqueue(fetch_job(1)), timeout=1) _ = await asyncio.wait_for(fetch.started.wait(), timeout=0.2) - await worker.enqueue(passthrough_job("full-payload")) + await asyncio.wait_for( + worker.enqueue(passthrough_job("full-payload")), timeout=1 + ) assert outcomes == [] @@ -237,6 +240,7 @@ async def scenario() -> None: asyncio.run(scenario()) +@pytest.mark.order(0) def test_postgres_fetch_worker_aborts_obsolete_jobs_without_waiting() -> None: async def scenario() -> None: fetch = BlockingRowFetch() @@ -244,9 +248,9 @@ async def scenario() -> None: outcomes = deliver.items worker = PostgresFetchWorker(fetch, deliver, queue_limit=1) - await worker.enqueue(fetch_job(1)) + await asyncio.wait_for(worker.enqueue(fetch_job(1)), timeout=1) _ = await asyncio.wait_for(fetch.started.wait(), timeout=0.2) - await worker.enqueue(fetch_job(2)) + await asyncio.wait_for(worker.enqueue(fetch_job(2)), timeout=1) await asyncio.wait_for(worker.abort(), timeout=0.2) @@ -274,8 +278,8 @@ async def scenario() -> None: deliver, queue_limit=2, ) - await worker.enqueue(fetch_job(1)) - await worker.enqueue(fetch_job(2)) + await asyncio.wait_for(worker.enqueue(fetch_job(1)), timeout=1) + await asyncio.wait_for(worker.enqueue(fetch_job(2)), timeout=1) await asyncio.wait_for(worker.close(), timeout=0.2) assert outcomes == [ @@ -302,15 +306,16 @@ async def fail_delivery(_outcome: PostgresFetchOutcome[str]) -> None: fail_delivery, queue_limit=1, ) - await worker.enqueue(fetch_job(2)) + await asyncio.wait_for(worker.enqueue(fetch_job(2)), timeout=1) _ = await delivery_started.wait() release_delivery.set() await asyncio.sleep(0) with pytest.raises(RuntimeError) as raised: - await worker.enqueue(fetch_job(3)) + await asyncio.wait_for(worker.enqueue(fetch_job(3)), timeout=1) assert raised.value is failure + assert worker._queue.empty() asyncio.run(scenario()) @@ -322,9 +327,9 @@ async def scenario() -> None: outcomes = deliver.items worker = PostgresFetchWorker(fetch, deliver, queue_limit=1) - await worker.enqueue(fetch_job(1)) + await asyncio.wait_for(worker.enqueue(fetch_job(1)), timeout=1) _ = await asyncio.wait_for(fetch.started.wait(), timeout=0.2) - await worker.enqueue(fetch_job(2)) + await asyncio.wait_for(worker.enqueue(fetch_job(2)), timeout=1) interrupted_close = asyncio.create_task(worker.close()) await asyncio.sleep(0) @@ -362,9 +367,9 @@ async def fail_delivery(_outcome: PostgresFetchOutcome[str]) -> None: fail_delivery, queue_limit=1, ) - await worker.enqueue(fetch_job(2)) + await asyncio.wait_for(worker.enqueue(fetch_job(2)), timeout=1) _ = await delivery_started.wait() - await worker.enqueue(fetch_job(3)) + await asyncio.wait_for(worker.enqueue(fetch_job(3)), timeout=1) blocked_enqueue = asyncio.create_task(worker.enqueue(fetch_job(4))) await asyncio.sleep(0) assert not blocked_enqueue.done() diff --git a/tests/unit/test_storage_boundaries.py b/tests/unit/test_storage_boundaries.py index dbf2a5a0..a44194c2 100644 --- a/tests/unit/test_storage_boundaries.py +++ b/tests/unit/test_storage_boundaries.py @@ -31,6 +31,28 @@ from volcano_sdk.storage import StorageBucket +class EndOfUploadSource: + """Record the terminal read and fail if an upload retries EOF.""" + + def __init__(self) -> None: + self.sizes: list[int] = [] + + def read(self, size: int = -1, /) -> bytes: + if self.sizes: + message = "upload source read after end of stream" + raise AssertionError(message) + self.sizes.append(size) + return b"" + + +@pytest.mark.order(0) +def test_upload_part_stops_reading_at_end_of_stream() -> None: + source = EndOfUploadSource() + + assert _read_upload_part(source, 4) == b"" + assert source.sizes == [4] + + class UnavailableRawStream(RawIOBase): @override def readable(self) -> bool: diff --git a/tests/unit/test_token_bootstrap.py b/tests/unit/test_token_bootstrap.py index 90144fba..c764f22c 100644 --- a/tests/unit/test_token_bootstrap.py +++ b/tests/unit/test_token_bootstrap.py @@ -84,18 +84,22 @@ def make_transport(*, api_url: str, timeout: float) -> GeneratedTransport: assert observed == [("https://api.volcano.dev", 60.0)] -def test_token_bootstrap_is_local_until_profile_validation() -> None: +@pytest.mark.order(0) +@pytest.mark.parametrize("refresh_token", [None, "supplied-refresh"]) +def test_token_bootstrap_is_local_until_profile_validation( + refresh_token: str | None, +) -> None: requests: list[httpx.Request] = [] def handle(request: httpx.Request) -> httpx.Response: requests.append(request) return httpx.Response(200, json={"user": PROFILE}) - client = token_client(handle) + client = token_client(handle, refresh_token=refresh_token) initial = client.auth.get_session() assert initial is not None assert initial.access_token == "supplied-access" - assert initial.refresh_token is None + assert initial.refresh_token == refresh_token assert initial.user_id is None assert initial.user is None assert not requests @@ -105,7 +109,7 @@ def handle(request: httpx.Request) -> httpx.Response: assert current.user_id == USER_ID assert current.user == PROFILE assert current.access_token == initial.access_token - assert current.refresh_token is None + assert current.refresh_token == refresh_token assert initial.user_id is None assert requests[0].headers["authorization"] == "Bearer supplied-access"