"""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