Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
test_cte.py238 linesDownload Raw Back to suite
1# testing/suite/test_cte.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
8
9from .. import fixtures
10from ..assertions import eq_
11from ..schema import Column
12from ..schema import Table
13from ... import column
14from ... import ForeignKey
15from ... import Integer
16from ... import select
17from ... import String
18from ... import testing
19from ... import values
20
21
22class CTETest(fixtures.TablesTest):
23    __sparse_driver_backend__ = True
24    __requires__ = ("ctes",)
25
26    run_inserts = "each"
27    run_deletes = "each"
28
29    @classmethod
30    def define_tables(cls, metadata):
31        Table(
32            "some_table",
33            metadata,
34            Column("id", Integer, primary_key=True),
35            Column("data", String(50)),
36            Column("parent_id", ForeignKey("some_table.id")),
37        )
38
39        Table(
40            "some_other_table",
41            metadata,
42            Column("id", Integer, primary_key=True),
43            Column("data", String(50)),
44            Column("parent_id", Integer),
45        )
46
47    @classmethod
48    def insert_data(cls, connection):
49        connection.execute(
50            cls.tables.some_table.insert(),
51            [
52                {"id": 1, "data": "d1", "parent_id": None},
53                {"id": 2, "data": "d2", "parent_id": 1},
54                {"id": 3, "data": "d3", "parent_id": 1},
55                {"id": 4, "data": "d4", "parent_id": 3},
56                {"id": 5, "data": "d5", "parent_id": 3},
57            ],
58        )
59
60    def test_select_nonrecursive_round_trip(self, connection):
61        some_table = self.tables.some_table
62
63        cte = (
64            select(some_table)
65            .where(some_table.c.data.in_(["d2", "d3", "d4"]))
66            .cte("some_cte")
67        )
68        result = connection.execute(
69            select(cte.c.data).where(cte.c.data.in_(["d4", "d5"]))
70        )
71        eq_(result.fetchall(), [("d4",)])
72
73    def test_select_recursive_round_trip(self, connection):
74        some_table = self.tables.some_table
75
76        cte = (
77            select(some_table)
78            .where(some_table.c.data.in_(["d2", "d3", "d4"]))
79            .cte("some_cte", recursive=True)
80        )
81
82        cte_alias = cte.alias("c1")
83        st1 = some_table.alias()
84        # note that SQL Server requires this to be UNION ALL,
85        # can't be UNION
86        cte = cte.union_all(
87            select(st1).where(st1.c.id == cte_alias.c.parent_id)
88        )
89        result = connection.execute(
90            select(cte.c.data)
91            .where(cte.c.data != "d2")
92            .order_by(cte.c.data.desc())
93        )
94        eq_(
95            result.fetchall(),
96            [("d4",), ("d3",), ("d3",), ("d1",), ("d1",), ("d1",)],
97        )
98
99    def test_insert_from_select_round_trip(self, connection):
100        some_table = self.tables.some_table
101        some_other_table = self.tables.some_other_table
102
103        cte = (
104            select(some_table)
105            .where(some_table.c.data.in_(["d2", "d3", "d4"]))
106            .cte("some_cte")
107        )
108        connection.execute(
109            some_other_table.insert().from_select(
110                ["id", "data", "parent_id"], select(cte)
111            )
112        )
113        eq_(
114            connection.execute(
115                select(some_other_table).order_by(some_other_table.c.id)
116            ).fetchall(),
117            [(2, "d2", 1), (3, "d3", 1), (4, "d4", 3)],
118        )
119
120    @testing.requires.ctes_with_update_delete
121    @testing.requires.update_from
122    def test_update_from_round_trip(self, connection):
123        some_table = self.tables.some_table
124        some_other_table = self.tables.some_other_table
125
126        connection.execute(
127            some_other_table.insert().from_select(
128                ["id", "data", "parent_id"], select(some_table)
129            )
130        )
131
132        cte = (
133            select(some_table)
134            .where(some_table.c.data.in_(["d2", "d3", "d4"]))
135            .cte("some_cte")
136        )
137        connection.execute(
138            some_other_table.update()
139            .values(parent_id=5)
140            .where(some_other_table.c.data == cte.c.data)
141        )
142        eq_(
143            connection.execute(
144                select(some_other_table).order_by(some_other_table.c.id)
145            ).fetchall(),
146            [
147                (1, "d1", None),
148                (2, "d2", 5),
149                (3, "d3", 5),
150                (4, "d4", 5),
151                (5, "d5", 3),
152            ],
153        )
154
155    @testing.requires.ctes_with_update_delete
156    @testing.requires.delete_from
157    def test_delete_from_round_trip(self, connection):
158        some_table = self.tables.some_table
159        some_other_table = self.tables.some_other_table
160
161        connection.execute(
162            some_other_table.insert().from_select(
163                ["id", "data", "parent_id"], select(some_table)
164            )
165        )
166
167        cte = (
168            select(some_table)
169            .where(some_table.c.data.in_(["d2", "d3", "d4"]))
170            .cte("some_cte")
171        )
172        connection.execute(
173            some_other_table.delete().where(
174                some_other_table.c.data == cte.c.data
175            )
176        )
177        eq_(
178            connection.execute(
179                select(some_other_table).order_by(some_other_table.c.id)
180            ).fetchall(),
181            [(1, "d1", None), (5, "d5", 3)],
182        )
183
184    @testing.requires.ctes_with_update_delete
185    def test_delete_scalar_subq_round_trip(self, connection):
186        some_table = self.tables.some_table
187        some_other_table = self.tables.some_other_table
188
189        connection.execute(
190            some_other_table.insert().from_select(
191                ["id", "data", "parent_id"], select(some_table)
192            )
193        )
194
195        cte = (
196            select(some_table)
197            .where(some_table.c.data.in_(["d2", "d3", "d4"]))
198            .cte("some_cte")
199        )
200        connection.execute(
201            some_other_table.delete().where(
202                some_other_table.c.data
203                == select(cte.c.data)
204                .where(cte.c.id == some_other_table.c.id)
205                .scalar_subquery()
206            )
207        )
208        eq_(
209            connection.execute(
210                select(some_other_table).order_by(some_other_table.c.id)
211            ).fetchall(),
212            [(1, "d1", None), (5, "d5", 3)],
213        )
214
215    @testing.variation("values_named", [True, False])
216    @testing.variation("cte_named", [True, False])
217    @testing.variation("literal_binds", [True, False])
218    @testing.requires.ctes_with_values
219    def test_values_named_via_cte(
220        self, connection, values_named, cte_named, literal_binds
221    ):
222
223        cte1 = (
224            values(
225                column("col1", String),
226                column("col2", Integer),
227                literal_binds=bool(literal_binds),
228                name="some name" if values_named else None,
229            )
230            .data([("a", 2), ("b", 3)])
231            .cte("cte1" if cte_named else None)
232        )
233
234        stmt = select(cte1)
235
236        rows = connection.execute(stmt).all()
237        eq_(rows, [("a", 2), ("b", 3)])
238