codekingpro/portable-devtools
114k
1"""2Transaction context managers returned by Connection.transaction()3"""4 5# Copyright (C) 2020 The Psycopg Team6 7import logging8 9from types import TracebackType10from typing import Generic, Iterator, Optional, Type, Union, TypeVar, TYPE_CHECKING11 12from . import pq13from . import sql14from . import errors as e15from .abc import ConnectionType, PQGen16from .pq.misc import connection_summary17 18if TYPE_CHECKING:19 from typing import Any20 from .connection import Connection21 from .connection_async import AsyncConnection22 23IDLE = pq.TransactionStatus.IDLE24 25OK = pq.ConnStatus.OK26 27logger = logging.getLogger(__name__)28 29 30class Rollback(Exception):31 """32 Exit the current `Transaction` context immediately and rollback any changes33 made within this context.34 35 If a transaction context is specified in the constructor, rollback36 enclosing transactions contexts up to and including the one specified.37 """38 39 __module__ = "psycopg"40 41 def __init__(42 self,43 transaction: Union["Transaction", "AsyncTransaction", None] = None,44 ):45 self.transaction = transaction46 47 def __repr__(self) -> str:48 return f"{self.__class__.__qualname__}({self.transaction!r})"49 50 51class OutOfOrderTransactionNesting(e.ProgrammingError):52 """Out-of-order transaction nesting detected"""53 54 55class BaseTransaction(Generic[ConnectionType]):56 def __init__(57 self,58 connection: ConnectionType,59 savepoint_name: Optional[str] = None,60 force_rollback: bool = False,61 ):62 self._conn = connection63 self.pgconn = self._conn.pgconn64 self._savepoint_name = savepoint_name or ""65 self.force_rollback = force_rollback66 self._entered = self._exited = False67 self._outer_transaction = False68 self._stack_index = -169 70 @property71 def savepoint_name(self) -> Optional[str]:72 """73 The name of the savepoint; `!None` if handling the main transaction.74 """75 # Yes, it may change on __enter__. No, I don't care, because the76 # un-entered state is outside the public interface.77 return self._savepoint_name78 79 def __repr__(self) -> str:80 cls = f"{self.__class__.__module__}.{self.__class__.__qualname__}"81 info = connection_summary(self.pgconn)82 if not self._entered:83 status = "inactive"84 elif not self._exited:85 status = "active"86 else:87 status = "terminated"88 89 sp = f"{self.savepoint_name!r} " if self.savepoint_name else ""90 return f"<{cls} {sp}({status}) {info} at 0x{id(self):x}>"91 92 def _enter_gen(self) -> PQGen[None]:93 if self._entered:94 raise TypeError("transaction blocks can be used only once")95 self._entered = True96 97 self._push_savepoint()98 for command in self._get_enter_commands():99 yield from self._conn._exec_command(command)100 101 def _exit_gen(102 self,103 exc_type: Optional[Type[BaseException]],104 exc_val: Optional[BaseException],105 exc_tb: Optional[TracebackType],106 ) -> PQGen[bool]:107 if not exc_val and not self.force_rollback:108 yield from self._commit_gen()109 return False110 else:111 # try to rollback, but if there are problems (connection in a bad112 # state) just warn without clobbering the exception bubbling up.113 try:114 return (yield from self._rollback_gen(exc_val))115 except OutOfOrderTransactionNesting:116 # Clobber an exception happened in the block with the exception117 # caused by out-of-order transaction detected, so make the118 # behaviour consistent with _commit_gen and to make sure the119 # user fixes this condition, which is unrelated from120 # operational error that might arise in the block.121 raise122 except Exception as exc2:123 logger.warning("error ignored in rollback of %s: %s", self, exc2)124 return False125 126 def _commit_gen(self) -> PQGen[None]:127 ex = self._pop_savepoint("commit")128 self._exited = True129 if ex:130 raise ex131 132 for command in self._get_commit_commands():133 yield from self._conn._exec_command(command)134 135 def _rollback_gen(self, exc_val: Optional[BaseException]) -> PQGen[bool]:136 if isinstance(exc_val, Rollback):137 logger.debug(f"{self._conn}: Explicit rollback from: ", exc_info=True)138 139 ex = self._pop_savepoint("rollback")140 self._exited = True141 if ex:142 raise ex143 144 for command in self._get_rollback_commands():145 yield from self._conn._exec_command(command)146 147 if isinstance(exc_val, Rollback):148 if not exc_val.transaction or exc_val.transaction is self:149 return True # Swallow the exception150 151 return False152 153 def _get_enter_commands(self) -> Iterator[bytes]:154 if self._outer_transaction:155 yield self._conn._get_tx_start_command()156 157 if self._savepoint_name:158 yield (159 sql.SQL("SAVEPOINT {}")160 .format(sql.Identifier(self._savepoint_name))161 .as_bytes(self._conn)162 )163 164 def _get_commit_commands(self) -> Iterator[bytes]:165 if self._savepoint_name and not self._outer_transaction:166 yield (167 sql.SQL("RELEASE {}")168 .format(sql.Identifier(self._savepoint_name))169 .as_bytes(self._conn)170 )171 172 if self._outer_transaction:173 assert not self._conn._num_transactions174 yield b"COMMIT"175 176 def _get_rollback_commands(self) -> Iterator[bytes]:177 if self._savepoint_name and not self._outer_transaction:178 yield (179 sql.SQL("ROLLBACK TO {n}")180 .format(n=sql.Identifier(self._savepoint_name))181 .as_bytes(self._conn)182 )183 yield (184 sql.SQL("RELEASE {n}")185 .format(n=sql.Identifier(self._savepoint_name))186 .as_bytes(self._conn)187 )188 189 if self._outer_transaction:190 assert not self._conn._num_transactions191 yield b"ROLLBACK"192 193 # Also clear the prepared statements cache.194 if self._conn._prepared.clear():195 yield from self._conn._prepared.get_maintenance_commands()196 197 def _push_savepoint(self) -> None:198 """199 Push the transaction on the connection transactions stack.200 201 Also set the internal state of the object and verify consistency.202 """203 self._outer_transaction = self.pgconn.transaction_status == IDLE204 if self._outer_transaction:205 # outer transaction: if no name it's only a begin, else206 # there will be an additional savepoint207 assert not self._conn._num_transactions208 else:209 # inner transaction: it always has a name210 if not self._savepoint_name:211 self._savepoint_name = f"_pg3_{self._conn._num_transactions + 1}"212 213 self._stack_index = self._conn._num_transactions214 self._conn._num_transactions += 1215 216 def _pop_savepoint(self, action: str) -> Optional[Exception]:217 """218 Pop the transaction from the connection transactions stack.219 220 Also verify the state consistency.221 """222 self._conn._num_transactions -= 1223 if self._conn._num_transactions == self._stack_index:224 return None225 226 return OutOfOrderTransactionNesting(227 f"transaction {action} at the wrong nesting level: {self}"228 )229 230 231class Transaction(BaseTransaction["Connection[Any]"]):232 """233 Returned by `Connection.transaction()` to handle a transaction block.234 """235 236 __module__ = "psycopg"237 238 _Self = TypeVar("_Self", bound="Transaction")239 240 @property241 def connection(self) -> "Connection[Any]":242 """The connection the object is managing."""243 return self._conn244 245 def __enter__(self: _Self) -> _Self:246 with self._conn.lock:247 self._conn.wait(self._enter_gen())248 return self249 250 def __exit__(251 self,252 exc_type: Optional[Type[BaseException]],253 exc_val: Optional[BaseException],254 exc_tb: Optional[TracebackType],255 ) -> bool:256 if self.pgconn.status == OK:257 with self._conn.lock:258 return self._conn.wait(self._exit_gen(exc_type, exc_val, exc_tb))259 else:260 return False261 262 263class AsyncTransaction(BaseTransaction["AsyncConnection[Any]"]):264 """265 Returned by `AsyncConnection.transaction()` to handle a transaction block.266 """267 268 __module__ = "psycopg"269 270 _Self = TypeVar("_Self", bound="AsyncTransaction")271 272 @property273 def connection(self) -> "AsyncConnection[Any]":274 return self._conn275 276 async def __aenter__(self: _Self) -> _Self:277 async with self._conn.lock:278 await self._conn.wait(self._enter_gen())279 return self280 281 async def __aexit__(282 self,283 exc_type: Optional[Type[BaseException]],284 exc_val: Optional[BaseException],285 exc_tb: Optional[TracebackType],286 ) -> bool:287 if self.pgconn.status == OK:288 async with self._conn.lock:289 return await self._conn.wait(self._exit_gen(exc_type, exc_val, exc_tb))290 else:291 return False292 