Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
connection_async.py434 linesDownload Raw Back to psycopg
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 
codekingpro/portable-devtools · Team Ai