Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 42 additions & 0 deletions src/carbonfactor_parser/persistence/postgresql_execution.py
Original file line number Diff line number Diff line change
@@ -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")
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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),
Expand All @@ -154,18 +160,18 @@ 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),
)
) 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,
Expand All @@ -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)

Expand Down
90 changes: 90 additions & 0 deletions tests/test_postgresql_execution.py
Original file line number Diff line number Diff line change
@@ -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
Loading