Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 6 additions & 1 deletion .github/workflows/ci.yml
Original file line number Diff line number Diff line change
@@ -1,3 +1,3 @@
name: CI

on:
Expand All @@ -22,7 +22,12 @@
# (add a shard).
#
# `conftest.py` assigns whole test *files* to shards, deterministically,
# from the collection alone — nothing is exchanged between these jobs.
# from the collection and the measured seconds per file in
# `tests/shard_seconds.json` — nothing is exchanged between these jobs.
# Balancing item counts instead left one shard at 13 of its 15 minutes on
# `main` while another took 8, and any new test file reshuffled the rest
# (#904). When a shard nears the cap, re-measure with
# `scripts/measure_shard_seconds.py` before adding a shard.
# `tests/test_shard_partition.py` asserts the union of the shards is the
# suite, because a partition that silently drops a file leaves every job
# green while a test stops running.
Expand Down
70 changes: 60 additions & 10 deletions ci_sharding.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,17 +8,50 @@
from anywhere and has one definition.

Pure. It reads nothing, writes nothing, and takes no environment: the caller
supplies the collection and gets back an assignment. That is what lets
``tests/test_shard_partition.py`` assert the properties directly.
supplies the collection (and the measured seconds, when it has them) and gets
back an assignment. That is what lets ``tests/test_shard_partition.py`` assert
the properties directly.
"""

from __future__ import annotations

import json
import math
from collections.abc import Mapping
from pathlib import Path

#: Measured seconds per test file, written by
#: ``scripts/measure_shard_seconds.py`` from a full ``--junitxml`` run.
SECONDS_FILE = Path(__file__).resolve().parent / "tests" / "shard_seconds.json"

def shard_assignment(paths: Mapping[str, int], shards: int) -> dict[str, int]:
"""Assign whole test *files* to shards, balancing collected item counts.

def load_seconds(path: Path = SECONDS_FILE) -> dict[str, float]:
"""The measured seconds per file, or nothing when there is no measurement.

A missing or unreadable file balances on item counts alone, as before:
the measurement only ever improves the balance, never the coverage.
"""

try:
payload = json.loads(path.read_text(encoding="utf-8"))
except (OSError, ValueError):
return {}
files = payload.get("files") if isinstance(payload, dict) else None
if not isinstance(files, dict):
return {}
return {
str(name): float(value)
for name, value in files.items()
if isinstance(value, int | float) and not isinstance(value, bool) and value >= 0
}


def shard_assignment(
paths: Mapping[str, int],
shards: int,
seconds: Mapping[str, float] | None = None,
) -> dict[str, int]:
"""Assign whole test *files* to shards, balancing their expected time.

Two properties, and both are load-bearing.

Expand All @@ -32,15 +65,32 @@ def shard_assignment(paths: Mapping[str, int], shards: int) -> dict[str, int]:
is exchanged between jobs, so the union of the shards is exactly the suite
and no test can fall between two of them — which a test asserts.

Item count is a proxy for time, and an imperfect one; it is used because it
is free and needs no stored measurements to go stale. Balance is checked in
``tests/test_shard_partition.py`` rather than assumed.
**Measured time where there is one.** Item count alone was a poor proxy:
a file of forty git-fixture tests costs more than a file of four hundred
pure ones. Balancing counts left one shard at 13 of its 15 minutes on
``main``, while another took 8. Adding any test file reshuffled most files
between shards, so an unrelated PR could tip a shard past its cap (#904).
``seconds`` holds each file's measured time. A file without a measurement
(new, or renamed since) costs its item count times the measured seconds
per item. A stale measurement only unbalances; it never drops a file.
Balance is checked in ``tests/test_shard_partition.py`` rather than
assumed.
"""

load = [0] * shards
cost = _costs(paths, seconds or {})
load = [0.0] * shards
owner: dict[str, int] = {}
for path, count in sorted(paths.items(), key=lambda item: (-item[1], item[0])):
for path in sorted(paths, key=lambda item: (-cost[item], item)):
target = min(range(shards), key=lambda index: (load[index], index))
owner[path] = target
load[target] += count
load[target] += cost[path]
return owner


def _costs(paths: Mapping[str, int], seconds: Mapping[str, float]) -> dict[str, float]:
known = {path: seconds[path] for path in sorted(paths) if path in seconds}
known_items = sum(paths[path] for path in known)
# ``fsum`` over a sorted order: every shard computes the same rate, bit
# for bit, whatever order its collection listed the files in.
rate = math.fsum(known.values()) / known_items if known_items else 1.0
return {path: known.get(path, paths[path] * rate) for path in paths}
4 changes: 2 additions & 2 deletions conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,7 @@

import pytest # noqa: E402

from ci_sharding import shard_assignment # noqa: E402
from ci_sharding import load_seconds, shard_assignment # noqa: E402


@pytest.fixture(autouse=True)
Expand Down Expand Up @@ -90,7 +90,7 @@ def pytest_collection_modifyitems(config, items) -> None: # noqa: ANN001
counts: dict[str, int] = {}
for item in items:
counts[item.location[0]] = counts.get(item.location[0], 0) + 1
owner = shard_assignment(counts, shards)
owner = shard_assignment(counts, shards, load_seconds())
keep = [item for item in items if owner[item.location[0]] == index - 1]
dropped = [item for item in items if owner[item.location[0]] != index - 1]
if not keep:
Expand Down
78 changes: 78 additions & 0 deletions scripts/measure_shard_seconds.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,78 @@
"""Write ``tests/shard_seconds.json`` from a full-suite ``--junitxml`` report.

The CI suite is split into shards by measured time per test file
(``ci_sharding.py``). Re-measure when a shard nears its ``timeout-minutes``:

python -m pytest -n auto -m "not perf" --ignore=tests/test_adapter_static_only.py \\
--junitxml=junit.xml
python scripts/measure_shard_seconds.py junit.xml

The times are relative weights: what matters is how the files compare with
each other, so one machine's measurement balances another's runners.
"""

from __future__ import annotations

import argparse
import json
import subprocess
import sys
import xml.etree.ElementTree as ET
from collections import defaultdict
from pathlib import Path

REPO_ROOT = Path(__file__).resolve().parent.parent
OUTPUT = REPO_ROOT / "tests" / "shard_seconds.json"


def file_seconds(report: Path, root: Path = REPO_ROOT) -> dict[str, float]:
"""Seconds per test file: each test case's setup, call and teardown summed."""

totals: dict[str, float] = defaultdict(float)
for case in ET.parse(report).getroot().iter("testcase"):
name = case.get("file") or _file_from_classname(case.get("classname", ""), root)
if name:
totals[name] += float(case.get("time") or 0.0)
return {name: round(value, 1) for name, value in sorted(totals.items())}


def _file_from_classname(classname: str, root: Path) -> str | None:
"""``tests.test_x.TestY`` → ``tests/test_x.py``, the file the case is in."""

parts = classname.split(".")
for end in range(len(parts), 0, -1):
candidate = root.joinpath(*parts[:end]).with_suffix(".py")
if candidate.is_file():
return candidate.relative_to(root).as_posix()
return None


def main(argv: list[str] | None = None) -> int:
parser = argparse.ArgumentParser(description=__doc__.splitlines()[0])
parser.add_argument("report", type=Path, help="a --junitxml report of the full CI suite")
parser.add_argument(
"--root",
type=Path,
default=REPO_ROOT,
help="the checkout the report was made in (default: this one)",
)
args = parser.parse_args(argv)
files = file_seconds(args.report, args.root.resolve())
if not files:
print(f"{args.report} holds no test cases", file=sys.stderr)
return 1
commit = subprocess.run(
["git", "rev-parse", "--short=12", "HEAD"],
cwd=args.root,
capture_output=True,
text=True,
check=False,
).stdout.strip()
payload = {"measured_at": commit or None, "files": files}
OUTPUT.write_text(json.dumps(payload, indent=1, sort_keys=True) + "\n", encoding="utf-8")
print(f"wrote {len(files)} files, {sum(files.values()):.0f}s in total, to {OUTPUT}")
return 0


if __name__ == "__main__":
raise SystemExit(main())
Loading
Loading