codekingpro/portable-devtools
114k
1"""2psycopg async connection objects3"""4 5# Copyright (C) 2020 The Psycopg Team6 7import sys8import asyncio9import logging10from types import TracebackType11from typing import Any, AsyncGenerator, AsyncIterator, Dict, List, Optional12from typing import Type, TypeVar, Union, cast, overload, TYPE_CHECKING13from contextlib import asynccontextmanager14 15from . import pq16from . import errors as e17from . import waiting18from .abc import AdaptContext, Params, PQGen, PQGenConn, Query, RV19from ._tpc import Xid20from .rows import Row, AsyncRowFactory, tuple_row, TupleRow, args_row21from .adapt import AdaptersMap22from ._enums import IsolationLevel23from .conninfo import make_conninfo, conninfo_to_dict, resolve_hostaddr_async24from ._pipeline import AsyncPipeline25from ._encodings import pgconn_encoding26from .connection import BaseConnection, CursorRow, Notify27from .generators import notifies28from .transaction import AsyncTransaction29from .cursor_async import AsyncCursor30from .server_cursor import AsyncServerCursor31 32if TYPE_CHECKING:33 from .pq.abc import PGconn34 35TEXT = pq.Format.TEXT36BINARY = pq.Format.BINARY37 38IDLE = pq.TransactionStatus.IDLE39INTRANS = pq.TransactionStatus.INTRANS40 41logger = logging.getLogger("psycopg")42 43 44class AsyncConnection(BaseConnection[Row]):45 """46 Asynchronous wrapper for a connection to the database.47 """48 49 __module__ = "psycopg"50 51 cursor_factory: Type[AsyncCursor[Row]]52 server_cursor_factory: Type[AsyncServerCursor[Row]]53 row_factory: AsyncRowFactory[Row]54 _pipeline: Optional[AsyncPipeline]55 _Self = TypeVar("_Self", bound="AsyncConnection[Any]")56 57 def __init__(58 self,59 pgconn: "PGconn",60 row_factory: AsyncRowFactory[Row] = cast(AsyncRowFactory[Row], tuple_row),61 ):62 super().__init__(pgconn)63 self.row_factory = row_factory64 self.lock = asyncio.Lock()65 self.cursor_factory = AsyncCursor66 self.server_cursor_factory = AsyncServerCursor67 68 @overload69 @classmethod70 async def connect(71 cls,72 conninfo: str = "",73 *,74 autocommit: bool = False,75 prepare_threshold: Optional[int] = 5,76 row_factory: AsyncRowFactory[Row],77 cursor_factory: Optional[Type[AsyncCursor[Row]]] = None,78 context: Optional[AdaptContext] = None,79 **kwargs: Union[None, int, str],80 ) -> "AsyncConnection[Row]":81 # TODO: returned type should be _Self. See #308.82 ...83 84 @overload85 @classmethod86 async def connect(87 cls,88 conninfo: str = "",89 *,90 autocommit: bool = False,91 prepare_threshold: Optional[int] = 5,92 cursor_factory: Optional[Type[AsyncCursor[Any]]] = None,93 context: Optional[AdaptContext] = None,94 **kwargs: Union[None, int, str],95 ) -> "AsyncConnection[TupleRow]":96 ...97 98 @classmethod # type: ignore[misc] # https://github.com/python/mypy/issues/1100499 async def connect(100 cls,101 conninfo: str = "",102 *,103 autocommit: bool = False,104 prepare_threshold: Optional[int] = 5,105 context: Optional[AdaptContext] = None,106 row_factory: Optional[AsyncRowFactory[Row]] = None,107 cursor_factory: Optional[Type[AsyncCursor[Row]]] = None,108 **kwargs: Any,109 ) -> "AsyncConnection[Any]":110 if sys.platform == "win32":111 loop = asyncio.get_running_loop()112 if isinstance(loop, asyncio.ProactorEventLoop):113 raise e.InterfaceError(114 "Psycopg cannot use the 'ProactorEventLoop' to run in async"115 " mode. Please use a compatible event loop, for instance by"116 " setting 'asyncio.set_event_loop_policy"117 "(WindowsSelectorEventLoopPolicy())'"118 )119 120 params = await cls._get_connection_params(conninfo, **kwargs)121 conninfo = make_conninfo(**params)122 123 try:124 rv = await cls._wait_conn(125 cls._connect_gen(conninfo, autocommit=autocommit),126 timeout=params["connect_timeout"],127 )128 except e._NO_TRACEBACK as ex:129 raise ex.with_traceback(None)130 131 if row_factory:132 rv.row_factory = row_factory133 if cursor_factory:134 rv.cursor_factory = cursor_factory135 if context:136 rv._adapters = AdaptersMap(context.adapters)137 rv.prepare_threshold = prepare_threshold138 return rv139 140 async def __aenter__(self: _Self) -> _Self:141 return self142 143 async def __aexit__(144 self,145 exc_type: Optional[Type[BaseException]],146 exc_val: Optional[BaseException],147 exc_tb: Optional[TracebackType],148 ) -> None:149 if self.closed:150 return151 152 if exc_type:153 # try to rollback, but if there are problems (connection in a bad154 # state) just warn without clobbering the exception bubbling up.155 try:156 await self.rollback()157 except Exception as exc2:158 logger.warning(159 "error ignored in rollback on %s: %s",160 self,161 exc2,162 )163 else:164 await self.commit()165 166 # Close the connection only if it doesn't belong to a pool.167 if not getattr(self, "_pool", None):168 await self.close()169 170 @classmethod171 async def _get_connection_params(172 cls, conninfo: str, **kwargs: Any173 ) -> Dict[str, Any]:174 """Manipulate connection parameters before connecting.175 176 .. versionchanged:: 3.1177 Unlike the sync counterpart, perform non-blocking address178 resolution and populate the ``hostaddr`` connection parameter,179 unless the user has provided one themselves. See180 `~psycopg._dns.resolve_hostaddr_async()` for details.181 182 """183 params = conninfo_to_dict(conninfo, **kwargs)184 185 # Make sure there is an usable connect_timeout186 if "connect_timeout" in params:187 params["connect_timeout"] = int(params["connect_timeout"])188 else:189 params["connect_timeout"] = None190 191 # Resolve host addresses in non-blocking way192 params = await resolve_hostaddr_async(params)193 194 return params195 196 async def close(self) -> None:197 if self.closed:198 return199 self._closed = True200 201 # TODO: maybe send a cancel on close, if the connection is ACTIVE?202 203 self.pgconn.finish()204 205 @overload206 def cursor(self, *, binary: bool = False) -> AsyncCursor[Row]:207 ...208 209 @overload210 def cursor(211 self, *, binary: bool = False, row_factory: AsyncRowFactory[CursorRow]212 ) -> AsyncCursor[CursorRow]:213 ...214 215 @overload216 def cursor(217 self,218 name: str,219 *,220 binary: bool = False,221 scrollable: Optional[bool] = None,222 withhold: bool = False,223 ) -> AsyncServerCursor[Row]:224 ...225 226 @overload227 def cursor(228 self,229 name: str,230 *,231 binary: bool = False,232 row_factory: AsyncRowFactory[CursorRow],233 scrollable: Optional[bool] = None,234 withhold: bool = False,235 ) -> AsyncServerCursor[CursorRow]:236 ...237 238 def cursor(239 self,240 name: str = "",241 *,242 binary: bool = False,243 row_factory: Optional[AsyncRowFactory[Any]] = None,244 scrollable: Optional[bool] = None,245 withhold: bool = False,246 ) -> Union[AsyncCursor[Any], AsyncServerCursor[Any]]:247 """248 Return a new `AsyncCursor` to send commands and queries to the connection.249 """250 self._check_connection_ok()251 252 if not row_factory:253 row_factory = self.row_factory254 255 cur: Union[AsyncCursor[Any], AsyncServerCursor[Any]]256 if name:257 cur = self.server_cursor_factory(258 self,259 name=name,260 row_factory=row_factory,261 scrollable=scrollable,262 withhold=withhold,263 )264 else:265 cur = self.cursor_factory(self, row_factory=row_factory)266 267 if binary:268 cur.format = BINARY269 270 return cur271 272 async def execute(273 self,274 query: Query,275 params: Optional[Params] = None,276 *,277 prepare: Optional[bool] = None,278 binary: bool = False,279 ) -> AsyncCursor[Row]:280 try:281 cur = self.cursor()282 if binary:283 cur.format = BINARY284 285 return await cur.execute(query, params, prepare=prepare)286 287 except e._NO_TRACEBACK as ex:288 raise ex.with_traceback(None)289 290 async def commit(self) -> None:291 async with self.lock:292 await self.wait(self._commit_gen())293 294 async def rollback(self) -> None:295 async with self.lock:296 await self.wait(self._rollback_gen())297 298 @asynccontextmanager299 async def transaction(300 self,301 savepoint_name: Optional[str] = None,302 force_rollback: bool = False,303 ) -> AsyncIterator[AsyncTransaction]:304 """305 Start a context block with a new transaction or nested transaction.306 307 :rtype: AsyncTransaction308 """309 tx = AsyncTransaction(self, savepoint_name, force_rollback)310 if self._pipeline:311 async with self.pipeline(), tx, self.pipeline():312 yield tx313 else:314 async with tx:315 yield tx316 317 async def notifies(self) -> AsyncGenerator[Notify, None]:318 while True:319 async with self.lock:320 try:321 ns = await self.wait(notifies(self.pgconn))322 except e._NO_TRACEBACK as ex:323 raise ex.with_traceback(None)324 enc = pgconn_encoding(self.pgconn)325 for pgn in ns:326 n = Notify(pgn.relname.decode(enc), pgn.extra.decode(enc), pgn.be_pid)327 yield n328 329 @asynccontextmanager330 async def pipeline(self) -> AsyncIterator[AsyncPipeline]:331 """Context manager to switch the connection into pipeline mode."""332 async with self.lock:333 self._check_connection_ok()334 335 pipeline = self._pipeline336 if pipeline is None:337 # WARNING: reference loop, broken ahead.338 pipeline = self._pipeline = AsyncPipeline(self)339 340 try:341 async with pipeline:342 yield pipeline343 finally:344 if pipeline.level == 0:345 async with self.lock:346 assert pipeline is self._pipeline347 self._pipeline = None348 349 async def wait(self, gen: PQGen[RV], timeout: Optional[float] = 0.1) -> RV:350 try:351 return await waiting.wait_async(gen, self.pgconn.socket, timeout=timeout)352 except (asyncio.CancelledError, KeyboardInterrupt):353 # On Ctrl-C, try to cancel the query in the server, otherwise354 # the connection will remain stuck in ACTIVE state.355 self._try_cancel(self.pgconn)356 try:357 await waiting.wait_async(gen, self.pgconn.socket, timeout=timeout)358 except e.QueryCanceled:359 pass # as expected360 raise361 362 @classmethod363 async def _wait_conn(cls, gen: PQGenConn[RV], timeout: Optional[int]) -> RV:364 return await waiting.wait_conn_async(gen, timeout)365 366 def _set_autocommit(self, value: bool) -> None:367 self._no_set_async("autocommit")368 369 async def set_autocommit(self, value: bool) -> None:370 """Async version of the `~Connection.autocommit` setter."""371 async with self.lock:372 await self.wait(self._set_autocommit_gen(value))373 374 def _set_isolation_level(self, value: Optional[IsolationLevel]) -> None:375 self._no_set_async("isolation_level")376 377 async def set_isolation_level(self, value: Optional[IsolationLevel]) -> None:378 """Async version of the `~Connection.isolation_level` setter."""379 async with self.lock:380 await self.wait(self._set_isolation_level_gen(value))381 382 def _set_read_only(self, value: Optional[bool]) -> None:383 self._no_set_async("read_only")384 385 async def set_read_only(self, value: Optional[bool]) -> None:386 """Async version of the `~Connection.read_only` setter."""387 async with self.lock:388 await self.wait(self._set_read_only_gen(value))389 390 def _set_deferrable(self, value: Optional[bool]) -> None:391 self._no_set_async("deferrable")392 393 async def set_deferrable(self, value: Optional[bool]) -> None:394 """Async version of the `~Connection.deferrable` setter."""395 async with self.lock:396 await self.wait(self._set_deferrable_gen(value))397 398 def _no_set_async(self, attribute: str) -> None:399 raise AttributeError(400 f"'the {attribute!r} property is read-only on async connections:"401 f" please use 'await .set_{attribute}()' instead."402 )403 404 async def tpc_begin(self, xid: Union[Xid, str]) -> None:405 async with self.lock:406 await self.wait(self._tpc_begin_gen(xid))407 408 async def tpc_prepare(self) -> None:409 try:410 async with self.lock:411 await self.wait(self._tpc_prepare_gen())412 except e.ObjectNotInPrerequisiteState as ex:413 raise e.NotSupportedError(str(ex)) from None414 415 async def tpc_commit(self, xid: Union[Xid, str, None] = None) -> None:416 async with self.lock:417 await self.wait(self._tpc_finish_gen("commit", xid))418 419 async def tpc_rollback(self, xid: Union[Xid, str, None] = None) -> None:420 async with self.lock:421 await self.wait(self._tpc_finish_gen("rollback", xid))422 423 async def tpc_recover(self) -> List[Xid]:424 self._check_tpc()425 status = self.info.transaction_status426 async with self.cursor(row_factory=args_row(Xid._from_record)) as cur:427 await cur.execute(Xid._get_recover_query())428 res = await cur.fetchall()429 430 if status == IDLE and self.info.transaction_status == INTRANS:431 await self.rollback()432 433 return res434 