From 3fc34c896531966d37df20be11ecd340daff3092 Mon Sep 17 00:00:00 2001 From: Jonathan Zernik Date: Fri, 22 Oct 2021 19:16:37 -0500 Subject: [PATCH] Add test for set received offer paid db method (#1678) --- squeaknode/core/received_offer.py | 1 + squeaknode/core/squeak_core.py | 1 + squeaknode/db/squeak_db.py | 5 +++-- squeaknode/node/squeak_controller.py | 2 +- tests/conftest.py | 1 + tests/db/test_squeak_db.py | 22 ++++++++++++++++++++++ 6 files changed, 29 insertions(+), 3 deletions(-) diff --git a/squeaknode/core/received_offer.py b/squeaknode/core/received_offer.py index 2e89429b..21f50ac7 100644 --- a/squeaknode/core/received_offer.py +++ b/squeaknode/core/received_offer.py @@ -40,3 +40,4 @@ class ReceivedOffer(NamedTuple): destination: str lightning_address: LightningAddressHostPort peer_address: PeerAddress + paid: bool diff --git a/squeaknode/core/squeak_core.py b/squeaknode/core/squeak_core.py index 2d039f1d..6a6b76be 100644 --- a/squeaknode/core/squeak_core.py +++ b/squeaknode/core/squeak_core.py @@ -287,6 +287,7 @@ class SqueakCore: destination=destination, lightning_address=lightning_address, peer_address=peer_address, + paid=False, ) def pay_offer(self, received_offer: ReceivedOffer) -> SentPayment: diff --git a/squeaknode/db/squeak_db.py b/squeaknode/db/squeak_db.py index c9923d12..f2ba61d3 100644 --- a/squeaknode/db/squeak_db.py +++ b/squeaknode/db/squeak_db.py @@ -1086,11 +1086,11 @@ class SqueakDb: # with self.get_connection() as connection: # connection.execute(s) - def set_received_offer_paid(self, payment_hash: bytes, paid: bool) -> None: + def set_received_offer_paid(self, received_offer_id: int, paid: bool) -> None: """ Set a received offer is paid. """ stmt = ( self.received_offers.update() - .where(self.received_offers.c.payment_hash == payment_hash) + .where(self.received_offers.c.received_offer_id == received_offer_id) .values(paid=paid) ) with self.get_connection() as connection: @@ -1472,6 +1472,7 @@ class SqueakDb: host=row["peer_host"], port=row["peer_port"], ), + paid=row["paid"], ) def _parse_sent_payment(self, row) -> SentPayment: diff --git a/squeaknode/node/squeak_controller.py b/squeaknode/node/squeak_controller.py index 20a6e591..2b08095a 100644 --- a/squeaknode/node/squeak_controller.py +++ b/squeaknode/node/squeak_controller.py @@ -385,7 +385,7 @@ class SqueakController: # self.squeak_db.delete_offer(sent_payment.payment_hash) # Mark the received offer as paid self.squeak_db.set_received_offer_paid( - sent_payment.payment_hash, + received_offer_id, paid=True, ) self.unlock_squeak( diff --git a/tests/conftest.py b/tests/conftest.py index 2a12cd4f..f9551f4b 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -397,6 +397,7 @@ def received_offer( destination=seller_pubkey, lightning_address=lightning_address, peer_address=peer_address, + paid=False, ) diff --git a/tests/db/test_squeak_db.py b/tests/db/test_squeak_db.py index d311bc5a..cd088402 100644 --- a/tests/db/test_squeak_db.py +++ b/tests/db/test_squeak_db.py @@ -316,6 +316,12 @@ def duplicate_inserted_received_offer_id(squeak_db, inserted_received_offer_id, yield squeak_db.insert_received_offer(received_offer) +@pytest.fixture +def paid_received_offer_id(squeak_db, inserted_received_offer_id): + squeak_db.set_received_offer_paid(inserted_received_offer_id, True) + yield inserted_received_offer_id + + def test_init_with_retries(squeak_db): with mock.patch.object(squeak_db, 'init', autospec=True) as mock_init, \ mock.patch('squeaknode.db.squeak_db.time.sleep', autospec=True) as mock_sleep: @@ -1166,3 +1172,19 @@ def test_delete_expired_received_offers_none(squeak_db, inserted_received_offer_ num_deleted = squeak_db.delete_expired_received_offers() assert num_deleted == 0 + + +def test_get_received_offer_paid(squeak_db, paid_received_offer_id): + retrieved_received_offer = squeak_db.get_received_offer( + paid_received_offer_id, + ) + + assert retrieved_received_offer.paid + + +def test_get_received_offer_not_paid(squeak_db, inserted_received_offer_id): + retrieved_received_offer = squeak_db.get_received_offer( + inserted_received_offer_id, + ) + + assert not retrieved_received_offer.paid