diff --git a/lnbits/db.py b/lnbits/db.py index a93e4a4c5..0b36cdf4a 100644 --- a/lnbits/db.py +++ b/lnbits/db.py @@ -9,6 +9,7 @@ from contextlib import asynccontextmanager from enum import Enum from typing import Any, Generic, Literal, Optional, TypeVar +import shortuuid from loguru import logger from pydantic import BaseModel, ValidationError, root_validator from sqlalchemy import event @@ -263,7 +264,7 @@ class Database(Compat): dbapi_connection.run_async( lambda connection: connection.set_type_codec( "TIMESTAMP", - encoder=datetime.datetime.timestamp, + encoder=datetime.datetime, decoder=_parse_timestamp, schema="pg_catalog", ) @@ -421,7 +422,7 @@ class Filter(BaseModel, Generic[TFilterModel]): validated, errors = compare_field.validate(raw_value, {}, loc="none") if errors: raise ValidationError(errors=[errors], model=model) - values[field] = validated + values[f"{field}__{shortuuid.uuid()}"] = validated else: raise ValueError("Unknown filter field") @@ -429,16 +430,17 @@ class Filter(BaseModel, Generic[TFilterModel]): @property def statement(self): - if self.op in (Operator.INCLUDE, Operator.EXCLUDE) and self.values: - placeholders = [] - for key in self.values.keys(): - if self.model and self.model.__fields__[key].type_ == datetime.datetime: - placeholders.append(compat_timestamp_placeholder(key)) - else: - placeholders.append(f":{key}") - stmt = [f"{self.field} {self.op.as_sql} ({', '.join(placeholders)})"] - else: - stmt = [f"{self.field} {self.op.as_sql} :{self.field}"] + stmt = [] + for key in self.values.keys() if self.values else []: + clean_key = key.split("__")[0] + if ( + self.model + and self.model.__fields__[clean_key].type_ == datetime.datetime + ): + placeholder = compat_timestamp_placeholder(key) + else: + placeholder = f":{key}" + stmt.append(f"{clean_key} {self.op.as_sql} {placeholder}") return " OR ".join(stmt) diff --git a/tests/api/test_api.py b/tests/api/test_api.py index fa8a011d7..52b3b4506 100644 --- a/tests/api/test_api.py +++ b/tests/api/test_api.py @@ -367,7 +367,6 @@ async def test_get_payments_history(client, adminkey_headers_from, fake_payments assert response.status_code == 200 data = response.json() assert len(data) == 1 - print(data) assert data[0]["income"] == sum( [int(payment.amount * 1000) for payment in fake_data if not payment.out] ) diff --git a/tests/conftest.py b/tests/conftest.py index 070d03008..6e92e7748 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -22,7 +22,7 @@ from lnbits.core.crud import ( from lnbits.core.models import CreateInvoice, PaymentState from lnbits.core.services import update_wallet_balance from lnbits.core.views.payment_api import api_payments_create_invoice -from lnbits.db import Database +from lnbits.db import DB_TYPE, SQLITE, Database from lnbits.settings import settings from tests.helpers import ( get_random_invoice_data, @@ -182,7 +182,8 @@ async def invoice(to_wallet): async def fake_payments(client, adminkey_headers_from): # Because sqlite only stores timestamps with milliseconds # we have to wait a second to ensure a different timestamp than previous invoices - await asyncio.sleep(1) + if DB_TYPE == SQLITE: + await asyncio.sleep(1) ts = int(time()) fake_data = [ @@ -200,5 +201,5 @@ async def fake_payments(client, adminkey_headers_from): assert data["checking_id"] await update_payment_status(data["checking_id"], status=PaymentState.SUCCESS) - params = {"time[ge]": ts, "time[le]": int(time())} + params = {"time[ge]": ts, "time[le]": int(time()) + 1} return fake_data, params