Fix: 500 Error When Searching by Wallet ID in Users. (#3789)
Co-authored-by: Vlad Stan <stan.v.vlad@gmail.com>
This commit is contained in:
co-authored by
Vlad Stan
parent
8c184356ef
commit
75bae67446
+21
-11
@@ -4,7 +4,12 @@ from typing import Any
|
|||||||
from uuid import uuid4
|
from uuid import uuid4
|
||||||
|
|
||||||
from lnbits.core.crud.extensions import get_user_active_extensions_ids
|
from lnbits.core.crud.extensions import get_user_active_extensions_ids
|
||||||
from lnbits.core.crud.wallets import clear_wallet_cache, create_wallet, get_wallets
|
from lnbits.core.crud.wallets import (
|
||||||
|
clear_wallet_cache,
|
||||||
|
create_wallet,
|
||||||
|
get_standalone_wallet,
|
||||||
|
get_wallets,
|
||||||
|
)
|
||||||
from lnbits.core.db import db
|
from lnbits.core.db import db
|
||||||
from lnbits.core.models import UserAcls
|
from lnbits.core.models import UserAcls
|
||||||
from lnbits.db import Connection, Filters, Page
|
from lnbits.db import Connection, Filters, Page
|
||||||
@@ -52,17 +57,22 @@ async def get_accounts(
|
|||||||
) -> Page[AccountOverview]:
|
) -> Page[AccountOverview]:
|
||||||
where_clauses = []
|
where_clauses = []
|
||||||
values: dict[str, Any] = {}
|
values: dict[str, Any] = {}
|
||||||
|
filters = filters or Filters()
|
||||||
|
|
||||||
# Make wallet filter explicit
|
wallet_filter = filters.get_filter_by_field("wallet_id")
|
||||||
wallet_filter = (
|
|
||||||
next((f for f in filters.filters if f.field == "wallet_id"), None)
|
if wallet_filter and wallet_filter.values:
|
||||||
if filters
|
wallet_id_value = next(iter(wallet_filter.values.values()), None)
|
||||||
else None
|
wallet = (
|
||||||
)
|
await get_standalone_wallet(wallet_id_value, deleted=None, conn=conn)
|
||||||
if filters and wallet_filter and wallet_filter.values:
|
if wallet_id_value
|
||||||
where_clauses.append("wallets.id = :wallet_id")
|
else None
|
||||||
values = {**values, "wallet_id": next(iter(wallet_filter.values.values()))}
|
)
|
||||||
filters.filters = [f for f in filters.filters if f.field != "wallet_id"]
|
if not wallet:
|
||||||
|
return Page(data=[], total=0)
|
||||||
|
where_clauses.append("accounts.id = :account_id")
|
||||||
|
values = {**values, "account_id": wallet.user}
|
||||||
|
filters.remove_filter_by_field("wallet_id")
|
||||||
|
|
||||||
return await (conn or db).fetch_page(
|
return await (conn or db).fetch_page(
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -613,6 +613,12 @@ class Filters(BaseModel, Generic[TFilterModel]):
|
|||||||
for page_filter in self.filters:
|
for page_filter in self.filters:
|
||||||
page_filter.table_name = table_name
|
page_filter.table_name = table_name
|
||||||
|
|
||||||
|
def get_filter_by_field(self, field: str) -> Filter[TFilterModel] | None:
|
||||||
|
return next((f for f in self.filters if f.field == field), None)
|
||||||
|
|
||||||
|
def remove_filter_by_field(self, field: str) -> None:
|
||||||
|
self.filters = [f for f in self.filters if f.field != field]
|
||||||
|
|
||||||
|
|
||||||
class DbJsonEncoder(json.JSONEncoder):
|
class DbJsonEncoder(json.JSONEncoder):
|
||||||
def default(self, o):
|
def default(self, o):
|
||||||
|
|||||||
@@ -2,10 +2,16 @@ from uuid import uuid4
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from lnbits.core.crud.users import get_user_from_account
|
from lnbits.core.crud.users import (
|
||||||
from lnbits.core.crud.wallets import delete_wallet, get_wallets
|
create_account,
|
||||||
from lnbits.core.models.users import Account
|
delete_account,
|
||||||
from lnbits.core.services.users import create_user_account
|
get_accounts,
|
||||||
|
get_user_from_account,
|
||||||
|
)
|
||||||
|
from lnbits.core.crud.wallets import delete_wallet, force_delete_wallet, get_wallets
|
||||||
|
from lnbits.core.models.users import Account, AccountFilters
|
||||||
|
from lnbits.core.services.users import create_user_account, create_user_account_no_ckeck
|
||||||
|
from lnbits.db import Filter, Filters, Operator
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.anyio
|
@pytest.mark.anyio
|
||||||
@@ -36,3 +42,151 @@ async def test_get_user_from_account_is_wallet_created():
|
|||||||
assert (
|
assert (
|
||||||
len(user.wallets) == 1
|
len(user.wallets) == 1
|
||||||
), "A new wallet should be created for the user if none exist after deletion"
|
), "A new wallet should be created for the user if none exist after deletion"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_get_accounts_success_flow():
|
||||||
|
# Create a new account
|
||||||
|
username = f"user_{uuid4().hex[:8]}"
|
||||||
|
account = Account(
|
||||||
|
id=uuid4().hex,
|
||||||
|
username=username,
|
||||||
|
email=f"{username}@lnbits.com",
|
||||||
|
)
|
||||||
|
await create_account(account)
|
||||||
|
# Should return the created account
|
||||||
|
filters = Filters[AccountFilters](filters=[], model=AccountFilters)
|
||||||
|
filters.sortby = "created_at"
|
||||||
|
filters.direction = "desc"
|
||||||
|
page = await get_accounts(filters=filters)
|
||||||
|
assert page.total >= 1
|
||||||
|
found = any(a.username == username for a in page.data)
|
||||||
|
assert found
|
||||||
|
await delete_account(account.id)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_get_accounts_with_wallet_id_filter():
|
||||||
|
# Create account and wallet
|
||||||
|
username = f"user_{uuid4().hex[:8]}"
|
||||||
|
account = Account(
|
||||||
|
id=uuid4().hex,
|
||||||
|
username=username,
|
||||||
|
email=f"{username}@lnbits.com",
|
||||||
|
)
|
||||||
|
await create_user_account_no_ckeck(account)
|
||||||
|
|
||||||
|
wallets = await get_wallets(account.id, deleted=False)
|
||||||
|
assert wallets
|
||||||
|
wallet = wallets[0]
|
||||||
|
# Filter by wallet_id
|
||||||
|
filters = Filters[AccountFilters](
|
||||||
|
filters=[
|
||||||
|
Filter(
|
||||||
|
field="wallet_id",
|
||||||
|
op=Operator.EQ,
|
||||||
|
model=AccountFilters,
|
||||||
|
values={"wallet_id__0": wallet.id},
|
||||||
|
)
|
||||||
|
],
|
||||||
|
model=AccountFilters,
|
||||||
|
)
|
||||||
|
page = await get_accounts(filters=filters)
|
||||||
|
assert page.total == 1
|
||||||
|
assert page.data[0].id == account.id
|
||||||
|
await delete_account(account.id)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_get_accounts_wallet_id_not_found():
|
||||||
|
|
||||||
|
filters = Filters[AccountFilters](
|
||||||
|
filters=[
|
||||||
|
Filter(
|
||||||
|
field="wallet_id",
|
||||||
|
op=Operator.EQ,
|
||||||
|
model=AccountFilters,
|
||||||
|
values={"wallet_id__0": uuid4().hex},
|
||||||
|
)
|
||||||
|
],
|
||||||
|
model=AccountFilters,
|
||||||
|
)
|
||||||
|
page = await get_accounts(filters=filters)
|
||||||
|
assert page.total == 0
|
||||||
|
assert page.data == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_get_accounts_empty_filters():
|
||||||
|
# Should not raise, should return a Page
|
||||||
|
page = await get_accounts()
|
||||||
|
assert hasattr(page, "data")
|
||||||
|
assert hasattr(page, "total")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_get_accounts_with_deleted_wallet():
|
||||||
|
# Create account and wallet, then delete wallet
|
||||||
|
username = f"user_{uuid4().hex[:8]}"
|
||||||
|
account = Account(
|
||||||
|
id=uuid4().hex,
|
||||||
|
username=username,
|
||||||
|
email=f"{username}@lnbits.com",
|
||||||
|
)
|
||||||
|
await create_user_account_no_ckeck(account)
|
||||||
|
|
||||||
|
wallets = await get_wallets(account.id, deleted=False)
|
||||||
|
assert wallets
|
||||||
|
wallet = wallets[0]
|
||||||
|
await delete_wallet(user_id=account.id, wallet_id=wallet.id)
|
||||||
|
|
||||||
|
filters = Filters[AccountFilters](
|
||||||
|
filters=[
|
||||||
|
Filter(
|
||||||
|
field="wallet_id",
|
||||||
|
op=Operator.EQ,
|
||||||
|
model=AccountFilters,
|
||||||
|
values={"wallet_id__0": wallet.id},
|
||||||
|
)
|
||||||
|
],
|
||||||
|
model=AccountFilters,
|
||||||
|
)
|
||||||
|
page = await get_accounts(filters=filters)
|
||||||
|
assert page.total == 1
|
||||||
|
assert page.data[0].id == account.id
|
||||||
|
|
||||||
|
await force_delete_wallet(wallet_id=wallet.id)
|
||||||
|
|
||||||
|
filters = Filters[AccountFilters](
|
||||||
|
filters=[
|
||||||
|
Filter(
|
||||||
|
field="wallet_id",
|
||||||
|
op=Operator.EQ,
|
||||||
|
model=AccountFilters,
|
||||||
|
values={"wallet_id__0": wallet.id},
|
||||||
|
)
|
||||||
|
],
|
||||||
|
model=AccountFilters,
|
||||||
|
)
|
||||||
|
page = await get_accounts(filters=filters)
|
||||||
|
assert page.total == 0
|
||||||
|
assert page.data == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_get_accounts_group_by_and_pagination():
|
||||||
|
# Create multiple accounts
|
||||||
|
accounts = []
|
||||||
|
for _ in range(3):
|
||||||
|
username = f"user_{uuid4().hex[:8]}"
|
||||||
|
account = Account(
|
||||||
|
id=uuid4().hex,
|
||||||
|
username=username,
|
||||||
|
email=f"{username}@lnbits.com",
|
||||||
|
)
|
||||||
|
await create_user_account_no_ckeck(account)
|
||||||
|
accounts.append(account)
|
||||||
|
filters = Filters[AccountFilters](model=AccountFilters, limit=2, offset=0)
|
||||||
|
page = await get_accounts(filters=filters)
|
||||||
|
assert page.total >= 3
|
||||||
|
assert len(page.data) <= 2
|
||||||
|
|||||||
Reference in New Issue
Block a user