diff --git a/lnbits/db.py b/lnbits/db.py index 4bd136c05..87503994f 100644 --- a/lnbits/db.py +++ b/lnbits/db.py @@ -9,11 +9,19 @@ from contextlib import asynccontextmanager from datetime import datetime, timezone from enum import Enum from types import UnionType -from typing import Any, Generic, Literal, TypeVar, Union, get_args, get_origin +from typing import ( + Any, + ClassVar, + Generic, + Literal, + TypeVar, + Union, + get_args, + get_origin, +) from loguru import logger -from pydantic import BaseModel as BaseModelV2 -from pydantic.v1 import BaseModel, ValidationError, root_validator +from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, model_validator from sqlalchemy.ext.asyncio import AsyncConnection, AsyncEngine, create_async_engine from sqlalchemy.sql import text @@ -22,7 +30,7 @@ from lnbits.settings import settings POSTGRES = "POSTGRES" COCKROACH = "COCKROACH" SQLITE = "SQLITE" -PYDANTIC_MODEL_TYPES = (BaseModel, BaseModelV2) +PYDANTIC_MODEL_TYPES = (BaseModel,) DateTrunc = Literal["hour", "day", "month"] sqlite_formats = { @@ -38,10 +46,6 @@ def _is_pydantic_model_class(model: Any) -> bool: ) -def _is_v2_pydantic_model_class(model: Any) -> bool: - return isinstance(model, type) and issubclass(model, BaseModelV2) - - def _issubclass(candidate: Any, parent: Any) -> bool: return isinstance(candidate, type) and issubclass(candidate, parent) @@ -88,21 +92,17 @@ def _field_inner_type(field: Any) -> Any: def _field_extra(field: Any) -> dict[str, Any]: field_info = getattr(field, "field_info", None) if field_info is not None: - return field_info.extra + return getattr(field_info, "extra", {}) return getattr(field, "json_schema_extra", None) or {} -def _model_dump(model: BaseModel | BaseModelV2) -> dict[str, Any]: - if isinstance(model, BaseModelV2): - return model.model_dump() - return model.dict() +def _model_dump(model: BaseModel) -> dict[str, Any]: + return model.model_dump() def _validate_model(model: type[Any], values: dict[str, Any]) -> Any: - if _is_v2_pydantic_model_class(model): - return model.model_validate(values) - return model.parse_obj(values) + return model.model_validate(values) if settings.lnbits_database_url: @@ -524,12 +524,12 @@ class Operator(Enum): class FilterModel(BaseModel): - __search_fields__: list[str] = [] - __sort_fields__: list[str] | None = None + __search_fields__: ClassVar[list[str]] = [] + __sort_fields__: ClassVar[list[str] | None] = None T = TypeVar("T") -TModel = TypeVar("TModel", bound=BaseModel | BaseModelV2) +TModel = TypeVar("TModel", bound=BaseModel) TFilterModel = TypeVar("TFilterModel", bound=FilterModel) @@ -539,10 +539,12 @@ class Page(BaseModel, Generic[T]): class Filter(BaseModel, Generic[TFilterModel]): + model_config = ConfigDict(arbitrary_types_allowed=True) + table_name: str | None = None field: str op: Operator = Operator.EQ - model: type[TFilterModel] | None + model: type[TFilterModel] | None = Field(default=None, exclude=True) values: dict | None = None @classmethod @@ -564,16 +566,17 @@ class Filter(BaseModel, Generic[TFilterModel]): field = key op = Operator("eq") - if field in model.__fields__: - compare_field = model.__fields__[field] + model_fields = _model_fields(model) + if field in model_fields: + compare_field = model_fields[field] values: dict = {} if op in {Operator.EVERY, Operator.ANY, Operator.INCLUDE, Operator.EXCLUDE}: raw_values = [v for rv in raw_values for v in rv.split(",")] for index, raw_value in enumerate(raw_values): - validated, errors = compare_field.validate(raw_value, {}, loc="none") - if errors: - raise ValidationError(errors=[errors], model=model) + validated = TypeAdapter( + _field_annotation(compare_field) + ).validate_python(raw_value) values[f"{field}__{index}"] = validated else: raise ValueError("Unknown filter field") @@ -585,7 +588,10 @@ class Filter(BaseModel, Generic[TFilterModel]): prefix = f"{self.table_name}." if self.table_name else "" stmt = [] for key in self.values.keys() if self.values else []: - if self.model and self.model.__fields__[self.field].type_ == datetime: + if ( + self.model + and _field_inner_type(_model_fields(self.model)[self.field]) == datetime + ): placeholder = compat_timestamp_placeholder(key) stmt.append(f"{prefix}{self.field} {self.op.as_sql} {placeholder}") if self.op in {Operator.INCLUDE, Operator.EXCLUDE}: @@ -611,7 +617,9 @@ class Filters(BaseModel, Generic[TFilterModel]): the values can be validated. Otherwise, make sure to validate the inputs manually. """ - filters: list[Filter[TFilterModel]] = [] + model_config = ConfigDict(arbitrary_types_allowed=True) + + filters: list[Filter[TFilterModel]] = Field(default_factory=list) search: str | None = None offset: int | None = None @@ -619,18 +627,20 @@ class Filters(BaseModel, Generic[TFilterModel]): sortby: str | None = None direction: Literal["asc", "desc"] | None = None - model: type[TFilterModel] | None = None + model: type[TFilterModel] | None = Field(default=None, exclude=True) - table_name: str | None = None + table_name: str | None = Field(default=None, exclude=True) - @root_validator(pre=True) - def validate_sortby(cls, values): + @model_validator(mode="before") + @classmethod + def validate_sortby(cls, values: Any): + if not isinstance(values, dict): + return values sortby = values.get("sortby") model = values.get("model") if sortby and model: - model = values["model"] # if no sort fields are specified explicitly all fields are allowed - allowed = model.__sort_fields__ or model.__fields__ + allowed = model.__sort_fields__ or _model_fields(model).keys() if sortby not in allowed: raise ValueError("Invalid sort field") return values @@ -737,7 +747,7 @@ def update_query( return f"UPDATE {table_name} SET {query} {where}" # noqa: S608 -def model_to_dict(model: BaseModel | BaseModelV2) -> dict: +def model_to_dict(model: BaseModel) -> dict: """ Convert a Pydantic model to a dictionary with JSON-encoded nested models private fields starting with _ are ignored diff --git a/lnbits/helpers.py b/lnbits/helpers.py index 7d7a9f463..a519cd5da 100644 --- a/lnbits/helpers.py +++ b/lnbits/helpers.py @@ -12,7 +12,6 @@ import shortuuid from fastapi.routing import APIRoute from loguru import logger from packaging import version -from pydantic.v1.schema import field_schema from starlette.templating import Jinja2Templates from lnbits.settings import settings @@ -127,23 +126,26 @@ def generate_filter_params_openapi(model: type[FilterModel], keep_optional=False :param keep_optional: If false, all parameters will be optional, otherwise inferred from model """ - fields = list(model.__fields__.values()) + schema = model.model_json_schema() + properties = schema.get("properties", {}) + required = set(schema.get("required", [])) params = [] - for field in fields: - schema, _, _ = field_schema(field, model_name_map={}) + for field_name, field in model.model_fields.items(): + field_key = field.alias or field_name + field_schema = properties.get(field_key, {}) description = "Supports Filtering" if ( hasattr(model, "__search_fields__") - and field.name in model.__search_fields__ + and field_name in model.__search_fields__ ): description += ". Supports Search" parameter = { - "name": field.alias, + "name": field_key, "in": "query", - "required": field.required if keep_optional else False, - "schema": schema, + "required": field_key in required if keep_optional else False, + "schema": field_schema, "description": description, } params.append(parameter) diff --git a/tests/unit/test_core_models.py b/tests/unit/test_core_models.py index 6a539b9bb..1a5ee3d3b 100644 --- a/tests/unit/test_core_models.py +++ b/tests/unit/test_core_models.py @@ -1,8 +1,9 @@ import pytest -from pydantic import ValidationError +from pydantic import TypeAdapter, ValidationError from lnbits.core.models.extensions import InstallableExtension from lnbits.core.models.lnurl import StoredPayLink +from lnbits.core.models.payments import PaymentFilters from lnbits.core.models.users import UserLabel from lnbits.core.models.wallets import ( Wallet, @@ -10,7 +11,8 @@ from lnbits.core.models.wallets import ( WalletSharePermission, WalletShareStatus, ) -from lnbits.db import dict_to_model +from lnbits.db import Filter, Page, dict_to_model +from lnbits.nodes.base import NodePayment def test_user_label_uses_pydantic_v2_pattern_validation(): @@ -78,3 +80,32 @@ def test_db_dict_to_model_parses_optional_nested_pydantic_v2_models(): assert ext.meta is not None assert ext.meta.installed_release is not None assert ext.meta.installed_release.source_repo == "lnbits/example" + + +def test_page_generic_validates_through_pydantic_v2_type_adapter(): + page = Page[NodePayment]( + data=[ + NodePayment( + pending=False, + amount=1, + time=1, + preimage="preimage", + payment_hash="payment-hash", + ) + ], + total=1, + ) + + validated = TypeAdapter(Page[NodePayment]).validate_python( + page, from_attributes=True + ) + + assert validated.total == 1 + assert validated.data[0].payment_hash == "payment-hash" + + +def test_filter_parse_query_uses_pydantic_v2_field_validation(): + parsed = Filter.parse_query("amount[eq]", ["42"], PaymentFilters) + + assert parsed.field == "amount" + assert parsed.values == {"amount__0": 42}