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
4 changes: 2 additions & 2 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,7 @@ maintainers = [{ name = "Litestar Developers", email = "hello@litestar.dev" }]
name = "sqlspec"
readme = "README.md"
requires-python = ">=3.10, <4.0"
version = "0.58.2"
version = "0.58.3"

[project.urls]
Discord = "https://discord.gg/litestar"
Expand Down Expand Up @@ -331,7 +331,7 @@ opt_level = "3" # Maximum optimization (0-3)
allow_dirty = true
commit = false
commit_args = "--no-verify"
current_version = "0.58.2"
current_version = "0.58.3"
ignore_missing_files = false
ignore_missing_version = false
message = "chore(release): bump to v{new_version}"
Expand Down
129 changes: 103 additions & 26 deletions sqlspec/utils/schema.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,13 @@
"""Schema transformation utilities for converting data to various schema types."""

import datetime
from collections.abc import Callable, Sequence
from collections.abc import Callable, Mapping, Sequence
from decimal import Decimal, InvalidOperation
from enum import Enum
from functools import partial
from pathlib import Path, PurePath
from typing import Any, Final, TypeGuard, cast, overload
from types import UnionType
from typing import Annotated, Any, Final, TypeGuard, Union, cast, get_args, get_origin, overload
from uuid import UUID

from typing_extensions import TypeVar
Expand All @@ -15,12 +16,9 @@
from sqlspec.exceptions import SQLSpecError
from sqlspec.typing import CATTRS_INSTALLED, NUMPY_INSTALLED, MsgspecValidationError, SchemaT, convert, get_type_adapter
from sqlspec.utils.dispatch import TypeDispatcher
from sqlspec.utils.logging import get_logger
from sqlspec.utils.module_loader import import_optional_attr
from sqlspec.utils.serializers import from_json
from sqlspec.utils.text import camelize, kebabize, pascalize
from sqlspec.utils.type_guards import (
get_msgspec_rename_config,
is_attrs_instance,
is_attrs_schema,
is_dataclass,
Expand All @@ -46,16 +44,12 @@
DataT = TypeVar("DataT", default=dict[str, Any])
ValueT = TypeVar("ValueT")

logger = get_logger(__name__)

_DATETIME_TYPES: Final[set[type]] = {datetime.datetime, datetime.date, datetime.time}
_DATETIME_TYPE_TUPLE: Final[tuple[type, ...]] = (datetime.datetime, datetime.date, datetime.time)
_MSGSPEC_RENAME_CONVERTERS: Final[dict[str, Callable[[str], str]]] = {
"camel": camelize,
"kebab": kebabize,
"pascal": pascalize,
}
_MAPPING_TYPE_ARGUMENT_COUNT: Final = 2
_MSGSPEC_FIELD_CACHE: "dict[type, tuple[tuple[tuple[str, str, Any], ...], frozenset[str], frozenset[str]]]" = {}
_NUMPY_RECURSIVE_DISPATCHER: "TypeDispatcher[Callable[[Any], Any]] | None" = None
_NULLABLE_UNION_ARGUMENT_COUNT: Final = 2


# =============================================================================
Expand Down Expand Up @@ -356,20 +350,11 @@ def _get_numpy_recursive_dispatcher() -> "TypeDispatcher[Callable[[Any], Any]]":

def _convert_msgspec(data: Any, schema_type: Any) -> Any:
"""Convert data to msgspec Struct."""
rename_config = get_msgspec_rename_config(schema_type)

transformed_data = data
if (rename_config and is_dict(data)) or (isinstance(data, Sequence) and data and is_dict(data[0])):
try:
converter = _MSGSPEC_RENAME_CONVERTERS.get(rename_config) if rename_config else None
if converter:
transformed_data = (
[transform_dict_keys(item, converter) if is_dict(item) else item for item in data]
if isinstance(data, Sequence)
else (transform_dict_keys(data, converter) if is_dict(data) else data)
)
except Exception as e:
logger.debug("Field name transformation failed for msgspec schema: %s", e)
transformed_data = (
[_normalize_msgspec_input(item, schema_type) for item in data]
if isinstance(data, Sequence) and not isinstance(data, (str, bytes, bytearray))
else _normalize_msgspec_input(data, schema_type)
)

target_type = list[schema_type] if isinstance(transformed_data, Sequence) else schema_type

Expand All @@ -386,6 +371,98 @@ def _convert_msgspec(data: Any, schema_type: Any) -> Any:
)


def _normalize_msgspec_input(data: Any, target_type: Any) -> Any:
"""Normalize Struct field aliases according to the declared target type."""
target_type = _unwrap_msgspec_target(target_type)
if target_type is None:
return data

if is_msgspec_struct(target_type):
return _normalize_msgspec_struct(data, cast("type", target_type))

origin = get_origin(target_type)
args = get_args(target_type)
if origin is None or not args:
return data

if _is_mapping_origin(origin):
if not isinstance(data, Mapping) or len(args) < _MAPPING_TYPE_ARGUMENT_COUNT:
return data
value_type = args[1]
return {key: _normalize_msgspec_input(value, value_type) for key, value in data.items()}

if _is_sequence_origin(origin):
if not isinstance(data, Sequence) or isinstance(data, (str, bytes, bytearray)):
return data
if origin is tuple and len(args) > 1 and args[1] is not Ellipsis:
return [
_normalize_msgspec_input(value, args[index]) if index < len(args) else value
for index, value in enumerate(data)
]
item_type = next(iter(args), Any)
return [_normalize_msgspec_input(value, item_type) for value in data]

return data


def _normalize_msgspec_struct(data: Any, schema_type: type) -> Any:
if not isinstance(data, Mapping):
return data

fields, python_names, encoded_names = _msgspec_field_plan(schema_type)
normalized = {key: value for key, value in data.items() if key not in python_names and key not in encoded_names}

for name, encode_name, field_type in fields:
if encode_name in data:
normalized[encode_name] = _normalize_msgspec_input(data[encode_name], field_type)
elif name in data:
normalized[encode_name] = _normalize_msgspec_input(data[name], field_type)

return normalized


def _msgspec_field_plan(schema_type: type) -> "tuple[tuple[tuple[str, str, Any], ...], frozenset[str], frozenset[str]]":
try:
return _MSGSPEC_FIELD_CACHE[schema_type]
except KeyError:
from msgspec import structs

fields = tuple(
(field.name, field.encode_name, field.type) for field in structs.fields(cast("Any", schema_type))
)
plan = fields, frozenset(field[0] for field in fields), frozenset(field[1] for field in fields)
_MSGSPEC_FIELD_CACHE[schema_type] = plan
return plan


def _unwrap_msgspec_target(target_type: Any) -> Any:
while get_origin(target_type) is Annotated:
target_type = get_args(target_type)[0]

origin = get_origin(target_type)
if origin is Union or origin is UnionType:
args = get_args(target_type)
non_null = tuple(arg for arg in args if arg is not type(None))
if len(args) == _NULLABLE_UNION_ARGUMENT_COUNT and len(non_null) == 1:
return _unwrap_msgspec_target(non_null[0])
return None
return target_type


def _is_mapping_origin(origin: Any) -> bool:
try:
return issubclass(origin, Mapping)
except TypeError:
return False


def _is_sequence_origin(origin: Any) -> bool:
try:
return issubclass(origin, Sequence)
except TypeError:
return False


def _convert_pydantic(data: Any, schema_type: Any) -> Any:
"""Convert data to Pydantic model."""
if isinstance(data, Sequence):
Expand Down
153 changes: 153 additions & 0 deletions tests/unit/utils/test_schema_msgspec.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,153 @@
"""Tests for msgspec schema conversion."""

from typing import Annotated

import msgspec
import pytest

import sqlspec.utils.schema as schema_utils
from sqlspec.utils.schema import to_schema


def _upper_field(name: str) -> str:
return name.upper()


class CamelChild(msgspec.Struct, rename="camel"):
item_name: str
item_count: int


class SingleWordEnvelope(msgspec.Struct, rename="camel"):
rows: list[CamelChild]
total: int


class KebabEnvelope(msgspec.Struct, rename="kebab"):
child_rows: list[CamelChild]


class ExplicitAliasChild(msgspec.Struct):
display_name: str = msgspec.field(name="label")


class CallableAliasEnvelope(msgspec.Struct, rename=_upper_field):
child: ExplicitAliasChild
item_count: int


class OptionalCollectionEnvelope(msgspec.Struct):
child: Annotated[CamelChild | None, "nested child"]
children: tuple[CamelChild, ...]
children_by_key: dict[str, CamelChild]


class ArbitraryPayloadEnvelope(msgspec.Struct, rename="camel"):
metadata: dict[str, object]


class StrictEnvelope(msgspec.Struct, rename="camel", forbid_unknown_fields=True):
child: CamelChild


def test_msgspec_single_word_envelope_decodes_nested_python_field_names() -> None:
result = to_schema({"rows": [{"item_name": "a", "item_count": 1}], "total": 1}, schema_type=SingleWordEnvelope)

assert result == SingleWordEnvelope(rows=[CamelChild(item_name="a", item_count=1)], total=1)


def test_msgspec_single_word_envelope_decodes_batch() -> None:
result = to_schema(
[
{"rows": [{"item_name": "a", "item_count": 1}], "total": 1},
{"rows": [{"item_name": "b", "item_count": 2}], "total": 1},
],
schema_type=SingleWordEnvelope,
)

assert result == [
SingleWordEnvelope(rows=[CamelChild(item_name="a", item_count=1)], total=1),
SingleWordEnvelope(rows=[CamelChild(item_name="b", item_count=2)], total=1),
]


def test_msgspec_mixed_parent_and_child_rename_conventions() -> None:
result = to_schema({"child_rows": [{"item_name": "a", "item_count": 1}]}, schema_type=KebabEnvelope)

assert result == KebabEnvelope(child_rows=[CamelChild(item_name="a", item_count=1)])


def test_msgspec_callable_and_explicit_aliases() -> None:
result = to_schema({"child": {"display_name": "Ada"}, "item_count": 1}, schema_type=CallableAliasEnvelope)

assert result == CallableAliasEnvelope(child=ExplicitAliasChild(display_name="Ada"), item_count=1)


def test_msgspec_optional_and_collection_fields() -> None:
child = {"item_name": "a", "item_count": 1}
result = to_schema(
{"child": child, "children": [child], "children_by_key": {"arbitrary_key": child}},
schema_type=OptionalCollectionEnvelope,
)

expected = CamelChild(item_name="a", item_count=1)
assert result == OptionalCollectionEnvelope(
child=expected, children=(expected,), children_by_key={"arbitrary_key": expected}
)


def test_msgspec_optional_field_accepts_none() -> None:
result = to_schema({"child": None, "children": [], "children_by_key": {}}, schema_type=OptionalCollectionEnvelope)

assert result == OptionalCollectionEnvelope(child=None, children=(), children_by_key={})


def test_msgspec_already_encoded_keys_are_unchanged() -> None:
result = to_schema({"rows": [{"itemName": "a", "itemCount": 1}], "total": 1}, schema_type=SingleWordEnvelope)

assert result == SingleWordEnvelope(rows=[CamelChild(item_name="a", item_count=1)], total=1)


def test_msgspec_encoded_key_wins_over_python_alias() -> None:
result = to_schema(
{"rows": [{"item_name": "python", "itemName": "encoded", "item_count": 1, "itemCount": 2}], "total": 1},
schema_type=SingleWordEnvelope,
)

assert result.rows == [CamelChild(item_name="encoded", item_count=2)]


def test_msgspec_arbitrary_mapping_keys_are_preserved() -> None:
metadata = {"snake_case": {"nested_key": "value"}}
result = to_schema({"metadata": metadata}, schema_type=ArbitraryPayloadEnvelope)

assert result.metadata == metadata


def test_msgspec_unknown_fields_remain_validation_errors() -> None:
with pytest.raises(msgspec.ValidationError, match="unknown_field"):
to_schema({"child": {"item_name": "a", "item_count": 1}, "unknown_field": True}, schema_type=StrictEnvelope)


def test_msgspec_ambiguous_union_is_not_normalized() -> None:
payload = {"item_name": "a", "item_count": 1}

assert schema_utils._normalize_msgspec_input(payload, CamelChild | int) is payload


def test_msgspec_field_plan_is_cached(monkeypatch: pytest.MonkeyPatch) -> None:
class CachedStruct(msgspec.Struct, rename="camel"):
item_name: str

original_fields = msgspec.structs.fields
call_count = 0

def count_fields(schema_type: type) -> tuple[msgspec.structs.FieldInfo, ...]:
nonlocal call_count
call_count += 1
return original_fields(schema_type)

monkeypatch.setattr(msgspec.structs, "fields", count_fields)
assert to_schema({"item_name": "a"}, schema_type=CachedStruct) == CachedStruct(item_name="a")
assert to_schema({"item_name": "b"}, schema_type=CachedStruct) == CachedStruct(item_name="b")
assert call_count == 1
Loading
Loading