codekingpro/portable-devtools
115k
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
