diff --git a/lnbits/db.py b/lnbits/db.py index c9f28f10f..e71bc0b83 100644 --- a/lnbits/db.py +++ b/lnbits/db.py @@ -12,9 +12,8 @@ from typing import Any, Generic, Literal, Optional, TypeVar from loguru import logger from pydantic import BaseModel, ValidationError, root_validator -from sqlalchemy import create_engine -from sqlalchemy_aio.base import AsyncConnection -from sqlalchemy_aio.strategy import ASYNCIO_STRATEGY +from sqlalchemy.ext.asyncio import AsyncConnection, AsyncEngine, create_async_engine +from sqlalchemy.sql import text from lnbits.settings import settings @@ -133,9 +132,8 @@ class Compat: class Connection(Compat): - def __init__(self, conn: AsyncConnection, txn, typ, name, schema): + def __init__(self, conn: AsyncConnection, typ, name, schema): self.conn = conn - self.txn = txn self.type = typ self.name = name self.schema = schema @@ -168,16 +166,16 @@ class Connection(Compat): async def fetchall(self, query: str, values: tuple = ()) -> list: result = await self.conn.execute( - self.rewrite_query(query), self.rewrite_values(values) + text(self.rewrite_query(query)), self.rewrite_values(values) ) - return await result.fetchall() + return result.fetchall() async def fetchone(self, query: str, values: tuple = ()): result = await self.conn.execute( - self.rewrite_query(query), self.rewrite_values(values) + text(self.rewrite_query(query)), self.rewrite_values(values) ) - row = await result.fetchone() - await result.close() + row = result.fetchone() + result.close() return row async def fetch_page( @@ -239,7 +237,7 @@ class Connection(Compat): async def execute(self, query: str, values: tuple = ()): return await self.conn.execute( - self.rewrite_query(query), self.rewrite_values(values) + text(self.rewrite_query(query)), self.rewrite_values(values) ) @@ -253,7 +251,7 @@ class Database(Compat): self.path = os.path.join( settings.lnbits_data_folder, f"{self.name}.sqlite3" ) - database_uri = f"sqlite:///{self.path}" + database_uri = f"sqlite+aiosqlite:///{self.path}" else: database_uri = settings.lnbits_database_url @@ -262,8 +260,8 @@ class Database(Compat): else: self.schema = None - self.engine = create_engine( - database_uri, strategy=ASYNCIO_STRATEGY, echo=settings.debug_database + self.engine: AsyncEngine = create_async_engine( + database_uri, echo=settings.debug_database ) self.lock = asyncio.Lock() @@ -273,34 +271,34 @@ class Database(Compat): async def connect(self): await self.lock.acquire() try: - async with self.engine.connect() as conn: # type: ignore - async with conn.begin() as txn: - wconn = Connection(conn, txn, self.type, self.name, self.schema) + async with self.engine.connect() as conn: + if not conn: + raise Exception("Could not connect to the database") - if self.schema: - if self.type in {POSTGRES, COCKROACH}: - await wconn.execute( - f"CREATE SCHEMA IF NOT EXISTS {self.schema}" - ) - elif self.type == SQLITE: - await wconn.execute( - f"ATTACH '{self.path}' AS {self.schema}" - ) + wconn = Connection(conn, self.type, self.name, self.schema) - yield wconn + if self.schema: + if self.type in {POSTGRES, COCKROACH}: + await wconn.execute( + f"CREATE SCHEMA IF NOT EXISTS {self.schema}" + ) + elif self.type == SQLITE: + await wconn.execute(f"ATTACH '{self.path}' AS {self.schema}") + + yield wconn finally: self.lock.release() async def fetchall(self, query: str, values: tuple = ()) -> list: async with self.connect() as conn: result = await conn.execute(query, values) - return await result.fetchall() + return result.fetchall() async def fetchone(self, query: str, values: tuple = ()): async with self.connect() as conn: result = await conn.execute(query, values) - row = await result.fetchone() - await result.close() + row = result.fetchone() + result.close() return row async def fetch_page( diff --git a/poetry.lock b/poetry.lock index 9e8ccfca2..12117d39b 100644 --- a/poetry.lock +++ b/poetry.lock @@ -1,5 +1,23 @@ # This file is automatically @generated by Poetry 1.8.3 and should not be changed by hand. +[[package]] +name = "aiosqlite" +version = "0.20.0" +description = "asyncio bridge to the standard sqlite3 module" +optional = false +python-versions = ">=3.8" +files = [ + {file = "aiosqlite-0.20.0-py3-none-any.whl", hash = "sha256:36a1deaca0cac40ebe32aac9977a6e2bbc7f5189f23f4a54d5908986729e5bd6"}, + {file = "aiosqlite-0.20.0.tar.gz", hash = "sha256:6d35c8c256637f4672f843c31021464090805bf925385ac39473fb16eaaca3d7"}, +] + +[package.dependencies] +typing_extensions = ">=4.0" + +[package.extras] +dev = ["attribution (==1.7.0)", "black (==24.2.0)", "coverage[toml] (==7.4.1)", "flake8 (==7.0.0)", "flake8-bugbear (==24.2.6)", "flit (==3.9.0)", "mypy (==1.8.0)", "ufmt (==2.3.0)", "usort (==1.0.8.post1)"] +docs = ["sphinx (==7.2.6)", "sphinx-mdinclude (==0.5.3)"] + [[package]] name = "anyio" version = "4.4.0" @@ -3165,4 +3183,4 @@ liquid = ["wallycore"] [metadata] lock-version = "2.0" python-versions = "^3.10 | ^3.9" -content-hash = "b0aef5d221aed5b287da4e97da2da1053750f8cb1d8be01d3f5cc25f61a78cf9" +content-hash = "56cbae093e02e5165df2c73d62b054e71f9d431fb704726b625473720496d473" diff --git a/pyproject.toml b/pyproject.toml index dd52f0f5c..72e52871d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -59,6 +59,7 @@ wallycore = {version = "1.3.0", optional = true} # needed for breez funding source breez-sdk = {version = "0.5.2", optional = true} +aiosqlite = "^0.20.0" [tool.poetry.extras] breez = ["breez-sdk"] liquid = ["wallycore"]