diff --git a/src/carbonfactor_parser/persistence/postgresql_execution.py b/src/carbonfactor_parser/persistence/postgresql_execution.py new file mode 100644 index 0000000..683dc98 --- /dev/null +++ b/src/carbonfactor_parser/persistence/postgresql_execution.py @@ -0,0 +1,42 @@ +"""PostgreSQL connection execution helpers.""" + +from __future__ import annotations + + +def execute( + connection: object, + statement: str, + parameters: object | None = None, +) -> object: + """Execute a statement against a connection-like object.""" + + execute_method = getattr(connection, "execute") + if parameters is None: + return execute_method(statement) + return execute_method(statement, parameters) + + +def fetchone(cursor: object) -> object | None: + """Fetch one row from a cursor-like object.""" + + fetchone_method = getattr(cursor, "fetchone") + return fetchone_method() + + +def commit(connection: object) -> None: + """Commit a connection-like object when it supports commit.""" + + commit_method = getattr(connection, "commit", None) + if commit_method is not None: + commit_method() + + +def rollback(connection: object) -> None: + """Rollback a connection-like object when it supports rollback.""" + + rollback_method = getattr(connection, "rollback", None) + if rollback_method is not None: + rollback_method() + + +__all__ = ("commit", "execute", "fetchone", "rollback") diff --git a/src/carbonfactor_parser/persistence/postgresql_source_family_repository.py b/src/carbonfactor_parser/persistence/postgresql_source_family_repository.py index f26e3f4..e0dc6da 100644 --- a/src/carbonfactor_parser/persistence/postgresql_source_family_repository.py +++ b/src/carbonfactor_parser/persistence/postgresql_source_family_repository.py @@ -10,6 +10,12 @@ from carbonfactor_parser.persistence.parsed_factor_persistence_writer import ( persist_parsed_factor_records, ) +from carbonfactor_parser.persistence.postgresql_execution import ( + commit, + execute, + fetchone, + rollback, +) from carbonfactor_parser.persistence.postgresql_source_family_parameters import ( detail_parameters, master_parameters, @@ -144,8 +150,8 @@ def persist_source_family_records( for master in master_records: ensure_ingestion_run(self._connection, master) ensure_source_document(self._connection, master) - if _fetchone( - _execute( + if fetchone( + execute( self._connection, master_insert_sql(master.source_family), master_parameters(master), @@ -154,8 +160,8 @@ def persist_source_family_records( inserted_masters += 1 for detail in detail_records: - if _fetchone( - _execute( + if fetchone( + execute( self._connection, detail_insert_sql(detail.source_family), detail_parameters(detail), @@ -163,9 +169,9 @@ def persist_source_family_records( ) is not None: inserted_details += 1 - _commit(self._connection) + commit(self._connection) except Exception as exc: # pragma: no cover - driver type varies - _rollback(self._connection) + rollback(self._connection) return SourceFamilyRepositoryPersistResult( provider_name=self.provider_name, status=SourceFamilyRepositoryPersistStatus.FAILED_DATABASE, @@ -190,34 +196,6 @@ def persist_source_family_records( ) -def _execute( - connection: object, - statement: str, - parameters: object | None = None, -) -> object: - execute = getattr(connection, "execute") - if parameters is None: - return execute(statement) - return execute(statement, parameters) - - -def _fetchone(cursor: object) -> object | None: - fetchone = getattr(cursor, "fetchone") - return fetchone() - - -def _commit(connection: object) -> None: - commit = getattr(connection, "commit", None) - if commit is not None: - commit() - - -def _rollback(connection: object) -> None: - rollback = getattr(connection, "rollback", None) - if rollback is not None: - rollback() - - def _redact_sensitive_text(value: str) -> str: return redact_sensitive_text(value) diff --git a/tests/test_postgresql_execution.py b/tests/test_postgresql_execution.py new file mode 100644 index 0000000..546aed3 --- /dev/null +++ b/tests/test_postgresql_execution.py @@ -0,0 +1,90 @@ +from __future__ import annotations + +from carbonfactor_parser.persistence.postgresql_execution import ( + commit, + execute, + fetchone, + rollback, +) + + +class _RecordingConnection: + def __init__(self) -> None: + self.calls: list[tuple[object, ...]] = [] + self.committed = False + self.rolled_back = False + + def execute(self, *args: object) -> object: + self.calls.append(args) + return "cursor" + + def commit(self) -> None: + self.committed = True + + def rollback(self) -> None: + self.rolled_back = True + + +class _ConnectionWithoutTransactionMethods: + pass + + +class _RecordingCursor: + def __init__(self) -> None: + self.called = False + + def fetchone(self) -> tuple[str]: + self.called = True + return ("row",) + + +def test_execute_calls_connection_execute_without_parameters_when_none() -> None: + connection = _RecordingConnection() + + result = execute(connection, "SELECT 1") + + assert result == "cursor" + assert connection.calls == [("SELECT 1",)] + + +def test_execute_calls_connection_execute_with_parameters_when_provided() -> None: + connection = _RecordingConnection() + parameters = ("value",) + + result = execute(connection, "SELECT %s", parameters) + + assert result == "cursor" + assert connection.calls == [("SELECT %s", parameters)] + + +def test_fetchone_calls_cursor_fetchone() -> None: + cursor = _RecordingCursor() + + result = fetchone(cursor) + + assert result == ("row",) + assert cursor.called is True + + +def test_commit_noops_when_method_missing() -> None: + commit(_ConnectionWithoutTransactionMethods()) + + +def test_commit_calls_connection_commit_when_present() -> None: + connection = _RecordingConnection() + + commit(connection) + + assert connection.committed is True + + +def test_rollback_noops_when_method_missing() -> None: + rollback(_ConnectionWithoutTransactionMethods()) + + +def test_rollback_calls_connection_rollback_when_present() -> None: + connection = _RecordingConnection() + + rollback(connection) + + assert connection.rolled_back is True