diff --git a/squeaknode/config/config.py b/squeaknode/config/config.py index 93579559..272caebc 100644 --- a/squeaknode/config/config.py +++ b/squeaknode/config/config.py @@ -16,6 +16,7 @@ logger = logging.getLogger(__name__) DEFAULT_NETWORK = "testnet" DEFAULT_PRICE_MSAT = 10000 DEFAULT_LOG_LEVEL = "INFO" +DEFAULT_MAX_SQUEAKS = 10000 DEFAULT_MAX_SQUEAKS_PER_ADDRESS_PER_BLOCK = 100 DEFAULT_SERVER_RPC_HOST = "0.0.0.0" DEFAULT_SERVER_RPC_PORT = None @@ -91,14 +92,20 @@ class WebadminConfig(Config): @section('core') class CoreConfig(Config): - network = key(cast=str, required=False, default=DEFAULT_NETWORK) + network = key( + cast=str, required=False, default=DEFAULT_NETWORK) default_peer_rpc_port = key( cast=int, required=False, default=DEFAULT_SERVER_RPC_PORT) - price_msat = key(cast=int, required=False, default=DEFAULT_PRICE_MSAT) + price_msat = key( + cast=int, required=False, default=DEFAULT_PRICE_MSAT) + max_squeaks = key( + cast=int, required=False, default=DEFAULT_MAX_SQUEAKS) max_squeaks_per_address_per_block = key( cast=int, required=False, default=DEFAULT_MAX_SQUEAKS_PER_ADDRESS_PER_BLOCK) - sqk_dir_path = key(cast=str, required=False, default=DEFAULT_SQK_DIR_PATH) - log_level = key(cast=str, required=False, default=DEFAULT_LOG_LEVEL) + sqk_dir_path = key( + cast=str, required=False, default=DEFAULT_SQK_DIR_PATH) + log_level = key( + cast=str, required=False, default=DEFAULT_LOG_LEVEL) sent_offer_retention_s = key( cast=int, required=False, default=DEFAULT_SENT_OFFER_RETENTION_S) subscribe_invoices_retry_s = key( @@ -109,8 +116,8 @@ class CoreConfig(Config): cast=int, required=False, default=DEFAULT_SQUEAK_DELETION_INTERVAL_S) offer_deletion_interval_s = key( cast=int, required=False, default=DEFAULT_OFFER_DELETION_INTERVAL_S) - interest_block_interval = key(cast=int, required=False, - default=DEFAULT_INTEREST_BLOCK_INTERVAL) + interest_block_interval = key( + cast=int, required=False, default=DEFAULT_INTEREST_BLOCK_INTERVAL) @section('db') diff --git a/squeaknode/db/squeak_db.py b/squeaknode/db/squeak_db.py index b8cd4e40..e4b6adda 100644 --- a/squeaknode/db/squeak_db.py +++ b/squeaknode/db/squeak_db.py @@ -413,6 +413,20 @@ class SqueakDb: hashes = [bytes.fromhex(row["hash"]) for row in rows] return hashes + def get_number_of_squeaks(self) -> int: + """ Get total number of squeaks. """ + s = ( + select([ + func.count().label("num_squeaks"), + ]) + .select_from(self.squeaks) + ) + with self.get_connection() as connection: + result = connection.execute(s) + row = result.fetchone() + num_squeaks = row["num_squeaks"] + return num_squeaks + def number_of_squeaks_with_address_with_block( self, address: str, @@ -423,7 +437,7 @@ class SqueakDb: select([ func.count().label("num_squeaks"), ]) - .select_from(self.received_payments) + .select_from(self.squeaks) .where(self.squeaks.c.author_address == address) .where(self.squeaks.c.n_block_height == block_height) ) diff --git a/squeaknode/node/squeak_controller.py b/squeaknode/node/squeak_controller.py index daa71bbd..4e9c6541 100644 --- a/squeaknode/node/squeak_controller.py +++ b/squeaknode/node/squeak_controller.py @@ -57,23 +57,15 @@ class SqueakController: self.new_squeak_listener = NewSqueakListener() self.config = config - def save_squeak( - self, - squeak: CSqueak, - skip_rate_limit: bool = False, - ) -> bytes: + def save_squeak(self, squeak: CSqueak) -> bytes: # Check if squeak is valid. CheckSqueak(squeak, skipDecryptionCheck=True) + # Check if block hash is valid. squeak_entry = self.squeak_core.validate_squeak(squeak) - # TODO: Check if rate limit is violated. - # if not skip_rate_limit: - # if not self.squeak_rate_limiter.should_rate_limit_allow(squeak): - # raise Exception( - # "Exceeded allowed number of squeaks per address per block.") - # Save the squeak. - logger.info("Saving squeak: {}".format( - get_hash(squeak).hex(), - )) + # Check if limit exceeded. + if self.get_number_of_squeaks() >= self.config.core.max_squeaks: + raise Exception("Exceeded max number of squeaks.") + # Insert the squeak in db. inserted_squeak_hash = self.squeak_db.insert_squeak( squeak, squeak_entry.block_header) # Unlock the squeak if decryption key exists. @@ -83,14 +75,13 @@ class SqueakController: inserted_squeak_hash, decryption_key, ) + logger.info("Saved squeak: {}".format( + inserted_squeak_hash, + )) self.new_squeak_listener.handle_new_squeak(squeak) - # Return the squeak hash. return inserted_squeak_hash - def get_squeak( - self, - squeak_hash: bytes, - ) -> Optional[CSqueak]: + def get_squeak(self, squeak_hash: bytes) -> Optional[CSqueak]: squeak_entry = self.squeak_db.get_squeak_entry(squeak_hash) if squeak_entry is None: return None @@ -237,7 +228,7 @@ class SqueakController: squeak_profile = self.squeak_db.get_profile(profile_id) squeak_entry = self.squeak_core.make_squeak( squeak_profile, content_str, replyto_hash) - return self.save_squeak(squeak_entry.squeak, skip_rate_limit=True) + return self.save_squeak(squeak_entry.squeak) def delete_squeak(self, squeak_hash: bytes) -> None: num_deleted_offers = self.squeak_db.delete_offers_for_squeak( @@ -421,6 +412,9 @@ class SqueakController: max_block, ) + def get_number_of_squeaks(self) -> int: + return self.squeak_db.get_number_of_squeaks() + def save_offer(self, received_offer: ReceivedOffer) -> None: logger.info("Saving received offer: {}".format(received_offer)) try: