codekingpro/portable-devtools
114k
1"""2Helper object to transform values between Python and PostgreSQL3"""4 5# Copyright (C) 2020 The Psycopg Team6 7from typing import Any, Dict, List, Optional, Sequence, Tuple8from typing import DefaultDict, TYPE_CHECKING9from collections import defaultdict10from typing_extensions import TypeAlias11 12from . import pq13from . import postgres14from . import errors as e15from .abc import Buffer, LoadFunc, AdaptContext, PyFormat, DumperKey, NoneType16from .rows import Row, RowMaker17from .postgres import INVALID_OID, TEXT_OID18from ._encodings import pgconn_encoding19 20if TYPE_CHECKING:21 from .abc import Dumper, Loader22 from .adapt import AdaptersMap23 from .pq.abc import PGresult24 from .connection import BaseConnection25 26DumperCache: TypeAlias = Dict[DumperKey, "Dumper"]27OidDumperCache: TypeAlias = Dict[int, "Dumper"]28LoaderCache: TypeAlias = Dict[int, "Loader"]29 30TEXT = pq.Format.TEXT31PY_TEXT = PyFormat.TEXT32 33 34class Transformer(AdaptContext):35 """36 An object that can adapt efficiently between Python and PostgreSQL.37 38 The life cycle of the object is the query, so it is assumed that attributes39 such as the server version or the connection encoding will not change. The40 object have its state so adapting several values of the same type can be41 optimised.42 43 """44 45 __module__ = "psycopg.adapt"46 47 __slots__ = """48 types formats49 _conn _adapters _pgresult _dumpers _loaders _encoding _none_oid50 _oid_dumpers _oid_types _row_dumpers _row_loaders51 """.split()52 53 types: Optional[Tuple[int, ...]]54 formats: Optional[List[pq.Format]]55 56 _adapters: "AdaptersMap"57 _pgresult: Optional["PGresult"]58 _none_oid: int59 60 def __init__(self, context: Optional[AdaptContext] = None):61 self._pgresult = self.types = self.formats = None62 63 # WARNING: don't store context, or you'll create a loop with the Cursor64 if context:65 self._adapters = context.adapters66 self._conn = context.connection67 else:68 self._adapters = postgres.adapters69 self._conn = None70 71 # mapping fmt, class -> Dumper instance72 self._dumpers: DefaultDict[PyFormat, DumperCache]73 self._dumpers = defaultdict(dict)74 75 # mapping fmt, oid -> Dumper instance76 # Not often used, so create it only if needed.77 self._oid_dumpers: Optional[Tuple[OidDumperCache, OidDumperCache]]78 self._oid_dumpers = None79 80 # mapping fmt, oid -> Loader instance81 self._loaders: Tuple[LoaderCache, LoaderCache] = ({}, {})82 83 self._row_dumpers: Optional[List["Dumper"]] = None84 85 # sequence of load functions from value to python86 # the length of the result columns87 self._row_loaders: List[LoadFunc] = []88 89 # mapping oid -> type sql representation90 self._oid_types: Dict[int, bytes] = {}91 92 self._encoding = ""93 94 @classmethod95 def from_context(cls, context: Optional[AdaptContext]) -> "Transformer":96 """97 Return a Transformer from an AdaptContext.98 99 If the context is a Transformer instance, just return it.100 """101 if isinstance(context, Transformer):102 return context103 else:104 return cls(context)105 106 @property107 def connection(self) -> Optional["BaseConnection[Any]"]:108 return self._conn109 110 @property111 def encoding(self) -> str:112 if not self._encoding:113 conn = self.connection114 self._encoding = pgconn_encoding(conn.pgconn) if conn else "utf-8"115 return self._encoding116 117 @property118 def adapters(self) -> "AdaptersMap":119 return self._adapters120 121 @property122 def pgresult(self) -> Optional["PGresult"]:123 return self._pgresult124 125 def set_pgresult(126 self,127 result: Optional["PGresult"],128 *,129 set_loaders: bool = True,130 format: Optional[pq.Format] = None,131 ) -> None:132 self._pgresult = result133 134 if not result:135 self._nfields = self._ntuples = 0136 if set_loaders:137 self._row_loaders = []138 return139 140 self._ntuples = result.ntuples141 nf = self._nfields = result.nfields142 143 if not set_loaders:144 return145 146 if not nf:147 self._row_loaders = []148 return149 150 fmt: pq.Format151 fmt = result.fformat(0) if format is None else format # type: ignore152 self._row_loaders = [153 self.get_loader(result.ftype(i), fmt).load for i in range(nf)154 ]155 156 def set_dumper_types(self, types: Sequence[int], format: pq.Format) -> None:157 self._row_dumpers = [self.get_dumper_by_oid(oid, format) for oid in types]158 self.types = tuple(types)159 self.formats = [format] * len(types)160 161 def set_loader_types(self, types: Sequence[int], format: pq.Format) -> None:162 self._row_loaders = [self.get_loader(oid, format).load for oid in types]163 164 def dump_sequence(165 self, params: Sequence[Any], formats: Sequence[PyFormat]166 ) -> Sequence[Optional[Buffer]]:167 nparams = len(params)168 out: List[Optional[Buffer]] = [None] * nparams169 170 # If we have dumpers, it means set_dumper_types had been called, in171 # which case self.types and self.formats are set to sequences of the172 # right size.173 if self._row_dumpers:174 for i in range(nparams):175 param = params[i]176 if param is not None:177 out[i] = self._row_dumpers[i].dump(param)178 return out179 180 types = [self._get_none_oid()] * nparams181 pqformats = [TEXT] * nparams182 183 for i in range(nparams):184 param = params[i]185 if param is None:186 continue187 dumper = self.get_dumper(param, formats[i])188 out[i] = dumper.dump(param)189 types[i] = dumper.oid190 pqformats[i] = dumper.format191 192 self.types = tuple(types)193 self.formats = pqformats194 195 return out196 197 def as_literal(self, obj: Any) -> bytes:198 dumper = self.get_dumper(obj, PY_TEXT)199 rv = dumper.quote(obj)200 # If the result is quoted, and the oid not unknown or text,201 # add an explicit type cast.202 # Check the last char because the first one might be 'E'.203 oid = dumper.oid204 if oid and rv and rv[-1] == b"'"[0] and oid != TEXT_OID:205 try:206 type_sql = self._oid_types[oid]207 except KeyError:208 ti = self.adapters.types.get(oid)209 if ti:210 if oid < 8192:211 # builtin: prefer "timestamptz" to "timestamp with time zone"212 type_sql = ti.name.encode(self.encoding)213 else:214 type_sql = ti.regtype.encode(self.encoding)215 if oid == ti.array_oid:216 type_sql += b"[]"217 else:218 type_sql = b""219 self._oid_types[oid] = type_sql220 221 if type_sql:222 rv = b"%s::%s" % (rv, type_sql)223 224 if not isinstance(rv, bytes):225 rv = bytes(rv)226 return rv227 228 def get_dumper(self, obj: Any, format: PyFormat) -> "Dumper":229 """230 Return a Dumper instance to dump `!obj`.231 """232 # Normally, the type of the object dictates how to dump it233 key = type(obj)234 235 # Reuse an existing Dumper class for objects of the same type236 cache = self._dumpers[format]237 try:238 dumper = cache[key]239 except KeyError:240 # If it's the first time we see this type, look for a dumper241 # configured for it.242 try:243 dcls = self.adapters.get_dumper(key, format)244 except e.ProgrammingError as ex:245 raise ex from None246 else:247 cache[key] = dumper = dcls(key, self)248 249 # Check if the dumper requires an upgrade to handle this specific value250 key1 = dumper.get_key(obj, format)251 if key1 is key:252 return dumper253 254 # If it does, ask the dumper to create its own upgraded version255 try:256 return cache[key1]257 except KeyError:258 dumper = cache[key1] = dumper.upgrade(obj, format)259 return dumper260 261 def _get_none_oid(self) -> int:262 try:263 return self._none_oid264 except AttributeError:265 pass266 267 try:268 rv = self._none_oid = self._adapters.get_dumper(NoneType, PY_TEXT).oid269 except KeyError:270 raise e.InterfaceError("None dumper not found")271 272 return rv273 274 def get_dumper_by_oid(self, oid: int, format: pq.Format) -> "Dumper":275 """276 Return a Dumper to dump an object to the type with given oid.277 """278 if not self._oid_dumpers:279 self._oid_dumpers = ({}, {})280 281 # Reuse an existing Dumper class for objects of the same type282 cache = self._oid_dumpers[format]283 try:284 return cache[oid]285 except KeyError:286 # If it's the first time we see this type, look for a dumper287 # configured for it.288 dcls = self.adapters.get_dumper_by_oid(oid, format)289 cache[oid] = dumper = dcls(NoneType, self)290 291 return dumper292 293 def load_rows(self, row0: int, row1: int, make_row: RowMaker[Row]) -> List[Row]:294 res = self._pgresult295 if not res:296 raise e.InterfaceError("result not set")297 298 if not (0 <= row0 <= self._ntuples and 0 <= row1 <= self._ntuples):299 raise e.InterfaceError(300 f"rows must be included between 0 and {self._ntuples}"301 )302 303 records = []304 for row in range(row0, row1):305 record: List[Any] = [None] * self._nfields306 for col in range(self._nfields):307 val = res.get_value(row, col)308 if val is not None:309 record[col] = self._row_loaders[col](val)310 records.append(make_row(record))311 312 return records313 314 def load_row(self, row: int, make_row: RowMaker[Row]) -> Optional[Row]:315 res = self._pgresult316 if not res:317 return None318 319 if not 0 <= row < self._ntuples:320 return None321 322 record: List[Any] = [None] * self._nfields323 for col in range(self._nfields):324 val = res.get_value(row, col)325 if val is not None:326 record[col] = self._row_loaders[col](val)327 328 return make_row(record)329 330 def load_sequence(self, record: Sequence[Optional[Buffer]]) -> Tuple[Any, ...]:331 if len(self._row_loaders) != len(record):332 raise e.ProgrammingError(333 f"cannot load sequence of {len(record)} items:"334 f" {len(self._row_loaders)} loaders registered"335 )336 337 return tuple(338 (self._row_loaders[i](val) if val is not None else None)339 for i, val in enumerate(record)340 )341 342 def get_loader(self, oid: int, format: pq.Format) -> "Loader":343 try:344 return self._loaders[format][oid]345 except KeyError:346 pass347 348 loader_cls = self._adapters.get_loader(oid, format)349 if not loader_cls:350 loader_cls = self._adapters.get_loader(INVALID_OID, format)351 if not loader_cls:352 raise e.InterfaceError("unknown oid loader not found")353 loader = self._loaders[format][oid] = loader_cls(oid, self)354 return loader355 