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