From ce2911c51f6f6680cff8e95d9768a180dd9db0db Mon Sep 17 00:00:00 2001 From: Roel <75250264+RoelBollens-TomTom@users.noreply.github.com> Date: Wed, 12 Aug 2026 16:45:14 +0200 Subject: [PATCH] Add an extension mechanism with discovery and codegen support Signed-off-by: Roel <75250264+RoelBollens-TomTom@users.noreply.github.com> --- README.pydantic.md | 61 +- .../changelog.d/634.feature.md | 1 + .../src/overture/schema/cli/commands.py | 25 +- .../tests/test_resolve_types.py | 4 +- .../changelog.d/634.feature.md | 1 + .../src/overture/schema/codegen/cli.py | 13 +- .../codegen/extraction/model_extraction.py | 3 + .../schema/codegen/extraction/specs.py | 1 + .../codegen/extraction/type_analyzer.py | 35 +- .../codegen/extraction/union_extraction.py | 13 +- .../schema/codegen/markdown/renderer.py | 35 +- .../tests/codegen_test_support.py | 8 +- .../overture-schema-codegen/tests/test_cli.py | 4 +- .../tests/test_markdown_renderer.py | 54 ++ .../tests/test_model_extraction.py | 49 +- .../tests/test_model_extractor.py | 34 + .../tests/test_type_analyzer.py | 29 + .../tests/test_union_extraction.py | 16 +- packages/overture-schema-system/README.md | 10 +- .../changelog.d/634.breaking.md | 1 + .../changelog.d/634.feature.md | 1 + .../overture-schema-system/pyproject.toml | 1 + .../schema/system/discovery/__init__.py | 6 +- .../schema/system/discovery/discovery.py | 271 +++++-- .../schema/system/discovery/tag_providers.py | 29 + .../src/overture/schema/system/extension.py | 514 ++++++++++++ .../tests/test_discovery_extension.py | 348 +++++++++ .../tests/test_discovery_filter_models.py | 71 +- .../tests/test_discovery_selection.py | 113 +++ .../tests/test_extension.py | 737 ++++++++++++++++++ .../overture-schema/changelog.d/634.bugfix.md | 1 + .../changelog.d/634.feature.md | 1 + .../src/overture/schema/__init__.py | 8 +- 33 files changed, 2355 insertions(+), 143 deletions(-) create mode 100644 packages/overture-schema-cli/changelog.d/634.feature.md create mode 100644 packages/overture-schema-codegen/changelog.d/634.feature.md create mode 100644 packages/overture-schema-system/changelog.d/634.breaking.md create mode 100644 packages/overture-schema-system/changelog.d/634.feature.md create mode 100644 packages/overture-schema-system/src/overture/schema/system/extension.py create mode 100644 packages/overture-schema-system/tests/test_discovery_extension.py create mode 100644 packages/overture-schema-system/tests/test_discovery_selection.py create mode 100644 packages/overture-schema-system/tests/test_extension.py create mode 100644 packages/overture-schema/changelog.d/634.bugfix.md create mode 100644 packages/overture-schema/changelog.d/634.feature.md diff --git a/README.pydantic.md b/README.pydantic.md index 866e39670..39c758323 100644 --- a/README.pydantic.md +++ b/README.pydantic.md @@ -181,7 +181,7 @@ to filter the working set without importing every model: from overture.schema.system.discovery import ( TagSelector, discover_models, - filter_models, + select_models, ) models = discover_models() @@ -193,7 +193,7 @@ models = discover_models() # ... # } -buildings = filter_models( +buildings = select_models( models, TagSelector(include_any=("overture:theme=buildings",)), ) @@ -206,6 +206,63 @@ to attach custom tags during discovery. See the [`overture-schema-system` README](packages/overture-schema-system/README.md#tagging) for tag format, reserved namespaces, and provider authoring. +### Model Extensions + +A data producer can attach optional fields to models it does not own. An +extension is registered on the same `overture.models` entry-point group as any +other model; discovery recognizes it as an extension by its `Extends` metadata. +A model extension declares its targets with `@extends`: + +```python +from overture.schema.system.extension import extends + +@extends(Place) +class OperatingHours(BaseModel): + primary: list[str] +``` + +Non-model extensions (e.g. a scalar `NewType`) use `Extends(...)` inside +`Annotated` metadata instead. During discovery each extension is exposed as a +standalone one-field wrapper model (hidden by `select_models` unless the +`extension` tag is engaged), and the extension pass adds the field -- optional, +named after the entry point -- to every registered model the targets resolve +to. + +#### How Targets Resolve + +A target may be a model class or a type expression resolving to model classes: +unions, `Annotated`, `NewType`, and `RootModel`. Two rules govern resolution: + +- A union qualifies only if *every* arm resolves to models -- `Place | int` is + rejected as a target. +- A `RootModel` subclass is never a model leaf itself, even though it is a + `BaseModel` subclass. It is an alias for its root annotation, and resolution + recurses into the root -- at any nesting depth. + +The second rule cuts both ways: a `RootModel` over models is an alias for its +arms, while a `RootModel` over a scalar resolves to no model at all, even when +nested inside an otherwise valid expression: + +```python +class Segment(RootModel[RoadSegment | RailSegment]): + pass + +class Version(RootModel[int]): + pass + +Extends(Segment) # OK -- extends RoadSegment and RailSegment +Extends(Version) # TypeError -- scalar root resolves to no models +Extends(Segment | Version) # TypeError -- every union arm must resolve +``` + +The extension pass applies the same alias view to registered entries: a +registered `Segment` is rebuilt as a subclass whose root annotation carries the +extended arms, while a registered `Version` passes through unchanged, since its +scalar root contains nothing to extend. Container types (`list[Place]`, +`dict[str, Place]`) are opaque on both sides: models nested inside them are +neither valid targets nor rewritten. A self-referential root has no finite +shape and is rejected. + ## Development This project uses [uv](https://docs.astral.sh/uv/) for dependency management: diff --git a/packages/overture-schema-cli/changelog.d/634.feature.md b/packages/overture-schema-cli/changelog.d/634.feature.md new file mode 100644 index 000000000..5b99a1d06 --- /dev/null +++ b/packages/overture-schema-cli/changelog.d/634.feature.md @@ -0,0 +1 @@ +Listed standalone extension entries in `list-types` while keeping their permissive wrapper models out of validation unless explicitly selected; extension data validates through the feature models it extends. diff --git a/packages/overture-schema-cli/src/overture/schema/cli/commands.py b/packages/overture-schema-cli/src/overture/schema/cli/commands.py index 24ce225f1..a26c4b397 100644 --- a/packages/overture-schema-cli/src/overture/schema/cli/commands.py +++ b/packages/overture-schema-cli/src/overture/schema/cli/commands.py @@ -23,7 +23,7 @@ ModelKey, TagSelector, discover_models, - filter_models, + select_models, ) from overture.schema.system.discovery.tag import get_values_for_key from overture.schema.system.feature import Feature @@ -35,7 +35,10 @@ group_errors_by_discriminator, select_most_likely_errors, ) -from .tag_options import build_selector, tag_selection_options +from .tag_options import ( + build_selector, + tag_selection_options, +) from .type_analysis import StructuralTuple, get_item_index, introspect_union from .types import ErrorLocation, UnionType, ValidationErrorDict @@ -201,9 +204,15 @@ def resolve_types( *, type_names: tuple[str, ...] = (), ) -> UnionType: - """Resolve a TagSelector + type-names into a Pydantic union type.""" + """Resolve a TagSelector + type-names into a Pydantic union type. + + Uses the discovery layer's default-hidden policy (`select_models`): standalone + extension wrapper models are excluded unless the caller engages the ``extension`` + tag or names a type that only a hidden entry provides. The extension *fields* + they contribute remain available on the feature models they target. + """ models = discover_models() - models = filter_models(models, selector, type_names=type_names) + models = select_models(models, selector, type_names=type_names) if not models: raise ValueError("No models found matching the specified criteria") @@ -859,7 +868,13 @@ def list_types( """ try: models = discover_models() - models = filter_models(models, build_selector(tags, filters, excludes)) + # A listing is introspection, not selection: show every discoverable + # entry, extension wrappers included. + models = select_models( + models, + build_selector(tags, filters, excludes), + include_extension_entries=True, + ) if group_by: grouped_models: dict[str, set[ModelKey]] = {} diff --git a/packages/overture-schema-cli/tests/test_resolve_types.py b/packages/overture-schema-cli/tests/test_resolve_types.py index 55acfd762..7aaeccad1 100644 --- a/packages/overture-schema-cli/tests/test_resolve_types.py +++ b/packages/overture-schema-cli/tests/test_resolve_types.py @@ -1,6 +1,6 @@ -"""Tests for resolve_types — CLI glue between filter_models and union creation. +"""Tests for resolve_types — CLI glue between select_models and union creation. -The combinator algebra of filter_models itself is covered in +The selector combinator algebra itself is covered in `test_discovery_filter_models.py` in the system package. """ diff --git a/packages/overture-schema-codegen/changelog.d/634.feature.md b/packages/overture-schema-codegen/changelog.d/634.feature.md new file mode 100644 index 000000000..630999b76 --- /dev/null +++ b/packages/overture-schema-codegen/changelog.d/634.feature.md @@ -0,0 +1 @@ +Added extension-field provenance to extraction (`is_extension`) and rendered an *(extension)* tag on extension-contributed fields in generated markdown. diff --git a/packages/overture-schema-codegen/src/overture/schema/codegen/cli.py b/packages/overture-schema-codegen/src/overture/schema/codegen/cli.py index 667843692..0c4364ef2 100644 --- a/packages/overture-schema-codegen/src/overture/schema/codegen/cli.py +++ b/packages/overture-schema-codegen/src/overture/schema/codegen/cli.py @@ -6,10 +6,13 @@ import click -from overture.schema.cli.tag_options import build_selector, tag_selection_options +from overture.schema.cli.tag_options import ( + build_selector, + tag_selection_options, +) from overture.schema.system.discovery import ( discover_models, - filter_models, + select_models, ) from .extraction.specs import ModelSpec @@ -57,7 +60,9 @@ def cli() -> None: @cli.command("list") def list_models() -> None: """List all discovered models.""" - models = discover_models() + # A listing is introspection, not selection: show every discoverable + # entry, extension wrappers included. + models = select_models(discover_models(), include_extension_entries=True) names = sorted( model.__name__ if isinstance(model, type) else str(model) for model in models.values() @@ -103,7 +108,7 @@ def generate( all_models = discover_models() - models = filter_models(all_models, build_selector(tags, filters, excludes)) + models = select_models(all_models, build_selector(tags, filters, excludes)) if output_dir: output_dir.mkdir(parents=True, exist_ok=True) diff --git a/packages/overture-schema-codegen/src/overture/schema/codegen/extraction/model_extraction.py b/packages/overture-schema-codegen/src/overture/schema/codegen/extraction/model_extraction.py index ee8f47d00..84bb73909 100644 --- a/packages/overture-schema-codegen/src/overture/schema/codegen/extraction/model_extraction.py +++ b/packages/overture-schema-codegen/src/overture/schema/codegen/extraction/model_extraction.py @@ -8,6 +8,7 @@ from pydantic.fields import FieldInfo from pydantic_core import PydanticUndefined +from overture.schema.system.extension import applied_extensions from overture.schema.system.model_constraint import ModelConstraint from .docstring import clean_docstring @@ -153,6 +154,7 @@ def _extract_model_recursive( descendant_ancestors = ancestors | {model_class} model_resolver, union_resolver = _make_resolvers(cache, descendant_ancestors) + extensions = applied_extensions(model_class) fields: list[FieldSpec] = [] for field_name in _field_order(model_class): @@ -179,6 +181,7 @@ def _extract_model_recursive( description=field_info.description or ti_description, is_required=_is_field_required(field_info, is_optional), is_optional=is_optional, + is_extension=field_name in extensions, ) ) diff --git a/packages/overture-schema-codegen/src/overture/schema/codegen/extraction/specs.py b/packages/overture-schema-codegen/src/overture/schema/codegen/extraction/specs.py index 880e07d9a..404f6690c 100644 --- a/packages/overture-schema-codegen/src/overture/schema/codegen/extraction/specs.py +++ b/packages/overture-schema-codegen/src/overture/schema/codegen/extraction/specs.py @@ -133,6 +133,7 @@ class FieldSpec: description: str | None = None is_required: bool = True is_optional: bool = False + is_extension: bool = False @dataclass diff --git a/packages/overture-schema-codegen/src/overture/schema/codegen/extraction/type_analyzer.py b/packages/overture-schema-codegen/src/overture/schema/codegen/extraction/type_analyzer.py index 559ac62c7..ead5ffecd 100644 --- a/packages/overture-schema-codegen/src/overture/schema/codegen/extraction/type_analyzer.py +++ b/packages/overture-schema-codegen/src/overture/schema/codegen/extraction/type_analyzer.py @@ -51,6 +51,8 @@ from pydantic.fields import FieldInfo from typing_extensions import Sentinel, assert_never, evaluate_forward_ref +from overture.schema.system.extension import Extends + from .docstring import clean_docstring from .field import ( AnyScalar, @@ -294,18 +296,21 @@ def _recurse( return NewTypeShape(name=ctx.name, ref=ctx.ref, inner=inner), opt, desc if origin is Annotated: - args = get_args(annotation) - inner_annotation = args[0] + inner_annotation, *metadata = get_args(annotation) own_desc: str | None = None collected: list[ConstraintSource] = [] - for c in args[1:]: - if isinstance(c, FieldInfo): - if c.description is not None and own_desc is None: - own_desc = clean_docstring(c.description) - for m in c.metadata: - collected.append(_constraint_source(m, newtype_ctx)) + for item in metadata: + constraint_items: tuple[object, ...] + if isinstance(item, FieldInfo): + if own_desc is None and item.description is not None: + own_desc = clean_docstring(item.description) + constraint_items = tuple(item.metadata) else: - collected.append(_constraint_source(c, newtype_ctx)) + constraint_items = (item,) + for constraint in constraint_items: + # `Extends` is extension-target metadata, not a field constraint. + if not isinstance(constraint, Extends): + collected.append(_constraint_source(constraint, newtype_ctx)) # Pick the annotation to recurse into and the optionality this # Annotated layer contributes. A directly-wrapped union is peeled @@ -514,10 +519,18 @@ def attach_field_metadata(shape: FieldShape, field_info: FieldInfo) -> FieldShap wrapping applies here just as it does during normal annotation unwrapping: the constraints anchor at the topmost constraint-bearing layer. Returns *shape* unchanged when there is no metadata. + + `Extends` is excluded here as well as in the `Annotated` frame: pydantic + hoists a top-level `Annotated`'s metadata into `field_info.metadata`, so + extension-target declarations would otherwise re-enter as constraints. """ - if not field_info.metadata: + extra = tuple( + ConstraintSource(None, None, m) + for m in field_info.metadata + if not isinstance(m, Extends) + ) + if not extra: return shape - extra = tuple(ConstraintSource(None, None, m) for m in field_info.metadata) return attach_constraints(shape, extra) diff --git a/packages/overture-schema-codegen/src/overture/schema/codegen/extraction/union_extraction.py b/packages/overture-schema-codegen/src/overture/schema/codegen/extraction/union_extraction.py index 7ab582f35..ec3066ee3 100644 --- a/packages/overture-schema-codegen/src/overture/schema/codegen/extraction/union_extraction.py +++ b/packages/overture-schema-codegen/src/overture/schema/codegen/extraction/union_extraction.py @@ -97,7 +97,7 @@ def extract_discriminator( _TypeShape = tuple[object, ...] -_FieldKey = tuple[str, _TypeShape, frozenset[object]] +_FieldKey = tuple[str, _TypeShape, frozenset[object], bool] def _structural_fingerprint(spec: FieldSpec) -> _TypeShape: @@ -242,8 +242,15 @@ def extract_union( # that already handles a field present on only some arms # (`check_builder._field_checks_for_union`), and the renderer's # collision resolver already disambiguates multiple `Check`s - # landing on the same field label. - key = (fs.name, _structural_fingerprint(fs), _constraints_fingerprint(fs)) + # landing on the same field label. Provenance joins the key so a + # native field and an identically-shaped extension field never + # collapse into one row. + key = ( + fs.name, + _structural_fingerprint(fs), + _constraints_fingerprint(fs), + fs.is_extension, + ) existing = seen.get(key) prior_sources = existing.variant_sources or () if existing else () seen[key] = AnnotatedField( diff --git a/packages/overture-schema-codegen/src/overture/schema/codegen/markdown/renderer.py b/packages/overture-schema-codegen/src/overture/schema/codegen/markdown/renderer.py index 8f0911ab8..bfe64f793 100644 --- a/packages/overture-schema-codegen/src/overture/schema/codegen/markdown/renderer.py +++ b/packages/overture-schema-codegen/src/overture/schema/codegen/markdown/renderer.py @@ -271,6 +271,11 @@ def _expandable_list_suffix(field_spec: FieldSpec) -> str: return "[]" * depth if depth > 0 else "" +def _extension_tag(field_spec: FieldSpec) -> str: + """Return an italic ` *(extension)*` tag for extension-contributed fields, else "".""" + return " *(extension)*" if field_spec.is_extension else "" + + def _expand_sub_model( field_spec: FieldSpec, name: str, @@ -320,7 +325,13 @@ def _expand_model_fields( for field_spec in fields: row = _field_template_context(field_spec, ctx) name = f"{prefix}{field_spec.name}" if prefix else field_spec.name - row["name"] = f"{name}{_expandable_list_suffix(field_spec)}" + display_name = f"{name}{_expandable_list_suffix(field_spec)}" + tag = _extension_tag(field_spec) + if tag: + row["name"] = f"`{display_name}`{tag}" + row["pre_formatted"] = True + else: + row["name"] = display_name if not prefix: _annotate_field_constraints(row, field_spec, ctx) result.append(row) @@ -348,10 +359,10 @@ def _short_variant_name(class_name: str, union_name: str) -> str: return class_name -def _variant_tag(annotated: AnnotatedField, union_name: str) -> str | None: - """Return an italic variant tag like `*(Road, Water)*`, or None for shared fields.""" +def _variant_tag(annotated: AnnotatedField, union_name: str) -> str: + """Return an italic variant tag like `*(Road, Water)*`, or "" for shared fields.""" if annotated.variant_sources is None: - return None + return "" short_names = [ _short_variant_name(v.__name__, union_name) for v in annotated.variant_sources ] @@ -371,20 +382,20 @@ def _expand_union_fields( result: list[_FieldRow] = [] for annotated in spec.annotated_fields: field_spec = annotated.field_spec - row = _field_template_context(field_spec, ctx) name = field_spec.name - suffix = _expandable_list_suffix(field_spec) + display_name = f"{name}{_expandable_list_suffix(field_spec)}" + row = _field_template_context(field_spec, ctx) _annotate_field_constraints(row, field_spec, ctx) - if constraint_notes and field_spec.name in constraint_notes: - _annotate_constraint_notes(row, constraint_notes[field_spec.name]) + if constraint_notes and name in constraint_notes: + _annotate_constraint_notes(row, constraint_notes[name]) - tag = _variant_tag(annotated, spec.name) - if tag is not None: - row["name"] = f"`{name}{suffix}`{tag}" + tags = _variant_tag(annotated, spec.name) + _extension_tag(field_spec) + if tags: + row["name"] = f"`{display_name}`{tags}" row["pre_formatted"] = True else: - row["name"] = f"{name}{suffix}" + row["name"] = display_name result.append(row) _expand_sub_model(field_spec, name, ctx, result) diff --git a/packages/overture-schema-codegen/tests/codegen_test_support.py b/packages/overture-schema-codegen/tests/codegen_test_support.py index 19da210a5..b87eb8e1c 100644 --- a/packages/overture-schema-codegen/tests/codegen_test_support.py +++ b/packages/overture-schema-codegen/tests/codegen_test_support.py @@ -36,7 +36,7 @@ from overture.schema.system.discovery import ( TagSelector, discover_models, - filter_models, + select_models, ) from overture.schema.system.discovery.tag import get_values_for_key from overture.schema.system.doc import DocumentedEnum @@ -436,8 +436,10 @@ def flat_specs_from_discovery( """Build a flat list of RecordSpecs from discovery, with entry_point set.""" models = discover_models() if theme: - models = filter_models( - models, TagSelector(include_any=(f"overture:theme={theme}",)) + models = select_models( + models, + TagSelector(include_any=(f"overture:theme={theme}",)), + include_extension_entries=True, ) return [ spec diff --git a/packages/overture-schema-codegen/tests/test_cli.py b/packages/overture-schema-codegen/tests/test_cli.py index 07120d132..b193e891f 100644 --- a/packages/overture-schema-codegen/tests/test_cli.py +++ b/packages/overture-schema-codegen/tests/test_cli.py @@ -56,8 +56,8 @@ def test_generate_with_tag_filter(self, cli_runner: CliRunner) -> None: ) assert result.exit_code == 0 - assert "Building" in result.output - assert "Place" not in result.output + assert "# Building" in result.output + assert "# Place" not in result.output def test_generate_accepts_filter_flag(self, cli_runner: CliRunner) -> None: """generate accepts --filter without error. diff --git a/packages/overture-schema-codegen/tests/test_markdown_renderer.py b/packages/overture-schema-codegen/tests/test_markdown_renderer.py index 76f527fb8..2b9f3a84c 100644 --- a/packages/overture-schema-codegen/tests/test_markdown_renderer.py +++ b/packages/overture-schema-codegen/tests/test_markdown_renderer.py @@ -49,6 +49,12 @@ render_pydantic_type, ) from overture.schema.codegen.markdown.reverse_references import UsedByEntry, UsedByKind +from overture.schema.system.extension import ( + Extends, + create_extended_model, + extends, + wrap_extension, +) from overture.schema.system.field_constraint import ( CountryCodeAlpha2Constraint, JsonPointerConstraint, @@ -306,6 +312,54 @@ class ModelWithMultilineDesc(BaseModel): assert lines[i].startswith("|"), f"Table broken at line {i}: {lines[i]}" +class TestRenderFeatureExtensionFields: + """Extension-contributed fields are visibly tagged in rendered markdown.""" + + def test_extension_field_gets_badge(self) -> None: + class Target(BaseModel): + """A target model.""" + + name: str + + @extends(Target) + class Ext(BaseModel): + note: str + + wrapper = wrap_extension("ext", Ext) + assert wrapper is not None + extended = create_extended_model(Target, {"ext": wrapper}) + + spec = extract_model(extended) + result = render_model(spec) + + assert "`ext` *(extension)*" in result + # The native field is untagged. + assert "`name` *(extension)*" not in result + assert "| `name` |" in result + + def test_scalar_extension_does_not_render_extends_as_a_constraint(self) -> None: + """Regression test: a scalar extension's `Extends` metadata must not appear as a + rendered constraint (it previously leaked its docstring into the Constraints column). + """ + + class Place(BaseModel): + """A place.""" + + name: str + + Capacity = NewType("Capacity", Annotated[int, Field(ge=0), Extends(Place)]) + wrapper = wrap_extension("capacity", Capacity) + assert wrapper is not None + extended = create_extended_model(Place, {"capacity": wrapper}) + + spec = extract_model(extended) + result = render_model(spec) + + assert "`capacity` *(extension)*" in result + assert "Extends" not in result + assert "non-model extension targets" not in result + + class TestRenderFeatureWithThemeType: """Tests for rendering Feature-like models with theme/type.""" diff --git a/packages/overture-schema-codegen/tests/test_model_extraction.py b/packages/overture-schema-codegen/tests/test_model_extraction.py index e9b23a3ba..045da7917 100644 --- a/packages/overture-schema-codegen/tests/test_model_extraction.py +++ b/packages/overture-schema-codegen/tests/test_model_extraction.py @@ -2,6 +2,7 @@ from typing import Annotated, Optional +from annotated_types import Ge from codegen_test_support import FeatureWithRootModel from overture.schema.codegen.extraction.field import ( ArrayOf, @@ -14,7 +15,8 @@ from overture.schema.codegen.extraction.length_constraints import ArrayMinLen from overture.schema.codegen.extraction.model_extraction import extract_model from overture.schema.common.scoping.vehicle import VehicleSelector -from pydantic import BaseModel, Field +from overture.schema.system.extension import Extends +from pydantic import BaseModel, Field, RootModel def test_extract_model_populates_union_terminal() -> None: @@ -159,6 +161,51 @@ class Node(BaseModel): assert inner.starts_cycle is True +def test_extends_hoisted_into_field_metadata_is_not_a_constraint() -> None: + """`Extends` must be excluded from the `FieldInfo.metadata` path too. + + Pydantic strips a top-level `Annotated` and hoists its metadata into + `FieldInfo.metadata`, bypassing the Annotated-frame exclusion; without + a filter in `attach_field_metadata`, the extension-target declaration + would render as a constraint in generated docs. + """ + + class Place(BaseModel): + name: str + + class Venue(BaseModel): + capacity: Annotated[int, Field(ge=0), Extends(Place)] + + spec = extract_model(Venue) + cap = next(f for f in spec.fields if f.name == "capacity") + + assert isinstance(cap.shape, Primitive) + constraints = [cs.constraint for cs in cap.shape.constraints] + assert any(isinstance(c, Ge) for c in constraints) + assert not any(isinstance(c, Extends) for c in constraints) + + +def test_extends_in_rootmodel_root_metadata_is_not_a_constraint() -> None: + """The `RootModel` root's hoisted metadata path excludes `Extends` too.""" + + class Place(BaseModel): + name: str + + class CapRoot(RootModel[Annotated[int, Field(ge=0), Extends(Place)]]): + pass + + class Venue(BaseModel): + capacity: CapRoot + + spec = extract_model(Venue) + cap = next(f for f in spec.fields if f.name == "capacity") + + assert isinstance(cap.shape, Primitive) + constraints = [cs.constraint for cs in cap.shape.constraints] + assert any(isinstance(c, Ge) for c in constraints) + assert not any(isinstance(c, Extends) for c in constraints) + + def test_field_metadata_minlen_wrapped_as_array_min_len() -> None: """MinLen in field_info.metadata is wrapped to ArrayMinLen, not left as raw MinLen. diff --git a/packages/overture-schema-codegen/tests/test_model_extractor.py b/packages/overture-schema-codegen/tests/test_model_extractor.py index 6356d30d3..ddfad1370 100644 --- a/packages/overture-schema-codegen/tests/test_model_extractor.py +++ b/packages/overture-schema-codegen/tests/test_model_extractor.py @@ -19,6 +19,11 @@ terminal_of, ) from overture.schema.codegen.extraction.model_extraction import extract_model +from overture.schema.system.extension import ( + create_extended_model, + extends, + wrap_extension, +) from overture.schema.system.field_constraint import UniqueItemsConstraint from overture.schema.system.geometric import ( Geometry, @@ -549,3 +554,32 @@ class TestModel(BaseModel): spec = extract_model(TestModel) field = find_field(spec, "name") assert field.description is None + + +class TestExtensionFields: + """`FieldSpec.is_extension` propagation for extension-contributed fields.""" + + def test_extension_field_is_flagged(self) -> None: + class Target(BaseModel): + name: str + + @extends(Target) + class Ext(BaseModel): + note: str + + wrapper = wrap_extension("ext", Ext) + assert wrapper is not None + extended = create_extended_model(Target, {"ext": wrapper}) + + spec = extract_model(extended) + + assert find_field(spec, "name").is_extension is False + assert find_field(spec, "ext").is_extension is True + + def test_model_without_extensions_has_no_flagged_fields(self) -> None: + class Plain(BaseModel): + name: str + + spec = extract_model(Plain) + + assert all(not field.is_extension for field in spec.fields) diff --git a/packages/overture-schema-codegen/tests/test_type_analyzer.py b/packages/overture-schema-codegen/tests/test_type_analyzer.py index 1fecc4c6c..a1438ee27 100644 --- a/packages/overture-schema-codegen/tests/test_type_analyzer.py +++ b/packages/overture-schema-codegen/tests/test_type_analyzer.py @@ -40,6 +40,7 @@ unwrap_list, ) from overture.schema.common.scoping.vehicle import VehicleSelector +from overture.schema.system.extension import Extends from overture.schema.system.numeric import int32 from overture.schema.system.ref import Id from overture.schema.system.string import ( @@ -198,6 +199,34 @@ def test_non_field_metadata_collected(self) -> None: assert isinstance(shape, Primitive) assert shape.constraints[0].constraint == "just a description" + def test_annotated_extends_metadata_is_not_a_constraint(self) -> None: + """`Extends` declares an extension's targets, not a field constraint. + + Regression test: it must not leak into the shape's constraints, + where it would render in generated docs via its class docstring + (see `field_constraints.py`). + """ + + class Target(BaseModel): + name: str + + shape = _shape(Annotated[int, Field(ge=0), Extends(Target)]) + + assert isinstance(shape, Primitive) + assert len(shape.constraints) == 1 + assert isinstance(shape.constraints[0].constraint, Ge) + + def test_annotated_bare_extends_metadata_is_not_a_constraint(self) -> None: + """`Extends` as the sole (non-FieldInfo) metadata item is also excluded.""" + + class Target(BaseModel): + name: str + + shape = _shape(Annotated[int, Extends(Target)]) + + assert isinstance(shape, Primitive) + assert shape.constraints == () + def test_list_level_minlen_lands_on_arrayof(self) -> None: shape = _shape(Annotated[list[str], Field(min_length=1)]) assert isinstance(shape, ArrayOf) diff --git a/packages/overture-schema-codegen/tests/test_union_extraction.py b/packages/overture-schema-codegen/tests/test_union_extraction.py index a543c63c4..0b0de60c2 100644 --- a/packages/overture-schema-codegen/tests/test_union_extraction.py +++ b/packages/overture-schema-codegen/tests/test_union_extraction.py @@ -1,7 +1,7 @@ """Tests for union extraction.""" import re -from typing import Any +from typing import Annotated, Any import pytest from annotated_types import MinLen @@ -26,6 +26,7 @@ from overture.schema.codegen.extraction.specs import FieldSpec, UnionSpec from overture.schema.codegen.extraction.union_extraction import ( _constraints_fingerprint, + extract_discriminator, extract_union, ) from overture.schema.common.scoping.vehicle import VehicleSelector @@ -38,10 +39,21 @@ JsonPointerConstraint, PatternConstraint, ) -from pydantic import Field, GetCoreSchemaHandler +from pydantic import Discriminator, Field, GetCoreSchemaHandler from pydantic_core import core_schema +def test_extract_discriminator_resolves_discriminator_object() -> None: + """Field name resolution handles the `pydantic.Discriminator` object form.""" + ann = Annotated[ + RoadSegment | RailSegment, + Field(discriminator=Discriminator("subtype")), + ] + field_name, mapping = extract_discriminator(ann, [RoadSegment, RailSegment]) + assert field_name == "subtype" + assert mapping == {"road": RoadSegment, "rail": RailSegment} + + class TestExtractUnion: """Tests for extract_union function.""" diff --git a/packages/overture-schema-system/README.md b/packages/overture-schema-system/README.md index 84ff7b86f..835f7284e 100644 --- a/packages/overture-schema-system/README.md +++ b/packages/overture-schema-system/README.md @@ -157,12 +157,12 @@ Tags classify discovered models. A package registers [tag providers](#providers) from overture.schema.system.discovery import ( TagSelector, discover_models, - filter_models, + select_models, ) models = discover_models() -selected = filter_models( +selected = select_models( models, TagSelector(include_any=("feature",), exclude_any=("draft",)), ) @@ -233,13 +233,15 @@ When a provider attempts to set a reserved tag from an unauthorized package, dis ### Selecting Models by Tag -`filter_models(models, selector)` applies `TagSelector` predicates against each `ModelKey.tags`: +`select_models(models, selector)` applies `TagSelector` predicates against each `ModelKey.tags`: - `include_any` -- OR scope; at least one tag must match (empty: no scope filter) - `require_all` -- AND narrowing; every tag must be present (empty: no narrowing) - `exclude_any` -- OR-NOT subtraction; any match drops the model -An empty selector returns the input unchanged. +An empty selector returns the input unchanged, except that standalone extension +entries are hidden by default; pass `include_extension_entries=True` (or engage the +`extension` tag in the selector) to include them. ## Also Included diff --git a/packages/overture-schema-system/changelog.d/634.breaking.md b/packages/overture-schema-system/changelog.d/634.breaking.md new file mode 100644 index 000000000..568851401 --- /dev/null +++ b/packages/overture-schema-system/changelog.d/634.breaking.md @@ -0,0 +1 @@ +Replaced `filter_models` with `select_models`, which hides standalone extension entries by default (lift with `include_extension_entries=True` or by engaging the `extension` tag). diff --git a/packages/overture-schema-system/changelog.d/634.feature.md b/packages/overture-schema-system/changelog.d/634.feature.md new file mode 100644 index 000000000..b9d4f78ed --- /dev/null +++ b/packages/overture-schema-system/changelog.d/634.feature.md @@ -0,0 +1 @@ +Added a generic model extension mechanism: a package can contribute optional fields to models it does not own by declaring targets with `@extends` (model extensions) or `Extends(...)` metadata in `Annotated` (scalar extensions) and registering on the `overture.models` entry-point group; discovery wraps each extension into a standalone wrapper model and merges its field into every targeted model. diff --git a/packages/overture-schema-system/pyproject.toml b/packages/overture-schema-system/pyproject.toml index 761a3011c..2b3129150 100644 --- a/packages/overture-schema-system/pyproject.toml +++ b/packages/overture-schema-system/pyproject.toml @@ -56,3 +56,4 @@ per-file-ignores = {"__init__.py" = ["F401"]} [project.entry-points."overture.tag_providers"] feature = "overture.schema.system.discovery.tag_providers:feature_provider" +extension = "overture.schema.system.discovery.tag_providers:extension_provider" diff --git a/packages/overture-schema-system/src/overture/schema/system/discovery/__init__.py b/packages/overture-schema-system/src/overture/schema/system/discovery/__init__.py index c894f591e..45bfef8a0 100644 --- a/packages/overture-schema-system/src/overture/schema/system/discovery/__init__.py +++ b/packages/overture-schema-system/src/overture/schema/system/discovery/__init__.py @@ -2,8 +2,9 @@ from .discovery import ( TagSelector, discover_models, - filter_models, + extend_models, get_registered_model, + select_models, ) from .entry_point import ( entry_point_class_alias, @@ -21,7 +22,8 @@ "discover_models", "entry_point_class_alias", "entry_point_to_path", - "filter_models", + "extend_models", + "select_models", "get_registered_model", "resolve_entry_point_key", "split_entry_point", diff --git a/packages/overture-schema-system/src/overture/schema/system/discovery/discovery.py b/packages/overture-schema-system/src/overture/schema/system/discovery/discovery.py index 2425a7551..690b688ff 100644 --- a/packages/overture-schema-system/src/overture/schema/system/discovery/discovery.py +++ b/packages/overture-schema-system/src/overture/schema/system/discovery/discovery.py @@ -3,7 +3,7 @@ import importlib.metadata import logging from dataclasses import dataclass, replace -from typing import Any +from typing import TypeGuard from pydantic import BaseModel @@ -17,13 +17,24 @@ TagProviderDict, TagProviderKey, ) +from overture.schema.system.extension import ( + applied_extension_names, + create_extended_model, + extension_targets, + wrap_extension, +) from overture.schema.system.typing_util import collect_types log = logging.getLogger(__name__) +# Set by `extension_provider`; engaging it in a selector lifts `select_models`' +# default hiding. +_EXTENSION_TAG = "extension" + # Tags that are reserved and can only be set by specific packages. _RESERVED_TAGS: dict[str, set[str]] = { "feature": {"overture-schema-system"}, + _EXTENSION_TAG: {"overture-schema-system"}, } # Namespaces that are reserved and can only be set by specific packages. _RESERVED_NAMESPACES: dict[str, set[str]] = { @@ -33,7 +44,7 @@ def _generate_tags( - model_class: Any, # noqa: ANN401 + model_class: object, key: ModelKey, providers: TagProviderDict, ) -> set[str]: @@ -161,51 +172,160 @@ def discover_tag_providers( def discover_models( model_group: str = "overture.models", + *, + apply_extensions: bool = True, ) -> ModelDict: """Discover and load models via entry points, attaching tags from tag providers. + A three-stage pipeline -- load and wrap, merge extensions across the complete registry, + then generate tags from the final model values -- so tags always describe the classes + this function actually returns. + + Extension entry points are wrapped into standalone wrapper models at load time. When + `apply_extensions` is true (the default), each extension's field is also merged into the + target models it extends. + Parameters ---------- model_group : str, optional Entry point group to search (default: `"overture.models"`). + apply_extensions : bool, optional + Whether to merge discovered extensions into their target models (default: ``True``). Pass + ``False`` to obtain the raw (un-extended) model set. Returns ------- ModelDict Discovered models keyed by ModelKey. """ - models = {} - tag_providers = discover_tag_providers() + models = _load_models(model_group) + if apply_extensions: + models = extend_models(models) + return _attach_generated_tags(models) + + +def _load_models(model_group: str) -> ModelDict: + """Load and wrap model entry points, keyed with empty tag sets. + + Load and wrap failures are reported distinctly, each skipping only its entry. + """ + models: ModelDict = {} try: - for model in importlib.metadata.entry_points(group=model_group): - try: - model_class = model.load() - key = ModelKey( - name=model.name, - entry_point=model.value, - tags=frozenset(), - ) - try: - key = replace( - key, - tags=frozenset(_generate_tags(model_class, key, tag_providers)), - ) - except Exception as e: - log.warning(f"Could not resolve tags for model {model.name}: {e}") - models[key] = model_class - except Exception as e: - log.warning(f"Could not load model {model.name}: {e}") + entries = importlib.metadata.entry_points(group=model_group) except Exception as e: log.warning(f"Could not discover entry points: {e}") + return models + for entry in entries: + try: + loaded = entry.load() + except Exception as e: + log.warning(f"Could not load model {entry.name}: {e}") + continue + try: + wrapper = wrap_extension( + entry.name, loaded, module=entry.value.partition(":")[0] or None + ) + except Exception as e: + log.warning(f"Could not wrap extension entry {entry.name}: {e}") + continue + if wrapper is not None: + loaded = wrapper + models[ModelKey(name=entry.name, entry_point=entry.value, tags=frozenset())] = ( + loaded + ) return models +def _attach_generated_tags(models: ModelDict) -> ModelDict: + """Re-key *models* with tags generated from the final model values.""" + tag_providers = discover_tag_providers() + tagged: ModelDict = {} + for key, model in models.items(): + tags: frozenset[str] = frozenset() + try: + tags = frozenset(_generate_tags(model, key, tag_providers)) + except Exception as e: + log.warning(f"Could not resolve tags for model {key.name}: {e}") + tagged[replace(key, tags=tags)] = model + return tagged + + +def extend_models(models: ModelDict) -> ModelDict: + """Merge discovered extensions into the target models they extend. + + Extensions are detected structurally (a wrapper model carrying extension targets), so this works + on any `ModelDict` regardless of how its tags were filtered. Extension entries themselves are + left unchanged; every other model is replaced by an extended subclass if any extension targets + it. + + Parameters + ---------- + models : ModelDict + Models to process, as returned by `discover_models(apply_extensions=False)`. + + Returns + ------- + ModelDict + Models with extension fields merged into their targets. + """ + + def is_extension(model: object) -> TypeGuard[type[BaseModel]]: + return ( + isinstance(model, type) + and issubclass(model, BaseModel) + and bool(extension_targets(model)) + ) + + extension_candidates: dict[str, list[tuple[ModelKey, type[BaseModel]]]] = {} + for key, model in models.items(): + if is_extension(model): + extension_candidates.setdefault(key.name, []).append((key, model)) + + extensions: dict[str, type[BaseModel]] = {} + for field_name, candidates in extension_candidates.items(): + if len(candidates) > 1: + entry_points = ", ".join(sorted(key.entry_point for key, _ in candidates)) + log.warning( + f"Multiple extensions are registered for field '{field_name}' " + f"({entry_points}); skipping all of them." + ) + continue + extensions[field_name] = candidates[0][1] + if not extensions: + return models + # Skip by the entry's own shape, not by name: a non-extension model that merely + # shares an entry-point name with an extension must still be extended. A failure + # to extend one model must not abort the pass for every other model. + extended: ModelDict = {} + for key, model in models.items(): + if is_extension(model): + extended[key] = model + continue + try: + extended[key] = create_extended_model(model, extensions) + except Exception as e: + log.warning( + f"Could not apply extensions to model '{key.name}' " + f"({key.entry_point}): {e}; leaving it unextended." + ) + extended[key] = model + applied_names = frozenset().union( + *(applied_extension_names(model) for model in extended.values()) + ) + for field_name in sorted(set(extensions) - applied_names): + log.warning( + f"Extension '{field_name}' was not applied to any discovered model; " + "its targets may not be registered via entry points." + ) + return extended + + @dataclass(frozen=True, slots=True, kw_only=True) class TagSelector: - """Three tag tuples consumed by `filter_models`. + """Three tag tuples consumed by `select_models`. - See `filter_models` for predicate semantics, including how - empty tuples are interpreted. + See `matches` for predicate semantics, including how empty tuples + are interpreted. Attributes ---------- @@ -221,8 +341,16 @@ class TagSelector: require_all: tuple[str, ...] = () exclude_any: tuple[str, ...] = () + def matches(self, tags: frozenset[str]) -> bool: + """Whether *tags* satisfies every predicate tuple (empty tuples are no-ops).""" + return ( + (not self.include_any or not tags.isdisjoint(self.include_any)) + and tags.issuperset(self.require_all) + and tags.isdisjoint(self.exclude_any) + ) + -def filter_models( +def _filter_models( models: ModelDict, selector: TagSelector = TagSelector(), *, @@ -230,44 +358,65 @@ def filter_models( ) -> ModelDict: """Filter models by tag predicates and optional type-name match. - Each tuple in `selector` is a predicate over `key.tags`; a model - is kept only if it satisfies every predicate. Empty tuples are - no-ops — empty `include_any` imposes no scope, empty - `require_all` imposes no narrowing, empty `exclude_any` drops - nothing — so an empty selector returns `models` unchanged. + The predicate stage of `select_models`, the public entry point; + `select_models(..., include_extension_entries=True)` reproduces this + function exactly. + """ + names = frozenset(type_names) + return { + key: model + for key, model in models.items() + if (not names or key.name in names) and selector.matches(key.tags) + } - Parameters - ---------- - models - Models to filter. - selector - Tag predicates to apply. - type_names - If non-empty, only models whose `key.name` is in the list - are kept. Orthogonal to the tag predicate algebra. - Returns - ------- - ModelDict - Models satisfying every supplied predicate. - """ +def _is_extension_entry(model: object) -> bool: + """Whether a registry entry structurally declares extension targets.""" + return bool(extension_targets(model)) - def matches(key: ModelKey) -> bool: - if selector.include_any and not any( - t in key.tags for t in selector.include_any - ): - return False - if selector.require_all and not all( - t in key.tags for t in selector.require_all - ): - return False - if selector.exclude_any and any(t in key.tags for t in selector.exclude_any): - return False - if type_names and key.name not in type_names: - return False - return True - - return {k: m for k, m in models.items() if matches(k)} + +def select_models( + models: ModelDict, + selector: TagSelector = TagSelector(), + *, + type_names: tuple[str, ...] = (), + include_extension_entries: bool = False, +) -> ModelDict: + """Select models: user predicates plus default extension-entry hiding. + + `selector` and `type_names` apply first and are never overridden. On top + of that, standalone extension entries are dropped by default, detected + structurally (declared extension targets) rather than by the + ``extension`` tag, so an entry whose tag generation failed stays hidden. + The default is lifted: wholesale via ``include_extension_entries=True``; + query-wide when the selector engages the ``extension`` tag in + `include_any`/`require_all`; or per name, when every entry carrying a + requested type name is an extension (an ambiguous name keeps the + wrapper hidden). + """ + selected = _filter_models(models, selector, type_names=type_names) + if ( + include_extension_entries + or _EXTENSION_TAG in selector.include_any + or _EXTENSION_TAG in selector.require_all + ): + return selected + + # Resolved against the unfiltered dict: name ambiguity must consider + # entries the predicates already removed. + extension_keys = { + key for key, model in models.items() if _is_extension_entry(model) + } + extension_names = {key.name for key in extension_keys} + non_extension_names = {key.name for key in models if key not in extension_keys} + explicitly_requested = frozenset(type_names) & ( + extension_names - non_extension_names + ) + return { + key: model + for key, model in selected.items() + if key not in extension_keys or key.name in explicitly_requested + } def get_registered_model(model_name: str) -> type[BaseModel] | None: diff --git a/packages/overture-schema-system/src/overture/schema/system/discovery/tag_providers.py b/packages/overture-schema-system/src/overture/schema/system/discovery/tag_providers.py index 46fc507d8..d8e808d83 100644 --- a/packages/overture-schema-system/src/overture/schema/system/discovery/tag_providers.py +++ b/packages/overture-schema-system/src/overture/schema/system/discovery/tag_providers.py @@ -4,7 +4,9 @@ from pydantic import BaseModel +from overture.schema.system.discovery.discovery import _EXTENSION_TAG from overture.schema.system.discovery.types import ModelKey +from overture.schema.system.extension import extension_targets from overture.schema.system.feature import Feature @@ -33,3 +35,30 @@ def feature_provider( if any(issubclass(tp, Feature) for tp in types): tags.add("feature") return tags + + +def extension_provider( + types: Iterable[type[BaseModel]], + key: ModelKey, + tags: set[str], +) -> set[str]: + """Add the `"extension"` tag if any concrete type declares it extends a target model. + + Parameters + ---------- + types + Concrete `BaseModel` subclasses for the entry point. Extension entry points are wrapped + into a single wrapper model at load time, so this is the wrapper carrying the targets. + key + Key identifying the model. + tags + Current tags; may be extended. + + Returns + ------- + set[str] + Updated tags, with `"extension"` added if applicable. + """ + if any(extension_targets(tp) for tp in types): + tags.add(_EXTENSION_TAG) + return tags diff --git a/packages/overture-schema-system/src/overture/schema/system/extension.py b/packages/overture-schema-system/src/overture/schema/system/extension.py new file mode 100644 index 000000000..1537f53d0 --- /dev/null +++ b/packages/overture-schema-system/src/overture/schema/system/extension.py @@ -0,0 +1,514 @@ +"""Add optional extension fields to existing discovered Pydantic models. + +Model extensions declare their targets with `@extends`: + +>>> class Place(BaseModel): +... name: str +... +>>> @extends(Place) +... class OperatingHours(BaseModel): +... primary: list[str] + +Other types use `Extends` in `Annotated` metadata, optionally through +`NewType`: + +>>> Capacity = NewType("Capacity", Annotated[int, Extends(Place)]) + +Targets may be models or model-bearing unions, `Annotated`, `NewType`, and +`RootModel` expressions. A `RootModel` is treated as an alias for its root +in both directions: + +>>> class RoadSegment(BaseModel): ... +>>> class RailSegment(BaseModel): ... +>>> class Segment(RootModel): +... root: RoadSegment | RailSegment + +Targeting `Segment` with `@extends(Segment)` extends `RoadSegment` and +`RailSegment`, not `Segment` itself. Conversely, when `Segment` is a +registered entry, `create_extended_model` rebuilds it as a subclass whose +root annotation has the extended arms. The alias view also means a +`RootModel` over a non-model root is not a valid target: + +>>> class Version(RootModel): +... root: int +>>> Extends(Version) # doctest: +ELLIPSIS +Traceback (most recent call last): + ... +TypeError: `Extends` targets must be (or resolve to) pydantic `BaseModel` ... + +`create_extended_model` likewise leaves `Version` untouched: its scalar +root contains nothing to extend. Types nested inside containers are not +traversed. + +During discovery, each extension is exposed as a standalone one-field wrapper. +`create_extended_model` adds that field to every matching target model. +""" + +import keyword +import logging +import types +from collections.abc import Callable, Iterable, Iterator, Mapping +from typing import Annotated, Any, NewType, TypeVar, Union, get_args, get_origin + +from pydantic import AliasChoices, AliasPath, BaseModel, RootModel + +from overture.schema.system.create_model import create_model + +log = logging.getLogger(__name__) + +__all__ = [ + "Extends", + "SelfReferentialRootError", + "applied_extension_names", + "applied_extensions", + "create_extended_model", + "extends", + "extension_targets", + "wrap_extension", +] + + +class SelfReferentialRootError(TypeError): + """Raised when a `RootModel`'s root annotation reaches the model itself. + + Such a root has no finite shape, so it can be neither validated as an + extension target nor rewritten by extension application. A dedicated + class lets callers that must tolerate exactly this case (e.g. warning + aggregation) catch it without swallowing unrelated `TypeError`s. + """ + + +# Class attributes this mechanism sets: qualified non-identifier names (the +# `Metadata` facility's convention), reachable only via getattr/setattr. + +# On a wrapper model: the original (unwrapped) extension type expression. +_EXTENSION_ATTR = "_[overture.schema.system.extension]__extension" +# On an extended model: the names of the extensions it already carries. +_APPLIED_ATTR = "_[overture.schema.system.extension]__applied_extensions" +# On a model extension: the `Extends` metadata declared via `@extends`. +_EXTENDS_ATTR = "_[overture.schema.system.extension]__extends" + + +def applied_extensions(model_class: type[BaseModel]) -> frozenset[str]: + """Return the names of the extensions `create_extended_model` has merged into `model_class`. + + Empty for a model with no extensions applied. Callers that need to know *whether* a + particular field on a model came from an extension (e.g. to annotate generated docs) should + check for that field's name in the returned set. + """ + return getattr(model_class, _APPLIED_ATTR, frozenset()) + + +def applied_extension_names(obj: Any) -> frozenset[str]: # noqa: ANN401 + """Return the names of the extensions applied anywhere in a model-bearing type expression. + + Aggregates `applied_extensions` over every `BaseModel` leaf of `obj`, so it also covers + union/`Annotated`/`NewType`/`RootModel` registry entries whose arms were extended + individually. A self-referential `RootModel` yields the empty set instead of raising: + such an entry cannot have been extended (`create_extended_model` refuses it), and this + function feeds warning aggregation, which must not abort the extension pass. + """ + try: + classes = _unwrap_model_classes(obj) + except SelfReferentialRootError: + return frozenset() + names: frozenset[str] = frozenset() + for cls in classes: + names |= applied_extensions(cls) + return names + + +def _dedupe_targets(targets: Iterable[Any]) -> tuple[Any, ...]: + """Deduplicate targets by equality, preserving order (targets may be unhashable).""" + merged: list[Any] = [] + for target in targets: + if target not in merged: + merged.append(target) + return tuple(merged) + + +def _self_referential_root_error(tp: type) -> SelfReferentialRootError: + return SelfReferentialRootError( + f"self-referential `RootModel` `{tp.__name__}` has no finite root shape" + ) + + +def _unwrap_model_classes( + tp: Any, # noqa: ANN401 + _seen: frozenset[type] = frozenset(), +) -> tuple[type[BaseModel], ...]: + """Collect the concrete `BaseModel` classes a type expression resolves to. + + Unwraps `Annotated`, `Union` (including `X | Y`), `NewType`, and `RootModel` (an alias over + its root value, so its root's classes are collected instead of the RootModel itself). + Non-model leaves are ignored. + """ + origin = get_origin(tp) + if origin is Annotated: + return _unwrap_model_classes(get_args(tp)[0], _seen) + if origin is Union or origin is types.UnionType: + classes: list[type[BaseModel]] = [] + for arg in get_args(tp): + classes.extend(_unwrap_model_classes(arg, _seen)) + return tuple(classes) + if hasattr(tp, "__supertype__"): + return _unwrap_model_classes(tp.__supertype__, _seen) + if isinstance(tp, type) and issubclass(tp, RootModel): + if tp in _seen: + raise _self_referential_root_error(tp) + return _unwrap_model_classes(tp.model_fields["root"].annotation, _seen | {tp}) + if isinstance(tp, type) and issubclass(tp, BaseModel): + return (tp,) + return () + + +def _is_model_target( + tp: Any, # noqa: ANN401 + _seen: frozenset[type] = frozenset(), +) -> bool: + """Whether a type expression resolves *entirely* to extendable `BaseModel` subclasses. + + Unlike `_unwrap_model_classes`, a union must have *every* arm resolve to a model for the whole + expression to qualify — a partially-model union (e.g. `Place | int`) is not a valid target. + A `RootModel` qualifies exactly when its root does: it is an alias over its root value, so + `RootModel[Road | Rail]` targets the arms while `RootModel[int]` resolves to no model. + """ + origin = get_origin(tp) + if origin is Annotated: + return _is_model_target(get_args(tp)[0], _seen) + if origin is Union or origin is types.UnionType: + args = get_args(tp) + return bool(args) and all(_is_model_target(arg, _seen) for arg in args) + if hasattr(tp, "__supertype__"): + return _is_model_target(tp.__supertype__, _seen) + if isinstance(tp, type) and issubclass(tp, RootModel): + if tp in _seen: + raise _self_referential_root_error(tp) + return _is_model_target(tp.model_fields["root"].annotation, _seen | {tp}) + return isinstance(tp, type) and issubclass(tp, BaseModel) + + +def _validate_targets(name: str, targets: tuple[Any, ...]) -> None: + if not targets: + raise TypeError(f"`{name}` requires at least one target model") + for target in targets: + if not _is_model_target(target): + raise TypeError( + f"`{name}` targets must be (or resolve to) pydantic `BaseModel` subclasses, " + f"but {target!r} does not qualify. A `RootModel` target is an alias over its " + f"root value, so its root must itself resolve to `BaseModel` subclasses." + ) + + +class Extends: + """ + Metadata class for declaring, via `typing.Annotated`, which models a non-model extension targets. + + Use this for extensions that are not themselves Pydantic models (e.g. a scalar `NewType`). For + model extensions, prefer the `@extends` decorator. + """ + + def __init__(self, *targets: Any) -> None: # noqa: ANN401 + _validate_targets(type(self).__name__, targets) + self.__targets = targets + + @property + def extends(self) -> tuple[Any, ...]: + return self.__targets + + +ModelT = TypeVar("ModelT", bound=BaseModel) + + +def extends(*targets: Any) -> Callable[[type[ModelT]], type[ModelT]]: # noqa: ANN401 + """ + Decorate a Pydantic model to declare it is an extension of one or more target models. + + Parameters + ---------- + targets + One or more target model classes (or expressions resolving to `BaseModel` subclasses). + + Returns + ------- + Callable + A decorator that stashes the targets on the model as `Extends` metadata, + introspectable via `extension_targets`; the decorated class is returned unchanged. + Stacked decorators merge their target sets; a subclass's own declaration shadows + an inherited one. + """ + metadata = Extends(*targets) # also validates the targets + + def decorator(model_class: type[ModelT]) -> type[ModelT]: + if not (isinstance(model_class, type) and issubclass(model_class, BaseModel)): + raise TypeError( + f"`@{extends.__name__}` can only be applied to pydantic `BaseModel` subclasses, " + f"but {model_class!r} is not one" + ) + combined = metadata + existing = model_class.__dict__.get(_EXTENDS_ATTR) + if isinstance(existing, Extends): + # Decorators apply bottom-up: `existing` came from the inner + # (earlier) declaration, so its targets stay first. + combined = Extends(*_dedupe_targets((*existing.extends, *metadata.extends))) + setattr(model_class, _EXTENDS_ATTR, combined) + return model_class + + return decorator + + +def _find_extends_metadata(tp: Any) -> tuple[Any, ...]: # noqa: ANN401 + """Find `Extends` metadata attached to a `NewType`/`Annotated` expression. + + Multiple declarations in one `Annotated` frame merge; the nearest frame + declaring any `Extends` decides. + """ + if hasattr(tp, "__supertype__"): + return _find_extends_metadata(tp.__supertype__) + if get_origin(tp) is Annotated: + args = get_args(tp) + found = _dedupe_targets( + target + for meta in args[1:] + if isinstance(meta, Extends) + for target in meta.extends + ) + if found: + return found + return _find_extends_metadata(args[0]) + return () + + +def extension_targets(obj: Any) -> tuple[Any, ...]: # noqa: ANN401 + """ + Return the target models an extension declares, or `()` if `obj` is not an extension. + + Detects all three declaration forms: a `BaseModel` carrying `Extends` metadata (via `@extends`), + a `NewType` over `Annotated[..., Extends(...)]`, and a bare `Annotated[..., Extends(...)]`. + """ + if isinstance(obj, type) and issubclass(obj, BaseModel): + metadata = getattr(obj, _EXTENDS_ATTR, None) + return metadata.extends if isinstance(metadata, Extends) else () + return _find_extends_metadata(obj) + + +def _wrapper_name(name: str) -> str: + return "".join(part.capitalize() for part in name.split("_")) + "Extension" + + +def _validate_extension_name(name: str) -> None: + """Reject entry-point names that cannot safely become a Pydantic field name. + + Covers non-identifiers/keywords, underscore- and ``model_``-prefixed names + (private / protected in Pydantic), and `BaseModel` attribute names. + """ + if ( + not name.isidentifier() + or keyword.iskeyword(name) + or name.startswith("_") + or name.startswith("model_") + or hasattr(BaseModel, name) + ): + raise ValueError( + f"extension entry-point name {name!r} cannot be used as a field name: it must be " + "a valid Python identifier, must not be a keyword or a `BaseModel` attribute name, " + "and must not start with '_' or 'model_'" + ) + + +def wrap_extension( + name: str, + obj: Any, # noqa: ANN401 + *, + module: str | None = None, +) -> type[BaseModel] | None: + """ + Wrap an extension entry-point value into a standalone wrapper model. + + The wrapper is a `BaseModel` with a single optional field named `name` holding the + extension type. It carries `Extends` metadata and stashes the original (unwrapped) + extension type (`_EXTENSION_ATTR`) so `create_extended_model` can reproduce the exact + field annotation. Its `__module__` is the extension's defining module; a bare + `Annotated[...]` has none of its own, so *module* (the entry point's module) fills the + gap, falling back to this module. + + Returns `None` if `obj` is not an extension. Raises `ValueError` if `obj` is an extension + but `name` cannot be used as a field name. + """ + targets = extension_targets(obj) + if not targets: + return None + _validate_extension_name(name) + if get_origin(obj) is Annotated: + # Annotated proxies attribute access to the wrapped type. + owner_module = module or __name__ + else: + owner_module = getattr(obj, "__module__", None) or module or __name__ + wrapper: type[BaseModel] = create_model( + _wrapper_name(name), + __module__=owner_module, + __doc__=f"Standalone wrapper model for the `{name}` extension.", + **{name: (obj | None, None)}, # type: ignore[arg-type] + ) + setattr(wrapper, _EXTENDS_ATTR, Extends(*targets)) + setattr(wrapper, _EXTENSION_ATTR, obj) + return wrapper + + +def _alias_strings(alias: object) -> Iterator[str]: + """Yield the payload-level key names an alias declaration can claim.""" + if isinstance(alias, str): + yield alias + elif isinstance(alias, AliasPath): + first = alias.path[0] if alias.path else None + if isinstance(first, str): + yield first + elif isinstance(alias, AliasChoices): + for choice in alias.choices: + yield from _alias_strings(choice) + + +def _occupied_names(model: type[BaseModel]) -> set[str]: + """Names an extension field may not use on `model`. + + Covers declared fields and their aliases, computed fields and their aliases, and every + attribute defined anywhere on the class hierarchy (methods, properties, private + attributes) — adding a field over any of these would break validation or shadow + existing behavior. + """ + names: set[str] = set(model.model_fields) + for field_info in model.model_fields.values(): + for alias in ( + field_info.alias, + field_info.validation_alias, + field_info.serialization_alias, + ): + names.update(_alias_strings(alias)) + names.update(model.model_computed_fields) + for computed_info in model.model_computed_fields.values(): + names.update(_alias_strings(computed_info.alias)) + for klass in model.__mro__: + names.update(vars(klass)) + return names + + +def create_extended_model( + model: Any, # noqa: ANN401 + extensions: Mapping[str, type[BaseModel]], + _seen: frozenset[type] = frozenset(), +) -> Any: # noqa: ANN401 + """ + Apply extension wrappers to a model (or model-bearing type expression). + + Recurses through `Annotated`, `Union`, `NewType`, and `RootModel` (rewriting the root + annotation in place) so discriminated-union entry points have every arm extended. For each + concrete `BaseModel`, an optional field is added per extension whose targets the model is a + subclass of. Returns the original expression unchanged when nothing applies, so identity is + preserved for untouched types. Container types (`list[...]`, `dict[...]`) are opaque: models + inside them are not rewritten. + + Parameters + ---------- + model + The model or type expression to extend. + extensions + Mapping of extension field name to wrapper model (as produced by `wrap_extension`). + A raw `@extends` model class is also accepted and applies as its own field type. + """ + origin = get_origin(model) + + if origin is Annotated: + tp, *metadata = get_args(model) + extended = create_extended_model(tp, extensions, _seen) + if extended is tp: + return model + return Annotated.__class_getitem__((extended, *metadata)) # type: ignore[attr-defined] + + if origin is Union or origin is types.UnionType: + args = get_args(model) + extended_args = tuple( + create_extended_model(arg, extensions, _seen) for arg in args + ) + if all(new is old for new, old in zip(extended_args, args, strict=True)): + return model + # Rebuild via `Union[...]` rather than `reduce(or_, ...)`: `|` raises + # on arms that don't implement it (e.g. an unresolved `ForwardRef`), + # while `Union` accepts any type argument. + return Union[extended_args] # noqa: UP007 + + if hasattr(model, "__supertype__"): + supertype = model.__supertype__ + extended = create_extended_model(supertype, extensions, _seen) + if extended is supertype: + return model + extended_alias = NewType(model.__name__, extended) # type: ignore[misc, valid-type] + # `NewType` stamps `__module__` from the calling frame and resets + # `__qualname__`/`__doc__`; restore the original alias's identity so + # the rebuilt alias doesn't appear to originate here and keeps any + # custom docstring (which codegen renders). + extended_alias.__module__ = model.__module__ + extended_alias.__qualname__ = model.__qualname__ + extended_alias.__doc__ = model.__doc__ + return extended_alias + + if not (isinstance(model, type) and issubclass(model, BaseModel)): + return model + if issubclass(model, RootModel): + # A RootModel is an alias over its root value: rewrite the root annotation and + # rebuild as a subclass, reusing the root FieldInfo so its metadata, default, + # and requiredness carry over. + if model in _seen: + raise _self_referential_root_error(model) + root_field = model.model_fields["root"] + extended = create_extended_model( + root_field.annotation, extensions, _seen | {model} + ) + if extended is root_field.annotation: + return model + return create_model( + model.__name__, + __base__=model, + __module__=model.__module__, + __doc__=model.__doc__, + root=(extended, root_field), + ) + + applied: frozenset[str] = getattr(model, _APPLIED_ATTR, frozenset()) + occupied: set[str] | None = None + fields: dict[str, tuple[Any, None]] = {} + for field_name, wrapper in extensions.items(): + if field_name in applied: + continue + if not any( + issubclass(model, cls) + for target in extension_targets(wrapper) + for cls in _unwrap_model_classes(target) + ): + continue + if occupied is None: + occupied = _occupied_names(model) + if field_name in occupied: + log.warning( + "Extension '%s' collides with an existing field, alias, or attribute on " + "model '%s'; skipping.", + field_name, + model.__name__, + ) + continue + # A wrapper built by `wrap_extension` stashes the original extension type; + # a raw `@extends` model is itself the extension type. + extension_type = getattr(wrapper, _EXTENSION_ATTR, wrapper) + fields[field_name] = (extension_type | None, None) + + if not fields: + return model + + extended_model = create_model( + model.__name__, + __base__=model, + __doc__=model.__doc__, + __module__=model.__module__, + **fields, # type: ignore[arg-type] + ) + setattr(extended_model, _APPLIED_ATTR, applied | frozenset(fields)) + return extended_model diff --git a/packages/overture-schema-system/tests/test_discovery_extension.py b/packages/overture-schema-system/tests/test_discovery_extension.py new file mode 100644 index 000000000..ee144a153 --- /dev/null +++ b/packages/overture-schema-system/tests/test_discovery_extension.py @@ -0,0 +1,348 @@ +"""Tests for extension handling in the discovery pipeline (no entry points required).""" + +import logging +from collections.abc import Iterable +from typing import Annotated, Literal, cast + +import pytest +from pydantic import BaseModel, Field, Tag, TypeAdapter + +from overture.schema.system.discovery import discovery as discovery_module +from overture.schema.system.discovery.discovery import _generate_tags, extend_models +from overture.schema.system.discovery.keys import ModelKey, TagProviderKey +from overture.schema.system.discovery.tag_providers import extension_provider +from overture.schema.system.discovery.types import ModelDict +from overture.schema.system.extension import extends, wrap_extension +from overture.schema.system.typing_util import collect_types + + +class Target(BaseModel): + name: str + + +class Unrelated(BaseModel): + value: int + + +class RoadSeg(BaseModel): + kind: Literal["road"] = "road" + + +class RailSeg(BaseModel): + kind: Literal["rail"] = "rail" + + +# A discriminated-union alias entry point, like the transportation theme's `Segment`. +SegmentAlias = Annotated[ + Annotated[RoadSeg, Tag("road")] | Annotated[RailSeg, Tag("rail")], + Field(discriminator="kind"), +] + + +@extends(Target) +class Ext(BaseModel): + note: str + + +def _key(name: str, *tags: str) -> ModelKey: + return ModelKey(name=name, entry_point=f"m:{name}", tags=frozenset(tags)) + + +def _model(value: object) -> type[BaseModel]: + assert isinstance(value, type) and issubclass(value, BaseModel) + return value + + +def test_extend_models_merges_into_targets() -> None: + wrapper = wrap_extension("ext", Ext) + assert wrapper is not None + models: ModelDict = { + _key("ext", "extension"): wrapper, + _key("target", "feature"): Target, + _key("unrelated"): Unrelated, + } + result = extend_models(models) + + by_name = {key.name: model for key, model in result.items()} + # Extension entry is left untouched. + assert by_name["ext"] is wrapper + # Target gains the field; unrelated model is unchanged. + assert "ext" in _model(by_name["target"]).model_fields + assert by_name["unrelated"] is Unrelated + + +def test_extend_models_noop_without_extensions() -> None: + models: ModelDict = {_key("target", "feature"): Target} + assert extend_models(models) is models + + +def test_extend_models_extends_model_sharing_an_extension_name() -> None: + wrapper = wrap_extension("ext", Ext) + assert wrapper is not None + # A plain model registered under the same entry-point name as the extension + # (the keys differ by entry_point, so both coexist) must still be extended. + models: ModelDict = { + _key("ext", "extension"): wrapper, + ModelKey(name="ext", entry_point="other:Target", tags=frozenset()): Target, + } + result = extend_models(models) + + by_entry_point = {key.entry_point: model for key, model in result.items()} + assert by_entry_point["m:ext"] is wrapper + assert "ext" in _model(by_entry_point["other:Target"]).model_fields + + +def test_extend_models_skips_duplicate_field_names( + caplog: pytest.LogCaptureFixture, +) -> None: + @extends(Target) + class OtherExt(BaseModel): + other: int + + first = wrap_extension("ext", Ext) + second = wrap_extension("ext", OtherExt) + assert first is not None and second is not None + first_key = ModelKey( + name="ext", entry_point="first:Ext", tags=frozenset({"extension"}) + ) + second_key = ModelKey( + name="ext", entry_point="second:OtherExt", tags=frozenset({"extension"}) + ) + target_key = _key("target") + + with caplog.at_level(logging.WARNING): + result = extend_models( + {first_key: first, second_key: second, target_key: Target} + ) + + assert result[target_key] is Target + assert "Multiple extensions are registered for field 'ext'" in caplog.text + + +def test_extend_models_isolates_per_model_failures( + caplog: pytest.LogCaptureFixture, + monkeypatch: pytest.MonkeyPatch, +) -> None: + wrapper = wrap_extension("ext", Ext) + assert wrapper is not None + + real = discovery_module.create_extended_model + + def failing(model: object, extensions: object) -> object: + if model is Unrelated: + raise RuntimeError("boom") + return real(model, extensions) # type: ignore[arg-type] + + monkeypatch.setattr(discovery_module, "create_extended_model", failing) + models: ModelDict = { + _key("ext", "extension"): wrapper, + _key("target"): Target, + _key("unrelated"): Unrelated, + } + with caplog.at_level(logging.WARNING): + result = extend_models(models) + + by_name = {key.name: model for key, model in result.items()} + # One model failing must not abort the pass: the target is still extended and + # the failing model is kept, unextended, instead of being dropped. + assert "ext" in _model(by_name["target"]).model_fields + assert by_name["unrelated"] is Unrelated + assert "Could not apply extensions to model 'unrelated'" in caplog.text + + +def test_extend_models_warns_when_extension_matches_nothing( + caplog: pytest.LogCaptureFixture, +) -> None: + wrapper = wrap_extension("ext", Ext) + assert wrapper is not None + # Ext targets Target, which is absent from the registry -- the extension would + # otherwise be silently ignored. + models: ModelDict = { + _key("ext", "extension"): wrapper, + _key("unrelated"): Unrelated, + } + with caplog.at_level(logging.WARNING): + result = extend_models(models) + + assert result[_key("unrelated")] is Unrelated + assert "Extension 'ext' was not applied to any discovered model" in caplog.text + + +def test_extend_models_extends_arms_of_rootmodel_entry( + caplog: pytest.LogCaptureFixture, +) -> None: + from pydantic import RootModel + + class SegmentRoot( + RootModel[Annotated[RoadSeg | RailSeg, Field(discriminator="kind")]] + ): + pass + + @extends(RoadSeg) + class SegExt(BaseModel): + note: str + + wrapper = wrap_extension("seg_ext", SegExt) + assert wrapper is not None + models: ModelDict = { + _key("seg_ext", "extension"): wrapper, + _key("segment", "feature"): SegmentRoot, + } + with caplog.at_level(logging.WARNING): + result = extend_models(models) + + by_name = {key.name: model for key, model in result.items()} + rebuilt = by_name["segment"] + assert isinstance(rebuilt, type) and issubclass(rebuilt, SegmentRoot) + # Discriminator survives the rebuild and the targeted arm takes extension data. + assert rebuilt.model_fields["root"].discriminator == "kind" + v = rebuilt.model_validate({"kind": "road", "seg_ext": {"note": "hi"}}) + assert v.model_dump()["seg_ext"]["note"] == "hi" + # The application site is inside the root -- no spurious "not applied" warning. + assert "was not applied to any discovered model" not in caplog.text + + +def test_extend_models_contains_self_referential_rootmodel_entry( + caplog: pytest.LogCaptureFixture, +) -> None: + from pydantic import RootModel + + class Loop(RootModel): # type: ignore[type-arg] + root: "Loop | Target" + + Loop.model_rebuild() + + wrapper = wrap_extension("ext", Ext) + assert wrapper is not None + models: ModelDict = { + _key("ext", "extension"): wrapper, + _key("loop"): Loop, + _key("target"): Target, + } + with caplog.at_level(logging.WARNING): + result = extend_models(models) + + by_name = {key.name: model for key, model in result.items()} + # The self-referential entry is kept unextended with a warning; the rest of the + # pass -- including the warning aggregation -- is unaffected. + assert by_name["loop"] is Loop + assert "Could not apply extensions to model 'loop'" in caplog.text + assert "ext" in _model(by_name["target"]).model_fields + assert "was not applied to any discovered model" not in caplog.text + + +def test_extension_provider_tags_wrapper() -> None: + wrapper = wrap_extension("ext", Ext) + assert wrapper is not None + tags = extension_provider([wrapper], _key("ext"), set()) + assert "extension" in tags + + +def test_extension_provider_ignores_plain_models() -> None: + tags = extension_provider([Target], _key("target"), set()) + assert "extension" not in tags + + +def test_generate_tags_passes_only_models_for_union_alias_entry() -> None: + received: list[list[type[BaseModel]]] = [] + + def provider( + models: Iterable[type[BaseModel]], key: ModelKey, tags: set[str] + ) -> list[str]: + received.append(list(models)) + return [] + + provider_key = TagProviderKey( + name="probe", entry_point="m:probe", package_name="overture-schema-system" + ) + _generate_tags(SegmentAlias, _key("segment"), {provider_key: provider}) + + (models,) = received + assert set(models) == {RoadSeg, RailSeg} + assert all(isinstance(m, type) and issubclass(m, BaseModel) for m in models) + + +def test_extend_models_extends_every_arm_of_union_alias_target() -> None: + @extends(SegmentAlias) + class SegExt(BaseModel): + note: str + + wrapper = wrap_extension("seg_ext", SegExt) + assert wrapper is not None + # The alias entry is not a class, so it needs a cast into ModelDict's value + # type -- discovery stores union aliases verbatim at runtime. + models = cast( + ModelDict, + { + _key("seg_ext", "extension"): wrapper, + _key("segment", "feature"): SegmentAlias, + }, + ) + result = extend_models(models) + + by_name = {key.name: model for key, model in result.items()} + extended_alias = by_name["segment"] + # Every arm of the alias gains the extension field. + arms = [ + tp + for tp in collect_types(extended_alias) + if isinstance(tp, type) and issubclass(tp, BaseModel) + ] + assert arms + assert all("seg_ext" in model.model_fields for model in arms) + # The discriminated-union structure survives the rebuild. + road = TypeAdapter(extended_alias).validate_python({"kind": "road"}) + assert isinstance(road, RoadSeg) + + +def test_tags_are_generated_from_the_extended_models( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Tags must describe the classes `discover_models` actually returns. + + A provider inspecting model fields sees extension-contributed fields, + which requires tag generation to run after extension application. + """ + wrapper = wrap_extension("ext", Ext) + assert wrapper is not None + + class _Entry: + def __init__(self, name: str, value: object) -> None: + self.name = name + self.value = f"m:{name}" + self._loaded = value + + def load(self) -> object: + return self._loaded + + def fake_entry_points(*, group: str) -> list[_Entry]: + if group == "overture.models": + return [_Entry("target", Target), _Entry("ext", Ext)] + return [] + + def field_provider( + types: Iterable[type[BaseModel]], key: ModelKey, tags: set[str] + ) -> set[str]: + return { + f"has:{field}" + for cls in types + for field in cls.model_fields + if field == "ext" + } + + monkeypatch.setattr( + discovery_module.importlib.metadata, "entry_points", fake_entry_points + ) + monkeypatch.setattr( + discovery_module, + "discover_tag_providers", + lambda: { + TagProviderKey( + name="fields", entry_point="t:fields", package_name="test" + ): field_provider + }, + ) + models = discovery_module.discover_models() + target_key = next(k for k in models if k.name == "target") + assert "has:ext" in target_key.tags + assert "ext" in models[target_key].model_fields diff --git a/packages/overture-schema-system/tests/test_discovery_filter_models.py b/packages/overture-schema-system/tests/test_discovery_filter_models.py index c1c051393..95ff0a007 100644 --- a/packages/overture-schema-system/tests/test_discovery_filter_models.py +++ b/packages/overture-schema-system/tests/test_discovery_filter_models.py @@ -1,4 +1,8 @@ -"""Direct coverage of filter_models combinator algebra.""" +"""Direct coverage of the selector combinator algebra. + +Exercised through `select_models` with `include_extension_entries=True`, +which is documented to reproduce the raw predicate stage exactly -- so +these tests also pin that equivalence.""" from pydantic import BaseModel @@ -6,7 +10,7 @@ ModelDict, ModelKey, TagSelector, - filter_models, + select_models, ) @@ -57,6 +61,23 @@ class Sources(BaseModel): tags=frozenset({"overture"}), ) + +def _select_all( + models: ModelDict, + selector: TagSelector = TagSelector(), + *, + type_names: tuple[str, ...] = (), +) -> ModelDict: + """Run select_models with the extension-entry hiding lifted. + + `include_extension_entries=True` reproduces the raw predicate stage + exactly, so the algebra is tested through the public entry point. + """ + return select_models( + models, selector, type_names=type_names, include_extension_entries=True + ) + + ALL_MODELS: ModelDict = { BUILDING_KEY: Building, SEGMENT_KEY: Segment, @@ -73,23 +94,23 @@ def names(models: ModelDict) -> set[str]: class TestEmptySelector: def test_empty_selector_returns_all(self) -> None: - result = filter_models(ALL_MODELS) + result = _select_all(ALL_MODELS) assert names(result) == names(ALL_MODELS) def test_empty_selector_explicit(self) -> None: - result = filter_models(ALL_MODELS, TagSelector()) + result = _select_all(ALL_MODELS, TagSelector()) assert names(result) == names(ALL_MODELS) class TestIncludeAny: def test_single_tag(self) -> None: - result = filter_models( + result = _select_all( ALL_MODELS, TagSelector(include_any=("overture:theme=buildings",)) ) assert names(result) == {"building"} def test_multi_tag_or(self) -> None: - result = filter_models( + result = _select_all( ALL_MODELS, TagSelector( include_any=( @@ -101,13 +122,13 @@ def test_multi_tag_or(self) -> None: assert names(result) == {"building", "segment", "connector"} def test_no_match(self) -> None: - result = filter_models( + result = _select_all( ALL_MODELS, TagSelector(include_any=("overture:theme=nonexistent",)) ) assert result == {} def test_mixed_match(self) -> None: - result = filter_models( + result = _select_all( ALL_MODELS, TagSelector( include_any=("overture:theme=buildings", "overture:theme=nonexistent") @@ -118,23 +139,21 @@ def test_mixed_match(self) -> None: class TestRequireAll: def test_single_tag(self) -> None: - result = filter_models(ALL_MODELS, TagSelector(require_all=("feature",))) + result = _select_all(ALL_MODELS, TagSelector(require_all=("feature",))) assert names(result) == {"building", "segment", "connector", "place"} def test_multi_tag_and_match(self) -> None: - result = filter_models( + result = _select_all( ALL_MODELS, TagSelector(require_all=("feature", "overture")) ) assert names(result) == {"building", "segment", "connector", "place"} def test_multi_tag_and_one_fails(self) -> None: - result = filter_models( - ALL_MODELS, TagSelector(require_all=("feature", "draft")) - ) + result = _select_all(ALL_MODELS, TagSelector(require_all=("feature", "draft"))) assert names(result) == {"place"} def test_no_match(self) -> None: - result = filter_models( + result = _select_all( ALL_MODELS, TagSelector(require_all=("feature", "nonexistent")) ) assert result == {} @@ -142,14 +161,14 @@ def test_no_match(self) -> None: class TestExcludeAny: def test_single_tag(self) -> None: - result = filter_models( + result = _select_all( ALL_MODELS, TagSelector(exclude_any=("overture:theme=buildings",)) ) assert "building" not in names(result) assert names(result) == {"segment", "connector", "place", "sources"} def test_multi_tag_or(self) -> None: - result = filter_models( + result = _select_all( ALL_MODELS, TagSelector( exclude_any=( @@ -161,21 +180,21 @@ def test_multi_tag_or(self) -> None: assert names(result) == {"place", "sources"} def test_no_match_keeps_all(self) -> None: - result = filter_models(ALL_MODELS, TagSelector(exclude_any=("nonexistent",))) + result = _select_all(ALL_MODELS, TagSelector(exclude_any=("nonexistent",))) assert names(result) == names(ALL_MODELS) class TestTypeNames: def test_single(self) -> None: - result = filter_models(ALL_MODELS, type_names=("building",)) + result = _select_all(ALL_MODELS, type_names=("building",)) assert names(result) == {"building"} def test_multiple(self) -> None: - result = filter_models(ALL_MODELS, type_names=("building", "place")) + result = _select_all(ALL_MODELS, type_names=("building", "place")) assert names(result) == {"building", "place"} def test_none_match(self) -> None: - result = filter_models(ALL_MODELS, type_names=("nonexistent",)) + result = _select_all(ALL_MODELS, type_names=("nonexistent",)) assert result == {} @@ -183,7 +202,7 @@ class TestCrossCombinator: def test_include_then_require(self) -> None: # Scope to features (places, transportation), narrow to those # also tagged "draft" → only place qualifies. - result = filter_models( + result = _select_all( ALL_MODELS, TagSelector( include_any=( @@ -197,7 +216,7 @@ def test_include_then_require(self) -> None: def test_include_then_exclude(self) -> None: # Scope to all themed features, exclude buildings. - result = filter_models( + result = _select_all( ALL_MODELS, TagSelector( include_any=( @@ -213,7 +232,7 @@ def test_include_then_exclude(self) -> None: def test_all_three_combinators_plus_type_names(self) -> None: # Scope to features in either places or transportation, # require feature tag, exclude drafts, restrict to segment by name. - result = filter_models( + result = _select_all( ALL_MODELS, TagSelector( include_any=( @@ -233,12 +252,12 @@ def test_double_application(self) -> None: selector = TagSelector( include_any=("overture:theme=buildings", "overture:theme=places") ) - once = filter_models(ALL_MODELS, selector) - twice = filter_models(once, selector) + once = _select_all(ALL_MODELS, selector) + twice = _select_all(once, selector) assert names(once) == names(twice) class TestInputInvariance: def test_returns_new_dict(self) -> None: - result = filter_models(ALL_MODELS) + result = _select_all(ALL_MODELS) assert result is not ALL_MODELS diff --git a/packages/overture-schema-system/tests/test_discovery_selection.py b/packages/overture-schema-system/tests/test_discovery_selection.py new file mode 100644 index 000000000..17f7ea664 --- /dev/null +++ b/packages/overture-schema-system/tests/test_discovery_selection.py @@ -0,0 +1,113 @@ +"""Tests for `select_models`: user predicates plus default extension hiding.""" + +from pydantic import BaseModel + +from overture.schema.system.discovery import ( + ModelKey, + TagSelector, + select_models, +) +from overture.schema.system.discovery.types import ModelDict +from overture.schema.system.extension import extends, wrap_extension + + +class Building(BaseModel): + name: str + + +@extends(Building) +class Capacity(BaseModel): + capacity: int | None = None + + +@extends(Building) +class Hours(BaseModel): + hours: str | None = None + + +def _wrapper(name: str, obj: object) -> type[BaseModel]: + wrapper = wrap_extension(name, obj) + assert wrapper is not None + return wrapper + + +def _key(name: str, *tags: str, entry_point: str | None = None) -> ModelKey: + return ModelKey( + name=name, entry_point=entry_point or f"m:{name}", tags=frozenset(tags) + ) + + +BUILDING = _key("building", "feature") +CAPACITY = _key("capacity", "extension") +# Deliberately untagged: hiding is structural, so an entry whose tag +# generation failed must stay hidden rather than leak into selections. +HOURS_UNTAGGED = _key("hours") + +MODELS: ModelDict = { + BUILDING: Building, + CAPACITY: _wrapper("capacity", Capacity), + HOURS_UNTAGGED: _wrapper("hours", Hours), +} + + +def test_extensions_dropped_by_default() -> None: + assert select_models(MODELS) == {BUILDING: Building} + + +def test_hiding_is_structural_not_tag_based() -> None: + # The untagged wrapper is still hidden: fail-safe against tag failures. + assert HOURS_UNTAGGED not in select_models(MODELS) + + +def test_include_extension_entries_shows_everything() -> None: + assert select_models(MODELS, include_extension_entries=True) == MODELS + + +def test_engaging_the_tag_lifts_the_default() -> None: + # The tag is the *selection* vocabulary: engaging it lifts the default, + # and the predicate stage then matches tagged entries (the untagged wrapper is + # unreachable this way -- by design, its tags are broken). + selected = select_models(MODELS, TagSelector(include_any=("extension",))) + assert set(selected) == {CAPACITY} + selected = select_models(MODELS, TagSelector(require_all=("extension",))) + assert set(selected) == {CAPACITY} + + +def test_user_exclusion_always_wins() -> None: + # A user's own exclude is a predicate, not a default -- naming or engaging + # cannot override it. + selector = TagSelector(exclude_any=("extension",)) + assert select_models(MODELS, selector, type_names=("capacity",)) == {} + + +def test_extension_only_name_is_exempt_from_the_default() -> None: + # "capacity" resolves exclusively to a hidden entry: naming it was + # necessarily about that entry. Works for the untagged wrapper too. + assert set(select_models(MODELS, type_names=("capacity",))) == {CAPACITY} + assert set(select_models(MODELS, type_names=("hours",))) == {HOURS_UNTAGGED} + + +def test_ambiguous_name_keeps_the_wrapper_hidden() -> None: + # A model and an extension sharing an entry-point name: naming it selects + # the visible entry only, so a collision cannot smuggle the permissive + # wrapper into a selection. + collided_model = _key("capacity", "feature", entry_point="other:Capacity") + models: ModelDict = {**MODELS, collided_model: Building} + selected = select_models(models, type_names=("capacity",)) + assert set(selected) == {collided_model} + + +def test_ambiguous_name_with_engaged_tag_shows_both() -> None: + collided_model = _key("capacity", "feature", entry_point="other:Capacity") + models: ModelDict = {**MODELS, collided_model: Building} + selected = select_models( + models, + TagSelector(include_any=("extension", "feature")), + type_names=("capacity",), + ) + assert set(selected) == {CAPACITY, collided_model} + + +def test_type_names_still_conjoin_with_predicates() -> None: + assert select_models(MODELS, type_names=("building",)) == {BUILDING: Building} + assert select_models(MODELS, type_names=("nonexistent",)) == {} diff --git a/packages/overture-schema-system/tests/test_extension.py b/packages/overture-schema-system/tests/test_extension.py new file mode 100644 index 000000000..652191f8e --- /dev/null +++ b/packages/overture-schema-system/tests/test_extension.py @@ -0,0 +1,737 @@ +"""Unit tests for the generic model extension mechanism. + +These tests deliberately use local, non-Feature models to prove the mechanism is generic over any +pydantic `BaseModel`, not tied to Overture's `Feature` base class. +""" + +import logging +from typing import Annotated, ForwardRef, Literal, NewType, Union, get_args + +import pytest +from pydantic import BaseModel, Field, RootModel, ValidationError, computed_field + +from overture.schema.system.extension import ( + _EXTENSION_ATTR, + Extends, + SelfReferentialRootError, + applied_extension_names, + create_extended_model, + extends, + extension_targets, + wrap_extension, +) + + +class Target(BaseModel): + name: str + + +class OtherTarget(BaseModel): + label: str + + +class Unrelated(BaseModel): + value: int + + +# --------------------------------------------------------------------------- +# Target validation +# --------------------------------------------------------------------------- + + +def test_extends_accepts_basemodel_targets() -> None: + assert Extends(Target).extends == (Target,) + assert Extends(Target, OtherTarget).extends == (Target, OtherTarget) + + +def test_extends_accepts_union_and_newtype_targets() -> None: + Aliased = NewType("Aliased", Target) + assert Extends(Target | OtherTarget).extends == (Target | OtherTarget,) + assert Extends(Aliased).extends == (Aliased,) + + +def test_extends_rejects_non_model_targets() -> None: + with pytest.raises(TypeError): + Extends(int) + with pytest.raises(TypeError): + Extends("not a model") + with pytest.raises(TypeError): + Extends() # at least one target required + + +def test_extends_rejects_partial_model_union() -> None: + with pytest.raises(TypeError): + Extends(Target | int) + + +def test_extends_rejects_rootmodel_with_non_model_root() -> None: + class IntRoot(RootModel[int]): + pass + + # A RootModel target is an alias over its root; a scalar root resolves to no + # models, so there is nothing to extend. + with pytest.raises(TypeError): + Extends(IntRoot) + with pytest.raises(TypeError): + Extends(Target | IntRoot) + + +def test_extends_accepts_rootmodel_over_models() -> None: + class Pair(RootModel[Target | OtherTarget]): + pass + + # Alias semantics: a RootModel whose root resolves to models is a valid target. + assert Extends(Pair).extends == (Pair,) + + +# --------------------------------------------------------------------------- +# extension_targets detection across declaration forms +# --------------------------------------------------------------------------- + + +def test_extension_targets_on_decorated_model() -> None: + @extends(Target) + class Ext(BaseModel): + extra: str + + assert extension_targets(Ext) == (Target,) + + +def test_extension_targets_on_newtype_with_metadata() -> None: + Scalar = NewType("Scalar", Annotated[int, Field(ge=0), Extends(Target)]) + assert extension_targets(Scalar) == (Target,) + + +def test_extension_targets_on_bare_annotated() -> None: + bare = Annotated[int, Extends(Target, OtherTarget)] + assert extension_targets(bare) == (Target, OtherTarget) + + +def test_extension_targets_returns_empty_for_non_extension() -> None: + assert extension_targets(Target) == () + assert extension_targets(int) == () + assert extension_targets(Annotated[int, Field(ge=0)]) == () + + +# --------------------------------------------------------------------------- +# wrap_extension +# --------------------------------------------------------------------------- + + +def test_wrap_extension_model_case() -> None: + @extends(Target) + class OpeningInfo(BaseModel): + note: str + + wrapper = wrap_extension("opening_info", OpeningInfo) + assert wrapper is not None + assert wrapper.__name__ == "OpeningInfoExtension" + assert set(wrapper.model_fields) == {"opening_info"} + # Default is None and the original extension type is preserved for re-use. + assert wrapper().model_dump()["opening_info"] is None + assert getattr(wrapper, _EXTENSION_ATTR) is OpeningInfo + # Targets remain introspectable through the wrapper. + assert extension_targets(wrapper) == (Target,) + # Validates an extension-only payload. + validated = wrapper.model_validate({"opening_info": {"note": "hi"}}) + assert validated.model_dump()["opening_info"]["note"] == "hi" + + +def test_wrap_extension_scalar_case_preserves_constraints() -> None: + Scalar = NewType("Scalar", Annotated[int, Field(ge=0, le=10), Extends(Target)]) + wrapper = wrap_extension("scalar_ext", Scalar) + assert wrapper is not None + assert wrapper.__name__ == "ScalarExtExtension" + assert wrapper.model_validate({"scalar_ext": 5}).model_dump()["scalar_ext"] == 5 + # Field(ge/le) constraints survive the round-trip through the wrapper. + with pytest.raises(ValidationError): + wrapper.model_validate({"scalar_ext": 99}) + + +def test_wrap_extension_returns_none_for_non_extension() -> None: + assert wrap_extension("plain", Target) is None + assert wrap_extension("plain", int) is None + + +@pytest.mark.parametrize( + "name", ["_secret", "model_stuff", "not-an-identifier", "class", "copy"] +) +def test_wrap_extension_rejects_unsafe_field_names(name: str) -> None: + @extends(Target) + class Ext(BaseModel): + note: str + + # Underscore names would silently become private attributes (a wrapper with no + # field), `model_` names hit Pydantic's protected namespace, `BaseModel` attribute + # names would shadow inherited behavior on every model, and the rest cannot be + # referenced by generated code. + with pytest.raises(ValueError, match="cannot be used as a field name"): + wrap_extension(name, Ext) + + +# --------------------------------------------------------------------------- +# create_extended_model +# --------------------------------------------------------------------------- + + +def _wrap(name: str, obj: object) -> type[BaseModel]: + wrapper = wrap_extension(name, obj) + assert wrapper is not None + return wrapper + + +def test_create_extended_model_adds_field_to_target() -> None: + @extends(Target) + class Ext(BaseModel): + note: str + + extensions = {"ext": _wrap("ext", Ext)} + extended = create_extended_model(Target, extensions) + + assert extended is not Target + assert issubclass(extended, Target) + assert "ext" in extended.model_fields + # Field is optional with default None. + assert extended(name="x").ext is None + instance = extended(name="x", ext=Ext(note="hello")) + assert instance.ext.note == "hello" + + +def test_create_extended_model_identity_for_non_matching_target() -> None: + @extends(Target) + class Ext(BaseModel): + note: str + + extensions = {"ext": _wrap("ext", Ext)} + # Unrelated is not a Target subclass — returned unchanged. + assert create_extended_model(Unrelated, extensions) is Unrelated + + +def test_create_extended_model_double_application_is_noop() -> None: + @extends(Target) + class Ext(BaseModel): + note: str + + extensions = {"ext": _wrap("ext", Ext)} + once = create_extended_model(Target, extensions) + twice = create_extended_model(once, extensions) + # Already applied — no new subclass is created. + assert twice is once + + +def test_create_extended_model_skips_field_name_collision() -> None: + class HasName(BaseModel): + name: str + + @extends(HasName) + class Ext(BaseModel): + note: str + + # Extension field name collides with an existing field on the target. + extensions = {"name": _wrap("name", Ext)} + assert create_extended_model(HasName, extensions) is HasName + + +def test_create_extended_model_no_collision_warning_for_non_target( + caplog: pytest.LogCaptureFixture, +) -> None: + class HasNote(BaseModel): + note: str + + @extends(Target) + class Ext(BaseModel): + x: int + + # HasNote shares a field name with the extension but is not one of its targets -- + # it must pass through silently, without a spurious collision warning. + extensions = {"note": _wrap("note", Ext)} + with caplog.at_level(logging.WARNING): + assert create_extended_model(HasNote, extensions) is HasNote + assert "collides" not in caplog.text + + +def test_create_extended_model_skips_alias_collision( + caplog: pytest.LogCaptureFixture, +) -> None: + class Aliased(BaseModel): + internal: str = Field(alias="ext") + + @extends(Aliased) + class Ext(BaseModel): + x: int + + # The extension field name matches an existing field's payload-level alias. + extensions = {"ext": _wrap("ext", Ext)} + with caplog.at_level(logging.WARNING): + assert create_extended_model(Aliased, extensions) is Aliased + assert "collides" in caplog.text + + +def test_create_extended_model_skips_computed_field_collision() -> None: + class WithComputed(BaseModel): + first: str + + @computed_field # type: ignore[prop-decorator] + @property + def ext(self) -> str: + return self.first + + @extends(WithComputed) + class Ext(BaseModel): + x: int + + extensions = {"ext": _wrap("ext", Ext)} + assert create_extended_model(WithComputed, extensions) is WithComputed + + +def test_create_extended_model_skips_computed_field_alias_collision() -> None: + # A computed field's serialization alias occupies its payload key just as a + # regular field alias does; applying an extension over it would produce + # duplicate keys in `model_dump_json(by_alias=True)`. + class WithAliasedComputed(BaseModel): + first: str + + @computed_field(alias="ext") # type: ignore[prop-decorator] + @property + def something(self) -> str: + return self.first + + @extends(WithAliasedComputed) + class Ext(BaseModel): + x: int + + extensions = {"ext": _wrap("ext", Ext)} + assert create_extended_model(WithAliasedComputed, extensions) is WithAliasedComputed + + +def test_create_extended_model_preserves_forward_ref_union_arms() -> None: + # `X | Y` raises on arms that don't implement `|` (an unresolved + # ForwardRef); the union rebuild must extend the model arms without + # crashing on the unresolved ones. + @extends(Target) + class Ext(BaseModel): + note: str + + extensions = {"ext": _wrap("ext", Ext)} + alias = Union[Target, "NeverDefined"] # type: ignore[name-defined] # noqa: F821 + + extended = create_extended_model(alias, extensions) + assert extended is not alias + arms = get_args(extended) + extended_target = next( + a for a in arms if isinstance(a, type) and issubclass(a, Target) + ) + assert "ext" in extended_target.model_fields + assert any(isinstance(a, ForwardRef) for a in arms) + + +def test_create_extended_model_skips_class_attribute_collision() -> None: + class WithMethod(BaseModel): + first: str + + def ext(self) -> str: + return self.first + + @extends(WithMethod) + class Ext(BaseModel): + x: int + + # `ext` is a method on the target; a field with that name would shadow it. + extensions = {"ext": _wrap("ext", Ext)} + assert create_extended_model(WithMethod, extensions) is WithMethod + + +def test_create_extended_model_leaves_scalar_rootmodel_unchanged() -> None: + class IntRoot(RootModel[int]): + pass + + @extends(Target) + class Ext(BaseModel): + x: int + + # A scalar root contains nothing to extend: the exact same class comes back, so + # code constructing it via `root=` is untouched. + extensions = {"ext": _wrap("ext", Ext)} + assert create_extended_model(IntRoot, extensions) is IntRoot + + +def test_create_extended_model_rewrites_rootmodel_union_root() -> None: + class Road(BaseModel): + kind: Literal["road"] = "road" + + class Rail(BaseModel): + kind: Literal["rail"] = "rail" + + class Segment(RootModel[Annotated[Road | Rail, Field(discriminator="kind")]]): + def which(self) -> str: + return self.root.kind + + @extends(Road) + class Ext(BaseModel): + note: str + + extensions = {"ext": _wrap("ext", Ext)} + rebuilt = create_extended_model(Segment, extensions) + + # Rebuilt as a subclass: identity changes, construction and methods survive. + assert rebuilt is not Segment + assert issubclass(rebuilt, Segment) + assert rebuilt.__name__ == Segment.__name__ + assert rebuilt.model_fields["root"].discriminator == "kind" + # The targeted arm accepts extension data through the root; the other arm and + # `root=` construction still work. + v = rebuilt.model_validate({"kind": "road", "ext": {"note": "hi"}}) + assert v.which() == "road" + assert v.root.ext.note == "hi" + assert rebuilt(root=Rail()).which() == "rail" + assert isinstance(rebuilt(root=Rail()), Segment) + + +def test_create_extended_model_rootmodel_identity_when_nothing_applies() -> None: + class Pair(RootModel[Target | OtherTarget]): + pass + + @extends(Unrelated) + class Ext(BaseModel): + x: int + + # No arm is targeted -- the exact same class comes back. + extensions = {"ext": _wrap("ext", Ext)} + assert create_extended_model(Pair, extensions) is Pair + + +def test_create_extended_model_rootmodel_double_application_is_noop() -> None: + class Pair(RootModel[Target | OtherTarget]): + pass + + @extends(Target) + class Ext(BaseModel): + x: int + + extensions = {"ext": _wrap("ext", Ext)} + once = create_extended_model(Pair, extensions) + assert once is not Pair + twice = create_extended_model(once, extensions) + assert twice is once + + +def test_create_extended_model_rootmodel_target_applies_to_arms() -> None: + class Pair(RootModel[Target | OtherTarget]): + pass + + # Alias semantics: targeting the RootModel targets its root's models, wherever + # they appear -- including as individually registered entries. + @extends(Pair) + class Ext(BaseModel): + x: int + + extensions = {"ext": _wrap("ext", Ext)} + extended_target = create_extended_model(Target, extensions) + assert extended_target is not Target + assert "ext" in extended_target.model_fields + assert create_extended_model(Unrelated, extensions) is Unrelated + + +def test_create_extended_model_self_referential_rootmodel_raises() -> None: + class Loop(RootModel): # type: ignore[type-arg] + root: "Loop | Target" + + Loop.model_rebuild() + + @extends(Target) + class Ext(BaseModel): + x: int + + extensions = {"ext": _wrap("ext", Ext)} + with pytest.raises(TypeError, match="self-referential"): + create_extended_model(Loop, extensions) + + +def test_wrap_extension_accepts_rootmodel_extension() -> None: + # An extension may itself be a RootModel: the wrapper field holds it and it + # validates as its bare root value. + @extends(Target) + class Hours(RootModel[list[str]]): + pass + + wrapper = _wrap("hours", Hours) + validated = wrapper.model_validate({"hours": ["09:00-17:00"]}) + assert validated.model_dump()["hours"] == ["09:00-17:00"] + extended = create_extended_model(Target, {"hours": wrapper}) + v = extended.model_validate({"name": "x", "hours": ["09:00-17:00"]}) + assert v.hours.root == ["09:00-17:00"] + + +def test_create_extended_model_recurses_into_union() -> None: + @extends(Target) + class Ext(BaseModel): + note: str + + extensions = {"ext": _wrap("ext", Ext)} + union = Target | Unrelated + extended = create_extended_model(union, extensions) + + arms = {arm.__name__: arm for arm in get_args(extended)} + # The Target arm gains the field; the Unrelated arm is untouched. + assert "ext" in arms["Target"].model_fields # type: ignore[union-attr] + assert arms["Unrelated"] is Unrelated + + +def test_create_extended_model_recurses_into_newtype() -> None: + @extends(Target) + class Ext(BaseModel): + note: str + + extensions = {"ext": _wrap("ext", Ext)} + Aliased = NewType("Aliased", Target) + extended = create_extended_model(Aliased, extensions) + assert hasattr(extended, "__supertype__") + assert "ext" in extended.__supertype__.model_fields + + +def test_create_extended_model_scalar_extension_field_type() -> None: + Scalar = NewType("Scalar", Annotated[int, Field(ge=0, le=10), Extends(Target)]) + extensions = {"scalar": _wrap("scalar", Scalar)} + extended = create_extended_model(Target, extensions) + + assert "scalar" in extended.model_fields + assert extended(name="x", scalar=5).scalar == 5 + # The uint-like constraint carried by the NewType is enforced on the target. + with pytest.raises(ValidationError): + extended(name="x", scalar=99) + + +def test_standalone_wrapper_join_pattern() -> None: + """The wrapper validates a standalone payload that joins onto an extended target.""" + + @extends(Target) + class Ext(BaseModel): + note: str + + wrapper = _wrap("ext", Ext) + extended = create_extended_model(Target, {"ext": wrapper}) + + base = Target(name="x") + ext_payload = wrapper.model_validate({"ext": {"note": "joined"}}) + joined = extended.model_validate( + base.model_dump(exclude_unset=True) | ext_payload.model_dump(exclude_unset=True) + ) + assert joined.name == "x" + assert joined.ext.note == "joined" + + +class TestDuplicateDeclarationsMerge: + def test_stacked_extends_decorators_merge_targets(self) -> None: + @extends(OtherTarget) + @extends(Target) + class Both(BaseModel): + note: str = "" + + # Bottom-up application: the inner (earlier) declaration comes first, + # and neither is silently dropped. + assert extension_targets(Both) == (Target, OtherTarget) + + def test_stacked_decorators_deduplicate(self) -> None: + @extends(Target, OtherTarget) + @extends(Target) + class Dup(BaseModel): + note: str = "" + + assert extension_targets(Dup) == (Target, OtherTarget) + + def test_subclass_declaration_shadows_inherited(self) -> None: + @extends(Target) + class Parent(BaseModel): + note: str = "" + + @extends(OtherTarget) + class Child(Parent): + pass + + # A subclass speaks for itself; it does not merge with what it + # inherits (which stays introspectable on the parent). + assert extension_targets(Child) == (OtherTarget,) + assert extension_targets(Parent) == (Target,) + + def test_multiple_extends_in_one_annotated_frame_merge(self) -> None: + expr = Annotated[int, Extends(Target), Extends(OtherTarget), Extends(Target)] + assert extension_targets(expr) == (Target, OtherTarget) + + def test_nested_annotated_flattens_and_merges(self) -> None: + # typing flattens nested Annotated into one frame, inner metadata + # first -- both declarations merge in declaration order. + inner = Annotated[int, Extends(Target)] + outer = Annotated[inner, Extends(OtherTarget)] + assert extension_targets(outer) == (Target, OtherTarget) + + def test_nearest_frame_wins_over_deeper_frames(self) -> None: + # A NewType boundary keeps frames distinct: the nearest frame that + # declares any Extends decides, deeper frames are not consulted. + inner_alias = NewType("inner_alias", Annotated[int, Extends(Target)]) + outer = Annotated[inner_alias, Extends(OtherTarget)] # type: ignore[valid-type] + assert extension_targets(outer) == (OtherTarget,) + + +class TestUnhashableTargets: + def test_merge_deduplicates_by_equality_not_hash(self) -> None: + # A valid target need not be hashable: Annotated[Model, []] resolves + # to a model but carries unhashable metadata. Merging must not + # require hashing. + listy = Annotated[Target, []] + + @extends(listy, OtherTarget) # type: ignore[arg-type] + @extends(listy) # type: ignore[arg-type] + class Both(BaseModel): + note: str = "" + + assert extension_targets(Both) == (listy, OtherTarget) + + def test_annotated_frame_merge_accepts_unhashable(self) -> None: + listy = Annotated[Target, []] + expr = Annotated[int, Extends(listy), Extends(OtherTarget), Extends(listy)] + assert extension_targets(expr) == (listy, OtherTarget) + + +class TestWrapperModuleProvenance: + def test_model_extension_keeps_its_defining_module(self) -> None: + @extends(Target) + class Modular(BaseModel): + note: str = "" + + wrapper = wrap_extension("modular", Modular) + assert wrapper is not None + assert wrapper.__module__ == Modular.__module__ + + def test_bare_annotated_uses_entry_point_module(self) -> None: + # `Annotated[...]` reports typing's module, which says nothing about + # where the extension was declared; the entry point's module does. + expr = Annotated[int, Extends(Target)] + wrapper = wrap_extension("cap", expr, module="my.extension.pkg") + assert wrapper is not None + assert wrapper.__module__ == "my.extension.pkg" + + def test_bare_annotated_without_module_falls_back(self) -> None: + expr = Annotated[int, Extends(Target)] + wrapper = wrap_extension("cap", expr) + assert wrapper is not None + assert wrapper.__module__ == "overture.schema.system.extension" + + +class TestContractHardening: + def test_raw_extends_model_in_mapping_applies_as_itself(self) -> None: + # A raw @extends model passed directly (instead of its wrap_extension + # wrapper) is coherent: the model itself becomes the optional field's + # type. Locked in so it degrades neither into an AttributeError nor a + # silent no-op. + @extends(Target) + class RawExt(BaseModel): + note: str + + extended = create_extended_model(Target, {"raw_ext": RawExt}) + assert isinstance(extended, type) and issubclass(extended, Target) + annotation = extended.model_fields["raw_ext"].annotation + assert annotation == (RawExt | None) + + def test_rebuilt_newtype_preserves_identity_metadata(self) -> None: + @extends(Target) + class Ext(BaseModel): + note: str = "" + + wrapper = wrap_extension("cap", Ext) + assert wrapper is not None + Alias = NewType("Alias", Target) + # A custom qualname distinguishes restoration from what a bare + # NewType("Alias", ...) rebuild would stamp anyway. + Alias.__qualname__ = "SomeNamespace.Alias" # type: ignore[attr-defined] + Alias.__doc__ = "An alias with prose." + extended_alias = create_extended_model(Alias, {"cap": wrapper}) + assert extended_alias is not Alias + assert extended_alias.__module__ == Alias.__module__ + assert extended_alias.__qualname__ == "SomeNamespace.Alias" + assert extended_alias.__doc__ == "An alias with prose." + + def test_self_referential_root_raises_dedicated_error(self) -> None: + class Loop(RootModel): # type: ignore[type-arg] + root: "Loop | Target" + + Loop.model_rebuild() + + @extends(Target) + class Ext(BaseModel): + x: int + + wrapper = wrap_extension("loop_ext", Ext) + assert wrapper is not None + with pytest.raises(SelfReferentialRootError): + create_extended_model(Loop, {"loop_ext": wrapper}) + + def test_applied_extension_names_tolerates_self_referential_root(self) -> None: + # The aggregation helper feeds warning reporting, which must not + # abort the extension pass on an entry that cannot have been extended. + class Loop(RootModel): # type: ignore[type-arg] + root: "Loop | Target" + + Loop.model_rebuild() + assert applied_extension_names(Loop) == frozenset() + + def test_applied_extension_names_propagates_unrelated_type_errors(self) -> None: + # Only the self-referential root case is tolerated; any other + # TypeError raised while inspecting the input must propagate. + class Hostile: + @property + def __supertype__(self) -> object: + raise TypeError("not a NewType") + + with pytest.raises(TypeError, match="not a NewType"): + applied_extension_names(Hostile()) + + +class TestRootModelDefaultFidelity: + """Extending a RootModel must not change its root's requiredness.""" + + def _wrapper(self) -> type[BaseModel]: + @extends(Target) + class Ext(BaseModel): + note: str = "" + + wrapper = wrap_extension("fidelity_ext", Ext) + assert wrapper is not None + return wrapper + + def test_defaulted_root_stays_constructible(self) -> None: + class DefaultedRoot(RootModel[Target]): + root: Target = Target(name="d") + + extended = create_extended_model( + DefaultedRoot, {"fidelity_ext": self._wrapper()} + ) + assert extended is not DefaultedRoot + assert extended.model_fields["root"].is_required() is False + instance = extended() + assert instance.root.name == "d" + # The default is carried verbatim (defaults are not re-validated), + # while explicit payloads validate against the extended root. + validated = extended.model_validate( + {"name": "x", "fidelity_ext": {"note": "n"}} + ) + assert validated.root.fidelity_ext.note == "n" + + def test_default_factory_root_preserved(self) -> None: + class FactoryRoot(RootModel[Target]): + root: Target = Field(default_factory=lambda: Target(name="f")) + + extended = create_extended_model(FactoryRoot, {"fidelity_ext": self._wrapper()}) + assert extended is not FactoryRoot + assert extended().root.name == "f" + + def test_required_root_stays_required(self) -> None: + class RequiredRoot(RootModel[Target]): + pass + + extended = create_extended_model( + RequiredRoot, {"fidelity_ext": self._wrapper()} + ) + assert extended is not RequiredRoot + assert extended.model_fields["root"].is_required() is True + with pytest.raises(ValidationError): + extended() diff --git a/packages/overture-schema/changelog.d/634.bugfix.md b/packages/overture-schema/changelog.d/634.bugfix.md new file mode 100644 index 000000000..384c348c7 --- /dev/null +++ b/packages/overture-schema/changelog.d/634.bugfix.md @@ -0,0 +1 @@ +Fixed unwrapping of `Annotated` annotations on the `type` field when deriving type literals (`get_args(Annotated)` was called instead of `get_args(type_type)`). diff --git a/packages/overture-schema/changelog.d/634.feature.md b/packages/overture-schema/changelog.d/634.feature.md new file mode 100644 index 000000000..655f301c1 --- /dev/null +++ b/packages/overture-schema/changelog.d/634.feature.md @@ -0,0 +1 @@ +Built the validation union with `select_models`, so extension data validates through the extended feature models while the permissive standalone extension wrappers stay out of the union. diff --git a/packages/overture-schema/src/overture/schema/__init__.py b/packages/overture-schema/src/overture/schema/__init__.py index 3728b389a..11386521f 100644 --- a/packages/overture-schema/src/overture/schema/__init__.py +++ b/packages/overture-schema/src/overture/schema/__init__.py @@ -9,7 +9,7 @@ from pydantic import BaseModel, Field, Tag, TypeAdapter from overture.schema.common import OvertureFeature -from overture.schema.system.discovery import discover_models +from overture.schema.system.discovery import discover_models, select_models from overture.schema.system.feature import Feature @@ -73,7 +73,9 @@ def _union_type_adapter() -> TypeAdapter: Return a Pydantic type adapter that can validate the union of all models discovered using entry points. """ - models = discover_models() + # Extension data validates through the feature models it was merged into; the + # permissive standalone wrappers are dropped by `select_models`' default hiding. + models = select_models(discover_models()) if not models: raise RuntimeError("no registered models found via entry points") @@ -159,7 +161,7 @@ def _typeliteral(feature_class: type[OvertureFeature]) -> object: """ type_type = feature_class.model_fields["type"].annotation while get_origin(type_type) is Annotated: - type_type = get_args(Annotated)[0] + type_type = get_args(type_type)[0] if get_origin(type_type) is not Literal: return None literal = get_args(type_type)[0]