diff --git a/crud.py b/crud.py index 0ae3cfe..d66d50e 100644 --- a/crud.py +++ b/crud.py @@ -4,7 +4,16 @@ from typing import List, Optional from lnbits.db import Database from .execution_queue import enqueue -from .models import NWCBudget, NWCKey, CreateNWCKey +from .models import ( + NWCBudget, + NWCKey, + CreateNWCKey, + GetNWCKey, + GetWalletNWCKey, + GetBudgetsNWC, + TrackedSpendNWC, + DeleteNWC +) db = Database("ext_nwcprovider") @@ -30,46 +39,35 @@ async def create_nwc(data: CreateNWCKey) -> NWCKey: await db.insert("nwcprovider.budgets", budget_entry) return NWCKey(**data.dict()) - -async def delete_nwc(pubkey: str, wallet_id: str): - nwc = await get_nwc(pubkey, wallet_id) - if not nwc: - raise Exception("Public key does not exist") +async def delete_nwc(data:DeleteNWC) -> None: await db.execute( - """ - DELETE FROM nwcprovider.keys WHERE pubkey = ? AND wallet = ? - """, - (pubkey, wallet_id), + "DELETE FROM nwcprovider.keys WHERE pubkey = :pubkey AND wallet = :wallet", {"pubkey": data.pubkey, "wallet": data.wallet_id} ) - -async def get_wallet_nwcs( - wallet_id: str, include_expired: Optional[bool] = False -) -> List[NWCKey]: - rows = await db.fetchall( +async def get_wallet_nwcs(data: GetWalletNWCKey) -> List[NWCKey]: + return await db.fetchall( """ - SELECT * FROM nwcprovider.keys - WHERE wallet = ? AND (expires_at = 0 OR expires_at > ?) - """, - (wallet_id, int(time.time()) if not include_expired else -1), + SELECT * FROM nwcprovider.keys + WHERE wallet = :wallet AND (expires_at = 0 OR expires_at > :expires) + """, + {"wallet": data.wallet_id, "expires": int(time.time()) if not data.include_expired else -1}, + model=NWCKey, ) - return [NWCKey(**row) for row in rows] - async def get_nwc( - pubkey: str, - wallet_id: Optional[str] = None, - include_expired: Optional[bool] = False, - refresh_last_used: Optional[bool] = False, + data: GetNWCKey ) -> Optional[NWCKey]: # expires_at = 0 means it never expires - if wallet_id: + if data.wallet_id: row = await db.fetchone( """ - SELECT * FROM nwcprovider.keys - WHERE pubkey = ? AND wallet = ? AND (expires_at = 0 OR expires_at > ?) + SELECT * FROM nwcprovider.keys + WHERE pubkey = :pubkey AND wallet = :wallet + AND (expires_at = 0 OR expires_at > :expires) """, - (pubkey, wallet_id, int(time.time()) if not include_expired else -1), + {"pubkey": data.pubkey, "wallet": data.wallet_id, + "expires": int(time.time()) if not data.include_expired else -1}, + NWCKey, ) else: row = await db.fetchone( @@ -77,44 +75,52 @@ async def get_nwc( SELECT * FROM nwcprovider.keys WHERE pubkey = ? AND (expires_at = 0 OR expires_at > ?) """, - (pubkey, int(time.time()) if not include_expired else -1), + (data.pubkey, int(time.time()) if not data.include_expired else -1), + ) + row = await db.fetchone( + """ + SELECT * FROM nwcprovider.keys + WHERE pubkey = :pubkey AND (expires_at = 0 OR expires_at > :expires) + """, + {"pubkey": data.pubkey, "expires": int(time.time()) if not data.include_expired else -1}, + NWCKey, ) if not row: return None - if refresh_last_used: + if data.refresh_last_used: await db.execute( """ - UPDATE nwcprovider.keys SET last_used = ? WHERE pubkey = ? + UPDATE nwcprovider.keys SET last_used = :last_used WHERE pubkey = :pubkey """, - (int(time.time()), pubkey), + {"last_used":int(time.time()), "pubkey":data.pubkey}, ) return NWCKey(**row) -async def get_budgets_nwc(pubkey, calculate_spent=False): +async def get_budgets_nwc(data: GetBudgetsNWC) -> Optional[NWCBudget]: rows = await db.fetchall( - "SELECT * FROM nwcprovider.budgets WHERE pubkey = ?", (pubkey) + "SELECT * FROM nwcprovider.budgets WHERE pubkey = :pubkey", {"pubkey":data.pubkey} ) budgets = [NWCBudget(**row) for row in rows] - if calculate_spent: + if data.calculate_spent: for budget in budgets: last_cycle, next_cycle = budget.get_timestamp_range() tot_spent_in_range_msats = await db.fetchone( """ SELECT SUM(amount_msats) FROM nwcprovider.spent - WHERE pubkey = ? AND created_at >= ? AND created_at < ? + WHERE pubkey = :pubkey AND created_at >= :last_cycle AND created_at < next_cycle """, - (pubkey, last_cycle, next_cycle), + {"pubkey":data.pubkey, "last_cycle":last_cycle, "next_cycle":next_cycle}, ) tot_spent_in_range_msats = tot_spent_in_range_msats[0] or 0 budget.used_budget_msats = tot_spent_in_range_msats return budgets -async def tracked_spend_nwc(pubkey: str, amount_msats: int, action): +async def tracked_spend_nwc(data:TrackedSpendNWC, action): async def r(): created_at = int(time.time()) - budgets = await get_budgets_nwc(pubkey) + budgets = await get_budgets_nwc(data.pubkey) in_budget = True for budget in budgets: last_cycle, next_cycle = budget.get_timestamp_range() @@ -123,14 +129,14 @@ async def tracked_spend_nwc(pubkey: str, amount_msats: int, action): await db.fetchone( """ SELECT SUM(amount_msats) FROM nwcprovider.spent - WHERE pubkey = ? AND created_at >= ? AND created_at < ? + WHERE pubkey = :pubkey AND created_at >= :last_cycle AND created_at < :next_cycle """, - (pubkey, last_cycle, next_cycle), + {"pubkey":data.pubkey, "last_cycle":last_cycle, "next_cycle":next_cycle}, ) )[0] or 0 ) - if tot_spent_in_range_msats + amount_msats > budget.budget_msats: + if tot_spent_in_range_msats + data.amount_msats > budget.budget_msats: in_budget = False break if not in_budget: @@ -139,9 +145,9 @@ async def tracked_spend_nwc(pubkey: str, amount_msats: int, action): await db.execute( """ INSERT INTO nwcprovider.spent (pubkey, amount_msats, created_at) - VALUES (?, ?, ?) + VALUES (:pubkey, :amount_msats, :created_at) """, - (pubkey, amount_msats, created_at), + {"pubkey":data.pubkey, "amount_msats":data.amount_msats, "created_at":created_at}, ) return True, out @@ -149,29 +155,27 @@ async def tracked_spend_nwc(pubkey: str, amount_msats: int, action): async def get_config_nwc(key: str): - row = await db.fetchone("SELECT * FROM nwcprovider.config WHERE key = ?", (key,)) + row = await db.fetchone("SELECT * FROM nwcprovider.config WHERE key = :key", {"key":key}) if not row: return None return row["value"] - -async def get_all_config_nwc(): - rows = await db.fetchall("SELECT * FROM nwcprovider.config") - return {row["key"]: row["value"] for row in rows} - - async def set_config_nwc(key: str, value: str): await db.execute( """ DELETE FROM nwcprovider.config - WHERE key = ? + WHERE key = :key """, - (key,), + {"key":key}, ) await db.execute( """ INSERT INTO nwcprovider.config (key, value) - VALUES (?, ?) + VALUES (:key, :value) """, - (key, value), + {"key":key, "value":value}, ) + +async def get_all_config_nwc(): + rows = await db.fetchall("SELECT * FROM nwcprovider.config") + return {row["key"]: row["value"] for row in rows} \ No newline at end of file diff --git a/models.py b/models.py index dae61af..b99a52f 100644 --- a/models.py +++ b/models.py @@ -26,6 +26,27 @@ class NWCKey(BaseModel): def from_row(cls, row: Dict[str, Any]) -> "NWCKey": return cls(**row) +class GetNWCKey(BaseModel): + pubkey: str + wallet_id: Optional[str] = None + include_expired: Optional[bool] = False + refresh_last_used: Optional[bool] = False + +class GetWalletNWCKey(BaseModel): + wallet_id: Optional[str] = None + include_expired: Optional[bool] = False + +class GetBudgetsNWC(BaseModel): + pubkey: str + calculate_spent: Optional[bool] = False + +class TrackedSpendNWC(BaseModel): + pubkey: str + amount_msats: int + +class DeleteNWC(BaseModel): + pubkey: str + wallet_id: str class NWCBudget(BaseModel): id: int