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 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)
-1
View File
@@ -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]
)
+4 -3
View File
@@ -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