Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
asyncpg.py1263 linesDownload Raw Back to postgresql
1# dialects/postgresql/asyncpg.py
2# Copyright (C) 2005-2024 the SQLAlchemy authors and contributors <see AUTHORS
3# file>
4#
5# This module is part of SQLAlchemy and is released under
6# the MIT License: https://www.opensource.org/licenses/mit-license.php
7# mypy: ignore-errors
8
9r"""
10.. dialect:: postgresql+asyncpg
11    :name: asyncpg
12    :dbapi: asyncpg
13    :connectstring: postgresql+asyncpg://user:password@host:port/dbname[?key=value&key=value...]
14    :url: https://magicstack.github.io/asyncpg/
15
16The asyncpg dialect is SQLAlchemy's first Python asyncio dialect.
17
18Using a special asyncio mediation layer, the asyncpg dialect is usable
19as the backend for the :ref:`SQLAlchemy asyncio <asyncio_toplevel>`
20extension package.
21
22This dialect should normally be used only with the
23:func:`_asyncio.create_async_engine` engine creation function::
24
25    from sqlalchemy.ext.asyncio import create_async_engine
26    engine = create_async_engine("postgresql+asyncpg://user:pass@hostname/dbname")
27
28.. versionadded:: 1.4
29
30.. note::
31
32    By default asyncpg does not decode the ``json`` and ``jsonb`` types and
33    returns them as strings. SQLAlchemy sets default type decoder for ``json``
34    and ``jsonb`` types using the python builtin ``json.loads`` function.
35    The json implementation used can be changed by setting the attribute
36    ``json_deserializer`` when creating the engine with
37    :func:`create_engine` or :func:`create_async_engine`.
38
39.. _asyncpg_multihost:
40
41Multihost Connections
42--------------------------
43
44The asyncpg dialect features support for multiple fallback hosts in the
45same way as that of the psycopg2 and psycopg dialects.  The
46syntax is the same,
47using ``host=<host>:<port>`` combinations as additional query string arguments;
48however, there is no default port, so all hosts must have a complete port number
49present, otherwise an exception is raised::
50
51    engine = create_async_engine(
52        "postgresql+asyncpg://user:password@/dbname?host=HostA:5432&host=HostB:5432&host=HostC:5432"
53    )
54
55For complete background on this syntax, see :ref:`psycopg2_multi_host`.
56
57.. versionadded:: 2.0.18
58
59.. seealso::
60
61    :ref:`psycopg2_multi_host`
62
63.. _asyncpg_prepared_statement_cache:
64
65Prepared Statement Cache
66--------------------------
67
68The asyncpg SQLAlchemy dialect makes use of ``asyncpg.connection.prepare()``
69for all statements.   The prepared statement objects are cached after
70construction which appears to grant a 10% or more performance improvement for
71statement invocation.   The cache is on a per-DBAPI connection basis, which
72means that the primary storage for prepared statements is within DBAPI
73connections pooled within the connection pool.   The size of this cache
74defaults to 100 statements per DBAPI connection and may be adjusted using the
75``prepared_statement_cache_size`` DBAPI argument (note that while this argument
76is implemented by SQLAlchemy, it is part of the DBAPI emulation portion of the
77asyncpg dialect, therefore is handled as a DBAPI argument, not a dialect
78argument)::
79
80
81    engine = create_async_engine("postgresql+asyncpg://user:pass@hostname/dbname?prepared_statement_cache_size=500")
82
83To disable the prepared statement cache, use a value of zero::
84
85    engine = create_async_engine("postgresql+asyncpg://user:pass@hostname/dbname?prepared_statement_cache_size=0")
86
87.. versionadded:: 1.4.0b2 Added ``prepared_statement_cache_size`` for asyncpg.
88
89
90.. warning::  The ``asyncpg`` database driver necessarily uses caches for
91   PostgreSQL type OIDs, which become stale when custom PostgreSQL datatypes
92   such as ``ENUM`` objects are changed via DDL operations.   Additionally,
93   prepared statements themselves which are optionally cached by SQLAlchemy's
94   driver as described above may also become "stale" when DDL has been emitted
95   to the PostgreSQL database which modifies the tables or other objects
96   involved in a particular prepared statement.
97
98   The SQLAlchemy asyncpg dialect will invalidate these caches within its local
99   process when statements that represent DDL are emitted on a local
100   connection, but this is only controllable within a single Python process /
101   database engine.     If DDL changes are made from other database engines
102   and/or processes, a running application may encounter asyncpg exceptions
103   ``InvalidCachedStatementError`` and/or ``InternalServerError("cache lookup
104   failed for type <oid>")`` if it refers to pooled database connections which
105   operated upon the previous structures. The SQLAlchemy asyncpg dialect will
106   recover from these error cases when the driver raises these exceptions by
107   clearing its internal caches as well as those of the asyncpg driver in
108   response to them, but cannot prevent them from being raised in the first
109   place if the cached prepared statement or asyncpg type caches have gone
110   stale, nor can it retry the statement as the PostgreSQL transaction is
111   invalidated when these errors occur.
112
113.. _asyncpg_prepared_statement_name:
114
115Prepared Statement Name with PGBouncer
116--------------------------------------
117
118By default, asyncpg enumerates prepared statements in numeric order, which
119can lead to errors if a name has already been taken for another prepared
120statement. This issue can arise if your application uses database proxies
121such as PgBouncer to handle connections. One possible workaround is to
122use dynamic prepared statement names, which asyncpg now supports through
123an optional ``name`` value for the statement name. This allows you to
124generate your own unique names that won't conflict with existing ones.
125To achieve this, you can provide a function that will be called every time
126a prepared statement is prepared::
127
128    from uuid import uuid4
129
130    engine = create_async_engine(
131        "postgresql+asyncpg://user:pass@somepgbouncer/dbname",
132        poolclass=NullPool,
133        connect_args={
134            'prepared_statement_name_func': lambda:  f'__asyncpg_{uuid4()}__',
135        },
136    )
137
138.. seealso::
139
140   https://github.com/MagicStack/asyncpg/issues/837
141
142   https://github.com/sqlalchemy/sqlalchemy/issues/6467
143
144.. warning:: When using PGBouncer, to prevent a buildup of useless prepared statements in
145   your application, it's important to use the :class:`.NullPool` pool
146   class, and to configure PgBouncer to use `DISCARD <https://www.postgresql.org/docs/current/sql-discard.html>`_
147   when returning connections.  The DISCARD command is used to release resources held by the db connection,
148   including prepared statements. Without proper setup, prepared statements can
149   accumulate quickly and cause performance issues.
150
151Disabling the PostgreSQL JIT to improve ENUM datatype handling
152---------------------------------------------------------------
153
154Asyncpg has an `issue <https://github.com/MagicStack/asyncpg/issues/727>`_ when
155using PostgreSQL ENUM datatypes, where upon the creation of new database
156connections, an expensive query may be emitted in order to retrieve metadata
157regarding custom types which has been shown to negatively affect performance.
158To mitigate this issue, the PostgreSQL "jit" setting may be disabled from the
159client using this setting passed to :func:`_asyncio.create_async_engine`::
160
161    engine = create_async_engine(
162        "postgresql+asyncpg://user:password@localhost/tmp",
163        connect_args={"server_settings": {"jit": "off"}},
164    )
165
166.. seealso::
167
168    https://github.com/MagicStack/asyncpg/issues/727
169
170"""  # noqa
171
172from __future__ import annotations
173
174import collections
175import decimal
176import json as _py_json
177import re
178import time
179
180from . import json
181from . import ranges
182from .array import ARRAY as PGARRAY
183from .base import _DECIMAL_TYPES
184from .base import _FLOAT_TYPES
185from .base import _INT_TYPES
186from .base import ENUM
187from .base import INTERVAL
188from .base import OID
189from .base import PGCompiler
190from .base import PGDialect
191from .base import PGExecutionContext
192from .base import PGIdentifierPreparer
193from .base import REGCLASS
194from .base import REGCONFIG
195from .types import BIT
196from .types import BYTEA
197from .types import CITEXT
198from ... import exc
199from ... import pool
200from ... import util
201from ...engine import AdaptedConnection
202from ...engine import processors
203from ...sql import sqltypes
204from ...util.concurrency import asyncio
205from ...util.concurrency import await_fallback
206from ...util.concurrency import await_only
207
208
209class AsyncpgARRAY(PGARRAY):
210    render_bind_cast = True
211
212
213class AsyncpgString(sqltypes.String):
214    render_bind_cast = True
215
216
217class AsyncpgREGCONFIG(REGCONFIG):
218    render_bind_cast = True
219
220
221class AsyncpgTime(sqltypes.Time):
222    render_bind_cast = True
223
224
225class AsyncpgBit(BIT):
226    render_bind_cast = True
227
228
229class AsyncpgByteA(BYTEA):
230    render_bind_cast = True
231
232
233class AsyncpgDate(sqltypes.Date):
234    render_bind_cast = True
235
236
237class AsyncpgDateTime(sqltypes.DateTime):
238    render_bind_cast = True
239
240
241class AsyncpgBoolean(sqltypes.Boolean):
242    render_bind_cast = True
243
244
245class AsyncPgInterval(INTERVAL):
246    render_bind_cast = True
247
248    @classmethod
249    def adapt_emulated_to_native(cls, interval, **kw):
250        return AsyncPgInterval(precision=interval.second_precision)
251
252
253class AsyncPgEnum(ENUM):
254    render_bind_cast = True
255
256
257class AsyncpgInteger(sqltypes.Integer):
258    render_bind_cast = True
259
260
261class AsyncpgBigInteger(sqltypes.BigInteger):
262    render_bind_cast = True
263
264
265class AsyncpgJSON(json.JSON):
266    render_bind_cast = True
267
268    def result_processor(self, dialect, coltype):
269        return None
270
271
272class AsyncpgJSONB(json.JSONB):
273    render_bind_cast = True
274
275    def result_processor(self, dialect, coltype):
276        return None
277
278
279class AsyncpgJSONIndexType(sqltypes.JSON.JSONIndexType):
280    pass
281
282
283class AsyncpgJSONIntIndexType(sqltypes.JSON.JSONIntIndexType):
284    __visit_name__ = "json_int_index"
285
286    render_bind_cast = True
287
288
289class AsyncpgJSONStrIndexType(sqltypes.JSON.JSONStrIndexType):
290    __visit_name__ = "json_str_index"
291
292    render_bind_cast = True
293
294
295class AsyncpgJSONPathType(json.JSONPathType):
296    def bind_processor(self, dialect):
297        def process(value):
298            if isinstance(value, str):
299                # If it's already a string assume that it's in json path
300                # format. This allows using cast with json paths literals
301                return value
302            elif value:
303                tokens = [str(elem) for elem in value]
304                return tokens
305            else:
306                return []
307
308        return process
309
310
311class AsyncpgNumeric(sqltypes.Numeric):
312    render_bind_cast = True
313
314    def bind_processor(self, dialect):
315        return None
316
317    def result_processor(self, dialect, coltype):
318        if self.asdecimal:
319            if coltype in _FLOAT_TYPES:
320                return processors.to_decimal_processor_factory(
321                    decimal.Decimal, self._effective_decimal_return_scale
322                )
323            elif coltype in _DECIMAL_TYPES or coltype in _INT_TYPES:
324                # pg8000 returns Decimal natively for 1700
325                return None
326            else:
327                raise exc.InvalidRequestError(
328                    "Unknown PG numeric type: %d" % coltype
329                )
330        else:
331            if coltype in _FLOAT_TYPES:
332                # pg8000 returns float natively for 701
333                return None
334            elif coltype in _DECIMAL_TYPES or coltype in _INT_TYPES:
335                return processors.to_float
336            else:
337                raise exc.InvalidRequestError(
338                    "Unknown PG numeric type: %d" % coltype
339                )
340
341
342class AsyncpgFloat(AsyncpgNumeric, sqltypes.Float):
343    __visit_name__ = "float"
344    render_bind_cast = True
345
346
347class AsyncpgREGCLASS(REGCLASS):
348    render_bind_cast = True
349
350
351class AsyncpgOID(OID):
352    render_bind_cast = True
353
354
355class AsyncpgCHAR(sqltypes.CHAR):
356    render_bind_cast = True
357
358
359class _AsyncpgRange(ranges.AbstractSingleRangeImpl):
360    def bind_processor(self, dialect):
361        asyncpg_Range = dialect.dbapi.asyncpg.Range
362
363        def to_range(value):
364            if isinstance(value, ranges.Range):
365                value = asyncpg_Range(
366                    value.lower,
367                    value.upper,
368                    lower_inc=value.bounds[0] == "[",
369                    upper_inc=value.bounds[1] == "]",
370                    empty=value.empty,
371                )
372            return value
373
374        return to_range
375
376    def result_processor(self, dialect, coltype):
377        def to_range(value):
378            if value is not None:
379                empty = value.isempty
380                value = ranges.Range(
381                    value.lower,
382                    value.upper,
383                    bounds=f"{'[' if empty or value.lower_inc else '('}"  # type: ignore  # noqa: E501
384                    f"{']' if not empty and value.upper_inc else ')'}",
385                    empty=empty,
386                )
387            return value
388
389        return to_range
390
391
392class _AsyncpgMultiRange(ranges.AbstractMultiRangeImpl):
393    def bind_processor(self, dialect):
394        asyncpg_Range = dialect.dbapi.asyncpg.Range
395
396        NoneType = type(None)
397
398        def to_range(value):
399            if isinstance(value, (str, NoneType)):
400                return value
401
402            def to_range(value):
403                if isinstance(value, ranges.Range):
404                    value = asyncpg_Range(
405                        value.lower,
406                        value.upper,
407                        lower_inc=value.bounds[0] == "[",
408                        upper_inc=value.bounds[1] == "]",
409                        empty=value.empty,
410                    )
411                return value
412
413            return [to_range(element) for element in value]
414
415        return to_range
416
417    def result_processor(self, dialect, coltype):
418        def to_range_array(value):
419            def to_range(rvalue):
420                if rvalue is not None:
421                    empty = rvalue.isempty
422                    rvalue = ranges.Range(
423                        rvalue.lower,
424                        rvalue.upper,
425                        bounds=f"{'[' if empty or rvalue.lower_inc else '('}"  # type: ignore  # noqa: E501
426                        f"{']' if not empty and rvalue.upper_inc else ')'}",
427                        empty=empty,
428                    )
429                return rvalue
430
431            if value is not None:
432                value = ranges.MultiRange(to_range(elem) for elem in value)
433
434            return value
435
436        return to_range_array
437
438
439class PGExecutionContext_asyncpg(PGExecutionContext):
440    def handle_dbapi_exception(self, e):
441        if isinstance(
442            e,
443            (
444                self.dialect.dbapi.InvalidCachedStatementError,
445                self.dialect.dbapi.InternalServerError,
446            ),
447        ):
448            self.dialect._invalidate_schema_cache()
449
450    def pre_exec(self):
451        if self.isddl:
452            self.dialect._invalidate_schema_cache()
453
454        self.cursor._invalidate_schema_cache_asof = (
455            self.dialect._invalidate_schema_cache_asof
456        )
457
458        if not self.compiled:
459            return
460
461    def create_server_side_cursor(self):
462        return self._dbapi_connection.cursor(server_side=True)
463
464
465class PGCompiler_asyncpg(PGCompiler):
466    pass
467
468
469class PGIdentifierPreparer_asyncpg(PGIdentifierPreparer):
470    pass
471
472
473class AsyncAdapt_asyncpg_cursor:
474    __slots__ = (
475        "_adapt_connection",
476        "_connection",
477        "_rows",
478        "description",
479        "arraysize",
480        "rowcount",
481        "_cursor",
482        "_invalidate_schema_cache_asof",
483    )
484
485    server_side = False
486
487    def __init__(self, adapt_connection):
488        self._adapt_connection = adapt_connection
489        self._connection = adapt_connection._connection
490        self._rows = []
491        self._cursor = None
492        self.description = None
493        self.arraysize = 1
494        self.rowcount = -1
495        self._invalidate_schema_cache_asof = 0
496
497    def close(self):
498        self._rows[:] = []
499
500    def _handle_exception(self, error):
501        self._adapt_connection._handle_exception(error)
502
503    async def _prepare_and_execute(self, operation, parameters):
504        adapt_connection = self._adapt_connection
505
506        async with adapt_connection._execute_mutex:
507            if not adapt_connection._started:
508                await adapt_connection._start_transaction()
509
510            if parameters is None:
511                parameters = ()
512
513            try:
514                prepared_stmt, attributes = await adapt_connection._prepare(
515                    operation, self._invalidate_schema_cache_asof
516                )
517
518                if attributes:
519                    self.description = [
520                        (
521                            attr.name,
522                            attr.type.oid,
523                            None,
524                            None,
525                            None,
526                            None,
527                            None,
528                        )
529                        for attr in attributes
530                    ]
531                else:
532                    self.description = None
533
534                if self.server_side:
535                    self._cursor = await prepared_stmt.cursor(*parameters)
536                    self.rowcount = -1
537                else:
538                    self._rows = await prepared_stmt.fetch(*parameters)
539                    status = prepared_stmt.get_statusmsg()
540
541                    reg = re.match(
542                        r"(?:SELECT|UPDATE|DELETE|INSERT \d+) (\d+)", status
543                    )
544                    if reg:
545                        self.rowcount = int(reg.group(1))
546                    else:
547                        self.rowcount = -1
548
549            except Exception as error:
550                self._handle_exception(error)
551
552    async def _executemany(self, operation, seq_of_parameters):
553        adapt_connection = self._adapt_connection
554
555        self.description = None
556        async with adapt_connection._execute_mutex:
557            await adapt_connection._check_type_cache_invalidation(
558                self._invalidate_schema_cache_asof
559            )
560
561            if not adapt_connection._started:
562                await adapt_connection._start_transaction()
563
564            try:
565                return await self._connection.executemany(
566                    operation, seq_of_parameters
567                )
568            except Exception as error:
569                self._handle_exception(error)
570
571    def execute(self, operation, parameters=None):
572        self._adapt_connection.await_(
573            self._prepare_and_execute(operation, parameters)
574        )
575
576    def executemany(self, operation, seq_of_parameters):
577        return self._adapt_connection.await_(
578            self._executemany(operation, seq_of_parameters)
579        )
580
581    def setinputsizes(self, *inputsizes):
582        raise NotImplementedError()
583
584    def __iter__(self):
585        while self._rows:
586            yield self._rows.pop(0)
587
588    def fetchone(self):
589        if self._rows:
590            return self._rows.pop(0)
591        else:
592            return None
593
594    def fetchmany(self, size=None):
595        if size is None:
596            size = self.arraysize
597
598        retval = self._rows[0:size]
599        self._rows[:] = self._rows[size:]
600        return retval
601
602    def fetchall(self):
603        retval = self._rows[:]
604        self._rows[:] = []
605        return retval
606
607
608class AsyncAdapt_asyncpg_ss_cursor(AsyncAdapt_asyncpg_cursor):
609    server_side = True
610    __slots__ = ("_rowbuffer",)
611
612    def __init__(self, adapt_connection):
613        super().__init__(adapt_connection)
614        self._rowbuffer = None
615
616    def close(self):
617        self._cursor = None
618        self._rowbuffer = None
619
620    def _buffer_rows(self):
621        new_rows = self._adapt_connection.await_(self._cursor.fetch(50))
622        self._rowbuffer = collections.deque(new_rows)
623
624    def __aiter__(self):
625        return self
626
627    async def __anext__(self):
628        if not self._rowbuffer:
629            self._buffer_rows()
630
631        while True:
632            while self._rowbuffer:
633                yield self._rowbuffer.popleft()
634
635            self._buffer_rows()
636            if not self._rowbuffer:
637                break
638
639    def fetchone(self):
640        if not self._rowbuffer:
641            self._buffer_rows()
642            if not self._rowbuffer:
643                return None
644        return self._rowbuffer.popleft()
645
646    def fetchmany(self, size=None):
647        if size is None:
648            return self.fetchall()
649
650        if not self._rowbuffer:
651            self._buffer_rows()
652
653        buf = list(self._rowbuffer)
654        lb = len(buf)
655        if size > lb:
656            buf.extend(
657                self._adapt_connection.await_(self._cursor.fetch(size - lb))
658            )
659
660        result = buf[0:size]
661        self._rowbuffer = collections.deque(buf[size:])
662        return result
663
664    def fetchall(self):
665        ret = list(self._rowbuffer) + list(
666            self._adapt_connection.await_(self._all())
667        )
668        self._rowbuffer.clear()
669        return ret
670
671    async def _all(self):
672        rows = []
673
674        # TODO: looks like we have to hand-roll some kind of batching here.
675        # hardcoding for the moment but this should be improved.
676        while True:
677            batch = await self._cursor.fetch(1000)
678            if batch:
679                rows.extend(batch)
680                continue
681            else:
682                break
683        return rows
684
685    def executemany(self, operation, seq_of_parameters):
686        raise NotImplementedError(
687            "server side cursor doesn't support executemany yet"
688        )
689
690
691class AsyncAdapt_asyncpg_connection(AdaptedConnection):
692    __slots__ = (
693        "dbapi",
694        "isolation_level",
695        "_isolation_setting",
696        "readonly",
697        "deferrable",
698        "_transaction",
699        "_started",
700        "_prepared_statement_cache",
701        "_prepared_statement_name_func",
702        "_invalidate_schema_cache_asof",
703        "_execute_mutex",
704    )
705
706    await_ = staticmethod(await_only)
707
708    def __init__(
709        self,
710        dbapi,
711        connection,
712        prepared_statement_cache_size=100,
713        prepared_statement_name_func=None,
714    ):
715        self.dbapi = dbapi
716        self._connection = connection
717        self.isolation_level = self._isolation_setting = "read_committed"
718        self.readonly = False
719        self.deferrable = False
720        self._transaction = None
721        self._started = False
722        self._invalidate_schema_cache_asof = time.time()
723        self._execute_mutex = asyncio.Lock()
724
725        if prepared_statement_cache_size:
726            self._prepared_statement_cache = util.LRUCache(
727                prepared_statement_cache_size
728            )
729        else:
730            self._prepared_statement_cache = None
731
732        if prepared_statement_name_func:
733            self._prepared_statement_name_func = prepared_statement_name_func
734        else:
735            self._prepared_statement_name_func = self._default_name_func
736
737    async def _check_type_cache_invalidation(self, invalidate_timestamp):
738        if invalidate_timestamp > self._invalidate_schema_cache_asof:
739            await self._connection.reload_schema_state()
740            self._invalidate_schema_cache_asof = invalidate_timestamp
741
742    async def _prepare(self, operation, invalidate_timestamp):
743        await self._check_type_cache_invalidation(invalidate_timestamp)
744
745        cache = self._prepared_statement_cache
746        if cache is None:
747            prepared_stmt = await self._connection.prepare(
748                operation, name=self._prepared_statement_name_func()
749            )
750            attributes = prepared_stmt.get_attributes()
751            return prepared_stmt, attributes
752
753        # asyncpg uses a type cache for the "attributes" which seems to go
754        # stale independently of the PreparedStatement itself, so place that
755        # collection in the cache as well.
756        if operation in cache:
757            prepared_stmt, attributes, cached_timestamp = cache[operation]
758
759            # preparedstatements themselves also go stale for certain DDL
760            # changes such as size of a VARCHAR changing, so there is also
761            # a cross-connection invalidation timestamp
762            if cached_timestamp > invalidate_timestamp:
763                return prepared_stmt, attributes
764
765        prepared_stmt = await self._connection.prepare(
766            operation, name=self._prepared_statement_name_func()
767        )
768        attributes = prepared_stmt.get_attributes()
769        cache[operation] = (prepared_stmt, attributes, time.time())
770
771        return prepared_stmt, attributes
772
773    def _handle_exception(self, error):
774        if self._connection.is_closed():
775            self._transaction = None
776            self._started = False
777
778        if not isinstance(error, AsyncAdapt_asyncpg_dbapi.Error):
779            exception_mapping = self.dbapi._asyncpg_error_translate
780
781            for super_ in type(error).__mro__:
782                if super_ in exception_mapping:
783                    translated_error = exception_mapping[super_](
784                        "%s: %s" % (type(error), error)
785                    )
786                    translated_error.pgcode = translated_error.sqlstate = (
787                        getattr(error, "sqlstate", None)
788                    )
789                    raise translated_error from error
790            else:
791                raise error
792        else:
793            raise error
794
795    @property
796    def autocommit(self):
797        return self.isolation_level == "autocommit"
798
799    @autocommit.setter
800    def autocommit(self, value):
801        if value:
802            self.isolation_level = "autocommit"
803        else:
804            self.isolation_level = self._isolation_setting
805
806    def ping(self):
807        try:
808            _ = self.await_(self._async_ping())
809        except Exception as error:
810            self._handle_exception(error)
811
812    async def _async_ping(self):
813        if self._transaction is None and self.isolation_level != "autocommit":
814            # create a tranasction explicitly to support pgbouncer
815            # transaction mode.   See #10226
816            tr = self._connection.transaction()
817            await tr.start()
818            try:
819                await self._connection.fetchrow(";")
820            finally:
821                await tr.rollback()
822        else:
823            await self._connection.fetchrow(";")
824
825    def set_isolation_level(self, level):
826        if self._started:
827            self.rollback()
828        self.isolation_level = self._isolation_setting = level
829
830    async def _start_transaction(self):
831        if self.isolation_level == "autocommit":
832            return
833
834        try:
835            self._transaction = self._connection.transaction(
836                isolation=self.isolation_level,
837                readonly=self.readonly,
838                deferrable=self.deferrable,
839            )
840            await self._transaction.start()
841        except Exception as error:
842            self._handle_exception(error)
843        else:
844            self._started = True
845
846    def cursor(self, server_side=False):
847        if server_side:
848            return AsyncAdapt_asyncpg_ss_cursor(self)
849        else:
850            return AsyncAdapt_asyncpg_cursor(self)
851
852    def rollback(self):
853        if self._started:
854            try:
855                self.await_(self._transaction.rollback())
856            except Exception as error:
857                self._handle_exception(error)
858            finally:
859                self._transaction = None
860                self._started = False
861
862    def commit(self):
863        if self._started:
864            try:
865                self.await_(self._transaction.commit())
866            except Exception as error:
867                self._handle_exception(error)
868            finally:
869                self._transaction = None
870                self._started = False
871
872    def close(self):
873        self.rollback()
874
875        self.await_(self._connection.close())
876
877    def terminate(self):
878        if util.concurrency.in_greenlet():
879            # in a greenlet; this is the connection was invalidated
880            # case.
881            try:
882                # try to gracefully close; see #10717
883                # timeout added in asyncpg 0.14.0 December 2017
884                self.await_(self._connection.close(timeout=2))
885            except (
886                asyncio.TimeoutError,
887                OSError,
888                self.dbapi.asyncpg.PostgresError,
889            ):
890                # in the case where we are recycling an old connection
891                # that may have already been disconnected, close() will
892                # fail with the above timeout.  in this case, terminate
893                # the connection without any further waiting.
894                # see issue #8419
895                self._connection.terminate()
896        else:
897            # not in a greenlet; this is the gc cleanup case
898            self._connection.terminate()
899        self._started = False
900
901    @staticmethod
902    def _default_name_func():
903        return None
904
905
906class AsyncAdaptFallback_asyncpg_connection(AsyncAdapt_asyncpg_connection):
907    __slots__ = ()
908
909    await_ = staticmethod(await_fallback)
910
911
912class AsyncAdapt_asyncpg_dbapi:
913    def __init__(self, asyncpg):
914        self.asyncpg = asyncpg
915        self.paramstyle = "numeric_dollar"
916
917    def connect(self, *arg, **kw):
918        async_fallback = kw.pop("async_fallback", False)
919        creator_fn = kw.pop("async_creator_fn", self.asyncpg.connect)
920        prepared_statement_cache_size = kw.pop(
921            "prepared_statement_cache_size", 100
922        )
923        prepared_statement_name_func = kw.pop(
924            "prepared_statement_name_func", None
925        )
926
927        if util.asbool(async_fallback):
928            return AsyncAdaptFallback_asyncpg_connection(
929                self,
930                await_fallback(creator_fn(*arg, **kw)),
931                prepared_statement_cache_size=prepared_statement_cache_size,
932                prepared_statement_name_func=prepared_statement_name_func,
933            )
934        else:
935            return AsyncAdapt_asyncpg_connection(
936                self,
937                await_only(creator_fn(*arg, **kw)),
938                prepared_statement_cache_size=prepared_statement_cache_size,
939                prepared_statement_name_func=prepared_statement_name_func,
940            )
941
942    class Error(Exception):
943        pass
944
945    class Warning(Exception):  # noqa
946        pass
947
948    class InterfaceError(Error):
949        pass
950
951    class DatabaseError(Error):
952        pass
953
954    class InternalError(DatabaseError):
955        pass
956
957    class OperationalError(DatabaseError):
958        pass
959
960    class ProgrammingError(DatabaseError):
961        pass
962
963    class IntegrityError(DatabaseError):
964        pass
965
966    class DataError(DatabaseError):
967        pass
968
969    class NotSupportedError(DatabaseError):
970        pass
971
972    class InternalServerError(InternalError):
973        pass
974
975    class InvalidCachedStatementError(NotSupportedError):
976        def __init__(self, message):
977            super().__init__(
978                message + " (SQLAlchemy asyncpg dialect will now invalidate "
979                "all prepared caches in response to this exception)",
980            )
981
982    # pep-249 datatype placeholders.  As of SQLAlchemy 2.0 these aren't
983    # used, however the test suite looks for these in a few cases.
984    STRING = util.symbol("STRING")
985    NUMBER = util.symbol("NUMBER")
986    DATETIME = util.symbol("DATETIME")
987
988    @util.memoized_property
989    def _asyncpg_error_translate(self):
990        import asyncpg
991
992        return {
993            asyncpg.exceptions.IntegrityConstraintViolationError: self.IntegrityError,  # noqa: E501
994            asyncpg.exceptions.PostgresError: self.Error,
995            asyncpg.exceptions.SyntaxOrAccessError: self.ProgrammingError,
996            asyncpg.exceptions.InterfaceError: self.InterfaceError,
997            asyncpg.exceptions.InvalidCachedStatementError: self.InvalidCachedStatementError,  # noqa: E501
998            asyncpg.exceptions.InternalServerError: self.InternalServerError,
999        }
1000
1001    def Binary(self, value):
1002        return value
1003
1004
1005class PGDialect_asyncpg(PGDialect):
1006    driver = "asyncpg"
1007    supports_statement_cache = True
1008
1009    supports_server_side_cursors = True
1010
1011    render_bind_cast = True
1012    has_terminate = True
1013
1014    default_paramstyle = "numeric_dollar"
1015    supports_sane_multi_rowcount = False
1016    execution_ctx_cls = PGExecutionContext_asyncpg
1017    statement_compiler = PGCompiler_asyncpg
1018    preparer = PGIdentifierPreparer_asyncpg
1019
1020    colspecs = util.update_copy(
1021        PGDialect.colspecs,
1022        {
1023            sqltypes.String: AsyncpgString,
1024            sqltypes.ARRAY: AsyncpgARRAY,
1025            BIT: AsyncpgBit,
1026            CITEXT: CITEXT,
1027            REGCONFIG: AsyncpgREGCONFIG,
1028            sqltypes.Time: AsyncpgTime,
1029            sqltypes.Date: AsyncpgDate,
1030            sqltypes.DateTime: AsyncpgDateTime,
1031            sqltypes.Interval: AsyncPgInterval,
1032            INTERVAL: AsyncPgInterval,
1033            sqltypes.Boolean: AsyncpgBoolean,
1034            sqltypes.Integer: AsyncpgInteger,
1035            sqltypes.BigInteger: AsyncpgBigInteger,
1036            sqltypes.Numeric: AsyncpgNumeric,
1037            sqltypes.Float: AsyncpgFloat,
1038            sqltypes.JSON: AsyncpgJSON,
1039            sqltypes.LargeBinary: AsyncpgByteA,
1040            json.JSONB: AsyncpgJSONB,
1041            sqltypes.JSON.JSONPathType: AsyncpgJSONPathType,
1042            sqltypes.JSON.JSONIndexType: AsyncpgJSONIndexType,
1043            sqltypes.JSON.JSONIntIndexType: AsyncpgJSONIntIndexType,
1044            sqltypes.JSON.JSONStrIndexType: AsyncpgJSONStrIndexType,
1045            sqltypes.Enum: AsyncPgEnum,
1046            OID: AsyncpgOID,
1047            REGCLASS: AsyncpgREGCLASS,
1048            sqltypes.CHAR: AsyncpgCHAR,
1049            ranges.AbstractSingleRange: _AsyncpgRange,
1050            ranges.AbstractMultiRange: _AsyncpgMultiRange,
1051        },
1052    )
1053    is_async = True
1054    _invalidate_schema_cache_asof = 0
1055
1056    def _invalidate_schema_cache(self):
1057        self._invalidate_schema_cache_asof = time.time()
1058
1059    @util.memoized_property
1060    def _dbapi_version(self):
1061        if self.dbapi and hasattr(self.dbapi, "__version__"):
1062            return tuple(
1063                [
1064                    int(x)
1065                    for x in re.findall(
1066                        r"(\d+)(?:[-\.]?|$)", self.dbapi.__version__
1067                    )
1068                ]
1069            )
1070        else:
1071            return (99, 99, 99)
1072
1073    @classmethod
1074    def import_dbapi(cls):
1075        return AsyncAdapt_asyncpg_dbapi(__import__("asyncpg"))
1076
1077    @util.memoized_property
1078    def _isolation_lookup(self):
1079        return {
1080            "AUTOCOMMIT": "autocommit",
1081            "READ COMMITTED": "read_committed",
1082            "REPEATABLE READ": "repeatable_read",
1083            "SERIALIZABLE": "serializable",
1084        }
1085
1086    def get_isolation_level_values(self, dbapi_connection):
1087        return list(self._isolation_lookup)
1088
1089    def set_isolation_level(self, dbapi_connection, level):
1090        dbapi_connection.set_isolation_level(self._isolation_lookup[level])
1091
1092    def set_readonly(self, connection, value):
1093        connection.readonly = value
1094
1095    def get_readonly(self, connection):
1096        return connection.readonly
1097
1098    def set_deferrable(self, connection, value):
1099        connection.deferrable = value
1100
1101    def get_deferrable(self, connection):
1102        return connection.deferrable
1103
1104    def do_terminate(self, dbapi_connection) -> None:
1105        dbapi_connection.terminate()
1106
1107    def create_connect_args(self, url):
1108        opts = url.translate_connect_args(username="user")
1109        multihosts, multiports = self._split_multihost_from_url(url)
1110
1111        opts.update(url.query)
1112
1113        if multihosts:
1114            assert multiports
1115            if len(multihosts) == 1:
1116                opts["host"] = multihosts[0]
1117                if multiports[0] is not None:
1118                    opts["port"] = multiports[0]
1119            elif not all(multihosts):
1120                raise exc.ArgumentError(
1121                    "All hosts are required to be present"
1122                    " for asyncpg multiple host URL"
1123                )
1124            elif not all(multiports):
1125                raise exc.ArgumentError(
1126                    "All ports are required to be present"
1127                    " for asyncpg multiple host URL"
1128                )
1129            else:
1130                opts["host"] = list(multihosts)
1131                opts["port"] = list(multiports)
1132        else:
1133            util.coerce_kw_type(opts, "port", int)
1134        util.coerce_kw_type(opts, "prepared_statement_cache_size", int)
1135        return ([], opts)
1136
1137    def do_ping(self, dbapi_connection):
1138        dbapi_connection.ping()
1139        return True
1140
1141    @classmethod
1142    def get_pool_class(cls, url):
1143        async_fallback = url.query.get("async_fallback", False)
1144
1145        if util.asbool(async_fallback):
1146            return pool.FallbackAsyncAdaptedQueuePool
1147        else:
1148            return pool.AsyncAdaptedQueuePool
1149
1150    def is_disconnect(self, e, connection, cursor):
1151        if connection:
1152            return connection._connection.is_closed()
1153        else:
1154            return isinstance(
1155                e, self.dbapi.InterfaceError
1156            ) and "connection is closed" in str(e)
1157
1158    async def setup_asyncpg_json_codec(self, conn):
1159        """set up JSON codec for asyncpg.
1160
1161        This occurs for all new connections and
1162        can be overridden by third party dialects.
1163
1164        .. versionadded:: 1.4.27
1165
1166        """
1167
1168        asyncpg_connection = conn._connection
1169        deserializer = self._json_deserializer or _py_json.loads
1170
1171        def _json_decoder(bin_value):
1172            return deserializer(bin_value.decode())
1173
1174        await asyncpg_connection.set_type_codec(
1175            "json",
1176            encoder=str.encode,
1177            decoder=_json_decoder,
1178            schema="pg_catalog",
1179            format="binary",
1180        )
1181
1182    async def setup_asyncpg_jsonb_codec(self, conn):
1183        """set up JSONB codec for asyncpg.
1184
1185        This occurs for all new connections and
1186        can be overridden by third party dialects.
1187
1188        .. versionadded:: 1.4.27
1189
1190        """
1191
1192        asyncpg_connection = conn._connection
1193        deserializer = self._json_deserializer or _py_json.loads
1194
1195        def _jsonb_encoder(str_value):
1196            # \x01 is the prefix for jsonb used by PostgreSQL.
1197            # asyncpg requires it when format='binary'
1198            return b"\x01" + str_value.encode()
1199
1200        deserializer = self._json_deserializer or _py_json.loads

Showing the first 1,200 of 1263 lines. Download the file for the rest.

codekingpro/portable-devtools · Team Ai