Files
mizan/protocol/mizan-codegen/python/pydantic_to_rust/walker.py
2026-07-30 21:12:29 -04:00

286 lines
9.9 KiB
Python

"""Pydantic class -> `ModelDecl`.
Walks `model_fields`, lowers each annotation to a `TypeShape`, and captures
defaults. Annotations this does not recognize lower to `Primitive.STRING`,
matching what Mizan's own IR does with an unknown shape.
"""
from __future__ import annotations
import enum
import re
import sys
import types
from typing import Any, Literal, Union, get_args, get_origin
from pydantic import BaseModel
from pydantic_core import PydanticUndefined
from pydantic_to_rust.idents import to_rust_variant_ident
from pydantic_to_rust.ir import (
DefaultValue,
FieldDecl,
ModelDecl,
Primitive,
TypeShape,
)
# NewTypes named i8..i64 / u8..u64 / f16..f64, recognized only when the
# caller opts in.
_WIDTH_NEWTYPE_RE = re.compile(r"^([iuf])(8|16|32|64)$")
def walk_pydantic_model(
model: type[BaseModel],
*,
newtype_width_detection: bool = False,
) -> ModelDecl:
"""Lower one Pydantic class.
With `newtype_width_detection`, width-named NewTypes carry their width
and signedness. Off, every NewType lowers to a `ref`.
"""
fields: list[FieldDecl] = []
for field_name, field_info in model.model_fields.items():
required = field_info.is_required()
default: DefaultValue | None = None
if not required:
# A factory wins over a literal — it is the canonical Pydantic
# path for compound defaults.
factory = getattr(field_info, "default_factory", None)
if factory is not None:
default = _capture_factory_default(_call_factory(factory, field_name))
else:
default = _capture_default(field_info.default)
fields.append(
FieldDecl(
name=field_name,
shape=lower_annotation(
field_info.annotation,
newtype_width_detection=newtype_width_detection,
),
required=required,
default=default,
)
)
return ModelDecl(
name=model.__name__,
fields=tuple(fields),
docstring=(model.__doc__ or None),
)
def _call_factory(factory, field_name: str):
"""A factory that raises leaves the field with no captured default."""
try:
return factory()
except Exception as exc:
print(
f"pydantic_to_rust: default_factory for {field_name!r} raised "
f"{type(exc).__name__}: {exc}; emitting no default",
file=sys.stderr,
)
return None
def lower_annotation(
annotation: Any,
*,
newtype_width_detection: bool = False,
) -> TypeShape:
"""Lower a bare annotation, for callers that do not have a whole model."""
return _lower_annotation(
annotation,
newtype_width_detection=newtype_width_detection,
)
def _capture_factory_default(value: Any) -> DefaultValue | None:
"""Capture a factory's result. Empty collections map onto Rust's own
`Default`; a model instance is serialized so the emitter can name it."""
if value is None:
return DefaultValue(kind="null", literal=None)
if isinstance(value, (list, tuple, set, frozenset)) and len(value) == 0:
return DefaultValue(kind="empty_seq", literal=None)
if isinstance(value, dict) and len(value) == 0:
return DefaultValue(kind="empty_map", literal=None)
if isinstance(value, BaseModel):
return DefaultValue(
kind="compound_model",
literal={"type": type(value).__name__, "fields": value.model_dump()},
)
if isinstance(value, bool):
return DefaultValue(kind="boolean", literal=value)
if isinstance(value, int):
return DefaultValue(kind="integer", literal=value)
if isinstance(value, float):
return DefaultValue(kind="number", literal=value)
if isinstance(value, str):
return DefaultValue(kind="string", literal=value)
return DefaultValue(kind="compound", literal=value)
def _capture_default(raw: Any) -> DefaultValue | None:
"""Filter Pydantic's sentinels, then record the literal's kind."""
if raw is None:
return DefaultValue(kind="null", literal=None)
if raw is PydanticUndefined or raw is ...:
return None
if isinstance(raw, enum.Enum):
# Store the Rust variant name, not the Python member name, so the
# default body and the enum declaration agree.
return DefaultValue(
kind="enum_variant",
literal={
"type": type(raw).__name__,
"variant": to_rust_variant_ident(raw.name),
},
)
if isinstance(raw, bool):
return DefaultValue(kind="boolean", literal=raw)
if isinstance(raw, int):
return DefaultValue(kind="integer", literal=raw)
if isinstance(raw, float):
return DefaultValue(kind="number", literal=raw)
if isinstance(raw, str):
return DefaultValue(kind="string", literal=raw)
if isinstance(raw, tuple):
return DefaultValue(kind="tuple", literal=tuple(raw))
return DefaultValue(kind="compound", literal=raw)
def _lower_annotation(
annotation: Any,
*,
newtype_width_detection: bool = False,
) -> TypeShape:
# NewType carries `__supertype__` and `__name__`, and keeps its named
# identity on the Rust side.
if hasattr(annotation, "__supertype__") and hasattr(annotation, "__name__"):
if newtype_width_detection:
match = _WIDTH_NEWTYPE_RE.match(annotation.__name__)
if match:
kind, bits = match.group(1), int(match.group(2))
if kind == "f":
return TypeShape(primitive=Primitive.NUMBER, float_width=bits)
return TypeShape(
primitive=Primitive.INTEGER,
int_width=bits,
int_signed=(kind == "i"),
)
return TypeShape(ref=annotation.__name__)
inner, is_optional = _extract_optional(annotation)
if is_optional:
return TypeShape(
optional_inner=_lower_annotation(
inner, newtype_width_detection=newtype_width_detection
)
)
origin = get_origin(annotation)
if origin is Union or isinstance(annotation, types.UnionType):
branches = [a for a in get_args(annotation) if a is not type(None)]
if len(branches) > 1:
return TypeShape(
union_branches=tuple(
_lower_annotation(
b, newtype_width_detection=newtype_width_detection
)
for b in branches
)
)
# `tuple[T1..TN]` with no ellipsis becomes `[T; N]`, but only when every
# element shares a shape; mixed tuples fall through to the list path.
if origin is tuple:
args = get_args(annotation)
if args and args[-1] is not Ellipsis:
first = _lower_annotation(
args[0], newtype_width_detection=newtype_width_detection
)
uniform = all(
_lower_annotation(a, newtype_width_detection=newtype_width_detection)
== first
for a in args[1:]
)
if uniform:
return TypeShape(array_element=first, array_length=len(args))
element = _extract_list_element(annotation)
if element is not None:
return TypeShape(
list_inner=_lower_annotation(
element, newtype_width_detection=newtype_width_detection
)
)
if origin is dict:
args = get_args(annotation)
if len(args) == 2:
key = _lower_annotation(
args[0], newtype_width_detection=newtype_width_detection
)
value = _lower_annotation(
args[1], newtype_width_detection=newtype_width_detection
)
return TypeShape(
map_value=value,
map_key=key if key.ref is not None else None,
)
if origin is Literal:
args = get_args(annotation)
if all(isinstance(a, str) for a in args):
return TypeShape(enum_variants=tuple(args))
if len(args) == 1:
sole = args[0]
if isinstance(sole, bool):
return TypeShape(primitive=Primitive.BOOLEAN)
if isinstance(sole, int):
return TypeShape(primitive=Primitive.INTEGER)
if isinstance(sole, float):
return TypeShape(primitive=Primitive.NUMBER)
# An Enum or model class is referenced by name; the caller emits or
# hand-writes the corresponding Rust declaration.
if isinstance(annotation, type) and issubclass(annotation, enum.Enum):
return TypeShape(ref=annotation.__name__)
if isinstance(annotation, type) and issubclass(annotation, BaseModel):
return TypeShape(ref=annotation.__name__)
if annotation is int:
return TypeShape(primitive=Primitive.INTEGER)
if annotation is float:
return TypeShape(primitive=Primitive.NUMBER)
if annotation is bool:
return TypeShape(primitive=Primitive.BOOLEAN)
if annotation is str:
return TypeShape(primitive=Primitive.STRING)
return TypeShape(primitive=Primitive.STRING)
def _extract_optional(ann: Any) -> tuple[Any, bool]:
"""`(T, True)` for `Optional[T]` or `T | None`, else `(ann, False)`."""
origin = get_origin(ann)
if origin is Union or isinstance(ann, types.UnionType):
args = get_args(ann)
non_none = [a for a in args if a is not type(None)]
if len(non_none) == 1 and type(None) in args:
return non_none[0], True
return ann, False
def _extract_list_element(ann: Any) -> Any | None:
"""`T` for `list[T]`, `tuple[T, ...]`, `set[T]`, `frozenset[T]`."""
origin = get_origin(ann)
if origin in (list, tuple, set, frozenset):
args = get_args(ann)
if origin is tuple and len(args) >= 2 and args[1] is Ellipsis:
return args[0]
if args:
return args[0]
return None