Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
test_cte.py212 linesDownload Raw Back to suite
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 
codekingpro/portable-devtools · Team Ai