fixup!
This commit is contained in:
+6
-6
@@ -155,7 +155,7 @@ class Connection(Compat):
|
|||||||
if not row:
|
if not row:
|
||||||
return []
|
return []
|
||||||
if model:
|
if model:
|
||||||
return [_dict_to_model(r, model) for r in row]
|
return [dict_to_model(r, model) for r in row]
|
||||||
return row
|
return row
|
||||||
|
|
||||||
async def fetchone(
|
async def fetchone(
|
||||||
@@ -166,18 +166,18 @@ class Connection(Compat):
|
|||||||
row = result.mappings().first()
|
row = result.mappings().first()
|
||||||
result.close()
|
result.close()
|
||||||
if model and row:
|
if model and row:
|
||||||
return _dict_to_model(row, model)
|
return dict_to_model(row, model)
|
||||||
return row
|
return row
|
||||||
|
|
||||||
async def update(self, table_name: str, model: BaseModel, where: str = "id = :id"):
|
async def update(self, table_name: str, model: BaseModel, where: str = "id = :id"):
|
||||||
await self.conn.execute(
|
await self.conn.execute(
|
||||||
text(update_query(table_name, model, where)), _model_to_dict(model)
|
text(update_query(table_name, model, where)), model_to_dict(model)
|
||||||
)
|
)
|
||||||
await self.conn.commit()
|
await self.conn.commit()
|
||||||
|
|
||||||
async def insert(self, table_name: str, model: BaseModel):
|
async def insert(self, table_name: str, model: BaseModel):
|
||||||
await self.conn.execute(
|
await self.conn.execute(
|
||||||
text(insert_query(table_name, model)), _model_to_dict(model)
|
text(insert_query(table_name, model)), model_to_dict(model)
|
||||||
)
|
)
|
||||||
await self.conn.commit()
|
await self.conn.commit()
|
||||||
|
|
||||||
@@ -584,7 +584,7 @@ def update_query(
|
|||||||
return f"UPDATE {table_name} SET {query} {where}"
|
return f"UPDATE {table_name} SET {query} {where}"
|
||||||
|
|
||||||
|
|
||||||
def _model_to_dict(model: BaseModel) -> dict:
|
def model_to_dict(model: BaseModel) -> dict:
|
||||||
_dict = model.dict()
|
_dict = model.dict()
|
||||||
for key, value in _dict.items():
|
for key, value in _dict.items():
|
||||||
if key.startswith("_"):
|
if key.startswith("_"):
|
||||||
@@ -595,7 +595,7 @@ def _model_to_dict(model: BaseModel) -> dict:
|
|||||||
return _dict
|
return _dict
|
||||||
|
|
||||||
|
|
||||||
def _dict_to_model(_dict: dict, model: TModel) -> TModel:
|
def dict_to_model(_dict: dict, model: TModel) -> TModel:
|
||||||
for key, value in _dict.items():
|
for key, value in _dict.items():
|
||||||
type_ = model.__fields__[key].type_
|
type_ = model.__fields__[key].type_
|
||||||
if type_ is BaseModel:
|
if type_ is BaseModel:
|
||||||
|
|||||||
@@ -2,6 +2,8 @@ import random
|
|||||||
import string
|
import string
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
|
from pydantic import BaseModel
|
||||||
|
|
||||||
from lnbits.db import FromRowModel
|
from lnbits.db import FromRowModel
|
||||||
from lnbits.wallets import get_funding_source, set_funding_source
|
from lnbits.wallets import get_funding_source, set_funding_source
|
||||||
|
|
||||||
@@ -16,6 +18,19 @@ class DbTestModel(FromRowModel):
|
|||||||
value: Optional[str] = None
|
value: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
|
class DbTestModelInner(BaseModel):
|
||||||
|
id: int
|
||||||
|
label: str
|
||||||
|
description: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
|
class DbTestModel2(BaseModel):
|
||||||
|
id: int
|
||||||
|
name: str
|
||||||
|
value: Optional[str] = None
|
||||||
|
child: DbTestModelInner
|
||||||
|
|
||||||
|
|
||||||
def get_random_string(iterations: int = 10):
|
def get_random_string(iterations: int = 10):
|
||||||
return "".join(
|
return "".join(
|
||||||
random.SystemRandom().choice(string.ascii_uppercase + string.digits)
|
random.SystemRandom().choice(string.ascii_uppercase + string.digits)
|
||||||
|
|||||||
@@ -1,28 +1,52 @@
|
|||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from lnbits.db import (
|
from lnbits.db import (
|
||||||
|
dict_to_model,
|
||||||
insert_query,
|
insert_query,
|
||||||
|
model_to_dict,
|
||||||
update_query,
|
update_query,
|
||||||
)
|
)
|
||||||
from tests.helpers import DbTestModel
|
from tests.helpers import DbTestModel2, DbTestModelInner
|
||||||
|
|
||||||
test = DbTestModel(id=1, name="test", value="yes")
|
test_data = DbTestModel2(
|
||||||
|
id=1,
|
||||||
|
name="test",
|
||||||
|
value="myvalue",
|
||||||
|
child=DbTestModelInner(id=2, label="mylabel", description="mydesc"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_helpers_insert_query():
|
async def test_helpers_insert_query():
|
||||||
q = insert_query("test_helpers_query", test)
|
q = insert_query("test_helpers_query", test_data)
|
||||||
assert (
|
assert (
|
||||||
q == "INSERT INTO test_helpers_query (id, name, value) "
|
q == "INSERT INTO test_helpers_query (id, name, value, child) "
|
||||||
"VALUES (:id, :name, :value)"
|
"VALUES (:id, :name, :value, :child)"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_helpers_update_query():
|
async def test_helpers_update_query():
|
||||||
q = update_query("test_helpers_query", test)
|
q = update_query("test_helpers_query", test_data)
|
||||||
assert (
|
assert (
|
||||||
q == "UPDATE test_helpers_query "
|
q == "UPDATE test_helpers_query "
|
||||||
"SET id = :id, name = :name, value = :value "
|
"SET id = :id, name = :name, value = :value, child = :child "
|
||||||
"WHERE id = :id"
|
"WHERE id = :id"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_helpers_model_to_dict():
|
||||||
|
d = model_to_dict(test_data)
|
||||||
|
assert d == {
|
||||||
|
"id": 1,
|
||||||
|
"name": "test",
|
||||||
|
"value": "myvalue",
|
||||||
|
"child": {"id": 2, "label": "mylabel", "description": "mydesc"},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_helpers_dict_to_model():
|
||||||
|
m = dict_to_model(model_to_dict(test_data), DbTestModel2)
|
||||||
|
assert m == test_data
|
||||||
|
|||||||
Reference in New Issue
Block a user