diff --git a/glances/exports/glances_timescaledb/__init__.py b/glances/exports/glances_timescaledb/__init__.py index 580307b9..6e6577f5 100644 --- a/glances/exports/glances_timescaledb/__init__.py +++ b/glances/exports/glances_timescaledb/__init__.py @@ -9,6 +9,7 @@ """TimescaleDB interface class.""" import sys +import threading import time from datetime import datetime, timezone from platform import node @@ -60,6 +61,10 @@ class Export(GlancesExport): # Init the TimescaleDB client self.client = self.init() + # A psycopg connection has a single transaction state. Glances may start + # another export thread before the previous one has completed, so keep + # transactions on this persistent connection strictly serialized. + self._client_lock = threading.Lock() def init(self): """Init the connection to the TimescaleDB server.""" @@ -171,66 +176,72 @@ class Export(GlancesExport): """Export the stats to the TimescaleDB server.""" logger.debug(f"Export {plugin} stats to TimescaleDB") - with self.client.cursor() as cur: - # Is the table exists? - cur.execute( - "SELECT EXISTS(SELECT * FROM information_schema.tables WHERE table_name=%s)", - [plugin], - ) - if not cur.fetchone()[0]: - # Create the table if it does not exist - # https://github.com/timescale/timescaledb/blob/main/README.md#create-a-hypertable - # Build CREATE TABLE using sql.Identifier for column names (prevents injection) - # Each item in creation_list is "colname TYPE [NULL|NOT NULL]" - fields = sql.SQL(', ').join( - sql.SQL("{} {}").format(sql.Identifier(item.split(' ')[0]), sql.SQL(' '.join(item.split(' ')[1:]))) - for item in creation_list - ) - create_query = sql.SQL( - "CREATE TABLE {table} ({fields}) WITH (" - "timescaledb.hypertable, " - "timescaledb.partition_column='time', " - "timescaledb.segmentby = {segmentby});" - ).format( - table=sql.Identifier(plugin), - fields=fields, - segmentby=sql.Literal(', '.join(segmented_by)), - ) - logger.debug(f"Create table: {create_query}") - try: - cur.execute(create_query) - except Exception as e: - logger.error(f"Cannot create table {plugin}: {e}") - return - - # Insert the data using parameterized queries (prevents injection) - # https://github.com/timescale/timescaledb/blob/main/README.md#insert-and-query-data - col_names = [item.split(' ')[0] for item in creation_list] - cols = sql.SQL(', ').join(sql.Identifier(c) for c in col_names) - placeholders = sql.SQL(', ').join(sql.Placeholder() for _ in col_names) - insert_query = sql.SQL("INSERT INTO {table} ({cols}) VALUES ({vals})").format( - table=sql.Identifier(plugin), - cols=cols, - vals=placeholders, - ) - logger.debug(f"Insert data into table: {insert_query}") + operation = 'check table' + with self._client_lock: try: - cur.executemany(insert_query, values_list) - except Exception as e: - logger.error(f"Cannot insert data into table {plugin}: {e}") - return + # The transaction context commits on success and, critically, + # rolls back before propagating an error. Catch outside the + # context so a suppressed error cannot leave the connection in + # PostgreSQL's failed-transaction state. + with self.client.transaction(): + with self.client.cursor() as cur: + cur.execute( + "SELECT EXISTS(SELECT * FROM information_schema.tables WHERE table_name=%s)", + [plugin], + ) + if not cur.fetchone()[0]: + operation = 'create table' + # Create the table if it does not exist + # https://github.com/timescale/timescaledb/blob/main/README.md#create-a-hypertable + # Build CREATE TABLE using sql.Identifier for column names (prevents injection) + # Each item in creation_list is "colname TYPE [NULL|NOT NULL]" + fields = sql.SQL(', ').join( + sql.SQL("{} {}").format( + sql.Identifier(item.split(' ')[0]), sql.SQL(' '.join(item.split(' ')[1:])) + ) + for item in creation_list + ) + create_query = sql.SQL( + "CREATE TABLE {table} ({fields}) WITH (" + "timescaledb.hypertable, " + "timescaledb.partition_column='time', " + "timescaledb.segmentby = {segmentby});" + ).format( + table=sql.Identifier(plugin), + fields=fields, + segmentby=sql.Literal(', '.join(segmented_by)), + ) + logger.debug(f"Create table: {create_query}") + cur.execute(create_query) - # Commit the changes (for every plugin or to be done at the end ?) - self.client.commit() + operation = 'insert data into table' + # Insert the data using parameterized queries (prevents injection) + # https://github.com/timescale/timescaledb/blob/main/README.md#insert-and-query-data + col_names = [item.split(' ')[0] for item in creation_list] + cols = sql.SQL(', ').join(sql.Identifier(c) for c in col_names) + placeholders = sql.SQL(', ').join(sql.Placeholder() for _ in col_names) + insert_query = sql.SQL("INSERT INTO {table} ({cols}) VALUES ({vals})").format( + table=sql.Identifier(plugin), + cols=cols, + vals=placeholders, + ) + logger.debug(f"Insert data into table: {insert_query}") + cur.executemany(insert_query, values_list) + except Exception as e: + logger.error(f"Cannot {operation} {plugin}: {e}") + return False + + return True def exit(self): """Close the TimescaleDB export module.""" - # Force last write - self.client.commit() + with self._client_lock: + # Force last write + self.client.commit() - # Close the TimescaleDB client - time.sleep(3) # Wait a bit to ensure all data is written - self.client.close() + # Close the TimescaleDB client + time.sleep(3) # Wait a bit to ensure all data is written + self.client.close() # Call the father method super().exit() diff --git a/tests/test_export_timescaledb_list.py b/tests/test_export_timescaledb_list.py index e9fb2cc5..64c161b2 100755 --- a/tests/test_export_timescaledb_list.py +++ b/tests/test_export_timescaledb_list.py @@ -13,8 +13,11 @@ Tests cover: - The number of generated columns matches the number of values (issue #3592) - The 'key' field is exported once, as key_id - No stat field is dropped from the exported row +- Failed database writes are rolled back before the next export """ +import threading + import pytest try: @@ -25,10 +28,85 @@ except ImportError: from glances.exports.glances_timescaledb import Export +class FakeTransaction: + """Model the commit/rollback contract of psycopg's transaction context.""" + + def __init__(self, connection): + """Track transaction outcomes on the fake connection.""" + self.connection = connection + + def __enter__(self): + """Enter the fake transaction context.""" + return self + + def __exit__(self, exc_type, exc_value, traceback): + """Commit successful work or roll back a failed transaction.""" + if exc_type is None: + self.connection.commit_count += 1 + else: + self.connection.rollback_count += 1 + self.connection.transaction_failed = False + return False + + +class FakeCursor: + """Simulate the cursor behavior needed by the exporter.""" + + def __init__(self, connection): + """Use transaction state from the fake connection.""" + self.connection = connection + + def __enter__(self): + """Enter the fake cursor context.""" + return self + + def __exit__(self, exc_type, exc_value, traceback): + """Leave the fake cursor context without suppressing exceptions.""" + return False + + def execute(self, query, parameters=None): + """Reject statements while the transaction is in a failed state.""" + if self.connection.transaction_failed: + raise RuntimeError('current transaction is aborted') + + def fetchone(self): + """Report that the table already exists.""" + return (True,) + + def executemany(self, query, values): + """Fail the first insert and record rows from later inserts.""" + if self.connection.fail_next_insert: + self.connection.fail_next_insert = False + self.connection.transaction_failed = True + raise RuntimeError('simulated insert failure') + self.connection.insert_count += len(values) + + +class FakeConnection: + """Track transaction recovery and successfully inserted rows.""" + + def __init__(self): + """Configure the first insert to fail.""" + self.fail_next_insert = True + self.transaction_failed = False + self.commit_count = 0 + self.rollback_count = 0 + self.insert_count = 0 + + def transaction(self): + """Return a transaction context associated with this connection.""" + return FakeTransaction(self) + + def cursor(self): + """Return a cursor associated with this connection.""" + return FakeCursor(self) + + class FakeStats: """Minimal stats stub exposing a single list plugin.""" def getAllExportsAsDict(self, plugin_list=None): + """Return representative network statistics.""" return { 'network': [ {'key': 'interface_name', 'interface_name': 'eth0', 'bytes_recv': 10, 'bytes_sent': 20}, @@ -37,6 +115,7 @@ class FakeStats: } def getAllLimitsAsDict(self, plugin_list=None): + """Return empty limits for the network plugin.""" return {'network': {}} @@ -59,6 +138,7 @@ class TestTimescaleDBListPlugin: @pytest.fixture def captured(self): + """Capture the normalized values passed to the database layer.""" captured = {} build_export(captured).update(FakeStats()) return captured @@ -84,3 +164,33 @@ class TestTimescaleDBListPlugin: assert row['interface_name'] == 'eth0' assert row['bytes_recv'] == 10 assert row['bytes_sent'] == 20 + + +def test_failed_insert_is_rolled_back_before_next_export(): + """A database error must not poison the persistent connection.""" + connection = FakeConnection() + export = object.__new__(Export) + export.client = connection + export._client_lock = threading.Lock() + export_args = ( + 'cpu', + ['time TIMESTAMPTZ NOT NULL', 'hostname_id TEXT NOT NULL', 'value BIGINT NULL'], + ['hostname_id'], + [[object(), 'testhost', 1]], + ) + + first_result = export.export(*export_args) + if first_result is not False: + pytest.fail(f'Expected failed export to return False, got {first_result!r}') + if connection.rollback_count != 1: + pytest.fail(f'Expected one rollback, got {connection.rollback_count}') + if connection.transaction_failed is not False: + pytest.fail('Expected the failed transaction to be cleared') + + second_result = export.export(*export_args) + if second_result is not True: + pytest.fail(f'Expected recovered export to return True, got {second_result!r}') + if connection.commit_count != 1: + pytest.fail(f'Expected one commit, got {connection.commit_count}') + if connection.insert_count != 1: + pytest.fail(f'Expected one inserted row, got {connection.insert_count}')