manual testing works

This commit is contained in:
dni ⚡
2026-03-25 11:14:53 +01:00
parent b0b8ea95f2
commit 034ba44241
3 changed files with 87 additions and 44 deletions
+44 -34
View File
@@ -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
+10 -8
View File
@@ -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)
+33 -2
View File
@@ -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}