fixup!
This commit is contained in:
+14
-12
@@ -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)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user