Skip to content
Merged
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
61 changes: 45 additions & 16 deletions src/backtest_engine/production.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,7 @@
QueuedRunProjection,
)
from backtest_engine.wiring import (
JobNotSatisfiable,
OrchestratorJobHandler,
PersistenceExecutionKeyStore,
PersistenceStorageObjectWritePort,
Expand Down Expand Up @@ -161,20 +162,44 @@ def by_checksum(self, checksum: str) -> Mapping[str, Any] | None:
return dict(document)


def _dataset_id(object_keys: list[str]) -> str:
values = {
segment.removeprefix("dataset=")
for key in object_keys
for segment in key.split("/")
if segment.startswith("dataset=")
}
if len(values) != 1:
raise ConfigurationError("dataset manifest object keys must bind one logical dataset= UUID")
def _unavailable_manifest(message: str) -> JobNotSatisfiable:
return JobNotSatisfiable(message, reason_code="REQUIRED_INPUT_UNAVAILABLE")


def _dataset_id(object_keys: list[str], manifest_id: uuid.UUID) -> str:
bindings: list[tuple[str, str]] = []
for key in object_keys:
key_bindings = [
(name, segment.removeprefix(f"{name}="))
for segment in key.split("/")
for name in ("dataset", "manifest_id")
if segment.startswith(f"{name}=")
]
if len(key_bindings) != 1:
raise _unavailable_manifest(
"dataset manifest object keys must each bind exactly one dataset= or manifest_id= UUID"
)
bindings.extend(key_bindings)

conventions = {name for name, _value in bindings}
values = {value for _name, value in bindings}
if len(conventions) != 1 or len(values) != 1:
raise _unavailable_manifest(
"dataset manifest object keys must use one binding convention and one UUID"
)
convention = conventions.pop()
value = values.pop()
try:
return str(uuid.UUID(value))
resolved = uuid.UUID(value)
except ValueError as exc:
raise ConfigurationError(f"dataset object key contains invalid dataset id: {value}") from exc
raise _unavailable_manifest(
f"dataset manifest object key contains an invalid {convention} UUID: {value}"
) from exc
if convention == "manifest_id" and resolved != manifest_id:
raise _unavailable_manifest(
f"legacy dataset object key binds manifest {resolved}, expected {manifest_id}"
)
return str(resolved)


class PostgresDatasetManifestSource:
Expand Down Expand Up @@ -212,10 +237,14 @@ def by_id(self, manifest_id: uuid.UUID) -> Mapping[str, Any] | None:
if row is None:
return None
object_rows = list(connection.execute(objects_sql, {"manifest_id": manifest_id}).mappings())
if row["status"] != "AVAILABLE":
raise _unavailable_manifest(f"dataset manifest {manifest_id} is not AVAILABLE")
if not object_rows:
raise ConfigurationError(f"dataset manifest {manifest_id} has no objects")
raise _unavailable_manifest(f"dataset manifest {manifest_id} has no objects")
if any(item["status"] != "AVAILABLE" for item in object_rows):
raise ConfigurationError(f"dataset manifest {manifest_id} references a non-AVAILABLE object")
raise _unavailable_manifest(
f"dataset manifest {manifest_id} references a non-AVAILABLE object"
)
objects = [
{
"storage_object_id": str(item["storage_object_id"]),
Expand All @@ -236,21 +265,21 @@ def by_id(self, manifest_id: uuid.UUID) -> Mapping[str, Any] | None:
]
available_at = row["available_at"]
if available_at is None:
raise ConfigurationError(f"dataset manifest {manifest_id} has no available_at evidence")
raise _unavailable_manifest(f"dataset manifest {manifest_id} has no available_at evidence")
raw_resolution = str(row["feed_resolution"])
try:
resolution = bar_resolution(raw_resolution)
except ValueError:
resolution = raw_resolution
if resolution not in {"30m", "1h", "4h", "1d"}:
raise ConfigurationError(
raise _unavailable_manifest(
f"dataset manifest {manifest_id} has unsupported production resolution {raw_resolution}"
)
return {
"contract_id": "com06.dataset-manifest",
"schema_version": 1,
"manifest_id": str(row["id"]),
"dataset_id": _dataset_id([item["object_key"] for item in objects]),
"dataset_id": _dataset_id([item["object_key"] for item in objects], manifest_id),
"revision": int(row["revision_number"]),
"status": str(row["status"]),
"dataset_hash": str(row["dataset_hash"]),
Expand Down
101 changes: 101 additions & 0 deletions tests/test_production.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
ConfigurationError,
JwtAuthenticator,
PostgresCompiledPlanSource,
PostgresDatasetManifestSource,
PostgresFeatureMaterializationSource,
PostgresOwnerDirectory,
PostgresQueuedRunSource,
Expand All @@ -29,6 +30,7 @@
orchestrator_job_handler,
service_endpoint,
)
from backtest_engine.wiring import JobNotSatisfiable


ACCOUNT_ID = UUID("aaaaaaaa-aaaa-4aaa-8aaa-aaaaaaaaaaaa")
Expand Down Expand Up @@ -88,6 +90,9 @@ def first(self) -> dict[str, object] | None:
def all(self) -> list[dict[str, object]]:
return self._rows

def __iter__(self):
return iter(self._rows)


class _Connection:
def __init__(self, rows: list[list[dict[str, object]]]) -> None:
Expand Down Expand Up @@ -137,6 +142,102 @@ def test_compiled_plan_source_returns_the_immutable_launch_contract_document() -
assert source.by_checksum(plan["planChecksum"]) == plan


def _dataset_manifest_source(
manifest_id: UUID,
object_keys: list[str],
) -> PostgresDatasetManifestSource:
manifest = {
"id": manifest_id,
"revision_number": 1,
"status": "AVAILABLE",
"dataset_hash": "sha256:" + "a" * 64,
"schema_version": "market-bars/1",
"period_start": datetime(2024, 1, 1, tzinfo=UTC),
"period_end": datetime(2024, 2, 1, tzinfo=UTC),
"available_at": datetime(2026, 8, 9, tzinfo=UTC),
"feed_resolution": "30m",
}
objects = [
{
"storage_object_id": UUID(f"00000000-0000-4000-8000-{index:012d}"),
"object_key": object_key,
"content_hash": "sha256:" + f"{index:x}" * 64,
"object_kind": "MARKET_BARS",
"partition_granularity": "YEAR",
"partition_start": date(2024, 1, 1),
"partition_end": date(2025, 1, 1),
"period_start": datetime(2024, 1, 1, tzinfo=UTC),
"period_end": datetime(2024, 2, 1, tzinfo=UTC),
"shard_key": f"{index:02d}-of-{len(object_keys):02d}",
"part_number": 1,
"row_count": 10,
"schema_version": "market-bars/1",
"status": "AVAILABLE",
}
for index, object_key in enumerate(object_keys, start=1)
]
return PostgresDatasetManifestSource(_Engine(manifest, objects)) # type: ignore[arg-type]


def test_dataset_manifest_source_accepts_the_deployed_legacy_loader_binding() -> None:
manifest_id = UUID("7f7113c9-3b02-4098-97ec-0baa07e2b3b0")
prefix = (
"historical/provider=alpaca/feed=sip/adjustment=all/session=regular/"
"resolution=30m/revision=00000001/year=2024"
)
object_keys = [
f"{prefix}/shard={shard:02d}-of-08/manifest_id={manifest_id}/part-00001.parquet"
for shard in range(8)
]

resolved = _dataset_manifest_source(manifest_id, object_keys).by_id(manifest_id)

assert resolved is not None
assert resolved["manifest_id"] == str(manifest_id)
assert resolved["dataset_id"] == str(manifest_id)
assert [item["object_key"] for item in resolved["objects"]] == object_keys


def test_dataset_manifest_source_preserves_the_canonical_logical_dataset_binding() -> None:
manifest_id = UUID("7f7113c9-3b02-4098-97ec-0baa07e2b3b0")
dataset_id = UUID("11111111-1111-4111-8111-111111111111")
object_keys = [
f"market-data/provider=ALPACA/feed=SIP/dataset={dataset_id}/revision=1/part-{part:05d}.parquet"
for part in range(1, 3)
]

resolved = _dataset_manifest_source(manifest_id, object_keys).by_id(manifest_id)

assert resolved is not None
assert resolved["dataset_id"] == str(dataset_id)


@pytest.mark.parametrize(
"object_keys",
[
[
"historical/dataset=11111111-1111-4111-8111-111111111111/part-00001.parquet",
"historical/manifest_id=7f7113c9-3b02-4098-97ec-0baa07e2b3b0/part-00002.parquet",
],
[
"historical/manifest_id=22222222-2222-4222-8222-222222222222/part-00001.parquet",
],
["historical/dataset=not-a-uuid/part-00001.parquet"],
["historical/revision=00000001/part-00001.parquet"],
],
ids=["mixed-conventions", "wrong-legacy-manifest", "invalid-uuid", "missing-binding"],
)
def test_dataset_manifest_source_classifies_invalid_catalog_bindings_as_terminal(
object_keys: list[str],
) -> None:
manifest_id = UUID("7f7113c9-3b02-4098-97ec-0baa07e2b3b0")

with pytest.raises(JobNotSatisfiable) as failure:
_dataset_manifest_source(manifest_id, object_keys).by_id(manifest_id)

assert failure.value.reason_code == "REQUIRED_INPUT_UNAVAILABLE"


def test_feature_materialization_source_returns_definition_manifest_and_object_evidence() -> None:
materialization_id = UUID("10000000-0000-4000-8000-000000000001")
objects = [{"object_key": "features/rsi.parquet", "provider_version_id": "v1"}]
Expand Down