From 63140f9ac77cb4732819c44d865f421ea310fce5 Mon Sep 17 00:00:00 2001 From: Anionix Date: Fri, 24 Jul 2026 16:00:12 +0900 Subject: [PATCH] feat: type adapter workload contracts --- pyproject.toml | 1 + src/format_bench/adapter_contract.py | 3 +- src/format_bench/fair.py | 8 ++-- src/format_bench/formats/__init__.py | 18 +++++++++ src/format_bench/model.py | 32 +++++++++++---- src/format_bench/workload_contract.py | 41 +++++++++++++++++++ src/format_bench/workloads.py | 2 + tests/test_adapter_contract.py | 39 +++++++++++++++++- tests/test_workloads.py | 57 +++++++++++++++------------ 9 files changed, 162 insertions(+), 39 deletions(-) create mode 100644 src/format_bench/workload_contract.py diff --git a/pyproject.toml b/pyproject.toml index f07f36a..8f6d058 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -70,6 +70,7 @@ include = [ "src/format_bench/model.py", "src/format_bench/nyc_snapshot.py", "src/format_bench/research.py", + "src/format_bench/workload_contract.py", "src/format_bench/workflow_contract.py", "src/format_bench/robustness/evidence.py", "src/format_bench/robustness/mutations.py", diff --git a/src/format_bench/adapter_contract.py b/src/format_bench/adapter_contract.py index 6dc9e78..7aad730 100644 --- a/src/format_bench/adapter_contract.py +++ b/src/format_bench/adapter_contract.py @@ -1,6 +1,7 @@ from typing import Literal, NotRequired, TypeAlias, TypedDict from .contracts import NormalizedColumn +from .workload_contract import WorkloadDeclarations class AdapterColumn(TypedDict): @@ -17,7 +18,7 @@ class AdapterManifest(TypedDict): columns: AdapterColumns canonical_hash: str expected_counts: dict[str, int] - workloads: NotRequired[dict[str, object]] + workloads: NotRequired[WorkloadDeclarations] class VerificationResult(TypedDict): diff --git a/src/format_bench/fair.py b/src/format_bench/fair.py index 796a045..9b73a08 100644 --- a/src/format_bench/fair.py +++ b/src/format_bench/fair.py @@ -8,7 +8,7 @@ import pyarrow.dataset as ds from .canonical import canonical_hash, order_insensitive_hash -from .model import WorkloadKind +from .model import WorkloadKind, WorkloadSpec from .workloads import apply_workload, expected_workload_rows, load_workloads @@ -57,7 +57,7 @@ def _operation_name(operation: Operation) -> str: def workload_for( operation: Operation, manifest: Mapping[str, object] | None = None -): +) -> WorkloadSpec: return load_workloads(manifest or {})[_operation_name(operation)] @@ -70,7 +70,7 @@ def columns_for( def arrow_filter( operation: Operation, manifest: Mapping[str, object] | None = None -): +) -> ds.Expression | None: spec = workload_for(operation, manifest) if spec.kind.value != "filter": return None @@ -117,7 +117,7 @@ def apply_arrow( return apply_workload(table, workload_for(operation, manifest)) -def expected_rows(operation: Operation, manifest: dict) -> int: +def expected_rows(operation: Operation, manifest: Mapping[str, object]) -> int: workloads = load_workloads(manifest) name = _operation_name(operation) spec = workloads[name] diff --git a/src/format_bench/formats/__init__.py b/src/format_bench/formats/__init__.py index 388b7b2..64fc821 100644 --- a/src/format_bench/formats/__init__.py +++ b/src/format_bench/formats/__init__.py @@ -4,6 +4,16 @@ AdapterManifest, VerificationResult, ) +from format_bench.workload_contract import ( + ComparisonOperator, + FilterWorkload, + HeadWorkload, + ProjectionWorkload, + ReadAllWorkload, + WorkloadDeclaration, + WorkloadDeclarations, + WorkloadScalar, +) from .base import ( Artifact, @@ -28,16 +38,21 @@ "ArrowIpcAdapter", "AvroAdapter", "CborAdapter", + "ComparisonOperator", "CsvAdapter", "DuckDbAdapter", "FeatherV2Adapter", + "FilterWorkload", "FormatAdapter", "FormatDescription", + "HeadWorkload", "LanceAdapter", "MessagePackAdapter", "ObjectJsonlAdapter", "OrcAdapter", "ParquetAdapter", + "ProjectionWorkload", + "ReadAllWorkload", "SqliteAdapter", "build_fts", "lance_components", @@ -46,4 +61,7 @@ "TsvAdapter", "VortexAdapter", "VerificationResult", + "WorkloadDeclaration", + "WorkloadDeclarations", + "WorkloadScalar", ] diff --git a/src/format_bench/model.py b/src/format_bench/model.py index d6c0d53..ed234ec 100644 --- a/src/format_bench/model.py +++ b/src/format_bench/model.py @@ -7,6 +7,12 @@ from types import MappingProxyType from typing import cast +from .workload_contract import ( + ComparisonOperator, + WorkloadScalar, + is_comparison_operator, +) + class Lane(StrEnum): FAIR = "fair" @@ -153,8 +159,8 @@ class WorkloadSpec: kind: WorkloadKind columns: tuple[str, ...] = () column: str | None = None - operator: str | None = None - value: object | None = None + operator: ComparisonOperator | None = None + value: WorkloadScalar | None = None limit: int | None = None expected_rows: int | None = None @@ -175,11 +181,19 @@ def from_mapping(cls, operation: str, payload: Mapping[str, object]) -> "Workloa raise ValueError(f"workload columns for {operation} must contain strings") columns = tuple(cast(list[str] | tuple[str, ...], untyped_columns)) column = payload.get("column") - operator = payload.get("operator") + raw_operator = payload.get("operator") if column is not None and not isinstance(column, str): raise ValueError("workload column must be a string") - if operator is not None and not isinstance(operator, str): - raise ValueError("workload operator must be a string") + operator: ComparisonOperator | None = None + if raw_operator is not None: + if not isinstance(raw_operator, str): + raise ValueError("workload operator must be a string") + if not is_comparison_operator(raw_operator): + raise ValueError("workload operator must be supported") + operator = raw_operator + value = payload.get("value") + if value is not None and not isinstance(value, (str, int, float, bool)): + raise ValueError("workload filter value must be a scalar") expected = payload.get("expected_rows") limit = payload.get("limit") @@ -196,7 +210,7 @@ def optional_int(value: object, field: str) -> int | None: columns=columns, column=column, operator=operator, - value=payload.get("value"), + value=value, limit=optional_int(limit, "limit"), expected_rows=optional_int(expected, "expected_rows"), ) @@ -209,7 +223,11 @@ def validate(self) -> None: if self.kind is WorkloadKind.PROJECTION and not self.columns: raise ValueError(f"projection workload {self.operation} needs columns") if self.kind is WorkloadKind.FILTER: - if not self.column or self.operator not in {"eq", "gt", "gte", "lt", "lte"}: + if ( + not self.column + or self.operator is None + or self.value is None + ): raise ValueError(f"filter workload {self.operation} needs a supported predicate") if self.kind is WorkloadKind.HEAD and (self.limit is None or self.limit <= 0): raise ValueError(f"head workload {self.operation} needs a positive limit") diff --git a/src/format_bench/workload_contract.py b/src/format_bench/workload_contract.py new file mode 100644 index 0000000..c5e4098 --- /dev/null +++ b/src/format_bench/workload_contract.py @@ -0,0 +1,41 @@ +from typing import Literal, NotRequired, TypeAlias, TypeGuard, TypedDict + + +ComparisonOperator: TypeAlias = Literal["eq", "gt", "gte", "lt", "lte"] +WorkloadScalar: TypeAlias = str | int | float | bool +_COMPARISON_OPERATORS = frozenset({"eq", "gt", "gte", "lt", "lte"}) + + +def is_comparison_operator(value: object) -> TypeGuard[ComparisonOperator]: + return isinstance(value, str) and value in _COMPARISON_OPERATORS + + +class ReadAllWorkload(TypedDict): + kind: Literal["read_all"] + expected_rows: NotRequired[int] + + +class ProjectionWorkload(TypedDict): + kind: Literal["projection"] + columns: list[str] + expected_rows: NotRequired[int] + + +class FilterWorkload(TypedDict): + kind: Literal["filter"] + column: str + operator: ComparisonOperator + value: WorkloadScalar + expected_rows: NotRequired[int] + + +class HeadWorkload(TypedDict): + kind: Literal["head"] + limit: int + expected_rows: NotRequired[int] + + +WorkloadDeclaration: TypeAlias = ( + ReadAllWorkload | ProjectionWorkload | FilterWorkload | HeadWorkload +) +WorkloadDeclarations: TypeAlias = dict[str, WorkloadDeclaration] diff --git a/src/format_bench/workloads.py b/src/format_bench/workloads.py index b5f80ae..3319278 100644 --- a/src/format_bench/workloads.py +++ b/src/format_bench/workloads.py @@ -62,6 +62,8 @@ def load_workloads(manifest: Mapping[str, object]) -> dict[str, WorkloadSpec]: return _legacy_workloads() if not isinstance(raw, Mapping): raise ValueError("manifest workloads must be an object") + # LLM contract: DISCOVERED -> ENCODED -> ROUNDTRIP_VERIFIED -> BENCHMARKED -> REPORTED. + # Parsing validates an ENCODED declaration; it does not advance evidence state. workloads: dict[str, WorkloadSpec] = {} for operation, payload in raw.items(): operation, normalized_payload = normalized_workload_entry(operation, payload) diff --git a/tests/test_adapter_contract.py b/tests/test_adapter_contract.py index ead299d..ed2d689 100644 --- a/tests/test_adapter_contract.py +++ b/tests/test_adapter_contract.py @@ -1,14 +1,23 @@ import tomllib from pathlib import Path -from typing import NotRequired, get_args, get_origin, get_type_hints +from typing import Literal, NotRequired, get_args, get_origin, get_type_hints from format_bench.contracts import NormalizedColumn from format_bench.formats import ( AdapterColumn, AdapterManifest, + ComparisonOperator, + FilterWorkload, FormatAdapter, + HeadWorkload, + ProjectionWorkload, + ReadAllWorkload, VerificationResult, + WorkloadDeclaration, + WorkloadDeclarations, + WorkloadScalar, ) +from format_bench.model import WorkloadSpec from format_bench.registry import adapters @@ -42,6 +51,7 @@ def test_adapter_contract_keys_are_explicit() -> None: "workloads", } assert get_origin(manifest_hints["workloads"]) is NotRequired + assert get_args(manifest_hints["workloads"]) == (WorkloadDeclarations,) assert set(get_type_hints(VerificationResult)) == { "canonical_hash", "counts", @@ -56,6 +66,32 @@ def test_adapter_contract_keys_are_explicit() -> None: assert AdapterManifest.__optional_keys__ == {"workloads"} +def test_workload_contract_is_explicit_and_discriminated() -> None: + assert get_args(ComparisonOperator) == ("eq", "gt", "gte", "lt", "lte") + assert set(get_args(WorkloadScalar)) == {str, int, float, bool} + assert set(get_args(WorkloadDeclaration)) == { + ReadAllWorkload, + ProjectionWorkload, + FilterWorkload, + HeadWorkload, + } + assert get_args(WorkloadDeclarations) == (str, WorkloadDeclaration) + assert get_type_hints(ReadAllWorkload)["kind"] == Literal["read_all"] + assert get_type_hints(ProjectionWorkload)["kind"] == Literal["projection"] + assert get_type_hints(FilterWorkload)["kind"] == Literal["filter"] + assert get_type_hints(HeadWorkload)["kind"] == Literal["head"] + assert get_type_hints(FilterWorkload)["operator"] == ComparisonOperator + assert get_type_hints(WorkloadSpec)["operator"] == ComparisonOperator | None + assert ReadAllWorkload.__required_keys__ == {"kind"} + assert ReadAllWorkload.__optional_keys__ == {"expected_rows"} + assert ProjectionWorkload.__required_keys__ == {"columns", "kind"} + assert ProjectionWorkload.__optional_keys__ == {"expected_rows"} + assert FilterWorkload.__required_keys__ == {"column", "kind", "operator", "value"} + assert FilterWorkload.__optional_keys__ == {"expected_rows"} + assert HeadWorkload.__required_keys__ == {"kind", "limit"} + assert HeadWorkload.__optional_keys__ == {"expected_rows"} + + def test_first_party_adapters_implement_the_named_contract() -> None: for adapter in adapters(): adapter_type = type(adapter) @@ -75,4 +111,5 @@ def test_adapter_contract_is_in_the_blocking_strict_frontier() -> None: assert pyright["typeCheckingMode"] == "strict" assert "src/format_bench/adapter_contract.py" in pyright["include"] assert "src/format_bench/registry.py" in pyright["include"] + assert "src/format_bench/workload_contract.py" in pyright["include"] assert "tests/typecheck/adapter_manifest_normalized.py" in pyright["include"] diff --git a/tests/test_workloads.py b/tests/test_workloads.py index bd71ec8..e6eed0d 100644 --- a/tests/test_workloads.py +++ b/tests/test_workloads.py @@ -5,38 +5,40 @@ from format_bench.fair import FairOperation, operations_for from format_bench.datasets import validate_manifest +from format_bench.workload_contract import WorkloadDeclarations from format_bench.workloads import apply_workload, expected_workload_rows, load_workloads def test_manifest_workload_is_not_tied_to_stars_columns() -> None: + declarations: WorkloadDeclarations = { + "read_all": {"kind": "read_all", "expected_rows": 3}, + "project_two": {"kind": "projection", "columns": ["name", "amount"]}, + "filter_ai_llm": { + "kind": "filter", + "column": "name", + "operator": "eq", + "value": "b", + "expected_rows": 1, + }, + "filter_repo_stars_gt_100000": { + "kind": "filter", + "column": "amount", + "operator": "gt", + "value": 10, + "expected_rows": 2, + }, + "exact_match": { + "kind": "filter", + "column": "name", + "operator": "eq", + "value": "c", + "expected_rows": 1, + }, + "head_10": {"kind": "head", "limit": 10}, + } manifest = { "rows": 3, - "workloads": { - "read_all": {"kind": "read_all", "expected_rows": 3}, - "project_two": {"kind": "projection", "columns": ["name", "amount"]}, - "filter_ai_llm": { - "kind": "filter", - "column": "name", - "operator": "eq", - "value": "b", - "expected_rows": 1, - }, - "filter_repo_stars_gt_100000": { - "kind": "filter", - "column": "amount", - "operator": "gt", - "value": 10, - "expected_rows": 2, - }, - "exact_match": { - "kind": "filter", - "column": "name", - "operator": "eq", - "value": "c", - "expected_rows": 1, - }, - "head_10": {"kind": "head", "limit": 10}, - }, + "workloads": declarations, } table = pa.table({"name": ["a", "b", "c"], "amount": [1, 11, 21]}) workloads = load_workloads(manifest) @@ -84,6 +86,9 @@ def test_manifest_rejects_malformed_boundary_mappings(manifest: dict) -> None: {"kind": "head", "limit": "10"}, {"kind": "read_all", "expected_rows": "3"}, {"kind": "filter", "column": 7, "operator": "eq", "value": 1}, + {"kind": "filter", "column": "name", "operator": "contains", "value": "a"}, + {"kind": "filter", "column": "name", "operator": "eq"}, + {"kind": "filter", "column": "name", "operator": "eq", "value": ["a"]}, ], ) def test_workload_rejects_coerced_boundary_values(