Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
util.py1487 linesDownload Raw Back to sql
1# sql/util.py
2# Copyright (C) 2005-2024 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: allow-untyped-defs, allow-untyped-calls
8
9"""High level utilities which build upon other modules here.
10
11"""
12from __future__ import annotations
13
14from collections import deque
15import copy
16from itertools import chain
17import typing
18from typing import AbstractSet
19from typing import Any
20from typing import Callable
21from typing import cast
22from typing import Collection
23from typing import Dict
24from typing import Iterable
25from typing import Iterator
26from typing import List
27from typing import Optional
28from typing import overload
29from typing import Sequence
30from typing import Tuple
31from typing import TYPE_CHECKING
32from typing import TypeVar
33from typing import Union
34
35from . import coercions
36from . import operators
37from . import roles
38from . import visitors
39from ._typing import is_text_clause
40from .annotation import _deep_annotate as _deep_annotate  # noqa: F401
41from .annotation import _deep_deannotate as _deep_deannotate  # noqa: F401
42from .annotation import _shallow_annotate as _shallow_annotate  # noqa: F401
43from .base import _expand_cloned
44from .base import _from_objects
45from .cache_key import HasCacheKey as HasCacheKey  # noqa: F401
46from .ddl import sort_tables as sort_tables  # noqa: F401
47from .elements import _find_columns as _find_columns
48from .elements import _label_reference
49from .elements import _textual_label_reference
50from .elements import BindParameter
51from .elements import ClauseElement
52from .elements import ColumnClause
53from .elements import ColumnElement
54from .elements import Grouping
55from .elements import KeyedColumnElement
56from .elements import Label
57from .elements import NamedColumn
58from .elements import Null
59from .elements import UnaryExpression
60from .schema import Column
61from .selectable import Alias
62from .selectable import FromClause
63from .selectable import FromGrouping
64from .selectable import Join
65from .selectable import ScalarSelect
66from .selectable import SelectBase
67from .selectable import TableClause
68from .visitors import _ET
69from .. import exc
70from .. import util
71from ..util.typing import Literal
72from ..util.typing import Protocol
73
74if typing.TYPE_CHECKING:
75    from ._typing import _EquivalentColumnMap
76    from ._typing import _LimitOffsetType
77    from ._typing import _TypeEngineArgument
78    from .elements import BinaryExpression
79    from .elements import TextClause
80    from .selectable import _JoinTargetElement
81    from .selectable import _SelectIterable
82    from .selectable import Selectable
83    from .visitors import _TraverseCallableType
84    from .visitors import ExternallyTraversible
85    from .visitors import ExternalTraversal
86    from ..engine.interfaces import _AnyExecuteParams
87    from ..engine.interfaces import _AnyMultiExecuteParams
88    from ..engine.interfaces import _AnySingleExecuteParams
89    from ..engine.interfaces import _CoreSingleExecuteParams
90    from ..engine.row import Row
91
92_CE = TypeVar("_CE", bound="ColumnElement[Any]")
93
94
95def join_condition(
96    a: FromClause,
97    b: FromClause,
98    a_subset: Optional[FromClause] = None,
99    consider_as_foreign_keys: Optional[AbstractSet[ColumnClause[Any]]] = None,
100) -> ColumnElement[bool]:
101    """Create a join condition between two tables or selectables.
102
103    e.g.::
104
105        join_condition(tablea, tableb)
106
107    would produce an expression along the lines of::
108
109        tablea.c.id==tableb.c.tablea_id
110
111    The join is determined based on the foreign key relationships
112    between the two selectables.   If there are multiple ways
113    to join, or no way to join, an error is raised.
114
115    :param a_subset: An optional expression that is a sub-component
116        of ``a``.  An attempt will be made to join to just this sub-component
117        first before looking at the full ``a`` construct, and if found
118        will be successful even if there are other ways to join to ``a``.
119        This allows the "right side" of a join to be passed thereby
120        providing a "natural join".
121
122    """
123    return Join._join_condition(
124        a,
125        b,
126        a_subset=a_subset,
127        consider_as_foreign_keys=consider_as_foreign_keys,
128    )
129
130
131def find_join_source(
132    clauses: List[FromClause], join_to: FromClause
133) -> List[int]:
134    """Given a list of FROM clauses and a selectable,
135    return the first index and element from the list of
136    clauses which can be joined against the selectable.  returns
137    None, None if no match is found.
138
139    e.g.::
140
141        clause1 = table1.join(table2)
142        clause2 = table4.join(table5)
143
144        join_to = table2.join(table3)
145
146        find_join_source([clause1, clause2], join_to) == clause1
147
148    """
149
150    selectables = list(_from_objects(join_to))
151    idx = []
152    for i, f in enumerate(clauses):
153        for s in selectables:
154            if f.is_derived_from(s):
155                idx.append(i)
156    return idx
157
158
159def find_left_clause_that_matches_given(
160    clauses: Sequence[FromClause], join_from: FromClause
161) -> List[int]:
162    """Given a list of FROM clauses and a selectable,
163    return the indexes from the list of
164    clauses which is derived from the selectable.
165
166    """
167
168    selectables = list(_from_objects(join_from))
169    liberal_idx = []
170    for i, f in enumerate(clauses):
171        for s in selectables:
172            # basic check, if f is derived from s.
173            # this can be joins containing a table, or an aliased table
174            # or select statement matching to a table.  This check
175            # will match a table to a selectable that is adapted from
176            # that table.  With Query, this suits the case where a join
177            # is being made to an adapted entity
178            if f.is_derived_from(s):
179                liberal_idx.append(i)
180                break
181
182    # in an extremely small set of use cases, a join is being made where
183    # there are multiple FROM clauses where our target table is represented
184    # in more than one, such as embedded or similar.   in this case, do
185    # another pass where we try to get a more exact match where we aren't
186    # looking at adaption relationships.
187    if len(liberal_idx) > 1:
188        conservative_idx = []
189        for idx in liberal_idx:
190            f = clauses[idx]
191            for s in selectables:
192                if set(surface_selectables(f)).intersection(
193                    surface_selectables(s)
194                ):
195                    conservative_idx.append(idx)
196                    break
197        if conservative_idx:
198            return conservative_idx
199
200    return liberal_idx
201
202
203def find_left_clause_to_join_from(
204    clauses: Sequence[FromClause],
205    join_to: _JoinTargetElement,
206    onclause: Optional[ColumnElement[Any]],
207) -> List[int]:
208    """Given a list of FROM clauses, a selectable,
209    and optional ON clause, return a list of integer indexes from the
210    clauses list indicating the clauses that can be joined from.
211
212    The presence of an "onclause" indicates that at least one clause can
213    definitely be joined from; if the list of clauses is of length one
214    and the onclause is given, returns that index.   If the list of clauses
215    is more than length one, and the onclause is given, attempts to locate
216    which clauses contain the same columns.
217
218    """
219    idx = []
220    selectables = set(_from_objects(join_to))
221
222    # if we are given more than one target clause to join
223    # from, use the onclause to provide a more specific answer.
224    # otherwise, don't try to limit, after all, "ON TRUE" is a valid
225    # on clause
226    if len(clauses) > 1 and onclause is not None:
227        resolve_ambiguity = True
228        cols_in_onclause = _find_columns(onclause)
229    else:
230        resolve_ambiguity = False
231        cols_in_onclause = None
232
233    for i, f in enumerate(clauses):
234        for s in selectables.difference([f]):
235            if resolve_ambiguity:
236                assert cols_in_onclause is not None
237                if set(f.c).union(s.c).issuperset(cols_in_onclause):
238                    idx.append(i)
239                    break
240            elif onclause is not None or Join._can_join(f, s):
241                idx.append(i)
242                break
243
244    if len(idx) > 1:
245        # this is the same "hide froms" logic from
246        # Selectable._get_display_froms
247        toremove = set(
248            chain(*[_expand_cloned(f._hide_froms) for f in clauses])
249        )
250        idx = [i for i in idx if clauses[i] not in toremove]
251
252    # onclause was given and none of them resolved, so assume
253    # all indexes can match
254    if not idx and onclause is not None:
255        return list(range(len(clauses)))
256    else:
257        return idx
258
259
260def visit_binary_product(
261    fn: Callable[
262        [BinaryExpression[Any], ColumnElement[Any], ColumnElement[Any]], None
263    ],
264    expr: ColumnElement[Any],
265) -> None:
266    """Produce a traversal of the given expression, delivering
267    column comparisons to the given function.
268
269    The function is of the form::
270
271        def my_fn(binary, left, right)
272
273    For each binary expression located which has a
274    comparison operator, the product of "left" and
275    "right" will be delivered to that function,
276    in terms of that binary.
277
278    Hence an expression like::
279
280        and_(
281            (a + b) == q + func.sum(e + f),
282            j == r
283        )
284
285    would have the traversal::
286
287        a <eq> q
288        a <eq> e
289        a <eq> f
290        b <eq> q
291        b <eq> e
292        b <eq> f
293        j <eq> r
294
295    That is, every combination of "left" and
296    "right" that doesn't further contain
297    a binary comparison is passed as pairs.
298
299    """
300    stack: List[BinaryExpression[Any]] = []
301
302    def visit(element: ClauseElement) -> Iterator[ColumnElement[Any]]:
303        if isinstance(element, ScalarSelect):
304            # we don't want to dig into correlated subqueries,
305            # those are just column elements by themselves
306            yield element
307        elif element.__visit_name__ == "binary" and operators.is_comparison(
308            element.operator  # type: ignore
309        ):
310            stack.insert(0, element)  # type: ignore
311            for l in visit(element.left):  # type: ignore
312                for r in visit(element.right):  # type: ignore
313                    fn(stack[0], l, r)
314            stack.pop(0)
315            for elem in element.get_children():
316                visit(elem)
317        else:
318            if isinstance(element, ColumnClause):
319                yield element
320            for elem in element.get_children():
321                yield from visit(elem)
322
323    list(visit(expr))
324    visit = None  # type: ignore  # remove gc cycles
325
326
327def find_tables(
328    clause: ClauseElement,
329    *,
330    check_columns: bool = False,
331    include_aliases: bool = False,
332    include_joins: bool = False,
333    include_selects: bool = False,
334    include_crud: bool = False,
335) -> List[TableClause]:
336    """locate Table objects within the given expression."""
337
338    tables: List[TableClause] = []
339    _visitors: Dict[str, _TraverseCallableType[Any]] = {}
340
341    if include_selects:
342        _visitors["select"] = _visitors["compound_select"] = tables.append
343
344    if include_joins:
345        _visitors["join"] = tables.append
346
347    if include_aliases:
348        _visitors["alias"] = _visitors["subquery"] = _visitors[
349            "tablesample"
350        ] = _visitors["lateral"] = tables.append
351
352    if include_crud:
353        _visitors["insert"] = _visitors["update"] = _visitors["delete"] = (
354            lambda ent: tables.append(ent.table)
355        )
356
357    if check_columns:
358
359        def visit_column(column):
360            tables.append(column.table)
361
362        _visitors["column"] = visit_column
363
364    _visitors["table"] = tables.append
365
366    visitors.traverse(clause, {}, _visitors)
367    return tables
368
369
370def unwrap_order_by(clause: Any) -> Any:
371    """Break up an 'order by' expression into individual column-expressions,
372    without DESC/ASC/NULLS FIRST/NULLS LAST"""
373
374    cols = util.column_set()
375    result = []
376    stack = deque([clause])
377
378    # examples
379    # column -> ASC/DESC == column
380    # column -> ASC/DESC -> label == column
381    # column -> label -> ASC/DESC -> label == column
382    # scalar_select -> label -> ASC/DESC == scalar_select -> label
383
384    while stack:
385        t = stack.popleft()
386        if isinstance(t, ColumnElement) and (
387            not isinstance(t, UnaryExpression)
388            or not operators.is_ordering_modifier(t.modifier)  # type: ignore
389        ):
390            if isinstance(t, Label) and not isinstance(
391                t.element, ScalarSelect
392            ):
393                t = t.element
394
395                if isinstance(t, Grouping):
396                    t = t.element
397
398                stack.append(t)
399                continue
400            elif isinstance(t, _label_reference):
401                t = t.element
402
403                stack.append(t)
404                continue
405            if isinstance(t, (_textual_label_reference)):
406                continue
407            if t not in cols:
408                cols.add(t)
409                result.append(t)
410
411        else:
412            for c in t.get_children():
413                stack.append(c)
414    return result
415
416
417def unwrap_label_reference(element):
418    def replace(
419        element: ExternallyTraversible, **kw: Any
420    ) -> Optional[ExternallyTraversible]:
421        if isinstance(element, _label_reference):
422            return element.element
423        elif isinstance(element, _textual_label_reference):
424            assert False, "can't unwrap a textual label reference"
425        return None
426
427    return visitors.replacement_traverse(element, {}, replace)
428
429
430def expand_column_list_from_order_by(collist, order_by):
431    """Given the columns clause and ORDER BY of a selectable,
432    return a list of column expressions that can be added to the collist
433    corresponding to the ORDER BY, without repeating those already
434    in the collist.
435
436    """
437    cols_already_present = {
438        col.element if col._order_by_label_element is not None else col
439        for col in collist
440    }
441
442    to_look_for = list(chain(*[unwrap_order_by(o) for o in order_by]))
443
444    return [col for col in to_look_for if col not in cols_already_present]
445
446
447def clause_is_present(clause, search):
448    """Given a target clause and a second to search within, return True
449    if the target is plainly present in the search without any
450    subqueries or aliases involved.
451
452    Basically descends through Joins.
453
454    """
455
456    for elem in surface_selectables(search):
457        if clause == elem:  # use == here so that Annotated's compare
458            return True
459    else:
460        return False
461
462
463def tables_from_leftmost(clause: FromClause) -> Iterator[FromClause]:
464    if isinstance(clause, Join):
465        yield from tables_from_leftmost(clause.left)
466        yield from tables_from_leftmost(clause.right)
467    elif isinstance(clause, FromGrouping):
468        yield from tables_from_leftmost(clause.element)
469    else:
470        yield clause
471
472
473def surface_selectables(clause):
474    stack = [clause]
475    while stack:
476        elem = stack.pop()
477        yield elem
478        if isinstance(elem, Join):
479            stack.extend((elem.left, elem.right))
480        elif isinstance(elem, FromGrouping):
481            stack.append(elem.element)
482
483
484def surface_selectables_only(clause):
485    stack = [clause]
486    while stack:
487        elem = stack.pop()
488        if isinstance(elem, (TableClause, Alias)):
489            yield elem
490        if isinstance(elem, Join):
491            stack.extend((elem.left, elem.right))
492        elif isinstance(elem, FromGrouping):
493            stack.append(elem.element)
494        elif isinstance(elem, ColumnClause):
495            if elem.table is not None:
496                stack.append(elem.table)
497            else:
498                yield elem
499        elif elem is not None:
500            yield elem
501
502
503def extract_first_column_annotation(column, annotation_name):
504    filter_ = (FromGrouping, SelectBase)
505
506    stack = deque([column])
507    while stack:
508        elem = stack.popleft()
509        if annotation_name in elem._annotations:
510            return elem._annotations[annotation_name]
511        for sub in elem.get_children():
512            if isinstance(sub, filter_):
513                continue
514            stack.append(sub)
515    return None
516
517
518def selectables_overlap(left: FromClause, right: FromClause) -> bool:
519    """Return True if left/right have some overlapping selectable"""
520
521    return bool(
522        set(surface_selectables(left)).intersection(surface_selectables(right))
523    )
524
525
526def bind_values(clause):
527    """Return an ordered list of "bound" values in the given clause.
528
529    E.g.::
530
531        >>> expr = and_(
532        ...    table.c.foo==5, table.c.foo==7
533        ... )
534        >>> bind_values(expr)
535        [5, 7]
536    """
537
538    v = []
539
540    def visit_bindparam(bind):
541        v.append(bind.effective_value)
542
543    visitors.traverse(clause, {}, {"bindparam": visit_bindparam})
544    return v
545
546
547def _quote_ddl_expr(element):
548    if isinstance(element, str):
549        element = element.replace("'", "''")
550        return "'%s'" % element
551    else:
552        return repr(element)
553
554
555class _repr_base:
556    _LIST: int = 0
557    _TUPLE: int = 1
558    _DICT: int = 2
559
560    __slots__ = ("max_chars",)
561
562    max_chars: int
563
564    def trunc(self, value: Any) -> str:
565        rep = repr(value)
566        lenrep = len(rep)
567        if lenrep > self.max_chars:
568            segment_length = self.max_chars // 2
569            rep = (
570                rep[0:segment_length]
571                + (
572                    " ... (%d characters truncated) ... "
573                    % (lenrep - self.max_chars)
574                )
575                + rep[-segment_length:]
576            )
577        return rep
578
579
580def _repr_single_value(value):
581    rp = _repr_base()
582    rp.max_chars = 300
583    return rp.trunc(value)
584
585
586class _repr_row(_repr_base):
587    """Provide a string view of a row."""
588
589    __slots__ = ("row",)
590
591    def __init__(self, row: Row[Any], max_chars: int = 300):
592        self.row = row
593        self.max_chars = max_chars
594
595    def __repr__(self) -> str:
596        trunc = self.trunc
597        return "(%s%s)" % (
598            ", ".join(trunc(value) for value in self.row),
599            "," if len(self.row) == 1 else "",
600        )
601
602
603class _long_statement(str):
604    def __str__(self) -> str:
605        lself = len(self)
606        if lself > 500:
607            lleft = 250
608            lright = 100
609            trunc = lself - lleft - lright
610            return (
611                f"{self[0:lleft]} ... {trunc} "
612                f"characters truncated ... {self[-lright:]}"
613            )
614        else:
615            return str.__str__(self)
616
617
618class _repr_params(_repr_base):
619    """Provide a string view of bound parameters.
620
621    Truncates display to a given number of 'multi' parameter sets,
622    as well as long values to a given number of characters.
623
624    """
625
626    __slots__ = "params", "batches", "ismulti", "max_params"
627
628    def __init__(
629        self,
630        params: Optional[_AnyExecuteParams],
631        batches: int,
632        max_params: int = 100,
633        max_chars: int = 300,
634        ismulti: Optional[bool] = None,
635    ):
636        self.params = params
637        self.ismulti = ismulti
638        self.batches = batches
639        self.max_chars = max_chars
640        self.max_params = max_params
641
642    def __repr__(self) -> str:
643        if self.ismulti is None:
644            return self.trunc(self.params)
645
646        if isinstance(self.params, list):
647            typ = self._LIST
648
649        elif isinstance(self.params, tuple):
650            typ = self._TUPLE
651        elif isinstance(self.params, dict):
652            typ = self._DICT
653        else:
654            return self.trunc(self.params)
655
656        if self.ismulti:
657            multi_params = cast(
658                "_AnyMultiExecuteParams",
659                self.params,
660            )
661
662            if len(self.params) > self.batches:
663                msg = (
664                    " ... displaying %i of %i total bound parameter sets ... "
665                )
666                return " ".join(
667                    (
668                        self._repr_multi(
669                            multi_params[: self.batches - 2],
670                            typ,
671                        )[0:-1],
672                        msg % (self.batches, len(self.params)),
673                        self._repr_multi(multi_params[-2:], typ)[1:],
674                    )
675                )
676            else:
677                return self._repr_multi(multi_params, typ)
678        else:
679            return self._repr_params(
680                cast(
681                    "_AnySingleExecuteParams",
682                    self.params,
683                ),
684                typ,
685            )
686
687    def _repr_multi(
688        self,
689        multi_params: _AnyMultiExecuteParams,
690        typ: int,
691    ) -> str:
692        if multi_params:
693            if isinstance(multi_params[0], list):
694                elem_type = self._LIST
695            elif isinstance(multi_params[0], tuple):
696                elem_type = self._TUPLE
697            elif isinstance(multi_params[0], dict):
698                elem_type = self._DICT
699            else:
700                assert False, "Unknown parameter type %s" % (
701                    type(multi_params[0])
702                )
703
704            elements = ", ".join(
705                self._repr_params(params, elem_type) for params in multi_params
706            )
707        else:
708            elements = ""
709
710        if typ == self._LIST:
711            return "[%s]" % elements
712        else:
713            return "(%s)" % elements
714
715    def _get_batches(self, params: Iterable[Any]) -> Any:
716        lparams = list(params)
717        lenparams = len(lparams)
718        if lenparams > self.max_params:
719            lleft = self.max_params // 2
720            return (
721                lparams[0:lleft],
722                lparams[-lleft:],
723                lenparams - self.max_params,
724            )
725        else:
726            return lparams, None, None
727
728    def _repr_params(
729        self,
730        params: _AnySingleExecuteParams,
731        typ: int,
732    ) -> str:
733        if typ is self._DICT:
734            return self._repr_param_dict(
735                cast("_CoreSingleExecuteParams", params)
736            )
737        elif typ is self._TUPLE:
738            return self._repr_param_tuple(cast("Sequence[Any]", params))
739        else:
740            return self._repr_param_list(params)
741
742    def _repr_param_dict(self, params: _CoreSingleExecuteParams) -> str:
743        trunc = self.trunc
744        (
745            items_first_batch,
746            items_second_batch,
747            trunclen,
748        ) = self._get_batches(params.items())
749
750        if items_second_batch:
751            text = "{%s" % (
752                ", ".join(
753                    f"{key!r}: {trunc(value)}"
754                    for key, value in items_first_batch
755                )
756            )
757            text += f" ... {trunclen} parameters truncated ... "
758            text += "%s}" % (
759                ", ".join(
760                    f"{key!r}: {trunc(value)}"
761                    for key, value in items_second_batch
762                )
763            )
764        else:
765            text = "{%s}" % (
766                ", ".join(
767                    f"{key!r}: {trunc(value)}"
768                    for key, value in items_first_batch
769                )
770            )
771        return text
772
773    def _repr_param_tuple(self, params: Sequence[Any]) -> str:
774        trunc = self.trunc
775
776        (
777            items_first_batch,
778            items_second_batch,
779            trunclen,
780        ) = self._get_batches(params)
781
782        if items_second_batch:
783            text = "(%s" % (
784                ", ".join(trunc(value) for value in items_first_batch)
785            )
786            text += f" ... {trunclen} parameters truncated ... "
787            text += "%s)" % (
788                ", ".join(trunc(value) for value in items_second_batch),
789            )
790        else:
791            text = "(%s%s)" % (
792                ", ".join(trunc(value) for value in items_first_batch),
793                "," if len(items_first_batch) == 1 else "",
794            )
795        return text
796
797    def _repr_param_list(self, params: _AnySingleExecuteParams) -> str:
798        trunc = self.trunc
799        (
800            items_first_batch,
801            items_second_batch,
802            trunclen,
803        ) = self._get_batches(params)
804
805        if items_second_batch:
806            text = "[%s" % (
807                ", ".join(trunc(value) for value in items_first_batch)
808            )
809            text += f" ... {trunclen} parameters truncated ... "
810            text += "%s]" % (
811                ", ".join(trunc(value) for value in items_second_batch)
812            )
813        else:
814            text = "[%s]" % (
815                ", ".join(trunc(value) for value in items_first_batch)
816            )
817        return text
818
819
820def adapt_criterion_to_null(crit: _CE, nulls: Collection[Any]) -> _CE:
821    """given criterion containing bind params, convert selected elements
822    to IS NULL.
823
824    """
825
826    def visit_binary(binary):
827        if (
828            isinstance(binary.left, BindParameter)
829            and binary.left._identifying_key in nulls
830        ):
831            # reverse order if the NULL is on the left side
832            binary.left = binary.right
833            binary.right = Null()
834            binary.operator = operators.is_
835            binary.negate = operators.is_not
836        elif (
837            isinstance(binary.right, BindParameter)
838            and binary.right._identifying_key in nulls
839        ):
840            binary.right = Null()
841            binary.operator = operators.is_
842            binary.negate = operators.is_not
843
844    return visitors.cloned_traverse(crit, {}, {"binary": visit_binary})
845
846
847def splice_joins(
848    left: Optional[FromClause],
849    right: Optional[FromClause],
850    stop_on: Optional[FromClause] = None,
851) -> Optional[FromClause]:
852    if left is None:
853        return right
854
855    stack: List[Tuple[Optional[FromClause], Optional[Join]]] = [(right, None)]
856
857    adapter = ClauseAdapter(left)
858    ret = None
859    while stack:
860        (right, prevright) = stack.pop()
861        if isinstance(right, Join) and right is not stop_on:
862            right = right._clone()
863            right.onclause = adapter.traverse(right.onclause)
864            stack.append((right.left, right))
865        else:
866            right = adapter.traverse(right)
867        if prevright is not None:
868            assert right is not None
869            prevright.left = right
870        if ret is None:
871            ret = right
872
873    return ret
874
875
876@overload
877def reduce_columns(
878    columns: Iterable[ColumnElement[Any]],
879    *clauses: Optional[ClauseElement],
880    **kw: bool,
881) -> Sequence[ColumnElement[Any]]: ...
882
883
884@overload
885def reduce_columns(
886    columns: _SelectIterable,
887    *clauses: Optional[ClauseElement],
888    **kw: bool,
889) -> Sequence[Union[ColumnElement[Any], TextClause]]: ...
890
891
892def reduce_columns(
893    columns: _SelectIterable,
894    *clauses: Optional[ClauseElement],
895    **kw: bool,
896) -> Collection[Union[ColumnElement[Any], TextClause]]:
897    r"""given a list of columns, return a 'reduced' set based on natural
898    equivalents.
899
900    the set is reduced to the smallest list of columns which have no natural
901    equivalent present in the list.  A "natural equivalent" means that two
902    columns will ultimately represent the same value because they are related
903    by a foreign key.
904
905    \*clauses is an optional list of join clauses which will be traversed
906    to further identify columns that are "equivalent".
907
908    \**kw may specify 'ignore_nonexistent_tables' to ignore foreign keys
909    whose tables are not yet configured, or columns that aren't yet present.
910
911    This function is primarily used to determine the most minimal "primary
912    key" from a selectable, by reducing the set of primary key columns present
913    in the selectable to just those that are not repeated.
914
915    """
916    ignore_nonexistent_tables = kw.pop("ignore_nonexistent_tables", False)
917    only_synonyms = kw.pop("only_synonyms", False)
918
919    column_set = util.OrderedSet(columns)
920    cset_no_text: util.OrderedSet[ColumnElement[Any]] = column_set.difference(
921        c for c in column_set if is_text_clause(c)  # type: ignore
922    )
923
924    omit = util.column_set()
925    for col in cset_no_text:
926        for fk in chain(*[c.foreign_keys for c in col.proxy_set]):
927            for c in cset_no_text:
928                if c is col:
929                    continue
930                try:
931                    fk_col = fk.column
932                except exc.NoReferencedColumnError:
933                    # TODO: add specific coverage here
934                    # to test/sql/test_selectable ReduceTest
935                    if ignore_nonexistent_tables:
936                        continue
937                    else:
938                        raise
939                except exc.NoReferencedTableError:
940                    # TODO: add specific coverage here
941                    # to test/sql/test_selectable ReduceTest
942                    if ignore_nonexistent_tables:
943                        continue
944                    else:
945                        raise
946                if fk_col.shares_lineage(c) and (
947                    not only_synonyms or c.name == col.name
948                ):
949                    omit.add(col)
950                    break
951
952    if clauses:
953
954        def visit_binary(binary):
955            if binary.operator == operators.eq:
956                cols = util.column_set(
957                    chain(
958                        *[c.proxy_set for c in cset_no_text.difference(omit)]
959                    )
960                )
961                if binary.left in cols and binary.right in cols:
962                    for c in reversed(cset_no_text):
963                        if c.shares_lineage(binary.right) and (
964                            not only_synonyms or c.name == binary.left.name
965                        ):
966                            omit.add(c)
967                            break
968
969        for clause in clauses:
970            if clause is not None:
971                visitors.traverse(clause, {}, {"binary": visit_binary})
972
973    return column_set.difference(omit)
974
975
976def criterion_as_pairs(
977    expression,
978    consider_as_foreign_keys=None,
979    consider_as_referenced_keys=None,
980    any_operator=False,
981):
982    """traverse an expression and locate binary criterion pairs."""
983
984    if consider_as_foreign_keys and consider_as_referenced_keys:
985        raise exc.ArgumentError(
986            "Can only specify one of "
987            "'consider_as_foreign_keys' or "
988            "'consider_as_referenced_keys'"
989        )
990
991    def col_is(a, b):
992        # return a is b
993        return a.compare(b)
994
995    def visit_binary(binary):
996        if not any_operator and binary.operator is not operators.eq:
997            return
998        if not isinstance(binary.left, ColumnElement) or not isinstance(
999            binary.right, ColumnElement
1000        ):
1001            return
1002
1003        if consider_as_foreign_keys:
1004            if binary.left in consider_as_foreign_keys and (
1005                col_is(binary.right, binary.left)
1006                or binary.right not in consider_as_foreign_keys
1007            ):
1008                pairs.append((binary.right, binary.left))
1009            elif binary.right in consider_as_foreign_keys and (
1010                col_is(binary.left, binary.right)
1011                or binary.left not in consider_as_foreign_keys
1012            ):
1013                pairs.append((binary.left, binary.right))
1014        elif consider_as_referenced_keys:
1015            if binary.left in consider_as_referenced_keys and (
1016                col_is(binary.right, binary.left)
1017                or binary.right not in consider_as_referenced_keys
1018            ):
1019                pairs.append((binary.left, binary.right))
1020            elif binary.right in consider_as_referenced_keys and (
1021                col_is(binary.left, binary.right)
1022                or binary.left not in consider_as_referenced_keys
1023            ):
1024                pairs.append((binary.right, binary.left))
1025        else:
1026            if isinstance(binary.left, Column) and isinstance(
1027                binary.right, Column
1028            ):
1029                if binary.left.references(binary.right):
1030                    pairs.append((binary.right, binary.left))
1031                elif binary.right.references(binary.left):
1032                    pairs.append((binary.left, binary.right))
1033
1034    pairs: List[Tuple[ColumnElement[Any], ColumnElement[Any]]] = []
1035    visitors.traverse(expression, {}, {"binary": visit_binary})
1036    return pairs
1037
1038
1039class ClauseAdapter(visitors.ReplacingExternalTraversal):
1040    """Clones and modifies clauses based on column correspondence.
1041
1042    E.g.::
1043
1044      table1 = Table('sometable', metadata,
1045          Column('col1', Integer),
1046          Column('col2', Integer)
1047          )
1048      table2 = Table('someothertable', metadata,
1049          Column('col1', Integer),
1050          Column('col2', Integer)
1051          )
1052
1053      condition = table1.c.col1 == table2.c.col1
1054
1055    make an alias of table1::
1056
1057      s = table1.alias('foo')
1058
1059    calling ``ClauseAdapter(s).traverse(condition)`` converts
1060    condition to read::
1061
1062      s.c.col1 == table2.c.col1
1063
1064    """
1065
1066    __slots__ = (
1067        "__traverse_options__",
1068        "selectable",
1069        "include_fn",
1070        "exclude_fn",
1071        "equivalents",
1072        "adapt_on_names",
1073        "adapt_from_selectables",
1074    )
1075
1076    def __init__(
1077        self,
1078        selectable: Selectable,
1079        equivalents: Optional[_EquivalentColumnMap] = None,
1080        include_fn: Optional[Callable[[ClauseElement], bool]] = None,
1081        exclude_fn: Optional[Callable[[ClauseElement], bool]] = None,
1082        adapt_on_names: bool = False,
1083        anonymize_labels: bool = False,
1084        adapt_from_selectables: Optional[AbstractSet[FromClause]] = None,
1085    ):
1086        self.__traverse_options__ = {
1087            "stop_on": [selectable],
1088            "anonymize_labels": anonymize_labels,
1089        }
1090        self.selectable = selectable
1091        self.include_fn = include_fn
1092        self.exclude_fn = exclude_fn
1093        self.equivalents = util.column_dict(equivalents or {})
1094        self.adapt_on_names = adapt_on_names
1095        self.adapt_from_selectables = adapt_from_selectables
1096
1097    if TYPE_CHECKING:
1098
1099        @overload
1100        def traverse(self, obj: Literal[None]) -> None: ...
1101
1102        # note this specializes the ReplacingExternalTraversal.traverse()
1103        # method to state
1104        # that we will return the same kind of ExternalTraversal object as
1105        # we were given.  This is probably not 100% true, such as it's
1106        # possible for us to swap out Alias for Table at the top level.
1107        # Ideally there could be overloads specific to ColumnElement and
1108        # FromClause but Mypy is not accepting those as compatible with
1109        # the base ReplacingExternalTraversal
1110        @overload
1111        def traverse(self, obj: _ET) -> _ET: ...
1112
1113        def traverse(
1114            self, obj: Optional[ExternallyTraversible]
1115        ) -> Optional[ExternallyTraversible]: ...
1116
1117    def _corresponding_column(
1118        self, col, require_embedded, _seen=util.EMPTY_SET
1119    ):
1120        newcol = self.selectable.corresponding_column(
1121            col, require_embedded=require_embedded
1122        )
1123        if newcol is None and col in self.equivalents and col not in _seen:
1124            for equiv in self.equivalents[col]:
1125                newcol = self._corresponding_column(
1126                    equiv,
1127                    require_embedded=require_embedded,
1128                    _seen=_seen.union([col]),
1129                )
1130                if newcol is not None:
1131                    return newcol
1132
1133        if (
1134            self.adapt_on_names
1135            and newcol is None
1136            and isinstance(col, NamedColumn)
1137        ):
1138            newcol = self.selectable.exported_columns.get(col.name)
1139        return newcol
1140
1141    @util.preload_module("sqlalchemy.sql.functions")
1142    def replace(
1143        self, col: _ET, _include_singleton_constants: bool = False
1144    ) -> Optional[_ET]:
1145        functions = util.preloaded.sql_functions
1146
1147        # TODO: cython candidate
1148
1149        if self.include_fn and not self.include_fn(col):  # type: ignore
1150            return None
1151        elif self.exclude_fn and self.exclude_fn(col):  # type: ignore
1152            return None
1153
1154        if isinstance(col, FromClause) and not isinstance(
1155            col, functions.FunctionElement
1156        ):
1157            if self.selectable.is_derived_from(col):
1158                if self.adapt_from_selectables:
1159                    for adp in self.adapt_from_selectables:
1160                        if adp.is_derived_from(col):
1161                            break
1162                    else:
1163                        return None
1164                return self.selectable  # type: ignore
1165            elif isinstance(col, Alias) and isinstance(
1166                col.element, TableClause
1167            ):
1168                # we are a SELECT statement and not derived from an alias of a
1169                # table (which nonetheless may be a table our SELECT derives
1170                # from), so return the alias to prevent further traversal
1171                # or
1172                # we are an alias of a table and we are not derived from an
1173                # alias of a table (which nonetheless may be the same table
1174                # as ours) so, same thing
1175                return col  # type: ignore
1176            else:
1177                # other cases where we are a selectable and the element
1178                # is another join or selectable that contains a table which our
1179                # selectable derives from, that we want to process
1180                return None
1181
1182        elif not isinstance(col, ColumnElement):
1183            return None
1184        elif not _include_singleton_constants and col._is_singleton_constant:
1185            # dont swap out NULL, TRUE, FALSE for a label name
1186            # in a SQL statement that's being rewritten,
1187            # leave them as the constant.  This is first noted in #6259,
1188            # however the logic to check this moved here as of #7154 so that
1189            # it is made specific to SQL rewriting and not all column
1190            # correspondence
1191
1192            return None
1193
1194        if "adapt_column" in col._annotations:
1195            col = col._annotations["adapt_column"]
1196
1197        if TYPE_CHECKING:
1198            assert isinstance(col, KeyedColumnElement)
1199
1200        if self.adapt_from_selectables and col not in self.equivalents:

Showing the first 1,200 of 1487 lines. Download the file for the rest.

codekingpro/portable-devtools · Team Ai