Skip to content

Example source: projects/getting-started/runners/local/run_spark.py

Source revision: 75f65139e26eb7079b7a897c877de97859c0120b

This page is the generated source; it shows the complete readable projection of the canonical file.

raw · View repository source

"""Run one lesson from the canonical local Spark onboarding project."""

from __future__ import annotations

import argparse
import json
import os
import sys
from decimal import Decimal
from pathlib import Path
from typing import Any

from datacoolie.core.models.run_config import DataCoolieRunConfig
from datacoolie.engines.spark_engine import SparkEngine
from datacoolie.metadata.file_provider import FileProvider
from datacoolie.orchestration.driver import DataCoolieDriver
from datacoolie.platforms.local_platform import LocalPlatform

from checks import (
    GuardError,
    latest_orders,
    parse_amount,
    parse_int,
    read_fixture,
    require_named_dataflows,
    require_delta_path,
    require_no_newer_rows,
    require_terminal_result,
    resolve_root,
)


LESSONS = ("orders", "customers", "multi-stage")


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(
        description="Run the canonical DataCoolie getting-started lesson with Spark."
    )
    parser.add_argument(
        "--lesson",
        choices=LESSONS,
        default="orders",
        help="Lesson to execute (default: orders).",
    )
    parser.add_argument(
        "--state-base-path",
        default=".runtime",
        help="Explicit runtime state and log root (default: .runtime).",
    )
    return parser.parse_args()


def _project_root() -> Path:
    return Path(__file__).resolve().parents[2]


def _orders_input(root: Path) -> tuple[Path, list[dict[str, str]]]:
    path = root / "data" / "input" / "orders" / "orders.csv"
    rows = read_fixture(path, ("order_id", "customer_id", "amount", "updated_at"))
    for row in rows:
        parse_int(row["order_id"], column="order_id")
        parse_int(row["customer_id"], column="customer_id")
        parse_amount(row["amount"])
    latest_orders(rows)
    return path, rows


def _customers_input(root: Path) -> tuple[Path, list[dict[str, str]]]:
    path = root / "data" / "input" / "customers" / "customers.csv"
    rows = read_fixture(path, ("customer_id", "name"))
    seen: set[int] = set()
    for row in rows:
        customer_id = parse_int(row["customer_id"], column="customer_id")
        if customer_id in seen:
            raise GuardError(f"Customers fixture contains duplicate customer_id: {customer_id}")
        if not row["name"]:
            raise GuardError(f"Customers fixture contains an empty name for customer_id {customer_id}")
        seen.add(customer_id)
    return path, rows


def _require_columns(frame: Any, required: dict[str, str], *, label: str) -> None:
    fields = {field.name: field.dataType.simpleString().lower() for field in frame.schema.fields}
    missing = sorted(set(required) - set(fields))
    if missing:
        raise GuardError(f"{label} is missing columns: {', '.join(missing)}")
    for column, kind in required.items():
        dtype = fields[column]
        if kind not in dtype:
            raise GuardError(f"{label}.{column} has type {dtype}; expected a {kind} type")


def _require_system_columns(frame: Any, *, label: str) -> None:
    required = {"__created_at", "__updated_at", "__updated_by"}
    missing = sorted(required - {field.name for field in frame.schema.fields})
    if missing:
        raise GuardError(f"{label} is missing framework columns: {', '.join(missing)}")


def _read_delta(spark: Any, path: Path, *, label: str) -> Any:
    require_delta_path(path, label=label)
    try:
        return spark.read.format("delta").load(str(path))
    except Exception as exc:
        raise GuardError(f"Cannot read {label} Delta output: {path}") from exc


def _validate_orders_output(
    spark: Any,
    path: Path,
    rows: list[dict[str, str]],
    *,
    label: str,
) -> dict[str, Any]:
    frame = _read_delta(spark, path, label=label)
    _require_columns(
        frame,
        {
            "order_id": "bigint",
            "customer_id": "bigint",
            "amount": "decimal",
            "updated_at": "timestamp",
        },
        label=label,
    )
    _require_system_columns(frame, label=label)
    expected = latest_orders(rows)
    actual_rows = frame.select("order_id", "customer_id", "amount").collect()
    actual_ids = {int(row["order_id"]) for row in actual_rows}
    expected_ids = set(expected)
    if actual_ids != expected_ids or len(actual_rows) != len(expected_ids):
        raise GuardError(
            f"{label} IDs/row count mismatch: actual_ids={sorted(actual_ids)} "
            f"rows={len(actual_rows)}, expected_ids={sorted(expected_ids)}"
        )
    actual = {int(row["order_id"]): row.asDict() for row in actual_rows}
    for order_id, source in expected.items():
        output = actual[order_id]
        if int(output["customer_id"]) != parse_int(source["customer_id"], column="customer_id"):
            raise GuardError(f"{label} customer_id mismatch for order_id {order_id}")
        if Decimal(str(output["amount"])).quantize(Decimal("0.01")) != parse_amount(source["amount"]):
            raise GuardError(f"{label} amount mismatch for order_id {order_id}")
    return {"rows": len(actual_rows), "order_ids": sorted(actual_ids)}


def _validate_customers_output(spark: Any, path: Path, rows: list[dict[str, str]]) -> dict[str, Any]:
    frame = _read_delta(spark, path, label="Customers")
    _require_columns(frame, {"customer_id": "bigint", "name": "string"}, label="Customers")
    _require_system_columns(frame, label="Customers")
    expected = {
        parse_int(row["customer_id"], column="customer_id"): row["name"]
        for row in rows
    }
    actual_rows = frame.select("customer_id", "name").collect()
    actual = {int(row["customer_id"]): str(row["name"]) for row in actual_rows}
    if actual != expected:
        raise GuardError(f"Customers output mismatch: actual={actual}, expected={expected}")
    return {"rows": len(actual_rows), "customer_ids": sorted(actual)}


def _validate_silver_output(spark: Any, path: Path, rows: list[dict[str, str]]) -> dict[str, Any]:
    frame = _read_delta(spark, path, label="Silver")
    _require_columns(
        frame,
        {
            "order_id": "bigint",
            "customer_id": "bigint",
            "amount": "decimal",
            "updated_at": "timestamp",
            "order_date": "date",
        },
        label="Silver",
    )
    result = _validate_orders_output(spark, path, rows, label="Silver")
    actual_dates = {
        int(row["order_id"]): str(row["order_date"])
        for row in frame.select("order_id", "order_date").collect()
    }
    expected_dates = {
        order_id: source["updated_at"][:10]
        for order_id, source in latest_orders(rows).items()
    }
    if actual_dates != expected_dates:
        raise GuardError(f"Silver order_date mismatch: actual={actual_dates}, expected={expected_dates}")
    partition_dirs = sorted(
        item.name
        for item in path.rglob("order_date=*")
        if item.is_dir()
    )
    if not partition_dirs:
        raise GuardError(f"Silver output has no order_date partitions: {path}")
    result["partitions"] = partition_dirs
    result["order_dates"] = actual_dates
    return result


def _spark_session() -> Any:
    from delta import configure_spark_with_delta_pip
    from pyspark.sql import SparkSession

    builder = (
        SparkSession.builder.appName("datacoolie-getting-started")
        .master("local[2]")
        .config("spark.sql.parquet.outputTimestampType", "TIMESTAMP_MICROS")
        .config("spark.sql.extensions", "io.delta.sql.DeltaSparkSessionExtension")
        .config("spark.sql.catalog.spark_catalog", "org.apache.spark.sql.delta.catalog.DeltaCatalog")
    )
    return configure_spark_with_delta_pip(builder).getOrCreate()


def _dataflows_for(
    metadata: FileProvider, root: Path, stage: str, expected_name: str
) -> list[Any]:
    """Load one stage and make local paths explicit for Spark's Delta SQL checks."""

    dataflows = require_named_dataflows(
        metadata.get_dataflows(stage=stage, active_only=True),
        expected_name=expected_name,
        stage=stage,
    )
    for dataflow in dataflows:
        for connection in (dataflow.source.connection, dataflow.destination.connection):
            base_path = connection.configure.get("base_path")
            if not isinstance(base_path, str) or not base_path.strip():
                continue
            if "://" in base_path or Path(base_path).is_absolute():
                continue
            connection.configure["base_path"] = str((root / base_path).resolve())
    return dataflows


def _run_lesson(args: argparse.Namespace) -> dict[str, Any]:
    root = _project_root()
    os.chdir(root)
    state_root = resolve_root(root, args.state_base_path).resolve()
    spark = _spark_session()
    platform = LocalPlatform()
    metadata = FileProvider(
        metadata_base_path=str(root / "metadata"),
        platform=platform,
        watermark_base_path=str(state_root / "watermarks"),
    )
    config = DataCoolieRunConfig(
        job_id="getting-started-local-spark",
        max_workers=1,
        stop_on_error=True,
        allowed_function_prefixes=[],
    )
    engine = SparkEngine(spark_session=spark, platform=platform)
    orders_path = root / "data" / "output" / "bronze" / "orders"
    silver_path = root / "data" / "output" / "silver" / "orders"
    customers_path = root / "data" / "output" / "customers" / "customers"

    try:
        with DataCoolieDriver(
            engine=engine,
            platform=platform,
            metadata_provider=metadata,
            state_base_path=str(state_root),
            log_base_path=str(state_root / "logs"),
            config=config,
        ) as driver:
            if args.lesson == "customers":
                _, customer_rows = _customers_input(root)
                counts = require_terminal_result(
                    driver.run(
                        dataflows=_dataflows_for(
                            metadata,
                            root,
                            "customers_full_refresh",
                            "customers_full_refresh",
                        )
                    ),
                    lesson="customers_full_refresh",
                )
                details = _validate_customers_output(spark, customers_path, customer_rows)
                return {
                    "lesson": args.lesson,
                    "ok": True,
                    "results": [{"stage": "customers_full_refresh", **counts}],
                    "output": details,
                }

            _, order_rows = _orders_input(root)
            bronze_counts = require_terminal_result(
                driver.run(
                    dataflows=_dataflows_for(
                        metadata, root, "ingest2bronze", "orders_to_bronze"
                    )
                ),
                lesson="orders_to_bronze",
                allow_skip=True,
            )
            watermark = None
            if bronze_counts["skipped"]:
                watermark = require_no_newer_rows(
                    order_rows,
                    runtime_root=state_root,
                    output_path=orders_path,
                ).isoformat()
            bronze_details = _validate_orders_output(spark, orders_path, order_rows, label="Bronze")

            if args.lesson == "orders":
                return {
                    "lesson": args.lesson,
                    "ok": True,
                    "results": [{"stage": "ingest2bronze", **bronze_counts}],
                    "output": bronze_details,
                    "no_change_watermark": watermark,
                }

            silver_counts = require_terminal_result(
                driver.run(
                    dataflows=_dataflows_for(
                        metadata, root, "bronze2silver", "orders_to_silver"
                    )
                ),
                lesson="orders_to_silver",
            )
            silver_details = _validate_silver_output(spark, silver_path, order_rows)
            return {
                "lesson": args.lesson,
                "ok": True,
                "results": [
                    {"stage": "ingest2bronze", **bronze_counts},
                    {"stage": "bronze2silver", **silver_counts},
                ],
                "bronze": bronze_details,
                "silver": silver_details,
                "no_change_watermark": watermark,
            }
    finally:
        spark.stop()


def main() -> int:
    args = parse_args()
    try:
        payload = _run_lesson(args)
    except GuardError as exc:
        print(json.dumps({"lesson": args.lesson, "ok": False, "error": str(exc)}), file=sys.stderr)
        return 2
    except Exception as exc:  # pragma: no cover - keeps CLI failures machine-readable
        print(
            json.dumps(
                {"lesson": args.lesson, "ok": False, "error": f"{type(exc).__name__}: {exc}"}
            ),
            file=sys.stderr,
        )
        return 1
    print(json.dumps(payload, sort_keys=True))
    return 0


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