This commit is contained in:
dni ⚡
2024-09-05 12:38:59 +02:00
parent 9f60b14745
commit cd7fcd1c71
3 changed files with 18 additions and 16 deletions
+14 -12
View File
@@ -9,6 +9,7 @@ from contextlib import asynccontextmanager
from enum import Enum from enum import Enum
from typing import Any, Generic, Literal, Optional, TypeVar from typing import Any, Generic, Literal, Optional, TypeVar
import shortuuid
from loguru import logger from loguru import logger
from pydantic import BaseModel, ValidationError, root_validator from pydantic import BaseModel, ValidationError, root_validator
from sqlalchemy import event from sqlalchemy import event
@@ -263,7 +264,7 @@ class Database(Compat):
dbapi_connection.run_async( dbapi_connection.run_async(
lambda connection: connection.set_type_codec( lambda connection: connection.set_type_codec(
"TIMESTAMP", "TIMESTAMP",
encoder=datetime.datetime.timestamp, encoder=datetime.datetime,
decoder=_parse_timestamp, decoder=_parse_timestamp,
schema="pg_catalog", schema="pg_catalog",
) )
@@ -421,7 +422,7 @@ class Filter(BaseModel, Generic[TFilterModel]):
validated, errors = compare_field.validate(raw_value, {}, loc="none") validated, errors = compare_field.validate(raw_value, {}, loc="none")
if errors: if errors:
raise ValidationError(errors=[errors], model=model) raise ValidationError(errors=[errors], model=model)
values[field] = validated values[f"{field}__{shortuuid.uuid()}"] = validated
else: else:
raise ValueError("Unknown filter field") raise ValueError("Unknown filter field")
@@ -429,16 +430,17 @@ class Filter(BaseModel, Generic[TFilterModel]):
@property @property
def statement(self): def statement(self):
if self.op in (Operator.INCLUDE, Operator.EXCLUDE) and self.values: stmt = []
placeholders = [] for key in self.values.keys() if self.values else []:
for key in self.values.keys(): clean_key = key.split("__")[0]
if self.model and self.model.__fields__[key].type_ == datetime.datetime: if (
placeholders.append(compat_timestamp_placeholder(key)) self.model
else: and self.model.__fields__[clean_key].type_ == datetime.datetime
placeholders.append(f":{key}") ):
stmt = [f"{self.field} {self.op.as_sql} ({', '.join(placeholders)})"] placeholder = compat_timestamp_placeholder(key)
else: else:
stmt = [f"{self.field} {self.op.as_sql} :{self.field}"] placeholder = f":{key}"
stmt.append(f"{clean_key} {self.op.as_sql} {placeholder}")
return " OR ".join(stmt) return " OR ".join(stmt)
-1
View File
@@ -367,7 +367,6 @@ async def test_get_payments_history(client, adminkey_headers_from, fake_payments
assert response.status_code == 200 assert response.status_code == 200
data = response.json() data = response.json()
assert len(data) == 1 assert len(data) == 1
print(data)
assert data[0]["income"] == sum( assert data[0]["income"] == sum(
[int(payment.amount * 1000) for payment in fake_data if not payment.out] [int(payment.amount * 1000) for payment in fake_data if not payment.out]
) )
+4 -3
View File
@@ -22,7 +22,7 @@ from lnbits.core.crud import (
from lnbits.core.models import CreateInvoice, PaymentState from lnbits.core.models import CreateInvoice, PaymentState
from lnbits.core.services import update_wallet_balance from lnbits.core.services import update_wallet_balance
from lnbits.core.views.payment_api import api_payments_create_invoice 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 lnbits.settings import settings
from tests.helpers import ( from tests.helpers import (
get_random_invoice_data, get_random_invoice_data,
@@ -182,7 +182,8 @@ async def invoice(to_wallet):
async def fake_payments(client, adminkey_headers_from): async def fake_payments(client, adminkey_headers_from):
# Because sqlite only stores timestamps with milliseconds # Because sqlite only stores timestamps with milliseconds
# we have to wait a second to ensure a different timestamp than previous invoices # 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()) ts = int(time())
fake_data = [ fake_data = [
@@ -200,5 +201,5 @@ async def fake_payments(client, adminkey_headers_from):
assert data["checking_id"] assert data["checking_id"]
await update_payment_status(data["checking_id"], status=PaymentState.SUCCESS) 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 return fake_data, params