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