diff --git a/README.md b/README.md index 2db346e..8e9945f 100644 --- a/README.md +++ b/README.md @@ -51,8 +51,10 @@ If you want to see a fully worked Postgres example, check out the [Postgres Quic ### Install +**NB:** Pydantic is optional, see the docs on [using Embar without Pydantic](https://embar.rdrn.me/no-pydantic). + ```bash -uv add embar +uv add embar pydantic ``` ### Set up database models diff --git a/docs/no-pydantic.md b/docs/no-pydantic.md new file mode 100644 index 0000000..313f12f --- /dev/null +++ b/docs/no-pydantic.md @@ -0,0 +1,132 @@ +# Without Pydantic + +Pydantic is an optional dependency. If you don't need validation or coercion, +you can skip it entirely — embar will load query results into plain Python +objects instead. + +Install without pydantic: + +```bash +uv add embar +``` + +Install with pydantic: + +```bash +uv add "embar[pydantic]" +``` + +## Define your schema + +Table definitions are identical regardless of whether pydantic is installed. + +```python +import sqlite3 +from typing import Annotated + +from embar.column.common import Integer, Text, integer, text +from embar.config import EmbarConfig +from embar.db.sqlite import SqliteDb +from embar.table import Table + + +class User(Table): + embar_config: EmbarConfig = EmbarConfig(table_name="users") + id: Integer = integer(primary=True) + email: Text = text("user_email", not_null=True) + + +class Message(Table): + id: Integer = integer() + user_id: Integer = integer(fk=lambda: User.id) + content: Text = text() + + +conn = sqlite3.connect(":memory:") +db = SqliteDb(conn) +db.migrate([User, Message]).run() +``` + +## Insert and select all columns + +Pass `use_pydantic=False` to get plain dataclass objects back with no validation. +`Table.all()` defaults to `use_pydantic=True`; opt out explicitly: + +```{.python continuation} +user = User(id=1, email="alice@example.com") +message = Message(id=1, user_id=1, content="Hello!") +db.insert(User).values(user).run() +db.insert(Message).values(message).run() + +results = db.select(User.all(use_pydantic=False)).from_(User).run() +assert results[0].id == 1 +assert results[0].email == "alice@example.com" +``` + +## Query with a plain model class + +Define a plain class with `Annotated` fields — no `BaseModel` required. +embar reads the annotations to build the SQL and to load results: + +```{.python continuation} +from embar.query.where import Eq + + +class UserSel: + id: Annotated[int, User.id] + email: Annotated[str, User.email] + + +results = db.select(UserSel).from_(User).where(Eq(User.id, 1)).run() +assert results[0].email == "alice@example.com" +``` + +## Nested results + +Nested tables work the same way — the plain loader parses the JSON +produced by the DB and builds the nested objects recursively: + +```{.python continuation} +class UserWithMessages: + id: Annotated[int, User.id] + messages: Annotated[list[Message], Message.many()] + + +results = ( + db.select(UserWithMessages) + .from_(User) + .left_join(Message, Eq(User.id, Message.user_id)) + .group_by(User.id) + .run() +) +assert results[0].messages[0].content == "Hello!" +``` + +## Insert with returning + +`.returning()` also accepts `use_pydantic=False`. +This example uses a table with no custom column names so the returned fields map directly: + +```{.python continuation} +class Tag(Table): + id: Integer = integer(primary=True) + name: Text = text() + + +db.migrate([Tag]).run() + +tag = Tag(id=1, name="python") +inserted = db.insert(Tag).values(tag).returning(use_pydantic=False).run() +assert inserted[0].name == "python" +``` + +## What you give up + +Without pydantic: + +- No type coercion — values are stored as-is from the database driver. +- No field validators or `BeforeValidator` transforms. +- No `ValidationError` on bad data — invalid values pass through silently. + +If you need any of these, install `embar[pydantic]` and use the default +`use_pydantic=True` (or omit the argument entirely). diff --git a/mkdocs.yml b/mkdocs.yml index 483d9c3..1fecf80 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -30,6 +30,7 @@ nav: - About: index.md - Quickstart: quickstart.md - Postgres Quickstart: postgres-quickstart.md + - Without Pydantic: no-pydantic.md - Schemas: - Basics: schemas/basics.md diff --git a/pyproject.toml b/pyproject.toml index 349474e..f15d020 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -40,9 +40,11 @@ requires-python = ">=3.14" dependencies = [ "psycopg[binary]>=3.2.11", "psycopg-pool>=3.3.0", - "pydantic>=2.12.4", ] +[project.optional-dependencies] +pydantic = ["pydantic~=2.10"] + [project.urls] homepage = "https://github.com/carderne/embar" repository = "https://github.com/carderne/embar" @@ -52,6 +54,7 @@ embar = "embar.tools.commands:main" [dependency-groups] dev = [ + "pydantic~=2.10", "mkdocs-gen-files>=0.5.0", "mkdocs-literate-nav>=0.6.2", "mkdocs-material>=9.6.23", diff --git a/src/embar/custom_types.py b/src/embar/custom_types.py index 5325c2f..a715053 100644 --- a/src/embar/custom_types.py +++ b/src/embar/custom_types.py @@ -4,8 +4,6 @@ from decimal import Decimal from typing import Any, TypeAliasType -from pydantic import Json - Undefined: Any = ... @@ -32,18 +30,5 @@ def __bool__(self) -> bool: # All the types that are allowed to ser/de to/from the DB. type PyType = ( - str - | int - | float - | Decimal - | bool - | bytes - | date - | time - | datetime - | timedelta - | dict[str, Any] - | Json[Any] - | list[PyType] - | None + str | int | float | Decimal | bool | bytes | date | time | datetime | timedelta | dict[str, Any] | list[Any] | None ) diff --git a/src/embar/db/pg.py b/src/embar/db/pg.py index 226fa68..c04687f 100644 --- a/src/embar/db/pg.py +++ b/src/embar/db/pg.py @@ -20,12 +20,12 @@ from psycopg import AsyncConnection, AsyncTransaction, Connection, Transaction from psycopg.types.json import Json from psycopg_pool import AsyncConnectionPool, ConnectionPool -from pydantic import BaseModel from embar.column.base import EnumBase from embar.db._util import get_migration_defs, merge_ddls from embar.db.base import AsyncDbBase, DbBase from embar.migration import Migration, MigrationDefs +from embar.model import DataModel from embar.query.delete import DeleteQueryReady from embar.query.insert import InsertQuery from embar.query.query import QueryMany, QuerySingle @@ -134,13 +134,13 @@ def transaction(self) -> PgDbTransaction: """ return PgDbTransaction(self) - def select[M: BaseModel](self, model: type[M]) -> SelectQuery[M, Self]: + def select[M: DataModel](self, model: type[M]) -> SelectQuery[M, Self]: """ Create a SELECT query. """ return SelectQuery[M, Self](db=self, model=model) - def select_distinct[M: BaseModel](self, model: type[M]) -> SelectDistinctQuery[M, Self]: + def select_distinct[M: DataModel](self, model: type[M]) -> SelectDistinctQuery[M, Self]: """ Create a SELECT query. """ @@ -358,13 +358,13 @@ def transaction(self) -> AsyncPgDbTransaction: """ return AsyncPgDbTransaction(self) - def select[M: BaseModel](self, model: type[M]) -> SelectQuery[M, Self]: + def select[M: DataModel](self, model: type[M]) -> SelectQuery[M, Self]: """ Create a SELECT query. """ return SelectQuery[M, Self](db=self, model=model) - def select_distinct[M: BaseModel](self, model: type[M]) -> SelectDistinctQuery[M, Self]: + def select_distinct[M: DataModel](self, model: type[M]) -> SelectDistinctQuery[M, Self]: """ Create a SELECT query. """ diff --git a/src/embar/db/sqlite.py b/src/embar/db/sqlite.py index 7ea8a2a..f7553a3 100644 --- a/src/embar/db/sqlite.py +++ b/src/embar/db/sqlite.py @@ -13,12 +13,11 @@ override, ) -from pydantic import BaseModel - from embar.column.base import EnumBase from embar.db._util import get_migration_defs, merge_ddls from embar.db.base import DbBase from embar.migration import Migration, MigrationDefs +from embar.model import DataModel from embar.query.delete import DeleteQueryReady from embar.query.insert import InsertQuery from embar.query.query import QueryMany, QuerySingle @@ -60,13 +59,13 @@ def transaction(self) -> SqliteDbTransaction: db_copy._commit_after_execute = False return SqliteDbTransaction(db_copy) - def select[M: BaseModel](self, model: type[M]) -> SelectQuery[M, Self]: + def select[M: DataModel](self, model: type[M]) -> SelectQuery[M, Self]: """ Create a SELECT query. """ return SelectQuery[M, Self](db=self, model=model) - def select_distinct[M: BaseModel](self, model: type[M]) -> SelectDistinctQuery[M, Self]: + def select_distinct[M: DataModel](self, model: type[M]) -> SelectDistinctQuery[M, Self]: """ Create a SELECT query. """ diff --git a/src/embar/model.py b/src/embar/model.py index d954375..11fef30 100644 --- a/src/embar/model.py +++ b/src/embar/model.py @@ -1,15 +1,38 @@ import json +from dataclasses import field, make_dataclass from typing import ( + TYPE_CHECKING, Annotated, Any, + ClassVar, Literal, + Protocol, cast, get_args, get_origin, get_type_hints, + overload, ) -from pydantic import BaseModel, BeforeValidator, Field, create_model +try: + from pydantic import BaseModel + + _PYDANTIC_AVAILABLE = True +except ImportError: + # Minimal stub so that `class SelectAllPydantic(BaseModel)` and + # `isinstance(x, BaseModel)` work at runtime even without pydantic. + # The stub is never used for actual validation — that path is guarded + # by _require_pydantic(). + class BaseModel: + """Stub used when pydantic is not installed.""" + + pass + + _PYDANTIC_AVAILABLE = False + +if TYPE_CHECKING: + # Re-import the real thing for the type checker. + from pydantic import BaseModel from embar.column.base import ColumnBase from embar.db.base import DbType @@ -18,18 +41,52 @@ from embar.table_base import TableBase -class SelectAll(BaseModel): +def _require_pydantic(feature: str) -> None: + """Raise a clear ImportError if pydantic is not installed.""" + try: + import pydantic # noqa: F401 + except ImportError: + raise ImportError(f"{feature} requires pydantic. Install it with: pip install 'embar[pydantic]'") from None + + +class DataclassType(Protocol): + """Protocol for plain (non-pydantic) dataclass models.""" + + __dataclass_fields__: ClassVar[dict[str, Any]] + + +class HasAnnotations(Protocol): + """ + Protocol satisfied by any class that carries `__annotations__` — i.e. every + Python class that declares at least one field-level type hint. + + This is the minimal structural requirement for `to_sql_columns` and + `load_dataclass` to work: they only need `get_type_hints()` to succeed. """ - `SelectAll` tells the query engine to get all fields from the `from()` table ONLY. - Ideally it could get fields from joined tables too, but no way for that to work (from a typing POV) - Not recommended for public use, users should rather use their table's `all()` method. + __annotations__: ClassVar[dict[str, Any]] + + +type DataModel = BaseModel | DataclassType | HasAnnotations + + +class SelectAllPydantic(BaseModel): + """ + `SelectAll` version that validates with Pydantic. """ ... -def to_sql_columns(model: type[BaseModel], db_type: DbType) -> str: +class SelectAllDataclass: + """ + `SelectAll` version that doesn't validate (plain dataclass). + """ + + __dataclass_fields__: ClassVar[dict[str, Any]] = {} + + +def to_sql_columns(model: type[DataModel], db_type: DbType) -> str: parts: list[str] = [] hints = get_type_hints(model, include_extras=True) for field_name, field_type in hints.items(): @@ -105,14 +162,13 @@ def _get_source_expr(field_name: str, field_type: type, db_type: DbType, hints: def _convert_annotation( field_type: type, + use_pydantic: bool, ) -> Annotated[Any, Any] | Literal[False]: """ Extract complex annotated types from `Annotated[int, MyTable.my_col]` expressions. If the annotated type is a column reference then this does nothing and returns false. - Only used by `embar.query.Select` but more at home here with the context where it's used. - ```python from typing import Annotated from pydantic import BaseModel @@ -123,6 +179,7 @@ class MyTable(Table): my_col: Text = text() class MyModel(BaseModel): my_col: Annotated[str, MyTable.my_col] + ``` """ if get_origin(field_type) is Annotated: annotations = get_args(field_type) @@ -131,35 +188,52 @@ class MyModel(BaseModel): if isinstance(annotation, ManyTable): many_table = cast(ManyTable[type[TableBase]], annotation) inner_type = many_table.of - dc = generate_model(inner_type) + dc = generate_model(inner_type, use_pydantic) new_type = Annotated[list[dc], annotation] return new_type if isinstance(annotation, OneTable): one_table = cast(OneTable[type[TableBase]], annotation) inner_type = one_table.of - dc = generate_model(inner_type) + dc = generate_model(inner_type, use_pydantic) new_type = Annotated[dc, annotation] return new_type return False -def generate_model(cls: type[TableBase]) -> type[BaseModel]: +@overload +def generate_model(cls: type[TableBase], use_pydantic: Literal[True]) -> type[BaseModel]: ... +@overload +def generate_model(cls: type[TableBase], use_pydantic: Literal[False]) -> type[DataclassType]: ... +@overload +def generate_model(cls: type[TableBase], use_pydantic: bool) -> type[DataModel]: ... + + +def generate_model(cls: type[TableBase], use_pydantic: bool) -> type[DataModel]: + if use_pydantic: + return generate_pydantic_model(cls) + return generate_dataclass_model(cls) + + +def generate_pydantic_model(cls: type[TableBase]) -> type[BaseModel]: """ - Create a model based on a `Table`. + Create a Pydantic model based on a `Table`. - Note the new table has the same exact name, maybe something to revisit. + Note the new model has the same exact name, maybe something to revisit. ```python from embar.table import Table - from embar.model import generate_model + from embar.model import generate_pydantic_model class MyTable(Table): ... - generate_model(MyTable) + generate_pydantic_model(MyTable) ``` """ + _require_pydantic("generate_pydantic_model") + from pydantic import BeforeValidator, create_model + from pydantic import Field as PydanticField fields_dict: dict[str, Any] = {} - for field_name, column in cls._fields.items(): + for field_name, column in cls._fields.items(): # pyright:ignore[reportPrivateUsage] field_type = column.info.py_type if column.info.col_type == "VECTOR": @@ -167,7 +241,7 @@ class MyTable(Table): ... fields_dict[field_name] = ( Annotated[field_type, column], - Field(default_factory=lambda a=column: column.info.fqn()), + PydanticField(default_factory=lambda a=column: column.info.fqn()), ) model = create_model(cls.__name__, **fields_dict) @@ -175,21 +249,158 @@ class MyTable(Table): ... return model -def upgrade_model_nested_fields[B: BaseModel](model: type[B]) -> type[B]: +def generate_dataclass_model(cls: type[TableBase]) -> type[DataclassType]: + """ + Create a plain dataclass based on a `Table` (no Pydantic validation). + + Fields are typed as `Annotated[py_type, column]` so that `to_sql_columns` + can discover the SQL column reference. + + Note the new dataclass has the same exact name, maybe something to revisit. + + ```python + from embar.table import Table + from embar.model import generate_dataclass_model + class MyTable(Table): ... + generate_dataclass_model(MyTable) + ``` + """ + dc_fields: list[Any] = [] + for field_name, column in cls._fields.items(): # pyright:ignore[reportPrivateUsage] + field_type = column.info.py_type + # Use Annotated so to_sql_columns can find the column reference + annotated_type = Annotated[field_type, column] + dc_fields.append( + ( + field_name, + annotated_type, + field(default=None), + ) + ) + + data_class = make_dataclass(cls.__name__, dc_fields) + return data_class + + +def upgrade_model_nested_fields[B: DataModel](model: type[B], use_pydantic: bool) -> type[B]: + """ + Upgrade a model so that nested `ManyTable`/`OneTable` fields are resolved to concrete models. + + For Pydantic models, creates a new subclass via `create_model`. + For plain dataclasses/annotated classes, creates a new dataclass with upgraded field types. + + ``use_pydantic`` controls whether nested table models are generated as Pydantic models + or plain dataclasses, and must be supplied explicitly by the caller. + """ type_hints = get_type_hints(model, include_extras=True) - fields_dict: dict[str, Any] = {} + # Without pydantic, BaseModel is the stub class; no real subclass of it can exist, + # so this branch is only reachable when _PYDANTIC_AVAILABLE is True anyway. + if isinstance(model, type) and issubclass(model, BaseModel): + from pydantic import create_model + + fields_dict: dict[str, Any] = {} + for field_name, field_type in type_hints.items(): + new_type = _convert_annotation(field_type, use_pydantic=True) + if new_type: + fields_dict[field_name] = (new_type, None) + else: + fields_dict[field_name] = (field_type, None) + + new_class = create_model(model.__name__, __base__=model, **fields_dict) + new_class.model_rebuild() + return new_class + + # Plain dataclass / annotated-class path + dc_fields: list[Any] = [] for field_name, field_type in type_hints.items(): - new_type = _convert_annotation(field_type) - if new_type: - fields_dict[field_name] = (new_type, None) - else: - fields_dict[field_name] = (field_type, None) + new_type = _convert_annotation(field_type, use_pydantic=use_pydantic) + resolved_type = new_type if new_type else field_type + dc_fields.append((field_name, resolved_type, field(default=None))) + + new_class = make_dataclass(model.__name__, dc_fields) + return new_class # type: ignore[return-value] + + +def load_results[T](model: type[T], data: list[dict[str, Any]]) -> list[T]: + """ + Load query result rows into model instances. + + Dispatches between Pydantic validation (for ``BaseModel`` subclasses) and + the plain dict→dataclass loader for everything else. + """ + if _PYDANTIC_AVAILABLE and isinstance(model, type) and issubclass(model, BaseModel): + from pydantic import TypeAdapter + + adapter = TypeAdapter(list[model]) + return adapter.validate_python(data) + return load_dataclass(model, data) + + +def load_dataclass[T](model: type[T], data: list[dict[str, Any]]) -> list[T]: + """ + Load a list of row dicts into plain dataclass/annotated-class instances + (no Pydantic validation). + + Handles nested dataclasses (from ManyTable/OneTable annotations) by recursively + loading JSON objects/arrays from the database into the appropriate types. + """ + return [_load_one(model, row) for row in data] + + +def _load_one[T](model: type[T], row: dict[str, Any]) -> T: + """ + Load a single row dict into a dataclass instance. + """ + hints = get_type_hints(model, include_extras=True) + kwargs: dict[str, Any] = {} + for field_name, field_type in hints.items(): + raw = row.get(field_name) + kwargs[field_name] = _coerce_field(field_type, raw) + return model(**kwargs) + + +def _coerce_field(field_type: type, value: Any) -> Any: + """ + Coerce a raw value into the expected Python type for a dataclass field. + + Handles nested dataclasses (list[SomeDataclass] or SomeDataclass) by parsing + JSON strings/dicts from the database. + """ + if value is None: + return None + + origin = get_origin(field_type) + args = get_args(field_type) + + # Unwrap Annotated[T, ...] + if origin is Annotated: + return _coerce_field(args[0], value) + + # list[SomeDataclass] — parameterised list, e.g. list[Message] + if origin is list and args: + inner = args[0] + if _is_plain_dataclass(inner): + items = value if isinstance(value, list) else json.loads(value) + return [_load_one(inner, item) for item in items] + return value + + # Bare list (no type args) — used by VECTOR columns whose py_type is plain `list`. + # The DB returns either a Python list (postgres array) or a JSON string (sqlite). + if field_type is list: + return value if isinstance(value, list) else json.loads(value) + + # SomeDataclass + if _is_plain_dataclass(field_type): + data = value if isinstance(value, dict) else json.loads(value) + return _load_one(field_type, data) + + return value - new_class = create_model(model.__name__, __base__=model, **fields_dict) - new_class.model_rebuild() - return new_class +def _is_plain_dataclass(t: Any) -> bool: + """Return True if `t` is a plain (non-Pydantic) dataclass type.""" + return isinstance(t, type) and hasattr(t, "__dataclass_fields__") and not issubclass(t, BaseModel) def _parse_json_list(v: Any): diff --git a/src/embar/query/delete.py b/src/embar/query/delete.py index 84e295e..5e6ef3a 100644 --- a/src/embar/query/delete.py +++ b/src/embar/query/delete.py @@ -4,12 +4,12 @@ from textwrap import dedent from typing import Any, Self, cast -from pydantic import BaseModel, TypeAdapter - from embar.column.base import ColumnBase from embar.db.base import AllDbBase, AsyncDbBase, DbBase from embar.model import ( + DataModel, generate_model, + load_results, ) from embar.query.clause_base import ClauseBase from embar.query.order_by import Asc, BareColumn, Desc, OrderBy, RawSqlOrder @@ -41,8 +41,15 @@ def __init__(self, table: type[T], db: Db): self.table = table self._db = db - def returning(self) -> DeleteQueryReturning[T, Db]: - return DeleteQueryReturning(self.table, self._db, self._where_clause, self._order_clause, self._limit_value) + def returning(self, use_pydantic: bool = True) -> DeleteQueryReturning[T, Db]: + return DeleteQueryReturning( + table=self.table, + db=self._db, + use_pydantic=use_pydantic, + where_clause=self._where_clause, + order_clause=self._order_clause, + limit_value=self._limit_value, + ) def where(self, where_clause: ClauseBase) -> Self: """ @@ -112,7 +119,8 @@ def __await__(self): The overrides provide for a few different cases: - A Model was passed, in which case that's the return type - - `SelectAll` was passed, in which case the return type is the `Table` + - `SelectAllPydantic` or `SelectAllDataclass` was passed (via `returning()`), + in which case the return type is the `Table` - This is called with an async db, in which case an error is returned. """ query = self.sql() @@ -183,6 +191,7 @@ class DeleteQueryReturning[T: Table, Db: AllDbBase]: table: type[T] _db: Db + _use_pydantic: bool _where_clause: ClauseBase | None = None _order_clause: OrderBy | None = None @@ -192,15 +201,17 @@ def __init__( self, table: type[T], db: Db, + use_pydantic: bool, where_clause: ClauseBase | None, order_clause: OrderBy | None, limit_value: int | None, ): """ - Create a new SelectQueryReady instance. + Create a new DeleteQueryReturning instance. """ self.table = table self._db = db + self._use_pydantic = use_pydantic self._where_clause = where_clause self._order_clause = order_clause self._limit_value = limit_value @@ -214,13 +225,13 @@ def __await__(self) -> Generator[Any, None, list[T]]: The overrides provide for a few different cases: - A Model was passed, in which case that's the return type - - `SelectAll` was passed, in which case the return type is the `Table` + - `SelectAllPydantic` or `SelectAllDataclass` was passed (via `returning()`), + in which case the return type is the `Table` - This is called with an async db, in which case an error is returned. """ query = self.sql() model = self._get_model() model = cast(type[T], model) - adapter = TypeAdapter(list[model]) async def awaitable(): db = self._db @@ -229,7 +240,7 @@ async def awaitable(): else: db = cast(DbBase, self._db) data = db.fetch(query) - results = adapter.validate_python(data) + results = load_results(model, data) return results return awaitable().__await__() @@ -244,10 +255,9 @@ def run(self) -> list[T]: query = self.sql() model = self._get_model() model = cast(type[T], model) - adapter = TypeAdapter(list[model]) db = cast(DbBase, self._db) data = db.fetch(query) - results = adapter.validate_python(data) + results = load_results(model, data) return results def sql(self) -> QuerySingle: @@ -290,9 +300,9 @@ def get_count() -> int: return QuerySingle(sql, params=params) - def _get_model(self) -> type[BaseModel]: + def _get_model(self) -> type[DataModel]: """ Generate the dataclass that will be used to deserialize (and validate) the query results. """ - model = generate_model(self.table) + model = generate_model(self.table, self._use_pydantic) return model diff --git a/src/embar/query/insert.py b/src/embar/query/insert.py index 50fd3bd..61eacef 100644 --- a/src/embar/query/insert.py +++ b/src/embar/query/insert.py @@ -3,11 +3,9 @@ from collections.abc import Generator, Sequence from typing import Any, Self, cast -from pydantic import BaseModel, TypeAdapter - from embar.custom_types import PyType from embar.db.base import AllDbBase, AsyncDbBase, DbBase -from embar.model import generate_model +from embar.model import DataModel, generate_model, load_results from embar.query.conflict import OnConflict, OnConflictDoNothing, OnConflictDoUpdate, TupleAtLeastOne from embar.query.query import QueryMany from embar.table import Table @@ -65,8 +63,14 @@ def __init__(self, table: type[T], db: Db, items: Sequence[T]): self._db = db self.items = items - def returning(self) -> InsertQueryReturning[T, Db]: - return InsertQueryReturning(self.table, self._db, self.items, on_conflict=self.on_conflict) + def returning(self, use_pydantic: bool = True) -> InsertQueryReturning[T, Db]: + return InsertQueryReturning( + table=self.table, + db=self._db, + use_pydantic=use_pydantic, + items=self.items, + on_conflict=self.on_conflict, + ) def on_conflict_do_nothing(self, target: TupleAtLeastOne | None = None) -> Self: self.on_conflict = OnConflictDoNothing(target) @@ -152,16 +156,18 @@ class InsertQueryReturning[T: Table, Db: AllDbBase]: """ _db: Db + _use_pydantic: bool table: type[T] items: Sequence[T] on_conflict: OnConflict | None - def __init__(self, table: type[T], db: Db, items: Sequence[T], on_conflict: OnConflict | None): + def __init__(self, table: type[T], db: Db, use_pydantic: bool, items: Sequence[T], on_conflict: OnConflict | None): """ Create a new InsertQueryReturning instance. """ self.table = table self._db = db + self._use_pydantic = use_pydantic self.items = items self.on_conflict = on_conflict @@ -174,7 +180,6 @@ def __await__(self) -> Generator[Any, None, Sequence[T]]: query = self.sql() model = self._get_model() model = cast(type[T], model) - adapter = TypeAdapter(list[model]) async def awaitable(): db = self._db @@ -183,7 +188,7 @@ async def awaitable(): else: db = cast(DbBase, self._db) data = db.fetch(query) - results = adapter.validate_python(data) + results = load_results(model, data) return results return awaitable().__await__() @@ -198,10 +203,9 @@ def run(self) -> list[T]: query = self.sql() model = self._get_model() model = cast(type[T], model) - adapter = TypeAdapter(list[model]) db = cast(DbBase, self._db) data = db.fetch(query) - results = adapter.validate_python(data) + results = load_results(model, data) return results def sql(self) -> QueryMany: @@ -231,9 +235,9 @@ def get_count() -> int: sql += " RETURNING *" return QueryMany(sql, many_params=values) - def _get_model(self) -> type[BaseModel]: + def _get_model(self) -> type[DataModel]: """ Generate the dataclass that will be used to deserialize (and validate) the query results. """ - model = generate_model(self.table) + model = generate_model(self.table, self._use_pydantic) return model diff --git a/src/embar/query/select.py b/src/embar/query/select.py index 64adfc5..b3ead98 100644 --- a/src/embar/query/select.py +++ b/src/embar/query/select.py @@ -5,13 +5,15 @@ from typing import Any, Self, cast, overload from warnings import deprecated -from pydantic import BaseModel, TypeAdapter - from embar.column.base import ColumnBase from embar.db.base import AllDbBase, AsyncDbBase, DbBase from embar.model import ( - SelectAll, + BaseModel, + DataModel, + SelectAllDataclass, + SelectAllPydantic, generate_model, + load_results, to_sql_columns, upgrade_model_nested_fields, ) @@ -25,7 +27,7 @@ from embar.table import Table -class SelectQuery[M: BaseModel, Db: AllDbBase]: +class SelectQuery[M: DataModel, Db: AllDbBase]: """ `SelectQuery` is returned by Db.select and exposes one method that produced the `SelectQueryReady`. """ @@ -54,7 +56,7 @@ def from_[T: Table](self, table: type[T]) -> SelectQueryReady[M, T, Db]: return SelectQueryReady[M, T, Db](model=self.model, table=table, db=self._db, distinct=False) -class SelectDistinctQuery[M: BaseModel, Db: AllDbBase]: +class SelectDistinctQuery[M: DataModel, Db: AllDbBase]: """ `SelectDistinctQuery` is returned by Db.select and exposes one method that produced the `SelectQueryReady`. @@ -85,7 +87,7 @@ def from_[T: Table](self, table: type[T]) -> SelectQueryReady[M, T, Db]: return SelectQueryReady[M, T, Db](model=self.model, table=table, db=self._db, distinct=True) -class SelectQueryReady[M: BaseModel, T: Table, Db: AllDbBase]: +class SelectQueryReady[M: DataModel, T: Table, Db: AllDbBase]: """ `SelectQueryReady` is used to insert data into a table. @@ -301,7 +303,9 @@ class User(Table): return self @overload - def __await__(self: SelectQueryReady[SelectAll, T, Db]) -> Generator[Any, None, Sequence[T]]: ... + def __await__(self: SelectQueryReady[SelectAllPydantic, T, Db]) -> Generator[Any, None, Sequence[T]]: ... + @overload + def __await__(self: SelectQueryReady[SelectAllDataclass, T, Db]) -> Generator[Any, None, Sequence[T]]: ... @overload def __await__(self: SelectQueryReady[M, T, Db]) -> Generator[Any, None, Sequence[M]]: ... @@ -314,13 +318,12 @@ def __await__(self) -> Generator[Any, None, Sequence[T | M]]: The overrides provide for a few different cases: - A Model was passed, in which case that's the return type - - `SelectAll` was passed, in which case the return type is the `Table` + - `SelectAllPydantic` or `SelectAllDataclass` was passed, in which case the return type is the `Table` - This is called with an async db, in which case an error is returned. """ query = self.sql() model = self._get_model() model = cast(type[T] | type[M], model) - adapter = TypeAdapter(list[model]) async def awaitable(): db = self._db @@ -329,13 +332,15 @@ async def awaitable(): else: db = cast(DbBase, self._db) data = db.fetch(query) - results = adapter.validate_python(data) + results = load_results(model, data) return results return awaitable().__await__() @overload - def run(self: SelectQueryReady[SelectAll, T, Db]) -> Sequence[T]: ... + def run(self: SelectQueryReady[SelectAllPydantic, T, Db]) -> Sequence[T]: ... + @overload + def run(self: SelectQueryReady[SelectAllDataclass, T, Db]) -> Sequence[T]: ... @overload def run(self) -> Sequence[M]: ... @@ -349,24 +354,32 @@ def run(self) -> Sequence[M | T]: query = self.sql() model = self._get_model() model = cast(type[T] | type[M], model) - adapter = TypeAdapter(list[model]) db = cast(DbBase, self._db) data = db.fetch(query) - results = adapter.validate_python(data) + results = load_results(model, data) return results - def _get_model(self) -> type[BaseModel] | type[M]: + def _get_model(self) -> type[DataModel] | type[M]: """ Generate the dataclass that will be used to deserialize (and validate) the query results. - If the model is `SelectAll`, we generate a dataclass based on the `Table`, - otherwise the model itself - is used. + If the model is `SelectAllPydantic` or `SelectAllDataclass`, we generate a model + based on the `Table`, otherwise the model itself is used. Extra processing is done to check for nested children that are Tables themselves. """ - model = generate_model(self.table) if self.model is SelectAll else self.model - upgraded = upgrade_model_nested_fields(model) + + if self.model is SelectAllPydantic: + model = generate_model(self.table, use_pydantic=True) + use_pydantic = True + elif self.model is SelectAllDataclass: + model = generate_model(self.table, use_pydantic=False) + use_pydantic = False + else: + model = self.model + use_pydantic = isinstance(model, type) and issubclass(model, BaseModel) + + upgraded = upgrade_model_nested_fields(model, use_pydantic=use_pydantic) return upgraded def sql(self) -> QuerySingle: diff --git a/src/embar/query/update.py b/src/embar/query/update.py index d09c69a..9d0c074 100644 --- a/src/embar/query/update.py +++ b/src/embar/query/update.py @@ -3,10 +3,8 @@ from collections.abc import Generator, Mapping, Sequence from typing import Any, Self, cast -from pydantic import BaseModel, TypeAdapter - from embar.db.base import AllDbBase, AsyncDbBase, DbBase -from embar.model import generate_model +from embar.model import DataModel, generate_model, load_results from embar.query.clause_base import ClauseBase from embar.query.query import QuerySingle from embar.table import Table @@ -70,8 +68,14 @@ def where(self, where_clause: ClauseBase) -> Self: self._where_clause = where_clause return self - def returning(self) -> UpdateQueryReturning[T, Db]: - return UpdateQueryReturning(self.table, self._db, self.data, self._where_clause) + def returning(self, use_pydantic: bool = True) -> UpdateQueryReturning[T, Db]: + return UpdateQueryReturning( + table=self.table, + db=self._db, + use_pydantic=use_pydantic, + data=self.data, + where_clause=self._where_clause, + ) def __await__(self): """ @@ -144,15 +148,24 @@ class UpdateQueryReturning[T: Table, Db: AllDbBase]: table: type[T] _db: Db + _use_pydantic: bool data: Mapping[str, Any] _where_clause: ClauseBase | None = None - def __init__(self, table: type[T], db: Db, data: Mapping[str, Any], where_clause: ClauseBase | None): + def __init__( + self, + table: type[T], + db: Db, + use_pydantic: bool, + data: Mapping[str, Any], + where_clause: ClauseBase | None, + ): """ Create a new UpdateQueryReturning instance. """ self.table = table self._db = db + self._use_pydantic = use_pydantic self.data = data self._where_clause = where_clause @@ -165,7 +178,6 @@ def __await__(self) -> Generator[Any, None, Sequence[T]]: query = self.sql() model = self._get_model() model = cast(type[T], model) - adapter = TypeAdapter(list[model]) async def awaitable(): db = self._db @@ -174,7 +186,7 @@ async def awaitable(): else: db = cast(DbBase, self._db) data = db.fetch(query) - results = adapter.validate_python(data) + results = load_results(model, data) return results return awaitable().__await__() @@ -189,10 +201,9 @@ def run(self) -> list[T]: query = self.sql() model = self._get_model() model = cast(type[T], model) - adapter = TypeAdapter(list[model]) db = cast(DbBase, self._db) data = db.fetch(query) - results = adapter.validate_python(data) + results = load_results(model, data) return results def sql(self) -> QuerySingle: @@ -231,9 +242,9 @@ def get_count() -> int: return QuerySingle(sql, params) - def _get_model(self) -> type[BaseModel]: + def _get_model(self) -> type[DataModel]: """ Generate the dataclass that will be used to deserialize (and validate) the query results. """ - model = generate_model(self.table) + model = generate_model(self.table, self._use_pydantic) return model diff --git a/src/embar/query/where.py b/src/embar/query/where.py index f136213..514bdbc 100644 --- a/src/embar/query/where.py +++ b/src/embar/query/where.py @@ -183,7 +183,7 @@ def sql(self, get_count: GetCount) -> QuerySingle: # String matching operators class Like[T: PyType](ClauseBase): left: ColumnInfo - right: PyType | ColumnInfo + right: T | ColumnInfo def __init__(self, left: Column[T], right: T | Column[T]): self.left = left.info @@ -205,7 +205,7 @@ class Ilike[T: PyType](ClauseBase): """ left: ColumnInfo - right: PyType | ColumnInfo + right: T | ColumnInfo def __init__(self, left: Column[T], right: T | Column[T]): self.left = left.info @@ -227,7 +227,7 @@ class NotLike[T: PyType](ClauseBase): """ left: ColumnInfo - right: PyType | ColumnInfo + right: T | ColumnInfo def __init__(self, left: Column[T], right: T | Column[T]): self.left = left.info diff --git a/src/embar/sql_db.py b/src/embar/sql_db.py index d77b50c..bfa354d 100644 --- a/src/embar/sql_db.py +++ b/src/embar/sql_db.py @@ -4,10 +4,8 @@ from string.templatelib import Template from typing import Any, Self, cast -from pydantic import BaseModel, TypeAdapter - from embar.db.base import AllDbBase, AsyncDbBase, DbBase -from embar.model import upgrade_model_nested_fields +from embar.model import BaseModel, DataModel, load_results, upgrade_model_nested_fields from embar.query.query import QuerySingle from embar.sql import Sql @@ -27,7 +25,7 @@ def __init__(self, template: Template, db: Db): self._sql = Sql(template) self._db = db - def model[M: BaseModel](self, model: type[M]) -> DbSqlReturning[M, Db]: + def model[M: DataModel](self, model: type[M]) -> DbSqlReturning[M, Db]: """ Specify a model for parsing results. """ @@ -68,7 +66,7 @@ def run(self) -> Self: return self -class DbSqlReturning[M: BaseModel, Db: AllDbBase]: +class DbSqlReturning[M: DataModel, Db: AllDbBase]: """ Used to run raw SQL queries and return a value. """ @@ -95,7 +93,6 @@ def __await__(self) -> Generator[Any, None, Sequence[M]]: sql = self._sql.sql() query = QuerySingle(sql) model = self._get_model() - adapter = TypeAdapter(list[model]) async def awaitable(): db = self._db @@ -105,7 +102,7 @@ async def awaitable(): else: db = cast(DbBase, self._db) data = db.fetch(query) - results = adapter.validate_python(data) + results = load_results(model, data) return results return awaitable().__await__() @@ -119,15 +116,14 @@ def run(self) -> Sequence[M]: sql = self._sql.sql() query = QuerySingle(sql) model = self._get_model() - adapter = TypeAdapter(list[model]) db = cast(DbBase, self._db) data = db.fetch(query) - self.model.__init_subclass__() - results = adapter.validate_python(data) + results = load_results(model, data) return results def _get_model(self) -> type[M]: """ Generate the dataclass that will be used to deserialize (and validate) the query results. """ - return upgrade_model_nested_fields(self.model) + use_pydantic = isinstance(self.model, type) and issubclass(self.model, BaseModel) + return upgrade_model_nested_fields(self.model, use_pydantic=use_pydantic) diff --git a/src/embar/table.py b/src/embar/table.py index cc99de5..d759898 100644 --- a/src/embar/table.py +++ b/src/embar/table.py @@ -5,9 +5,17 @@ """ from textwrap import dedent, indent -from typing import Any, Self, dataclass_transform, get_args, get_origin +from typing import TYPE_CHECKING, Any, Literal, Self, dataclass_transform, get_args, get_origin, overload -from pydantic_core import core_schema +if TYPE_CHECKING: + from pydantic_core import core_schema as _core_schema + +try: + from pydantic_core import core_schema + + _PYDANTIC_AVAILABLE = True +except ImportError: + _PYDANTIC_AVAILABLE = False from embar.column.base import ColumnBase from embar.column.common import Column, Null, float_col, integer, text @@ -34,7 +42,7 @@ ) from embar.config import EmbarConfig from embar.custom_types import Undefined -from embar.model import SelectAll +from embar.model import SelectAllDataclass, SelectAllPydantic from embar.query.many import ManyTable, OneTable from embar.table_base import TableBase @@ -159,13 +167,15 @@ def __init__(self, **kwargs: Any) -> None: if missing: raise TypeError(f"Missing required fields: {missing}") - @classmethod - def __get_pydantic_core_schema__( - cls, - source_type: Any, - handler: Any, - ) -> core_schema.CoreSchema: - return core_schema.any_schema() + if _PYDANTIC_AVAILABLE: + + @classmethod + def __get_pydantic_core_schema__( + cls, + source_type: Any, + handler: Any, + ) -> "_core_schema.CoreSchema": + return core_schema.any_schema() @classmethod def many(cls) -> ManyTable[type[Self]]: @@ -213,20 +223,38 @@ def ddl(cls) -> str: return sql + @overload + @classmethod + def all(cls) -> type[SelectAllPydantic]: ... + @overload + @classmethod + def all(cls, use_pydantic: Literal[True]) -> type[SelectAllPydantic]: ... + @overload + @classmethod + def all(cls, use_pydantic: Literal[False]) -> type[SelectAllDataclass]: ... + @classmethod - def all(cls) -> type[SelectAll]: + def all(cls, use_pydantic: bool = True) -> type[SelectAllPydantic] | type[SelectAllDataclass]: """ Generate a Select query model that returns all the table's fields. ```python - from embar.model import SelectAll + from embar.model import SelectAllPydantic from embar.table import Table class MyTable(Table): ... model = MyTable.all() - assert model == SelectAll + assert model == SelectAllPydantic ``` """ - return SelectAll + if use_pydantic: + if not _PYDANTIC_AVAILABLE: + raise ImportError( + "Table.all() requires pydantic when use_pydantic=True (the default). " + "Either install it with: pip install 'embar[pydantic]' " + "or opt in to the plain-dataclass path with: MyTable.all(use_pydantic=False)" + ) + return SelectAllPydantic + return SelectAllDataclass def value_dict(self) -> dict[str, Any]: """ diff --git a/tests/test_plain_models.py b/tests/test_plain_models.py new file mode 100644 index 0000000..76c448b --- /dev/null +++ b/tests/test_plain_models.py @@ -0,0 +1,437 @@ +"""Tests for non-pydantic (plain class / plain dataclass) model support.""" + +from typing import Annotated + +import pytest + +from embar.column.common import Integer, Text, integer, text +from embar.config import EmbarConfig +from embar.db.pg import PgDb +from embar.model import ( + SelectAllDataclass, + SelectAllPydantic, + generate_dataclass_model, + generate_pydantic_model, + load_dataclass, + load_results, + upgrade_model_nested_fields, +) +from embar.table import Table + +# --------------------------------------------------------------------------- +# Test schema +# --------------------------------------------------------------------------- + + +class Author(Table): + embar_config: EmbarConfig = EmbarConfig(table_name="authors") + id: Integer = integer(primary=True) + name: Text = text() + + +class Book(Table): + embar_config: EmbarConfig = EmbarConfig(table_name="books") + id: Integer = integer(primary=True) + title: Text = text() + author_id: Integer = integer(fk=lambda: Author.id) + + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + + +@pytest.fixture(scope="module") +def db_dummy() -> PgDb: + return PgDb(None) # ty: ignore[invalid-argument-type] + + +# --------------------------------------------------------------------------- +# Table.all() tests +# --------------------------------------------------------------------------- + + +def test_table_all_default_returns_pydantic(): + """Table.all() with no args returns SelectAllPydantic by default (backward-compatible).""" + assert Author.all() is SelectAllPydantic + + +def test_table_all_use_pydantic_false_returns_dataclass(): + """Table.all(use_pydantic=False) explicitly returns SelectAllDataclass.""" + assert Author.all(use_pydantic=False) is SelectAllDataclass + + +def test_table_all_use_pydantic_true_returns_pydantic(): + """Table.all(use_pydantic=True) returns SelectAllPydantic.""" + assert Author.all(use_pydantic=True) is SelectAllPydantic + + +# --------------------------------------------------------------------------- +# generate_dataclass_model tests +# --------------------------------------------------------------------------- + + +def test_generate_dataclass_model_creates_dataclass(): + """generate_dataclass_model returns a plain dataclass with the right fields.""" + dc = generate_dataclass_model(Author) + assert hasattr(dc, "__dataclass_fields__") + assert "id" in dc.__dataclass_fields__ + assert "name" in dc.__dataclass_fields__ + + +def test_generate_pydantic_model_creates_pydantic(): + """generate_pydantic_model returns a Pydantic BaseModel with the right fields.""" + from pydantic import BaseModel + + m = generate_pydantic_model(Author) + assert issubclass(m, BaseModel) + assert "id" in m.model_fields + assert "name" in m.model_fields + + +# --------------------------------------------------------------------------- +# load_dataclass tests +# --------------------------------------------------------------------------- + + +def test_load_dataclass_simple(): + """load_dataclass populates plain dataclass fields from row dicts.""" + dc = generate_dataclass_model(Author) + rows = [{"id": 1, "name": "Alice"}, {"id": 2, "name": "Bob"}] + results = load_dataclass(dc, rows) + assert len(results) == 2 + assert results[0].id == 1 + assert results[0].name == "Alice" + assert results[1].id == 2 + assert results[1].name == "Bob" + + +def test_load_dataclass_missing_field_becomes_none(): + """load_dataclass sets missing fields to None rather than raising.""" + dc = generate_dataclass_model(Author) + rows = [{"id": 42}] + results = load_dataclass(dc, rows) + assert results[0].id == 42 + assert results[0].name is None + + +# --------------------------------------------------------------------------- +# Plain (non-pydantic) model in SELECT query — SQL generation +# --------------------------------------------------------------------------- + + +def test_select_plain_model_sql_generation(db_dummy: PgDb): + """A plain class with Annotated fields can drive SQL generation.""" + db = db_dummy + + class AuthorSel: + id: Annotated[int, Author.id] + name: Annotated[str, Author.name] + + query = db.select(AuthorSel).from_(Author) + sql_result = query.sql() + + assert '"authors"."id" AS "id"' in sql_result.sql + assert '"authors"."name" AS "name"' in sql_result.sql + assert "FROM" in sql_result.sql + + +def test_select_all_dataclass_sql_generation(db_dummy: PgDb): + """SelectAllDataclass (opt-in via use_pydantic=False) drives correct SQL generation.""" + db = db_dummy + + query = db.select(Author.all(use_pydantic=False)).from_(Author) + sql_result = query.sql() + + assert '"authors"."id" AS "id"' in sql_result.sql + assert '"authors"."name" AS "name"' in sql_result.sql + + +def test_select_all_pydantic_sql_generation(db_dummy: PgDb): + """SelectAllPydantic (Table.all(use_pydantic=True)) drives correct SQL generation.""" + db = db_dummy + + query = db.select(Author.all(use_pydantic=True)).from_(Author) + sql_result = query.sql() + + assert '"authors"."id" AS "id"' in sql_result.sql + assert '"authors"."name" AS "name"' in sql_result.sql + + +# --------------------------------------------------------------------------- +# upgrade_model_nested_fields with plain dataclass +# --------------------------------------------------------------------------- + + +def test_upgrade_nested_fields_plain_dataclass_no_nesting(): + """upgrade_model_nested_fields on a plain dataclass with no nested tables is a no-op.""" + dc = generate_dataclass_model(Author) + upgraded = upgrade_model_nested_fields(dc, use_pydantic=False) + assert hasattr(upgraded, "__dataclass_fields__") + assert "id" in upgraded.__dataclass_fields__ + assert "name" in upgraded.__dataclass_fields__ + + +def test_upgrade_nested_fields_pydantic_no_nesting(): + """upgrade_model_nested_fields on a Pydantic model with no nesting is a no-op.""" + from pydantic import BaseModel + + m = generate_pydantic_model(Author) + upgraded = upgrade_model_nested_fields(m, use_pydantic=True) + assert issubclass(upgraded, BaseModel) + assert "id" in upgraded.model_fields + assert "name" in upgraded.model_fields + + +# --------------------------------------------------------------------------- +# Data round-trip: generate model then load data +# --------------------------------------------------------------------------- + + +def test_dataclass_model_round_trip(): + """Generating a dataclass model then loading data produces correct objects.""" + dc = generate_dataclass_model(Book) + rows = [{"id": 10, "title": "Dune", "author_id": 5}] + results = load_dataclass(dc, rows) + assert results[0].id == 10 + assert results[0].title == "Dune" + assert results[0].author_id == 5 + + +# --------------------------------------------------------------------------- +# Plain class model with nested ManyTable — SQL generation +# --------------------------------------------------------------------------- + + +def test_select_with_nested_many_sql(db_dummy: PgDb): + """A plain class model with a nested ManyTable annotation generates correct SQL.""" + from embar.query.where import Eq + + db = db_dummy + + class AuthorWithBooks: + id: Annotated[int, Author.id] + books: Annotated[list[Book], Book.many()] + + query = db.select(AuthorWithBooks).from_(Author).left_join(Book, Eq(Author.id, Book.author_id)) + # We just verify the sql() call doesn't blow up and has the table reference + sql_result = query.sql() + assert '"authors"."id" AS "id"' in sql_result.sql + + +# --------------------------------------------------------------------------- +# E2E: load data via SQLite with plain dataclass model +# --------------------------------------------------------------------------- + + +def test_e2e_plain_model_sqlite(): + """End-to-end test: insert and select using SQLite and a plain dataclass model (opt-in).""" + import sqlite3 + + from embar.db.sqlite import SqliteDb + + conn = sqlite3.connect(":memory:") + db = SqliteDb(conn) + db.migrate([Author]).run() + + author = Author(id=1, name="Alice") + db.insert(Author).values(author).run() + + # Opt in to plain dataclass path with use_pydantic=False + results = db.select(Author.all(use_pydantic=False)).from_(Author).run() + + assert len(results) == 1 + row = results[0] + assert row.id == 1 + assert row.name == "Alice" + + +def test_e2e_pydantic_model_sqlite(): + """End-to-end test: select using SQLite and SelectAllPydantic.""" + import sqlite3 + + from embar.db.sqlite import SqliteDb + + conn = sqlite3.connect(":memory:") + db = SqliteDb(conn) + db.migrate([Author]).run() + + author = Author(id=2, name="Bob") + db.insert(Author).values(author).run() + + results = db.select(Author.all(use_pydantic=True)).from_(Author).run() + + assert len(results) == 1 + row = results[0] + assert row.id == 2 + assert row.name == "Bob" + + +def test_e2e_plain_class_model_sqlite(): + """End-to-end test: select using a user-defined plain class (not Table, not BaseModel).""" + import sqlite3 + + from embar.db.sqlite import SqliteDb + + conn = sqlite3.connect(":memory:") + db = SqliteDb(conn) + db.migrate([Author]).run() + + author = Author(id=3, name="Carol") + db.insert(Author).values(author).run() + + class AuthorSel: + id: Annotated[int, Author.id] + name: Annotated[str, Author.name] + + results = db.select(AuthorSel).from_(Author).run() + + assert len(results) == 1 + row = results[0] + assert row.id == 3 + assert row.name == "Carol" + + +def test_e2e_nested_many_plain_model_sqlite(): + """End-to-end: nested ManyTable with plain class model loads correctly via SQLite.""" + import sqlite3 + + from embar.db.sqlite import SqliteDb + from embar.query.where import Eq + + conn = sqlite3.connect(":memory:") + db = SqliteDb(conn) + db.migrate([Author, Book]).run() + + author = Author(id=1, name="Alice") + book = Book(id=1, title="Dune", author_id=1) + db.insert(Author).values(author).run() + db.insert(Book).values(book).run() + + class AuthorWithBooks: + id: Annotated[int, Author.id] + name: Annotated[str, Author.name] + books: Annotated[list[Book], Book.many()] + + results = ( + db.select(AuthorWithBooks) + .from_(Author) + .left_join(Book, Eq(Author.id, Book.author_id)) + .group_by(Author.id) + .run() + ) + + assert len(results) == 1 + row = results[0] + assert row.id == 1 + assert row.name == "Alice" + nested = row.books + assert len(nested) == 1 + assert nested[0].title == "Dune" + assert nested[0].id == 1 + + +# --------------------------------------------------------------------------- +# Pydantic validation is real: coercion and error tests +# --------------------------------------------------------------------------- + + +def test_load_results_pydantic_coerces_types(): + """load_results with a Pydantic model coerces compatible raw values (e.g. str→int).""" + from typing import Any, cast + + m = generate_pydantic_model(Author) + # SQLite can return numeric columns as strings in some edge cases; + # Pydantic should coerce "1" → 1 for an int field. + rows = [{"id": "1", "name": "Alice"}] + results = cast(list[Any], load_results(m, rows)) + assert results[0].id == 1 + assert isinstance(results[0].id, int) + + +def test_load_results_pydantic_rejects_invalid_data(): + """load_results with a Pydantic model raises ValidationError for bad data.""" + from pydantic import ValidationError + + m = generate_pydantic_model(Author) + rows = [{"id": "not-an-int", "name": "Alice"}] + with pytest.raises(ValidationError): + load_results(m, rows) + + +def test_load_results_plain_dataclass_does_not_validate(): + """load_results with a plain dataclass passes through bad data without raising.""" + dc = generate_dataclass_model(Author) + rows = [{"id": "not-an-int", "name": "Alice"}] + # Should not raise — no validation + results = load_results(dc, rows) + assert results[0].id == "not-an-int" + + +def test_all_default_uses_pydantic_validation_sqlite(): + """Table.all() default (Pydantic) validates; use_pydantic=False skips validation.""" + import sqlite3 + + from embar.db.sqlite import SqliteDb + + conn = sqlite3.connect(":memory:") + db = SqliteDb(conn) + db.migrate([Author]).run() + + author = Author(id=5, name="Dave") + db.insert(Author).values(author).run() + + # Default path — Pydantic, returns a Pydantic model instance + from pydantic import BaseModel + + results_pydantic = db.select(Author.all()).from_(Author).run() + assert len(results_pydantic) == 1 + assert isinstance(results_pydantic[0], BaseModel) + assert results_pydantic[0].id == 5 + + # Opt-in plain path — plain dataclass, NOT a BaseModel instance + results_plain = db.select(Author.all(use_pydantic=False)).from_(Author).run() + assert len(results_plain) == 1 + assert not isinstance(results_plain[0], BaseModel) + assert results_plain[0].id == 5 + + +def test_returning_default_uses_pydantic_sqlite(): + """returning() default (Pydantic) returns Pydantic-validated instances.""" + import sqlite3 + + from pydantic import BaseModel + + from embar.db.sqlite import SqliteDb + + conn = sqlite3.connect(":memory:") + db = SqliteDb(conn) + db.migrate([Author]).run() + + author = Author(id=10, name="Eve") + + # Default returning() — Pydantic + results = db.insert(Author).values(author).returning().run() + assert len(results) == 1 + assert isinstance(results[0], BaseModel) + assert results[0].id == 10 + + +def test_returning_plain_skips_validation_sqlite(): + """returning(use_pydantic=False) returns plain dataclass instances.""" + import sqlite3 + + from pydantic import BaseModel + + from embar.db.sqlite import SqliteDb + + conn = sqlite3.connect(":memory:") + db = SqliteDb(conn) + db.migrate([Author]).run() + + author = Author(id=11, name="Frank") + + results = db.insert(Author).values(author).returning(use_pydantic=False).run() + assert len(results) == 1 + assert not isinstance(results[0], BaseModel) + assert results[0].id == 11 diff --git a/tests/test_select_clauses.py b/tests/test_select_clauses.py index 10c72b0..07e08e9 100644 --- a/tests/test_select_clauses.py +++ b/tests/test_select_clauses.py @@ -5,7 +5,6 @@ from pydantic import BaseModel from embar.db.pg import PgDb -from embar.model import SelectAll from embar.query.order_by import Asc, Desc from embar.query.where import Gt from embar.sql import Sql @@ -193,7 +192,7 @@ def test_limit_with_offset(db_dummy: PgDb): # fmt: off query = ( - db.select(SelectAll) + db.select(User.all()) .from_(User) .limit(5) .offset(10) diff --git a/uv.lock b/uv.lock index f25b047..d1d4ae1 100644 --- a/uv.lock +++ b/uv.lock @@ -130,6 +130,10 @@ source = { editable = "." } dependencies = [ { name = "psycopg", extra = ["binary"] }, { name = "psycopg-pool" }, +] + +[package.optional-dependencies] +pydantic = [ { name = "pydantic" }, ] @@ -141,6 +145,7 @@ dev = [ { name = "mkdocs-section-index" }, { name = "mkdocstrings", extra = ["python"] }, { name = "poethepoet" }, + { name = "pydantic" }, { name = "pytest" }, { name = "pytest-asyncio" }, { name = "pytest-cov" }, @@ -154,8 +159,9 @@ dev = [ requires-dist = [ { name = "psycopg", extras = ["binary"], specifier = ">=3.2.11" }, { name = "psycopg-pool", specifier = ">=3.3.0" }, - { name = "pydantic", specifier = ">=2.12.4" }, + { name = "pydantic", marker = "extra == 'pydantic'", specifier = "~=2.10" }, ] +provides-extras = ["pydantic"] [package.metadata.requires-dev] dev = [ @@ -165,6 +171,7 @@ dev = [ { name = "mkdocs-section-index", specifier = ">=0.3.10" }, { name = "mkdocstrings", extras = ["python"], specifier = ">=0.30.1" }, { name = "poethepoet", specifier = ">=0.37.0" }, + { name = "pydantic", specifier = "~=2.10" }, { name = "pytest", specifier = ">=8.4.2" }, { name = "pytest-asyncio", specifier = ">=1.2.0" }, { name = "pytest-cov", specifier = ">=7.0.0" },