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 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.models import UserAcls
|
||||
from lnbits.db import Connection, Filters, Page
|
||||
@@ -52,17 +57,22 @@ async def get_accounts(
|
||||
) -> Page[AccountOverview]:
|
||||
where_clauses = []
|
||||
values: dict[str, Any] = {}
|
||||
filters = filters or Filters()
|
||||
|
||||
# Make wallet filter explicit
|
||||
wallet_filter = (
|
||||
next((f for f in filters.filters if f.field == "wallet_id"), None)
|
||||
if filters
|
||||
else None
|
||||
)
|
||||
if filters and wallet_filter and wallet_filter.values:
|
||||
where_clauses.append("wallets.id = :wallet_id")
|
||||
values = {**values, "wallet_id": next(iter(wallet_filter.values.values()))}
|
||||
filters.filters = [f for f in filters.filters if f.field != "wallet_id"]
|
||||
wallet_filter = filters.get_filter_by_field("wallet_id")
|
||||
|
||||
if wallet_filter and wallet_filter.values:
|
||||
wallet_id_value = next(iter(wallet_filter.values.values()), None)
|
||||
wallet = (
|
||||
await get_standalone_wallet(wallet_id_value, deleted=None, conn=conn)
|
||||
if wallet_id_value
|
||||
else None
|
||||
)
|
||||
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(
|
||||
"""
|
||||
|
||||
@@ -613,6 +613,12 @@ class Filters(BaseModel, Generic[TFilterModel]):
|
||||
for page_filter in self.filters:
|
||||
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):
|
||||
def default(self, o):
|
||||
|
||||
@@ -2,10 +2,16 @@ from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
|
||||
from lnbits.core.crud.users import get_user_from_account
|
||||
from lnbits.core.crud.wallets import delete_wallet, get_wallets
|
||||
from lnbits.core.models.users import Account
|
||||
from lnbits.core.services.users import create_user_account
|
||||
from lnbits.core.crud.users import (
|
||||
create_account,
|
||||
delete_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
|
||||
@@ -36,3 +42,151 @@ async def test_get_user_from_account_is_wallet_created():
|
||||
assert (
|
||||
len(user.wallets) == 1
|
||||
), "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