diff --git a/.github/workflows/pages.yml b/.github/workflows/pages.yml index 2f2bf927c..9090b0f44 100644 --- a/.github/workflows/pages.yml +++ b/.github/workflows/pages.yml @@ -1,4 +1,4 @@ -name: Dashboard +name: Documentation on: push: @@ -19,7 +19,7 @@ concurrency: jobs: build: - name: Generate Dashboard + name: Build Documentation runs-on: ubuntu-latest steps: - uses: actions/checkout@v6 @@ -33,9 +33,9 @@ jobs: uses: actions/cache@v5 with: path: ~/.cache/pip - key: pip-dashboard-${{ hashFiles('pyproject.toml') }} + key: pip-docs-${{ hashFiles('pyproject.toml') }} restore-keys: | - pip-dashboard- + pip-docs- - name: Install PyTorch CPU run: pip install torch --index-url https://download.pytorch.org/whl/cpu @@ -46,17 +46,8 @@ jobs: pip install -r docs/requirements.txt pip install -e '.[testing]' - - name: Generate dashboard - run: | - mkdir -p _site - python scripts/generate_dashboard.py \ - --output _site/index.html \ - --commit $(git rev-parse --short HEAD) - - - name: Build Sphinx docs - run: | - python docs/_generate_models.py - sphinx-build docs _site/docs + - name: Build documentation + run: sphinx-build docs _site - name: Upload Pages artifact uses: actions/upload-pages-artifact@v4 diff --git a/.gitignore b/.gitignore index 0382198ef..46fc20d0f 100644 --- a/.gitignore +++ b/.gitignore @@ -222,3 +222,4 @@ cache_dir/** .flightdeck/shared/** .worktrees/ dashboard_preview.html +docs/feature-flags.md diff --git a/README.md b/README.md index edce233d0..428ca3e0d 100644 --- a/README.md +++ b/README.md @@ -14,7 +14,7 @@ Supports building ONNX models from HuggingFace model IDs with automatic weight downloading, dtype casting (including bfloat16 via `ir.LazyTensor`), and multi-component export for pipelines. -๐Ÿ“– **[Documentation](https://onnxruntime.github.io/mobius/docs/)** ยท ๐Ÿ“ฆ **[Supported Models](https://onnxruntime.github.io/mobius/docs/models/index.html)** +๐Ÿ“– **[Documentation](https://onnxruntime.github.io/mobius/)** ยท ๐Ÿ“ฆ **[Supported Models](https://onnxruntime.github.io/mobius/models/index.html)** ## Highlighted Models @@ -34,7 +34,7 @@ multi-component export for pipelines. Supports **130 Transformers model types** and **5 Diffusers component types** across **14 task types** and **56+ reusable components**. -See the [model documentation](https://onnxruntime.github.io/mobius/docs/models/index.html) for the complete list. +See the [model documentation](https://onnxruntime.github.io/mobius/models/index.html) for the complete list. ## Installation @@ -79,7 +79,7 @@ mobius build --model openai/whisper-tiny output_dir/ mobius build --model meta-llama/Llama-3.2-1B output_dir/ --dtype f16 ``` -See the [CLI reference](https://onnxruntime.github.io/mobius/docs/cli.html) for all options. +See the [CLI reference](https://onnxruntime.github.io/mobius/cli.html) for all options. ### Examples @@ -113,7 +113,7 @@ The package is organised into four layers: - **Tasks** โ€” Define the ONNX graph I/O contract (inputs, outputs, KV cache) - **Registry** โ€” Maps HuggingFace `model_type` strings to model classes -See the [design document](https://onnxruntime.github.io/mobius/docs/design.html) for details. +See the [design document](https://onnxruntime.github.io/mobius/design.html) for details. ## Development @@ -133,7 +133,7 @@ lintrunner f --all-files ### Adding a new model -See the [AI-assisted model support strategy](https://onnxruntime.github.io/mobius/docs/ai-model-support-strategy.html) +See the [AI-assisted model support strategy](https://onnxruntime.github.io/mobius/ai-model-support-strategy.html) and the developer skills in `.github/skills/`: | Skill | Use when | diff --git a/docs/_ext/dashboard.py b/docs/_ext/dashboard.py new file mode 100644 index 000000000..7d496caed --- /dev/null +++ b/docs/_ext/dashboard.py @@ -0,0 +1,97 @@ +"""Sphinx extension: generate the confidence dashboard after build. + +Hooks into ``build-finished`` to run ``scripts/generate_dashboard.py`` +and place the output in the build directory under ``/dashboard/``. +Also creates a redirect from the old ``/docs/`` path to the new root. +""" + +from __future__ import annotations + +import logging +import subprocess +import sys +from pathlib import Path +from typing import Any + +from sphinx.application import Sphinx + +logger = logging.getLogger(__name__) + +_REDIRECT_HTML = """\ + +\ +\ +Redirecting... +Click here +""" + + +def generate_dashboard(app: Sphinx, exception: Exception | None) -> None: + """Generate dashboard HTML into the build output directory.""" + if exception is not None: + # Don't generate dashboard if Sphinx build failed + return + + repo_root = Path(app.srcdir).parent + script = repo_root / "scripts" / "generate_dashboard.py" + + if not script.exists(): + logger.warning("Dashboard script not found: %s", script) + return + + outdir = Path(app.outdir) + dashboard_dir = outdir / "dashboard" + dashboard_dir.mkdir(parents=True, exist_ok=True) + + # Determine current git commit for display in the dashboard + commit = _git_short_hash(repo_root) + + result = subprocess.run( + [ + sys.executable, + str(script), + "--output", + str(dashboard_dir / "index.html"), + "--commit", + commit, + ], + cwd=str(repo_root), + capture_output=True, + text=True, + check=False, + ) + + if result.returncode != 0: + raise RuntimeError(f"Dashboard generation failed:\n{result.stderr}") + + if result.stdout: + logger.info(result.stdout.strip()) + + # Add redirect from old /docs/ path to root + docs_redirect_dir = outdir / "docs" + docs_redirect_dir.mkdir(parents=True, exist_ok=True) + (docs_redirect_dir / "index.html").write_text(_REDIRECT_HTML) + + +def _git_short_hash(repo_root: Path) -> str: + """Return the short git commit hash, or 'unknown' on failure.""" + try: + result = subprocess.run( + ["git", "rev-parse", "--short", "HEAD"], + cwd=str(repo_root), + capture_output=True, + text=True, + check=True, + ) + return result.stdout.strip() + except (subprocess.CalledProcessError, FileNotFoundError): + return "unknown" + + +def setup(app: Sphinx) -> dict[str, Any]: + app.connect("build-finished", generate_dashboard) + return { + "version": "0.1", + "parallel_read_safe": True, + "parallel_write_safe": True, + } diff --git a/docs/_ext/flags_gen.py b/docs/_ext/flags_gen.py new file mode 100644 index 000000000..c0925e111 --- /dev/null +++ b/docs/_ext/flags_gen.py @@ -0,0 +1,58 @@ +"""Sphinx extension: generate feature flags documentation at build time. + +Hooks into ``builder-inited`` to generate ``docs/feature-flags.md`` from +the ``_Flags`` dataclass. This extension is conditional โ€” it only runs +when the flags module exists (the feature-flags script lives on a +separate PR branch and may not be merged yet). +""" + +from __future__ import annotations + +import logging +import subprocess +import sys +from pathlib import Path +from typing import Any + +from sphinx.application import Sphinx + +logger = logging.getLogger(__name__) + + +def generate_flags_docs(app: Sphinx) -> None: + """Generate feature-flags.md if the generator script exists.""" + docs_dir = Path(app.srcdir) + + # Check both possible locations for the script + script = docs_dir / "_generate_flags_docs.py" + if not script.exists(): + # Fallback to the scripts/ directory (legacy location) + script = docs_dir.parent / "scripts" / "generate_flags_docs.py" + + if not script.exists(): + # Not an error โ€” feature-flags PR may not be merged yet + logger.debug("Flags docs generator not found; skipping") + return + + result = subprocess.run( + [sys.executable, str(script)], + cwd=str(docs_dir.parent), + capture_output=True, + text=True, + check=False, + ) + + if result.returncode != 0: + raise RuntimeError(f"Feature flags doc generation failed:\n{result.stderr}") + + if result.stdout: + logger.info(result.stdout.strip()) + + +def setup(app: Sphinx) -> dict[str, Any]: + app.connect("builder-inited", generate_flags_docs) + return { + "version": "0.1", + "parallel_read_safe": True, + "parallel_write_safe": True, + } diff --git a/docs/_ext/models_gen.py b/docs/_ext/models_gen.py new file mode 100644 index 000000000..61e6d67f3 --- /dev/null +++ b/docs/_ext/models_gen.py @@ -0,0 +1,52 @@ +"""Sphinx extension: generate model documentation pages at build time. + +Hooks into ``builder-inited`` to run the existing model page generator +(``docs/_generate_models.py``) before Sphinx processes source files. +""" + +from __future__ import annotations + +import logging +import subprocess +import sys +from pathlib import Path +from typing import Any + +from sphinx.application import Sphinx + +logger = logging.getLogger(__name__) + + +def generate_model_pages(app: Sphinx) -> None: + """Generate model .md pages from the registry.""" + docs_dir = Path(app.srcdir) + script = docs_dir / "_generate_models.py" + + if not script.exists(): + logger.warning("Model generation script not found: %s", script) + return + + # Run as subprocess to avoid polluting Sphinx's import state. + # The script handles its own sys.path manipulation. + result = subprocess.run( + [sys.executable, str(script)], + cwd=str(docs_dir.parent), + capture_output=True, + text=True, + check=False, + ) + + if result.returncode != 0: + raise RuntimeError(f"Model page generation failed:\n{result.stderr}") + + if result.stdout: + logger.info(result.stdout.strip()) + + +def setup(app: Sphinx) -> dict[str, Any]: + app.connect("builder-inited", generate_model_pages) + return { + "version": "0.1", + "parallel_read_safe": True, + "parallel_write_safe": True, + } diff --git a/docs/_generate_flags_docs.py b/docs/_generate_flags_docs.py new file mode 100644 index 000000000..73072f059 --- /dev/null +++ b/docs/_generate_flags_docs.py @@ -0,0 +1,234 @@ +#!/usr/bin/env python3 +# Copyright (c) ONNX Project Contributors +# SPDX-License-Identifier: Apache-2.0 + +"""Auto-generate docs/feature-flags.md from the _Flags dataclass. + +Run this script whenever a flag is added or modified:: + + python docs/_generate_flags_docs.py + +The generated file is committed to the repository. A CI check (see +``.github/workflows/check-flags-docs.yml``) verifies it is up to date. + +Extraction strategy +------------------- +For each field in :class:`mobius._flags._Flags`: + +* **Name** โ€” from ``dataclasses.fields()``. +* **Env var** โ€” parsed from the ``_env_bool("ENV_VAR", default)`` call + inside the field's ``default_factory`` lambda via the source AST. +* **Default** โ€” the second argument to ``_env_bool`` in the same AST node. +* **Docstring** โ€” the string-literal ``Expr`` node immediately following the + field's ``AnnAssign`` node in the class body. +""" + +from __future__ import annotations + +import ast +import inspect +import sys +import textwrap +from pathlib import Path + +# Resolve paths relative to repo root (this script lives in scripts/). +_REPO_ROOT = Path(__file__).parent.parent +_OUTPUT = _REPO_ROOT / "docs" / "feature-flags.md" + +# Add src/ to the path so we can import mobius without an editable install. +sys.path.insert(0, str(_REPO_ROOT / "src")) + +from mobius._flags import _Flags # noqa: E402 (after sys.path manipulation) + +# --------------------------------------------------------------------------- +# AST extraction helpers +# --------------------------------------------------------------------------- + + +def _parse_class_ast(cls: type) -> ast.ClassDef: + """Return the ``ClassDef`` AST node for *cls*.""" + source = textwrap.dedent(inspect.getsource(cls)) + tree = ast.parse(source) + return next(n for n in ast.walk(tree) if isinstance(n, ast.ClassDef)) + + +def _extract_field_docstrings(class_def: ast.ClassDef) -> dict[str, str]: + """Map field names โ†’ per-field docstring (string literal after the field).""" + docs: dict[str, str] = {} + body = class_def.body + for i, node in enumerate(body): + if not (isinstance(node, ast.AnnAssign) and isinstance(node.target, ast.Name)): + continue + name = node.target.id + if i + 1 < len(body): + nxt = body[i + 1] + if ( + isinstance(nxt, ast.Expr) + and isinstance(nxt.value, ast.Constant) + and isinstance(nxt.value.value, str) + ): + docs[name] = nxt.value.value.strip() + return docs + + +def _extract_env_info(class_def: ast.ClassDef) -> dict[str, tuple[str, bool]]: + """Map field names โ†’ (env_var_name, default_value) from _env_bool calls.""" + info: dict[str, tuple[str, bool]] = {} + for node in class_def.body: + if not (isinstance(node, ast.AnnAssign) and isinstance(node.target, ast.Name)): + continue + name = node.target.id + if node.value is None: + continue + # Walk node.value looking for Lambda โ†’ _env_bool(env_var, default) + for sub in ast.walk(node.value): + if not isinstance(sub, ast.Lambda): + continue + body = sub.body + if ( + isinstance(body, ast.Call) + and isinstance(body.func, ast.Name) + and body.func.id == "_env_bool" + and len(body.args) >= 2 + and isinstance(body.args[0], ast.Constant) + and isinstance(body.args[1], ast.Constant) + ): + env_var: str = body.args[0].value + default: bool = bool(body.args[1].value) + info[name] = (env_var, default) + break + return info + + +# --------------------------------------------------------------------------- +# Markdown generation +# --------------------------------------------------------------------------- + +_STATIC_PREAMBLE = """\ + + +# Feature Flags + +mobius exposes runtime feature flags that control experimental or +environment-specific behaviour. Flags live in `mobius._flags` and are +for **internal use only** โ€” external callers configure them via environment +variables (see below). + +""" + +_STATIC_BODY = """ +## Setting flags via environment variables + +Environment variables are read when the global `flags` singleton is constructed +at import time. Set them before importing mobius (e.g., in a shell or `.env`): + +```bash +# Example: disable warning suppression to see all deduplication warnings +export MOBIUS_SUPPRESS_DEDUP_WARNING=0 + +python -c "import mobius; mobius.build('Qwen/Qwen2.5-0.5B-Instruct')" +``` + +Accepted truthy values: `1`, `true`, `yes` (case-insensitive).\\ +Accepted falsy values: `0`, `false`, `no` (case-insensitive).\\ +Any other value falls back to the field default. + +## Setting flags programmatically (internal code only) + +Internal modules can import `flags` directly and assign to it: + +```python +from mobius._flags import flags + +flags.suppress_dedup_warning = False # disable +flags.suppress_dedup_warning = True # re-enable +``` + +## Using `override_flags()` in tests + +For tests that need a temporary flag value, use the `override_flags` context +manager. It restores the original values on exit, even if the test raises: + +```python +from mobius._flags import override_flags + +def test_build_with_warnings(tmp_path): + with override_flags(suppress_dedup_warning=False): + pkg = mobius.build("Qwen/Qwen2.5-0.5B-Instruct") + # suppress_dedup_warning is restored here +``` + +`override_flags` raises `ValueError` for unknown flag names, catching typos early. + +> **Thread safety:** `override_flags` is not thread-safe โ€” concurrent calls +> in different threads may interleave the save/restore cycle (TOCTOU). For +> pytest, this is safe with `pytest -n auto` because xdist workers run in +> separate processes with independent flag singletons. + +## Listing all flags + +```python +from mobius._flags import list_flags + +print(list_flags()) +``` + +## Adding new flags + +1. Add a field to `_Flags` in `src/mobius/_flags.py`: + + ```python + my_new_flag: bool = dataclasses.field( + default_factory=lambda: _env_bool("MOBIUS_MY_NEW_FLAG", False) + ) + \"\"\"Short description of what my_new_flag controls.\"\"\" + ``` + +2. Wire the flag into the code path it controls. + +3. Regenerate this page: + + ```bash + python docs/_generate_flags_docs.py + ``` + +4. Add tests in `src/mobius/_flags_test.py` following existing patterns. +""" + + +def generate(output: Path = _OUTPUT) -> str: + """Generate the feature-flags markdown and write it to *output*.""" + import dataclasses as dc + + class_def = _parse_class_ast(_Flags) + field_docs = _extract_field_docstrings(class_def) + env_info = _extract_env_info(class_def) + + # Build the flags reference table. + table_rows: list[str] = [] + for field in dc.fields(_Flags): + name = field.name + env_var, default = env_info.get(name, (f"MOBIUS_{name.upper()}", "?")) + doc = field_docs.get(name, "") + # Collapse multi-line docstrings to a single line for the table cell. + doc_single = " ".join(doc.split()) + default_str = str(default).lower() # "true" / "false" + table_rows.append(f"| `{name}` | `{env_var}` | `{default_str}` | {doc_single} |") + + table = ( + "## Available flags\n\n" + "| Flag | Environment variable | Default | Description |\n" + "|------|---------------------|---------|-------------|\n" + + "\n".join(table_rows) + + "\n" + ) + + content = _STATIC_PREAMBLE + table + _STATIC_BODY + output.write_text(content, encoding="utf-8") + print(f"Generated {output}") + return content + + +if __name__ == "__main__": + generate() diff --git a/docs/conf.py b/docs/conf.py index 7f07581e3..7409b0365 100644 --- a/docs/conf.py +++ b/docs/conf.py @@ -5,8 +5,12 @@ from __future__ import annotations +import os import sys +# Make local extensions importable +sys.path.insert(0, os.path.abspath("_ext")) + # -- Project information ----------------------------------------------------- project = "mobius" @@ -22,6 +26,10 @@ "sphinx.ext.viewcode", "sphinx.ext.intersphinx", "sphinx.ext.napoleon", + # Local extensions that generate content during build + "models_gen", + "flags_gen", + "dashboard", ] myst_enable_extensions = [ diff --git a/docs/index.md b/docs/index.md index 1cd631cca..9abbce373 100644 --- a/docs/index.md +++ b/docs/index.md @@ -4,6 +4,8 @@ ONNX model definitions for generative AI architectures using the [onnxscript](ht Build ONNX models directly from HuggingFace model IDs with automatic weight downloading, dtype casting, and multi-component export. +๐Ÿ“Š [Model Support Dashboard](dashboard/index.html) + ```{toctree} :maxdepth: 2 :caption: Contents @@ -14,4 +16,5 @@ models/index design test-architecture ai-model-support-strategy +feature-flags ``` diff --git a/scripts/templates/dashboard.html.j2 b/scripts/templates/dashboard.html.j2 index de5fdb9dc..051d9b83f 100644 --- a/scripts/templates/dashboard.html.j2 +++ b/scripts/templates/dashboard.html.j2 @@ -265,7 +265,10 @@ document.documentElement.setAttribute('data-theme',t);}())

🧪 Testing Confidence Dashboard

- +
+ 📖 Docs + +

Generated {{ timestamp }} · Commit {{ commit }} · {{ total_models }} registered model types diff --git a/src/mobius/_builder.py b/src/mobius/_builder.py index d00727379..c801b888a 100644 --- a/src/mobius/_builder.py +++ b/src/mobius/_builder.py @@ -34,6 +34,7 @@ from mobius._configs import ( BaseModelConfig, ) +from mobius._flags import flags from mobius._model_package import ModelPackage from mobius._registry import registry from mobius._weight_loading import _download_weights @@ -121,7 +122,10 @@ def call(self, model: ir.Model) -> ir.passes.PassResult: def _optimize(model: ir.Model) -> None: """Apply default optimization passes to a model in-place.""" pass_ = ir.passes.PassManager(_DEFAULT_PASSES, steps=2) - with _suppress_dedup_empty_initializer_warnings(): + if flags.suppress_dedup_warning: + with _suppress_dedup_empty_initializer_warnings(): + pass_(model) + else: pass_(model) diff --git a/src/mobius/_flags.py b/src/mobius/_flags.py new file mode 100644 index 000000000..6043139c2 --- /dev/null +++ b/src/mobius/_flags.py @@ -0,0 +1,139 @@ +# Copyright (c) ONNX Project Contributors +# SPDX-License-Identifier: Apache-2.0 + +"""Runtime feature flags for mobius. + +Flags control experimental or environment-specific behaviour. Each flag can be +set via an environment variable (``MOBIUS_``) or programmatically +by assigning to the :data:`flags` singleton. + +Environment variable values are read each time a :class:`_Flags` instance is +constructed. The global :data:`flags` singleton is constructed at import time, +so env vars should be set before importing mobius. Valid truthy strings are +``1``, ``true``, ``yes``; falsy are ``0``, ``false``, ``no`` +(case-insensitive). Any other value falls back to the field default. + +**Adding new flags:** add a field to :class:`_Flags` with a +``dataclasses.field(default_factory=...)`` that calls :func:`_env_bool`, +plus a docstring string literal immediately after the field for documentation +generation. + +Example:: + + from mobius import flags, override_flags + + # Check a flag + if flags.suppress_dedup_warning: + ... + + # Programmatic override (persists until changed) + flags.suppress_dedup_warning = False + + # Scoped override for tests + with override_flags(suppress_dedup_warning=False): + ... +""" + +from __future__ import annotations + +import dataclasses +import os +from collections.abc import Iterator +from contextlib import contextmanager + + +def _env_bool(name: str, default: bool) -> bool: + """Read a boolean from an environment variable. + + Returns *default* if the variable is unset or has an unrecognised value. + """ + val = os.environ.get(name, "") + if val.lower() in ("1", "true", "yes"): + return True + if val.lower() in ("0", "false", "no"): + return False + return default + + +@dataclasses.dataclass +class _Flags: + """Runtime feature flags singleton. + + Each flag maps to a ``MOBIUS_`` environment variable read when + a :class:`_Flags` instance is constructed. The global :data:`flags` + singleton is constructed at import time. Flags can be overridden + programmatically at any point or scoped temporarily with + :func:`override_flags`. + + **Available flags** + + .. list-table:: + :header-rows: 1 + + * - Flag + - Env var + - Default + - Description + * - ``suppress_dedup_warning`` + - ``MOBIUS_SUPPRESS_DEDUP_WARNING`` + - ``True`` + - Suppress "has no constant value" warnings from the initializer + deduplication pass. + """ + + suppress_dedup_warning: bool = dataclasses.field( + default_factory=lambda: _env_bool("MOBIUS_SUPPRESS_DEDUP_WARNING", True) + ) + """Suppress "has no constant value" warnings from the initializer-deduplication pass. + + These warnings are expected noise when optimisation passes run before weights + are loaded. Set ``MOBIUS_SUPPRESS_DEDUP_WARNING=0`` to see all warnings. + """ + + +# Global singleton โ€” import and use this directly. +flags = _Flags() + + +def list_flags() -> dict[str, object]: + """Return the current value of all flags as a plain dict snapshot.""" + return dataclasses.asdict(flags) + + +@contextmanager +def override_flags(**kwargs: bool) -> Iterator[None]: + """Temporarily override one or more flags within a ``with`` block. + + Restores the original values on exit, even if an exception is raised. + Intended for use in tests. + + .. note:: + **Thread safety:** ``override_flags`` is not thread-safe โ€” concurrent + calls in different threads may interleave the save/restore cycle + (TOCTOU). For pytest, this is safe when running with ``-n auto`` + because xdist spawns separate worker *processes* (not threads), so + each worker has its own copy of the flag singleton. + + Raises: + ValueError: If any key in *kwargs* is not a known flag name. + + Example:: + + with override_flags(suppress_dedup_warning=False): + build(model_id) + """ + valid = {f.name for f in dataclasses.fields(_Flags)} + unknown = sorted(set(kwargs) - valid) + if unknown: + available = ", ".join(sorted(valid)) + raise ValueError( + f"Unknown flag name(s): {', '.join(unknown)}. Available flags: {available}" + ) + old = {k: getattr(flags, k) for k in kwargs} + for k, v in kwargs.items(): + setattr(flags, k, v) + try: + yield + finally: + for k, v in old.items(): + setattr(flags, k, v) diff --git a/src/mobius/_flags_test.py b/src/mobius/_flags_test.py new file mode 100644 index 000000000..50975f7a5 --- /dev/null +++ b/src/mobius/_flags_test.py @@ -0,0 +1,110 @@ +# Copyright (c) ONNX Project Contributors +# SPDX-License-Identifier: Apache-2.0 + +"""Tests for the feature flag system.""" + +from __future__ import annotations + +import pytest + +from mobius import _flags + + +class TestDefaultValues: + """Flags have the expected defaults when no env vars are set.""" + + def test_suppress_dedup_warning_default_on(self, monkeypatch): + monkeypatch.delenv("MOBIUS_SUPPRESS_DEDUP_WARNING", raising=False) + f = _flags._Flags() + assert f.suppress_dedup_warning is True + + +class TestEnvVarOverride: + """Flags are read from environment variables at construction time.""" + + @pytest.mark.parametrize("val", ["1", "true", "True", "TRUE", "yes", "YES"]) + def test_truthy_values(self, monkeypatch, val): + monkeypatch.setenv("MOBIUS_SUPPRESS_DEDUP_WARNING", val) + f = _flags._Flags() + assert f.suppress_dedup_warning is True + + @pytest.mark.parametrize("val", ["0", "false", "False", "FALSE", "no", "NO"]) + def test_falsy_values(self, monkeypatch, val): + monkeypatch.setenv("MOBIUS_SUPPRESS_DEDUP_WARNING", val) + f = _flags._Flags() + assert f.suppress_dedup_warning is False + + def test_unknown_value_falls_back_to_default(self, monkeypatch): + monkeypatch.setenv("MOBIUS_SUPPRESS_DEDUP_WARNING", "maybe") + f = _flags._Flags() + assert f.suppress_dedup_warning is True # default + + +class TestProgrammaticOverride: + """Flags can be assigned directly on the singleton.""" + + def test_assign_and_restore(self): + original = _flags.flags.suppress_dedup_warning + try: + _flags.flags.suppress_dedup_warning = not original + assert _flags.flags.suppress_dedup_warning is not original + finally: + _flags.flags.suppress_dedup_warning = original + + +class TestOverrideFlagsContextManager: + """override_flags() restores original values on exit.""" + + def test_single_flag_restored(self): + original = _flags.flags.suppress_dedup_warning + with _flags.override_flags(suppress_dedup_warning=not original): + assert _flags.flags.suppress_dedup_warning is not original + assert _flags.flags.suppress_dedup_warning is original + + def test_restored_on_exception(self): + original = _flags.flags.suppress_dedup_warning + exc = RuntimeError("boom") + try: + with _flags.override_flags(suppress_dedup_warning=not original): + raise exc + except RuntimeError: + pass + assert _flags.flags.suppress_dedup_warning is original + + def test_nested_overrides(self): + original = _flags.flags.suppress_dedup_warning + with _flags.override_flags(suppress_dedup_warning=not original): + assert _flags.flags.suppress_dedup_warning is not original + with _flags.override_flags(suppress_dedup_warning=original): + assert _flags.flags.suppress_dedup_warning is original + assert _flags.flags.suppress_dedup_warning is not original + assert _flags.flags.suppress_dedup_warning is original + + def test_unknown_flag_raises_value_error(self): + with ( + pytest.raises(ValueError, match="Unknown flag"), + _flags.override_flags(nonexistent_flag=True), + ): + pass # pragma: no cover + + +class TestListFlags: + """list_flags() returns a plain dict of all current flag values.""" + + def test_returns_dict(self): + result = _flags.list_flags() + assert isinstance(result, dict) + + def test_contains_suppress_dedup_warning(self): + assert "suppress_dedup_warning" in _flags.list_flags() + + def test_values_match_singleton(self): + result = _flags.list_flags() + assert result["suppress_dedup_warning"] == _flags.flags.suppress_dedup_warning + + def test_returns_snapshot_not_live_view(self): + """list_flags() returns a copy, not a live reference.""" + snapshot = _flags.list_flags() + original = _flags.flags.suppress_dedup_warning + with _flags.override_flags(suppress_dedup_warning=not original): + assert snapshot["suppress_dedup_warning"] == original