diff --git a/.github/scripts/check_release_tag.py b/.github/scripts/check_release_tag.py new file mode 100644 index 0000000..893c4ce --- /dev/null +++ b/.github/scripts/check_release_tag.py @@ -0,0 +1,67 @@ +"""Reject release tags unless stable metadata and fetched main ancestry agree. + +Run with Python 3.11+ after fetching origin/main. This script does not publish. +""" + +import argparse +import ast +import re +import subprocess +from pathlib import Path + +STABLE_TAG = re.compile(r"refs/tags/v(0|[1-9][0-9]*)\.(0|[1-9][0-9]*)\.(0|[1-9][0-9]*)") + + +def git(root: Path, *arguments: str) -> str: + """Read an identity from the checked-out repository.""" + return subprocess.check_output(["git", "-C", str(root), *arguments], text=True, stderr=subprocess.PIPE).strip() + + +def validate_release(root: Path, event: str, ref: str, sha: str) -> str: + """Return the stable version only for a matching checkout on fetched main.""" + import tomllib + + if event != "push" or STABLE_TAG.fullmatch(ref) is None: + raise ValueError("Publication requires a push of an exact stable vMAJOR.MINOR.PATCH tag.") + version = ref.removeprefix("refs/tags/v") + project = tomllib.loads((root / "pyproject.toml").read_text())["project"] + if project["version"] != version: + raise ValueError("Release tag does not match project.version.") + module = root / "src" / project["name"].replace("-", "_") / "__init__.py" + source_versions = [ + ast.literal_eval(node.value) + for node in ast.parse(module.read_text()).body + if isinstance(node, ast.Assign) + and any(isinstance(target, ast.Name) and target.id == "__version__" for target in node.targets) + ] + if source_versions != [version]: + raise ValueError("Release tag does not match the single source __version__ literal.") + commit = git(root, "rev-parse", "--verify", f"{sha}^{{commit}}") + if git(root, "rev-parse", "HEAD") != commit: + raise ValueError("Release checkout does not match the event commit.") + subprocess.run( + ["git", "-C", str(root), "merge-base", "--is-ancestor", commit, "refs/remotes/origin/main"], + check=True, + capture_output=True, + text=True, + ) + return version + + +def main() -> None: + """Check explicit event inputs; a nonzero exit prevents release steps.""" + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--root", type=Path, default=Path.cwd()) + parser.add_argument("--event", required=True) + parser.add_argument("--ref", required=True) + parser.add_argument("--sha", required=True) + args = parser.parse_args() + try: + version = validate_release(args.root, args.event, args.ref, args.sha) + except (ValueError, KeyError, OSError, subprocess.SubprocessError) as error: + parser.exit(1, f"Release rejected: {error}\n") + print(f"Verified stable release {version}: metadata, source version, checkout and main ancestry agree.") + + +if __name__ == "__main__": + main() diff --git a/.github/workflows/ci-full.yml b/.github/workflows/ci-full.yml index 7c46eb7..bc4fb00 100644 --- a/.github/workflows/ci-full.yml +++ b/.github/workflows/ci-full.yml @@ -45,8 +45,8 @@ jobs: - name: Generate test summary if: always() run: | - echo "## Test Results - Python ${{ matrix.python }}" >> $GITHUB_STEP_SUMMARY - python3 .github/scripts/test_summary.py >> $GITHUB_STEP_SUMMARY + echo "## Test Results - Python ${{ matrix.python }}" >> "$GITHUB_STEP_SUMMARY" + python3 .github/scripts/test_summary.py >> "$GITHUB_STEP_SUMMARY" - name: Upload test results if: always() @@ -65,7 +65,7 @@ jobs: - name: Upload coverage to Codecov if: matrix.python == '3.12' - uses: codecov/codecov-action@v5.4.3 + uses: codecov/codecov-action@303a32d7a59b442fa8d48b6a1cc6825c09c847a5 # v7.1.1 with: files: ./coverage.xml fail_ci_if_error: true diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index d191d93..2d741e5 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -75,7 +75,7 @@ jobs: - name: Generate test summary if: always() - run: python3 .github/scripts/test_summary.py >> $GITHUB_STEP_SUMMARY + run: python3 .github/scripts/test_summary.py >> "$GITHUB_STEP_SUMMARY" - name: Upload test results if: always() @@ -92,7 +92,7 @@ jobs: run: uv run coverage xml - name: Upload coverage to Codecov - uses: codecov/codecov-action@v5.4.3 + uses: codecov/codecov-action@303a32d7a59b442fa8d48b6a1cc6825c09c847a5 # v7.1.1 with: files: ./coverage.xml fail_ci_if_error: true diff --git a/.github/workflows/publish.yml b/.github/workflows/publish.yml index 3e03588..3d9ec33 100644 --- a/.github/workflows/publish.yml +++ b/.github/workflows/publish.yml @@ -15,7 +15,7 @@ jobs: name: Build the package runs-on: ubuntu-latest timeout-minutes: 10 - if: startsWith(github.ref, 'refs/tags') || github.event_name == 'workflow_dispatch' + if: (github.event_name == 'push' && startsWith(github.ref, 'refs/tags/')) || github.event_name == 'workflow_dispatch' permissions: contents: read @@ -28,14 +28,16 @@ jobs: with: install-deps: 'false' - - name: Verify tag is on main branch - if: startsWith(github.ref, 'refs/tags') + - name: Verify stable release tag + if: github.event_name == 'push' && startsWith(github.ref, 'refs/tags/') + env: + RELEASE_EVENT: ${{ github.event_name }} + RELEASE_REF: ${{ github.ref }} + RELEASE_SHA: ${{ github.sha }} run: | - git fetch origin main - if ! git merge-base --is-ancestor ${{ github.sha }} origin/main; then - echo "Error: Tag is not on the main branch" - exit 1 - fi + git fetch --no-tags origin +refs/heads/main:refs/remotes/origin/main + uv run --no-project --python 3.12 python .github/scripts/check_release_tag.py \ + --event "$RELEASE_EVENT" --ref "$RELEASE_REF" --sha "$RELEASE_SHA" - name: Build a binary wheel and a source tarball run: uv build @@ -49,6 +51,7 @@ jobs: publish: name: Publish the package needs: build + if: github.event_name == 'push' && startsWith(github.ref, 'refs/tags/') runs-on: ubuntu-latest timeout-minutes: 10 permissions: diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index 9d8aaac..94dd52d 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -13,20 +13,27 @@ jobs: name: Create GitHub Release runs-on: ubuntu-latest timeout-minutes: 5 - if: startsWith(github.ref, 'refs/tags') + if: github.event_name == 'push' && startsWith(github.ref, 'refs/tags/') steps: - uses: actions/checkout@v4 with: fetch-depth: 0 - - name: Verify tag is on main branch + - uses: ./.github/actions/setup + with: + install-deps: 'false' + + - name: Verify stable release tag + if: github.event_name == 'push' && startsWith(github.ref, 'refs/tags/') + env: + RELEASE_EVENT: ${{ github.event_name }} + RELEASE_REF: ${{ github.ref }} + RELEASE_SHA: ${{ github.sha }} run: | - git fetch origin main - if ! git merge-base --is-ancestor ${{ github.sha }} origin/main; then - echo "Error: Tag is not on the main branch" - exit 1 - fi + git fetch --no-tags origin +refs/heads/main:refs/remotes/origin/main + uv run --no-project --python 3.12 python .github/scripts/check_release_tag.py \ + --event "$RELEASE_EVENT" --ref "$RELEASE_REF" --sha "$RELEASE_SHA" - name: Create Release uses: softprops/action-gh-release@v2 diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 4720213..cd8f393 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -18,14 +18,16 @@ brew install just # Or see https://github.com/casey/just for other platforms ``` -### 2. Clone and Setup +### 2. Set up the matching source + +Follow the public clone and editable commands in [installation](docs/getting-started/installation.md#development-setup). They select an immutable reviewed snapshot; use the explicitly agreed branch/revision for new work. From that clone, run: ```bash -git clone https://github.com/project-lighter/sparkwheel.git -cd sparkwheel just setup ``` +This uses or updates the clone's `.venv` and adds contributor dependencies. Activate that environment when running its tools directly. + This will: - Install all dependencies (dev, test, doc groups) - Set up pre-commit hooks @@ -53,6 +55,30 @@ We use: Pre-commit hooks will automatically run on commit. +## Keep documentation useful + +Use one documentation set for people and agents. Keep behavior in ordinary +Python, teach one complete first example, and link to the canonical semantic +explanation instead of maintaining parallel copies. + +- State the working directory, prerequisites, exact inputs and expected outputs + for a runnable example. Mark excerpts through their surrounding explanation. +- Preserve the distinction between source inspection and execution, and between + a shared result and a copied definition. +- Run changed complete examples with the documented install; record which checks + actually ran. Review other affected snippets against source. +- Build the site strictly and check affected links/anchors. With documentation + dependencies installed and the environment active, use: + +```bash +python -m mkdocs build --strict +``` + +Keep explanatory Markdown under `docs/`, entrypoint guidance in README, and the +Python API reference generated from source. Prefer existing paths and descriptive +headings. New documentation machinery needs a concrete unmet need. A passing +build does not execute every code example. + ## Pull Request Process 1. Fork the repository @@ -86,3 +112,36 @@ When reporting issues, please include: ## License By contributing, you agree that your contributions will be licensed under the Apache License 2.0. + +## Version preparation + +Version maintenance uses the pinned `bump-my-version==0.30.1` executable through +`uvx`, independently of runtime dependency installation. The supported recipe +parts are `major`, `minor`, `patch`, and `release`: + +```bash +just bump-dry release # Preview 0.1.0.dev0 -> 0.1.0 +just bump-dry patch # A stable 0.1.0 previews 0.1.1 +``` + +`release` removes an existing `.devN` suffix without advancing the numeric +version. Numeric bumps advance that component and reset subordinate components +to stable; `patch` from `0.1.0.dev0` therefore previews `0.1.1`, not promotion to +`0.1.0`. Use `release` for that promotion. Unsupported parts, including the raw +`dev` counter, fail before invoking the tool. Starting another development cycle +requires a separately reviewed explicit version choice. + +Dry runs disable file writes, commits and tags. After release authorization, +`just bump ` updates project metadata, the source version constant, the +bump configuration and only Sparkwheel's editable-root version in `uv.lock`; it +also creates the configured local commit and tag. Review the diff and run the +ordinary package checks before any separately authorized push or publication. +Current development snapshots have not been tagged or published by this work. + +Previewing or preparing a version locally does not authorize its remote release. + +## Stable release guard + +Use the existing pinned `bump-my-version` recipes to preview or promote a version. Actual bump commands can create a commit and tag; `just bump-dry release` previews without those changes. For a rehearsal, work in disposable copies and explicitly disable commits and tags. + +Publication accepts only a pushed `vMAJOR.MINOR.PATCH` tag matching both project metadata and the source version, at a commit on freshly fetched `main`. Development/prerelease tags and mismatches fail before publishing or creating a stable GitHub Release. Manual dispatch of Publish builds and retains workflow artifacts only. Current token authentication is unchanged. Building an artifact does not publish it or establish registry availability. diff --git a/README.md b/README.md index 3e703b8..b4e2628 100644 --- a/README.md +++ b/README.md @@ -1,65 +1,56 @@ -
-
- -
-

-

- CI - Coverage - PyPI - License - Documentation -

+# Sparkwheel -

YAML configuration meets Python

-

Define Python objects in YAML. Reference, compose, and instantiate them effortlessly.

-
+Compose configuration, then build ordinary Python objects. -## Quick Start +Sparkwheel loads YAML or Python dictionaries, combines settings, and calls your +classes and functions. Your application keeps its normal Python code. `@` shares +a resolved object; `%` copies a definition for separate construction. -```bash -pip install sparkwheel -``` - -```yaml -# config.yaml -dataset: - num_classes: 10 - batch_size: 32 +**Development version:** this checkout is `0.1.0.dev0`, an unpublished development +package. The documentation describes this source, not necessarily the version +available on PyPI. Python 3.10 or newer is required; the installation below uses +Python 3.11. Torch and Lightning are not required. Use the source documentation +linked below for this pair; the hosted site may describe an earlier release. -model: - _target_: torch.nn.Linear - in_features: 784 - out_features: "%dataset::num_classes" # Reference +Install the immutable reviewed public snapshot in a new directory (Python 3.11 and Git required): -training: - steps_per_epoch: "$10000 // @dataset::batch_size" # Expression +```bash +python3.11 -m venv .venv +source .venv/bin/activate +python -m pip install "sparkwheel @ git+https://github.com/project-lighter/sparkwheel.git@b73e786e8716d11a77206fb3482d21a621a4ba81" +python -m pip check ``` +See [installation](docs/getting-started/installation.md) for Windows activation +and editable-development instructions. This full commit pin does not follow later branch changes. + ```python from sparkwheel import Config -config = Config() -config.update("config.yaml") +config = Config().update( + { + "words": ["red", "blue", "red"], + "counts": {"_target_": "collections.Counter", "_args_": ["@words"]}, + } +) -model = config.resolve("model") # Actual torch.nn.Linear(784, 10) +print(dict(config.resolve("counts"))) # {'red': 2, 'blue': 1} ``` -## Features - -- **Declarative Objects** - Instantiate any Python class with `_target_` -- **Smart References** - `@` for resolved values, `%` for raw YAML -- **Composition by Default** - Dicts merge, lists extend automatically -- **Explicit Control** - `=` to replace, `~` to delete -- **Python Expressions** - Dynamic values with `$` -- **Schema Validation** - Type-check with dataclasses - -**[Get Started](https://project-lighter.github.io/sparkwheel/getting-started/quickstart/)** · **[Documentation](https://project-lighter.github.io/sparkwheel/)** · **[Quick Reference](https://project-lighter.github.io/sparkwheel/user-guide/quick-reference/)** - -## Community - -- [Discord](https://discord.gg/zJcnp6KrUp) · [YouTube](https://www.youtube.com/channel/UCef1oTpv2QEBrD2pZtrdk1Q) · [Issues](https://github.com/project-lighter/sparkwheel/issues) - -## About - -Sparkwheel is a hard fork of [MONAI Bundle](https://github.com/Project-MONAI/MONAI/tree/dev/monai/bundle)'s config system, with the goal of making a more general-purpose configuration library for Python projects. It combines the best of MONAI Bundle and [Hydra](http://hydra.cc/)/[OmegaConf](https://omegaconf.readthedocs.io/), while introducing new features and improvements not found in either. +`get()` reads the authored configuration. `resolve()` can import modules, evaluate +expressions and construct objects. Use configuration from sources you trust. + +| I want to… | Start here | +|---|---| +| Run a complete Python/YAML example | [Quick start](docs/getting-started/quickstart.md) | +| Understand source, objects and edits | [Configuration model](docs/user-guide/basics.md) | +| Share objects or copy definitions | [References](docs/user-guide/references.md) | +| Use my own classes and functions | [Python authoring](docs/user-guide/instantiation.md) | +| Combine files and command-line changes | [Composition](docs/user-guide/operators.md) · [CLI](docs/user-guide/cli.md) | +| Diagnose a configuration | [Troubleshooting](docs/user-guide/troubleshooting.md) | +| Look up syntax or APIs | [Quick reference](docs/user-guide/quick-reference.md) · [Documentation](https://project-lighter.github.io/sparkwheel/) | + +Sparkwheel grew from MONAI Bundle's configuration system and now serves +standalone Python projects and [Lighter](https://github.com/project-lighter/lighter). +It is licensed under Apache 2.0. Report issues on +[GitHub](https://github.com/project-lighter/sparkwheel/issues). diff --git a/docs/getting-started/installation.md b/docs/getting-started/installation.md index 1bff49b..903b2f5 100644 --- a/docs/getting-started/installation.md +++ b/docs/getting-started/installation.md @@ -1,78 +1,48 @@ # Installation -Sparkwheel requires Python 3.10 or higher. +Sparkwheel requires Python 3.10 or newer and PyYAML. Torch and Lightning are not required. The examples below use Python 3.11 and a fresh environment. -## Install from PyPI - -The simplest way to install Sparkwheel: - -```bash -pip install sparkwheel -``` +The current source version is **0.1.0.dev0**. Its source is public; older PyPI releases may have different behavior. The command below selects the **immutable reviewed snapshot** `b73e786e8716d11a77206fb3482d21a621a4ba81`. It includes the current runtime and diagnostic corrections; it does not follow subsequent branch changes. Release handoffs name the exact revision they qualify. ## Install from Source -For the latest development version: +With Python 3.11 and Git installed, start in a new working directory: ```bash -git clone https://github.com/project-lighter/sparkwheel.git -cd sparkwheel -pip install -e . +python3.11 -m venv .venv +source .venv/bin/activate +python -m pip install "sparkwheel @ git+https://github.com/project-lighter/sparkwheel.git@b73e786e8716d11a77206fb3482d21a621a4ba81" +python -m pip check ``` -## Development Setup +This installs a normal package from the public Git revision. It needs network access but no supplied checkout or GitHub account. On Windows, use `py -3.11 -m venv .venv` and activate `.venv\Scripts\Activate.ps1` in PowerShell, then run the same pip commands. -If you want to contribute to Sparkwheel, we use [uv](https://github.com/astral-sh/uv) and [just](https://github.com/casey/just) for development: +Keep the environment active for the examples. The full commit pin does not follow branch updates, and changes in an unrelated local clone do not change the installed package. -### Install uv +## Verify Installation ```bash -curl -LsSf https://astral.sh/uv/install.sh | sh +python -c "import sparkwheel; print(sparkwheel.__version__); print(sparkwheel.__file__)" ``` -### Install just +Expect `0.1.0.dev0` and a path inside this environment's `site-packages/sparkwheel`. If the version or path differs, check `python -m pip show sparkwheel` and the active environment before continuing. pip's distribution metadata also records the installed VCS commit ID. -=== "macOS" - ```bash - brew install just - ``` +## Install from PyPI -=== "Linux" - ```bash - apt install just - ``` +For an application targeting a released version, install that release and use its matching documentation. The explicit source revision above provides the development behavior described here while the corresponding stable release is being prepared. -=== "Windows" - ```powershell - winget install --id Casey.Just --exact - ``` +## Development Setup -### Setup Development Environment +Only contributors need an editable clone. From a new working directory with the environment active: ```bash -git clone https://github.com/project-lighter/sparkwheel.git -cd sparkwheel -just setup +git clone https://github.com/project-lighter/sparkwheel.git sparkwheel +git -C sparkwheel checkout --detach b73e786e8716d11a77206fb3482d21a621a4ba81 +python -m pip install --editable ./sparkwheel ``` -Check out the [`justfile`](https://github.com/project-lighter/sparkwheel/blob/main/justfile) for other available commands. - -This will: - -- Install all dependencies (including dev, test, and doc groups) -- Set up pre-commit hooks -- Configure your development environment - -## Verify Installation - -Test that Sparkwheel is installed correctly: - -```python -import sparkwheel -print(sparkwheel.__version__) -``` +Edits in this clone now affect imports. Create a working branch before committing changes. For tests, docs and hooks, install `uv` and `just`, then use `just setup` from the clone. See [Contributing](https://github.com/project-lighter/sparkwheel/blob/b5138c192092989910c6e62c59245865928fc51f/CONTRIBUTING.md). Contributor tools are not required for package use. ## Next Steps -- [Quick Start](quickstart.md) - Learn the basics -- [User Guide](../user-guide/basics.md) - Deep dive into features +Continue to the [quick start](quickstart.md), which runs independently of a machine-learning framework. diff --git a/docs/getting-started/quickstart.md b/docs/getting-started/quickstart.md index 4bcf488..05aae90 100644 --- a/docs/getting-started/quickstart.md +++ b/docs/getting-started/quickstart.md @@ -1,133 +1,109 @@ # Quick Start -Get productive with Sparkwheel in 5 minutes. - -## Installation - -```bash -pip install sparkwheel -``` +Build a word counter from configuration, inspect its definition, then change its +input. This complete example uses only Sparkwheel and Python's standard library. +First complete [installation](installation.md) and keep that environment active. ## Your First Config -Create a file `config.yaml`: +The ordinary Python version is: -```yaml -# config.yaml -dataset: - path: "/data/train" - num_classes: 10 - batch_size: 32 - -model: - _target_: torch.nn.Linear - in_features: 784 - out_features: "@dataset::num_classes" # Reference! +```python +from collections import Counter -training: - epochs: 10 - learning_rate: 0.001 - steps_per_epoch: "$10000 // @dataset::batch_size" # Expression! +counts = Counter(["red", "blue", "red"]) +print(dict(counts)) # {'red': 2, 'blue': 1} ``` -Load and use it in Python: +To select the input through configuration, create a new example directory: -```python -from sparkwheel import Config +```bash +mkdir sparkwheel-demo +cd sparkwheel-demo +``` -# Load the config -config = Config() -config.update("config.yaml") +Create these two files in that directory. -# Access values with path notation -batch_size = config["dataset::batch_size"] # 32 -epochs = config["training::epochs"] # 10 +```yaml title="config.yaml" +words: [red, blue, red] +counts: + _target_: collections.Counter + _args_: ["@words"] +``` -# Resolve references and expressions -steps = config.resolve("training::steps_per_epoch") # 312 (10000 // 32) +```python title="run.py" +import sys +from sparkwheel import Config -# Instantiate objects -model = config.resolve("model") # Actual torch.nn.Linear(784, 10) instance! +config = Config().update("config.yaml") +for override in sys.argv[1:]: + config.update(override) -print(f"Training for {epochs} epochs with batch size {batch_size}") -print(f"Model: {model}") +print("Definition:", config.get("counts")["_target_"]) +counts = config.resolve("counts") +print("Counts:", dict(counts)) ``` -**That's it!** You just: +Run from `sparkwheel-demo`: -- ✓ Loaded a YAML config -- ✓ Referenced resolved values with `@` (gets instantiated/computed results) -- ✓ Computed values with `$` (Python expressions) -- ✓ Instantiated a Python object from config with `_target_` +```bash +python run.py +``` -!!! tip "Two Types of References" - - `@` = **Resolved reference** - gets the final instantiated/evaluated value - - `%` = **Raw reference** - copies unprocessed YAML content (from same or external file) +Expected output: -## Experiment Without Copying +```text +Definition: collections.Counter +Counts: {'red': 2, 'blue': 1} +``` -Create a variant without duplicating the base config (merges automatically!): +`_target_` names the Python class to call. `_args_` supplies positional arguments. +`@words` supplies the resolved value of the `words` node. `get()` reads the +configuration; `resolve()` constructs the counter. Nothing is written to disk by +this program. -```yaml -# experiment_large.yaml -model: # Merges by default - no operator needed! - in_features: 1568 # Override just this - # out_features is still @dataset::num_classes +## CLI Overrides -training: - learning_rate: 0.0001 # Lower learning rate - # epochs and steps_per_epoch inherited from base -``` +From the same directory, replace the input list: -Load both configs: +```bash +python run.py '=words=[green, green, red]' +``` -```python -config = (Config() - .update("config.yaml") - .update("experiment_large.yaml")) +Expected output: -model = config.resolve("model") # Linear(1568, 10) - merged automatically! -lr = config["training::learning_rate"] # 0.0001 -epochs = config["training::epochs"] # 10 (inherited) +```text +Definition: collections.Counter +Counts: {'green': 2, 'red': 1} ``` -**Sparkwheel composes by default!** Dicts merge and lists extend - no operators needed for the common case. +The leading `=` means **replace**. Without it, list composition appends entries. +Single quotes keep the override together as one shell argument. Values are parsed +as YAML; see [CLI overrides](../user-guide/cli.md) for booleans, nulls and mappings. -## CLI Overrides +## Experiment Without Copying -Override values from the command line without editing files: +Create an overlay beside the other files: -```python -# train.py -from sparkwheel import Config -import sys -import ast - -config = Config() -config.update("config.yaml") - -# Parse CLI overrides (simple 3-line pattern) -for arg in sys.argv[1:]: - if "=" in arg: - key, value = arg.split("=", 1) - try: - value = ast.literal_eval(value) # Parse numbers, lists, etc. - except (ValueError, SyntaxError): - pass # Keep as string - config.set(key, value) -# ... use config ... +```yaml title="extra.yaml" +words: [green] ``` -Run with overrides: +The same program accepts a file path as well as an override: ```bash -python train.py training::learning_rate=0.01 dataset::batch_size=64 +python run.py extra.yaml ``` -## Next Steps +Expected counts are `{'red': 2, 'blue': 1, 'green': 1}`. The definition line stays +`Definition: collections.Counter`. Dictionaries merge and lists extend by default; +[composition](../user-guide/operators.md) explains how to replace or delete them. -Now that you've seen the basics: +## Next Steps -- **[Core Concepts](../user-guide/basics.md)** - Learn more about references, expressions, and instantiation -- **[Composition & Operators](../user-guide/operators.md)** - Master config composition with `=` and `~` -- **[Schema Validation](../user-guide/schema-validation.md)** - Validate configs with dataclasses +- Read the [configuration model](../user-guide/basics.md) before editing objects + after construction. +- Use [`@` and `%`](../user-guide/references.md) to choose shared objects versus + copied definitions. +- Replace the standard-library target with [your own Python class](../user-guide/instantiation.md). +- If a command fails, start with [troubleshooting](../user-guide/troubleshooting.md). diff --git a/docs/index.md b/docs/index.md index f782043..006fda2 100644 --- a/docs/index.md +++ b/docs/index.md @@ -1,284 +1,54 @@ ---- -title: Sparkwheel ---- - - -# - - - - - -
- - -
-
- -```bash -pip install sparkwheel -``` -
-
- - -
- -- :material-cog-outline:{ .lg .middle } __Declarative Configuration__ - - --- - - Define complex Python objects in clean YAML files. Replace boilerplate instantiation code with simple `_target_` declarations. - -- :material-link-variant:{ .lg .middle } __Smart References__ - - --- - - Use `@` for **resolved references** (instantiated objects, computed values) or `%` for **raw references** (unprocessed YAML). Keep configurations DRY and maintainable. - -- :material-puzzle-outline:{ .lg .middle } __Flexible Composition__ - - --- - - Configs compose naturally by default (merge dicts, extend lists). Use `=` to replace or `~` to delete. Build modular configs for experiments and environments. - -- :material-function-variant:{ .lg .middle } __Python Expressions__ - - --- - - Execute code with `$` prefix. Compute values, call functions, and create dynamic configurations on the fly. - -- :material-shield-check-outline:{ .lg .middle } __Schema Validation__ - - --- - - Validate configs with Python dataclasses. Continuous validation catches errors immediately at mutation time with type checking, coercion, and required field validation. - -- :material-console:{ .lg .middle } __CLI Overrides__ - - --- - - Override any config value from command line. Perfect for hyperparameter sweeps and quick experiments. - -
- -## Python Objects from YAML - -If you're tired of **hardcoding parameters** and want **configuration-driven workflows**, Sparkwheel makes it effortless. Define components in YAML, reference and compose them freely, then instantiate in Python. - -=== "Config" - - ```yaml title="config.yaml" - dataset: - path: "/data/train" - num_classes: 10 - batch_size: 32 - - model: - _target_: torch.nn.Sequential - _args_: - - _target_: torch.nn.Linear - in_features: 784 - out_features: "@dataset::num_classes" # Reference! - - _target_: torch.nn.ReLU - - training: - epochs: 10 - learning_rate: 0.001 - steps_per_epoch: "$10000 // @dataset::batch_size" # Expression! - ``` - -=== "Python" - - ```python title="train.py" - from sparkwheel import Config - - # Load config (or multiple configs!) - config = Config() - config.update("config.yaml") - - # Access raw values - batch_size = config["dataset::batch_size"] # 32 - - # Resolve references and expressions - steps = config.resolve("training::steps_per_epoch") # 312 - - # Instantiate Python objects automatically - model = config.resolve("model") # Actual torch.nn.Sequential! - ``` - -=== "Experiment Override" - - ```yaml title="experiment_large.yaml" - # Override specific values, keep the rest (merges by default!) - model: - _args_: - - 0: # Override first layer - out_features: 20 # More classes - - training: - learning_rate: 0.0001 # Lower LR - # epochs inherited from base! - ``` - - ```python - from sparkwheel import Config - import sys - - # Load base + experiment (composes automatically!) - config = (Config() - .update("config.yaml") - .update("experiment_large.yaml")) - - # Or override from CLI (parse args yourself) - config = Config() - config.update("config.yaml") - for arg in sys.argv[1:]: - if "=" in arg: - key, value = arg.split("=", 1) - # Simple parsing - use ast.literal_eval for type conversion - try: - import ast - value = ast.literal_eval(value) - except (ValueError, SyntaxError): - pass # Keep as string - config.set(key, value) - ``` - -## Understanding References - -Sparkwheel has two types of references with distinct purposes: - -!!! abstract "@ - Resolved References" - - **Get the final, computed value** after instantiation and evaluation. - - ```yaml - model: - _target_: torch.nn.Linear - in_features: 784 - out_features: 10 - - # @ follows the reference and gets the instantiated object - trained_model: "@model" # Gets the actual torch.nn.Linear instance - ``` - - Use `@` when you want the **result** of computation. - -!!! abstract "% - Raw References" - - **Get the unprocessed YAML content** before any resolution. - - ```yaml - # base.yaml - defaults: - learning_rate: 0.001 - - # config.yaml - # % copies the raw YAML definition (can be from external files or same file) - optimizer: - lr: "%base.yaml::defaults::learning_rate" # Gets raw value: 0.001 - - # Or reference within same file - backup_defaults: "%defaults" # Gets the entire defaults dict as-is - ``` - - Use `%` when you want to **copy/import raw YAML** (like copy-paste). - -## Why Sparkwheel? - -!!! tip "Familiar, But More Powerful" - - If you've used **Hydra** or **OmegaConf**, you'll feel right at home. Sparkwheel adds: - - - **Composition-by-default** - Configs merge/extend naturally, no operators needed for common case - - **List extension** - Lists extend by default (unique vs Hydra!) - - **`=` replace operator** - Explicit control when you need replacement - - **`~` delete operator** - Remove inherited keys explicitly - - **Python expressions with `$`** - Compute values dynamically - - **Dataclass validation** - Type-safe configs without boilerplate - - **Dual reference system** - `@` for resolved values, `%` for raw YAML - - **Simpler API** - Less magic, clearer behavior - - ```yaml - # Merges by default - no operator needed! - model: - hidden_size: 1024 # Override just this - ~dropout: null # Remove dropout - # Other fields preserved automatically! - ``` - -## Start Learning - -
- -- :material-rocket-launch-outline:{ .lg .middle } __Quick Start__ - - --- - - Get productive in 5 minutes with a hands-on tutorial - - [:octicons-arrow-right-24: Quick Start](getting-started/quickstart.md) - -- :material-book-open-page-variant:{ .lg .middle } __User Guide__ - - --- - - Deep dive into references, expressions, and composition - - [:octicons-arrow-right-24: Core Concepts](user-guide/basics.md) - -- :material-code-tags:{ .lg .middle } __API Reference__ - - --- - - Complete API documentation and reference - - [:octicons-arrow-right-24: Browse API](reference/) - -
- -## Feature Deep Dives - -
- -- :material-link:{ .lg .middle } [**References**](user-guide/references.md) - - Link config values with `@` to eliminate duplication - -- :material-code-braces:{ .lg .middle } [**Expressions**](user-guide/expressions.md) - - Execute Python code in configs with `$` - -- :material-merge:{ .lg .middle } [**Composition & Operators**](user-guide/operators.md) - - Composition-by-default with `=` (replace) and `~` (delete) operators - -- :material-check-circle-outline:{ .lg .middle } [**Schema Validation**](user-guide/schema-validation.md) - - Validate with Python dataclasses - -- :material-console-line:{ .lg .middle } [**CLI Support**](user-guide/cli.md) - - Override configs from command line - -- :material-cog-transfer:{ .lg .middle } [**Instantiation**](user-guide/instantiation.md) - - Create Python objects with `_target_` - -
+# Sparkwheel + +Compose configuration, then build ordinary Python objects. + +Keep behavior in Python. Use YAML or dictionaries to choose objects, supply their +arguments and share dependencies. Sparkwheel handles composition and construction; +it does not run your application or replace its libraries. + +**This documentation describes the unpublished `0.1.0.dev0` source checkout.** +Follow [source installation](getting-started/installation.md) before the examples. +Sparkwheel requires Python 3.10+ and PyYAML; the first example needs no Torch, +Lightning, data download or service. + +## Start here + +The [quick start](getting-started/quickstart.md) builds a word counter, inspects +its definition and changes its input from the command line. Then read the +[configuration model](user-guide/basics.md) to understand what is stored and when +Python code runs. + +| Task | Guide | +|---|---| +| Install the current source | [Installation](getting-started/installation.md) | +| Run the first example | [Quick start](getting-started/quickstart.md) | +| Read and edit configuration | [Configuration model](user-guide/basics.md) | +| Share an object or create another instance | [References: `@` and `%`](user-guide/references.md) | +| Call my classes and functions | [Python authoring](user-guide/instantiation.md) | +| Merge, replace or delete settings | [Composition](user-guide/operators.md) | +| Accept command-line overrides | [CLI](user-guide/cli.md) | +| Check values against a schema | [Validation](user-guide/schema-validation.md) | +| Diagnose an error or unexpected result | [Troubleshooting](user-guide/troubleshooting.md) | +| Find syntax and public entrypoints | [Quick reference](user-guide/quick-reference.md) | +| Integrate construction into a framework | [Advanced contracts](user-guide/advanced.md) | + +## The essential distinction + +A configuration is a **definition**. Resolving it may create **live Python objects**. +`@counts` shares the resolved value of `counts` within the current resolution +mode and generation. `%counts` copies its definition; a copied `_target_` can +construct another object. Copies can themselves contain `@` references to shared +dependencies. + +`get()` reads source without configured execution. `resolve()` can execute Python. +Successful `set()` or `update()` calls start a new resolution generation; they do +not rewrite objects already returned to your application. ## About -Sparkwheel is a hard fork of [MONAI Bundle](https://github.com/Project-MONAI/MONAI/tree/dev/monai/bundle)'s configuration system, refined and expanded for general-purpose use. We're deeply grateful to the MONAI team for their excellent foundation. - -Sparkwheel powers [Lighter](https://project-lighter.github.io/lighter/), a configuration-driven deep learning framework built on PyTorch Lightning. - -**Ready to contribute?** [:octicons-mark-github-16: View on GitHub](https://github.com/project-lighter/sparkwheel) +Sparkwheel originated from MONAI Bundle's configuration system. It supports +standalone Python applications and powers configuration in +[Lighter](https://project-lighter.github.io/lighter/). -
+[Source and issues](https://github.com/project-lighter/sparkwheel) · +[API reference](reference/index.md) diff --git a/docs/user-guide/advanced.md b/docs/user-guide/advanced.md index 737d483..b8fda1f 100644 --- a/docs/user-guide/advanced.md +++ b/docs/user-guide/advanced.md @@ -1,354 +1,226 @@ # Advanced Features -## Frozen Configs - -Prevent modifications after initialization: - -```python -from sparkwheel import Config - -config = Config(schema=MySchema) -config.update("config.yaml") - -# Freeze to make immutable -config.freeze() - -# Mutations now raise FrozenConfigError -try: - config.set("model::lr", 0.001) -except FrozenConfigError as e: - print(f"Error: {e}") # Cannot modify frozen config - -# Read operations still work -value = config.get("model::lr") -resolved = config.resolve() - -# Check if frozen -if config.is_frozen(): - print("Config is frozen!") - -# Unfreeze if needed -config.unfreeze() -config.set("model::lr", 0.001) # Now works -``` +This reference is for integrations that manage construction and object ownership. +Ordinary applications can use `Config.update()` and `Config.resolve()` with their +own Python classes. Start with the [configuration model](basics.md) if you only +need to load, inspect or edit a configuration. -**Use cases:** -- Prevent accidental modifications in production -- Ensure config consistency across app lifecycle -- Debug configuration issues by freezing after initial setup +## Source inspection and resolution generations -## MISSING Sentinel +`get()` and subscript access read source without configured imports, targets or +expressions. Resolution compiles a separate working tree. Local `%` strings, +relative references and `_imports_` remain in source after resolution; external +`%file.yaml` includes are loaded snapshots from `update()`. -Support partial configs with required-but-not-yet-set values: +Successful `set()`/`update()` calls invalidate all runtime modes. Objects already +returned to the application are not rewritten. Unfrozen source-container views +remain mutable, so write through the supported edit APIs to invalidate caches. +Frozen reads return detached containers. Opaque Python leaves keep identity and +mutability; neither freezing nor retention deeply snapshots those objects. -```python -from sparkwheel import Config, MISSING -from dataclasses import dataclass - -@dataclass -class APIConfigSchema: - api_key: str - endpoint: str - timeout: int = 30 - -# Build config incrementally with MISSING values -config = Config(schema=APIConfigSchema, allow_missing=True) -config.update({ - "api_key": MISSING, # Will be set later - "endpoint": "https://api.example.com", - "timeout": 60 -}) - -# Fill in missing values from environment -import os -config.set("api_key", os.getenv("API_KEY")) - -# Validate that nothing is MISSING anymore -config.validate(APIConfigSchema) # Uses allow_missing=False by default - -# Freeze for production use -config.freeze() -``` +Migration: code that formerly read expanded local `%` values from `get()` after +resolution must use `resolve(path)` when it wants the runtime value. Source +inspection and runtime access now stay distinct. -**MISSING vs None:** -- `None` is a valid value that satisfies `Optional[T]` fields -- `MISSING` indicates a required field that hasn't been set yet -- `MISSING` raises ValidationError unless `allow_missing=True` +## Resolution modes and object sharing -**Common patterns:** +Caches are separate for each `(instantiate, eval_expr)` pair. Switching modes +cannot return the representation cached by another mode; switching back reuses +that mode's existing values. Within one mode, repeated resolution and `@` +references share identity. ```python -# Template configs with placeholders -base_config = { - "database::host": MISSING, - "database::port": MISSING, - "database::name": "myapp", - "api_key": MISSING -} - -# Environment-specific configs fill in MISSING values -config = Config(schema=MySchema, allow_missing=True) -config.update(base_config) -config.set("database::host", os.getenv("DB_HOST")) -config.set("database::port", int(os.getenv("DB_PORT"))) -config.set("api_key", os.getenv("API_KEY")) -config.validate(MySchema) # Ensure complete -``` - -## Macros (`%`) - -Load **raw YAML values** from external files using `%`: - -```yaml -# base.yaml -defaults: - learning_rate: 0.001 - -# experiment.yaml -training: - lr: "%base.yaml::defaults::learning_rate" -``` - -**Important:** `%` references get the raw YAML definition (not instantiated), while `@` references get the resolved/instantiated object from the current config. - -## Special Keys - -Sparkwheel recognizes these special keys in configuration: - -- `_target_`: Class or function path to instantiate (e.g., `"torch.nn.Linear"`) -- `_disabled_`: Skip instantiation if `true` (removed from parent). See [Instantiation](instantiation.md#_disabled_-skip-instantiation) for details. -- `_mode_`: Operating mode for instantiation (see below) -- `_imports_`: Declare imports available to all expressions (see [Imports](#imports-for-expressions) below) - -### `_mode_` - Instantiation Modes - -The `_mode_` key controls how the target is instantiated: - -- **`"default"`** (default): Returns `component(**kwargs)` - normal instantiation -- **`"callable"`**: Returns the component itself, or `functools.partial(component, **kwargs)` if kwargs provided -- **`"debug"`**: Returns `pdb.runcall(component, **kwargs)` - runs in debugger +from sparkwheel import Config -```yaml -# Example: Get a callable instead of instance -model_class: - _target_: torch.nn.Linear - _mode_: "callable" - in_features: 784 - out_features: 10 - # This returns functools.partial(torch.nn.Linear, in_features=784, out_features=10) - # instead of an instantiated Linear object - -# Example: Debug mode -buggy_component: - _target_: mymodule.BuggyClass - _mode_: "debug" # Will run in pdb debugger - param: value +config = Config().update({"object": {"_target_": "builtins.dict", "value": 8}}) +definition = config.resolve("object", instantiate=False) +instance = config.resolve("object") +assert instance == {"value": 8} +assert config.resolve("object") is instance +assert config.resolve("object", instantiate=False) is definition ``` -## Composition & Operators +`instantiate=False` returns `Component` wrappers; `eval_expr=False` returns +`Expression` wrappers. References inside wrappers still resolve in the selected +mode. Different modes may construct separate objects or evaluate expressions +separately. These options are not static inspection: configured imports may run, +and either enabled option may execute Python. -Sparkwheel uses **composition-by-default**: configs naturally merge (dicts) or extend (lists). Use operators for explicit control: +Successful source edits or `resolve(lazy=False)` start a fresh generation for all +modes. Published lambdas/generators retain their original definitions, imports, +schema policy and runtime cache; calling them later does not switch to edited +source. Caller-owned opaque objects retain their normal mutability. -### Default Behavior: Composition +For low-level resolver users, `Resolver.get_item()` retains its registered Item +definition even after `get_item(resolve=True)`. Use `resolve()` for runtime values. -By default, configs compose naturally - no operators needed: +## Retained recipes and construction scopes -```yaml -# base.yaml -model: - hidden_size: 512 - activation: "relu" - dropout: 0.1 -``` +A retained recipe captures source containers, locations and caller-provided import +bindings without copying runtime caches. Create one with `config.retain()`, then +create independent runtime scopes when an integration needs different ownership +or construction boundaries. This requires no additional YAML syntax. -```yaml -# override.yaml -model: # Merges by default! - hidden_size: 1024 # Update this - # activation and dropout are preserved! -``` +The following standalone example builds a state object in one scope and explicitly +supplies it to a second scope: ```python from sparkwheel import Config -config = Config() -config.update("base.yaml") -config.update("override.yaml") - -# Result: -# model: -# hidden_size: 1024 (updated) -# activation: "relu" (preserved - composition!) -# dropout: 0.1 (preserved - composition!) -``` - -### Replace Operator (`=`) - -Use `=key` when you need to completely replace instead of merge: - -```yaml -# override.yaml -=model: # Replace entire model dict - hidden_size: 1024 - # activation and dropout are GONE! -``` - -See [Composition & Operators](operators.md) for full details on composition-by-default and the `=` operator. - -### Delete Directive (`~`) +config = Config().update( + { + "state": {"_target_": "builtins.dict", "value": 1}, + "consumer": {"_target_": "builtins.dict", "state": "@state"}, + } +) +recipe = config.retain() +state_scope = recipe.scope(blocked_paths={"consumer"}) +state = state_scope.resolve("state") +consumer_scope = recipe.scope(bindings={"state": state}) +consumer = consumer_scope.resolve("consumer") +assert consumer["state"] is state +assert recipe.scope().resolve("state") is not state +assert set(consumer_scope.materialized_components()) == {"consumer"} +``` + +Creating a recipe or scope executes no configured imports or targets. Resolving a +scope uses a fresh namespace/cache. `@` shares within that scope; local `%` copies +are independent recipes. New scopes share only explicit bindings or caller-owned +opaque inputs. Retention is not a serializer. + +### Inspect a retained definition + +`recipe.definition(path="")` returns a detached source-container snapshot of one +node. It follows only the local `%` chain needed to select that node; it does not +expand children, execute imports, resolve `@`, evaluate expressions or construct +targets. Missing paths raise `ConfigKeyError`; opaque leaves retain identity. + +Use `Config.update()` to load external includes before retaining. Unexpanded +external includes supplied through direct `Config(data=...)` do not become later +scope file reads. Prefer `update()` for normal authoring; direct constructor data +bypasses the normal loading/validation path. + +### Bindings and blocked paths + +Bindings are authoritative **exact source paths**, checked before their definitions +or dependencies are traversed. A bound parent suppresses its subtree when resolving +that parent or an ancestor. Explicitly resolving an unbound descendant still +builds that descendant's source recipe; it does not read an attribute from the +bound object. Binding `model` and `model::network`, for example, still permits +resolving an unbound `model::optimizer` from its definition. + +Blocked paths prohibit their nodes and descendants. Explicit reachable references +are checked before imports or construction; bound subtrees are pruned. Binding a +parent and blocking a descendant is valid. Binding at or below a blocked path +fails. This is graph checking, not static analysis of arbitrary Python code. + +`scope.bind(path, value)` adds a binding for a valid source path. Rebinding the same +identity is allowed. A different identity cannot replace a value already resolved +or consumed by a cached ancestor: create a fresh scope. Import declarations are +not runtime binding paths. Scope list paths use non-negative canonical indices. + +After an expression publishes a result, all its lexical reference paths reserve +binding authority, including references Python did not visit. A new binding that +touches an unresolved reserved path requires a fresh scope. Resolving the path +normally and binding the same resulting identity is allowed. This preserves +delayed values without inspecting arbitrary returned object graphs. + +### Guards and preflight + +A statically true `_disabled_` guard, including an exact literal guard binding, +prunes its payload before preflight. The guard itself remains subject to bindings +and blocks. Dynamic guards retain conservative preflight: a blocked argument fails +before the guard executes even if the guard would later disable the node. +Ordinary `Config` has no declared block set and evaluates a dynamic guard first. + +Scopes conservatively check explicit expression references, including unselected +branches. Missing paths, blocked paths and static cycles can fail before runtime +branch selection. Neither scopes nor disabled nodes sandbox arbitrary Python. + +### Construction inventory and limits + +`scope.materialized_components()` returns a defensive mapping from `_target_` +paths to objects constructed in that scope. It excludes supplied bindings, +disabled components, expressions and iterator outputs, including iterators returned +by a target. The mapped objects remain live; editing the returned dictionary does +not change bookkeeping. This inventory records construction, not subsequent use +or application correctness. + +Attached schemas are explicitly unsupported by `retain()`; use schema-aware +`Config.resolve()` instead. A scope must not silently discard a schema policy. +Ordinary `Config.resolve()` retains its own cache and ownership behavior. + +## Atomic edits and container paths + +`set()` and `update()` stage the whole edit, check its raw schema and publish only +on success. Edit, loading or validation failure preserves prior source, locations, +expression namespace and cached identities. Custom validator effects and file reads +are outside the transaction. Opaque inputs retain identity. + +Paths respect dictionaries and existing list indices. Traversing a scalar or an +out-of-range list index fails. Missing dictionary paths can be created. Single +indexed deletion and batch deletion are both supported; see +[composition](operators.md#editing-nested-paths). -Use `~key: null` to delete a key, or `~key: [items]` to delete specific items from lists/dicts: - -```yaml -# override.yaml -~model::dropout: null # Remove entire key - -# Remove specific list items by index -~plugins: [0, 2, 4] # Remove items at indices 0, 2, 4 - -# Remove specific dict keys -~dataloaders: ["train", "test"] # Remove these keys +## Frozen Configs -# Negative indices work too -~plugins: [-1] # Remove last item -``` +`freeze()` detaches source containers from earlier views and blocks supported +edits. Frozen reads produce detached ordinary dictionaries/lists suitable for +source export. Opaque leaves and runtime objects are not deeply frozen. Use +`unfreeze()` followed by `set()`/`update()` for later source changes. -```python -config = Config() -config.update("base.yaml") -config.update({"~model::dropout": None}) # Remove entire key -config.update({"~plugins": [0, 2]}) # Remove list items -config.update({"~dataloaders": ["train", "test"]}) # Remove dict keys -``` +## Imports for Expressions -### Programmatic Updates +`_imports_` declares serializable root imports. It remains in `get()` and source +export but is omitted from the resolved runtime dictionary. After an edit, +resolution gets a fresh namespace: removed YAML imports and `$import ...` aliases +do not remain available. Imports supplied through `Config(imports=...)` remain +available across generations; string import paths passed to that constructor are +loaded at construction time. -Apply operators programmatically: +This complete example exports source for replay: ```python -config = Config() -config.update("config.yaml") - -# Set individual values -config.set("model::hidden_size", 1024) - -# Use operators -config.update({ - "optimizer": {"lr": 0.01}, # Compose (merge by default) - "=database": {"host": "prod.db"}, # Replace - "~training::old_param": None, # Delete -}) -``` - -## Relative ID References - -Use relative references to navigate the config hierarchy: - -```yaml -model: - encoder: - hidden_size: 512 - activation: "relu" - decoder: - # Reference sibling section - hidden_size: "@::encoder::hidden_size" # Same level (model) - # Reference parent level - loss_fn: "@::::training::loss" # Go up to root, then to training -``` - -**Syntax:** -- `@::` - Same level (sibling) -- `@::::` - Parent level -- Add more `::` to go up more levels - -## Enhanced Error Messages - -Sparkwheel provides helpful error messages with suggestions: +from sparkwheel import Config -```python -from sparkwheel import Config, ConfigKeyError - -config = Config() -config.update({ - "model": {"hidden_size": 512, "num_layers": 4}, - "training": {"batch_size": 32} -}) - -try: - # Typo in key name - value = config.resolve("model::hiden_size") -except ConfigKeyError as e: - print(e) - # Output: - # Config ID 'model::hiden_size' not found - # - # Did you mean one of these? - # - model::hidden_size - # - model::num_layers +config = Config().update({"_imports_": {"math": "math"}, "rounded": "$math.ceil(2.4)"}) +assert config.resolve("rounded") == 3 +Config.export_config_file(config.get(), "snapshot.yaml") +replayed = Config().update("snapshot.yaml") +assert replayed.resolve("rounded") == 3 ``` -Color output is auto-detected and respects `NO_COLOR` environment variable. +The output `snapshot.yaml` preserves configuration data, not Python module state +or arbitrary prebuilt objects. Use import paths in `_imports_` rather than only +caller-supplied bindings when a standalone YAML export needs those imports. -## Imports for Expressions - -Make modules available to all expressions. There are two ways to do this: - -### Method 1: `_imports_` Key in YAML +## Relative ID References -Declare imports directly in your config file: +`@::name` moves from the current value path to its parent and selects `name`. +Each extra `::` moves one more level up. For example: ```yaml -# config.yaml -_imports_: - torch: torch - np: numpy - Path: pathlib.Path - -# Now use them in expressions -device: "$torch.device('cuda' if torch.cuda.is_available() else 'cpu')" -data: "$np.array([1, 2, 3])" -save_path: "$Path('/data/models')" +group: + value: 3 + same: "@::value" # @group::value + nested: + same: "@::::value" # @group::value ``` -The `_imports_` key is removed from the config after processing—it won't appear in your resolved config. - -### Method 2: `imports` Parameter in Python - -Pass imports when creating the Config: +Going above the root is an error. The same relative path machinery supports `%` +copies. Prefer absolute paths when they make ownership clearer. -```python -from sparkwheel import Config - -# Pre-import modules for all expressions -config = Config(imports={"torch": "torch", "np": "numpy"}) -config.update("config.yaml") - -# Now expressions can use torch and np without importing -``` - -### Combining Both Methods - -You can use both approaches together—they merge: - -```python -from collections import Counter - -config = Config(imports={"Counter": Counter}) -config.update({ - "_imports_": {"json": "json"}, - "data": '$json.dumps({"a": 1})', - "counts": "$Counter([1, 1, 2])" -}) -``` +## MISSING Sentinel -## Type Hints +`MISSING` represents a required value not yet supplied; `None` is a value that must +be permitted by the field's type. See [schema validation](schema-validation.md#missing-sentinel) +for incremental configuration and completeness checking. -```python -from sparkwheel import Config +## Enhanced Error Messages -config: Config = Config() -config.update("config.yaml") -resolved: dict = config.resolve() -``` +Source locations and native exception causes are part of useful diagnostics. +[The troubleshooting guide](troubleshooting.md) provides a staged workflow rather +than recommending full resolution as a static check. -For complete details, see the [API Reference](../reference/). +For normal entrypoints versus lower-level implementation machinery, use the +[public API map](quick-reference.md#python-api). diff --git a/docs/user-guide/basics.md b/docs/user-guide/basics.md index 7a36ddb..9999007 100644 --- a/docs/user-guide/basics.md +++ b/docs/user-guide/basics.md @@ -1,463 +1,142 @@ # Configuration Basics -Learn the fundamentals of Sparkwheel configuration files. +Sparkwheel separates the configuration you author from the objects your program +uses. Understanding that distinction makes inspection, sharing and edits predictable. ## Configuration File Format -Sparkwheel uses YAML for configuration: +A configuration contains ordinary values plus a few construction directives: ```yaml -# config.yaml -name: "My Project" -version: 1.0 -settings: - debug: true - timeout: 30 +words: [red, blue, red] +counts: + _target_: collections.Counter + _args_: ["@words"] ``` -YAML provides excellent readability and native support for comments, making it ideal for configuration files. +This is a recipe for calling `Counter(words)`. Your Python application decides +when to resolve it and what to do with the result. ## Loading Configurations -### Basic Loading +Use `update()` for files, dictionaries and CLI overrides. Later inputs compose +with earlier ones; dictionaries merge, lists extend and other values replace. ```python from sparkwheel import Config -# Load from file -config = Config() -config.update("config.yaml") +config = Config().update({"settings": {"limit": 3}, "words": ["red"]}) +config.update({"settings": {"title": "Report"}}) +config.update("settings::limit=5") +assert config.get("settings") == {"limit": 5, "title": "Report"} ``` -### Loading from Dictionary - -```python -config_dict = { - "name": "Test", - "value": 42 -} - -# Load from dict -config = Config() -config.update(config_dict) -``` - -### Loading Multiple Files - -```python -# Load and merge multiple config files (method chaining!) -config = (Config() - .update("base.yaml") - .update("override.yaml")) -``` +For files, the equivalent loading pattern is +`Config().update("base.yaml").update("experiment.yaml")`. +Read [composition](operators.md) for explicit replacement and deletion. ## Accessing Configuration Values -Sparkwheel provides two equivalent syntaxes for accessing nested configuration values: - -### Two Ways to Access Nested Values - -```python -config = Config() -config.update("config.yaml") - -# Method 1: Standard nested dictionary access -name = config["name"] -debug = config["settings"]["debug"] -lr = config["model"]["optimizer"]["lr"] - -# Method 2: Path notation with :: separator -debug = config["settings::debug"] -lr = config["model::optimizer::lr"] - -# Both methods work identically! -assert config["settings"]["debug"] == config["settings::debug"] -``` - -**When to use each:** - -- **Nested access** (`config["a"]["b"]`) - Familiar Python syntax, works like any dict -- **Path notation** (`config["a::b"]`) - More concise for deeply nested values, easier to pass as strings +`::` addresses dictionary keys and list indices: `settings::limit`, `words::0`. +Unresolved source containers also support ordinary Python reads such as +`config["settings"]["limit"]`. ### Using get() and resolve() -The same two syntaxes work with `get()` and `resolve()`: - -```python -# Method 1: Nested access -raw_value = config.get("model")["optimizer"]["lr"] - -# Method 2: Path notation (more convenient) -raw_value = config.get("model::optimizer::lr") - -# Both work with resolve() too -debug_mode = config.resolve("settings::debug") -debug_mode = config.resolve("settings")["debug"] # Also works - -# Resolve entire config -all_config = config.resolve() - -# Resolve specific section -training_config = config.resolve("training") -``` - -**Key difference:** -- `get()` returns raw values (references like `"@model::lr"` are not resolved) -- `resolve()` resolves references, evaluates expressions, and instantiates objects - -## Choosing Between Syntaxes - -Both syntaxes have their place: - -### Use Path Notation (`::`) When: - -```python -# 1. Passing paths as function arguments -def get_param(config, path: str): - return config.get(path) - -lr = get_param(config, "model::optimizer::lr") - -# 2. Working with very deep nesting (more readable) -value = config["a::b::c::d::e"] - -# 3. Setting values programmatically -config.set("model::optimizer::lr", 0.001) - -# 4. Matching reference syntax in YAML -# YAML: lr: "@model::optimizer::base_lr" -base_lr = config.get("model::optimizer::base_lr") -``` - -### Use Standard Dict Access When: - -```python -# 1. You want to work with intermediate sections -model_config = config["model"] -model_config["dropout"] = 0.1 -model_config["lr"] = 0.001 - -# 2. Iterating over config sections -for key in config["training"].keys(): - print(key, config["training"][key]) - -# 3. It feels more natural for your use case -settings = config["app"]["settings"] -if settings["debug"]: - print("Debug mode enabled") -``` - -## Configuration Structure - -### Nested Structures - -```yaml -project: - name: "Sparkwheel Demo" - version: 1.0 - - database: - host: "localhost" - port: 5432 - credentials: - username: "admin" - password: "secret" - - features: - authentication: true - logging: true -``` - -Access nested values with either syntax: - -```python -# Path notation (concise) -db_host = config.resolve("project::database::host") -username = config.resolve("project::database::credentials::username") - -# Standard dict access (also works) -db_host = config.resolve("project")["database"]["host"] -username = config["project"]["database"]["credentials"]["username"] -``` - -### Lists and Arrays - -```yaml -colors: - - red - - green - - blue - -matrix: - - [1, 2, 3] - - [4, 5, 6] - - [7, 8, 9] -``` - -Access list elements with either syntax: - -```python -# Path notation -first_color = config.resolve("colors::0") # "red" -matrix_row = config.resolve("matrix::1") # [4, 5, 6] - -# Standard list access -first_color = config["colors"][0] # "red" -matrix_row = config["matrix"][1] # [4, 5, 6] -``` - -## Configuration Sections - -### Organizing Large Configs - -Break large configurations into logical sections: - -```yaml -# Application settings -app: - name: "My App" - version: "2.0.0" - debug: false - -# Database configuration -database: - host: "localhost" - port: 5432 - pool_size: 10 - -# Logging configuration -logging: - level: "INFO" - format: "%(asctime)s - %(name)s - %(levelname)s - %(message)s" - handlers: - - console - - file - -# Training configuration -training: - batch_size: 32 - epochs: 100 - learning_rate: 0.001 -``` - -## Configuration Validation +| Operation | Meaning | Can execute configured Python? | +|---|---|---| +| `config.get(path)` or `config[path]` | Read source at a path | No | +| `config.resolve(path)` | Resolve that node and its dependencies | Yes | +| `config.resolve()` | Resolve the whole configuration | Yes | +| `config.validate(Schema)` | Check source against a dataclass schema | Custom validators may execute Python | -### Schema Validation with Dataclasses +`get()` keeps local `%` strings, `@` references, expressions and `_imports_` +declarations even after resolution. External `%file.yaml` includes are already +loaded into the source when `update()` processes them. -Sparkwheel supports automatic validation using Python dataclasses with **continuous validation** - errors are caught immediately when you mutate the config: +Resolving may import modules, evaluate expressions and call constructors. It is +not a static validation pass. Use trusted configurations, just as you would use +trusted Python modules. Partial resolution avoids unrelated runtime targets, +but configuration-wide import declarations still have setup behavior. -```python -from dataclasses import dataclass -from sparkwheel import Config - -@dataclass -class AppConfigSchema: - name: str - version: str - port: int - debug: bool = False - -# Continuous validation - validates on every update/set! -config = Config(schema=AppConfigSchema) -config.update("config.yaml") - -# This will raise ValidationError immediately -config.set("port", "not a number") # ✗ Error caught at mutation time! - -# Or validate explicitly after mutations -config = Config() -config.update("config.yaml") -config.validate(AppConfigSchema) -``` +## Editing Configuration -Schema validation provides: -- **Continuous validation**: Errors caught immediately at mutation time (when schema provided to `Config()`) -- **Type checking**: Ensures values have the correct types -- **Type coercion**: Automatically converts compatible types (e.g., `"8080"` → `8080`) -- **Required fields**: Catches missing configuration -- **Clear errors**: Points directly to the problem with helpful messages - -See the [Schema Validation Guide](schema-validation.md) for complete details. - -### Manual Validation - -You can also validate manually: +Use `set()` or `update()` for every intended source change: ```python from sparkwheel import Config -# Load config -config = Config() -config.update("config.yaml") - -# Validate required keys -required_keys = ["name", "version", "settings"] -for key in required_keys: - if key not in config: - raise ValueError(f"Missing required key: {key}") - -# Validate by attempting resolution -try: - resolved = config.resolve() - print("Config resolved successfully!") -except Exception as e: - print(f"Config validation failed: {e}") +config = Config().update({"settings": {"limit": 3}, "double": "$@settings::limit * 2"}) +assert config.resolve("double") == 6 +config.set("settings::limit", 5) +assert config.resolve("double") == 10 ``` -## Best Practices +Unfrozen raw dictionaries/lists are mutable views. Writing through +`config["settings"]["limit"]` bypasses transactional validation and cache +invalidation. Use that syntax for reads; use `set("settings::limit", value)` for +edits. Existing list indices work with `set()` too; sparse list assignment does not. -### 1. Use Descriptive Keys +A successful edit publishes the complete candidate source and clears runtime +caches. If composition, loading or schema validation fails, the prior source and +cached identities remain intact. Multi-field `update()` calls validate together. +Custom validator side effects and file reads are outside that transaction. -```yaml -# Good -database_connection_pool_size: 10 -max_retry_attempts: 3 +## Object Lifetimes -# Avoid -db_pool: 10 -retries: 3 -``` +Repeated `resolve(path)` calls reuse the same value within a resolution mode and +generation. `@` references participate in that sharing. Successful `set()` or +`update()` calls, or `resolve(lazy=False)`, start a fresh generation. -### 2. Group Related Settings - -```yaml -# Good - grouped by feature -email: - smtp_host: "smtp.gmail.com" - smtp_port: 587 - from_address: "noreply@example.com" - -# Avoid - scattered -smtp_host: "smtp.gmail.com" -smtp_port: 587 -email_from: "noreply@example.com" -``` - -### 3. Use Comments - -```yaml -training: - batch_size: 32 # Optimal for 16GB GPU - learning_rate: 0.001 # Recommended by paper X - - # Experimental: improved convergence - warmup_steps: 1000 -``` - -### 4. Separate Environment-Specific Config - -```yaml -# base_config.yaml -common: - app_name: "My App" - features: - caching: true - -# dev_config.yaml -environment: development -debug: true -database: - host: "localhost" - -# prod_config.yaml -environment: production -debug: false -database: - host: "prod-db.example.com" -``` - -## Configuration Inheritance - -Load and merge multiple config files: +Objects already returned to your program remain the old objects. Editing the +recipe does not change a running counter, model, callback or optimizer. This +complete example demonstrates the boundary: ```python from sparkwheel import Config -import ast - -# Method 1: Chain updates (recommended!) -config = (Config() - .update("base_config.yaml") - .update("prod_config.yaml")) - -# Method 2: Sequential updates -config = Config() -config.update("base_config.yaml") -config.update("prod_config.yaml") - -# Method 3: With CLI overrides (manual parsing) -config = Config() -config.update("override.yaml") -# Parse CLI args yourself - simple! -for arg in ["model::lr=0.001"]: - if "=" in arg: - key, value = arg.split("=", 1) - try: - value = ast.literal_eval(value) - except (ValueError, SyntaxError): - pass - config.set(key, value) - -# Later configs override earlier ones -resolved = config.resolve() -``` - -See [Composition & Operators](operators.md) for details on composition-by-default, replace (`=`), and delete (`~`) operators. - -## Special Keys -Sparkwheel reserves certain keys with special meaning: - -- `_target_`: Specifies a class to instantiate -- `_disabled_`: Skip instantiation if true -- `_mode_`: Instantiation mode (default, callable, debug) -- `_imports_`: Declare imports available to all expressions - -These are covered in detail in [Instantiation Guide](instantiation.md) and [Advanced Features](advanced.md). +config = Config().update( + { + "words": ["red", "red"], + "counts": {"_target_": "collections.Counter", "_args_": ["@words"]}, + } +) +old = config.resolve("counts") +assert config.resolve("counts") is old +config.set("words", ["blue"]) +new = config.resolve("counts") +assert new is not old +assert dict(old) == {"red": 2} +assert dict(new) == {"blue": 1} +``` -## Common Patterns +Different `instantiate`/`eval_expr` modes have separate caches. Retained callbacks +and generators keep their original generation; caller-owned opaque Python objects +keep ordinary identity and mutability. See [advanced contracts](advanced.md) for +the precise boundaries. -### Default Values +## Frozen Configs -```yaml -defaults: - timeout: 30 - retries: 3 - debug: false - -# Override specific values -api: - timeout: "@defaults::timeout" - retries: 5 # Override default - debug: "@defaults::debug" -``` +`config.freeze()` prevents `set()` and `update()`. Frozen raw reads return detached +container snapshots, and freezing detaches views handed out earlier. Use +`config.unfreeze()` before another supported edit. -### Feature Flags +Freezing protects configuration containers. It does not freeze live resolved +objects or arbitrary Python objects supplied by the caller. -```yaml -features: - authentication: true - rate_limiting: true - caching: false - analytics: true - -# Reference in other parts -api: - enable_auth: "@features::authentication" - enable_cache: "@features::caching" -``` +## Duplicate YAML keys -### Environment Variables +Sparkwheel rejects duplicate authored YAML keys at every mapping level and reports +both locations. Keys YAML considers equal, such as `true` and `1`, also collide. +YAML `<<` merges remain legal: explicit values override inherited ones; earlier +mappings in a merge sequence take precedence. -```yaml -database: - # Use environment variable with fallback - host: "$import os; os.getenv('DB_HOST', 'localhost')" - port: "$import os; int(os.getenv('DB_PORT', '5432'))" -``` +Compose separate files or use `update()` for deliberate changes. For migration +only, `SPARKWHEEL_STRICT_KEYS=0` restores warning/last-wins behavior; +`SPARKWHEEL_STRICT_KEYS=1` explicitly enables strict handling. A rejected file or +include leaves the previous source generation intact. ## Next Steps -- [References](references.md) - Link configuration values -- [Expressions](expressions.md) - Execute Python code -- [Instantiation](instantiation.md) - Create objects from config -- [Advanced Features](advanced.md) - Power user techniques +[Sharing and copying](references.md) · [Composition](operators.md) · +[Validation](schema-validation.md) · [Troubleshooting](troubleshooting.md) diff --git a/docs/user-guide/cli.md b/docs/user-guide/cli.md index f32f8a4..99b99bc 100644 --- a/docs/user-guide/cli.md +++ b/docs/user-guide/cli.md @@ -1,193 +1,96 @@ # CLI Overrides -Override configuration values from the command line with automatic file and override detection. +Sparkwheel accepts input strings; your application owns its command line. The +[quick-start program](../getting-started/quickstart.md) is a complete working CLI. +Its integration is simply: -## Auto-Detection Pattern - -!!! success "Dead Simple CLI Integration" - **Just loop over CLI arguments** - `config.update()` automatically detects whether each string is a file path or an override! - - - Strings **with** `=` → Parsed as overrides (e.g., `key=value`, `=key=value`, `~key`) - - Strings **without** `=` → Loaded as file paths - - No manual separation needed! - -## Quick Start - -=== "argparse" - - ```python - import argparse - from sparkwheel import Config - - parser = argparse.ArgumentParser() - parser.add_argument("inputs", nargs="+") - args = parser.parse_args() - - config = Config() - for item in args.inputs: - config.update(item) - - model = config.resolve("model") - ``` - - ```bash - python train.py base.yaml exp.yaml optimizer::lr=0.01 model::dropout=0.1 - ``` - -=== "Click" - - ```python - import click - from sparkwheel import Config - - @click.command() - @click.argument("inputs", nargs=-1, required=True) - def train(inputs): - config = Config() - for item in inputs: - config.update(item) - - model = config.resolve("model") - - if __name__ == "__main__": - train() - ``` - - ```bash - python train.py base.yaml exp.yaml optimizer::lr=0.01 model::dropout=0.1 - ``` - -=== "Typer" - - ```python - import typer - from sparkwheel import Config - - app = typer.Typer() - - @app.command() - def train(inputs: list[str] = typer.Argument(None)): - config = Config() - for item in inputs or []: - config.update(item) - - model = config.resolve("model") - - if __name__ == "__main__": - app() - ``` - - ```bash - python train.py base.yaml exp.yaml optimizer::lr=0.01 model::dropout=0.1 - ``` +```python +for value in sys.argv[1:]: + config.update(value) +``` -=== "Fire" +This fragment belongs after creating the config and loading the base file. +It supports file overlays and overrides in their supplied order. Use argparse, +Click or another CLI library when the application needs its own options; pass +only the configuration inputs to `update()`. - ```python - import fire - from sparkwheel import Config +## Auto-Detection Pattern - class TrainCLI: - def train(self, *inputs): - config = Config() - for item in inputs: - config.update(item) +- A string containing `=` is parsed as an override. +- A string beginning with `~` is a delete override, even without `=`. +- Other strings are file paths. - model = config.resolve("model") +There is no `--key value` configuration syntax. A filename containing `=` or +starting with `~` is ambiguous as a string; pass a `pathlib.Path` when loading it +through Python. - if __name__ == "__main__": - fire.Fire(TrainCLI) - ``` +## Override Syntax - ```bash - python train.py train base.yaml exp.yaml optimizer::lr=0.01 model::dropout=0.1 - ``` +| Input | Meaning | +|---|---| +| `report::title=Example` | Set a nested value | +| `words=[green]` | Append to an existing list | +| `=words=[green]` | Replace the list | +| `=report={title: Example}` | Replace the whole mapping | +| `~report::limit` | Delete one key | +| `~words::0` | Delete one list entry | -## Override Syntax +Quote complete overrides containing spaces or shell punctuation. From the +quick-start directory, these are complete commands: -Three operators for fine-grained control: +```bash +python run.py 'words=[green]' +python run.py '=words=[green, green]' +``` -| Operator | Syntax | Behavior | Example | -|----------|--------|----------|---------| -| **Compose** (default) | `key=value` | Merges dicts, extends lists | `model::lr=0.001` | -| **Replace** | `=key=value` | Completely replaces value | `=model={'_target_': 'ResNet'}` | -| **Delete** | `~key` | Removes key (errors if missing) | `~debug` | +Their counts are respectively `{'red': 2, 'blue': 1, 'green': 1}` and +`{'green': 2}`. In both cases the program also prints the unchanged target name. -!!! info "Type Inference" - Values are automatically typed using `ast.literal_eval()`: +## Value Types - - `lr=0.001` → `float` - - `epochs=100` → `int` - - `debug=True` → `bool` - - `devices=[0,1,2]` → `list` - - `config={'lr':0.001}` → `dict` - - `name=resnet50` → `str` (fallback) +Override values use **YAML parsing**, not Python `ast.literal_eval`: -!!! note "The `=` Dual Purpose" - - In `key=value` → Assignment operator (CLI syntax) - - In `=key=value` → Replace operator prefix (config operator) +| Input value | Parsed value | +|---|---| +| `3`, `0.5` | Integer, float | +| `true`, `false` | Boolean | +| `yes`, `no`, `on`, `off` | Boolean under the current YAML parser | +| `null`, `~` | `None` | +| `[red, blue]` | List of strings | +| `{limit: 3}` | Dictionary | +| `None` | The string `"None"` | -!!! tip "Adding Your Own Flags" - The examples above show minimal integration. You can add your own flags (e.g., `--verbose`, `--device`) alongside the config inputs - Sparkwheel only cares about the arguments you pass to `config.update()`! +To preserve a YAML-looking value as text, preserve the inner quotes through the +shell: `config.update('label="true"')` stores a string. The equivalent shell +argument is `'label="true"'`. Shell quoting and YAML quoting solve different +problems: the outer quotes keep one argument; the inner quotes control its type. ## Advanced: Manual Override Parsing -If you need to separate override parsing from application, use `parse_overrides()`: +Use `parse_overrides()` when parsing and application must be separate: ```python from sparkwheel import Config, parse_overrides -# Manually parse overrides -overrides = parse_overrides(["model::lr=0.001", "=optimizer={'type':'sgd'}", "~debug"]) -# Result: {"model::lr": 0.001, "=optimizer": {"type": "sgd"}, "~debug": None} - -config = Config() -config.update("base.yaml") +overrides = parse_overrides(["limit=3", "enabled=true", "=words=[blue]", "~old"]) +assert overrides == {"limit": 3, "enabled": True, "=words": ["blue"], "~old": None} +config = Config().update({"old": 1, "words": ["red"]}) config.update(overrides) +assert config.get("words") == ["blue"] +assert config.get("old") is None ``` -!!! warning "parse_overrides() Syntax" - `parse_overrides()` **only** supports `key=value` syntax (no `--key value` flag style). +The parser handles overrides, not filenames. Avoid maintaining a second parser +with different boolean, null or operator rules. ## Schema Validation -Add continuous validation with dataclasses: - -```python -import argparse -from dataclasses import dataclass -from sparkwheel import Config - -@dataclass -class TrainingConfig: - model: dict - optimizer: dict - trainer: dict - -parser = argparse.ArgumentParser() -parser.add_argument("inputs", nargs="+") -args = parser.parse_args() - -# Validates on every update! -config = Config(schema=TrainingConfig) -for item in args.inputs: - config.update(item) # Raises ValidationError if invalid - -config.freeze() # Lock the config -model = config.resolve("model") -``` - -**Usage:** -```bash -python train.py base.yaml optimizer::lr=0.001 trainer::epochs=100 -``` - -!!! warning "Validation Errors" - Invalid overrides raise `ValidationError` immediately - helps catch config errors early! +A schema-enabled config validates each complete `update()` transaction. Deferred +references and expressions are checked when resolved. Compose a complete initial +mapping when required fields are needed together; see +[validation timing](schema-validation.md#validation-timing). ## Next Steps -- **[Configuration Basics](basics.md)** - Loading and accessing configs -- **[Operators](operators.md)** - Composition, replacement, and deletion -- **[Schema Validation](schema-validation.md)** - Type-safe configs with dataclasses -- **[API Reference](reference/)** - Full API documentation +[Composition](operators.md) · [Troubleshooting](troubleshooting.md) · +[Quick reference](quick-reference.md) diff --git a/docs/user-guide/expressions.md b/docs/user-guide/expressions.md index 6468efa..c52fb2a 100644 --- a/docs/user-guide/expressions.md +++ b/docs/user-guide/expressions.md @@ -1,250 +1,140 @@ # Expressions -Execute Python code directly in your configuration files using the `$` prefix. +Use `$` for a small Python expression whose value depends on configuration. Keep +control flow, exception handling and substantial application logic in +[ordinary Python](instantiation.md). ## Basic Expressions -```yaml -# Simple math -result: "$2 + 2" # 4 -square: "$10 ** 2" # 100 - -# String operations -message: "$'Hello, ' + 'World!'" - -# Lists -numbers: "$[1, 2, 3, 4, 5]" -squares: "$[x**2 for x in range(5)]" -``` - -## Combining with References +This complete example evaluates arithmetic and a derived count: -Expressions can use references: +```python +from sparkwheel import Config -```yaml -training: - batch_size: 32 - total_samples: 10000 - steps_per_epoch: "$@training::total_samples // @training::batch_size" +config = Config().update( + { + "words": ["red", "blue", "red"], + "total": "$len(@words)", + "double": "$@total * 2", + } +) +assert config.resolve("total") == 3 +assert config.resolve("double") == 6 ``` -## Importing Modules - -Import Python libraries in expressions: - -```yaml -# Math operations -pi: "$import math; math.pi" -sqrt_2: "$import math; math.sqrt(2)" +The YAML equivalent of the derived node is `total: "$len(@words)"`. +Lists, comprehensions, conditional expressions and lambdas follow Python's rules. +Assignment and arbitrary statement blocks such as `try/except` are not supported. -# Check for GPU -device: "$import torch; 'cuda' if torch.cuda.is_available() else 'cpu'" +## Combining with References -# Get environment variable -db_host: "$import os; os.getenv('DB_HOST', 'localhost')" -``` +Python resolves expression references when their values are used. In ordinary +`Config` resolution, `$@a if True else @bad` uses `a` without evaluating or +constructing `bad`. The same applies to `and`, `or`, comprehensions, lambdas and +generators. Whole-value `@path` references preserve shared identity in the selected +resolution mode. -## Complex Expressions +Quoted strings, raw strings, triple-quoted strings and comments preserve literal +`@` text. `$'me@example.com'` returns the email address. Python matrix multiplication +remains an operator. Whole-value references such as `"@learning-rate"` can address +hyphenated keys; expression references use the `@word::path` grammar. For a +punctuated nested key use indexing, such as `$@settings["learning-rate"]`. -### Multi-line Logic +F-string literal text is preserved, including `@`, but configuration references +inside f-string replacement fields are not supported. Use +`$"value={}".format(@value)` instead. -```yaml -learning_rate: "$ - 0.001 if @training::batch_size < 64 - else 0.0001 if @training::batch_size < 128 - else 0.00001 -" -``` +Retained construction scopes conservatively check explicit references before +execution, including unselected branches. Missing or blocked paths and static +cycles can therefore fail in a scope even when ordinary lazy execution would not +visit them. This distinction is described in [advanced contracts](advanced.md). -### List Comprehensions +## Importing Modules -```yaml -# Generate range -values: "$list(range(10))" +Declare reusable expression imports at the configuration root. This is a complete +standalone example: -# Transform data -scaled: "$[x / 255.0 for x in @raw_values]" +```python +from sparkwheel import Config -# Filter -evens: "$[x for x in @numbers if x % 2 == 0]" +config = Config().update( + { + "_imports_": {"math": "math"}, + "rounded": "$math.ceil(2.4)", + } +) +assert config.get("rounded") == "$math.ceil(2.4)" +assert config.resolve("rounded") == 3 ``` -### Function Definitions +Alternatively, prefix an expression with imports: ```yaml -# Define and call function -processed: "$ - (lambda x: x ** 2 + 2 * x + 1)(@input_value) -" +rounded: "$import math; math.ceil(2.4)" +home: "$import os; os.getenv('HOME', '')" ``` -## Calling Object Methods +An import prefix may contain one or more `import`/`from` statements followed by +one final Python expression. All aliases follow native Python import rules; the +final expression supplies the result. A bare import returns its first imported +binding. Arbitrary assignments or statement blocks are not part of this form. -Reference object methods when instantiation is involved: +Import declarations are made available to the generation before ordinary +expression evaluation. A trailing expression runs only when its node is requested +or used as a dependency. Root `_imports_` declarations are setup, not lazy target +construction. See [disabled payloads](instantiation.md#_disabled_-skip-instantiation) +for their more specific runtime boundary. -```yaml -model: - _target_: torch.nn.Linear - in_features: 784 - out_features: 10 - -optimizer: - _target_: torch.optim.Adam - lr: 0.001 - params: "$@model.parameters()" # Call model's parameters() method -``` +Migration: older versions could discard aliases or the expression after an import. +All listed aliases now import, and the final expression executes when requested; +its failures are not hidden. ## Expression Scope -### Global Scope - -Imports are added to global scope: - -```yaml -setup: "$import numpy as np" # np is now available globally - -data: - array: "$np.array([1, 2, 3])" # Can use np here -``` - -### Local Variables - -Access config values as variables: - -```yaml -value_a: 10 -value_b: 20 -sum: "$value_a + value_b" # Error: use @value_a instead -``` - -Use references (`@`) to access config values. - -## Common Patterns - -### Environment Detection - -```yaml -config: - is_production: "$import os; os.getenv('ENV') == 'production'" - debug: "$not @config::is_production" -``` - -### Conditional Paths - -```yaml -paths: - base: "/data" - train: "$@paths::base + '/train' if @mode == 'train' else @paths::base + '/val'" -``` - -### Dynamic Imports +Configuration keys are accessed with `@`, not injected as Python variables. +`$value_a + value_b` does not read those config keys; use `$@value_a + @value_b`. +Imported aliases are available within the configuration generation. -```yaml -backend: "torch" +Returned lambdas and generators retain the generation that built them. After a +successful source edit or `resolve(lazy=False)`, an older callable still uses its +original definitions, imports, schema policy and runtime cache. This is not a +deep snapshot of arbitrary caller-owned Python objects. -tensor_fn: "$ - __import__('torch').tensor if @backend == 'torch' - else __import__('tensorflow').constant -" -``` +Within a retained scope, an expression's lexical reference paths reserve binding +authority after it publishes a result, including references in unselected +branches. Adding a binding that touches an unresolved reserved path requires a +fresh scope. Resolving that path normally and binding the same resulting identity +is allowed. No inspection of arbitrary returned object graphs is promised. -### Calculate Derived Values +## Calling Object Methods -```yaml -model: - input_shape: [3, 224, 224] - input_size: "$@model::input_shape[0] * @model::input_shape[1] * @model::input_shape[2]" -``` +Expressions may call methods of resolved objects. For the quick-start definition, +`$@counts.most_common(1)` calls the counter's method. That operation can have the +same side effects as a native Python call. Use a Python helper when a method chain +or condition becomes difficult to read. ## Error Handling -### Syntax Errors - -```yaml -# Bad: Python syntax error -result: "$2 +" # SyntaxError -``` - -### Runtime Errors - -```yaml -# Bad: NameError -result: "$undefined_variable" # Will raise error - -# Good: Use references -value: 10 -result: "$@value * 2" -``` - -### Safe Evaluation +Expression parsing/evaluation failures raise contextual `EvaluationError` with +the original exception as their cause. Lazy dependency and schema failures retain +their structured error types. Read the cause and referenced path; do not replace +all errors with a silent default. -Check before using: - -```yaml -# Check if module exists -has_torch: "$ - try: - import torch - True - except ImportError: - False -" - -device: "$'cuda' if @has_torch and torch.cuda.is_available() else 'cpu'" -``` - -## Best Practices - -### 1. Keep Expressions Simple - -```yaml -# Good -steps: "$@samples // @batch_size" - -# Avoid -steps: "$ - sum([1 for _ in range(@samples)]) // - (lambda x: x if x > 0 else 1)(@batch_size) -" -``` - -### 2. Use Comments - -```yaml -# Calculate learning rate based on batch size -# Formula from paper: lr = base_lr * sqrt(batch_size) -learning_rate: "$0.001 * (@training::batch_size ** 0.5)" -``` - -### 3. Validate Expressions - -```python -# In your Python code -from sparkwheel import Config - -try: - config = Config() - config.update("config.yaml") - resolved = config.resolve() -except SyntaxError as e: - print(f"Expression syntax error: {e}") -except Exception as e: - print(f"Expression evaluation error: {e}") -``` +For optional imports or exception recovery, put `try/except` in a Python module +and call that helper from configuration. The [custom Python example](instantiation.md) +shows how to make such a module importable. ## Security Considerations -!!! warning "Expression Safety" - Expressions execute arbitrary Python code. Only load configurations from trusted sources. - - ```yaml - # Dangerous if config is from untrusted source! - dangerous: "$__import__('os').system('rm -rf /')" - ``` +Expressions and targets execute Python code. Use trusted configurations; lazy +branches, disabled components and `eval_expr=False` are not sandboxing mechanisms. +Configuration imports may execute even when expression evaluation is disabled. - Always validate configuration sources in production. +Migration for 0.1: expressions no longer eagerly construct every `@` dependency. +Move side effects that must always occur into explicitly requested nodes. Root +imports and composition-time external file loading retain their setup behavior. ## Next Steps -- [Instantiation](instantiation.md) - Create objects with expressions -- [Advanced Features](advanced.md) - Complex expression patterns +[Source versus execution](basics.md#accessing-configuration-values) · +[Troubleshooting](troubleshooting.md) · [Advanced ownership](advanced.md) diff --git a/docs/user-guide/instantiation.md b/docs/user-guide/instantiation.md index b3f17c6..ba388e2 100644 --- a/docs/user-guide/instantiation.md +++ b/docs/user-guide/instantiation.md @@ -1,229 +1,156 @@ # Instantiation -Create Python objects directly from configuration using the `_target_` key. +Keep application behavior in ordinary Python classes and functions. Sparkwheel +calls them with configured arguments; they do not need a Sparkwheel base class, +decorator or factory. ## Basic Instantiation -```yaml -model: - _target_: torch.nn.Linear - in_features: 784 - out_features: 10 -``` - -```python -from sparkwheel import Config +The [quick start](../getting-started/quickstart.md) calls `collections.Counter`. +This next complete example uses your own class. With the installed environment +active, create a separate `custom-report` directory and put all three files below +inside it. -config = Config() -config.update("config.yaml") +```python title="reports.py" +from collections import Counter -# Instantiate the object -model = config.resolve("model") -# model is now a torch.nn.Linear(784, 10) instance! -``` -## The `_target_` Key +class WordReport: + def __init__(self, words, title="Word counts"): + self.words = list(words) + self.title = title -`_target_` specifies the full Python path to a class or function: + def render(self): + counts = Counter(self.words) + entries = ", ".join(f"{word}={count}" for word, count in sorted(counts.items())) + return f"{self.title}: {entries}" +``` -```yaml -examples: - # Class instantiation - linear: - _target_: torch.nn.Linear - in_features: 100 - out_features: 10 - - # Class with multiple parameters - adam: - _target_: torch.optim.Adam - params: "$@model.parameters()" - lr: 0.001 - betas: [0.9, 0.999] - - # Custom class - custom: - _target_: myproject.models.CustomModel - hidden_size: 256 +```yaml title="config.yaml" +words: [red, blue, red] +report: + _target_: reports.WordReport + words: "@words" + title: Example ``` -## Positional Arguments with `_args_` +```python title="run.py" +from sparkwheel import Config +from reports import WordReport + +# Ordinary Python and configured construction call the same class. +plain = WordReport(["red", "blue", "red"], title="Example") +config = Config().update("config.yaml") +configured = config.resolve("report") +assert type(configured) is WordReport +assert configured.render() == plain.render() +print(configured.render()) +``` -Use `_args_` to pass positional arguments to classes or functions that require them: +Run from `custom-report`: -```yaml -# Basic example with list() -my_list: - _target_: builtins.list - _args_: - - [1, 2, 3, 4, 5] - # Equivalent to: list([1, 2, 3, 4, 5]) - -# torch.nn.Sequential requires positional args -model: - _target_: torch.nn.Sequential - _args_: - - _target_: torch.nn.Linear - in_features: 784 - out_features: 128 - - _target_: torch.nn.ReLU - - _target_: torch.nn.Linear - in_features: 128 - out_features: 10 - # Equivalent to: nn.Sequential(Linear(784, 128), ReLU(), Linear(128, 10)) +```bash +python run.py ``` -### Mixing `_args_` and Keyword Arguments +Expected output is `Example: blue=1, red=2`. Python can import `reports` because +`run.py` and `reports.py` are together. In a larger application, install your +package into the same environment and use its full import path. -You can combine positional arguments with keyword arguments: +## The `_target_` Key -```yaml -pipeline: - _target_: sklearn.pipeline.Pipeline - _args_: - - - scaler: - _target_: sklearn.preprocessing.StandardScaler - - model: - _target_: sklearn.linear_model.LogisticRegression - memory: null # Keyword argument - verbose: true # Keyword argument -``` +`_target_` is the full import path of a class or function. Other ordinary keys +become keyword arguments. A function target returns the function's result. +No automatic project discovery is performed. -!!! note "List Requirement" - `_args_` must always be a list, even if you're only passing a single positional argument. Each item in the list becomes a positional argument in order. +In Python-authored dictionaries, `_target_` may also be the callable itself. From +the custom example directory, this works without a string import path: -## Nested Instantiation +```python +from sparkwheel import Config +from reports import WordReport + +config = Config().update( + { + "report": {"_target_": WordReport, "words": ["blue"], "title": "Python"}, + } +) +assert config.resolve("report").render() == "Python: blue=1" +``` -Instantiate objects within objects: +Use serializable import paths when the source must be exported as standalone +YAML. Caller-owned Python objects keep normal identity and mutability; storing +one in a config does not turn it into a portable recipe. -```yaml -# Nested components -transform: - _target_: torchvision.transforms.Compose - transforms: - - _target_: torchvision.transforms.Resize - size: [224, 224] - - _target_: torchvision.transforms.ToTensor - - _target_: torchvision.transforms.Normalize - mean: [0.485, 0.456, 0.406] - std: [0.229, 0.224, 0.225] -``` +## Positional Arguments with `_args_` -## Complex Example +`_args_` must be a list. Each entry becomes one positional argument, and ordinary +keys supply keyword arguments. This definition calls `Counter(["red", "red"])`: ```yaml -# Complete training setup -dataset: - path: "/data/cifar10" - -transform: - _target_: torchvision.transforms.Compose - transforms: - - _target_: torchvision.transforms.ToTensor - - _target_: torchvision.transforms.Normalize - mean: [0.5, 0.5, 0.5] - std: [0.5, 0.5, 0.5] - -dataloader: - _target_: torch.utils.data.DataLoader - dataset: "@dataset" - batch_size: 32 - shuffle: true - -model: - _target_: torch.nn.Linear - in_features: 784 - out_features: 10 - -optimizer: - _target_: torch.optim.Adam - params: "$@model.parameters()" - lr: 0.001 +counts: + _target_: collections.Counter + _args_: [[red, red]] ``` +Nested `_target_` definitions construct nested objects. Use `@path` when a +constructor must receive an existing shared object; use `%path` for a separately +resolved copied definition. See [references](references.md). + ## Special Keys ### `_mode_` - Instantiation Modes -Control how objects are instantiated with `_mode_`: +| Mode | Result | +|---|---| +| `default` | Call the target with resolved arguments | +| `callable` | Return the target, or a `functools.partial` when arguments are supplied | +| `debug` | Call the target under `pdb` | -```yaml -# Default: instantiate normally -model: - _target_: torch.nn.Linear - _mode_: "default" # Optional, this is the default - in_features: 784 - out_features: 10 - # Returns: Linear(in_features=784, out_features=10) - -# Callable: return the class/function, not an instance -model_factory: - _target_: torch.nn.Linear - _mode_: "callable" - in_features: 784 - # Returns: functools.partial(torch.nn.Linear, in_features=784) - -# Debug: run in debugger -debug_component: - _target_: mymodule.MyClass - _mode_: "debug" - # Runs in pdb debugger -``` +`callable` is optional advanced control, useful when an API explicitly expects a +callable. It is not required for ordinary authoring. Arguments still resolve +before the callable/partial is returned. ### `_disabled_` - Skip Instantiation -Skip instantiation of a component without removing it from config: +`_disabled_: true` skips a target and its runtime argument dependencies. Inline +disabled components are removed from resolved parent lists/dicts. Resolving the +disabled path directly, or referring to it with `@`, returns `None`. The source +remains available for later re-enabling. -```yaml -callbacks: - - _target_: pytorch_lightning.callbacks.EarlyStopping - monitor: val_loss - patience: 3 - - _target_: pytorch_lightning.callbacks.ModelCheckpoint - _disabled_: true # This callback is removed from the list - save_top_k: 3 -``` +The guard accepts booleans, case-insensitive `"true"`/`"false"` strings, or a +resolved dynamic value. A dynamic guard runs its own dependencies first. A true +guard skips nested targets, argument expressions, unused `@` dependencies and +import statements inside its payload. Explicitly resolving a child path still +resolves that child's definition independently of the parent guard. -**Behavior:** +With `instantiate=False`, a disabled node returns a `Component` wrapper containing +its untouched payload. With `eval_expr=False`, an unresolved guard also returns a +wrapper instead of guessing its truth value. These modes have separate caches. -- When `_disabled_: true`, the component is skipped entirely -- **Inline in lists/dicts**: Disabled components are **removed** from the parent structure -- **Direct resolution**: `config.resolve("disabled_component")` returns `None` -- **References**: `@disabled_component` resolves to `None` -- Default is `false` (component is enabled) -- Accepts boolean values (`true`/`false`) or strings (`"true"`/`"false"`, case-insensitive) -- The config is preserved—you can re-enable by setting `_disabled_: false` +This is a runtime construction boundary: root `_imports_` declarations belong to +compilation, external `%file.yaml` includes load during `update()`, and ordinary +`Config` still expands/checks local copy syntax during compilation. A disabled +node does not bypass those stages. Retained scopes have stricter preflight rules, +described in [advanced contracts](advanced.md#retained-recipes-and-construction-scopes). -**Use cases:** +Migration: older versions could execute payload dependencies before checking the +guard. Put intentionally required setup outside disabled payloads. -```yaml -# Temporarily disable a feature for debugging -scheduler: - _target_: torch.optim.lr_scheduler.CosineAnnealingLR - _disabled_: true # Disable while debugging optimizer issues - optimizer: "@optimizer" - T_max: 100 - -# Environment-specific components -profiler: - _target_: pytorch_lightning.profilers.PyTorchProfiler - _disabled_: "$not os.environ.get('ENABLE_PROFILER')" # Expression support - -# A/B testing configurations -augmentation: - _target_: torchvision.transforms.RandomErasing - _disabled_: false # Toggle between experiments - p: 0.5 -``` +### Other Special Keys -!!! tip "Disabled vs Deleted" - Use `_disabled_` when you want to keep the config for future use. Use the delete operator (`~key: null`) when you want to permanently remove a key. +`_target_`, `_args_`, `_mode_` and `_disabled_` are construction metadata rather +than constructor keyword arguments. `_imports_` is a root configuration declaration +for expression imports; see [expressions](expressions.md#importing-modules). -### Other Special Keys +## When a Constructor Fails -- `_target_`: Class or function path to instantiate (required) -- `_args_`: List of positional arguments to pass -- `_disabled_`: Skip instantiation if `true` (removed from parent) -- `_mode_`: Instantiation mode (`"default"`, `"callable"`, or `"debug"`) +Check the native constructor signature and resolved arguments. Missing targets +raise `TargetNotFoundError`; exceptions raised while calling a target are wrapped +in `InstantiationError`. The original exception remains in the `__cause__` chain, +which may include another Sparkwheel wrapper. Follow the chain to the native +error or read the complete traceback; Sparkwheel does not repair application inputs. -For complete details, see the [Advanced Features](advanced.md) and [API Reference](../reference/). +[Validation](schema-validation.md) can check configured arguments before a target +is called. [Troubleshooting](troubleshooting.md) explains how to inspect source +without accidentally constructing the whole application. diff --git a/docs/user-guide/operators.md b/docs/user-guide/operators.md index 9b03f63..8381404 100644 --- a/docs/user-guide/operators.md +++ b/docs/user-guide/operators.md @@ -1,541 +1,127 @@ # Composition & Operators -Sparkwheel uses **composition-by-default**: configs merge naturally with just 2 operators (`=`, `~`) for explicit control. +Use composition to change an experiment or environment without copying its whole +configuration. Apply inputs in order with `config.update(...)`. ## Composition Decision Flow -!!! abstract "How Sparkwheel Merges Configs" +| Update | Existing value | Result | +|---|---|---| +| Ordinary key with a dictionary | Dictionary | Merge recursively | +| Ordinary key with a list | List | Append the new entries | +| Ordinary key with another value/type | Any | Replace with the new value | +| `=key` | Any | Replace the whole selected value | +| `~key: null` or CLI `~key` | Existing key/index | Delete it | +| `~key: [items]` | List or dictionary | Delete listed indices or keys | - When merging two configs, Sparkwheel follows this decision tree: - - **1. Key exists in both configs?** - - - ❌ **No** → Simply add the new key-value pair - - ✅ **Yes** → Continue to step 2 - - **2. Does the key have an operator?** - - - **`~key`** → 🗑️ Delete the key (highest priority) - - **`=key`** → 🔄 Replace completely (overwrite everything) - - **No operator** → Continue to step 3 - - **3. What's the value type?** (Default behavior) - - - **Dict** → ✅ **Merge recursively** (combine keys) - - **List** → ✅ **Extend** (append items) - - **Other** → Replace with new value - -!!! tip "Priority Order" - Delete (`~`) > Replace (`=`) > Type-based default (merge/extend) +Deleting a missing key is an error. In YAML, operators are prefixes on keys; +in Python dictionaries, quote the complete key string. Operators are not Python +assignment syntax. ## Composition by Default -By default, configs compose naturally - dicts merge, lists extend: - -**Dicts merge automatically:** - -```yaml -# base.yaml -model: - hidden_size: 512 - activation: "relu" - dropout: 0.1 -``` - -```yaml -# override.yaml -model: - hidden_size: 1024 # Update this field - # Other fields preserved! -``` +This standalone example shows dictionary merging and list extension: ```python -config = (Config() - .update("base.yaml") - .update("override.yaml")) -# Result: -# model: -# hidden_size: 1024 (updated) -# activation: "relu" (preserved) -# dropout: 0.1 (preserved) -``` - -**Lists extend automatically:** - -```yaml -# base.yaml -plugins: - - logger - - metrics - -# override.yaml -plugins: - - cache # Adds to the list! +from sparkwheel import Config -# Result: [logger, metrics, cache] +config = Config().update( + { + "report": {"title": "Words", "limit": 3}, + "words": ["red", "blue"], + } +) +config.update({"report": {"limit": 5}, "words": ["green"]}) +assert config.get("report") == {"title": "Words", "limit": 5} +assert config.get("words") == ["red", "blue", "green"] ``` -!!! success "Natural Composition" - **No operators needed for the common case!** Sparkwheel merges dicts and extends lists by default, matching how you naturally think about config layering. +Files behave the same way. The [quick start](../getting-started/quickstart.md#experiment-without-copying) +provides a complete file-overlay example. ## The `=` Operator: Explicit Replace -When you need to completely replace something, use `=key`: - -```yaml -# override.yaml -=model: # Replace the entire model dict - hidden_size: 1024 - # Old fields (activation, dropout) are GONE! -``` +Continuing the example above: ```python -config = (Config() - .update("base.yaml") - .update("override.yaml")) -# Result: -# model: -# hidden_size: 1024 (only this remains) -``` - -### When to Use `=` - -Use `=` when you want to: -- Replace an entire section with a fresh start -- Change the type of a value (e.g., dict → list) -- Clear out all previous settings - -```yaml -# Replace list entirely (no extension) -=plugins: [redis, cache] - -# Replace nested section -training: - =optimizer: # Replace optimizer, but merge training - type: "sgd" - lr: 0.1 +config.update({"=report": {"title": "New"}, "=words": ["blue"]}) +assert config.get("report") == {"title": "New"} +assert config.get("words") == ["blue"] ``` -!!! tip "Quoting in YAML Files" - When using `=` in YAML files, you can quote the key (`'=model'`) for clarity, but it's not required. In Python code, no quoting is needed. +Use replacement when switching a component's entire definition, clearing inherited +arguments, or replacing a list. Merely changing `_target_` in a merged mapping +leaves its other arguments in place; they may be invalid for the new target. -## The `~` Operator: Delete - -Remove keys or list items with `~key`: - -### Delete Entire Keys +The YAML form is: ```yaml -# Remove keys explicitly -~old_param: null -~debug_settings: null +=report: + title: New +=words: [blue] ``` -!!! warning "Key Must Exist" - The delete operator will raise an error if the key doesn't exist. This helps catch typos and configuration mistakes. - -### Delete Dict Keys +A leading `=` on a CLI override does the same thing: +`config.update("=words=[blue]")`. See [CLI syntax](cli.md#override-syntax). -Use path notation for nested keys: +## Editing Nested Paths -```yaml -# Path notation -~model::dropout: null -~training::old_params: null -``` - -Or structural notation: - -```yaml -# Structural notation (works without parent operator!) -model: - lr: 0.01 # Update - ~dropout: null # Delete - ~batch_norm: null # Delete -``` - -!!! success "No Parent Context Required!" - With composition-by-default, nested `~` just works - no special parent operator needed! - -### Delete from Lists - -Remove items by index (batch syntax): - -```yaml -# base.yaml -plugins: - - logger # 0 - - metrics # 1 - - cache # 2 - - auth # 3 - - debug # 4 - -# override.yaml - Remove by indices -~plugins: [0, 2, 4] # Remove indices 0, 2, 4 - -# Result: [metrics, auth] -``` - -**Negative indices work too:** - -```yaml -~plugins: [-1] # Remove last item -~plugins: [0, -1] # Remove first and last -``` - -### Delete from Dicts - -Remove nested dict keys by name: - -```yaml -# base.yaml -dataloaders: - train: {batch_size: 32} - val: {batch_size: 16} - test: {batch_size: 8} - -# override.yaml -~dataloaders: ["train", "test"] - -# Result: -# dataloaders: -# val: {batch_size: 16} -``` - -!!! warning "Removing List Items" - To remove items from a list, **you must use the batch syntax** `~key: [indices]`: - - ```yaml - # ✓ CORRECT - Batch deletion syntax - ~plugins: [0, 2, 4] - ``` - - ```yaml - # ✗ WRONG - Path notation doesn't work for list items! - ~plugins::0: null - ``` - - **Why?** Path notation is designed for dict keys, not list indices. The batch syntax handles index normalization and processes deletions correctly (high to low order). - -## Combining Operators - -Mix composition, replace, and delete: - -```yaml -# base.yaml -application: - name: "MyApp" - version: 1.0 - features: - auth: enabled - cache: enabled - debug: enabled - plugins: [logger, metrics] - database: - host: localhost - port: 5432 - pool_size: 10 - -# production.yaml -application: - version: 1.1 # Compose: update (default) - features: # Compose: merge (default) - cache: redis # Update - ~debug: null # Delete - plugins: [monitor] # Compose: extend (default!) - =database: # Replace: fresh db config - host: prod.example.com - port: 5432 - ssl: true - -# Result: -# application: -# name: "MyApp" (preserved) -# version: 1.1 (updated) -# features: -# auth: enabled (preserved) -# cache: redis (updated) -# # debug removed -# plugins: [logger, metrics, monitor] (extended!) -# database: (replaced entirely) -# host: prod.example.com -# port: 5432 -# ssl: true -``` - -## Programmatic Usage - -Apply operators in Python: +`::` paths traverse existing dictionaries and list indices. Use a top-level path +key to update a list element; a dictionary placed inside an incoming list is an +ordinary appended item, not an index-edit directive. ```python from sparkwheel import Config -config = Config() -config.update("base.yaml") - -# Compose (merge dict) - default behavior -config.update({"model": {"hidden_size": 1024}}) - -# Replace explicitly -config.update({"=optimizer": {"type": "sgd", "lr": 0.1}}) - -# Delete keys -config.update({ - "~training::old_param": None, - "~model::dropout": None -}) - -# Combine operations -config.update({ - "model": { # Merge - "hidden_size": 1024, # Update - "~dropout": None # Delete - }, - "=database": { # Replace - "host": "prod.example.com" - } -}) - -# Remove list items by index -config.update({"~plugins": [0, 2, 4]}) - -# Remove dict keys -config.update({"~dataloaders": ["train", "test"]}) +config = Config().update({"layers": [{"width": 8}, {"width": 16}]}) +config.update({"layers::0::width": 4}) +config.set("layers::1::width", 32) +assert config.get("layers") == [{"width": 4}, {"width": 32}] ``` -### Merging Config Instances +The corresponding YAML overlay is `layers::0::width: 4`. Missing dictionary paths +can be created; out-of-range list indices and traversal through scalars fail. +Extend a list through ordinary composition instead of sparse assignment. -Configs compose when merged: - -```python -base = Config() -base.update("base.yaml") -override = Config() -override.update("override.yaml") - -# Merge one Config into another (composes by default!) -base.update(override) -``` - -## Common Patterns - -### Environment-Specific Configs - -```yaml -# base.yaml -database: - host: "localhost" - port: 5432 - pool_size: 10 - ssl: false - -# production.yaml (merges automatically!) -database: - host: "prod-db.example.com" - ssl: true - pool_size: 50 - # Other settings inherited from base -``` - -### Experiment Variations +## The `~` Operator: Delete -```yaml -# base_model.yaml -model: - hidden_size: 512 - num_layers: 6 - dropout: 0.1 - -# experiment_large.yaml (merges automatically!) -model: - hidden_size: 1024 - num_layers: 12 - -# experiment_no_dropout.yaml (merges automatically, deletes dropout) -model: - ~dropout: null -``` +### Delete Entire Keys -### Feature Flags +Both `config.update({"~report": None})` and `config.update("~report")` remove +`report`. The key must exist. Nested deletion uses a path such as `~report::limit`, +or a structural update such as `{"report": {"~limit": None}}`. -```yaml -# base.yaml -plugins: - - logger - - metrics - - profiler - - debugger - - test_reporter - -# production.yaml - Remove debug/test plugins -~plugins: [2, 3, 4] # Remove profiler, debugger, test_reporter - -# Result: [logger, metrics] -``` +### Delete from Lists -### Layered Configuration +Single-index deletion is supported. Batch deletion is useful when several indices +refer to the same original list: ```python -# Build configs in layers (all compose naturally!) -config = (Config() - .update("defaults.yaml") - .update("models/resnet50.yaml") - .update("datasets/imagenet.yaml") - .update("experiments/exp_042.yaml") - .update("env/production.yaml")) -``` - -## Best Practices - -### Leverage Composition - -```yaml -# Good - natural composition (no operators!) -optimizer: - lr: 0.01 - -# Unnecessary - = not needed for simple updates -=optimizer: - lr: 0.01 -``` - -### Use `=` Only When Needed - -```yaml -# Use = when completely replacing -=optimizer: # Start fresh, discard all old settings - type: "sgd" - lr: 0.1 - -# Default composition is usually what you want -optimizer: # Keep other settings, update lr - lr: 0.01 -``` - -### Choose Path vs Structural Notation - -**Use path notation** for single, independent operations: - -```yaml -# Quick single updates/deletes -~model::dropout: null -~training::old_param: null -``` - -**Use structural notation** for bulk related operations: - -```yaml -# Multiple related changes -model: - hidden_size: 1024 - num_layers: 12 - ~dropout: null - ~batch_norm: null -``` - -### Write Reusable Configs - -!!! warning "Delete Requires Key Existence" - The delete operator (`~`) is **strict** - it raises an error if the key doesn't exist. This helps catch typos and configuration mistakes. - -When writing configs that should work with different base configurations, you have a few options: - -**Option 1: Document required keys** -```yaml -# production.yaml -# Requires: base config must have debug_mode and verbose_logging -~debug_mode: null -~verbose_logging: null -database: - pool_size: 100 - ssl: true -``` - -**Option 2: Use composition order** -```yaml -# production.yaml - override instead of delete -debug_mode: false # Overrides if exists, sets if not -verbose_logging: false -database: - pool_size: 100 - ssl: true -``` - -**Option 3: Conditional deletion with lists** -```yaml -# Delete multiple optional keys - fails only if ALL are missing -~: [debug_mode, verbose_logging] # At least one must exist -database: - pool_size: 100 -``` - -## Common Mistakes - -### Using `=` When Not Needed - -```yaml -# Unnecessary - composition merges by default! -=model: - hidden_size: 1024 - -# Better - let it compose naturally -model: - hidden_size: 1024 -``` - -### Expecting List Replacement by Default - -```yaml -# This EXTENDS the list (doesn't replace) -plugins: [cache] - -# Use = to replace -=plugins: [cache] -``` - -### Wrong List Deletion Syntax - -```yaml -# Wrong - path notation doesn't work for list indices -~plugins::0: null - -# Correct - use batch syntax -~plugins: [0] -``` - -### Forgetting Quotes for Operators - -```yaml -# Wrong - YAML might misinterpret -=model: - lr: 0.001 +from sparkwheel import Config -# Safer - quote operators (optional but clearer) -'=model': - lr: 0.001 +config = Config().update({"words": ["red", "blue", "green", "yellow"]}) +config.update("~words::1") +assert config.get("words") == ["red", "green", "yellow"] +config.update({"~words": [0, 2]}) +assert config.get("words") == ["green"] ``` -## Comparison with Other Systems +Negative indices are supported. A batch such as `{"~words": [0, -1]}` removes the +first and last entries of the list as it existed before that batch, without index +shift changing the meaning. Sequential calls each operate on the current list. -### vs Hydra +### Delete from Dicts -| Feature | Hydra | Sparkwheel | -|---------|-------|------------| -| Dict merge default | Yes ✅ | Yes ✅ | -| List extend default | No ❌ | **Yes** ✅ | -| Operators in YAML | CLI-only | **Yes** ✅ (YAML + CLI) | -| Operator count | 4 (`=`, `+`, `++`, `~`) | **2** (`=`, `~`) ✅ | -| Delete dict keys | CLI-only (`~foo.bar`) | **Yes** ✅ (YAML + CLI) | -| Delete list items | No ❌ | **Yes** ✅ (by index) | +A batch on a dictionary names keys rather than indices. For example, +`{"~report": ["title", "limit"]}` deletes those two keys inside `report` while +keeping the report mapping. `{"~report": None}` deletes the entire mapping. -Sparkwheel differs from Hydra: -- **Full composition philosophy**: Both dicts AND lists compose by default -- **Operators in YAML files**: Not just CLI overrides -- **Simpler operator set**: Just 2 operators (`=`, `~`) vs 4 (`=`, `+`, `++`, `~`) -- **List deletion**: Delete items by index with `~plugins: [0, 2]` -- **Flexible delete**: Use `~` anywhere (YAML, CLI, programmatic) +## Atomic Edits -## Next Steps +An `update()` call is one source transaction. Composition, external loading or +schema failure preserves the previous source and runtime cache. Use one update +for related fields that must validate together. Successful changes start a new +resolution generation; already-returned objects remain unchanged. -- **[Configuration Basics](basics.md)** - Core config management -- **[Advanced Features](advanced.md)** - Macros and power features +Read [the configuration model](basics.md#editing-configuration) for mutation and +lifetime rules, and [troubleshooting](troubleshooting.md) for unexpected merges. diff --git a/docs/user-guide/quick-reference.md b/docs/user-guide/quick-reference.md index 774bd6b..b2d6263 100644 --- a/docs/user-guide/quick-reference.md +++ b/docs/user-guide/quick-reference.md @@ -1,297 +1,105 @@ # Quick Reference -A one-page cheat sheet for Sparkwheel syntax and features. +A lookup page for syntax and entrypoints. Start with the +[quick start](../getting-started/quickstart.md) for a complete program. ## Core Syntax -=== "References" - - | Syntax | Type | Returns | Example | - |--------|------|---------|---------| - | `@key` | Resolved reference | Final computed value | `lr: "@defaults::learning_rate"` | - | `%key` | Raw reference | Unprocessed YAML | `config: "%base.yaml::model"` | - | `@key::nested` | Nested access | Nested value | `@dataset::train::batch_size` | - | `@list::0` | List indexing | List element | `@transforms::0` | - -=== "Expressions" - - | Expression | Description | Example | - |------------|-------------|---------| - | `$(@a + @b)` | Math operations | `$(@lr * 0.1)` | - | `$(@name + "_v2")` | String concatenation | `$(@model_name + "_trained")` | - | `$(@debug ? "dev" : "prod")` | Ternary conditional | `$(@is_training ? 0.5 : 0.0)` | - | `$(@items[0])` | Dynamic indexing | `$(@datasets[@mode])` | - | `$len(@items)` | Built-in functions | `$len(@layers)` | - -=== "Operators" - - | Operator | Purpose | Example | Result | - |----------|---------|---------|--------| - | `=key` | Replace (don't merge) | `=optimizer: sgd` | Replaces entire dict/value | - | `~key` | Delete key | `~debug: true` | Removes the key | - | `~list: [0, 2]` | Delete list items | `~layers: [1, 3]` | Removes items at indices 1 and 3 | - | (none) | Default: Merge | `model: {size: 512}` | Merges with existing dict | - -=== "Instantiation" - - | Key | Type | Purpose | Example | - |-----|------|---------|---------| - | `_target_` | str | Class/function to instantiate | `torch.optim.Adam` | - | `_args_` | list | Positional arguments | `[arg1, arg2]` | - | `_disabled_` | bool | Skip instantiation | `true` | - | `_mode_` | str | Instantiation mode | `"default"` / `"callable"` / `"debug"` | - -## Common Patterns - -### Single Source of Truth - -```yaml -defaults: - learning_rate: 0.001 - -optimizer: - lr: "@defaults::learning_rate" - -scheduler: - base_lr: "@defaults::learning_rate" -``` - -### Computed Values - -```yaml -dataset: - samples: 10000 - batch_size: 32 - -training: - steps_per_epoch: "$@dataset::samples // @dataset::batch_size" -``` - -### Conditional Configuration - -```yaml -environment: "production" - -database: - prod_host: "prod.db.com" - dev_host: "localhost" - host: "$@database::prod_host if @environment == 'production' else @database::dev_host" -``` - -### Object Instantiation - -```yaml -# Basic instantiation -model: - _target_: torch.nn.Linear - in_features: 784 - out_features: 10 - -# With positional arguments -sequential: - _target_: torch.nn.Sequential - _args_: - - _target_: torch.nn.Linear - in_features: 784 - out_features: 128 - - _target_: torch.nn.ReLU - -# With optimizer -optimizer: - _target_: torch.optim.Adam - params: "$@model.parameters()" - lr: 0.001 -``` - -### Config Composition - -```yaml title="base.yaml" -model: - hidden_size: 512 - dropout: 0.1 -``` - -```yaml title="override.yaml" -# Merge by default -model: - hidden_size: 1024 # Updates only this field - -# Or replace completely -=model: - hidden_size: 1024 # Replaces entire dict -``` - -## Type Coercion - -| From | To | Supported | Notes | -|------|----|-----------| ------| -| str → int | ✅ | `"42"` → `42` | -| str → float | ✅ | `"3.14"` → `3.14` | -| str → bool | ✅ | `"true"` → `True` | Accepts: true/false, yes/no, 1/0 | -| int → float | ✅ | `42` → `42.0` | -| float → int | ✅ | `3.14` → `3` | Truncates decimal | -| Any → str | ✅ | Universal | +| Syntax | Meaning | +|---|---| +| `@path` | Share the resolved value at a path within a mode/generation | +| `%path` | Copy a local definition, then resolve the copy normally | +| `%file.yaml::path` | Load a definition from an external file during composition | +| `$expression` | Evaluate a Python expression | +| `parent::child`, `items::0` | Dictionary path, list index | +| `_target_: module.Class` | Call an importable class or function | +| `_args_: [a, b]` | Positional arguments; ordinary keys become keyword arguments | +| `_disabled_: true` | Skip runtime target construction; preserve its source | +| `_mode_: callable` | Return the callable or a partial with resolved arguments | +| `_imports_: {alias: module}` | Root expression-import declaration | + +`@` shares a result; `%` copies a recipe, which may still contain shared `@` +dependencies. See [references](references.md) for identity examples. + +## Composition + +| Input | Behavior | +|---|---| +| Dictionary + dictionary | Recursive merge | +| List + list | Append incoming entries | +| `=key: value` | Replace the whole selected value | +| `~key: null` | Delete a key or indexed path | +| `~list: [0, 2]` | Delete indices from the original list for that batch | +| `~mapping: [a, b]` | Delete keys inside a mapping | +| `layers::0::width: 4` | Update an existing list element through a top-level path key | + +[Composition rules](operators.md) include missing paths, batch deletion and atomic +failure behavior. Duplicate authored YAML keys are rejected by default. ## CLI Overrides -=== "Direct Assignment" - - ```bash - python train.py learning_rate=0.01 batch_size=64 - ``` - -=== "Nested Keys" - - ```bash - python train.py model.hidden_size=1024 optimizer.lr=0.001 - ``` - -=== "List Values" - - ```bash - python train.py layers=[128,256,512] - ``` - -=== "Replace vs Merge" - - ```bash - # Merge (default) - python train.py model.dropout=0.2 - - # Replace entire section - python train.py =model={hidden_size:1024} - ``` - -## Schema Validation - -```python -from dataclasses import dataclass -from sparkwheel import Config, validator - -@dataclass -class AppConfigSchema: - name: str - port: int - debug: bool = False - - @validator - def check_port(self): - if not (1024 <= self.port <= 65535): - raise ValueError(f"Invalid port: {self.port}") - -# Continuous validation (validates on every mutation) -config = Config(schema=AppConfigSchema) - -# Or explicit validation -config = Config() -config.update("config.yaml") -config.validate(AppConfigSchema) -``` +`config.update(input_string)` detects files and overrides. Values use YAML types. +Quote the entire override for the shell, and retain inner YAML quotes when a value +must stay text. + +| Argument | Effect | +|---|---| +| `'limit=3'` | Integer assignment | +| `'enabled=true'` | Boolean assignment | +| `'label="true"'` | String assignment | +| `'=words=[blue]'` | List replacement | +| `'~words::0'` | Single-index deletion | + +Read the [CLI guide](cli.md) for a complete integration and type boundaries. + +## Python API + +These entrypoints cover normal application code: + +| API | Purpose | +|---|---| +| `Config().update(source)` | Compose a file, mapping, Config or override; return the Config | +| `config.get(path)` / `config[path]` | Read authored source without configured execution | +| `config.set(path, value)` | Transactionally set a value; returns `None` | +| `config.resolve(path)` | Resolve that node and dependencies; may execute Python | +| `config.resolve()` | Resolve the complete tree | +| `Config(schema=Schema, ...)` | Attach dataclass validation and optional coercion policy | +| `config.validate(Schema)` | Validate source explicitly without coercion | +| `config.freeze()` / `unfreeze()` / `is_frozen()` | Control supported source edits | +| `Config.export_config_file(config.get(), path)` | Write serializable source to YAML | +| `parse_overrides(arguments)` | Parse override strings separately from application | +| `validate(mapping, Schema)` / `validator` / `MISSING` | Standalone schema tools | +| `config.locations.get(path)` | Inspect recorded source location when available | + +For integrations, `config.retain()` produces a retained definition with +`definition()` and `scope()`; see [advanced contracts](advanced.md). +The [generated API reference](../reference/index.md) also exposes low-level +wrappers, resolver machinery and private helpers. Their visibility is useful for +reading the implementation; they are not required authoring interfaces or a +blanket promise of stable extension contracts. Prefer the entrypoints above. ## Resolution Order -References are resolved in dependency order: - -```yaml -a: 10 -b: "@a" # Resolved first (depends on a) -c: "$@a + @b" # Resolved after a and b -d: "$@c * 2" # Resolved last (depends on c) -``` - -!!! danger "Avoid Circular References" - ```yaml - # ❌ This will fail! - a: "@b" - b: "@a" - ``` - -## Best Practices - -!!! success "Do This" - - ✅ Use `@` references for DRY config - - ✅ Enable schema validation for type safety - - ✅ Leverage composition-by-default (no operators needed) - - ✅ Use expressions for computed values - - ✅ Keep configs simple and readable - -!!! warning "Avoid This" - - ❌ Don't create circular references - - ❌ Don't overuse expressions (hurts readability) - - ❌ Don't use operators when default composition works - - ❌ Don't put complex logic in configs - -## Common Gotchas - -| Issue | Problem | Solution | -|-------|---------|----------| -| **Reference not found** | `@key` doesn't exist | Check spelling and nesting | -| **Circular reference** | `a: "@b"`, `b: "@a"` | Restructure to break cycle | -| **Type mismatch** | Schema expects `int`, got `str` | Enable coercion or fix type | -| **Expression error** | Invalid Python in `$()` | Check syntax and references | -| **Unexpected merge** | Dict merged when you wanted replace | Use `=key` to replace | -| **List not extending** | List replaced instead of extended | This is default for scalars, expected | - -## File Organization - -### Small Projects - -``` -project/ -├── config.yaml # Single config file -└── train.py -``` - -### Medium Projects - -``` -project/ -├── configs/ -│ ├── defaults.yaml # Shared defaults -│ ├── dev.yaml # Development -│ └── prod.yaml # Production -└── train.py -``` - -### Large Projects - -``` -project/ -├── configs/ -│ ├── base/ -│ │ ├── model.yaml -│ │ ├── dataset.yaml -│ │ └── training.yaml -│ ├── experiments/ -│ │ ├── baseline.yaml -│ │ └── improved.yaml -│ └── env/ -│ ├── dev.yaml -│ ├── staging.yaml -│ └── prod.yaml -└── train.py -``` +References resolve on access. Python expressions access the branches Python +selects in ordinary `Config` resolution. Retained scopes additionally preflight +explicit graph references. Circular dependencies raise errors; see +[expressions](expressions.md) for lazy-branch and import details. ## Performance Tips -!!! info "Expression Evaluation" - Expressions are evaluated at **access time**. For frequently accessed values: - - ```python - # Slow: Re-evaluates expression each time - for i in range(1000): - x = config.resolve("computed_value") +Repeated resolution in one `(instantiate, eval_expr)` mode reuses cached values. +Successful supported source edits or `resolve(lazy=False)` reset all modes. +Already-returned objects, lambdas and generators retain their original generation. +Use `get()` for source inspection; disabled processing modes are not a static or +security boundary. - # Fast: Evaluate once - computed = config.resolve("computed_value") - for i in range(1000): - x = computed - ``` +## Common Gotchas -## Next Steps +- Use `set()`/`update()` for edits; raw nested writes bypass cache invalidation. +- Use `=` when replacing a list or whole component recipe. +- A copied `_target_` definition still instantiates when resolved. +- Schema checks do not prove application behavior; partial resolution is partial validation. +- Caller-owned Python objects are not deeply frozen or made serializable. -- **[Configuration Basics](basics.md)** - Learn config fundamentals -- **[References](references.md)** - Deep dive into `@` and `%` -- **[Expressions](expressions.md)** - Master `$()` expressions -- **[Operators](operators.md)** - Composition with `=` and `~` -- **[Schema Validation](schema-validation.md)** - Type-safe configs +The [troubleshooting guide](troubleshooting.md) maps symptoms to the next useful +check. [Advanced contracts](advanced.md) retain the precise integration boundaries. diff --git a/docs/user-guide/references.md b/docs/user-guide/references.md index 0c46aa6..cfcda62 100644 --- a/docs/user-guide/references.md +++ b/docs/user-guide/references.md @@ -1,339 +1,127 @@ # References -Sparkwheel provides two types of references for linking configuration values: - -- **`@` - Resolved References**: Get the final, instantiated/evaluated value -- **`%` - Raw References**: Get the unprocessed YAML content +Use `@` when consumers should receive the same resolved value. Use `%` when you +want to reuse a definition under another path and resolve it separately. ## Quick Comparison -| Feature | `@ref` (Resolved) | `%ref` (Raw) | `$expr` (Expression) | -|---------|-------------------|--------------|----------------------| -| **Returns** | Final computed value | Raw YAML content | Evaluated expression result | -| **When processed** | Lazy (`resolve()`) | External: Eager / Local: Lazy | Lazy (`resolve()`) | -| **Instantiates objects** | ✅ Yes | ❌ No | ✅ Yes (if referenced) | -| **Evaluates expressions** | ✅ Yes | ❌ No | ✅ Yes | -| **Use in dataclass validation** | ✅ Yes | ⚠️ Limited | ✅ Yes | -| **CLI override compatible** | ✅ Yes | ✅ Yes | ❌ No | -| **Cross-file references** | ✅ Yes | ✅ Yes | ❌ No | -| **When to use** | Get computed results | Copy config structures | Compute new values | - -## Two-Phase Processing Model - -Sparkwheel processes raw references (`%`) in two phases to support CLI overrides: - -!!! abstract "When References Are Processed" - - **Phase 1: Eager Processing (during `update()`)** - - - **External file raw refs (`%file.yaml::key`)** are expanded immediately - - External files are frozen—their content won't change based on CLI overrides - - Enables copy-then-delete workflows with external files - - **Phase 2: Lazy Processing (during `resolve()`)** - - - **Local raw refs (`%key`)** are expanded after all composition is complete - - **Resolved References (`@`)** are processed on-demand - - **Expressions (`$`)** are evaluated when needed - - **Components (`_target_`)** are instantiated only when requested - - CLI overrides can affect local `%` refs - -**Why two phases?** - -This design ensures CLI overrides work intuitively with local raw references: - -```yaml -# base.yaml -vars: - features_path: null # Default, will be overridden - -# model.yaml -dataset: - path: "%vars::features_path" # Local ref - sees CLI override -``` - -```python -config = Config() -config.update("base.yaml") -config.update("model.yaml") -config.update("vars::features_path=/data/features.npz") # CLI override - -# Local % ref sees the override! -path = config.resolve("dataset::path") # "/data/features.npz" -``` - -!!! tip "External vs Local Raw References" - - | Type | Example | When Expanded | Use Case | - |------|---------|---------------|----------| - | **External** | `%file.yaml::key` | Eager (update) | Import from frozen files | - | **Local** | `%vars::key` | Lazy (resolve) | Reference config values | - - External files are "frozen"—their content is fixed at load time. - Local config values may be overridden via CLI, so local refs see the final state. - -## Resolution Flow +| Syntax | What happens | Typical use | +|---|---|---| +| `@counts` | Resolve `counts` and share its value in this mode/generation | One dependency used by several consumers | +| `%counts` | Copy the definition, then resolve the copy normally | Another independently constructed component | +| `%base.yaml::counts` | Load a definition from a file during composition | Reusable configuration fragment | +| `$len(@words)` | Evaluate a Python expression using a resolved dependency | A small derived value | -!!! abstract "How References Are Resolved" - - **Step 1: Load Configs** → During `update()` - - - Parse YAML files - - Expand external `%file.yaml::key` refs immediately - - Keep local `%key` refs as strings - - **Step 2: Apply Overrides** → During `update()` calls - - - CLI overrides modify local config values - - Local `%` refs still see the string form - - **Step 3: Resolve** → During `resolve()` - - - Expand local `%key` refs (now sees final values) - - Resolve `@` dependencies in order - - Check for circular references → ❌ **Error if found** - - Evaluate expressions and instantiate objects - - Return final computed value ✅ +A `%` copy of a `_target_` **does construct an object when resolved**. It is not a +way to request an unevaluated dictionary. Use `get()` to inspect source; for a +retained, selectively expanded definition, see [advanced inspection](advanced.md#retained-recipes-and-construction-scopes). ## Resolved References (`@`) -Use `@` followed by the key path with `::` separator to reference **resolved values** (after instantiation, expression evaluation, etc.): - -```yaml title="config.yaml" hl_lines="7 10" -dataset: - path: "/data/images" - num_classes: 10 - batch_size: 32 - -model: - num_outputs: "@dataset::num_classes" # (1)! - -training: - batch: "@dataset::batch_size" # (2)! -``` - -1. References the resolved value of `dataset.num_classes` (10) -2. Uses `::` separator for nested key access +This standalone example distinguishes shared identity from a copied definition: -```python title="main.py" -config = Config() -config.update("config.yaml") - -# References are resolved when you call resolve() -num_outputs = config.resolve("model::num_outputs") # 10 -batch = config.resolve("training::batch") # 32 +```python +from sparkwheel import Config + +config = Config().update( + { + "words": ["red", "blue", "red"], + "counts": {"_target_": "collections.Counter", "_args_": ["@words"]}, + "shared": "@counts", + "independent": "%counts", + } +) +counts = config.resolve("counts") +shared = config.resolve("shared") +independent = config.resolve("independent") +assert shared is counts +assert independent is not counts +assert independent == counts +counts["red"] += 1 +assert shared["red"] == 3 +assert independent["red"] == 2 +assert config.get("independent") == "%counts" ``` -!!! tip "Single Source of Truth" - References prevent copy-paste errors by maintaining a single source of truth for shared values across your configuration. +Sharing applies to objects, containers and computed results. It is scoped to the +current resolution mode and generation, not global across all `Config` instances. +Read [object lifetimes](basics.md#object-lifetimes) before rebuilding dependencies. ## List References -Reference list elements by index (0-based): - -```yaml -transforms: - - resize - - normalize - - augment - -first_transform: "@transforms::0" # "resize" -last_transform: "@transforms::2" # "augment" -``` +Use `::` to address nested dictionary keys or a list index. For the quick-start +configuration, `@words::0` resolves to `"red"`. References are whole strings; +quote them in YAML. ## Nested References -References can reference other references: - -```yaml -base: - value: 100 - -derived: - double: "$@base::value * 2" # 200 - -final: - quad: "$@derived::double * 2" # 400 -``` - -## Resolution Order - -Sparkwheel resolves references in dependency order: - -```yaml -a: 10 -b: "@a" # Resolved first -c: "$@a + @b" # Resolved after a and b -d: "$@c * 2" # Resolved last -``` +A reference may point to another reference. Expression dependencies are resolved +when Python uses them: `$@a if True else @bad` does not evaluate `bad` in ordinary +`Config` resolution. For dynamic selection use normal indexing, such as +`$@datasets[@mode]`. ### Circular References -!!! danger "Avoid Circular References" - Circular references will cause a resolution error and must be avoided: - - ```yaml - # ❌ This will fail! - a: "@b" - b: "@a" - ``` - - Sparkwheel detects circular dependencies during resolution and raises a descriptive error to help you identify the cycle. - -## Advanced Patterns - -### Conditional References - -```yaml -environment: "production" - -database: - prod_host: "prod.db.example.com" - dev_host: "localhost" - host: "$@database::prod_host if @environment == 'production' else @database::dev_host" -``` - -### Dynamic Selection - -```yaml -datasets: - train: "/data/train" - test: "/data/test" - val: "/data/val" - -mode: "train" -current_dataset: "$@datasets[@mode]" # Dynamically select based on mode -``` - -**Note:** This requires Python expression evaluation. +A requested cycle such as `a: "@b"` and `b: "@a"` raises a resolution error. +Construction scopes apply additional conservative graph checks; see +[expressions](expressions.md#combining-with-references). ## Raw References (`%`) -Use `%` to reference **raw YAML content** (unprocessed, before instantiation/evaluation). Works with both external files and within the same file: - -### External File Raw References - -```yaml -# base.yaml -defaults: - learning_rate: 0.001 - batch_size: 32 - -model: - _target_: torch.nn.Linear - in_features: 784 - out_features: 10 - -# experiment.yaml -training: - lr: "%base.yaml::defaults::learning_rate" # Gets raw value: 0.001 - batch: "%base.yaml::defaults::batch_size" # Gets raw value: 32 - -# Gets the raw dict definition (with _target_), NOT the instantiated object -model_template: "%base.yaml::model" -``` - ### Local Raw References -Local raw references are expanded lazily during `resolve()`, which means CLI overrides can affect them: +A local `%counts` stays in the authored source until a resolution tree is compiled. +Its copied definition therefore sees source overrides applied before that +compilation. The copy gets its own path and normal resolution cache entry. -```yaml -# config.yaml -defaults: - timeout: 30 - retries: 3 +Copying a definition does not guarantee independence of every dependency. For +example, the copied counter above still contains `@words`, so both constructors +receive the shared resolved input. Counter creates its own count state from that +input. A custom class that stores a shared mutable dependency would retain it. -# Copy raw YAML from same file -api_config: - timeout: "%defaults::timeout" # Gets raw value: 30 - -# Copy entire section -backup_defaults: "%defaults" # Gets the whole defaults dict -``` - -!!! tip "CLI Overrides Work with Local Raw Refs" - - ```python - config = Config() - config.update("config.yaml") - config.update("defaults::timeout=60") # CLI override - - # Local % ref sees the override! - config.resolve("api_config::timeout") # 60 - ``` - -### Key Distinction - -!!! abstract "@ vs % - When to Use Each" - - | Reference Type | Symbol | What You Get | When To Use | - |----------------|--------|--------------|-------------| - | **Resolved Reference** | `@` | Final value after instantiation/evaluation | When you want the computed result or object instance | - | **Raw Reference** | `%` | Unprocessed YAML content | When you want to copy/reuse configuration definitions | - -**Example showing the difference:** - -```yaml title="config.yaml" hl_lines="8 11" -model: - _target_: torch.nn.Linear - in_features: 784 - out_features: 10 +### External File Raw References -# Resolved reference - gets the actual instantiated torch.nn.Linear object -trained_model: "@model" # (1)! +Create these two files in one directory and run the Python block from that directory: -# Raw reference - gets the raw dict with _target_, in_features, out_features -model_config_copy: "%model" # (2)! +```yaml title="base.yaml" +counts: + _target_: collections.Counter + _args_: [[red, blue, red]] ``` -1. ✅ Returns an actual `torch.nn.Linear` instance -2. ✅ Returns a dictionary: `{"_target_": "torch.nn.Linear", "in_features": 784, "out_features": 10}` - -See [Advanced Features](advanced.md) for more on raw references. - -## Common Use Cases - -### Shared Hyperparameters - -```yaml -# Single source of truth -model_config: - hidden_size: 512 +```yaml title="experiment.yaml" +counts: "%base.yaml::counts" +``` -encoder: - size: "@model_config::hidden_size" +```python +from sparkwheel import Config -decoder: - size: "@model_config::hidden_size" +config = Config().update("experiment.yaml") +assert config.get("counts")["_target_"] == "collections.Counter" +assert dict(config.resolve("counts")) == {"red": 2, "blue": 1} ``` -### Computed Values +External `%file.yaml::path` includes are expanded during `update()`. Their content +is a loaded snapshot; later source-file edits are not watched. Local `%path` +copies expand when resolution compiles the composed source. Nested local `%` copies +inside an external include are expanded against the included file during loading. +Remaining `@` references and expressions resolve later in the composed configuration. -```yaml -dataset: - samples: 10000 - batch_size: 32 +### Key Distinction -training: - steps: "$@dataset::samples // @dataset::batch_size" # 312 -``` +`@` shares a **result**. `%` copies a **definition**. Whether the resulting objects +share their internals follows the copied definition and ordinary Python behavior. -### Object Parameters +## Two-Phase Processing Model -```yaml -model: - _target_: torch.nn.Linear - in_features: 784 - out_features: 10 +1. `update()` loads and composes source, expanding external `%file.yaml` includes. +2. Resolution compiles a separate tree, expands local `%path` copies, and resolves + requested references, expressions and targets. The authored source is retained. -optimizer: - _target_: torch.optim.Adam - params: "$@model.parameters()" # Call model's method - lr: 0.001 -``` +External file access is a loading effect. Configured imports and target execution +are separate runtime effects; neither raw copies nor lazy resolution form a sandbox. ## Next Steps -- **[Expressions](expressions.md)** - Execute Python code in configs -- **[Instantiation](instantiation.md)** - Create objects with references -- **[Advanced Features](advanced.md)** - Complex reference patterns +[Python authoring](instantiation.md) · [Expressions](expressions.md) · +[Advanced ownership contracts](advanced.md) diff --git a/docs/user-guide/schema-validation.md b/docs/user-guide/schema-validation.md index b63d6a2..3451dfc 100644 --- a/docs/user-guide/schema-validation.md +++ b/docs/user-guide/schema-validation.md @@ -1,617 +1,241 @@ # Schema Validation -Validate configurations at runtime using Python dataclasses with **continuous validation** - errors caught immediately when you mutate the config. - -## Type Coercion Matrix - -Sparkwheel automatically converts compatible types when coercion is enabled (default: `True`): - -| From ↓ To → | `int` | `float` | `str` | `bool` | `list` | `dict` | -|-------------|-------|---------|-------|--------|--------|--------| -| **int** | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ | -| **float** | ✅* | ✅ | ✅ | ❌ | ❌ | ❌ | -| **str** | ✅** | ✅** | ✅ | ✅*** | ❌ | ❌ | -| **bool** | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ | -| **list** | ❌ | ❌ | ✅ | ❌ | ✅ | ❌ | -| **dict** | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ | - -\* Truncates decimal part (e.g., `3.14` → `3`) -\*\* Requires valid format (e.g., `"42"` for int, `"3.14"` for float) -\*\*\* Accepts: `"true"`, `"false"`, `"1"`, `"0"`, `"yes"`, `"no"` (case-insensitive) - -!!! success "Default Behavior" - Type coercion is **enabled by default** to handle common cases like environment variables and CLI arguments (which are always strings). - -!!! warning "Disable for Strict Validation" - Set `coerce=False` for strict type checking: - ```python - config = Config(schema=AppConfigSchema, coerce=False) - ``` +Use a dataclass schema when configuration values need an explicit contract. +Concrete source values are checked when edited; deferred references and expressions +are checked when resolved. A schema checks arguments and values, not whether the +application achieves its intended result. ## Quick Start -Define a schema with dataclasses: +This standalone example checks and coerces a numeric option, then preserves the +previous configuration when an edit fails: -```python title="app.py" hl_lines="10 14 15" +```python from dataclasses import dataclass -from sparkwheel import Config +from sparkwheel import Config, ValidationError -@dataclass -class AppConfigSchema: - name: str - port: int - debug: bool = False - -# Continuous validation - validates on every update/set! -config = Config(schema=AppConfigSchema) # (1)! -config.update("config.yaml") - -# Errors caught immediately at mutation time -config.set("port", "8080") # (2)! -config.set("port", "not a number") # (3)! - -# Or validate explicitly after loading -config = Config() -config.update("config.yaml") -config.validate(AppConfigSchema) # (4)! -``` -1. ✅ Enable continuous validation - errors caught on every mutation -2. ✅ Auto-coerced to `int(8080)` (coercion enabled by default) -3. ❌ Raises `ValidationError` immediately - invalid type conversion -4. ✅ Alternative: validate explicitly after loading all config +@dataclass +class ReportOptions: + title: str + limit: int -With **type coercion** enabled by default, compatible types are automatically converted: -```python -# config.yaml: -# name: "myapp" -# port: "8080" # String value -# debug: "true" # String value - -config = Config(schema=AppConfigSchema, coerce=True) -config.update("config.yaml") -# ✓ port coerced to int(8080) -# ✓ debug coerced to bool(True) +config = Config(schema=ReportOptions).update({"title": "Words", "limit": "3"}) +assert config.get("limit") == 3 +try: + config.set("limit", "not an integer") +except ValidationError: + pass +else: + raise AssertionError("Expected validation to reject the edit") +assert config.get("limit") == 3 ``` -If validation fails, you get clear errors: - -```python -# With coercion disabled -config = Config(schema=AppConfigSchema, coerce=False) -config.update({"port": "8080"}) -# ValidationError: Validation error at 'port': Type mismatch -# Expected type: int -# Actual type: str -# Actual value: '8080' -``` +Supply related required fields together in one update. Each update validates its +complete staged source; successive files are not one combined transaction. ## Defining Schemas -Schemas are Python dataclasses with type hints. +Schemas are Python dataclasses with type hints. Nested dataclasses describe nested +mappings; annotated lists and dictionaries describe their contents. Postponed +annotations resolve in the defining module. Unresolved annotation names produce +an error rather than silently disabling checks. -### Basic Types - -```python -@dataclass -class ConfigSchema: - text: str - count: int - ratio: float - enabled: bool - items: list[str] - mapping: dict[str, int] -``` +Dataclass defaults make fields optional for validation; they are **not inserted** +into authored source or constructor arguments. Type-only validation does not call +default factories. A target's Python signature determines runtime defaults. +Custom validators instantiate the schema, so its defaults, `default_factory` and +`__post_init__` may execute. ### Optional Fields -```python -from typing import Optional - -@dataclass -class ConfigSchema: - required: str - optional_with_none: Optional[int] = None - optional_with_default: int = 42 -``` - -### Nested Dataclasses - -```python -@dataclass -class DatabaseConfigSchema: - host: str - port: int - pool_size: int = 10 - -@dataclass -class AppConfigSchema: - database: DatabaseConfigSchema # Nested - secret_key: str -``` - -Corresponding YAML: - -```yaml -database: - host: localhost - port: 5432 - # pool_size uses default - -secret_key: my-secret -``` - -### Lists of Dataclasses - -```python -@dataclass -class PluginConfigSchema: - name: str - enabled: bool = True - -@dataclass -class AppConfigSchema: - plugins: list[PluginConfigSchema] -``` +`int | None` permits a null value. It does not itself provide a default or make a +missing field optional. Write `limit: int | None = None` when omission is allowed. +`None` and `MISSING` have different meanings; see below. -```yaml -plugins: - - name: logger - enabled: true - - name: metrics - - name: cache - enabled: false -``` +## Type Coercion -### Dictionaries with Dataclass Values +`Config(schema=Schema, coerce=True)` normalizes compatible concrete values during +transactional edits and resolved argument values during resolution. `coerce=True` +is the default; use `coerce=False` for type checking without normalization. -```python -@dataclass -class ModelConfigSchema: - hidden_size: int - dropout: float +| Conversion | Boundary | +|---|---| +| Numeric text → int/float | Must parse as the requested numeric type | +| Integral float → int | Fractional values such as `3.14` are rejected | +| int → float | Must be exactly representable | +| Recognized text → bool | true/false, yes/no, 1/0, ignoring case and surrounding whitespace | +| Integer 0/1 → bool | Other integers are not accepted by this conversion | +| int/float/bool → str | Arbitrary objects and containers are not stringified | -@dataclass -class ConfigSchema: - models: dict[str, ModelConfigSchema] -``` +Booleans do not satisfy numeric schema fields. Coercion recurses through annotated +containers and dataclasses; dictionary keys are checked without coercion. An +already-valid union alternative wins, so `int | str` preserves `"002"` as text. +Ambiguous coercion, such as `"2"` to `int | float`, raises `ValidationError`. -```yaml -models: - small: - hidden_size: 128 - dropout: 0.1 - large: - hidden_size: 512 - dropout: 0.2 -``` +`config.validate(Schema)` and standalone `validate(mapping, Schema)` check source +without coercing it or executing configured targets/expressions. Custom validators +can still execute Python. ## Custom Validation -Add validation logic with `@validator`: +Use `@validator` for domain constraints beyond types: ```python -from sparkwheel import validator - -@dataclass -class TrainingConfigSchema: - lr: float - batch_size: int - - @validator - def check_lr(self): - """Validate learning rate.""" - if not (0 < self.lr < 1): - raise ValueError(f"lr must be between 0 and 1, got {self.lr}") - - @validator - def check_batch_size(self): - """Validate batch size is power of 2.""" - if self.batch_size <= 0: - raise ValueError("batch_size must be positive") - if self.batch_size & (self.batch_size - 1) != 0: - raise ValueError("batch_size must be power of 2") -``` - -### Cross-Field Validation +from dataclasses import dataclass +from sparkwheel import Config, ValidationError, validator -Validators can check relationships between fields: -```python @dataclass -class ConfigSchema: - start_epoch: int - end_epoch: int - warmup_epochs: int +class ReportOptions: + limit: int @validator - def check_epochs(self): - """Ensure epoch configuration is valid.""" - if self.end_epoch <= self.start_epoch: - raise ValueError("end_epoch must be > start_epoch") - if self.warmup_epochs >= (self.end_epoch - self.start_epoch): - raise ValueError("warmup_epochs too large") -``` + def positive_limit(self): + if self.limit <= 0: + raise ValueError("limit must be positive") -### With Optional Fields -```python -@dataclass -class ConfigSchema: - value: float - max_value: Optional[float] = None - - @validator - def check_max(self): - """Check value doesn't exceed max if specified.""" - if self.max_value is not None and self.value > self.max_value: - raise ValueError(f"value ({self.value}) exceeds max_value ({self.max_value})") +config = Config(schema=ReportOptions).update({"limit": 3}) +try: + config.set("limit", 0) +except ValidationError: + pass +else: + raise AssertionError("Expected the positive-limit check to fail") +assert config.get("limit") == 3 ``` -**Note:** Validators run after type checking. If types are wrong, validation stops there. +Validators run after type checking. Construction metadata such as `_target_` and +lenient extra keys are excluded from the schema instance initializer. Setup failures +and unexpected validator exceptions become contextual `ValidationError` with their +original cause, including inside unions. Prefer pure checks: source rollback +cannot undo external effects. Nested argument dictionaries remain dictionaries; +schemas do not recursively instantiate nested dataclasses for validator methods. ## Discriminated Unions -Use tagged unions for type-safe variants: +A shared `Literal` field can select a variant. This complete example chooses a +limited report without accepting fields from another branch: ```python -from typing import Literal, Union - -@dataclass -class SGDOptimizerSchema: - type: Literal["sgd"] # Discriminator - lr: float - momentum: float = 0.9 +from dataclasses import dataclass +from typing import Literal +from sparkwheel import Config -@dataclass -class AdamOptimizerSchema: - type: Literal["adam"] # Discriminator - lr: float - beta1: float = 0.9 @dataclass -class ConfigSchema: - optimizer: Union[SGDOptimizerSchema, AdamOptimizerSchema] -``` - -YAML: - -```yaml -optimizer: - type: sgd # Selects SGDOptimizer - lr: 0.01 - momentum: 0.95 -``` - -Sparkwheel detects `type` as a discriminator and validates against the matching schema. - -**Error examples:** - -```python -# Missing discriminator -{"optimizer": {"lr": 0.01}} -# ValidationError: Missing discriminator field 'type' - -# Invalid value -{"optimizer": {"type": "rmsprop", "lr": 0.01}} -# ValidationError: Invalid discriminator value 'rmsprop'. Valid: 'sgd', 'adam' - -# Wrong fields for type -{"optimizer": {"type": "adam", "momentum": 0.9}} -# ValidationError: Missing required field 'lr' -``` - -## With Sparkwheel Features +class AllWords: + kind: Literal["all"] -Validation works with references, expressions, and instantiation. -## Type Coercion +@dataclass +class TopWords: + kind: Literal["top"] + limit: int -Sparkwheel automatically converts compatible types when `coerce=True` (default): -```python @dataclass -class ServerConfigSchema: - port: int - timeout: float - enabled: bool - -# Coercion enabled by default -config = Config(schema=ServerConfigSchema) -config.update({ - "port": "8080", # str → int - "timeout": "30.5", # str → float - "enabled": "true" # str → bool -}) - -print(config["port"]) # 8080 (int, not str!) -print(config["timeout"]) # 30.5 (float) -print(config["enabled"]) # True (bool) -``` - -**Supported coercions:** -- `str → int` (e.g., `"42"` → `42`) -- `str → float` (e.g., `"3.14"` → `3.14`) -- `str → bool` (e.g., `"true"` → `True`, `"false"` → `False`) -- `int → float` (e.g., `42` → `42.0`) -- Recursive coercion through lists, dicts, and nested dataclasses +class Options: + report: AllWords | TopWords -**Disable coercion if needed:** -```python -config = Config(schema=ServerConfigSchema, coerce=False) -config.update({ - "port": "8080" # ValidationError: expected int, got str -}) +config = Config(schema=Options).update({"report": {"kind": "top", "limit": 3}}) +assert config.resolve("report") == {"kind": "top", "limit": 3} ``` -## Strict vs Lenient Mode - -Control whether extra fields are rejected: +A missing or invalid discriminator fails validation. A deferred discriminator +postpones branch checks until resolution identifies the branch. Strictness and +missing-value policies also apply within the selected branch. -```python -@dataclass -class MySchema: - required_field: int - -# Strict mode (default) - rejects extra fields -config = Config(schema=MySchema, strict=True) -config.update({ - "required_field": 42, - "extra_field": "oops" # ✗ ValidationError! -}) - -# Lenient mode - allows extra fields -config = Config(schema=MySchema, strict=False) -config.update({ - "required_field": 42, - "extra_field": "ok" # ✓ Allowed -}) -``` +## Strict vs Lenient Mode -Use lenient mode for: -- Development/prototyping -- Gradual schema migration -- Configs with experimental fields +`strict=True` rejects unexpected keys and is the default. Set `strict=False` when +you intentionally permit fields outside the schema; it does not turn off checks +for fields the schema does describe. Both policies apply recursively. ## MISSING Sentinel -Support partial configs with required-but-not-yet-set values: +Use an explicit sentinel when a required value will be supplied later: ```python +from dataclasses import dataclass from sparkwheel import Config, MISSING -@dataclass -class APIConfigSchema: - api_key: str - endpoint: str - timeout: int = 30 - -# Partial config - api_key not set yet -config = Config(schema=APIConfigSchema, allow_missing=True) -config.update({ - "api_key": MISSING, - "endpoint": "https://api.example.com" -}) - -# Later, fill in the missing value -import os -config.set("api_key", os.getenv("API_KEY")) - -# Now validate that nothing is MISSING -config.validate(APIConfigSchema) # Uses allow_missing=False by default -``` - -## Frozen Configs -Prevent modifications after initialization: - -```python -config = Config(schema=MySchema) -config.update("config.yaml") -config.freeze() - -# Mutations now raise FrozenConfigError -config.set("model::lr", 0.001) # ✗ FrozenConfigError! -config.update({"new": "data"}) # ✗ FrozenConfigError! - -# Read operations still work -value = config.get("model::lr") -resolved = config.resolve() - -# Unfreeze if needed -config.unfreeze() -config.set("model::lr", 0.001) # ✓ Now works -``` - -## With Sparkwheel Features - -Validation works with references, expressions, and instantiation. - -### References - -```python @dataclass -class ConfigSchema: - base_lr: float - optimizer_lr: float # Can be a reference - -config = Config(schema=ConfigSchema) -config.update({ - "base_lr": 0.001, - "optimizer_lr": "@base_lr" # Reference allowed -}) -``` +class Options: + limit: int -### Expressions -```python -@dataclass -class ConfigSchema: - batch_size: int - total_steps: int # Computed - -config = Config(schema=ConfigSchema) -config.update({ - "batch_size": 32, - "total_steps": "$@batch_size * 100" # Expression allowed -}) +config = Config(schema=Options, allow_missing=True).update({"limit": MISSING}) +config.set("limit", 3) +config.validate(Options) ``` -### Instantiation - -Special keys like `_target_` are automatically ignored: - -```python -@dataclass -class OptimizerConfigSchema: - lr: float - momentum: float = 0.9 - -config = Config(schema=OptimizerConfigSchema) -config.update({ - "_target_": "torch.optim.SGD", # Ignored by validation - "lr": 0.001, - "momentum": 0.95 -}) -``` +The final explicit `validate()` uses its normal completeness policy and rejects a +remaining `MISSING`. `None` is different: it must satisfy the declared type. -## Error Messages +## With Sparkwheel Features -### Type Mismatch +Schemas on `_target_` nodes describe **constructor arguments**, not the class of +the constructed result. This standalone example validates a deferred argument +before passing it to `dict`: ```python -# Expected int, got str -# ValidationError: Validation error at 'port': Type mismatch -# Expected type: int -# Actual type: str -# Actual value: '8080' -``` +from dataclasses import dataclass +from sparkwheel import Config -### Missing Field -```python -# ValidationError: Validation error at 'required_field': -# Missing required field 'required_field' -# Expected type: str -``` +@dataclass +class Arguments: + limit: int -### Unexpected Field -```python -# ValidationError: Validation error at 'unexpected': -# Unexpected field 'unexpected' not in schema ConfigSchema +config = Config(schema=Arguments).update({"_target_": "builtins.dict", "limit": "$2 + 1"}) +assert config.get("limit") == "$2 + 1" +assert config.resolve() == {"limit": 3} ``` -### Nested Errors - -```python -# ValidationError: Validation error at 'database.port': Type mismatch -# Expected type: int -# Actual type: str -# Actual value: 'wrong' -``` +Nested components and aliases are checked through their resolved argument mappings +while constructors receive native objects. Shared container/component identities +are preserved. An alias cannot coerce an already-built object's argument mapping +into a different schema. Expression coercion changes the runtime value while +preserving the authored expression string. ## Validation Timing -### Continuous (Recommended) - -```python -# Validates on every update() and set() -config = Config(schema=MySchema) -config.update("config.yaml") -config.set("port", "8080") # Validates immediately! -``` - -### Explicit +| Operation | What is checked | +|---|---| +| `set()` / `update()` with a schema | Staged source; unresolved `@`, `$`, `%` remain deferred | +| `resolve(path)` | That node, ordinary dependencies and any needed union discriminator | +| `resolve()` | Complete resolved argument tree, including root cross-field checks | +| `config.validate(Schema)` | Current source, without coercion or configured resolution | -```python -# Load without schema, validate later -config = Config() -config.update("config.yaml") -# ... maybe modify ... -config.validate(MySchema) -``` +Partial resolution does not construct unrelated nodes merely to run a root +cross-field validator. A successful partial check is not validation of the whole +application. Invalid resolved arguments fail before their target is called, but +dependencies, imports and expressions may already have executed; resolution does +not roll back those effects. -### Standalone Function +`instantiate=False` still checks resolved arguments. `eval_expr=False` keeps +expression wrappers and postpones resolved-value checks. `retain()` explicitly +rejects attached schemas; use schema-aware `Config.resolve()` for those configs. -```python -from sparkwheel import validate - -# Validate a dict directly -validate(config_dict, AppConfigSchema) -``` - -## Complete Example - -```python -from dataclasses import dataclass -from typing import Optional -from sparkwheel import Config, validator - -@dataclass -class DatabaseConfigSchema: - host: str - port: int - database: str - username: str - password: str - pool_size: int = 10 - timeout: int = 30 - -@dataclass -class APIConfigSchema: - host: str = "0.0.0.0" - port: int = 8000 - workers: int = 4 - - @validator - def check_port(self): - if not (1024 <= self.port <= 65535): - raise ValueError(f"port must be 1024-65535, got {self.port}") - -@dataclass -class AppConfigSchema: - app_name: str - environment: str - debug: bool = False - api: APIConfigSchema - database: DatabaseConfigSchema - -# Load and validate continuously -config = Config(schema=AppConfigSchema) -config.update("production.yaml") - -# Access validated config -print(f"Starting {config['app_name']} on port {config['api::port']}") - -# Freeze to prevent modifications -config.freeze() -``` - -The YAML: +## Frozen Configs -```yaml -app_name: "My API" -environment: production -debug: false +Freeze after configuring when later source edits should fail. Frozen reads remain +available, but runtime objects and opaque Python inputs are not deeply frozen. +See [configuration lifetimes](basics.md#frozen-configs). -api: - port: 3000 - workers: 8 +## Error Messages -database: - host: db.example.com - port: 5432 - database: myapp - username: "$import os; os.getenv('DB_USER')" - password: "$import os; os.getenv('DB_PASSWORD')" - pool_size: 20 -``` +`ValidationError` identifies the field and expected/actual value or type. Source +locations and native causes help distinguish a type mismatch from a validator's +own failure. See [troubleshooting](troubleshooting.md) for a staged diagnosis. ## Next Steps -- **[Configuration Basics](basics.md)** - Learn config management -- **[References](references.md)** - Link values with @ -- **[Expressions](expressions.md)** - Compute values with $ +[Configuration model](basics.md) · [Python authoring](instantiation.md) · +[Quick reference](quick-reference.md) diff --git a/docs/user-guide/troubleshooting.md b/docs/user-guide/troubleshooting.md new file mode 100644 index 0000000..5a50e46 --- /dev/null +++ b/docs/user-guide/troubleshooting.md @@ -0,0 +1,96 @@ +# Troubleshooting + +Start with the smallest failing input and the operation that failed. Loading, +source validation and object construction are different stages. + +## Inspect Before Executing + +In the quick-start directory, this reads and prints the composed source without +constructing the counter: + +```python +from sparkwheel import Config + +config = Config().update("config.yaml") +print(config.get()) +print(config.get("counts")) +``` + +`get()` preserves references and expressions. External `%file.yaml` includes have +already been loaded during `update()`. This inspection shows definitions, not +whether imports, constructors or application behavior will succeed. + +When you intend to construct a component, resolve its specific path first: +`config.resolve("counts")`. Resolving the root builds the whole requested runtime +tree and may trigger unrelated application work. Neither `instantiate=False` nor +`eval_expr=False` alone makes resolution static; imports and enabled processing +can still execute Python. + +## Choose the Next Check + +| Symptom | Check | Next action | +|---|---|---| +| File not found | Working directory and supplied filename | Use the intended file path; run the example from its stated directory | +| Duplicate YAML key | Both locations in the error | Keep one definition, or compose separate files deliberately | +| Missing reference | Exact `::` path in the composed source | Correct the path or load the definition before resolution | +| Cannot locate target | Full import path and active Python environment | Import the same module directly with that Python | +| Constructor error | Native exception cause and constructor signature | Correct the arguments; use `=component` when replacing an inherited recipe | +| Unexpected list entries | List composition versus replacement | Use `=list` to replace, or indexed deletion to remove entries | +| An edit seems ignored | Raw nested writes or an already-returned object | Use `set()`/`update()` and explicitly obtain the new generation | +| Shared state changes elsewhere | `@` references and mutable constructor inputs | Use `%` for a copied definition and review any nested `@` dependencies | +| String becomes boolean/null | YAML override parsing | Preserve inner YAML quotes for values intended as text | +| A schema accepts a deferred value | Raw versus resolved validation stage | Resolve the intended node when execution is appropriate | + +## Check Your Python Module + +For the [custom report example](instantiation.md), run from `custom-report` with +the installed environment active: + +```bash +python -c "from reports import WordReport; print(WordReport(['blue']).render())" +``` + +Expected output is `Word counts: blue=1`. If this fails, fix the Python import or +class first. Sparkwheel does not search arbitrary project directories for targets. + +A target lookup failure raises `TargetNotFoundError`. A failure while calling it +raises `InstantiationError`; the native error remains in its `__cause__` chain. +There may be another Sparkwheel wrapper before the native error, so follow the +whole chain or read the complete traceback. Expression failures similarly retain +their cause through `EvaluationError`. Use the native error to distinguish a +configuration mistake from a failure in application code. + +## Unexpected Debugger or Missing Expression Result + +`SPARKWHEEL_DEBUG=1` enables interactive debugging. The flag is read when Sparkwheel +is imported; set it before starting Python, and restart after changing it. +Construction calls `breakpoint()`, while expressions run through `pdb.run()` and +return `None` instead of their normal result. This is debugger behavior, not +verbose logging. Leave the flag unset or `0` for ordinary or unattended execution. + +An individual component can also request `_mode_: debug` or `_debug_: true`; +check those settings if a debugger appears with the environment flag disabled. +`get()` remains source inspection and does not trigger configured construction or +expression debugging. + +## Validate at the Right Stage + +`Config(schema=Schema)` checks concrete source values during edits and resolved +arguments before calling their target. Deferred values cannot be fully checked +without resolution. `config.validate(Schema)` checks source without executing +configured targets or expressions, but custom validators themselves are Python +and may have effects. + +Partial resolution does not run unrelated root cross-field checks. See +[validation timing](schema-validation.md#validation-timing) before interpreting a +successful check as validation of the entire application. + +## Report a Reproducible Issue + +Include the package version and import path, the smallest relevant Python/YAML +example, the exact operation, the full exception and its cause, and the expected +result. Redact secrets and private data. State whether the problem occurs during +loading, raw validation or resolution, and whether a source edit preceded it. + +[Configuration model](basics.md) · [Composition](operators.md) · +[Advanced contracts](advanced.md) diff --git a/justfile b/justfile index a648c24..7115424 100644 --- a/justfile +++ b/justfile @@ -47,10 +47,12 @@ docs port="8000": uv run --only-group doc mkdocs serve --dev-addr=localhost:{{port}} bump part="patch": - uvx bump-my-version bump {{part}} --verbose + @case {{quote(part)}} in major|minor|patch|release) ;; *) echo "Supported version parts: major, minor, patch, release" >&2; exit 2 ;; esac + uvx --from bump-my-version==0.30.1 bump-my-version bump {{quote(part)}} --verbose bump-dry part="patch": - uvx bump-my-version bump {{part}} --dry-run --verbose --allow-dirty + @case {{quote(part)}} in major|minor|patch|release) ;; *) echo "Supported version parts: major, minor, patch, release" >&2; exit 2 ;; esac + uvx --from bump-my-version==0.30.1 bump-my-version bump {{quote(part)}} --dry-run --no-commit --no-tag --verbose --allow-dirty push: git push && git push --tags diff --git a/mkdocs.yml b/mkdocs.yml index 5dc01d1..3474b88 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -67,20 +67,22 @@ plugins: nav: - Home: index.md - - Getting Started: + - Start: - Installation: getting-started/installation.md - Quick Start: getting-started/quickstart.md - - User Guide: - - Configuration Basics: user-guide/basics.md + - Guides: + - Configuration Model: user-guide/basics.md + - Sharing and Copying: user-guide/references.md + - Python Authoring: user-guide/instantiation.md + - Compose and Edit: user-guide/operators.md + - CLI Overrides: user-guide/cli.md + - Validate Values: user-guide/schema-validation.md + - Troubleshooting: user-guide/troubleshooting.md + - Reference: - Quick Reference: user-guide/quick-reference.md - - References: user-guide/references.md - Expressions: user-guide/expressions.md - - Instantiation: user-guide/instantiation.md - - Merging & Deleting: user-guide/operators.md - - Schema Validation: user-guide/schema-validation.md - - CLI Support: user-guide/cli.md - - Advanced Features: user-guide/advanced.md - - API: reference/ + - Advanced Contracts: user-guide/advanced.md + - API: reference/ markdown_extensions: - pymdownx.highlight: diff --git a/pyproject.toml b/pyproject.toml index 5d52396..9f307a9 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "uv_build" [project] name = "sparkwheel" -version = "0.0.11" +version = "0.1.0.dev0" description = "A powerful YAML-based configuration system with references, expressions, and dynamic instantiation" authors = [ { name = "Project Lighter", email = "noreply@lighter.com" }, @@ -12,6 +12,7 @@ authors = [ requires-python = ">=3.10" readme = "README.md" license = "Apache-2.0" +license-files = ["LICENSE"] keywords = ["configuration", "yaml", "config", "deep-learning", "machine-learning"] classifiers = [ "Development Status :: 3 - Alpha", @@ -59,6 +60,7 @@ maintain = ["bump-my-version==0.30.1"] quality = ["ruff>=0.5.0"] types = [ "mypy>=1.14.1", + "types-PyYAML>=6.0.12", "typing-extensions>=4.4.0", ] test = [ @@ -131,9 +133,16 @@ select = ["B", "E", "F", "I", "W"] ignore = ["E203", "E501", "N813"] [tool.bumpversion] -current_version = "0.0.11" +current_version = "0.1.0.dev0" commit = true tag = true +parse = '^(?P\d+)\.(?P\d+)\.(?P\d+)(?:\.(?Pdev)(?P\d+))?$' +serialize = ["{major}.{minor}.{patch}.{release}{dev}", "{major}.{minor}.{patch}"] + +[tool.bumpversion.parts.release] +values = ["dev", "final"] +optional_value = "final" +first_value = "final" [[tool.bumpversion.files]] filename = "src/sparkwheel/__init__.py" @@ -144,3 +153,13 @@ replace = "__version__ = \"{new_version}\"" filename = "pyproject.toml" search = 'version = "{current_version}"' replace = 'version = "{new_version}"' + +# Change only this project's editable root, never a dependency version. +[[tool.bumpversion.files]] +filename = "uv.lock" +search = '''name = "sparkwheel" +version = "{current_version}" +source = {{ editable = "." }}''' +replace = '''name = "sparkwheel" +version = "{new_version}" +source = {{ editable = "." }}''' diff --git a/src/sparkwheel/__init__.py b/src/sparkwheel/__init__.py index 067eae5..a0c8718 100644 --- a/src/sparkwheel/__init__.py +++ b/src/sparkwheel/__init__.py @@ -22,7 +22,7 @@ TargetNotFoundError, ) -__version__ = "0.0.11" +__version__ = "0.1.0.dev0" __all__ = [ "__version__", diff --git a/src/sparkwheel/config.py b/src/sparkwheel/config.py index 87f3423..c43d1bb 100644 --- a/src/sparkwheel/config.py +++ b/src/sparkwheel/config.py @@ -95,12 +95,20 @@ See Config class docstring for full API details. """ +from copy import copy, deepcopy from pathlib import Path -from typing import Any +from typing import TYPE_CHECKING, Any from .loader import Loader from .locations import LocationRegistry -from .operators import MergeContext, _validate_delete_operator, apply_operators, validate_operators +from .operators import ( + MergeContext, + _copy_containers, + _normalize_update_value, + _validate_delete_operator, + apply_operators, + validate_operators, +) from .parser import Parser from .path_utils import get_by_id, split_id from .preprocessor import Preprocessor @@ -109,6 +117,9 @@ from .utils.constants import ID_SEP_KEY, REMOVE_KEY, REPLACE_KEY from .utils.exceptions import ConfigKeyError, build_missing_key_error +if TYPE_CHECKING: + from .construction import RetainedConfig + __all__ = ["Config", "parse_overrides"] @@ -186,7 +197,12 @@ def __init__( """ self._data: dict[str, Any] = data or {} # Start with provided data or empty self._locations = LocationRegistry() - self._resolver = Resolver() + if schema is None: + self._resolver = Resolver() + else: + from .schema_resolver import _SchemaResolver + + self._resolver = _SchemaResolver(self) self._is_parsed = False self._frozen = False # Set via freeze() method later @@ -197,10 +213,11 @@ def __init__( self._allow_missing: bool = allow_missing # Process imports (import string module paths) - self._imports: dict[str, Any] = {} + self._provided_imports: dict[str, Any] = {} if isinstance(imports, dict): for k, v in imports.items(): - self._imports[k] = optional_import(v)[0] if isinstance(v, str) else v + self._provided_imports[k] = optional_import(v)[0] if isinstance(v, str) else v + self._imports = dict(self._provided_imports) self._loader = Loader() self._preprocessor = Preprocessor(self._loader, self._imports) @@ -213,7 +230,9 @@ def get(self, id: str = "", default: Any = None) -> Any: - Returns `@` references as strings (e.g., "@model::lr") - Returns `$` expressions as strings (e.g., "$@lr * 2") - - Returns `%` raw references already expanded (eager expansion during update()) + - Retains local `%` definitions and `_imports_`, even after resolution + - External `%file.yaml` includes are snapshots expanded during update() + - Does not import, instantiate, or evaluate expressions - Fast, no resolution overhead Use this when you need to: @@ -258,7 +277,8 @@ def get(self, id: str = "", default: Any = None) -> Any: 999 """ try: - return self._get_by_id(id) + value = self._get_by_id(id) + return _copy_containers(value) if self._frozen else value except (KeyError, IndexError, TypeError): return default @@ -284,28 +304,66 @@ def set(self, id: str, value: Any) -> None: if self._frozen: raise FrozenConfigError("Cannot modify frozen config", field_path=id) + candidate = self._edit_candidate() + candidate._set_inplace(id, value) + self._publish_edit(candidate) + + def _edit_candidate(self) -> "Config": + """Stage source/metadata without borrowing the live resolver or namespace.""" + candidate = Config(data=_copy_containers(self._data)) + candidate._locations = deepcopy(self._locations) + return candidate + + def _publish_edit(self, candidate: "Config") -> None: + if self._schema: + from .schema import _coerce_field + from .schema import validate as validate_schema + + if self._coerce: + candidate._data = _coerce_field( + candidate._data, + self._schema, + metadata=candidate._locations, + allow_missing=self._allow_missing, + strict=self._strict, + ) + validate_schema( + candidate._data, + self._schema, + metadata=candidate._locations, + allow_missing=self._allow_missing, + strict=self._strict, + ) + self._data = candidate._data + self._locations = candidate._locations + self._invalidate_resolution() + + def _set_inplace(self, id: str, value: Any) -> None: + """Apply one path assignment to a staged source tree.""" + value = _copy_containers(value) if id == "": self._data = value self._invalidate_resolution() return - keys = split_id(id) - - # Ensure root is dict if not isinstance(self._data, dict): self._data = {} # type: ignore[unreachable] - - # Create missing intermediate paths - current = self._data - for k in keys[:-1]: - if k not in current: - current[k] = {} - elif not isinstance(current[k], dict): - current[k] = {} - current = current[k] - - # Set final value - current[keys[-1]] = value + current: Any = self._data + for key in keys[:-1]: + if isinstance(current, dict): + if key not in current: + current[key] = {} + current = current[key] + elif isinstance(current, list): + current = current[int(key)] + else: + raise TypeError(f"Cannot traverse scalar while setting '{id}'") + if isinstance(current, dict): + current[keys[-1]] = value + elif isinstance(current, list): + current[int(keys[-1])] = value + else: + raise TypeError(f"Cannot assign a child of a scalar while setting '{id}'") self._invalidate_resolution() def validate(self, schema: type) -> None: @@ -340,13 +398,17 @@ def freeze(self) -> None: - set() raises FrozenConfigError - update() raises FrozenConfigError - resolve() still works (read-only) - - get() still works (read-only) + - get()/subscript access return detached container snapshots + - previously returned source containers cannot mutate the frozen source + - opaque Python leaves and runtime objects remain caller-owned Example: >>> config = Config(schema=MySchema).update("config.yaml") >>> config.freeze() >>> config.set("model::lr", 0.001) # Raises FrozenConfigError """ + # Detach containers already handed out before freezing. + self._data = _copy_containers(self._data) self._frozen = True def unfreeze(self) -> None: @@ -426,6 +488,13 @@ def update(self, source: PathLike | dict[str, Any] | "Config" | str) -> "Config" if self._frozen: raise FrozenConfigError("Cannot update frozen config") + candidate = self._edit_candidate() + candidate._update_inplace(source) + self._publish_edit(candidate) + return self + + def _update_inplace(self, source: PathLike | dict[str, Any] | "Config" | str) -> None: + """Apply a complete update to a staged candidate before validation.""" if isinstance(source, Config): self._update_from_config(source) elif isinstance(source, dict): @@ -447,21 +516,6 @@ def update(self, source: PathLike | dict[str, Any] | "Config" | str) -> "Config" self._data, self._data, id="", locations=self._locations, external_only=True ) - # Validate after raw ref expansion if schema exists - # This validates the final structure, not intermediate raw reference strings - if self._schema: - from .schema import validate as validate_schema - - validate_schema( - self._data, - self._schema, - metadata=self._locations, - allow_missing=self._allow_missing, - strict=self._strict, - ) - - return self # Enable chaining - def _update_from_config(self, source: "Config") -> None: """Update from another Config instance.""" context = MergeContext(locations=source.locations) @@ -477,13 +531,15 @@ def _apply_path_updates(self, source: dict[str, Any]) -> None: """Apply nested path updates (e.g., model::lr=value, =model=replace, ~old::param=null).""" for key, value in source.items(): if not isinstance(key, str): - self.set(str(key), value) # type: ignore[unreachable] + context = MergeContext(locations=self._locations, current_path=str(key)) # type: ignore[unreachable] + self._set_inplace(str(key), _normalize_update_value(value, context)) continue if key.startswith(REPLACE_KEY): # Replace operator: =key (explicit override) actual_key = key[1:] - self.set(actual_key, value) + context = MergeContext(locations=self._locations, current_path=actual_key) + self._set_inplace(actual_key, _normalize_update_value(value, context)) elif key.startswith(REMOVE_KEY): # Delete operator: ~key @@ -516,19 +572,26 @@ def _apply_path_updates(self, source: dict[str, Any]) -> None: available_keys = list(self._data.keys()) if isinstance(self._data, dict) else [] raise build_missing_key_error(error_key, available_keys, source_location, parent_key=parent_key_name) - self._delete_nested_key(actual_key) + if isinstance(value, list): + context = MergeContext(locations=self._locations, current_path=actual_key) + remaining = apply_operators({"value": self._get_by_id(actual_key)}, {"~value": value}, context=context) + self._set_inplace(actual_key, remaining["value"]) + else: + self._delete_nested_key(actual_key) else: # Default: compose (merge dict or extend list) if key in self and isinstance(self[key], dict) and isinstance(value, dict): context = MergeContext(locations=self._locations, current_path=key) merged = apply_operators(self[key], value, context=context) - self.set(key, merged) + self._set_inplace(key, merged) elif key in self and isinstance(self[key], list) and isinstance(value, list): - self.set(key, self[key] + value) + context = MergeContext(locations=self._locations, current_path=key) + self._set_inplace(key, self[key] + _normalize_update_value(value, context)) else: # Normal set (handles nested paths with ::) - self.set(key, value) + context = MergeContext(locations=self._locations, current_path=key) + self._set_inplace(key, _normalize_update_value(value, context)) def _delete_nested_key(self, key: str) -> None: """Delete a key, supporting nested paths with ::.""" @@ -538,6 +601,8 @@ def _delete_nested_key(self, key: str) -> None: parent = self[parent_id] if parent_id else self._data if isinstance(parent, dict) and keys[-1] in parent: del parent[keys[-1]] + elif isinstance(parent, list): + del parent[int(keys[-1])] else: # Top-level key if isinstance(self._data, dict) and key in self._data: @@ -574,6 +639,16 @@ def _update_from_override_string(self, override: str) -> None: overrides_dict = parse_overrides([override]) self._apply_path_updates(overrides_dict) + def retain(self) -> "RetainedConfig": + """Capture source and caller imports for independent construction scopes. + + This does not resolve, import, or construct configured objects. Opaque + Python inputs retain their identity; existing runtime caches are excluded. + """ + from .construction import RetainedConfig + + return RetainedConfig(self) + def resolve( self, id: str = "", @@ -591,10 +666,12 @@ def resolve( 4. Caches results in a separate resolution cache (`_resolver._resolved`) Unlike get(), which always returns raw `_data`, resolve() performs full processing - and uses a separate cache for efficiency. + and uses a separate cache for each (instantiate, eval_expr) pair. Repeating + a mode preserves its shared objects even after resolving in another mode. + Editing the config or passing lazy=False invalidates all modes. Processing stages: - - `%` raw references: Already expanded during update() (eager) + - `%` raw references: External includes expanded on update; local copies compiled here - `@` resolved references: Resolved now (lazy, supports circular deps) - `$` expressions: Evaluated now (lazy) - `_target_` components: Instantiated now (lazy) @@ -642,8 +719,8 @@ def resolve( >>> type(optimizer).__name__ 'Adam' - >>> # Disable instantiation (useful for inspection) - >>> config.resolve("optimizer", instantiate=False) + >>> # Disable instantiation (returns a Component wrapper) + >>> config.resolve("optimizer", instantiate=False).get_config() {'_target_': 'torch.optim.Adam', 'lr': 0.001} >>> # With default value @@ -681,31 +758,33 @@ def _parse(self, reset: bool = True) -> None: """ # Reset resolver if requested if reset: - self._resolver.reset() + self._replace_resolver() - # Process _imports_ key if present in config data - # This allows YAML-based imports that become available to all expressions - self._process_imports_key() + # Compile a separate tree. Source imports, local copies and relative + # references must remain available for export and the next generation. + working_data = _copy_containers(self._data) + self._imports = dict(self._provided_imports) + self._process_imports_key(working_data) # Phase 2: Expand local raw references (%key) now that all composition is complete # CLI overrides have been applied, so local refs will see final values - self._data = self._preprocessor.process_raw_refs( - self._data, self._data, id="", locations=self._locations, external_only=False + working_data = self._preprocessor.process_raw_refs( + working_data, working_data, id="", locations=self._locations, external_only=False ) # Stage 1: Preprocess (@:: relative resolved IDs) - self._data = self._preprocessor.process(self._data, self._data, id="") + working_data = self._preprocessor.process(working_data, working_data, id="") # Stage 2: Parse config tree to create Items parser = Parser(globals=self._imports, metadata=self._locations) - items = parser.parse(self._data) + items = parser.parse(working_data) # Stage 3: Add items to resolver self._resolver.add_items(items) self._is_parsed = True - def _process_imports_key(self) -> None: + def _process_imports_key(self, working_data: dict[str, Any]) -> None: """Process _imports_ key from config data. The _imports_ key allows declaring imports directly in YAML: @@ -721,13 +800,14 @@ def _process_imports_key(self) -> None: ``` These imports become available to all expressions in the config. - The _imports_ key is removed from the data after processing. + The _imports_ key is removed only from the working compilation tree. + The authored declaration remains available through get() and export. """ imports_key = "_imports_" - if imports_key not in self._data: + if imports_key not in working_data: return - imports_config = self._data.pop(imports_key) + imports_config = working_data.pop(imports_key) if not isinstance(imports_config, dict): return @@ -766,10 +846,23 @@ def _get_by_id(self, id: str) -> Any: """ return get_by_id(self._data, id) + def _replace_resolver(self) -> None: + """Start a generation without mutating resolvers held by escaped values.""" + if self._schema is None: + self._resolver = Resolver() + else: + from .schema_resolver import _SchemaResolver + + # Schema policy and source locations belong to this generation too. + # Copy the owner record, not opaque inputs or runtime graph objects. + self._resolver = _SchemaResolver(copy(self)) + def _invalidate_resolution(self) -> None: """Invalidate cached resolution (called when config changes).""" self._is_parsed = False - self._resolver.reset() + self._replace_resolver() + # Replacing the namespace preserves objects held from older generations. + self._imports = dict(self._provided_imports) def __getitem__(self, id: str) -> Any: """Get config value by ID (subscript access). @@ -785,7 +878,8 @@ def __getitem__(self, id: str) -> Any: >>> config["model::lr"] 0.001 """ - return self._get_by_id(id) + value = self._get_by_id(id) + return _copy_containers(value) if self._frozen else value def __setitem__(self, id: str, value: Any) -> None: """Set config value by ID (subscript access). @@ -828,7 +922,7 @@ def export_config_file(config: dict[str, Any], filepath: PathLike, **kwargs: Any filepath: Target file path kwargs: Additional arguments for yaml.safe_dump """ - import yaml # type: ignore[import-untyped] + import yaml filepath_str = str(Path(filepath)) with open(filepath_str, "w") as f: diff --git a/src/sparkwheel/construction.py b/src/sparkwheel/construction.py new file mode 100644 index 0000000..c9f532d --- /dev/null +++ b/src/sparkwheel/construction.py @@ -0,0 +1,356 @@ +"""Retained definitions and isolated resolution scopes for framework integrations.""" + +from collections.abc import Iterator +from copy import deepcopy +from dataclasses import dataclass +from typing import Any + +from .config import Config +from .items import Component, Expression +from .operators import _copy_containers +from .parser import Parser +from .path_utils import get_by_id, resolve_relative_ids, split_file_and_id, split_id +from .resolver import Resolver +from .utils.exceptions import CircularReferenceError, ConfigKeyError + + +class BlockedPathError(ValueError): + """A requested path enters a registered blocked construction subtree.""" + + def __init__(self, requested_path: str, blocked_path: str): + self.requested_path = requested_path + self.blocked_path = blocked_path + super().__init__(f"Cannot resolve '{requested_path}': construction path '{blocked_path}' is blocked") + + def __reduce__(self) -> tuple[Any, ...]: + return type(self), (self.requested_path, self.blocked_path), self.__dict__ + + +def _within(path: str, parent: str) -> bool: + return path == parent or not parent or path.startswith(parent + "::") + + +def _child(parent: str, key: Any) -> str: + return f"{parent}::{key}" if parent else str(key) + + +@dataclass(frozen=True) +class _Boundary: + """Internal placeholder that dependency scanning can see without traversing.""" + + path: str + + +class RetainedConfig: + """A container snapshot without runtime caches; obtain through Config.retain().""" + + def __init__(self, config: Config): + if config._schema is not None: + raise NotImplementedError("Retained scopes do not yet support attached schemas; use Config.resolve()") + self._source = _copy_containers(config.get()) + self._imports = dict(config._provided_imports) + self._locations = deepcopy(config._locations) + + def definition(self, path: str = "") -> Any: + """Inspect one retained definition without running configured Python. + + Follow only local raw-copy chains needed to select the node. References + and expressions inside the returned container remain authored strings. + Containers are detached; opaque Python leaves retain their identity. + """ + if not isinstance(path, str): + raise TypeError("Definition paths must be strings") + try: + value, origin = self._definition(path) + value, _ = self._unwrap(value, origin, frozenset()) + except (KeyError, TypeError, IndexError) as error: + raise ConfigKeyError(f"Unknown definition path '{path}'", missing_key=path) from error + return _copy_containers(value) + + def scope( + self, + *, + bindings: dict[str, Any] | None = None, + blocked_paths: set[str] | None = None, + ) -> "ResolutionScope": + """Create a fresh scope without executing configured imports or targets.""" + return ResolutionScope(self, bindings or {}, blocked_paths or set()) + + def _definition(self, path: str, seen: frozenset[str] = frozenset()) -> tuple[Any, str]: + """Find a source path, following local copies without executing Python.""" + node = self._source + origin = "" + for key in split_id(path) if path else []: + node, origin = self._unwrap(node, origin, seen) + if isinstance(node, list) and (not key.isdigit() or str(int(key)) != key): + raise KeyError(f"Construction list index '{key}' must be a non-negative canonical index") + node = get_by_id(node, key) + origin = _child(origin, key) + return node, origin + + def _unwrap(self, value: Any, origin: str, seen: frozenset[str]) -> tuple[Any, str]: + while isinstance(value, str) and value.startswith("%"): + reference = resolve_relative_ids(origin, value) + if reference in seen: + raise CircularReferenceError(f"Circular raw reference at '{origin}': {reference}") + seen = seen | {reference} + file, path = split_file_and_id(reference[1:]) + if file: + raise ValueError("Load external includes with Config.update() before retaining their descendant paths") + value, origin = self._definition(path, seen) + return value, origin + + +class ResolutionScope: + """One runtime graph with exact bindings and explicit forbidden source paths.""" + + def __init__(self, recipe: RetainedConfig, bindings: dict[str, Any], blocked_paths: set[str]): + self._recipe = recipe + self._bindings: dict[str, Any] = {} + self._blocked = set(blocked_paths) + self._consumed: dict[str, Any] = {} + self._components: dict[str, Any] = {} + self._active_paths: set[str] = set() + self._reserved_paths: set[str] = set() + self._config = Config() + self._config._imports = dict(recipe._imports) + self._imports_loaded = False + self._resolver = _ScopedResolver(self) + for path in self._blocked: + self._validate_path(path) + for path, value in bindings.items(): + self.bind(path, value) + + def _validate_path(self, path: str) -> None: + if not isinstance(path, str): + raise TypeError("Construction paths must be strings") + try: + self._recipe._definition(path) + except (KeyError, TypeError, IndexError) as error: + raise ConfigKeyError(f"Unknown construction path '{path}'", missing_key=path) from error + if path == "_imports_" or path.startswith("_imports_::"): + raise ValueError("Import declarations are not runtime binding paths") + + def _check_block(self, path: str) -> None: + for blocked in self._blocked: + if _within(path, blocked): + raise BlockedPathError(path, blocked) + + def bind(self, path: str, value: Any) -> None: + """Bind a valid exact source path, rejecting incompatible prior use.""" + self._validate_path(path) + self._check_block(path) + known = any(path in values for values in (self._bindings, self._resolver._resolved, self._consumed)) + if not known and any(_within(path, reserved) or _within(reserved, path) for reserved in self._reserved_paths): + raise ValueError(f"Cannot bind '{path}' after use in a published expression; create a fresh scope") + for values in (self._bindings, self._resolver._resolved, self._consumed): + if path in values: + if values[path] is not value: + raise ValueError(f"Cannot rebind '{path}' to a different identity after use; create a fresh scope") + break + self._bindings[path] = value + self._resolver._resolved[path] = value + # Existing compiled ancestors may contain this path. Substitution observes + # bindings directly; already consumed ancestors were checked above. + + def resolve(self, path: str = "") -> Any: + """Resolve a source path after checking explicit dependency boundaries.""" + self._validate_path(path) + self._active_paths = set() + self._preflight(path, set()) + if path in self._resolver._resolved: + return self._resolver._resolved[path] + if not self._imports_loaded: + self._load_imports() + self._imports_loaded = True + return self._resolver.resolve(path) + + def materialized_components(self) -> dict[str, Any]: + """Return constructed _target_ identities, excluding bindings/expressions.""" + return dict(self._components) + + def _preflight(self, path: str, visiting: set[str]) -> None: + self._check_block(path) + self._active_paths.add(path) + if path in self._bindings or path in self._resolver._resolved: + return + if path in visiting: + raise CircularReferenceError(f"Circular reference while resolving '{path}'") + self._resolver._prepare_item(path) + visiting.add(path) + item = self._resolver._items[path] + for dependency in self._resolver._dependency_refs(item.get_config(), path): + self._preflight(dependency, visiting) + visiting.remove(path) + + def _literal_guard(self, value: Any, path: str) -> tuple[bool | None, Any]: + if not Component.is_instantiable(value): + return False, value + guard_path = _child(path, "_disabled_") + if any(_within(guard_path, blocked) for blocked in self._blocked): + return None, value + if guard_path in self._bindings: + value = {**value, "_disabled_": self._bindings[guard_path]} + return Resolver._literal_disabled(value), value + + def _compile(self, value: Any, path: str, origin: str, seen: frozenset[str] = frozenset()) -> Any: + if path in self._bindings or any(_within(path, blocked) for blocked in self._blocked): + # Represent boundaries as pure references; normal Resolver scans and + # substitution then preserve the boundary without walking inside it. + return _Boundary(path) + if isinstance(value, str) and value.startswith("%"): + reference = resolve_relative_ids(origin, value) + if reference in seen: + raise CircularReferenceError(f"Circular raw reference at '{path}': {reference}") + file, target = split_file_and_id(reference[1:]) + if file: + # Config.update freezes external includes before retain. Do not + # silently turn an unexpanded include into a later file read. + raise ValueError("Load external includes with Config.update() before retaining") + copied, copied_origin = self._recipe._definition(target) + return self._compile(copied, path, copied_origin, seen | {reference}) + disabled, guarded = self._literal_guard(value, path) + if disabled is True: + return _copy_containers(guarded) + if isinstance(value, dict): + return { + key: self._compile(item, _child(path, key), _child(origin, key), seen) + for key, item in value.items() + if path or key != "_imports_" + } + if isinstance(value, list): + return [self._compile(item, _child(path, i), _child(origin, i), seen) for i, item in enumerate(value)] + return resolve_relative_ids(path, value) if isinstance(value, str) else value + + def _load_imports(self) -> None: + declaration = self._recipe._source.get("_imports_") + self._config._process_imports_key({"_imports_": declaration}) + + def collect(value: Any, path: str) -> None: + if path in self._bindings or any(_within(path, blocked) for blocked in self._blocked): + return + if self._literal_guard(value, path)[0] is True: + return + if isinstance(value, dict): + for key, child in value.items(): + if path or key != "_imports_": + collect(child, _child(path, key)) + elif isinstance(value, list): + for i, child in enumerate(value): + collect(child, _child(path, i)) + elif isinstance(value, str) and Expression.is_import_statement(value): + self._resolver._prepare_item(path) + + collect(self._recipe._source, "") + + +class _ScopedResolver(Resolver): + def __init__(self, scope: ResolutionScope): + super().__init__() + self._scope = scope + + def _prepare_item(self, id: str) -> None: + if id in self._items: + return + self._scope._validate_path(id) + value, origin = self._scope._recipe._definition(id) + compiled = self._scope._compile(value, id, origin) + # Boundaries must be dependencies even when their source is a scalar. + compiled = self._reference_boundaries(compiled) + parser = Parser(globals=self._scope._config._imports, metadata=self._scope._recipe._locations) + items = parser.parse(compiled, id_prefix=id) + disabled_paths = { + item.get_id() + for item in items + if isinstance(item, Component) and self._literal_disabled(item.get_config()) is True + } + for item in items: + path = item.get_id() + # Preserve definitions for a later explicit descendant request: a + # disabled parent's untouched children have not been compiled yet. + if any(path != parent and _within(path, parent) for parent in disabled_paths): + continue + if path not in self._items: + self.add_item(item) + + @classmethod + def _reference_boundaries(cls, value: Any) -> Any: + if isinstance(value, _Boundary): + return "@" + value.path + if isinstance(value, dict): + return {key: cls._reference_boundaries(item) for key, item in value.items()} + if isinstance(value, list): + return [cls._reference_boundaries(item) for item in value] + return value + + def _resolve_imports(self, resolved: dict[str, Any], eval_expr: bool) -> None: + for id, item in self._items.items(): + if any(_within(id, parent) for parent in self._scope._blocked) or id in self._scope._bindings: + continue + if id not in self._scope._active_paths and any(_within(id, parent) for parent in self._scope._bindings): + continue + if ( + id not in resolved + and isinstance(item, Expression) + and item.is_import_statement(item.get_config()) + and self._import_visible(id) + ): + self._resolve_import_item(item, id, resolved, eval_expr) + + def _publish_expression_references(self, paths: set[str]) -> None: + # Returned Python values may retain their reference mapping (for example + # lambdas or generators). Reserve lexical paths without running them or + # inspecting arbitrary returned object graphs. + self._scope._reserved_paths.update(paths) + + def _substitute_refs(self, config: Any, id: str, refs: dict[str, Any]) -> Any: + if id in self._scope._bindings: + result = self._scope._bindings[id] + elif isinstance(config, (dict, list)): + result = type(config)() + for key, path, value in self.iter_subconfigs(id, config): + if path in self._scope._bindings: + resolved = self._scope._bindings[path] + elif Component.is_instantiable(value) or Expression.is_expression(value): + resolved = refs[path] + if Component.is_instantiable(value) and resolved is None: + continue + else: + resolved = self._substitute_refs(value, path, refs) + self._scope._consumed[path] = resolved + if isinstance(result, dict): + result[key] = resolved + else: + result.append(resolved) + else: + result = self.update_config_with_refs(config, id, refs) + self._scope._consumed[id] = result + return result + + def _resolve_one_item( + self, + id: str, + waiting_list: set[str] | None = None, + _depth: int = 0, + instantiate: bool = True, + eval_expr: bool = True, + default: Any = None, + ) -> Any: + cached = id in self._resolved + result = super()._resolve_one_item( + id, waiting_list=waiting_list, _depth=_depth, instantiate=instantiate, eval_expr=eval_expr, default=default + ) + item = self._items.get(id) + if not cached and id not in self._scope._bindings and isinstance(item, Component): + guard_path = _child(id, "_disabled_") + if "_disabled_" in item.get_config(): + self._scope._consumed[guard_path] = self._resolved.get(guard_path, item.get_config()["_disabled_"]) + if ( + not cached + and id not in self._scope._bindings + and isinstance(self._items.get(id), Component) + and result is not None + and not isinstance(result, Iterator) + ): + self._scope._components[id] = result + return result diff --git a/src/sparkwheel/items.py b/src/sparkwheel/items.py index 300feec..547a27c 100644 --- a/src/sparkwheel/items.py +++ b/src/sparkwheel/items.py @@ -5,9 +5,10 @@ from pprint import pformat from typing import Any -from .utils import CompInitMode, first, instantiate, optional_import, run_debug, run_eval +from .path_utils import replace_references, scan_references +from .utils import CompInitMode, instantiate, run_debug, run_eval from .utils.constants import EXPR_KEY -from .utils.exceptions import EvaluationError, InstantiationError, Location, TargetNotFoundError +from .utils.exceptions import BaseError, EvaluationError, InstantiationError, Location, TargetNotFoundError __all__ = ["Item", "Expression", "Component", "Instantiable"] @@ -145,7 +146,7 @@ def resolve_module_name(self): No automatic module discovery is performed. Returns: - str or callable: The module path or callable from _target_ + target (Any): The module path or callable from _target_. """ config = dict(self.get_config()) target = config.get("_target_") @@ -160,7 +161,7 @@ def resolve_args(self): Utility function used in `instantiate()` to resolve the arguments from current config content. Returns: - tuple: A tuple of (args, kwargs) where: + arguments (tuple[list[Any], dict[str, Any]]): A tuple of (args, kwargs) where: - args is a list of positional arguments from _args_ (or empty list if not present) - kwargs is a dict of keyword arguments (excluding special keys) """ @@ -314,25 +315,50 @@ def __init__( ) -> None: super().__init__(config=config, id=id, source_location=source_location) self.globals = globals if globals is not None else {} + self._imports_loaded = False + self._import_result: Any = None + + @staticmethod + def _syntax_tree(value: str) -> ast.Module: + # Parse reference-bearing code without resolving or evaluating its inputs. + substituted = replace_references(value, dict.fromkeys(scan_references(value))) + return ast.parse(substituted[1:].lstrip()) def _parse_import_string(self, import_string: str) -> Any | None: - """parse single import statement such as "from pathlib import Path" """ - node = first(ast.iter_child_nodes(ast.parse(import_string))) - if not isinstance(node, (ast.Import, ast.ImportFrom)): - return None - if len(node.names) < 1: + """Execute an import prefix, binding every alias with native Python rules.""" + tree = self._syntax_tree("$" + import_string) + imports: list[ast.stmt] = [] + for node in tree.body: + if isinstance(node, (ast.Import, ast.ImportFrom)): + imports.append(node) + else: + break + if not imports: return None - if len(node.names) > 1: - warnings.warn(f"ignoring multiple import alias '{import_string}'.", stacklevel=2) - name, asname = f"{node.names[0].name}", node.names[0].asname - asname = name if asname is None else f"{asname}" - if isinstance(node, ast.ImportFrom): - self.globals[asname], _ = optional_import(f"{node.module}", name=f"{name}") - return self.globals[asname] - if isinstance(node, ast.Import): - self.globals[asname], _ = optional_import(f"{name}") - return self.globals[asname] - return None # type: ignore[unreachable] + if len(tree.body) > len(imports) and (len(tree.body) != len(imports) + 1 or not isinstance(tree.body[-1], ast.Expr)): + raise SyntaxError("An import prefix must be followed by one final Python expression") + exec(compile(ast.Module(body=imports, type_ignores=[]), "", "exec"), self.globals) + first_import = imports[0] + assert isinstance(first_import, (ast.Import, ast.ImportFrom)) + alias = first_import.names[0] + name = alias.asname or (alias.name.split(".")[0] if isinstance(first_import, ast.Import) else alias.name) + return self.globals.get(name) + + def load_imports(self, *, force: bool = False) -> Any: + """Bind this expression's imports without evaluating its trailing value.""" + if force or not self._imports_loaded: + try: + self._import_result = self._parse_import_string(self.get_config()[len(self.prefix) :]) + except Exception as error: + raise EvaluationError( + f"Failed to load expression imports: '{self.get_config()[len(self.prefix) :]}'", + source_location=self.source_location, + ) from error + self._imports_loaded = True + return self._import_result + + def is_import_only(self) -> bool: + return all(isinstance(node, (ast.Import, ast.ImportFrom)) for node in self._syntax_tree(self.get_config()).body) def evaluate(self, globals: dict[str, Any] | None = None, locals: dict[str, Any] | None = None) -> str | Any | None: """Evaluate the expression and return the result. @@ -352,9 +378,23 @@ def evaluate(self, globals: dict[str, Any] | None = None, locals: dict[str, Any] value = self.get_config() if not Expression.is_expression(value): return None - optional_module = self._parse_import_string(value[len(self.prefix) :]) - if optional_module is not None: - return optional_module + try: + tree = self._syntax_tree(value) + if tree.body and isinstance(tree.body[0], (ast.Import, ast.ImportFrom)): + imported = self.load_imports(force=True) + if self.is_import_only(): + return imported + final = tree.body[-1] + assert isinstance(final, ast.Expr) + expression = ast.Expression(final.value) + else: + expression = ast.parse(value[len(self.prefix) :].lstrip(), mode="eval") + except EvaluationError: + raise + except Exception as error: + raise EvaluationError( + f"Failed to parse expression: '{value[len(self.prefix) :]}'", source_location=self.source_location + ) from error if not self.run_eval: return f"{value[len(self.prefix) :]}" globals_ = dict(self.globals) @@ -365,7 +405,10 @@ def evaluate(self, globals: dict[str, Any] | None = None, locals: dict[str, Any] globals_[k] = v if not run_debug: try: - return eval(value[len(self.prefix) :], globals_, locals) + return eval(compile(expression, "", "eval"), globals_, locals) + except BaseError: + # Lazy graph lookups retain structured dependency/schema errors. + raise except Exception as e: raise EvaluationError( f"Failed to evaluate expression: '{value[len(self.prefix) :]}'", @@ -399,8 +442,10 @@ def is_import_statement(cls, config: dict[str, Any] | list[Any] | str) -> bool: Args: config: input config content to check. """ - if not cls.is_expression(config): + if not cls.is_expression(config) or "import" not in config: return False - if "import" not in config: + try: + tree = cls._syntax_tree(str(config)) + except (SyntaxError, ValueError): return False - return isinstance(first(ast.iter_child_nodes(ast.parse(f"{config[len(cls.prefix) :]}"))), (ast.Import, ast.ImportFrom)) # type: ignore[index] + return bool(tree.body) and isinstance(tree.body[0], (ast.Import, ast.ImportFrom)) diff --git a/src/sparkwheel/loader.py b/src/sparkwheel/loader.py index 748f2f1..cc802e2 100644 --- a/src/sparkwheel/loader.py +++ b/src/sparkwheel/loader.py @@ -5,7 +5,7 @@ from pathlib import Path from typing import Any -import yaml # type: ignore[import-untyped] +import yaml from .locations import LocationRegistry from .path_utils import is_yaml_file @@ -43,40 +43,29 @@ def construct_mapping(self, node, deep=False): ) self.registry.register(current_id, location) - # For non-deep construction, we construct children manually to track paths - if not deep: - mapping = {} - for key_node, value_node in node.value: - # Construct key - key = self.construct_object(key_node, deep=False) - - # Push key onto path stack before constructing value - self.id_path_stack.append(str(key)) - - # Register source location for this specific key - # This allows us to track where each key was defined - key_id = ID_SEP_KEY.join(self.id_path_stack) if self.id_path_stack else "" + # Check authored duplicates before flattening YAML merge keys. The same + # path-aware construction applies at every depth, including list entries. + self.flatten_mapping(node) + mapping = {} + for key_node, value_node in node.value: + key = self.construct_object(key_node, deep=True) + self.id_path_stack.append(str(key)) + try: + key_id = ID_SEP_KEY.join(self.id_path_stack) if key_node.start_mark: - key_location = Location( - filepath=self.filepath, - line=key_node.start_mark.line + 1, - column=key_node.start_mark.column + 1, - id=key_id, + self.registry.register( + key_id, + Location( + filepath=self.filepath, + line=key_node.start_mark.line + 1, + column=key_node.start_mark.column + 1, + id=key_id, + ), ) - self.registry.register(key_id, key_location) - - # Construct value with updated path - value = self.construct_object(value_node, deep=True) - - # Pop key from path stack + mapping[key] = self.construct_object(value_node, deep=True) + finally: self.id_path_stack.pop() - - mapping[key] = value - - return mapping - else: - # Use parent's deep construction - return super().construct_mapping(node, deep=True) + return mapping def construct_sequence(self, node, deep=False): """Override to track source locations for list nodes.""" @@ -92,25 +81,14 @@ def construct_sequence(self, node, deep=False): ) self.registry.register(current_id, location) - # For non-deep construction, construct children manually to track paths - if not deep: - sequence = [] - for idx, child_node in enumerate(node.value): - # Push index onto path stack - self.id_path_stack.append(str(idx)) - - # Construct child with updated path - value = self.construct_object(child_node, deep=True) - - # Pop index from path stack + sequence = [] + for idx, child_node in enumerate(node.value): + self.id_path_stack.append(str(idx)) + try: + sequence.append(self.construct_object(child_node, deep=True)) + finally: self.id_path_stack.pop() - - sequence.append(value) - - return sequence - else: - # Use parent's deep construction - return super().construct_sequence(node, deep=True) + return sequence class Loader: @@ -179,7 +157,7 @@ def _load_yaml_with_metadata(self, stream, filepath: str, registry: LocationRegi """Load YAML and populate metadata registry during construction. Args: - stream: File stream to load from + stream (Any): File stream to load from filepath: Path string for error messages registry: LocationRegistry to populate diff --git a/src/sparkwheel/operators.py b/src/sparkwheel/operators.py index e2e9313..5681e2b 100644 --- a/src/sparkwheel/operators.py +++ b/src/sparkwheel/operators.py @@ -1,6 +1,5 @@ """Configuration merging with composition-by-default and operators (=, ~).""" -from copy import deepcopy from dataclasses import dataclass from typing import TYPE_CHECKING, Any @@ -14,6 +13,15 @@ __all__ = ["apply_operators", "validate_operators", "_validate_delete_operator", "MergeContext"] +def _copy_containers(value: Any) -> Any: + """Copy source containers while preserving caller-owned opaque inputs.""" + if isinstance(value, dict): + return {key: _copy_containers(item) for key, item in value.items()} + if isinstance(value, list): + return [_copy_containers(item) for item in value] + return value + + @dataclass class MergeContext: """Context for configuration merging operations. @@ -183,6 +191,15 @@ def validate_operators(config: dict[str, Any], parent_key: str = "") -> None: validate_operators(value, full_key) +def _normalize_update_value(value: Any, context: MergeContext) -> Any: + """Apply operators in newly introduced containers without touching old data.""" + if isinstance(value, dict): + return apply_operators({}, value, context=context) + if isinstance(value, list): + return [_normalize_update_value(item, context.child_path(str(index))) for index, item in enumerate(value)] + return value + + def apply_operators( base: dict[str, Any], override: dict[str, Any], @@ -249,19 +266,19 @@ def apply_operators( context = MergeContext() if not isinstance(base, dict) or not isinstance(override, dict): - return deepcopy(override) # type: ignore[unreachable] + return _copy_containers(override) # type: ignore[unreachable] - result = deepcopy(base) + result: dict[str, Any] = _copy_containers(base) for key, value in override.items(): if not isinstance(key, str): - result[key] = deepcopy(value) # type: ignore[unreachable] + result[key] = _normalize_update_value(value, context.child_path(str(key))) # type: ignore[unreachable] continue # Process replace operator (=key) if key.startswith(REPLACE_KEY): actual_key = key[1:] - result[actual_key] = deepcopy(value) + result[actual_key] = _normalize_update_value(value, context.child_path(actual_key)) continue # Process remove operator (~key) @@ -353,13 +370,13 @@ def apply_operators( # For lists: EXTEND (composition) if isinstance(base_val, list) and isinstance(value, list): - result[key] = base_val + value + result[key] = base_val + _normalize_update_value(value, context.child_path(key)) continue # For scalars: REPLACE # For type mismatches: REPLACE # Set/replace (for new keys or non-matching types) - result[key] = deepcopy(value) + result[key] = _normalize_update_value(value, context.child_path(key)) return result diff --git a/src/sparkwheel/path_utils.py b/src/sparkwheel/path_utils.py index ad07245..6641a93 100644 --- a/src/sparkwheel/path_utils.py +++ b/src/sparkwheel/path_utils.py @@ -5,7 +5,11 @@ and file path extraction from combined strings like "config.yaml::model::lr". """ +import io +import keyword import re +import token +import tokenize from typing import Any from .utils.constants import ID_SEP_KEY, RAW_REF_KEY, RESOLVED_REF_KEY @@ -23,6 +27,63 @@ ] +def _expression_matches(text: str, pattern: re.Pattern[str]) -> list[re.Match[str]]: + """Find reference syntax outside Python strings/comments and binary @. + + Tokenization is lexical only: no code is executed. Whole f-strings are kept + opaque on Python 3.10+; reference interpolation inside f-strings is not syntax. + """ + code = text[1:] + offsets = [0] + for line in code.splitlines(keepends=True): + offsets.append(offsets[-1] + len(line)) + + def offset(position: tuple[int, int]) -> int: + row, column = position + return 1 + offsets[min(row - 1, len(offsets) - 1)] + column + + protected: list[tuple[int, int]] = [] + binary_at: set[int] = set() + previous: tokenize.TokenInfo | None = None + fstring_start: int | None = None + fstring_depth = 0 + fstart, fend = getattr(token, "FSTRING_START", -1), getattr(token, "FSTRING_END", -1) + try: + for part in tokenize.generate_tokens(io.StringIO(code).readline): + if part.type == fstart: + if fstring_depth == 0: + fstring_start = offset(part.start) + fstring_depth += 1 + elif part.type == fend: + fstring_depth -= 1 + if fstring_depth == 0 and fstring_start is not None: + protected.append((fstring_start, offset(part.end))) + if part.type in (token.STRING, token.COMMENT): + protected.append((offset(part.start), offset(part.end))) + if part.string == "@" and previous is not None: + ends_value = previous.type in (token.NAME, token.NUMBER, token.STRING, fend) or previous.string in ( + ")", + "]", + "}", + ) + prefix_keyword = keyword.iskeyword(previous.string) and previous.string not in {"True", "False", "None"} + if ends_value and not prefix_keyword: + binary_at.add(offset(part.start)) + if part.type == token.NEWLINE: + previous = None + elif part.type not in (tokenize.NL, token.INDENT, token.DEDENT, token.COMMENT, token.ENDMARKER): + previous = part + except (tokenize.TokenError, IndentationError): + # Evaluation supplies the actual syntax error; retain lexical protection + # collected before the malformed token instead of executing any parser. + pass + return [ + match + for match in pattern.finditer(text) + if match.start() not in binary_at and not any(start <= match.start() < end for start, end in protected) + ] + + # ============================================================================ # YAML File Detection # ============================================================================ @@ -208,7 +269,9 @@ def find_absolute_references(cls, text: str) -> list[str]: if not (is_expr or is_pure_ref): return [] - return cls.ABSOLUTE_REFERENCE.findall(text) + if is_pure_ref: + return [text[1:]] if "@" not in text[1:] else [] + return [match.group(1) for match in _expression_matches(text, cls.ABSOLUTE_REFERENCE)] # ============================================================================ @@ -390,45 +453,29 @@ def resolve_relative_ids(current_id: str, value: str) -> str: Raises: ValueError: If relative reference goes beyond root """ - # Find all relative reference patterns using centralized regex - patterns = PathPatterns.find_relative_references(value) - - # Sort by length (longest first) to avoid partial replacements - # e.g., replace "@::::" before "@::" so we don't double-process - patterns = sorted(set(patterns), key=len, reverse=True) - + if not value.startswith(("$", "@", "%")): + return value + matches = ( + _expression_matches(value, PathPatterns.RELATIVE_REFERENCE) + if value.startswith("$") + else list(PathPatterns.RELATIVE_REFERENCE.finditer(value)) + ) current_parts = current_id.split(ID_SEP_KEY) if current_id else [] - - for pattern in patterns: - # Determine symbol (@ for resolved reference, % for raw reference) + for match in reversed(matches): + pattern = match.group(0) symbol = pattern[0] - - # Count :: pairs to determine how many levels to go up - # @:: = 1 level up, @:::: = 2 levels up levels_up = pattern[1:].count(ID_SEP_KEY) - - # Validate we don't go too far up the tree if levels_up > len(current_parts): raise ValueError( f"Relative reference '{pattern}' in '{value}' attempts to go " f"{levels_up} levels up, but current path '{current_id}' only " f"has {len(current_parts)} levels" ) - - # Calculate the absolute path - if levels_up == len(current_parts): - # Going to root level - absolute = symbol - else: - # Going to ancestor at specific level - ancestor_parts = current_parts[:-levels_up] if levels_up > 0 else current_parts - absolute = symbol + ID_SEP_KEY.join(ancestor_parts) - if ancestor_parts: # Add trailing separator if not at root - absolute += ID_SEP_KEY - - # Replace pattern in value - value = value.replace(pattern, absolute) - + ancestor_parts = current_parts[:-levels_up] + absolute = symbol + ID_SEP_KEY.join(ancestor_parts) + if ancestor_parts: + absolute += ID_SEP_KEY + value = value[: match.start()] + absolute + value[match.end() :] return value @@ -519,7 +566,10 @@ def replace_match(match): ref_id = match.group(1) if ref_id not in resolved_refs: raise KeyError(f"Reference '@{ref_id}' not found in resolved references") - return f"{local_var_name}['{ref_id}']" + separation = " " if match.start() > 0 and (text[match.start() - 1].isalnum() or text[match.start() - 1] == "_") else "" + return f"{separation}{local_var_name}['{ref_id}']" - result = PathPatterns.ABSOLUTE_REFERENCE.sub(replace_match, text) + result = text + for match in reversed(_expression_matches(text, PathPatterns.ABSOLUTE_REFERENCE)): + result = result[: match.start()] + replace_match(match) + result[match.end() :] return result diff --git a/src/sparkwheel/preprocessor.py b/src/sparkwheel/preprocessor.py index 0e4262a..d0420a9 100644 --- a/src/sparkwheel/preprocessor.py +++ b/src/sparkwheel/preprocessor.py @@ -67,7 +67,7 @@ def __init__(self, loader, globals: dict[str, Any] | None = None): # type: igno """Initialize preprocessor. Args: - loader: Loader instance for loading external raw reference files + loader (Loader): Loader instance for loading external raw reference files globals: Global context (unused here, kept for API consistency) """ self.loader = loader diff --git a/src/sparkwheel/resolver.py b/src/sparkwheel/resolver.py index 02328cc..407a45f 100644 --- a/src/sparkwheel/resolver.py +++ b/src/sparkwheel/resolver.py @@ -2,9 +2,11 @@ import warnings from collections.abc import Iterator +from copy import copy from typing import Any from .items import Component, Expression, Item +from .operators import _copy_containers from .path_utils import normalize_id, replace_references, scan_references from .utils import allow_missing_reference, look_up_option from .utils.constants import ID_SEP_KEY, RESOLVED_REF_KEY @@ -13,6 +15,33 @@ __all__ = ["Resolver"] +class _ExpressionReferences: + """Resolve Python-selected references within one immutable build generation.""" + + def __init__(self, resolver: "Resolver", path: str, waiting: set[str], depth: int, mode: tuple[bool, bool]): + self._resolver = resolver + self._path = path + self._waiting = waiting | {path} + self._depth = depth + self._mode = mode + self._cache = resolver._resolved_by_mode[mode] + + def __getitem__(self, path: str) -> Any: + if path not in self._cache and path in self._waiting: + item = self._resolver._items.get(self._path) + raise CircularReferenceError( + f"Circular reference detected: '{path}' references back to '{self._path}'", + source_location=item.source_location if item is not None else None, + ) + return self._resolver._resolve_one_item( + path, + waiting_list=set(self._waiting), + _depth=self._depth + 1, + instantiate=self._mode[0], + eval_expr=self._mode[1], + ) + + class Resolver: """Resolve references between Items. @@ -53,7 +82,13 @@ def __init__(self, items: list[Item] | None = None): items: Optional list of Items to add during initialization """ self._items: dict[str, Item] = {} + # Keep the default cache available for existing internal integrations. self._resolved: dict[str, Any] = {} + self._resolved_by_mode: dict[tuple[bool, bool], dict[str, Any]] = {(True, True): self._resolved} + self._guard_states: dict[tuple[bool, bool], dict[str, bool]] = {} + self._resolution_mode = (True, True) + self._current_path = "" + self._active_resolutions: dict[tuple[bool, bool], set[str]] = {} if items: for item in items: @@ -63,10 +98,13 @@ def reset(self) -> None: """Clear all items and resolved content.""" self._items = {} self._resolved = {} + self._resolved_by_mode = {(True, True): self._resolved} + self._guard_states = {} + self._active_resolutions = {} def is_resolved(self) -> bool: """Check if any items have been resolved.""" - return bool(self._resolved) + return any(self._resolved_by_mode.values()) def add_item(self, item: Item) -> None: """Add a Item to resolve. @@ -96,7 +134,7 @@ def add_items(self, items: list[Item]) -> None: self.add_item(item) def get_item(self, id: str, resolve: bool = False, **kwargs: Any) -> Item | None: - """Get Item by id, optionally resolved. + """Get the registered Item by id, optionally populating the resolution cache. Args: id: ID of the config item @@ -104,10 +142,10 @@ def get_item(self, id: str, resolve: bool = False, **kwargs: Any) -> Item | None **kwargs: Additional arguments for resolution Returns: - Item if found, None otherwise (or resolved value if resolve=True) + The original Item if found, None otherwise. Use resolve() for its resolved value. """ id = self.normalize_id(id) - if resolve and id not in self._resolved: + if resolve: self._resolve_one_item(id=id, **kwargs) return self._items.get(id) @@ -121,7 +159,9 @@ def resolve( """Resolve a config item and return the result. Resolves all references, instantiates components (if requested), and - evaluates expressions (if requested). Results are cached for efficiency. + evaluates expressions (if requested). Each (instantiate, eval_expr) pair + has a separate cache. Repeated calls in the same mode share resolved objects, + including after calls in a different mode. Args: id: ID of item to resolve (empty string for root) @@ -146,6 +186,33 @@ def _resolve_one_item( instantiate: bool = True, eval_expr: bool = True, default: Any = None, + ) -> Any: + """Track active graph construction, including escaped Python callbacks.""" + path = self.normalize_id(id) + mode = (instantiate, eval_expr) + if path in self._resolved_by_mode.get(mode, {}): + return self._resolve_item(path, waiting_list, _depth, instantiate, eval_expr, default) + active = self._active_resolutions.setdefault(mode, set()) + if path in active: + item = self._items.get(path) + raise CircularReferenceError( + f"Circular reference detected while constructing '{path}'", + source_location=item.source_location if item is not None else None, + ) + active.add(path) + try: + return self._resolve_item(path, waiting_list, _depth, instantiate, eval_expr, default) + finally: + active.discard(path) + + def _resolve_item( + self, + id: str, + waiting_list: set[str] | None = None, + _depth: int = 0, + instantiate: bool = True, + eval_expr: bool = True, + default: Any = None, ) -> Any: """Internal recursive resolution implementation. @@ -174,9 +241,14 @@ def _resolve_one_item( id = self.normalize_id(id) + self._resolution_mode = (instantiate, eval_expr) + resolved = self._resolved_by_mode.setdefault(self._resolution_mode, {}) + # Return cached result if available - if id in self._resolved: - return self._resolved[id] + if id in resolved: + return resolved[id] + + self._prepare_item(id) # Look up the item try: @@ -221,13 +293,55 @@ def _resolve_one_item( waiting_list = set() waiting_list.add(id) - # First, resolve any import expressions (they need to run first) - for t, v in self._items.items(): - if t not in self._resolved and isinstance(v, Expression) and v.is_import_statement(v.get_config()): - self._resolved[t] = v.evaluate() if eval_expr else v - - # Find all references in this item's config - refs = self.find_refs_in_config(config=item_config, id=id) + if isinstance(item, Component): + disabled = self._literal_disabled(item_config) + guard = item_config.get("_disabled_", False) + if disabled is None: + guard_id = f"{id}::_disabled_" if id else "_disabled_" + self._prepare_item(guard_id) + guard = self._resolve_one_item( + guard_id, waiting_list, _depth + 1, instantiate=instantiate, eval_expr=eval_expr + ) + waiting_list.discard(guard_id) + if not eval_expr and isinstance(guard, Item): + # The guard cannot be decided in this mode. Preserve its + # wrapper without constructing or touching the payload. + result = copy(item) + config = _copy_containers(item_config) + config["_disabled_"] = guard + result.update_config(config) + resolved[id] = result + return result + disabled = Component({"_disabled_": guard}).is_disabled() + self._guard_states.setdefault(self._resolution_mode, {})[id] = disabled + if disabled: + if instantiate: + resolved[id] = None + else: + result = copy(item) + config = _copy_containers(item_config) + config["_disabled_"] = guard + result.update_config(config) + resolved[id] = result + return resolved[id] + + self._prepare_enabled_item(id, waiting_list, _depth, instantiate, eval_expr) + # Import declarations in inactive component payloads stay inactive. + self._current_path = id + self._resolve_imports(resolved, eval_expr) + + if isinstance(item, Expression) and eval_expr: + references = scan_references(item_config) + prepared = copy(item) + prepared.update_config(replace_references(item_config, dict.fromkeys(references), self._vars)) + values = _ExpressionReferences(self, id, waiting_list, _depth, (instantiate, eval_expr)) + expression_result = prepared.evaluate(globals={self._vars: values}) + self._publish_expression_references(set(references)) + resolved[id] = expression_result + return expression_result + + # A nested component owns its argument dependencies and guard boundary. + refs = self._dependency_refs(config=item_config, id=id) # Resolve dependencies first for dep_id in refs.keys(): @@ -239,7 +353,8 @@ def _resolve_one_item( ) # Resolve dependency if not already resolved - if dep_id not in self._resolved: + if dep_id not in resolved: + self._prepare_item(dep_id) try: look_up_option(dep_id, self._items, print_all_options=False) except ValueError as err: @@ -266,18 +381,106 @@ def _resolve_one_item( waiting_list.discard(dep_id) # All dependencies resolved, now resolve this item - new_config = self.update_config_with_refs(config=item_config, id=id, refs=self._resolved) + new_config = self._substitute_refs(config=item_config, id=id, refs=resolved) + # Replacing references is mode-specific. Keep the registered definition + # intact so another mode never inherits instantiated objects or wrappers. + # A shallow copy retains opaque Python inputs without cloning them; + # update_config_with_refs has already rebuilt the config containers. + item = copy(item) item.update_config(config=new_config) # Generate final resolved value based on item type if isinstance(item, Component): - self._resolved[id] = item.instantiate() if instantiate else item + resolved[id] = item.instantiate() if instantiate else item elif isinstance(item, Expression): - self._resolved[id] = item.evaluate(globals={f"{self._vars}": self._resolved}) if eval_expr else item + resolved[id] = item.evaluate(globals={f"{self._vars}": resolved}) if eval_expr else item else: - self._resolved[id] = new_config + resolved[id] = new_config + + return resolved[id] + + def _publish_expression_references(self, paths: set[str]) -> None: + """Extension point for conservative binding ownership after publication.""" + + def _prepare_item(self, id: str) -> None: + """Extension point for scoped, on-demand definitions.""" + + def _prepare_enabled_item(self, id: str, waiting_list: set[str], depth: int, instantiate: bool, eval_expr: bool) -> None: + """Extension point after a component's guard permits argument work.""" + + def _resolve_imports(self, resolved: dict[str, Any], eval_expr: bool) -> None: + for id, item in self._items.items(): + if ( + id not in resolved + and isinstance(item, Expression) + and item.is_import_statement(item.get_config()) + and self._import_visible(id) + ): + self._resolve_import_item(item, id, resolved, eval_expr) + + @staticmethod + def _resolve_import_item(item: Expression, id: str, resolved: dict[str, Any], eval_expr: bool) -> None: + if not eval_expr: + # Each requested wrapper still needs ordinary dependency preparation; + # global declaration discovery must not cache untouched definitions. + return + imported = item.load_imports() + if item.is_import_only(): + resolved[id] = imported + + @staticmethod + def _literal_disabled(config: Any) -> bool | None: + """Return a literal guard state without calling opaque Python objects.""" + if not Component.is_instantiable(config): + return False + value = config.get("_disabled_", False) + if type(value) is str: + return None if value.startswith(("@", "$", "%")) else value.lower().strip() == "true" + if Component.is_instantiable(value): + return None + if value is None or type(value) in (bool, int, float, complex, tuple, list, dict): + return bool(value) + return None - return self._resolved[id] + @classmethod + def _dependency_refs(cls, config: Any, id: str) -> dict[str, int]: + """Discover immediate graph dependencies, preserving component guards.""" + if Component.is_instantiable(config) and cls._literal_disabled(config) is True: + return {} + if isinstance(config, str): + return cls.match_refs_pattern(config) + refs: dict[str, int] = {} + if isinstance(config, (dict, list)): + for _, path, value in cls.iter_subconfigs(id, config): + if Component.is_instantiable(value) or Expression.is_expression(value): + refs[path] = refs.get(path, 0) + 1 + else: + for dependency, count in cls._dependency_refs(value, path).items(): + refs[dependency] = refs.get(dependency, 0) + count + return refs + + def _import_visible(self, path: str) -> bool: + """Only collect imports from enabled components or the requested subtree.""" + parts = path.split(self.sep) + states = self._guard_states.get(self._resolution_mode, {}) + for length in range(len(parts)): + parent = self.sep.join(parts[:length]) + item = self._items.get(parent) + if not isinstance(item, Component): + continue + state = states.get(parent, self._literal_disabled(item.get_config())) + if state is False: + continue + # A requested descendant is independent of its disabled parent; + # imports must still belong to that explicitly requested subtree. + inside_parent = self._current_path != parent and (not parent or self._current_path.startswith(parent + self.sep)) + inside_request = path == self._current_path or path.startswith(self._current_path + self.sep) + if not (inside_parent and inside_request): + return False + return True + + def _substitute_refs(self, config: Any, id: str, refs: dict[str, Any]) -> Any: + return self.update_config_with_refs(config=config, id=id, refs=refs) @classmethod def normalize_id(cls, id: str | int) -> str: diff --git a/src/sparkwheel/schema.py b/src/sparkwheel/schema.py index 43a7ee5..de82d13 100644 --- a/src/sparkwheel/schema.py +++ b/src/sparkwheel/schema.py @@ -37,7 +37,7 @@ class ModelConfig: import dataclasses import types -from typing import Any, Union, get_args, get_origin +from typing import Any, Union, get_args, get_origin, get_type_hints from .utils.exceptions import BaseError, Location @@ -106,6 +106,25 @@ def check_range(self): return func +def _schema_hints(schema: type) -> dict[str, Any]: + try: + return get_type_hints(schema) + except (NameError, TypeError) as error: + raise TypeError(f"Cannot resolve annotations for schema {schema.__name__}: {error}") from error + + +def _has_deferred(value: Any) -> bool: + if isinstance(value, _MissingSentinel): + return True + if isinstance(value, str): + return value.startswith(("@", "$", "%")) + if isinstance(value, dict): + return any(_has_deferred(item) for item in value.values()) + if isinstance(value, list): + return any(_has_deferred(item) for item in value) + return False + + def _get_validators(schema_type: type) -> list[Any]: """Get all validator methods from a dataclass.""" validators = [] @@ -126,6 +145,7 @@ def _run_validators( schema: type, field_path: str = "", metadata: Any = None, + allow_unresolved: bool = True, ) -> None: """Run all @validator methods on a dataclass. @@ -142,21 +162,19 @@ def _run_validators( if not validators: return - # Skip validation for configs with references/expressions/macros - # They'll be validated after resolution - for value in config.values(): - if isinstance(value, str) and value.startswith(("@", "$", "%")): - # Has unresolved references - skip custom validation - return - - # Create instance to call validators on - try: - instance = schema(**config) - except Exception: - # Can't create instance - skip validation + # Framework metadata and lenient extra fields are not dataclass init args. + values = {field.name: config[field.name] for field in dataclasses.fields(schema) if field.init and field.name in config} + if allow_unresolved and _has_deferred(values): return - source_loc = _get_source_location(metadata, field_path) if metadata else None + try: + instance = schema(**values) + except Exception as error: + raise _ValidatorExecutionError( + f"Cannot construct validation instance for {schema.__name__}: {type(error).__name__}: {error}", + field_path=field_path, + source_location=source_loc, + ) from error for validator_method in validators: try: @@ -168,7 +186,7 @@ def _run_validators( source_location=source_loc, ) from e except Exception as e: - raise ValidationError( + raise _ValidatorExecutionError( f"Validator '{validator_method.__name__}' raised {type(e).__name__}: {e}", field_path=field_path, source_location=source_loc, @@ -222,6 +240,10 @@ def __init__( super().__init__(full_message, source_location=source_location) +class _ValidatorExecutionError(ValidationError): + """Unexpected user-code errors must not be treated as a union mismatch.""" + + def validate( config: dict[str, Any], schema: type, @@ -229,6 +251,8 @@ def validate( metadata: Any = None, allow_missing: bool = False, strict: bool = True, + allow_unresolved: bool = True, + _run_custom: bool = True, ) -> None: """Validate configuration against a dataclass schema. @@ -278,6 +302,7 @@ class AppConfig: # Get all fields from the dataclass schema_fields = {f.name: f for f in dataclasses.fields(schema)} + hints = _schema_hints(schema) # Check for required fields for field_name, field_info in schema_fields.items(): @@ -293,17 +318,20 @@ class AppConfig: raise ValidationError( f"Missing required field '{field_name}'", field_path=current_path, - expected_type=field_info.type, # type: ignore[arg-type] + expected_type=hints[field_name], source_location=source_loc, ) # Validate the field value _validate_field( config[field_name], - field_info.type, # type: ignore[arg-type] + hints[field_name], current_path, metadata, allow_missing=allow_missing, + strict=strict, + allow_unresolved=allow_unresolved, + _run_custom=_run_custom, ) # Check for unexpected fields - only if strict mode @@ -324,7 +352,8 @@ class AppConfig: ) # Run custom validators - _run_validators(config, schema, field_path, metadata) + if _run_custom: + _run_validators(config, schema, field_path, metadata, allow_unresolved=allow_unresolved) def _find_discriminator(union_types: tuple[Any, ...]) -> tuple[bool, str | None]: @@ -344,7 +373,7 @@ def _find_discriminator(union_types: tuple[Any, ...]) -> tuple[bool, str | None] from typing import Literal # Filter to dataclasses only - dataclass_types = [t for t in union_types if dataclasses.is_dataclass(t)] + dataclass_types = [t for t in union_types if isinstance(t, type) and dataclasses.is_dataclass(t)] if len(dataclass_types) < 2: return False, None @@ -352,10 +381,10 @@ def _find_discriminator(union_types: tuple[Any, ...]) -> tuple[bool, str | None] all_fields: dict[str, list[Any]] = {} for dc_type in dataclass_types: for f in dataclasses.fields(dc_type): - if get_origin(f.type) is Literal: + if get_origin(_schema_hints(dc_type)[f.name]) is Literal: if f.name not in all_fields: all_fields[f.name] = [] - literal_values = get_args(f.type) + literal_values = get_args(_schema_hints(dc_type)[f.name]) all_fields[f.name].append({"type": dc_type, "values": literal_values}) # Find a field present in all types with unique values @@ -387,6 +416,10 @@ def _validate_discriminated_union( discriminator_field: str, field_path: str, metadata: Any = None, + allow_missing: bool = False, + strict: bool = True, + allow_unresolved: bool = True, + _run_custom: bool = True, ) -> None: """Validate a discriminated union by checking the discriminator. @@ -412,7 +445,7 @@ def _validate_discriminated_union( # Check discriminator field exists if discriminator_field not in value: - dataclass_types = [t for t in union_types if dataclasses.is_dataclass(t)] + dataclass_types = [t for t in union_types if isinstance(t, type) and dataclasses.is_dataclass(t)] type_names = ", ".join(t.__name__ if isinstance(t, type) else type(t).__name__ for t in dataclass_types) raise ValidationError( f"Missing discriminator field '{discriminator_field}' (required for union of {type_names})", @@ -423,14 +456,18 @@ def _validate_discriminated_union( discriminator_value = value[discriminator_field] + if allow_unresolved and isinstance(discriminator_value, str) and discriminator_value.startswith(("@", "$", "%")): + # The selected branch is not known until resolution. + return + # Find matching type - dataclass_types = [t for t in union_types if dataclasses.is_dataclass(t)] + dataclass_types = [t for t in union_types if isinstance(t, type) and dataclasses.is_dataclass(t)] matching_type = None for dc_type in dataclass_types: for f in dataclasses.fields(dc_type): if f.name == discriminator_field: - literal_values = get_args(f.type) + literal_values = get_args(_schema_hints(dc_type)[f.name]) if discriminator_value in literal_values: matching_type = dc_type break @@ -443,7 +480,7 @@ def _validate_discriminated_union( for dc_type in dataclass_types: for f in dataclasses.fields(dc_type): if f.name == discriminator_field: - literal_values = get_args(f.type) + literal_values = get_args(_schema_hints(dc_type)[f.name]) for val in literal_values: type_name = dc_type.__name__ if isinstance(dc_type, type) else type(dc_type).__name__ valid_values.append(f"'{val}' ({type_name})") @@ -458,202 +495,110 @@ def _validate_discriminated_union( # Validate against the selected type assert isinstance(matching_type, type) - validate(value, matching_type, field_path, metadata, allow_missing=False, strict=True) + validate( + value, + matching_type, + field_path, + metadata, + allow_missing=allow_missing, + strict=strict, + allow_unresolved=allow_unresolved, + _run_custom=_run_custom, + ) def _validate_field( value: Any, - expected_type: type, + expected_type: Any, field_path: str, metadata: Any = None, allow_missing: bool = False, + strict: bool = True, + allow_unresolved: bool = True, + _run_custom: bool = True, ) -> None: - """Validate a single field value against its expected type. - - Args: - value: The value to validate - expected_type: The expected type (may be generic like list[int]) - field_path: Dot-separated path to this field - metadata: Optional metadata registry for source locations - allow_missing: If True, allow MISSING sentinel values for partial configs - - Raises: - ValidationError: If validation fails - """ + """Validate a value with the same policy at every nested boundary.""" source_loc = _get_source_location(metadata, field_path) if metadata else None - - # Handle MISSING values + policy = dict(allow_missing=allow_missing, strict=strict, allow_unresolved=allow_unresolved, _run_custom=_run_custom) if isinstance(value, _MissingSentinel): if allow_missing: - return # OK for partial configs - else: - raise ValidationError( - "Field has MISSING value but MISSING not allowed", - field_path=field_path, - expected_type=expected_type, - actual_value=value, - source_location=source_loc, - ) - - # Handle None values - origin = get_origin(expected_type) - args = get_args(expected_type) - - # Handle Optional[T] (which is Union[T, None]) + return + raise ValidationError( + "Field has MISSING value but MISSING not allowed", + field_path=field_path, + expected_type=expected_type, + actual_value=value, + source_location=source_loc, + ) + if allow_unresolved and isinstance(value, str) and value.startswith(("@", "$", "%")): + return + if expected_type is Any: + return + origin, args = get_origin(expected_type), get_args(expected_type) if _is_union_type(origin): - # Check for discriminated union first - has_discriminator, discriminator_field = _find_discriminator(args) - if has_discriminator and discriminator_field: - _validate_discriminated_union(value, args, discriminator_field, field_path, metadata) + if value is None and type(None) in args: return - - # Check if None is allowed - if type(None) in args: - if value is None: - return # None is valid - # Remove None from the union and validate against remaining types - non_none_types = [t for t in args if t is not type(None)] - if len(non_none_types) == 1: - # Simple Optional[T] case - recursively validate with the single type - _validate_field(value, non_none_types[0], field_path, metadata, allow_missing) + branches = tuple(branch for branch in args if branch is not type(None)) if type(None) in args else args + has_discriminator, discriminator = _find_discriminator(branches) + if ( + has_discriminator + and discriminator + and all(isinstance(branch, type) and dataclasses.is_dataclass(branch) for branch in branches) + ): + _validate_discriminated_union(value, branches, discriminator, field_path, metadata, **policy) + return + errors = [] + for branch in branches: + try: + _validate_field(value, branch, field_path, metadata, **policy) return - else: - # Union with multiple non-None types - try each and collect errors - errors = [] - for union_type in non_none_types: - try: - _validate_field(value, union_type, field_path, metadata, allow_missing) - return # Validation succeeded - except ValidationError as e: - type_name = getattr(union_type, "__name__", str(union_type)) - # Extract just the error message without field path prefix - error_msg = str(e).split("\n")[0] - if f"Validation error at '{field_path}': " in error_msg: - error_msg = error_msg.replace(f"Validation error at '{field_path}': ", "") - errors.append(f" Tried {type_name}: {error_msg}") - - # All failed - build comprehensive error message - union_str = _format_union_type(tuple(non_none_types)) - error_details = "\n".join(errors) - raise ValidationError( - f"Value doesn't match any type in {union_str}\n{error_details}", - field_path=field_path, - expected_type=expected_type, - actual_value=value, - source_location=source_loc, - ) - else: - # Non-Optional Union - try each type and collect errors - errors = [] - for union_type in args: - try: - _validate_field(value, union_type, field_path, metadata, allow_missing) - return # Validation succeeded - except ValidationError as e: - type_name = getattr(union_type, "__name__", str(union_type)) - # Extract just the error message without field path prefix - error_msg = str(e).split("\n")[0] - if f"Validation error at '{field_path}': " in error_msg: - error_msg = error_msg.replace(f"Validation error at '{field_path}': ", "") - errors.append(f" Tried {type_name}: {error_msg}") - - # All failed - build comprehensive error message - union_str = _format_union_type(args) - error_details = "\n".join(errors) - raise ValidationError( - f"Value doesn't match any type in {union_str}\n{error_details}", - field_path=field_path, - expected_type=expected_type, - actual_value=value, - source_location=source_loc, - ) - - # Handle references and expressions early - # Accept resolved references (@), raw references (%), and expressions ($) as strings - # since they'll be resolved/expanded later - we can't validate their type until resolution - if isinstance(value, str) and (value.startswith("@") or value.startswith("$") or value.startswith("%")): - return - - # Handle list[T] + except _ValidatorExecutionError: + raise + except ValidationError as error: + name = getattr(branch, "__name__", str(branch)) + message = str(error).split("\n")[0].removeprefix(f"Validation error at '{field_path}': ") + errors.append(f" Tried {name}: {message}") + raise ValidationError( + f"Value doesn't match any type in {_format_union_type(branches)}\n" + "\n".join(errors), + field_path=field_path, + expected_type=expected_type, + actual_value=value, + source_location=source_loc, + ) if origin is list: if not isinstance(value, list): raise ValidationError( - "Expected list", - field_path=field_path, - expected_type=list, - actual_value=value, - source_location=source_loc, + "Expected list", field_path=field_path, expected_type=list, actual_value=value, source_location=source_loc ) if args: - item_type = args[0] - # Skip validation for List[Any] - accept any item types - if item_type is not Any: - for i, item in enumerate(value): - _validate_field( - item, - item_type, - f"{field_path}[{i}]", - metadata, - allow_missing, - ) + for i, item in enumerate(value): + _validate_field(item, args[0], f"{field_path}[{i}]", metadata, **policy) return - - # Handle dict[K, V] if origin is dict: if not isinstance(value, dict): raise ValidationError( - "Expected dict", - field_path=field_path, - expected_type=dict, - actual_value=value, - source_location=source_loc, + "Expected dict", field_path=field_path, expected_type=dict, actual_value=value, source_location=source_loc ) - if args and len(args) == 2: - key_type, value_type = args - # For Dict[K, Any], only validate keys and allow arbitrary values - if value_type is Any: - for k in value.keys(): - if not isinstance(k, key_type): - raise ValidationError( - "Dict key has wrong type", - field_path=f"{field_path}[{k!r}]", - expected_type=key_type, - actual_value=k, - source_location=source_loc, - ) - return - # Otherwise validate both keys and values - for k, v in value.items(): - # Validate key type - if not isinstance(k, key_type): + if len(args) == 2: + for key, item in value.items(): + if args[0] is not Any and not isinstance(key, args[0]): raise ValidationError( "Dict key has wrong type", - field_path=f"{field_path}[{k!r}]", - expected_type=key_type, - actual_value=k, + field_path=f"{field_path}[{key!r}]", + expected_type=args[0], + actual_value=key, source_location=source_loc, ) - # Validate value type - _validate_field( - v, - value_type, - f"{field_path}[{k!r}]", - metadata, - allow_missing, - ) + _validate_field(item, args[1], f"{field_path}[{key!r}]", metadata, **policy) return - - # Handle nested dataclasses - if dataclasses.is_dataclass(expected_type): - validate(value, expected_type, field_path, metadata, allow_missing, strict=True) + if isinstance(expected_type, type) and dataclasses.is_dataclass(expected_type): + validate(value, expected_type, field_path, metadata, **policy) return - - # Handle Literal types from typing import Literal if origin is Literal: - if value not in args: - valid_values = ", ".join(repr(v) for v in args) + if not any(type(value) is type(candidate) and value == candidate for candidate in args): + valid_values = ", ".join(repr(candidate) for candidate in args) raise ValidationError( f"Value must be one of {valid_values}, got {value!r}", field_path=field_path, @@ -662,26 +607,129 @@ def _validate_field( source_location=source_loc, ) return - - # Handle Any type - accept any value - if expected_type == Any: + # Preserve integer compatibility for float annotations, without treating bool + # as a numeric parameter just because bool subclasses int in Python. + if expected_type is float and type(value) is int: return - - # Handle basic types (int, str, float, bool, etc.) - if not isinstance(value, expected_type): - # Special case: allow int for float - if expected_type is float and isinstance(value, int): - return - + numeric_bool = expected_type in (int, float) and type(value) is bool + if numeric_bool or not isinstance(value, expected_type): raise ValidationError( - "Type mismatch", - field_path=field_path, - expected_type=expected_type, - actual_value=value, - source_location=source_loc, + "Type mismatch", field_path=field_path, expected_type=expected_type, actual_value=value, source_location=source_loc ) +def _coerce_field( + value: Any, + expected_type: Any, + field_path: str = "", + metadata: Any = None, + *, + strict: bool = True, + allow_missing: bool = False, + allow_unresolved: bool = True, + _preserve: set[int] | None = None, +) -> Any: + """Conservative normalization; never execute targets, validators or defaults.""" + if _preserve and id(value) in _preserve: + return value + + options = dict(strict=strict, allow_missing=allow_missing, allow_unresolved=allow_unresolved) + if isinstance(value, _MissingSentinel) or expected_type is Any: + return value + if allow_unresolved and isinstance(value, str) and value.startswith(("@", "$", "%")): + return value + origin, args = get_origin(expected_type), get_args(expected_type) + if _is_union_type(origin): + _, discriminator = _find_discriminator(args) + if allow_unresolved and discriminator and isinstance(value, dict) and _has_deferred(value.get(discriminator)): + return value + # Preserve already-valid alternatives, especially a literal string in + # str|int. Type checks here never run user validators. + for branch in args: + try: + _validate_field(value, branch, field_path, metadata, _run_custom=False, **options) + except ValidationError: + continue + return _coerce_field(value, branch, field_path, metadata, **options, _preserve=_preserve) + candidates = [] + for branch in args: + try: + candidate = _coerce_field(value, branch, field_path, metadata, **options, _preserve=_preserve) + _validate_field(candidate, branch, field_path, metadata, _run_custom=False, **options) + candidates.append(candidate) + except ValidationError: + continue + if len(candidates) > 1: + raise ValidationError( + f"Ambiguous coercion to {_format_union_type(args)}; supply an explicitly typed value", + field_path=field_path, + actual_value=value, + source_location=_get_source_location(metadata, field_path), + ) + return candidates[0] if candidates else value + if isinstance(expected_type, type) and dataclasses.is_dataclass(expected_type) and isinstance(value, dict): + result = dict(value) + for name, annotation in _schema_hints(expected_type).items(): + if name in result: + path = f"{field_path}.{name}" if field_path else name + result[name] = _coerce_field(result[name], annotation, path, metadata, **options, _preserve=_preserve) + return result + if origin is list and isinstance(value, list) and args: + return [ + _coerce_field(item, args[0], f"{field_path}[{i}]", metadata, **options, _preserve=_preserve) + for i, item in enumerate(value) + ] + if origin is dict and isinstance(value, dict) and len(args) == 2: + # Keys retain their identity and spelling: coercing keys can collapse two + # distinct entries. Their declared types are validated normally. + return { + key: _coerce_field(item, args[1], f"{field_path}[{key!r}]", metadata, **options, _preserve=_preserve) + for key, item in value.items() + } + try: + if expected_type is int: + if type(value) is str: + return int(value) + if type(value) is float and value.is_integer(): + return int(value) + elif expected_type is float: + if type(value) is str: + return float(value) + if type(value) is int: + try: + converted = float(value) + except OverflowError as error: + raise ValidationError( + "Integer cannot be represented as float", + field_path=field_path, + actual_value=value, + source_location=_get_source_location(metadata, field_path), + ) from error + if int(converted) != value: + raise ValidationError( + "Integer cannot be represented exactly as float", + field_path=field_path, + actual_value=value, + source_location=_get_source_location(metadata, field_path), + ) + return converted + elif expected_type is bool: + if type(value) is str: + normalized = value.strip().lower() + if normalized in {"true", "yes", "1"}: + return True + if normalized in {"false", "no", "0"}: + return False + elif type(value) is int and value in (0, 1): + return bool(value) + elif expected_type is str and type(value) in (int, float, bool): + return str(value) + except (ValueError, OverflowError): + # Ordinary validation produces the source-aware type error. + pass + return value + + def _get_source_location(metadata: Any, field_path: str) -> Location | None: """Get source location from metadata registry. diff --git a/src/sparkwheel/schema_resolver.py b/src/sparkwheel/schema_resolver.py new file mode 100644 index 0000000..7e2786f --- /dev/null +++ b/src/sparkwheel/schema_resolver.py @@ -0,0 +1,205 @@ +"""Schema checks on resolved argument trees, before component construction.""" + +import dataclasses +from typing import TYPE_CHECKING, Any, Union, get_args, get_origin + +from .items import Component, Expression +from .resolver import Resolver +from .schema import ( + _coerce_field, + _find_discriminator, + _is_union_type, + _schema_hints, + _validate_field, +) +from .utils.exceptions import CircularReferenceError + +if TYPE_CHECKING: + from .config import Config + + +class _SchemaResolver(Resolver): + """Keep validation projections separate from live runtime objects. + + A component's projection is its resolved constructor-argument mapping. Its + runtime cache still contains the native object passed to other constructors. + Only requested nodes, their ordinary dependencies, and schema discriminator + dependencies needed to select a variant are resolved. + """ + + def __init__(self, config: "Config"): + super().__init__() + self._config = config + self._projections: dict[tuple[bool, bool], dict[str, Any]] = {} + self._mode = (True, True) + + def reset(self) -> None: + super().reset() + self._projections = {} + + def _annotation(self, path: str) -> Any: + annotation = self._config._schema + parent = "" + for key in path.split("::") if path else []: + item = self._items.get(parent) + source = item.get_config() if item is not None else None + if isinstance(source, dict): + cache = self._resolved_by_mode[self._mode] + source = {name: cache.get(f"{parent}::{name}" if parent else name, value) for name, value in source.items()} + annotation = self._child_annotation(annotation, key, source) + parent = f"{parent}::{key}" if parent else key + return annotation + + def _child_annotation(self, annotation: Any, key: str, source: Any) -> Any: + origin, args = get_origin(annotation), get_args(annotation) + if _is_union_type(origin): + branches = [branch for branch in args if branch is not type(None)] + _, discriminator = _find_discriminator(tuple(branches)) + if discriminator and isinstance(source, dict): + chosen = source.get(discriminator) + matching = [ + branch + for branch in branches + if isinstance(branch, type) + and dataclasses.is_dataclass(branch) + and chosen in get_args(_schema_hints(branch).get(discriminator)) + ] + if matching: + branches = matching + children = [self._child_annotation(branch, key, source) for branch in branches] + if not children or Any in children: + return Any + return children[0] if all(child == children[0] for child in children) else Union[tuple(children)] + if isinstance(annotation, type) and dataclasses.is_dataclass(annotation): + return _schema_hints(annotation).get(key, Any) + if origin is list and args: + return args[0] + if origin is dict and len(args) == 2: + return args[1] + return Any + + def _prepare_discriminators(self, path: str, waiting: set[str], depth: int) -> None: + """Resolve only discriminator dependencies needed to type this path.""" + parts = path.split("::") if path else [] + if "_disabled_" in parts: + # Guard evaluation must not pull in unrelated variant arguments. + return + for length in range(len(parts) + 1): + parent = "::".join(parts[:length]) + annotation = self._annotation(parent) + if not _is_union_type(get_origin(annotation)): + continue + _, discriminator = _find_discriminator(get_args(annotation)) + item = self._items.get(parent) + source = item.get_config() if item is not None else None + if not discriminator or not isinstance(source, dict): + continue + raw = source.get(discriminator) + if not isinstance(raw, str) or not raw.startswith(("@", "$")): + continue + dependency = f"{parent}::{discriminator}" if parent else discriminator + if dependency == path or dependency in self._projections.get(self._mode, {}): + continue + if dependency in waiting: + raise CircularReferenceError(f"Circular schema discriminator dependency at '{dependency}'") + self._resolve_one_item(dependency, waiting | {path}, depth + 1, instantiate=self._mode[0], eval_expr=self._mode[1]) + + def _normalize(self, value: Any, annotation: Any, path: str) -> Any: + if self._config._coerce: + # Constructor arguments already used to create native objects cannot + # be re-coerced through a differently annotated alias or ancestor. + cache = self._resolved_by_mode[self._mode] + projections = self._projections.setdefault(self._mode, {}) + preserve = { + identity + for key, item in self._items.items() + if key in cache + and key in projections + and (isinstance(item, Component) or isinstance(cache[key], (dict, list))) + for identity in (id(cache[key]), id(projections.get(key))) + } + return _coerce_field( + value, + annotation, + path.replace("::", "."), + self._config._locations, + strict=self._config._strict, + allow_missing=self._config._allow_missing, + allow_unresolved=False, + _preserve=preserve, + ) + return value + + def _check(self, value: Any, annotation: Any, path: str) -> Any: + normalized = self._normalize(value, annotation, path) + _validate_field( + normalized, + annotation, + path.replace("::", "."), + self._config._locations, + strict=self._config._strict, + allow_missing=self._config._allow_missing, + allow_unresolved=False, + ) + return normalized + + def _substitute_cached(self, config: Any, path: str, refs: dict[str, Any], *, root: bool = True) -> Any: + if not root and path in refs: + return refs[path] + if not isinstance(config, (dict, list)): + return self.update_config_with_refs(config, path, refs) + result: Any = type(config)() + for key, child_path, child in self.iter_subconfigs(path, config): + value = self._substitute_cached(child, child_path, refs, root=False) + if Component.is_instantiable(child) and value is None: + continue + if isinstance(result, dict): + result[key] = value + else: + result.append(value) + return result + + def _substitute_refs(self, config: Any, id: str, refs: dict[str, Any]) -> Any: + actual = self._substitute_cached(config, id, refs) + if not self._mode[1] or isinstance(self._items[id], Expression): + return actual + projections = self._projections.setdefault(self._mode, {}) + projected_refs = {key: projections.get(key, value) for key, value in refs.items()} + projected = self._substitute_cached(config, id, projected_refs) + annotation = self._annotation(id) + projections[id] = self._check(projected, annotation, id) + return self._normalize(actual, annotation, id) + + def _prepare_enabled_item(self, id: str, waiting_list: set[str], depth: int, instantiate: bool, eval_expr: bool) -> None: + if eval_expr and isinstance(self._items.get(id), Component): + self._prepare_discriminators(id, waiting_list, depth) + + def _resolve_one_item( + self, + id: str, + waiting_list: set[str] | None = None, + _depth: int = 0, + instantiate: bool = True, + eval_expr: bool = True, + default: Any = None, + ) -> Any: + id = self.normalize_id(id) + self._mode = (instantiate, eval_expr) + cache = self._resolved_by_mode.setdefault(self._mode, {}) + if ( + eval_expr + and id in self._items + and _depth < self.max_resolution_depth + and not isinstance(self._items[id], Component) + ): + self._prepare_discriminators(id, waiting_list or set(), _depth) + result = super()._resolve_one_item(id, waiting_list, _depth, instantiate, eval_expr, default) + if eval_expr and isinstance(self._items.get(id), Expression) and id not in self._projections.get(self._mode, {}): + try: + result = self._check(result, self._annotation(id), id) + except Exception: + cache.pop(id, None) + raise + cache[id] = result + self._projections.setdefault(self._mode, {})[id] = result + return result diff --git a/src/sparkwheel/utils/misc.py b/src/sparkwheel/utils/misc.py index f985955..bf36b0c 100644 --- a/src/sparkwheel/utils/misc.py +++ b/src/sparkwheel/utils/misc.py @@ -3,7 +3,9 @@ from collections.abc import Iterable from typing import Any, TypeVar -from yaml import SafeLoader # type: ignore[import-untyped] +from yaml import SafeLoader +from yaml.constructor import ConstructorError +from yaml.nodes import MappingNode __all__ = [ "first", @@ -51,8 +53,8 @@ def ensure_tuple(vals: Any) -> tuple[Any, ...]: def check_key_duplicates(ordered_pairs: list[tuple[Any, Any]]) -> dict[Any, Any]: """ Checks if there is a duplicated key in the sequence of `ordered_pairs`. - If there is - it will log a warning or raise ValueError - (if configured by environmental var `SPARKWHEEL_STRICT_KEYS==1`) + Duplicates raise ValueError by default. Set `SPARKWHEEL_STRICT_KEYS=0` + to retain warning/last-wins compatibility. Otherwise, it returns the dict made from this sequence. @@ -64,7 +66,7 @@ def check_key_duplicates(ordered_pairs: list[tuple[Any, Any]]) -> dict[Any, Any] keys = set() for k, _ in ordered_pairs: if k in keys: - if os.environ.get("SPARKWHEEL_STRICT_KEYS", "0") == "1": + if os.environ.get("SPARKWHEEL_STRICT_KEYS", "1") != "0": raise ValueError(f"Duplicate key: `{k}`") else: warnings.warn(f"Duplicate key: `{k}`", stacklevel=2) @@ -83,17 +85,56 @@ def __init__(self, stream): super().__init__(stream) # Store filename if available self.source_file = getattr(stream, "name", None) + self._checked_mappings: set[Any] = set() + + def flatten_mapping(self, node): + """Check authored keys before PyYAML expands legal merge overrides. + + PyYAML recursively calls this method for merge sources. Remember checked + nodes because aliases can revisit a mapping after it has been flattened. + """ + if node not in self._checked_mappings: + self._checked_mappings.add(node) + marks: dict[Any, Any] = {} + merge_key = object() + for key_node, _ in node.value: + if key_node.tag == "tag:yaml.org,2002:merge": + key = merge_key + label = "<<" + else: + # PyYAML treats a bare '=' key as text during flattening. + if key_node.tag == "tag:yaml.org,2002:value": + key_node.tag = "tag:yaml.org,2002:str" + key = self.construct_object(key_node, deep=True) + label = key + try: + duplicate = key in marks + except TypeError as error: + raise ConstructorError( + "while constructing a mapping", + node.start_mark, + "found unhashable key", + key_node.start_mark, + ) from error + if duplicate: + first_mark = marks[key] + repeated = key_node.start_mark + message = ( + f"Duplicate key: `{label}`; first defined at {first_mark.name}, " + f"line {first_mark.line + 1}, column {first_mark.column + 1}; " + f"repeated at {repeated.name}, line {repeated.line + 1}, " + f"column {repeated.column + 1}" + ) + if os.environ.get("SPARKWHEEL_STRICT_KEYS", "1") != "0": + raise ValueError(message) + warnings.warn(message, stacklevel=2) + else: + marks[key] = key_node.start_mark + return super().flatten_mapping(node) def construct_mapping(self, node, deep=False): - mapping = set() - for key_node, _ in node.value: - key = self.construct_object(key_node, deep=deep) - if key in mapping: - if os.environ.get("SPARKWHEEL_STRICT_KEYS", "0") == "1": - raise ValueError(f"Duplicate key: `{key}`") - else: - warnings.warn(f"Duplicate key: `{key}`", stacklevel=2) - mapping.add(key) + if not isinstance(node, MappingNode): + raise ConstructorError(None, None, "expected a mapping node", node.start_mark) return super().construct_mapping(node, deep) def construct_object(self, node, deep=False): diff --git a/tests/test_components.py b/tests/test_components.py index 98e87f7..f090fa3 100644 --- a/tests/test_components.py +++ b/tests/test_components.py @@ -323,14 +323,13 @@ def test_parse_import_string_not_import(self): result = expr._parse_import_string("1 + 1") assert result is None - def test_parse_import_string_multiple_imports_warning(self): - """Test _parse_import_string warns on multiple imports.""" + def test_parse_import_string_loads_all_import_aliases(self): + from collections import Counter, defaultdict + expr = Expression("$from collections import Counter, defaultdict", id="test") - with warnings.catch_warnings(record=True) as w: - warnings.simplefilter("always") - expr._parse_import_string("from collections import Counter, defaultdict") - assert len(w) > 0 - assert "multiple import" in str(w[0].message).lower() + assert expr._parse_import_string("from collections import Counter, defaultdict") is Counter + assert expr.globals["Counter"] is Counter + assert expr.globals["defaultdict"] is defaultdict def test_repr(self): """Test string representation.""" @@ -786,3 +785,25 @@ def test_resolve_one_item_with_non_config_item(self): if __name__ == "__main__": pytest.main([__file__, "-v"]) + + +@pytest.mark.parametrize("options", [{"instantiate": False}, {"eval_expr": False}]) +def test_resolver_reset_clears_nondefault_mode(options): + resolver = Resolver(items=[Component({"_target_": "builtins.dict", "value": 1}, id="component")]) + resolver.resolve("component", **options) + assert resolver.is_resolved() + resolver.reset() + assert not resolver.is_resolved() + assert resolver.get_item("component") is None + + +def test_get_item_resolve_populates_requested_mode_without_rewriting_source(): + source = {"_target_": "builtins.dict", "value": "@number"} + item = Component(source, id="component") + resolver = Resolver(items=[item, Expression("$2 * 4", id="number")]) + assert resolver.get_item("component", resolve=True, instantiate=False, eval_expr=False) is item + built = resolver.resolve("component") + assert built == {"value": 8} + assert resolver.get_item("component", resolve=True) is item + assert resolver.resolve("component") is built + assert item.get_config() == source diff --git a/tests/test_config.py b/tests/test_config.py index 6d1e51a..cd21bca 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -12,13 +12,17 @@ - Advanced features (lazy parsing, relative IDs, etc.) """ +import subprocess +import sys import tempfile +from copy import deepcopy +from itertools import permutations from pathlib import Path import pytest import yaml -from sparkwheel import Config, apply_operators +from sparkwheel import Component, Config, Expression, Item, apply_operators from sparkwheel.path_utils import resolve_relative_ids, split_file_and_id @@ -179,16 +183,17 @@ def test_imports_key_multiple_modules(self): assert config.resolve("sep") == os.sep assert config.resolve("path_type") is Path - def test_imports_key_removed_from_data(self): - """Test _imports_ key is removed from config data after processing.""" + def test_imports_key_preserved_in_source_but_not_runtime(self): + """Authored imports survive resolution for inspection and export.""" config = Config().update( { "_imports_": {"json": "json"}, "data": '$json.dumps({"a": 1})', } ) - config.resolve() # Trigger parsing - assert "_imports_" not in config._data + resolved = config.resolve() + assert config.get("_imports_") == {"json": "json"} + assert "_imports_" not in resolved def test_imports_key_combined_with_imports_parameter(self): """Test _imports_ key works with imports parameter.""" @@ -362,16 +367,20 @@ def test_basic_macro(self): """Test basic macro expansion with %.""" config = {"original": {"a": 1, "b": 2}, "copy": "%original"} parser = Config().update(config) - parser.resolve() - assert parser["copy"] == {"a": 1, "b": 2} - assert parser["copy"] is not parser["original"] + resolved = parser.resolve() + assert parser["copy"] == "%original" + assert resolved["copy"] == {"a": 1, "b": 2} + assert resolved["copy"] is not resolved["original"] def test_do_resolve_macro_from_config(self): """Test preprocessing with macro referencing same config.""" parser = Config({"template": {"a": 1, "b": 2}, "copy": "%template"}) parser._parse() - assert parser["copy"] == {"a": 1, "b": 2} - parser["copy"]["a"] = 99 + assert parser["copy"] == "%template" + resolved_copy = parser.resolve("copy") + assert resolved_copy == {"a": 1, "b": 2} + resolved_copy["a"] = 99 + assert parser.resolve("template")["a"] == 1 assert parser["template"]["a"] == 1 def test_do_resolve_macro_load(self): @@ -383,7 +392,8 @@ def test_do_resolve_macro_load(self): try: parser = Config({"local": f"%{filepath}::external"}) parser._parse() - assert parser["local"] == {"value": 42} + assert parser["local"] == f"%{filepath}::external" + assert parser.resolve("local") == {"value": 42} finally: Path(filepath).unlink() @@ -1018,12 +1028,8 @@ def test_delete_dict_keys_via_config_update(self): assert config["dataloaders"] == {"val": {}} def test_delete_list_items_batch_vs_individual(self): - """Test that batch deletion is the only way to delete list items. - - Path notation like ~plugins::0 doesn't work for lists - you MUST use - the batch syntax ~plugins: [0, 2] to delete list items. - """ - # Batch deletion - the correct way + """Batch and sequential deletions use indices in the current list.""" + # Batch deletion uses the list before any of these indices are removed. config1 = Config().update({"plugins": ["a", "b", "c", "d", "e"]}) config1.update({"~plugins": [0, 2]}) assert config1["plugins"] == ["b", "d", "e"] # Removed "a" and "c" @@ -1578,3 +1584,182 @@ def test_is_frozen(self): if __name__ == "__main__": pytest.main([__file__, "-v"]) + + +class TestResolutionModeIsolation: + """Resolving in one mode must not change another mode's representation.""" + + MODES = [(True, True), (True, False), (False, True), (False, False)] + + @staticmethod + def recipe(): + return { + "base": 4, + "number": "$@base * 2", + "component": { + "_target_": "builtins.dict", + "value": "@number", + "nested": {"_target_": "builtins.list", "_args_": [["@base"]]}, + }, + "alias": "@component", + } + + @classmethod + def representation(cls, value): + if isinstance(value, Item): + return (type(value), cls.representation(value.get_config())) + if isinstance(value, dict): + return {key: cls.representation(item) for key, item in value.items()} + if isinstance(value, list): + return [cls.representation(item) for item in value] + return value + + @pytest.mark.parametrize("first_mode,second_mode", list(permutations(MODES, 2))) + def test_modes_match_fresh_configs_in_both_orders(self, first_mode, second_mode): + source = self.recipe() + config = Config(data=deepcopy(source)) + results = [] + for instantiate, eval_expr in (first_mode, second_mode): + options = {"instantiate": instantiate, "eval_expr": eval_expr} + actual = config.resolve(**options) + fresh = Config(data=self.recipe()).resolve(**options) + assert self.representation(actual) == self.representation(fresh) + assert actual["alias"] is actual["component"] + if instantiate: + assert isinstance(actual["component"], dict) + else: + assert isinstance(actual["component"], Component) + if eval_expr: + assert actual["number"] == 8 + else: + assert isinstance(actual["number"], Expression) + results.append(actual) + + # Returning to a previous mode must retain its own shared objects. + for mode, result in zip((first_mode, second_mode), results, strict=True): + assert config.resolve(instantiate=mode[0], eval_expr=mode[1]) is result + assert config.resolve("alias", instantiate=mode[0], eval_expr=mode[1]) is result["component"] + assert config.get() == source + + @pytest.mark.parametrize("mode", MODES) + @pytest.mark.parametrize("edit", ["set", "update", "lazy_false"]) + def test_edit_invalidates_all_resolution_modes(self, mode, edit): + config = Config(data=self.recipe()) + before = {options: config.resolve(instantiate=options[0], eval_expr=options[1]) for options in self.MODES} + if edit == "set": + config.set("base", 7) + elif edit == "update": + config.update({"=base": 7}) + else: + config.resolve(lazy=False) + after = config.resolve(instantiate=mode[0], eval_expr=mode[1]) + fresh_source = self.recipe() + if edit != "lazy_false": + fresh_source["base"] = 7 + fresh = Config(data=fresh_source).resolve(instantiate=mode[0], eval_expr=mode[1]) + assert after is not before[mode] + assert self.representation(after) == self.representation(fresh) + + +class TestSourcePreservation: + """Compiling a runtime generation must preserve authored definitions.""" + + def test_source_preserves_imports_local_copies_and_relative_references(self): + source = { + "_imports_": {"m": "math"}, + "n": 25, + "copied": "%n", + "root": "$m.sqrt(@copied)", + "group": {"n": 9, "copied": "%::n", "alias": "@::n"}, + } + config = Config(data=deepcopy(source)) + assert config.resolve("root") == 5 + assert config.resolve("group") == {"n": 9, "copied": 9, "alias": 9} + assert config.get() == source + config.resolve(lazy=False) + assert config.get() == source + + def test_export_after_resolve_replays_in_fresh_process(self, tmp_path): + source = {"_imports_": {"m": "math"}, "n": 25, "copied": "%n", "root": "$m.sqrt(@copied)"} + config = Config(data=source) + assert config.resolve("root") == 5 + exported = tmp_path / "replay.yaml" + Config.export_config_file(config.get(), exported) + replay = subprocess.run( + [ + sys.executable, + "-c", + "from sparkwheel import Config; import sys; " + "c = Config().update(sys.argv[1]); assert c.resolve('root') == 5; " + "c.set('n', 36); assert c.resolve('root') == 6", + str(exported), + ], + check=False, + capture_output=True, + text=True, + ) + assert replay.returncode == 0, replay.stderr + + @pytest.mark.parametrize("edit", ["set", "update"]) + def test_local_raw_reference_rebuilds_after_edit_without_mutating_old_runtime(self, edit): + config = Config(data={"n": 2, "object": {"_target_": "builtins.dict", "value": "%n"}}) + first = config.resolve("object") + if edit == "set": + config.set("n", 7) + else: + config.update("n=7") + second = config.resolve("object") + assert first == {"value": 2} + assert second == {"value": 7} + assert first is not second + assert config.get("object::value") == "%n" + + @pytest.mark.parametrize("declaration", ["yaml", "expression"]) + def test_removed_import_does_not_survive_next_generation(self, declaration): + from sparkwheel.utils.exceptions import EvaluationError + + source = {"value": "$m.sqrt(4)"} + key = "_imports_" if declaration == "yaml" else "setup" + source[key] = {"m": "math"} if declaration == "yaml" else "$import math as m" + config = Config(data=source) + assert config.resolve("value") == 2 + config.update({f"~{key}": None}) + with pytest.raises(EvaluationError) as error: + config.resolve("value") + assert isinstance(error.value.__cause__, NameError) + + def test_caller_imports_survive_yaml_override_removal(self): + import math + + config = Config(data={"_imports_": {"lib": "json"}, "module": "$lib"}, imports={"lib": math}) + assert config.resolve("module").__name__ == "json" + config.update({"~_imports_": None}) + assert config.resolve("module") is math + + def test_raw_inspection_never_imports_or_constructs(self, monkeypatch): + def forbidden(*args, **kwargs): + raise AssertionError("Inspection executed Python") + + import sparkwheel.config as config_module + + monkeypatch.setattr(config_module, "optional_import", forbidden) + source = {"_imports_": {"custom": "uninstalled_package"}, "component": {"_target_": forbidden}, "x": "$1/0"} + config = Config(data=source) + assert config.get() == source + assert config["_imports_"] == {"custom": "uninstalled_package"} + assert config.get("component::_target_") is forbidden + assert config.get("x") == "$1/0" + + def test_compilation_keeps_opaque_python_inputs(self): + marker = object() + config = Config(data={"object": {"_target_": "builtins.dict", "marker": marker}}) + assert config.resolve("object")["marker"] is marker + assert config.get("object::marker") is marker + + def test_held_callable_keeps_its_import_generation(self): + config = Config(data={"_imports_": {"lib": "math"}, "function": "$lambda: lib.__name__"}) + old_function = config.resolve("function") + config.update({"=_imports_": {"lib": "json"}}) + new_function = config.resolve("function") + assert old_function() == "math" + assert new_function() == "json" diff --git a/tests/test_construction.py b/tests/test_construction.py new file mode 100644 index 0000000..b7e81bc --- /dev/null +++ b/tests/test_construction.py @@ -0,0 +1,388 @@ +"""Behavioral contracts for retained recipes and isolated construction scopes.""" + +import pickle + +import pytest + +from sparkwheel import CircularReferenceError, Config, ConfigKeyError +from sparkwheel.construction import BlockedPathError + + +def test_retaining_and_creating_scope_do_not_import_or_construct(monkeypatch): + import sparkwheel.config as config_module + + def forbidden(*args, **kwargs): + raise AssertionError("Creation executed Python") + + monkeypatch.setattr(config_module, "optional_import", forbidden) + config = Config(data={"_imports_": {"missing": "missing_package"}, "object": {"_target_": forbidden}}) + recipe = config.retain() + scope = recipe.scope() + assert scope.materialized_components() == {} + + +def test_bound_recipe_and_dependencies_are_suppressed_in_ancestor_scan(): + calls = [] + + def forbidden(**kwargs): + calls.append(kwargs) + raise AssertionError("Bound recipe was executed") + + config = Config( + data={ + "object": {"_target_": forbidden, "dependency": {"_target_": forbidden}}, + "alias": "@object", + } + ) + existing = object() + scope = config.retain().scope(bindings={"object": existing}) + result = scope.resolve() + assert result["object"] is existing + assert result["alias"] is existing + assert scope.materialized_components() == {} + assert calls == [] + + +def test_bound_recipe_prunes_broken_raw_references(): + existing = object() + scope = ( + Config(data={"object": {"_target_": "builtins.dict", "bad": "%missing"}}).retain().scope(bindings={"object": existing}) + ) + assert scope.resolve()["object"] is existing + + +@pytest.mark.parametrize("path", ["", "late", "late::dependency", "alias"]) +def test_blocked_graph_fails_before_import_or_dependency_effects(path, monkeypatch): + import sparkwheel.config as config_module + + calls = [] + + def side_effect(): + calls.append("constructor") + return object() + + def import_effect(*args, **kwargs): + calls.append("import") + raise AssertionError("Import must wait for preflight") + + monkeypatch.setattr(config_module, "optional_import", import_effect) + scope = ( + Config( + data={ + "_imports_": {"library": "missing_package"}, + "early": {"_target_": side_effect}, + "late": {"_target_": "builtins.dict", "dependency": {"_target_": side_effect}}, + "alias": "@late", + } + ) + .retain() + .scope(blocked_paths={"late"}) + ) + with pytest.raises(ValueError, match="late"): + scope.resolve(path) + assert calls == [] + + +def test_blocked_plain_child_is_checked_before_ancestor_construction(): + calls = [] + scope = ( + Config(data={"object": {"_target_": lambda **kwargs: calls.append(kwargs), "value": 1}}) + .retain() + .scope(blocked_paths={"object::value"}) + ) + with pytest.raises(ValueError, match="object::value"): + scope.resolve("object") + assert calls == [] + + +def test_bound_parent_with_blocked_descendant_and_exact_child_binding(): + config = Config(data={"model": {"network": {"_target_": "builtins.dict"}, "optimizer": 1}, "alias": "@model"}) + model, network = object(), object() + scope = config.retain().scope(bindings={"model": model, "model::network": network}, blocked_paths={"model::optimizer"}) + assert scope.resolve("model") is model + assert scope.resolve("model::network") is network + assert scope.resolve("alias") is model + with pytest.raises(ValueError, match="optimizer"): + scope.resolve("model::optimizer") + + +def test_scope_iterator_caches_are_fresh_and_aliases_share(): + recipe = Config(data={"iterator": "$iter([1, 2])", "alias": "@iterator"}).retain() + first, second = recipe.scope(), recipe.scope() + iterator = first.resolve("iterator") + assert first.resolve("alias") is iterator + assert list(iterator) == [1, 2] + assert list(first.resolve("iterator")) == [] + assert list(second.resolve("iterator")) == [1, 2] + assert first.materialized_components() == {} + + +def test_retained_source_is_independent_of_edits_and_runtime_cache(): + config = Config(data={"value": 2, "object": {"_target_": "builtins.dict", "value": "%value"}}) + old_runtime = config.resolve("object") + recipe = config.retain() + config.set("value", 7) + first, second = recipe.scope(), recipe.scope() + assert first.resolve("object") == {"value": 2} + assert first.resolve("object") is not old_runtime + assert first.resolve("object") is not second.resolve("object") + assert config.get("object::value") == "%value" + + +def test_same_identity_bind_is_allowed_but_rebinding_is_rejected(): + scope = Config(data={"object": {"_target_": "builtins.dict"}}).retain().scope() + value = scope.resolve("object") + scope.bind("object", value) + scope.bind("object", value) + with pytest.raises(ValueError, match="object"): + scope.bind("object", {}) + assert scope.resolve("object") is value + + +def test_unresolved_binding_works_and_never_enters_materialized_map(): + scope = Config(data={"object": {"_target_": "builtins.dict"}, "alias": "@object"}).retain().scope() + value = object() + scope.bind("object", value) + assert scope.resolve("alias") is value + assert scope.materialized_components() == {} + + +@pytest.mark.parametrize("kind", ["binding", "block"]) +def test_invalid_declaration_paths_are_rejected(kind): + recipe = Config(data={"object": 1}).retain() + with pytest.raises(ConfigKeyError, match="typo"): + recipe.scope(**({"bindings": {"typo": 2}} if kind == "binding" else {"blocked_paths": {"typo"}})) + + +def test_local_copy_child_is_a_valid_binding_path(): + source = {"template": {"child": {"_target_": "builtins.dict"}}, "copy": "%template"} + existing = object() + scope = Config(data=source).retain().scope(bindings={"copy::child": existing}) + assert scope.resolve("copy")["child"] is existing + assert scope.resolve("template")["child"] is not existing + + +def test_binding_under_blocked_ancestor_is_rejected(): + recipe = Config(data={"object": {"child": 1}}).retain() + with pytest.raises(ValueError, match="object"): + recipe.scope(bindings={"object::child": 2}, blocked_paths={"object"}) + + +def test_materialized_map_only_reports_constructed_components_and_is_defensive(): + scope = ( + Config( + data={ + "object": {"_target_": "builtins.dict", "child": {"_target_": "builtins.list"}}, + "iterator": "$iter([1, 2])", + "alias": "@object", + "disabled": {"_target_": "builtins.dict", "_disabled_": True}, + } + ) + .retain() + .scope() + ) + result = scope.resolve() + components = scope.materialized_components() + assert components == {"object": result["object"], "object::child": result["object"]["child"]} + assert components["object"] is result["alias"] + components.clear() + assert set(scope.materialized_components()) == {"object", "object::child"} + + +def test_scopes_get_fresh_import_namespaces_without_losing_caller_bindings(): + import math + + recipe = Config(data={"_imports_": {"j": "json"}, "value": "$j.dumps(m.sqrt(4))"}, imports={"m": math}).retain() + assert recipe.scope().resolve("value") == "2.0" + assert recipe.scope().resolve("value") == "2.0" + + +def test_bound_parent_keeps_unbound_descendant_recipe_resolvable(): + network = object() + model = object() + scope = ( + Config( + data={ + "model": { + "network": {"_target_": "builtins.dict"}, + "optimizer": {"_target_": "builtins.dict", "params": "@model::network"}, + }, + } + ) + .retain() + .scope(bindings={"model": model, "model::network": network}) + ) + assert scope.resolve("model") is model + assert scope.resolve("model::optimizer")["params"] is network + assert set(scope.materialized_components()) == {"model::optimizer"} + + +def test_late_binding_cannot_make_cached_ancestor_inconsistent(): + scope = Config(data={"container": {"plain": []}}).retain().scope() + container = scope.resolve("container") + with pytest.raises(ValueError, match="fresh scope"): + scope.bind("container::plain", []) + scope.bind("container::plain", container["plain"]) + assert scope.resolve("container::plain") is container["plain"] + + +def test_bound_parent_suppresses_unrelated_broken_sibling_when_resolving_child(): + scope = ( + Config(data={"model": {"bad": "%missing", "optimizer": {"_target_": "builtins.dict"}}}) + .retain() + .scope(bindings={"model": object()}) + ) + assert scope.resolve("model::optimizer") == {} + + +def test_bound_direct_resolution_does_not_import(monkeypatch): + import sparkwheel.config as config_module + + def forbidden(*args, **kwargs): + raise AssertionError("Bound-only resolution imported Python") + + monkeypatch.setattr(config_module, "optional_import", forbidden) + existing = object() + scope = ( + Config(data={"_imports_": {"lib": "missing"}, "object": {"_target_": "builtins.dict"}}) + .retain() + .scope(bindings={"object": existing}) + ) + assert scope.resolve("object") is existing + + +def test_no_boundary_scope_matches_config_for_nested_copies_relative_refs_and_imports(): + from copy import deepcopy + + source = { + "_imports_": {"m": "math"}, + "setup": "$import json as j", + "n": 9, + "template": {"n": 4, "number": "%::n", "alias": "@::n", "object": {"_target_": "builtins.list"}}, + "copy": "%template", + "second_copy": "%copy", + "same": "@copy::object", + "expression": "$j.dumps(m.sqrt(@n))", + } + config = Config(data=deepcopy(source)) + expected = config.resolve() + scope = Config(data=deepcopy(source)).retain().scope() + actual = scope.resolve() + assert actual == expected + assert actual["same"] is actual["copy"]["object"] + assert actual["template"]["object"] is not actual["copy"]["object"] + assert actual["second_copy"]["object"] is not actual["copy"]["object"] + + +def test_constructed_error_keeps_retained_source_location(tmp_path): + from sparkwheel import InstantiationError + + path = tmp_path / "recipe.yaml" + path.write_text("object:\n _target_: builtins.int\n _args_: [not-an-integer]\n") + config = Config().update(path) + recipe = config.retain() + config._locations._locations.clear() + with pytest.raises(InstantiationError) as error: + recipe.scope().resolve("object") + assert error.value.source_location.filepath == str(path) + assert error.value.source_location.line == 2 + + +@pytest.mark.parametrize("path", ["items::-1", "items::00"]) +def test_noncanonical_list_binding_paths_are_rejected(path): + recipe = Config(data={"items": [1]}).retain() + with pytest.raises(ConfigKeyError, match="Unknown construction path"): + recipe.scope(bindings={path: object()}) + + +def test_list_binding_is_authoritative_when_resolving_ancestor(): + existing = object() + scope = Config(data={"items": [{"_target_": "builtins.dict"}]}).retain().scope(bindings={"items::0": existing}) + assert scope.resolve("items")[0] is existing + assert scope.materialized_components() == {} + + +def test_import_expression_inside_explicit_descendant_of_bound_parent_runs_first(): + source = { + "model": { + "optimizer": { + "_target_": "builtins.dict", + "value": "$m.sqrt(4)", + "setup": "$import math as m", + } + } + } + scope = Config(data=source).retain().scope(bindings={"model": object()}) + assert scope.resolve("model::optimizer") == Config(data=source).resolve("model::optimizer") + + +def test_materialized_components_excludes_component_returned_iterators(): + scope = Config(data={"iterator": {"_target_": "builtins.iter", "_args_": [[1, 2]]}}).retain().scope() + assert list(scope.resolve("iterator")) == [1, 2] + assert scope.materialized_components() == {} + + +def test_definition_follows_selected_copy_chain_without_executing_or_expanding_siblings(): + opaque = object() + recipe = Config( + data={ + "_imports_": {"never": "missing.module"}, + "template": { + "_target_": "missing.Target", + "network": "%missing", + "opaque": opaque, + "alias": "@missing", + "expression": "$1 / 0", + }, + "model": "%template", + "copy": "%model", + "broken": "%broken", + } + ).retain() + definition = recipe.definition("copy") + assert definition == recipe.definition("template") + assert definition["opaque"] is opaque + assert definition["network"] == "%missing" + assert definition["alias"] == "@missing" + assert recipe.definition("copy::_target_") == "missing.Target" + definition.clear() + assert recipe.definition("model")["_target_"] == "missing.Target" + assert recipe.definition()["model"] == "%template" + + +def test_definition_handles_relative_copy_none_missing_and_cycles(): + recipe = Config(data={"group": {"value": None, "copy": "%::value"}, "bad": "%bad"}).retain() + assert recipe.definition("group::copy") is None + with pytest.raises(ConfigKeyError, match="Unknown definition path"): + recipe.definition("missing") + with pytest.raises(CircularReferenceError): + recipe.definition("bad") + + +@pytest.mark.parametrize("requested", ["model::optimizer", "model::optimizer::lr", "alias"]) +def test_blocked_path_error_preserves_registered_root_and_requested_descendant(requested): + config = Config(data={"model": {"optimizer": {"lr": 0.05}}, "alias": "@model::optimizer::lr"}) + scope = config.retain().scope(blocked_paths={"model::optimizer"}) + with pytest.raises(BlockedPathError) as caught: + scope.resolve(requested) + error = caught.value + expected_path = "model::optimizer::lr" if requested == "alias" else requested + assert isinstance(error, ValueError) + assert error.requested_path == expected_path + assert error.blocked_path == "model::optimizer" + assert str(error) == f"Cannot resolve '{expected_path}': construction path 'model::optimizer' is blocked" + assert error.__cause__ is None + + +@pytest.mark.parametrize("protocol", [pickle.DEFAULT_PROTOCOL, pickle.HIGHEST_PROTOCOL]) +def test_blocked_path_error_survives_pickle_round_trip(protocol): + original = BlockedPathError("model::optimizer::lr", "model::optimizer") + original.context = {"stage": "fit"} + restored = pickle.loads(pickle.dumps(original, protocol=protocol)) + assert type(restored) is BlockedPathError + assert isinstance(restored, ValueError) + assert restored.args == original.args + assert str(restored) == "Cannot resolve 'model::optimizer::lr': construction path 'model::optimizer' is blocked" + assert restored.requested_path == "model::optimizer::lr" + assert restored.blocked_path == "model::optimizer" + assert restored.context == {"stage": "fit"} diff --git a/tests/test_disabled_contract.py b/tests/test_disabled_contract.py new file mode 100644 index 0000000..dd57479 --- /dev/null +++ b/tests/test_disabled_contract.py @@ -0,0 +1,246 @@ +"""Disabled component guards prune unneeded runtime dependencies.""" + +from dataclasses import dataclass + +import pytest + +from sparkwheel import Component, Config, Expression + + +@pytest.mark.parametrize("scoped", [False, True]) +@pytest.mark.parametrize("path", ["off", ""]) +def test_static_disabled_component_prunes_runtime_payload(scoped, path): + calls = [] + source = { + "off": { + "_target_": lambda **kwargs: calls.append("target"), + "_disabled_": True, + "bad": "$1 / 0", + "missing": "@missing", + "nested": {"_target_": lambda: calls.append("nested")}, + "setup": "$import nonexistent_disabled_payload", + } + } + config = Config(data=source) + resolver = config.retain().scope() if scoped else config + assert resolver.resolve(path) == (None if path else {}) + assert calls == [] + + +@pytest.mark.parametrize("disabled", [False, True]) +@pytest.mark.parametrize("scoped", [False, True]) +def test_dynamic_guard_resolves_only_needed_dependencies_before_payload(disabled, scoped): + calls = [] + + def guard(): + calls.append("guard") + return disabled + + config = Config( + data={ + "off": {"_target_": lambda n: calls.append("target") or n, "_disabled_": "$guard()", "n": "$payload()"}, + }, + imports={"guard": guard, "payload": lambda: calls.append("payload") or 2}, + ) + resolver = config.retain().scope() if scoped else config + assert resolver.resolve("off") == (None if disabled else 2) + assert calls == (["guard"] if disabled else ["guard", "payload", "target"]) + resolver.resolve("off") + assert calls.count("guard") == 1 + + +@pytest.mark.parametrize("scoped", [False, True]) +def test_disabled_parent_does_not_prevent_explicit_child_request(scoped): + config = Config( + data={ + "off": { + "_target_": "builtins.dict", + "_disabled_": True, + "child": { + "_target_": "builtins.dict", + "value": "$m.sqrt(4)", + "setup": "$import math as m", + }, + } + } + ) + resolver = config.retain().scope() if scoped else config + assert resolver.resolve("off") is None + assert resolver.resolve("off::child")["value"] == 2 + + +def test_disabled_mode_controls_do_not_execute_guard_expressions_when_disabled(): + calls = [] + config = Config( + data={ + "off": { + "_target_": lambda n: calls.append("target"), + "_disabled_": "$guard()", + "n": "$payload()", + } + }, + imports={"guard": lambda: calls.append("guard") or True, "payload": lambda: calls.append("payload")}, + ) + deferred = config.resolve("off", eval_expr=False) + assert isinstance(deferred, Component) + assert isinstance(deferred.get_config()["_disabled_"], Expression) + assert calls == [] + assert config.resolve("off") is None + assert calls == ["guard"] + assert config.resolve("off", eval_expr=False) is deferred + + +def test_static_disabled_scoped_node_prunes_blocked_payload_without_effects(): + config = Config(data={"off": {"_target_": "builtins.dict", "_disabled_": True, "payload": "$1 / 0"}}) + scope = config.retain().scope(blocked_paths={"off::payload"}) + assert scope.resolve("off") is None + assert scope.materialized_components() == {} + + +def test_disabled_component_does_not_validate_unresolved_constructor_arguments(): + @dataclass + class Args: + n: int + + config = Config(schema=Args).update({"_target_": "builtins.dict", "_disabled_": True, "n": "$'invalid'"}) + assert config.resolve() is None + + +@pytest.mark.parametrize("scoped", [False, True]) +def test_dynamic_true_guard_suppresses_payload_imports_but_guard_imports_execute(scoped): + config = Config( + data={ + "off": { + "_target_": "missing.Target", + "_disabled_": "$import math; math.isfinite(2)", + "setup": "$import missing_disabled_payload_import", + "value": "$1 / 0", + } + } + ) + resolver = config.retain().scope() if scoped else config + assert resolver.resolve("off") is None + + +@pytest.mark.parametrize("scoped", [False, True]) +def test_guard_can_reference_an_import_prefixed_dependency(scoped): + config = Config( + data={ + "flag": "$import math; math.isfinite(2)", + "off": {"_target_": "missing.Target", "_disabled_": "@flag", "value": "$1 / 0"}, + } + ) + resolver = config.retain().scope() if scoped else config + assert resolver.resolve("off") is None + assert resolver.resolve("flag") is True + + +def test_dynamic_scoped_guard_does_not_bypass_conservative_block_preflight(): + calls = [] + config = Config( + data={ + "off": { + "_target_": "builtins.dict", + "_disabled_": "$guard()", + "payload": "$1 / 0", + } + }, + imports={"guard": lambda: calls.append("guard") or True}, + ) + scope = config.retain().scope(blocked_paths={"off::payload"}) + with pytest.raises(ValueError, match="blocked"): + scope.resolve("off") + assert calls == [] + + +def test_static_disabled_root_scan_keeps_copied_descendant_available_for_explicit_scope_request(): + config = Config( + data={ + "template": {"n": 3}, + "off": {"_target_": "builtins.dict", "_disabled_": True, "child": "%template"}, + } + ) + scope = config.retain().scope() + assert scope.resolve() == {"template": {"n": 3}} + assert scope.resolve("off::child") == {"n": 3} + + +def test_false_literal_guard_constructs_normally_and_unevaluated_static_guard_keeps_identity(): + config = Config( + data={ + "on": {"_target_": "builtins.dict", "_disabled_": False, "n": "$2"}, + "off": {"_target_": "builtins.dict", "_disabled_": True, "n": "$1 / 0"}, + } + ) + assert config.resolve("on") == {"n": 2} + wrapper = config.resolve("off", instantiate=False) + assert isinstance(wrapper, Component) + assert wrapper.get_config()["n"] == "$1 / 0" + assert config.resolve("off", instantiate=False) is wrapper + assert config.resolve("off") is None + + +def test_circular_disabled_guard_fails_before_payload_construction(): + from sparkwheel import CircularReferenceError + + calls = [] + config = Config(data={"off": {"_target_": lambda: calls.append("target"), "_disabled_": "$bool(@off)"}}) + with pytest.raises(CircularReferenceError): + config.resolve("off") + assert calls == [] + + +@pytest.mark.parametrize("guard", [True, "$guard()"]) +def test_schema_variant_discriminator_is_not_evaluated_for_disabled_payload(guard): + from typing import Literal + + @dataclass + class IntegerArgs: + kind: Literal["int"] + n: int + + @dataclass + class FloatArgs: + kind: Literal["float"] + n: float + + @dataclass + class Root: + choice: IntegerArgs | FloatArgs + + calls = [] + config = Config( + schema=Root, + imports={ + "guard": lambda: calls.append("guard") or True, + "kind": lambda: calls.append("kind") or "int", + }, + ).update({"choice": {"_target_": "builtins.dict", "_disabled_": guard, "kind": "$kind()", "n": "$1 / 0"}}) + assert config.resolve("choice") is None + assert calls == ([] if guard is True else ["guard"]) + + +def test_exact_guard_bindings_override_source_before_static_pruning(): + recipe = Config( + data={"off": {"_target_": "builtins.dict", "_disabled_": True, "n": "$m.sqrt(4)", "setup": "$import math as m"}} + ).retain() + enabled = recipe.scope(bindings={"off::_disabled_": False}) + assert enabled.resolve("off")["n"] == 2 + + recipe = Config(data={"off": {"_target_": "builtins.dict", "_disabled_": False, "n": "$1 / 0"}}).retain() + disabled = recipe.scope(bindings={"off::_disabled_": True}, blocked_paths={"off::n"}) + assert disabled.resolve("off") is None + + +def test_blocked_guard_is_checked_even_when_source_guard_is_statically_true(): + recipe = Config(data={"off": {"_target_": "builtins.dict", "_disabled_": True}}).retain() + with pytest.raises(ValueError, match="blocked"): + recipe.scope(blocked_paths={"off::_disabled_"}).resolve("off") + + +def test_late_guard_rebinding_cannot_change_an_already_consumed_disabled_decision(): + scope = Config(data={"off": {"_target_": "builtins.dict", "_disabled_": True}}).retain().scope() + assert scope.resolve("off") is None + with pytest.raises(ValueError, match="fresh scope"): + scope.bind("off::_disabled_", False) + scope.bind("off::_disabled_", True) diff --git a/tests/test_editing.py b/tests/test_editing.py new file mode 100644 index 0000000..86f231b --- /dev/null +++ b/tests/test_editing.py @@ -0,0 +1,158 @@ +"""Transactions, container-aware paths, and frozen-source behavior.""" + +from dataclasses import dataclass + +import pytest + +from sparkwheel import Config, ConfigMergeError, ValidationError, validator + + +@dataclass +class Settings: + n: int + payload: dict + + +@pytest.mark.parametrize("edit", ["set", "update", "root"]) +def test_rejected_schema_edit_preserves_source_namespace_and_runtime_identity(edit): + config = Config(schema=Settings).update({"n": 2, "payload": {"_target_": "builtins.dict"}}) + runtime = config.resolve("payload") + source = config.get() + imports = config._imports + resolver = config._resolver + with pytest.raises(ValidationError): + if edit == "set": + config.set("n", "invalid") + elif edit == "root": + config.set("", {"n": "invalid", "payload": {}}) + else: + config.update({"n": "invalid"}) + assert config.get() is source + assert config.get("n") == 2 + assert config.resolve("payload") is runtime + assert config._resolver is resolver + assert config._imports is imports + + +def test_partial_path_patch_failure_does_not_publish_earlier_edits(): + config = Config(data={"n": 2, "payload": {"_target_": "builtins.dict"}}) + runtime = config.resolve("payload") + source = config.get() + with pytest.raises(ConfigMergeError): + config.update({"n": 9, "~missing::child": None}) + assert config.get() is source + assert config.get("n") == 2 + assert config.resolve("payload") is runtime + + +def test_rejected_file_edit_preserves_source_locations(tmp_path): + good, bad = tmp_path / "good.yaml", tmp_path / "bad.yaml" + good.write_text("n: 2\npayload: {}\n") + bad.write_text("n: invalid\n") + config = Config(schema=Settings).update(good) + location = config.locations.get("n") + with pytest.raises(ValidationError): + config.update(bad) + assert config.locations.get("n") is location + assert location.filepath == str(good) + + +def test_list_index_set_preserves_other_elements_and_container_type(): + config = Config(data={"items": [{"n": 1}, {"n": 2}]}) + config.set("items::1::n", 7) + assert config.get("items") == [{"n": 1}, {"n": 7}] + config.update("=items::0={n: 9}") + assert config.get("items") == [{"n": 9}, {"n": 7}] + + +@pytest.mark.parametrize("path", ["items::3", "items::nope", "scalar::child"]) +def test_invalid_traversal_preserves_source_and_runtime(path): + config = Config(data={"items": [1, 2], "scalar": 3, "object": {"_target_": "builtins.dict"}}) + runtime = config.resolve("object") + source = config.get() + with pytest.raises((KeyError, IndexError, TypeError, ValueError)): + config.set(path, 8) + assert config.get() is source + assert config.resolve("object") is runtime + + +def test_path_deletion_and_structural_item_deletion_agree(): + first = Config(data={"group": {"items": [0, 1, 2], "options": {"a": 1, "b": 2}}}) + second = Config(data={"group": {"items": [0, 1, 2], "options": {"a": 1, "b": 2}}}) + first.update({"~group::items": [1], "~group::options": ["a"]}) + second.update({"group": {"~items": [1], "~options": ["a"]}}) + assert first.get() == second.get() == {"group": {"items": [0, 2], "options": {"b": 2}}} + first.update("~group::items::0") + assert first.get("group::items") == [2] + + +def test_failed_batch_item_removal_leaves_all_items(): + config = Config(data={"group": {"options": {"a": 1, "b": 2}}}) + with pytest.raises(ConfigMergeError): + config.update({"~group::options": ["a", "missing"]}) + assert config.get() == {"group": {"options": {"a": 1, "b": 2}}} + + +def test_freeze_detaches_old_views_and_returns_defensive_containers(): + original = {"nested": {"values": [1, 2]}} + config = Config(data=original) + old_view = config.get("nested") + cached = config.resolve("nested") + config.freeze() + old_view["values"].append(3) + original["nested"]["values"].append(4) + current_view = config.get("nested") + current_view["values"].append(5) + config["nested"]["values"].append(6) + assert config.get("nested") == {"values": [1, 2]} + assert config.resolve("nested") is cached + assert cached == {"values": [1, 2]} + config.unfreeze() + config.set("nested::values::0", 7) + assert config.get("nested::values") == [7, 2] + + +def test_updates_and_frozen_views_preserve_opaque_identity(): + import math + + marker = object() + config = Config(data={"module": math, "marker": marker}) + config.update({"n": 2}) + config.freeze() + assert config.get("module") is math + assert config.get("marker") is marker + + +def test_batch_schema_validation_happens_after_all_fields_are_updated(): + config = Config(schema=Settings).update({"n": 1, "payload::value": 2}) + assert config.get() == {"n": 1, "payload": {"value": 2}} + + +def test_cross_field_validator_observes_atomic_batch_and_rejects_atomically(): + @dataclass + class Bounds: + start: int + end: int + + @validator + def check_order(self): + if self.start >= self.end: + raise ValueError("start must precede end") + + config = Config(schema=Bounds).update({"start": 0, "end": 1}) + config.update({"start": 2, "end": 3}) + before = config.get() + with pytest.raises(ValidationError, match="start must precede end"): + config.update({"start": 8, "end": 4}) + assert config.get() is before + assert config.get() == {"start": 2, "end": 3} + + +def test_failed_external_include_update_preserves_cached_objects(): + config = Config(data={"n": 2, "object": {"_target_": "builtins.dict"}}) + runtime = config.resolve("object") + before = config.get() + with pytest.raises(FileNotFoundError): + config.update({"n": 3, "include": "%definitely_missing_recipe_file.yaml"}) + assert config.get() is before + assert config.resolve("object") is runtime diff --git a/tests/test_expression_contract.py b/tests/test_expression_contract.py new file mode 100644 index 0000000..e102f12 --- /dev/null +++ b/tests/test_expression_contract.py @@ -0,0 +1,166 @@ +"""Lexical references and import-prefixed Python expression semantics.""" + +import math + +import pytest + +from sparkwheel import Config + + +@pytest.mark.parametrize( + "expression, expected", + [ + ("$'me@example.com'", "me@example.com"), + ('$"literal @missing and @::relative"', "literal @missing and @::relative"), + ("$r'@missing\\path'", "@missing\\path"), + ("$'''literal @missing\n@::relative'''", "literal @missing\n@::relative"), + ("$f'me@example.com {1 + 1}'", "me@example.com 2"), + ("$1 # @missing", 1), + ], +) +def test_python_literal_and_comment_at_characters_are_not_references(expression, expected): + config = Config(data={"value": expression}) + assert config.resolve("value") == expected + assert config.retain().scope().resolve("value") == expected + assert config.get("value") == expression + + +def test_quoted_reference_lookalike_and_real_reference_coexist(): + config = Config(data={"name": "ok", "group": {"value": "$'@::name=' + @::::name"}}) + assert config.resolve("group::value") == "@::name=ok" + + +def test_whole_value_references_keep_hyphenated_keys(): + config = Config(data={"learning-rate": 2, "alias": "@learning-rate", "value": "$@alias - 1"}) + assert config.resolve("value") == 1 + + +def test_import_followed_by_expression_returns_final_value(): + config = Config(data={"pi": "$import math; math.pi"}) + assert config.resolve("pi") == math.pi + assert config.retain().scope().resolve("pi") == math.pi + + +def test_import_prefix_resolves_reference_dependencies_before_final_expression(): + config = Config(data={"value": "$import math as m; m.sqrt(@n)", "n": 9, "later": "$m.pi"}) + assert config.resolve("value") == 3.0 + assert config.resolve("later") == math.pi + assert config.retain().scope().resolve("value") == 3.0 + + +def test_all_import_aliases_are_available_and_multiple_import_statements_work(): + config = Config(data={"value": "$from math import sqrt, pi; import json as j; j.dumps(sqrt(4) + pi)"}) + assert config.resolve("value") == str(2 + math.pi) + + +def test_ordinary_python_matrix_operator_is_not_a_reference(): + class Matrix: + def __matmul__(self, other): + return (self, other) + + left, right = Matrix(), Matrix() + config = Config(data={"value": "$left@right"}, imports={"left": left, "right": right}) + assert config.resolve("value") == (left, right) + + +def test_import_declarations_are_available_without_executing_unrequested_final_value(): + config = Config(data={"setup": "$import json as j; 1 / 0", "value": "$j.dumps(2)"}) + assert config.resolve("value") == "2" + assert config.retain().scope().resolve("value") == "2" + + +def test_dotted_import_uses_python_binding_rules(): + config = Config(data={"value": "$import urllib.parse; urllib.parse.quote('hello world')"}) + assert config.resolve("value") == "hello%20world" + + +def test_matrix_operator_with_configuration_reference_operands(): + class Matrix: + def __matmul__(self, other): + return (self, other) + + left, right = Matrix(), Matrix() + config = Config(data={"left": left, "right": right, "value": "$@left @ @right"}) + assert config.resolve("value") == (left, right) + + +def test_import_prefix_and_bare_import_modes_do_not_contaminate_each_other(): + from sparkwheel import Expression + + config = Config(data={"setup": "$import math", "value": "$import math; math.sqrt(@n)", "n": 4}) + deferred = config.resolve("value", eval_expr=False) + assert isinstance(deferred, Expression) + assert config.resolve("value") == 2 + assert config.resolve("setup") is math + assert config.resolve("value", eval_expr=False) is deferred + + +def test_unsupported_statements_after_import_fail_instead_of_silently_truncating(): + from sparkwheel import EvaluationError + + config = Config(data={"value": "$import math; n = 2; n"}) + with pytest.raises(EvaluationError) as error: + config.resolve("value") + assert isinstance(error.value.__cause__, SyntaxError) + + +def test_reference_string_formatting_uses_normal_python_expression(): + config = Config(data={"n": 2, "value": "$'value={}'.format(@n)"}) + assert config.resolve("value") == "value=2" + + +def test_fstring_replacement_does_not_add_configuration_reference_syntax(): + from sparkwheel import EvaluationError + + config = Config(data={"n": 2, "value": "$f'value={@n}'"}) + with pytest.raises(EvaluationError) as error: + config.resolve("value") + assert isinstance(error.value.__cause__, SyntaxError) + + +def test_each_import_prefixed_expression_uses_its_own_declared_alias(): + config = Config(data={"real": "$import math as m; m.sqrt(4)", "complex": "$import cmath as m; m.sqrt(4)"}) + assert type(config.resolve("real")) is float + assert type(config.resolve("complex")) is complex + + +def test_expression_imports_follow_direct_item_updates(): + import json + + from sparkwheel import Expression + + expression = Expression("$import math") + assert expression.evaluate() is math + expression.update_config("$import json") + assert expression.evaluate() is json + + +def test_whitespace_after_expression_prefix_is_accepted(): + assert Config(data={"value": "$ @n + 1", "n": 2}).resolve("value") == 3 + + +@pytest.mark.parametrize("expression", ["$import math\n@x + 1", "$import math #comment\n@x + 1"]) +def test_reference_at_new_statement_is_not_matrix_multiplication(expression): + config = Config(data={"x": 2, "value": expression}) + assert config.resolve("value") == 3 + assert config.retain().scope().resolve("value") == 3 + + +def test_reference_after_yield_from_and_adjacent_python_keyword(): + config = Config(data={"xs": [1, 2], "flag": False, "value": "$(lambda: (yield from @xs))()", "negated": "$not@flag"}) + assert list(config.resolve("value")) == [1, 2] + assert config.resolve("negated") is True + scope = config.retain().scope() + assert list(scope.resolve("value")) == [1, 2] + assert scope.resolve("negated") is True + + +def test_unevaluated_import_expression_wrapper_does_not_depend_on_prior_resolution_order(): + source = {"a": "$1", "b": "$import math; math.sqrt(@x)", "x": 4} + first = Config(data=source) + first.resolve("a", eval_expr=False) + after_other = first.resolve("b", eval_expr=False) + direct = Config(data=source).resolve("b", eval_expr=False) + assert after_other.get_config() == direct.get_config() == "$import math; math.sqrt(__local_refs['x'])" + assert first.resolve("b", eval_expr=False) is after_other + assert first.resolve("b") == 2.0 diff --git a/tests/test_items.py b/tests/test_items.py index e2595fa..f143dc8 100644 --- a/tests/test_items.py +++ b/tests/test_items.py @@ -383,17 +383,12 @@ def test_evaluate_regular_import(self): assert result == json assert "json" in expr.globals - def test_evaluate_multiple_imports_warning(self): - """Test warning for multiple imports in one statement.""" - expr = Expression(config="$from os import path, environ", id="test") - - with pytest.warns(UserWarning, match="ignoring multiple import"): - result = expr.evaluate() - + def test_evaluate_multiple_imports_loads_every_alias(self): import os - # Should import the first one - assert result == os.path + expr = Expression(config="$from os import path, environ", id="test") + assert expr.evaluate() is os.path + assert expr.globals["environ"] is os.environ def test_evaluate_non_expression_returns_none(self): """Test evaluating non-expression returns None.""" @@ -549,14 +544,8 @@ def test_component_suggestion_exception_handling(self): component.instantiate() def test_expression_multiple_import_aliases(self): - """Test Expression with multiple import aliases (should warn).""" - import warnings + from collections import Counter, defaultdict expr = Expression(config="$from collections import Counter, defaultdict", id="test") - - # Should warn about multiple imports - with warnings.catch_warnings(record=True) as w: - warnings.simplefilter("always") - expr.evaluate() - # Check that a warning was issued - assert len(w) >= 1 + assert expr.evaluate() is Counter + assert expr.globals["defaultdict"] is defaultdict diff --git a/tests/test_lazy_expression_contract.py b/tests/test_lazy_expression_contract.py new file mode 100644 index 0000000..c9938ac --- /dev/null +++ b/tests/test_lazy_expression_contract.py @@ -0,0 +1,314 @@ +from dataclasses import dataclass +from typing import Any + +import pytest + +from sparkwheel import Config +from sparkwheel.items import Expression +from sparkwheel.schema import ValidationError +from sparkwheel.utils.exceptions import CircularReferenceError + + +@pytest.mark.parametrize( + "expression, expected", + [ + ("$@safe if True else @bad", 2), + ("$True or @bad", True), + ("$False and @bad", False), + ("$[@safe for _ in range(1)]", [2]), + ], +) +def test_native_selection(expression, expected): + config = Config({"safe": 2, "bad": "$1/0", "selected": expression}) + result = config.resolve("selected") + assert result == expected + assert type(result) is type(expected) + assert "bad" not in config._resolver._resolved + + +def test_selected_dependency_executes_once_and_shares_identity(): + calls = [] + + def create(): + calls.append(1) + return object() + + config = Config({"good": {"_target_": create}, "bad": "$1/0", "selected": "$@good if True else @bad"}) + result = config.resolve("selected") + assert result is config.resolve("selected") is config.resolve("good") + assert calls == [1] + + +def test_evaluated_cycle_has_error_cause(): + config = Config({"a": "$@b", "b": "$@a"}) + with pytest.raises(CircularReferenceError): + config.resolve("a") + + +def test_missing_unselected_reference_is_unused(): + assert Config({"x": "$1 if True else @missing"}).resolve("x") == 1 + + +def test_schema_boundary_still_rejects_before_constructor(): + @dataclass + class Root: + value: int + + calls = [] + + def create(value): + calls.append(value) + return object() + + config = Config(schema=Root).update({"_target_": create, "value": "$'bad' if True else @missing"}) + with pytest.raises(ValidationError): + config.resolve() + assert not calls + + +def test_schema_coercion_and_shared_alias(): + @dataclass + class Root: + n: int + alias: int + unused: Any + + config = Config(schema=Root, coerce=True).update({"n": "$'2' if True else @unused", "alias": "@n", "unused": "$1/0"}) + assert config.resolve("n") == config.resolve("alias") == 2 + + +def test_retained_scope_lazy_runtime_preserves_static_blocks(): + recipe = Config({"safe": 2, "bad": "$1/0", "selected": "$@safe if True else @bad"}).retain() + assert recipe.scope().resolve("selected") == 2 + with pytest.raises(ValueError, match="blocked"): + recipe.scope(blocked_paths={"bad"}).resolve("selected") + + +def test_modes_remain_distinct(): + config = Config({"n": 2, "x": "$@n + 1"}) + unevaluated = config.resolve("x", eval_expr=False) + assert isinstance(unevaluated, Expression) + assert config.resolve("x") == 3 + assert config.resolve("x", eval_expr=False) is unevaluated + + +def test_old_delayed_callable_keeps_prior_generation(): + config = Config({"n": 1, "f": "$lambda: @n"}) + old = config.resolve("f") + config.set("n", 2) + new = config.resolve("f") + assert old() == 1 + assert new() == 2 + + +def test_scope_delayed_callable_cannot_change_under_late_binding(): + scope = Config({"n": [1], "f": "$lambda: @n"}).retain().scope() + delayed = scope.resolve("f") + with pytest.raises(ValueError, match="after use"): + scope.bind("n", [2]) + assert delayed() == [1] + + +def test_old_delayed_generator_keeps_prior_generation(): + config = Config({"n": 1, "g": "$( @n for _ in range(2))"}) + old = config.resolve("g") + config.update({"n": 2}) + new = config.resolve("g") + assert list(old) == [1, 1] + assert list(new) == [2, 2] + + +def test_late_same_identity_binding_remains_valid(): + scope = Config({"n": [1], "f": "$lambda: @n"}).retain().scope() + delayed = scope.resolve("f") + value = scope.resolve("n") + scope.bind("n", value) + assert delayed() is value + + +def test_schema_delayed_generation_keeps_old_value_and_cache(): + @dataclass + class Root: + n: int + f: Any + + config = Config(schema=Root, coerce=True).update({"n": "$'1'", "f": "$lambda: @n"}) + old = config.resolve("f") + config.update({"n": "$'2'"}) + new = config.resolve("f") + assert old() == 1 + assert new() == 2 + + +def test_scope_unselected_reference_binding_is_conservatively_reserved(): + scope = Config({"a": 1, "b": 2, "x": "$@a if True else @b"}).retain().scope() + assert scope.resolve("x") == 1 + with pytest.raises(ValueError, match="after use"): + scope.bind("b", 3) + assert scope.resolve("b") == 2 + + +@pytest.mark.parametrize("kind", ["lambda: @n", "(@n for _ in range(2))"]) +def test_rejected_edit_preserves_delayed_value_generation(kind): + @dataclass + class Root: + n: int + delayed: Any + + config = Config(schema=Root).update({"n": 1, "delayed": "$" + kind}) + old = config.resolve("delayed") + resolver = config._resolver + with pytest.raises(ValidationError): + config.update({"n": "invalid"}) + assert config._resolver is resolver + assert config.resolve("delayed") is old + assert (old() if callable(old) else list(old)) in (1, [1, 1]) + + +def test_delayed_import_namespace_is_owned_by_generation(): + config = Config({"_imports_": {"ops": "math"}, "n": 4, "f": "$lambda: ops.sqrt(@n)"}) + old = config.resolve("f") + config.update({"_imports_": {"ops": "cmath"}, "n": 9}) + new = config.resolve("f") + assert old() == 2.0 and type(old()) is float + assert new() == 3.0 and type(new()) is complex + + +@pytest.mark.parametrize("order", [(False, True), (True, False)]) +def test_delayed_references_keep_instantiation_mode_identity(order): + from sparkwheel.items import Component + + calls = [] + + def create(): + calls.append(1) + return object() + + config = Config({"object": {"_target_": create}, "callback": "$lambda: @object"}) + callbacks = {mode: config.resolve("callback", instantiate=mode) for mode in order} + assert calls == [] + native = callbacks[True]() + wrapper = callbacks[False]() + assert calls == [1] + assert isinstance(wrapper, Component) + assert callbacks[True]() is config.resolve("object") is native + assert callbacks[False]() is config.resolve("object", instantiate=False) is wrapper + + +def test_forced_rebuild_does_not_reset_a_published_callback(): + calls = [] + + def create(): + obj = object() + calls.append(obj) + return obj + + config = Config({"object": {"_target_": create}, "callback": "$lambda: @object"}) + old = config.resolve("callback") + new = config.resolve("callback", lazy=False) + assert old is not new + assert old() is not new() + assert old() is calls[0] + assert new() is calls[1] + assert len(calls) == 2 + + +def test_schema_error_in_old_callback_keeps_original_file_location(tmp_path): + @dataclass + class Root: + n: int + callback: Any + + first = tmp_path / "first.yaml" + second = tmp_path / "second.yaml" + first.write_text("n: \"$'invalid'\"\ncallback: '$lambda: @n'\n") + second.write_text("n: '$2'\n") + config = Config(schema=Root).update(first) + old = config.resolve("callback") + config.update(second) + assert config.resolve("callback")() == 2 + with pytest.raises(ValidationError) as caught: + old() + assert caught.value.source_location.filepath == str(first) + + +def test_reservation_covers_nested_changes_and_same_identity_bindings(): + source = {"group": {"n": [1]}, "callback": "$lambda: @group"} + scope = Config(source).retain().scope() + callback = scope.resolve("callback") + with pytest.raises(ValueError, match="fresh scope"): + scope.bind("group::n", [2]) + group = scope.resolve("group") + scope.bind("group", group) + assert callback() is group + + +def test_unselected_cycle_does_not_execute_in_ordinary_config(): + config = Config({"result": "$1 if True else @result"}) + assert config.resolve("result") == 1 + # Scopes retain their conservative static graph contract. + with pytest.raises(CircularReferenceError): + config.retain().scope().resolve("result") + + +@pytest.mark.parametrize("first", ["a", "b"]) +def test_published_callback_cycle_is_detected_in_either_order(first): + config = Config({"a": "$lambda: @b", "b": "$@a()"}) + if first == "a": + config.resolve("a") + with pytest.raises(CircularReferenceError): + config.resolve("b") + assert not any(config._resolver._active_resolutions.values()) + + +def test_published_callback_detects_active_constructor_cycle(): + from sparkwheel.utils.exceptions import InstantiationError + + def create(callback): + return callback() + + config = Config({"a": "$lambda: @b", "b": {"_target_": create, "callback": "@a"}}) + config.resolve("a") + with pytest.raises(InstantiationError) as caught: + config.resolve("b") + cause = caught.value.__cause__ + while isinstance(cause, InstantiationError): + cause = cause.__cause__ + assert isinstance(cause, CircularReferenceError) + assert not any(config._resolver._active_resolutions.values()) + + +def test_cached_callback_can_recurse_with_ordinary_python_termination(): + config = Config({"a": "$lambda n: 0 if n <= 0 else @a(n - 1) + 1"}) + callback = config.resolve("a") + assert callback(4) == 4 + assert config.resolve("a") is callback + + +def test_failed_expression_clears_active_state_for_retry(): + attempts = [] + + def sometimes(): + attempts.append(1) + if len(attempts) == 1: + raise ValueError("try again") + return 3 + + from sparkwheel.utils.exceptions import EvaluationError + + config = Config({"result": "$sometimes()"}, imports={"sometimes": sometimes}) + with pytest.raises(EvaluationError): + config.resolve("result") + assert config.resolve("result") == 3 + assert len(attempts) == 2 + assert not any(config._resolver._active_resolutions.values()) + + +def test_lazy_cycle_reports_the_owning_expression_location(tmp_path): + authored = tmp_path / "authored.yaml" + authored.write_text("a: '$@b'\nb: '$@a'\n") + config = Config().update(authored) + with pytest.raises(CircularReferenceError) as caught: + config.resolve("a") + assert caught.value.source_location.filepath == str(authored) + assert caught.value.source_location.line == 2 diff --git a/tests/test_new_subtree_operators.py b/tests/test_new_subtree_operators.py new file mode 100644 index 0000000..261d29c --- /dev/null +++ b/tests/test_new_subtree_operators.py @@ -0,0 +1,79 @@ +"""Update operators have the same meaning in existing and newly authored trees.""" + +import pytest + +from sparkwheel import Config +from sparkwheel.utils.exceptions import ConfigMergeError + + +@pytest.mark.parametrize( + "source", + [ + {"group": {"=n": 1}}, + {"=group": {"=n": 1}}, + {"group": {"=n": 1}, "items::kept": 2}, + {"=group": {"=n": 1}, "items::kept": 2}, + ], +) +def test_new_subtree_operators_are_applied(source): + config = Config({"items": {"kept": 1}}).update(source) + assert config.get("group") == {"n": 1} + assert config.resolve("group") == {"n": 1} + assert source["group" if "group" in source else "=group"] == {"=n": 1} + + +@pytest.mark.parametrize("initial", [{}, {"group": 0}, {"group": {"n": 0}}]) +def test_nested_operator_meaning_does_not_depend_on_existing_parent(initial): + config = Config(initial).update({"group": {"=n": 1}}) + assert config.get("group") == {"n": 1} + + +@pytest.mark.parametrize( + "patch", + [ + {"values": [{"=n": 1}]}, + {"=values": [{"=n": 1}]}, + {"values": [{"=n": 1}], "other::n": 3}, + ], +) +def test_nested_operators_in_new_list_entries(patch): + config = Config({"values": [{"kept": 2}], "other": {"n": 0}}).update(patch) + assert config.get("values")[-1] == {"n": 1} + if "=values" not in patch: + assert config.get("values")[0] == {"kept": 2} + + +@pytest.mark.parametrize( + "patch", + [ + {"fresh": {"~missing": None}}, + {"=fresh": {"~missing": None}}, + {"fresh": [{"~missing": None}]}, + {"items::kept": 2, "fresh": {"~missing": None}}, + ], +) +def test_missing_removal_in_new_subtree_rolls_back(patch): + config = Config({"items": {"kept": 1}, "object": {"_target_": "builtins.dict"}}) + instance = config.resolve("object") + source = config.get() + imports = config._imports + with pytest.raises(ConfigMergeError, match="missing"): + config.update(patch) + assert config.get() is source + assert config._imports is imports + assert config.resolve("object") is instance + assert config.get("items::kept") == 1 + + +def test_direct_set_remains_exact_assignment(): + config = Config() + config.set("payload", {"=literal": 1}) + assert config.get("payload") == {"=literal": 1} + config.update({"other": 2}) + assert config.get("payload") == {"=literal": 1} + + +def test_nonstring_key_in_mixed_path_update_normalizes_its_value(): + config = Config({"other": {"n": 0}}).update({1: {"=n": 1}, "other::n": 2}) + assert config.get("1") == {"n": 1} + assert config.get("other::n") == 2 diff --git a/tests/test_quick_reference_contract.py b/tests/test_quick_reference_contract.py new file mode 100644 index 0000000..5e8e2ab --- /dev/null +++ b/tests/test_quick_reference_contract.py @@ -0,0 +1,92 @@ +"""Execute current documentation examples with independent expected results.""" + +import re +import shlex +from pathlib import Path + +import yaml + +from sparkwheel import Config + +GUIDES = Path(__file__).resolve().parents[1] / "docs/user-guide" + + +def _section(page, heading): + matches = re.findall(rf"^## {re.escape(heading)}\s*\n(.*?)(?=^## |\Z)", page.read_text(), re.MULTILINE | re.DOTALL) + assert len(matches) == 1, f"{page}: expected exactly one section {heading!r}, found {len(matches)}" + return matches[0] + + +def _fence(section, language, location): + matches = re.findall(rf"^```{language}\s*\n(.*?)^```\s*$", section, re.MULTILINE | re.DOTALL) + assert len(matches) == 1, f"{location}: expected exactly one {language} fence, found {len(matches)}" + return matches[0] + + +def _table_inputs(section, count, location): + inputs = re.findall(r"^\|\s*`([^`]+)`\s*\|", section, re.MULTILINE) + assert len(inputs) == count, f"{location}: expected {count} input rows, found {len(inputs)}" + return inputs + + +def test_documented_basic_expressions(): + page = GUIDES / "expressions.md" + code = _fence(_section(page, "Basic Expressions"), "python", page) + namespace = {} + exec(compile(code, str(page), "exec"), namespace) + assert "config" in namespace, f"{page}: basic example must expose its configuration" + config = namespace["config"] + for name, expected in (("total", 3), ("double", 6)): + result = config.resolve(name) + assert result == expected + assert type(result) is int + + +def test_quick_reference_cli_arguments_preserve_types_and_shell_words(): + page = GUIDES / "quick-reference.md" + arguments = _table_inputs(_section(page, "CLI Overrides"), 5, page) + expected_states = [ + ({"words": ["red"], "limit": 3}, "limit", int), + ({"words": ["red"], "enabled": True}, "enabled", bool), + ({"words": ["red"], "label": "true"}, "label", str), + ({"words": ["blue"]}, "words", list), + ({"words": []}, "words", list), + ] + for argument, (expected, key, expected_type) in zip(arguments, expected_states, strict=True): + tokens = shlex.split(argument) + assert len(tokens) == 1, f"{page}: override must remain one shell word: {argument!r}" + config = Config({"words": ["red"]}).update(tokens[0]) + assert config.get() == expected + assert type(config.get(key)) is expected_type + + +def test_cli_guide_preserves_nested_paths(): + page = GUIDES / "cli.md" + arguments = _table_inputs(_section(page, "Override Syntax"), 6, page) + config = Config({"report": {"title": "Before", "limit": 3}}).update(arguments[0]) + assert config.get() == {"report": {"title": "Example", "limit": 3}} + assert config.get("report::title") == "Example" + + +def test_cli_commands_append_and_replace_independently(): + page = GUIDES / "cli.md" + commands = _fence(_section(page, "Override Syntax"), "bash", page).strip().splitlines() + assert len(commands) == 2, f"{page}: expected both append and replacement commands" + expected_words = [["red", "blue", "red", "green"], ["green", "green"]] + for command, expected in zip(commands, expected_words, strict=True): + tokens = shlex.split(command) + assert tokens[:2] == ["python", "run.py"], f"{page}: unexpected example command {command!r}" + assert len(tokens) == 3, f"{page}: override must remain one shell word: {command!r}" + config = Config({"words": ["red", "blue", "red"]}).update(tokens[2]) + assert config.get() == {"words": expected} + + +def test_delete_example_uses_valid_null_value(): + page = GUIDES / "quick-reference.md" + section = _section(page, "Composition") + matches = re.findall(r"^\|\s*`(~key: [^`]+)`\s*\|", section, re.MULTILINE) + assert len(matches) == 1, f"{page}: expected exactly one whole-key deletion example" + deletion = yaml.safe_load(matches[0]) + assert deletion == {"~key": None} + config = Config({"key": True, "keep": "sentinel"}).update(deletion) + assert config.get() == {"keep": "sentinel"} diff --git a/tests/test_release_tag_contract.py b/tests/test_release_tag_contract.py new file mode 100644 index 0000000..ec63334 --- /dev/null +++ b/tests/test_release_tag_contract.py @@ -0,0 +1,142 @@ +"""Network-free checks for the actual release guard and both workflow callers.""" + +import importlib.util +import subprocess +import sys +from pathlib import Path + +import pytest +import yaml + +pytestmark = pytest.mark.skipif(sys.version_info < (3, 11), reason="Release tooling explicitly uses Python 3.12") +ROOT = Path(__file__).resolve().parents[1] +SCRIPT = ROOT / ".github/scripts/check_release_tag.py" +spec = importlib.util.spec_from_file_location("release_tag_guard", SCRIPT) +guard = importlib.util.module_from_spec(spec) +spec.loader.exec_module(guard) +TAG_PUSH = "github.event_name == 'push' && startsWith(github.ref, 'refs/tags/')" + + +def git(root, *args): + return subprocess.check_output(["git", "-C", str(root), *args], text=True, stderr=subprocess.PIPE).strip() + + +@pytest.fixture +def repository(tmp_path): + git(tmp_path, "init", "--initial-branch=main") + git(tmp_path, "config", "user.email", "release-fixture@example.invalid") + git(tmp_path, "config", "user.name", "Release fixture") + (tmp_path / "src/example").mkdir(parents=True) + (tmp_path / "pyproject.toml").write_text('[project]\nname = "example"\nversion = "0.1.0"\n') + (tmp_path / "src/example/__init__.py").write_text('__version__ = "0.1.0"\nraise RuntimeError("must not import")\n') + git(tmp_path, "add", ".") + git(tmp_path, "commit", "-m", "stable fixture") + sha = git(tmp_path, "rev-parse", "HEAD") + git(tmp_path, "update-ref", "refs/remotes/origin/main", sha) + return tmp_path, sha + + +def test_matching_stable_tag_without_importing_package(repository): + root, sha = repository + assert guard.validate_release(root, "push", "refs/tags/v0.1.0", sha) == "0.1.0" + + +@pytest.mark.parametrize( + "ref", + [ + "refs/tags/v0.1.1", + "refs/tags/v0.1.0.dev0", + "refs/tags/v0.1.0rc1", + "refs/tags/v0.1.0.post1", + "refs/tags/v0.1.0+local", + "refs/tags/anything", + "refs/tags/0.1.0", + "refs/tags/v00.1.0", + "refs/heads/main", + ], +) +def test_reject_wrong_or_nonstable_tag(repository, ref): + root, sha = repository + with pytest.raises(ValueError): + guard.validate_release(root, "push", ref, sha) + + +def test_reject_manual_dispatch_even_on_stable_tag(repository): + root, sha = repository + with pytest.raises(ValueError, match="requires a push"): + guard.validate_release(root, "workflow_dispatch", "refs/tags/v0.1.0", sha) + + +@pytest.mark.parametrize("literal", ['"0.1.1"', '"0.1.0.dev0"', "None"]) +def test_reject_source_version_mismatch(repository, literal): + root, sha = repository + (root / "src/example/__init__.py").write_text(f"__version__ = {literal}\n") + with pytest.raises(ValueError, match="source __version__"): + guard.validate_release(root, "push", "refs/tags/v0.1.0", sha) + + +def test_tag_may_be_ancestor_of_main(repository): + root, sha = repository + git(root, "commit", "--allow-empty", "-m", "later main") + git(root, "update-ref", "refs/remotes/origin/main", "HEAD") + git(root, "checkout", "--detach", sha) + assert guard.validate_release(root, "push", "refs/tags/v0.1.0", sha) == "0.1.0" + + +def test_reject_commit_outside_main(repository): + root, _ = repository + git(root, "checkout", "-b", "unmerged") + git(root, "commit", "--allow-empty", "-m", "unmerged") + with pytest.raises(subprocess.CalledProcessError): + guard.validate_release(root, "push", "refs/tags/v0.1.0", git(root, "rev-parse", "HEAD")) + + +def test_reject_wrong_checkout(repository): + root, sha = repository + git(root, "commit", "--allow-empty", "-m", "other checkout") + with pytest.raises(ValueError, match="checkout"): + guard.validate_release(root, "push", "refs/tags/v0.1.0", sha) + + +def test_missing_main_fails_closed(repository): + root, sha = repository + git(root, "update-ref", "-d", "refs/remotes/origin/main") + with pytest.raises(subprocess.CalledProcessError): + guard.validate_release(root, "push", "refs/tags/v0.1.0", sha) + + +def test_cli_failure_and_success(repository): + root, sha = repository + command = [sys.executable, str(SCRIPT), "--root", str(root), "--event", "push", "--sha", sha, "--ref"] + good = subprocess.run([*command, "refs/tags/v0.1.0"], capture_output=True, text=True) + bad = subprocess.run([*command, "refs/tags/v0.1.0rc1"], capture_output=True, text=True) + assert good.returncode == 0 and "Verified stable release 0.1.0" in good.stdout + assert bad.returncode == 1 and "Release rejected" in bad.stderr + + +def test_workflows_gate_publication_and_share_the_guard(): + publish = yaml.load((ROOT / ".github/workflows/publish.yml").read_text(), Loader=yaml.BaseLoader) + release = yaml.load((ROOT / ".github/workflows/release.yml").read_text(), Loader=yaml.BaseLoader) + assert "workflow_dispatch" in publish["on"] + assert publish["jobs"]["publish"]["if"] == TAG_PUSH + assert publish["jobs"]["publish"]["needs"] == "build" + assert release["jobs"]["release"]["if"] == TAG_PUSH + assert "workflow_dispatch" not in release["on"] + for steps, destination in [ + (publish["jobs"]["build"]["steps"], "Build a binary wheel and a source tarball"), + (release["jobs"]["release"]["steps"], "Create Release"), + ]: + checks = [step for step in steps if step.get("name") == "Verify stable release tag"] + assert len(checks) == 1 + check = checks[0] + assert check["if"] == TAG_PUSH + assert "git fetch --no-tags origin +refs/heads/main:refs/remotes/origin/main" in check["run"] + assert "--no-project --python 3.12 python .github/scripts/check_release_tag.py" in check["run"] + assert '--event "$RELEASE_EVENT" --ref "$RELEASE_REF" --sha "$RELEASE_SHA"' in check["run"] + assert check["env"] == { + "RELEASE_EVENT": "${{ github.event_name }}", + "RELEASE_REF": "${{ github.ref }}", + "RELEASE_SHA": "${{ github.sha }}", + } + assert steps.index(check) < next(i for i, step in enumerate(steps) if step.get("name") == destination) + assert "secrets.PYPI_TOKEN" in publish["jobs"]["publish"]["steps"][-1]["run"] diff --git a/tests/test_schema.py b/tests/test_schema.py index 1fe7810..521fcf6 100644 --- a/tests/test_schema.py +++ b/tests/test_schema.py @@ -1217,8 +1217,10 @@ def check_value(self): # This validator won't run if __init__ fails assert self.value > 0 - # Should still validate the types even if instance creation fails - validate({"value": -5}, Config) + # Validator setup must not silently turn a failure into success. + with pytest.raises(ValidationError, match="Value must be positive") as error: + validate({"value": -5}, Config) + assert isinstance(error.value.__cause__, ValueError) class TestUnionValidationSuccess: diff --git a/tests/test_schema_contract.py b/tests/test_schema_contract.py new file mode 100644 index 0000000..c458904 --- /dev/null +++ b/tests/test_schema_contract.py @@ -0,0 +1,496 @@ +"""Schema policy checks across edits and value/constructor boundaries.""" + +from dataclasses import dataclass, field +from typing import Literal + +import pytest + +from sparkwheel import MISSING, Config, ValidationError, validate, validator + + +@dataclass +class ChoiceA: + kind: Literal["a"] + count: int + + +@dataclass +class ChoiceB: + kind: Literal["b"] + count: int + + +@dataclass +class ChoiceContainer: + choice: ChoiceA | ChoiceB | None + + +def test_optional_discriminated_union_accepts_none(): + validate({"choice": None}, ChoiceContainer) + + +def test_nested_missing_and_leniency_propagate_to_selected_union_branch(): + validate({"choice": {"kind": "a", "count": MISSING, "extra": 1}}, ChoiceContainer, allow_missing=True, strict=False) + + +def test_postponed_annotations_resolve_with_module_context(monkeypatch): + import sys + import types + + module = types.ModuleType("_test_sparkwheel_schema_annotations") + monkeypatch.setitem(sys.modules, module.__name__, module) + exec( + "from __future__ import annotations\nfrom dataclasses import dataclass\n@dataclass\nclass Child:\n n: int\n@dataclass\nclass Parent:\n child: Child\n", + module.__dict__, + ) + validate({"child": {"n": 2}}, module.Parent) + with pytest.raises(ValidationError): + validate({"child": {"n": "invalid"}}, module.Parent) + + +def test_framework_keys_do_not_suppress_custom_validator(): + @dataclass + class Positive: + n: int + + @validator + def check_positive(self): + if self.n <= 0: + raise ValueError("n must be positive") + + with pytest.raises(ValidationError, match="n must be positive"): + Config(schema=Positive).update({"_target_": "builtins.dict", "n": -1}) + + +def test_validator_setup_failure_preserves_original_cause(): + @dataclass + class BadSetup: + n: int + + def __post_init__(self): + raise RuntimeError("setup failed") + + @validator + def check_n(self): + raise AssertionError("unreachable") + + with pytest.raises(ValidationError, match="setup failed") as error: + validate({"n": 2}, BadSetup) + assert isinstance(error.value.__cause__, RuntimeError) + + +def test_defaults_are_not_inserted_and_factories_are_not_called_for_type_validation(): + calls = [] + + @dataclass + class Defaults: + n: int = 2 + values: list = field(default_factory=lambda: calls.append("default") or []) + + config = Config(schema=Defaults).update({}) + assert config.get() == {} + assert calls == [] + + +@pytest.mark.parametrize("annotation", [int, float, Literal[1]]) +def test_bool_does_not_pass_as_a_numeric_field(annotation): + @dataclass + class Numeric: + n: annotation + + with pytest.raises(ValidationError): + validate({"n": True}, Numeric) + + +@dataclass +class Coercible: + count: int + ratio: float + enabled: bool + numbers: list[int] + + +def test_schema_coercion_applies_to_staged_source_recursively(): + config = Config(schema=Coercible).update({"count": "2", "ratio": "0.5", "enabled": "false", "numbers": ["1", "3"]}) + assert config.get() == {"count": 2, "ratio": 0.5, "enabled": False, "numbers": [1, 3]} + config.set("count", "4") + assert config.get("count") == 4 + + +def test_disabled_coercion_rejects_strings_transactionally(): + config = Config(schema=Coercible, coerce=False).update({"count": 2, "ratio": 0.5, "enabled": False, "numbers": [1]}) + before = config.get() + with pytest.raises(ValidationError): + config.set("count", "4") + assert config.get() is before + + +@pytest.mark.parametrize("value", [2.5, True, [1, 2]]) +def test_coercion_does_not_truncate_or_reinterpret_unrelated_types(value): + @dataclass + class Counter: + n: int + + with pytest.raises(ValidationError): + Config(schema=Counter).update({"n": value}) + + +def test_ambiguous_union_coercion_is_rejected_but_existing_string_branch_wins(): + @dataclass + class Ambiguous: + n: int | float + + with pytest.raises(ValidationError, match="Ambiguous"): + Config(schema=Ambiguous).update({"n": "2"}) + + @dataclass + class StringOrInt: + n: int | str + + config = Config(schema=StringOrInt).update({"n": "002"}) + assert config.get("n") == "002" + + +@dataclass +class CountArgs: + n: int + + +@dataclass +class NestedArgs: + child: CountArgs + label: str + + +@dataclass +class RuntimeConfig: + model: NestedArgs + untouched: CountArgs + + +def test_resolved_invalid_argument_fails_before_target_constructor(): + calls = [] + config = Config(schema=CountArgs).update({"_target_": lambda n: calls.append(n), "n": "$'invalid'"}) + with pytest.raises(ValidationError, match="n"): + config.resolve() + assert calls == [] + assert config.get("n") == "$'invalid'" + + +def test_resolved_argument_coercion_keeps_native_children_and_unrelated_nodes_unbuilt(): + calls = [] + + class Child: + def __init__(self, n): + calls.append(("child", n)) + self.n = n + + class Parent: + def __init__(self, child, label): + calls.append(("parent", label)) + assert isinstance(child, Child) + self.child = child + + config = Config(schema=RuntimeConfig).update( + { + "model": {"_target_": Parent, "child": {"_target_": Child, "n": "$'7'"}, "label": "native"}, + "untouched": {"_target_": lambda n: calls.append(("untouched", n)), "n": 1}, + } + ) + model = config.resolve("model") + assert model.child.n == 7 + assert calls == [("child", 7), ("parent", "native")] + assert config.resolve("model") is model + assert config.resolve("model::child") is model.child + assert config.get("model::child::n") == "$'7'" + + +def test_runtime_coercion_disabled_rejects_expression_result(): + config = Config(schema=CountArgs, coerce=False).update({"n": "$'2'"}) + with pytest.raises(ValidationError): + config.resolve("n") + + +def test_root_validator_runs_only_on_root_resolution_with_resolved_arguments(): + calls = [] + + @dataclass + class Bounds: + start: int + end: int + + @validator + def ascending(self): + calls.append((self.start, self.end)) + if self.start >= self.end: + raise ValueError("bounds must ascend") + + config = Config(schema=Bounds).update({"start": "$5", "end": "$1"}) + assert calls == [] + assert config.resolve("start") == 5 + assert calls == [] + with pytest.raises(ValidationError, match="bounds must ascend"): + config.resolve() + assert calls == [(5, 1)] + + +def test_component_aliases_validate_definition_arguments_without_reconstructing(): + @dataclass + class Aliased: + first: CountArgs + second: CountArgs + + calls = [] + config = Config(schema=Aliased).update( + { + "first": {"_target_": lambda n: calls.append(n) or object(), "n": "$2"}, + "second": "@first", + } + ) + values = config.resolve() + assert values["first"] is values["second"] + assert calls == [2] + + +def test_attached_schema_cannot_silently_disappear_in_retained_scope(): + config = Config(schema=CountArgs).update({"n": 1}) + with pytest.raises(NotImplementedError, match="schema"): + config.retain() + + +def test_native_container_component_identity_survives_parent_and_alias_coercion(): + @dataclass + class Shared: + first: CountArgs + second: CountArgs + + config = Config(schema=Shared).update( + { + "first": {"_target_": "builtins.dict", "n": "$'2'"}, + "second": "@first", + } + ) + first = config.resolve("first") + values = config.resolve() + assert first == {"n": 2} + assert values["first"] is first + assert values["second"] is first + assert config.resolve("second") is first + + +def test_alias_cannot_claim_coerced_arguments_after_component_was_constructed(): + @dataclass + class TextArgs: + n: str + + @dataclass + class Conflicting: + first: TextArgs + second: CountArgs + + config = Config(schema=Conflicting).update( + { + "first": {"_target_": "builtins.dict", "n": "2"}, + "second": "@first", + } + ) + assert config.resolve("first") == {"n": "2"} + with pytest.raises(ValidationError): + config.resolve("second") + + +def test_unexpected_validator_exception_is_not_swallowed_by_union_fallback(): + @dataclass + class Broken: + n: int + + @validator + def check(self): + raise RuntimeError("unexpected validation failure") + + @dataclass + class Alternative: + n: int + + @dataclass + class Root: + value: Broken | Alternative + + with pytest.raises(ValidationError, match="unexpected validation failure") as error: + validate({"value": {"n": 1}}, Root) + assert isinstance(error.value.__cause__, RuntimeError) + + +def test_overflowing_or_inexact_integer_is_not_coerced_to_float(): + @dataclass + class Ratio: + value: float + + for value in (10**400, 2**53 + 1): + with pytest.raises(ValidationError, match="represented"): + Config(schema=Ratio).update({"value": value}) + + +@dataclass +class IntChoice: + kind: Literal["int"] + n: int + + +@dataclass +class FloatChoice: + kind: Literal["float"] + n: float + + +@dataclass +class RuntimeChoice: + choice: IntChoice | FloatChoice | None + + +@pytest.mark.parametrize("kind", ["int", "$'int'"]) +def test_selected_union_child_uses_its_branch_and_constructor_receives_coercion(kind): + calls = [] + config = Config(schema=RuntimeChoice).update( + { + "choice": {"_target_": lambda **kw: calls.append(kw) or kw, "n": "$'2'", "kind": kind}, + } + ) + result = config.resolve("choice") + assert result == {"n": 2, "kind": "int"} + assert type(result["n"]) is int + assert calls == [result] + + +def test_local_copy_arguments_coerce_at_runtime_without_rewriting_source_template(): + @dataclass + class Copy: + template: dict + model: CountArgs + + config = Config(schema=Copy).update({"template": {"_target_": "builtins.dict", "n": "2"}, "model": "%template"}) + assert config.resolve("model") == {"n": 2} + assert config.get("model") == "%template" + assert config.get("template::n") == "2" + + +def test_runtime_string_prefix_does_not_suppress_validator(): + @dataclass + class Text: + value: str + + @validator + def check(self): + if self.value.startswith("$"): + raise ValueError("prefix is forbidden") + + config = Config(schema=Text).update({"value": "$'$literal'"}) + with pytest.raises(ValidationError, match="prefix is forbidden"): + config.resolve() + + +def test_plain_container_references_share_validated_identity(): + @dataclass + class Shared: + first: CountArgs + second: CountArgs + + config = Config(schema=Shared).update({"first": {"n": "$'2'"}, "second": "@first"}) + first = config.resolve("first") + second = config.resolve("second") + root = config.resolve() + assert first is second is root["first"] is root["second"] + assert first == {"n": 2} + + +@pytest.mark.parametrize("first_path", ["choice", "choice::n"]) +def test_dynamic_discriminator_keeps_child_cache_and_parent_consistent(first_path): + config = Config(schema=RuntimeChoice).update({"choice": {"n": "$'2'", "kind": "$'float'"}}) + config.resolve(first_path) + child = config.resolve("choice::n") + assert type(child) is float + assert config.resolve("choice")["n"] is child + + +def test_cached_import_expression_is_checked_on_first_schema_request(): + @dataclass + class Fields: + n: int + other: int + + config = Config(schema=Fields).update({"n": "$import math", "other": 1}) + assert config.resolve("other") == 1 + with pytest.raises(ValidationError): + config.resolve("n") + + +def test_discriminated_union_can_also_have_a_primitive_branch(): + @dataclass + class Mixed: + choice: IntChoice | FloatChoice | int + + validate({"choice": 2}, Mixed) + + +def test_circular_discriminator_reference_reports_cycle(): + from sparkwheel import CircularReferenceError + + config = Config(schema=RuntimeChoice).update({"choice": {"n": "$'2'", "kind": "@choice::n"}}) + with pytest.raises(CircularReferenceError, match="Circular"): + config.resolve("choice::n") + + +def test_expression_container_is_coerced_before_becoming_a_shared_dependency(): + @dataclass + class Root: + args: CountArgs + alias: CountArgs + + config = Config(schema=Root).update({"args": "$dict(n='2')", "alias": "@args"}) + args = config.resolve("args") + assert args == {"n": 2} + assert config.resolve("alias") is args + + +def test_schema_modes_keep_wrappers_separate_from_runtime_identity(): + from sparkwheel import Component, Expression + + config = Config(schema=CountArgs).update({"_target_": "builtins.dict", "n": "$'2'"}) + wrapper = config.resolve(instantiate=False) + assert isinstance(wrapper, Component) + assert wrapper.get_config()["n"] == 2 + runtime = config.resolve() + assert runtime == {"n": 2} + deferred = config.resolve(instantiate=False, eval_expr=False) + assert isinstance(deferred.get_config()["n"], Expression) + assert config.resolve() is runtime + assert config.resolve(instantiate=False) is wrapper + + +def test_direct_retained_recipe_cannot_drop_attached_schema(): + from sparkwheel.construction import RetainedConfig + + with pytest.raises(NotImplementedError, match="schema"): + RetainedConfig(Config(schema=CountArgs).update({"n": 1})) + + +def test_discriminated_union_can_also_have_a_plain_mapping_branch(): + @dataclass + class Mixed: + choice: IntChoice | FloatChoice | dict[str, int] + + validate({"choice": {"n": 2}}, Mixed) + + +def test_untagged_union_partial_child_rejects_ambiguous_coercion(): + @dataclass + class FloatArgs: + n: float + + @dataclass + class Root: + args: CountArgs | FloatArgs + + config = Config(schema=Root).update({"args": {"n": "$'2'"}}) + with pytest.raises(ValidationError, match="Ambiguous"): + config.resolve("args::n") diff --git a/tests/test_utils.py b/tests/test_utils.py index 1f7adaf..d531081 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -119,8 +119,9 @@ def test_check_key_duplicates_no_duplicates(self): result = check_key_duplicates(pairs) assert result == {"a": 1, "b": 2, "c": 3} - def test_check_key_duplicates_with_duplicates_warning(self): + def test_check_key_duplicates_with_duplicates_warning(self, monkeypatch): """Test check_key_duplicates warns on duplicates.""" + monkeypatch.setenv("SPARKWHEEL_STRICT_KEYS", "0") pairs = [("a", 1), ("b", 2), ("a", 3)] with warnings.catch_warnings(record=True) as w: warnings.simplefilter("always") @@ -143,8 +144,9 @@ def test_check_key_duplicates_strict_mode(self): else: os.environ["SPARKWHEEL_STRICT_KEYS"] = original - def test_yaml_loader_duplicate_warning(self): + def test_yaml_loader_duplicate_warning(self, monkeypatch): """Test CheckKeyDuplicatesYamlLoader warns on duplicates.""" + monkeypatch.setenv("SPARKWHEEL_STRICT_KEYS", "0") import yaml yaml_str = """ diff --git a/tests/test_yaml_duplicates_contract.py b/tests/test_yaml_duplicates_contract.py new file mode 100644 index 0000000..01a6a37 --- /dev/null +++ b/tests/test_yaml_duplicates_contract.py @@ -0,0 +1,112 @@ +"""Authored duplicate YAML keys fail before committing a configuration update.""" + +import pytest +import yaml + +from sparkwheel import Config +from sparkwheel.loader import Loader +from sparkwheel.utils import CheckKeyDuplicatesYamlLoader, check_key_duplicates + + +@pytest.mark.parametrize("strict", [None, "1"]) +@pytest.mark.parametrize("body", ["n: 1\nn: 2\n", "group:\n n: 1\n n: 2\n", "items:\n - n: 1\n n: 2\n"]) +def test_rejects_duplicates_at_every_depth(tmp_path, monkeypatch, body, strict): + if strict is None: + monkeypatch.delenv("SPARKWHEEL_STRICT_KEYS", raising=False) + else: + monkeypatch.setenv("SPARKWHEEL_STRICT_KEYS", strict) + path = tmp_path / "duplicate.yaml" + path.write_text(body) + with pytest.raises(ValueError, match="Duplicate key") as caught: + Loader().load_file(path) + message = str(caught.value) + assert str(path) in message + assert "first defined" in message and "repeated" in message + assert "line 2" in message or "line 3" in message + + +def test_permissive_compatibility_is_explicit(tmp_path, monkeypatch): + monkeypatch.setenv("SPARKWHEEL_STRICT_KEYS", "0") + path = tmp_path / "compat.yaml" + path.write_text("n: 1\nn: 2\ngroup:\n n: 3\n n: 4\n") + with pytest.warns(UserWarning, match="Duplicate key") as caught: + result, _ = Loader().load_file(path) + assert len(caught) == 2 + assert result == {"n": 2, "group": {"n": 4}} + + +def test_direct_loader_and_pairs_are_strict_by_default(monkeypatch): + monkeypatch.delenv("SPARKWHEEL_STRICT_KEYS", raising=False) + with pytest.raises(ValueError, match="Duplicate key"): + yaml.load("n: 1\nn: 2", CheckKeyDuplicatesYamlLoader) + with pytest.raises(ValueError, match="Duplicate key"): + check_key_duplicates([("n", 1), ("n", 2)]) + + +def test_yaml_merges_keep_native_precedence(tmp_path): + path = tmp_path / "merge.yaml" + path.write_text( + "first: &first {n: 1, a: 10}\nsecond: &second {n: 2, b: 20}\nmerged:\n <<: [*first, *second]\n n: 3\ninherited:\n <<: [*first, *second]\n" + ) + result, locations = Loader().load_file(path) + assert result["merged"] == {"n": 3, "a": 10, "b": 20} + assert result["inherited"] == {"n": 1, "a": 10, "b": 20} + assert locations.get("merged::n").line == 5 + + +def test_duplicate_in_merged_definition_is_rejected(tmp_path): + path = tmp_path / "merge.yaml" + path.write_text("base: &base {n: 1}\nmerged:\n <<: *base\n n: 2\n n: 3\n") + with pytest.raises(ValueError, match="Duplicate key"): + Loader().load_file(path) + + +@pytest.mark.parametrize("include", [False, True]) +def test_duplicate_edit_rolls_back_source_and_runtime(tmp_path, include): + config = Config({"n": 1, "object": {"_target_": "builtins.dict", "n": "@n"}}) + runtime = config.resolve("object") + source = config.get() + namespace = config._imports + path = tmp_path / "duplicate.yaml" + path.write_text("n: 2\nn: 3\n") + edit = {"other": f"%{path}::n"} if include else path + with pytest.raises(ValueError, match="Duplicate key"): + config.update(edit) + assert config.get() is source + assert config._imports is namespace + assert config.resolve("object") is runtime + assert config.get("n") == 1 + + +def test_separate_files_and_explicit_updates_can_override(tmp_path): + first = tmp_path / "first.yaml" + second = tmp_path / "second.yaml" + first.write_text("n: 1\n") + second.write_text("n: 2\n") + config = Config().update(first).update(second) + assert config.get("n") == 2 + config.update({"n": 3}) + assert config.resolve("n") == 3 + + +def test_nested_locations_are_kept(tmp_path): + path = tmp_path / "nested.yaml" + path.write_text("group:\n items:\n - n: 2\n") + _, locations = Loader().load_file(path) + assert locations.get("group::items::0::n").line == 3 + + +@pytest.mark.parametrize("body", ['n: 1\n"n": 2\n', "true: one\n1: two\n", "0x10: one\n16: two\n"]) +def test_equivalent_yaml_keys_are_duplicates(body): + with pytest.raises(ValueError, match="Duplicate key"): + yaml.load(body, CheckKeyDuplicatesYamlLoader) + + +def test_inline_merge_source_is_checked_before_flattening(): + with pytest.raises(ValueError, match="Duplicate key"): + yaml.load("group: {<<: {n: 1, n: 2}}", CheckKeyDuplicatesYamlLoader) + + +def test_quoted_merge_key_is_ordinary_text(): + value = yaml.load('base: &base {n: 1}\ngroup: {<<: *base, "<<": literal}', CheckKeyDuplicatesYamlLoader) + assert value["group"] == {"n": 1, "<<": "literal"} diff --git a/uv.lock b/uv.lock index 4f8a986..ec1876f 100644 --- a/uv.lock +++ b/uv.lock @@ -1714,7 +1714,7 @@ wheels = [ [[package]] name = "sparkwheel" -version = "0.0.9" +version = "0.1.0.dev0" source = { editable = "." } dependencies = [ { name = "pyyaml" }, @@ -1739,6 +1739,7 @@ dev = [ { name = "pytest-html" }, { name = "pytest-metadata" }, { name = "ruff" }, + { name = "types-pyyaml" }, { name = "typing-extensions" }, ] doc = [ @@ -1764,6 +1765,7 @@ test = [ ] types = [ { name = "mypy" }, + { name = "types-pyyaml" }, { name = "typing-extensions" }, ] @@ -1789,6 +1791,7 @@ dev = [ { name = "pytest-html", specifier = ">=3.2.0" }, { name = "pytest-metadata", specifier = ">=3.1.1" }, { name = "ruff", specifier = ">=0.5.0" }, + { name = "types-pyyaml", specifier = ">=6.0.12" }, { name = "typing-extensions", specifier = ">=4.4.0" }, ] doc = [ @@ -1810,6 +1813,7 @@ test = [ ] types = [ { name = "mypy", specifier = ">=1.14.1" }, + { name = "types-pyyaml", specifier = ">=6.0.12" }, { name = "typing-extensions", specifier = ">=4.4.0" }, ] @@ -1913,6 +1917,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/00/c0/8f5d070730d7836adc9c9b6408dec68c6ced86b304a9b26a14df072a6e8c/traitlets-5.14.3-py3-none-any.whl", hash = "sha256:b74e89e397b1ed28cc831db7aea759ba6640cb3de13090ca145426688ff1ac4f", size = 85359, upload-time = "2024-04-19T11:11:46.763Z" }, ] +[[package]] +name = "types-pyyaml" +version = "6.0.12.20260906" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/90/6e/abec85b9013db5b934b0280a6dd104904d84f7bcbaab2e2f3def87ac7463/types_pyyaml-6.0.12.20260906.tar.gz", hash = "sha256:f59c1cc05010b833d2d72287bbaa72610106b28d42d89a907313117faba85212", size = 18649, upload-time = "2026-09-06T06:35:35.362Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/15/c0/fc0644b7ddcfb969e95845837143cb5173ddd6e06ee4ba5fc493cd9329b7/types_pyyaml-6.0.12.20260906-py3-none-any.whl", hash = "sha256:bca893ff0d51df5c9053137d5d0e6ccd36e939a196356f1d5c16372422f5137b", size = 21282, upload-time = "2026-09-06T06:35:34.372Z" }, +] + [[package]] name = "typing-extensions" version = "4.15.0"