codekingpro/portable-devtools
115k
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 