Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@
import re
from collections.abc import Callable

from annotated_types import Ge, Gt, Interval, Le, Lt
from annotated_types import Ge, Gt, Interval, Le, Lt, MultipleOf

from overture.schema.system.primitive import GeometryTypeConstraint
from overture.schema.system.ref import Reference
Expand Down Expand Up @@ -104,6 +104,10 @@ def describe_field_constraint(
result = _first_bound(constraint)
if result is not None:
return result
if isinstance(constraint, MultipleOf):
if constraint.multiple_of == 1:
return "Must be a whole number"
return f"Must be a multiple of {constraint.multiple_of}"
if isinstance(constraint, (ArrayMinLen, ScalarMinLen)):
return f"Minimum length: {constraint.min_length}"
if isinstance(constraint, (ArrayMaxLen, ScalarMaxLen)):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -97,7 +97,7 @@ def extract_discriminator(


_TypeShape = tuple[object, ...]
_FieldKey = tuple[str, _TypeShape]
_FieldKey = tuple[str, _TypeShape, frozenset[object]]


def _structural_fingerprint(spec: FieldSpec) -> _TypeShape:
Expand Down Expand Up @@ -231,16 +231,20 @@ def extract_union(
for fs in member.spec.fields:
if fs.name in shared_field_names:
continue
key = (fs.name, _structural_fingerprint(fs))
# The key includes the constraints fingerprint alongside the
# structural one: two arms with the same name and shape but
# different constraints (e.g. VehicleAxleCountSelector's
# `ge=1, le=100, multiple_of=1` vs the other selectors' `ge=0`)
# must not collapse into one `AnnotatedField` sharing a single
# constraint set -- that would silently drop one arm's rules.
# Keeping them as separate rows, each gated to its own
# `variant_sources`, reuses the same per-arm `Guard` mechanism
# 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))
existing = seen.get(key)
if existing is not None:
existing_constraints = _constraints_fingerprint(existing.field_spec)
if _constraints_fingerprint(fs) != existing_constraints:
raise ValueError(
f"Union {name!r} field {fs.name!r} has the same structural "
f"shape across members but diverging constraints; dedup "
f"would silently drop one member's constraints"
)
prior_sources = existing.variant_sources or () if existing else ()
seen[key] = AnnotatedField(
field_spec=fs,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,7 @@

from overture.schema.system.field_path import ArrayPath, MapProjection

from .check_ir import Check, ModelCheck
from .check_ir import Check, Guard, ModelCheck
from .constraint_dispatch import ForbidIf, RequireIf, model_constraint_function

__all__ = [
Expand Down Expand Up @@ -277,13 +277,58 @@ def _symmetric_label_suffixes(keys: list[_K]) -> list[str]:
return [f"_{idx}" if total > 1 else "" for idx, total in _occurrence_indices(keys)]


def _field_label_suffixes(
keys: list[tuple[str, str, tuple[Guard, ...]]],
) -> list[str]:
"""Per-row field-label collision suffixes.

`keys` carries `(base_label, check_name, guards)` per emitted row. A
field label collides for two reasons, resolved differently:

- Across union arms -- the same field appears in several discriminator
arms, each carrying its own guard tuple. Every check gated to one
arm shares that arm's `_N` suffix (`N` its first-appearance order
among the label's arms), so a split field reports one consistent
label per arm. This includes a check unique to one arm (the axle
arm's `integer` check, absent from the dimension arms), which would
otherwise escape suffixing and report the bare label alongside its
`_N`-suffixed siblings.
- Within a single arm -- one field carries two same-named checks (a
lower- and upper-`bounds` pair emitted as separate checks). These
take a per-occurrence `_N` suffix keyed on `(label, name)`; a field
whose check names are all distinct stays bare.

A label reached by a single arm uses the occurrence rule (leaving
unsplit fields untouched); a label reached by several uses the arm
rule.
"""
arms_by_label: dict[str, list[tuple[Guard, ...]]] = {}
for label, _name, guards in keys:
arms = arms_by_label.setdefault(label, [])
if guards not in arms:
arms.append(guards)
occurrences = _occurrence_indices([(label, name) for label, name, _ in keys])
suffixes: list[str] = []
for (label, _name, guards), (occ_idx, occ_total) in zip(
keys, occurrences, strict=True
):
arms = arms_by_label[label]
if len(arms) > 1:
suffixes.append(f"_{arms.index(guards)}")
elif occ_total > 1:
suffixes.append(f"_{occ_idx}")
else:
suffixes.append("")
return suffixes


@dataclass(frozen=True, slots=True)
class FieldCheckRow:
"""One emitted field-check row, with its final derived strings.

The renderer emits one row per descriptor of each `Check`.
`field_check_rows` flattens the check list into these rows once,
computing both the symmetric `label` collision suffix and the
computing both the arm-grouped `label` collision suffix and the
asymmetric `func_name` disambiguation, so the renderer and test
renderer agree without each re-deriving them by a positional index.

Expand Down Expand Up @@ -341,15 +386,25 @@ def field_check_rows(field_checks: list[Check]) -> list[FieldCheckRow]:
raw_func_names.append(f"_{sanitize_field_name(label)}{func_suffix}_check")
flattened.append((check, desc_idx, label, name))
func_names = disambiguate(raw_func_names)
label_suffixes = _symmetric_label_suffixes(
[(label, name) for _check, _idx, label, name in flattened]
label_suffixes = _field_label_suffixes(
[(label, name, check.guards) for check, _idx, label, name in flattened]
)
return [
rows = [
FieldCheckRow(check, desc_idx, f"{label}{label_suffix}", name, func_name)
for (check, desc_idx, label, name), label_suffix, func_name in zip(
flattened, label_suffixes, func_names, strict=True
)
]
# Arm-grouped suffixing cannot distinguish two same-name checks that
# land in one arm of a split field; fail generation loudly if a schema
# ever produces that instead of emitting indistinguishable violations.
identities = [(row.label, row.name) for row in rows]
duplicates = {i for i in identities if identities.count(i) > 1}
if duplicates:
raise ValueError(
f"Duplicate violation identities in generated checks: {sorted(duplicates)}"
)
return rows


@dataclass(frozen=True, slots=True)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -89,6 +89,59 @@
]


_BOUND_ORDER = ("ge", "gt", "le", "lt")


def _coalesce_bounds(
descriptors: list[ExpressionDescriptor],
) -> list[ExpressionDescriptor]:
"""Merge multiple `check_bounds` descriptors into a single one.

A field with both a lower and an upper bound (`Field(ge=1, le=100)`)
yields separate `Ge` and `Le` constraints, each dispatched to its own
`check_bounds`. They target the same column and describe one range, so
they collapse into one `check_bounds(ge=1, le=100)`: one violation
identity instead of two same-name checks -- which a split union arm
cannot otherwise label distinctly (see `_render_common.field_check_rows`).

Merges only bounds of distinct kinds. Two constraints of the *same* kind
with different values (`ge=1` from a NewType, `ge=5` at the field) have no
unambiguous merge, so this raises rather than silently keeping one -- the
author should state the single intended bound.

Raises
------
ValueError
When two bounds of the same kind carry different values.
"""
bound_descs = [d for d in descriptors if d.function == "check_bounds"]
if len(bound_descs) <= 1:
return descriptors
merged: dict[str, object] = {}
for d in bound_descs:
for key, value in d.kwargs:
if key in merged and merged[key] != value:
raise ValueError(
f"conflicting {key} bounds on one field: "
f"{merged[key]!r} vs {value!r}; declare a single {key}"
)
merged[key] = value
merged_desc = ExpressionDescriptor(
function="check_bounds",
kwargs=tuple((k, merged[k]) for k in _BOUND_ORDER if k in merged),
check_nan=bound_descs[0].check_nan,
)
result: list[ExpressionDescriptor] = []
placed = False
for d in descriptors:
if d.function != "check_bounds":
result.append(d)
elif not placed:
result.append(merged_desc)
placed = True
return result


def _dispatch_layer_constraints(
constraints: tuple[ConstraintSource, ...],
base_type: str | None,
Expand All @@ -101,7 +154,7 @@ def _dispatch_layer_constraints(
desc = dispatch_constraint(cs.constraint, base_type=base_type)
if desc is not None:
descriptors.append(desc)
return descriptors
return _coalesce_bounds(descriptors)


def _literal_alternatives(shape: Scalar | MapOf) -> tuple[object, ...]:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@
from dataclasses import dataclass
from typing import Any, NamedTuple, TypeAlias

from annotated_types import Ge, Gt, Interval, Le, Lt
from annotated_types import Ge, Gt, Interval, Le, Lt, MultipleOf
from pydantic import Strict
from pydantic._internal._fields import PydanticMetadata

Expand Down Expand Up @@ -235,6 +235,21 @@ def _dispatch_bounds(
)


def _dispatch_multiple_of(
constraint: MultipleOf,
_base_type: str | None,
) -> ExpressionDescriptor:
"""Map `Field(multiple_of=n)` to a check_multiple_of descriptor.

`check_multiple_of(col, n)` tests `col % n == 0`; `multiple_of=1` is the
integral (whole-number) case. The divisor rides in `args`, so any positive
`n` dispatches without special-casing.
"""
return ExpressionDescriptor(
function="check_multiple_of", args=(constraint.multiple_of,)
)


def _dispatch_pattern(
constraint: PatternConstraint,
_base_type: str | None,
Expand Down Expand Up @@ -346,6 +361,7 @@ def _raw_pattern(constraint: object) -> str | None:
function="check_geometry_type", args=tuple(c.allowed_types)
),
),
(MultipleOf, _dispatch_multiple_of),
]


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -16,15 +16,14 @@
Scalar,
UnionRef,
)
from ..extraction.field_walk import enum_source, terminal_scalar
from ..extraction.field_walk import enum_source
from ..extraction.specs import FieldSpec, ModelSpec, UnionSpec
from ..extraction.type_registry import get_type_mapping

__all__ = [
"SHARED_TYPE_REFS",
"SchemaField",
"build_schema",
"spark_type_rank",
]

# Types whose base_type name maps to a _schema_structs.py StructType constant.
Expand Down Expand Up @@ -87,58 +86,45 @@ def _spark_for_scalar(scalar: Scalar) -> str:
return _spark_for_base(scalar.base_type, scalar.source_type)


# Spark numeric type widening precedence (higher rank = wider type).
_SPARK_TYPE_WIDENING: dict[str, int] = {
"IntegerType()": 0,
"LongType()": 1,
"DoubleType()": 2,
}


def spark_type_rank(field_spec: FieldSpec) -> int:
"""Return a widening rank for the field's resolved Spark type.

Fields with a higher rank are preferred when deduplicating union
members by name. Non-numeric types return -1 (no widening).
"""
scalar = terminal_scalar(field_spec.shape)
if not isinstance(scalar, Primitive):
return -1
expr = _spark_for_base(scalar.base_type, scalar.source_type)
return _SPARK_TYPE_WIDENING.get(expr, -1)


def _deduplicate_by_name(fields: list[FieldSpec]) -> list[FieldSpec]:
"""Keep one FieldSpec per name, widening the Spark type on conflict.

Union annotated_fields may contain the same field name with different
type shapes (e.g. `value` as uint8 in one variant and float64 in
another). Parquet stores one column per name, so the schema needs
exactly one entry. When two fields share a name, the one with the
wider numeric Spark type wins (matching Parquet's type-widening
behavior). Two same-named fields whose Spark types are non-numeric and
not identical cannot share a column, so the collision fails loudly
rather than silently keeping whichever arm came first.
"""Keep one FieldSpec per name, requiring every arm to agree on Spark type.

Union annotated_fields may contain the same field name declared by
multiple arms with different `FieldSpec`s -- a `Literal` discriminator
whose value differs per arm, or a field whose per-arm constraints
diverge (see `union_extraction.extract_union`). A columnar sink stores
one type per column name, so the schema needs exactly one entry. Two
same-named fields are compatible when they resolve to the SAME Spark
type -- the first-seen `FieldSpec`'s shape is kept (arbitrarily; the
column type is identical either way). Two same-named fields that resolve
to DIFFERENT Spark types cannot share one generated column, so this
always raises, whether the mismatch is numeric (a narrower int type vs a
float) or not.

Widening the two to their common type would often work in practice --
Spark and Parquet can promote a narrower numeric column to a wider one
(reading an int where the schema declares a double, say). It is forbidden
anyway: a widened column makes the union's type an implicit property
inferred from whichever arms happen to disagree, rather than a decision
stated in the model. Raising forces that decision to the surface at
generation instead of leaving it as a silent compatibility trap.
"""
seen: dict[str, FieldSpec] = {}
for f in fields:
existing = seen.get(f.name)
if existing is None:
seen[f.name] = f
continue
rank_f, rank_existing = spark_type_rank(f), spark_type_rank(existing)
if rank_f < 0 and rank_existing < 0:
spark_f = _shape_to_spark(f.shape)
spark_existing = _shape_to_spark(existing.shape)
if spark_f != spark_existing:
raise ValueError(
f"Union field {f.name!r} resolves to incompatible "
f"non-widening Spark types across arms "
f"({spark_existing} vs {spark_f}); a single Parquet "
"column cannot represent both."
)
if rank_f > rank_existing:
seen[f.name] = f
spark_f, spark_existing = (
_shape_to_spark(f.shape),
_shape_to_spark(existing.shape),
)
if spark_f != spark_existing:
raise ValueError(
f"Union field {f.name!r} resolves to incompatible Spark "
f"types across arms ({spark_existing} vs {spark_f}); a "
"single Parquet column cannot represent both."
)
return list(seen.values())


Expand Down
Loading