Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
_transform.py355 linesDownload Raw Back to psycopg
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 
codekingpro/portable-devtools · Team Ai