codekingpro/portable-devtools
114k
1"""2Utility module to manipulate queries3"""4 5# Copyright (C) 2020 The Psycopg Team6 7import re8from typing import Any, Callable, Dict, List, Mapping, Match, NamedTuple, Optional9from typing import Sequence, Tuple, Union, TYPE_CHECKING10from functools import lru_cache11from typing_extensions import TypeAlias12 13from . import pq14from . import errors as e15from .sql import Composable16from .abc import Buffer, Query, Params17from ._enums import PyFormat18from ._encodings import conn_encoding19 20if TYPE_CHECKING:21 from .abc import Transformer22 23MAX_CACHED_STATEMENT_LENGTH = 409624MAX_CACHED_STATEMENT_PARAMS = 5025 26 27class QueryPart(NamedTuple):28 pre: bytes29 item: Union[int, str]30 format: PyFormat31 32 33class PostgresQuery:34 """35 Helper to convert a Python query and parameters into Postgres format.36 """37 38 __slots__ = """39 query params types formats40 _tx _want_formats _parts _encoding _order41 """.split()42 43 def __init__(self, transformer: "Transformer"):44 self._tx = transformer45 46 self.params: Optional[Sequence[Optional[Buffer]]] = None47 # these are tuples so they can be used as keys e.g. in prepared stmts48 self.types: Tuple[int, ...] = ()49 50 # The format requested by the user and the ones to really pass Postgres51 self._want_formats: Optional[List[PyFormat]] = None52 self.formats: Optional[Sequence[pq.Format]] = None53 54 self._encoding = conn_encoding(transformer.connection)55 self._parts: List[QueryPart]56 self.query = b""57 self._order: Optional[List[str]] = None58 59 def convert(self, query: Query, vars: Optional[Params]) -> None:60 """61 Set up the query and parameters to convert.62 63 The results of this function can be obtained accessing the object64 attributes (`query`, `params`, `types`, `formats`).65 """66 if isinstance(query, str):67 bquery = query.encode(self._encoding)68 elif isinstance(query, Composable):69 bquery = query.as_bytes(self._tx)70 else:71 bquery = query72 73 if vars is not None:74 # Avoid caching queries extremely long or with a huge number of75 # parameters. They are usually generated by ORMs and have poor76 # cacheablility (e.g. INSERT ... VALUES (...), (...) with varying77 # numbers of tuples.78 # see https://github.com/psycopg/psycopg/discussions/62879 if (80 len(bquery) <= MAX_CACHED_STATEMENT_LENGTH81 and len(vars) <= MAX_CACHED_STATEMENT_PARAMS82 ):83 f: _Query2Pg = _query2pg84 else:85 f = _query2pg_nocache86 87 (self.query, self._want_formats, self._order, self._parts) = f(88 bquery, self._encoding89 )90 else:91 self.query = bquery92 self._want_formats = self._order = None93 94 self.dump(vars)95 96 def dump(self, vars: Optional[Params]) -> None:97 """98 Process a new set of variables on the query processed by `convert()`.99 100 This method updates `params` and `types`.101 """102 if vars is not None:103 params = _validate_and_reorder_params(self._parts, vars, self._order)104 assert self._want_formats is not None105 self.params = self._tx.dump_sequence(params, self._want_formats)106 self.types = self._tx.types or ()107 self.formats = self._tx.formats108 else:109 self.params = None110 self.types = ()111 self.formats = None112 113 114# The type of the _query2pg() and _query2pg_nocache() methods115_Query2Pg: TypeAlias = Callable[116 [bytes, str], Tuple[bytes, List[PyFormat], Optional[List[str]], List[QueryPart]]117]118 119 120def _query2pg_nocache(121 query: bytes, encoding: str122) -> Tuple[bytes, List[PyFormat], Optional[List[str]], List[QueryPart]]:123 """124 Convert Python query and params into something Postgres understands.125 126 - Convert Python placeholders (``%s``, ``%(name)s``) into Postgres127 format (``$1``, ``$2``)128 - placeholders can be %s, %t, or %b (auto, text or binary)129 - return ``query`` (bytes), ``formats`` (list of formats) ``order``130 (sequence of names used in the query, in the position they appear)131 ``parts`` (splits of queries and placeholders).132 """133 parts = _split_query(query, encoding)134 order: Optional[List[str]] = None135 chunks: List[bytes] = []136 formats = []137 138 if isinstance(parts[0].item, int):139 for part in parts[:-1]:140 assert isinstance(part.item, int)141 chunks.append(part.pre)142 chunks.append(b"$%d" % (part.item + 1))143 formats.append(part.format)144 145 elif isinstance(parts[0].item, str):146 seen: Dict[str, Tuple[bytes, PyFormat]] = {}147 order = []148 for part in parts[:-1]:149 assert isinstance(part.item, str)150 chunks.append(part.pre)151 if part.item not in seen:152 ph = b"$%d" % (len(seen) + 1)153 seen[part.item] = (ph, part.format)154 order.append(part.item)155 chunks.append(ph)156 formats.append(part.format)157 else:158 if seen[part.item][1] != part.format:159 raise e.ProgrammingError(160 f"placeholder '{part.item}' cannot have different formats"161 )162 chunks.append(seen[part.item][0])163 164 # last part165 chunks.append(parts[-1].pre)166 167 return b"".join(chunks), formats, order, parts168 169 170# Note: the cache size is 128 items, but someone has reported throwing ~12k171# queries (of type `INSERT ... VALUES (...), (...)` with a varying amount of172# records), and the resulting cache size is >100Mb. So, we will avoid to cache173# large queries or queries with a large number of params. See174# https://github.com/sqlalchemy/sqlalchemy/discussions/10270175_query2pg = lru_cache()(_query2pg_nocache)176 177 178class PostgresClientQuery(PostgresQuery):179 """180 PostgresQuery subclass merging query and arguments client-side.181 """182 183 __slots__ = ("template",)184 185 def convert(self, query: Query, vars: Optional[Params]) -> None:186 """187 Set up the query and parameters to convert.188 189 The results of this function can be obtained accessing the object190 attributes (`query`, `params`, `types`, `formats`).191 """192 if isinstance(query, str):193 bquery = query.encode(self._encoding)194 elif isinstance(query, Composable):195 bquery = query.as_bytes(self._tx)196 else:197 bquery = query198 199 if vars is not None:200 if (201 len(bquery) <= MAX_CACHED_STATEMENT_LENGTH202 and len(vars) <= MAX_CACHED_STATEMENT_PARAMS203 ):204 f: _Query2PgClient = _query2pg_client205 else:206 f = _query2pg_client_nocache207 208 (self.template, self._order, self._parts) = f(bquery, self._encoding)209 else:210 self.query = bquery211 self._order = None212 213 self.dump(vars)214 215 def dump(self, vars: Optional[Params]) -> None:216 """217 Process a new set of variables on the query processed by `convert()`.218 219 This method updates `params` and `types`.220 """221 if vars is not None:222 params = _validate_and_reorder_params(self._parts, vars, self._order)223 self.params = tuple(224 self._tx.as_literal(p) if p is not None else b"NULL" for p in params225 )226 self.query = self.template % self.params227 else:228 self.params = None229 230 231_Query2PgClient: TypeAlias = Callable[232 [bytes, str], Tuple[bytes, Optional[List[str]], List[QueryPart]]233]234 235 236def _query2pg_client_nocache(237 query: bytes, encoding: str238) -> Tuple[bytes, Optional[List[str]], List[QueryPart]]:239 """240 Convert Python query and params into a template to perform client-side binding241 """242 parts = _split_query(query, encoding, collapse_double_percent=False)243 order: Optional[List[str]] = None244 chunks: List[bytes] = []245 246 if isinstance(parts[0].item, int):247 for part in parts[:-1]:248 assert isinstance(part.item, int)249 chunks.append(part.pre)250 chunks.append(b"%s")251 252 elif isinstance(parts[0].item, str):253 seen: Dict[str, Tuple[bytes, PyFormat]] = {}254 order = []255 for part in parts[:-1]:256 assert isinstance(part.item, str)257 chunks.append(part.pre)258 if part.item not in seen:259 ph = b"%s"260 seen[part.item] = (ph, part.format)261 order.append(part.item)262 chunks.append(ph)263 else:264 chunks.append(seen[part.item][0])265 order.append(part.item)266 267 # last part268 chunks.append(parts[-1].pre)269 270 return b"".join(chunks), order, parts271 272 273_query2pg_client = lru_cache()(_query2pg_client_nocache)274 275 276def _validate_and_reorder_params(277 parts: List[QueryPart], vars: Params, order: Optional[List[str]]278) -> Sequence[Any]:279 """280 Verify the compatibility between a query and a set of params.281 """282 # Try concrete types, then abstract types283 t = type(vars)284 if t is list or t is tuple:285 sequence = True286 elif t is dict:287 sequence = False288 elif isinstance(vars, Sequence) and not isinstance(vars, (bytes, str)):289 sequence = True290 elif isinstance(vars, Mapping):291 sequence = False292 else:293 raise TypeError(294 "query parameters should be a sequence or a mapping,"295 f" got {type(vars).__name__}"296 )297 298 if sequence:299 if len(vars) != len(parts) - 1:300 raise e.ProgrammingError(301 f"the query has {len(parts) - 1} placeholders but"302 f" {len(vars)} parameters were passed"303 )304 if vars and not isinstance(parts[0].item, int):305 raise TypeError("named placeholders require a mapping of parameters")306 return vars # type: ignore[return-value]307 308 else:309 if vars and len(parts) > 1 and not isinstance(parts[0][1], str):310 raise TypeError(311 "positional placeholders (%s) require a sequence of parameters"312 )313 try:314 return [vars[item] for item in order or ()] # type: ignore[call-overload]315 except KeyError:316 raise e.ProgrammingError(317 "query parameter missing:"318 f" {', '.join(sorted(i for i in order or () if i not in vars))}"319 )320 321 322_re_placeholder = re.compile(323 rb"""(?x)324 % # a literal %325 (?:326 (?:327 \( ([^)]+) \) # or a name in (braces)328 . # followed by a format329 )330 |331 (?:.) # or any char, really332 )333 """334)335 336 337def _split_query(338 query: bytes, encoding: str = "ascii", collapse_double_percent: bool = True339) -> List[QueryPart]:340 parts: List[Tuple[bytes, Optional[Match[bytes]]]] = []341 cur = 0342 343 # pairs [(fragment, match], with the last match None344 m = None345 for m in _re_placeholder.finditer(query):346 pre = query[cur : m.span(0)[0]]347 parts.append((pre, m))348 cur = m.span(0)[1]349 if m:350 parts.append((query[cur:], None))351 else:352 parts.append((query, None))353 354 rv = []355 356 # drop the "%%", validate357 i = 0358 phtype = None359 while i < len(parts):360 pre, m = parts[i]361 if m is None:362 # last part363 rv.append(QueryPart(pre, 0, PyFormat.AUTO))364 break365 366 ph = m.group(0)367 if ph == b"%%":368 # unescape '%%' to '%' if necessary, then merge the parts369 if collapse_double_percent:370 ph = b"%"371 pre1, m1 = parts[i + 1]372 parts[i + 1] = (pre + ph + pre1, m1)373 del parts[i]374 continue375 376 if ph == b"%(":377 raise e.ProgrammingError(378 "incomplete placeholder:"379 f" '{query[m.span(0)[0]:].split()[0].decode(encoding)}'"380 )381 elif ph == b"% ":382 # explicit messasge for a typical error383 raise e.ProgrammingError(384 "incomplete placeholder: '%'; if you want to use '%' as an"385 " operator you can double it up, i.e. use '%%'"386 )387 elif ph[-1:] not in b"sbt":388 raise e.ProgrammingError(389 "only '%s', '%b', '%t' are allowed as placeholders, got"390 f" '{m.group(0).decode(encoding)}'"391 )392 393 # Index or name394 item: Union[int, str]395 item = m.group(1).decode(encoding) if m.group(1) else i396 397 if not phtype:398 phtype = type(item)399 elif phtype is not type(item):400 raise e.ProgrammingError(401 "positional and named placeholders cannot be mixed"402 )403 404 format = _ph_to_fmt[ph[-1:]]405 rv.append(QueryPart(pre, item, format))406 i += 1407 408 return rv409 410 411_ph_to_fmt = {412 b"s": PyFormat.AUTO,413 b"t": PyFormat.TEXT,414 b"b": PyFormat.BINARY,415}416 