1# dialects/postgresql/_psycopg_common.py
2# Copyright (C) 2005-2026 the SQLAlchemy authors and contributors
3# <see AUTHORS 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
8from __future__ import annotations
9
10import decimal
11
12from .array import ARRAY as PGARRAY
13from .base import _DECIMAL_TYPES
14from .base import _FLOAT_TYPES
15from .base import _INT_TYPES
16from .base import PGDialect
17from .base import PGExecutionContext
18from .hstore import HSTORE
19from .pg_catalog import _SpaceVector
20from .pg_catalog import INT2VECTOR
21from .pg_catalog import OIDVECTOR
22from ... import exc
23from ... import types as sqltypes
24from ... import util
25from ...engine import processors
26
27_server_side_id = util.counter()
28
29
30class _PsycopgNumericCommon(sqltypes.NumericCommon):
31 def bind_processor(self, dialect):
32 return None
33
34 def result_processor(self, dialect, coltype):
35 if self.asdecimal:
36 if coltype in _FLOAT_TYPES:
37 return processors.to_decimal_processor_factory(
38 decimal.Decimal, self._effective_decimal_return_scale
39 )
40 elif coltype in _DECIMAL_TYPES or coltype in _INT_TYPES:
41 # psycopg returns Decimal natively for 1700
42 return None
43 else:
44 raise exc.InvalidRequestError(
45 "Unknown PG numeric type: %d" % coltype
46 )
47 else:
48 if coltype in _FLOAT_TYPES:
49 # psycopg returns float natively for 701
50 return None
51 elif coltype in _DECIMAL_TYPES or coltype in _INT_TYPES:
52 return processors.to_float
53 else:
54 raise exc.InvalidRequestError(
55 "Unknown PG numeric type: %d" % coltype
56 )
57
58
59class _PsycopgNumeric(_PsycopgNumericCommon, sqltypes.Numeric):
60 pass
61
62
63class _PsycopgFloat(_PsycopgNumericCommon, sqltypes.Float):
64 pass
65
66
67class _PsycopgHStore(HSTORE):
68 def bind_processor(self, dialect):
69 if dialect._has_native_hstore:
70 return None
71 else:
72 return super().bind_processor(dialect)
73
74 def result_processor(self, dialect, coltype):
75 if dialect._has_native_hstore:
76 return None
77 else:
78 return super().result_processor(dialect, coltype)
79
80
81class _PsycopgARRAY(PGARRAY):
82 render_bind_cast = True
83
84
85class _PsycopgINT2VECTOR(_SpaceVector, INT2VECTOR):
86 pass
87
88
89class _PsycopgOIDVECTOR(_SpaceVector, OIDVECTOR):
90 pass
91
92
93class _PGExecutionContext_common_psycopg(PGExecutionContext):
94 def create_server_side_cursor(self):
95 # use server-side cursors:
96 # psycopg
97 # https://www.psycopg.org/psycopg3/docs/advanced/cursors.html#server-side-cursors
98 # psycopg2
99 # https://www.psycopg.org/docs/usage.html#server-side-cursors
100 ident = "c_%s_%s" % (hex(id(self))[2:], hex(_server_side_id())[2:])
101 return self._dbapi_connection.cursor(ident)
102
103
104class _PGDialect_common_psycopg(PGDialect):
105 supports_statement_cache = True
106 supports_server_side_cursors = True
107
108 default_paramstyle = "pyformat"
109
110 _has_native_hstore = True
111
112 colspecs = util.update_copy(
113 PGDialect.colspecs,
114 {
115 sqltypes.Numeric: _PsycopgNumeric,
116 sqltypes.Float: _PsycopgFloat,
117 HSTORE: _PsycopgHStore,
118 sqltypes.ARRAY: _PsycopgARRAY,
119 INT2VECTOR: _PsycopgINT2VECTOR,
120 OIDVECTOR: _PsycopgOIDVECTOR,
121 },
122 )
123
124 def retrieve_dbapi_version(self, dbapi):
125 return util.parse_version_string(getattr(dbapi, "__version__", None))
126
127 def __init__(
128 self,
129 client_encoding=None,
130 use_native_hstore=True,
131 **kwargs,
132 ):
133 PGDialect.__init__(self, **kwargs)
134 if not use_native_hstore:
135 self._has_native_hstore = False
136 self.use_native_hstore = use_native_hstore
137 self.client_encoding = client_encoding
138
139 def create_connect_args(self, url):
140 opts = url.translate_connect_args(username="user", database="dbname")
141
142 multihosts, multiports = self._split_multihost_from_url(url)
143
144 if opts or url.query:
145 if not opts:
146 opts = {}
147 if "port" in opts:
148 opts["port"] = int(opts["port"])
149 opts.update(url.query)
150
151 if multihosts:
152 opts["host"] = ",".join(multihosts)
153 comma_ports = ",".join(str(p) if p else "" for p in multiports)
154 if comma_ports:
155 opts["port"] = comma_ports
156 return ([], opts)
157 else:
158 # no connection arguments whatsoever; psycopg2.connect()
159 # requires that "dsn" be present as a blank string.
160 return ([""], opts)
161
162 def get_isolation_level_values(self, dbapi_connection):
163 return (
164 "AUTOCOMMIT",
165 "READ COMMITTED",
166 "READ UNCOMMITTED",
167 "REPEATABLE READ",
168 "SERIALIZABLE",
169 )
170
171 def set_deferrable(self, connection, value):
172 connection.deferrable = value
173
174 def get_deferrable(self, connection):
175 return connection.deferrable
176
177 def _do_autocommit(self, connection, value):
178 connection.autocommit = value
179
180 def detect_autocommit_setting(self, dbapi_connection):
181 return bool(dbapi_connection.autocommit)
182
183 def do_ping(self, dbapi_connection):
184 before_autocommit = dbapi_connection.autocommit
185
186 if not before_autocommit:
187 dbapi_connection.autocommit = True
188 cursor = dbapi_connection.cursor()
189 try:
190 cursor.execute(self._dialect_specific_select_one)
191 finally:
192 cursor.close()
193 if not before_autocommit and not dbapi_connection.closed:
194 dbapi_connection.autocommit = before_autocommit
195
196 return True
197
198 def do_begin_twophase(self, connection, xid):
199 connection.connection.tpc_begin(xid)
200
201 def do_prepare_twophase(self, connection, xid):
202 connection.connection.tpc_prepare()
203
204 def _do_twophase(self, dbapi_conn, operation, xid, recover=False):
205 if recover:
206 if not self._twophase_idle_check(dbapi_conn):
207 dbapi_conn.rollback()
208 operation(xid)
209 else:
210 operation()
211
212 def _twophase_idle_check(self, dbapi_conn):
213 raise NotImplementedError
214
215 def do_rollback_twophase(
216 self, connection, xid, is_prepared=True, recover=False
217 ):
218 dbapi_conn = connection.connection.dbapi_connection
219 self._do_twophase(
220 dbapi_conn, dbapi_conn.tpc_rollback, xid, recover=recover
221 )
222
223 def do_commit_twophase(
224 self, connection, xid, is_prepared=True, recover=False
225 ):
226 dbapi_conn = connection.connection.dbapi_connection
227 self._do_twophase(
228 dbapi_conn, dbapi_conn.tpc_commit, xid, recover=recover
229 )
230
231 def do_recover_twophase(self, connection):
232 return [str(row) for row in connection.connection.tpc_recover()]