Line data Source code
1 : // Copyright (C) 2025-2026 Free Software Foundation, Inc.
2 :
3 : // This file is part of GCC.
4 :
5 : // GCC is free software; you can redistribute it and/or modify it under
6 : // the terms of the GNU General Public License as published by the Free
7 : // Software Foundation; either version 3, or (at your option) any later
8 : // version.
9 :
10 : // GCC is distributed in the hope that it will be useful, but WITHOUT ANY
11 : // WARRANTY; without even the implied warranty of MERCHANTABILITY or
12 : // FITNESS FOR A PARTICULAR PURPOSE. See the GNU General Public License
13 : // for more details.
14 :
15 : // You should have received a copy of the GNU General Public License
16 : // along with GCC; see the file COPYING3. If not see
17 : // <http://www.gnu.org/licenses/>.
18 :
19 : #include "rust-derive-ord.h"
20 : #include "rust-ast.h"
21 : #include "rust-derive-cmp-common.h"
22 : #include "rust-derive.h"
23 : #include "rust-item.h"
24 : #include "rust-system.h"
25 :
26 : namespace Rust {
27 : namespace AST {
28 :
29 145 : DeriveOrd::DeriveOrd (Ordering ordering, location_t loc,
30 : Builder::Source item_source)
31 145 : : DeriveVisitor (loc, item_source), ordering (ordering)
32 145 : {}
33 :
34 : std::unique_ptr<Item>
35 145 : DeriveOrd::go (Item &item)
36 : {
37 145 : item.accept_vis (*this);
38 :
39 145 : return std::move (expanded);
40 : }
41 :
42 : std::unique_ptr<Expr>
43 289 : DeriveOrd::cmp_call (std::unique_ptr<Expr> &&self_expr,
44 : std::unique_ptr<Expr> &&other_expr)
45 : {
46 1734 : auto cmp_fn_path = builder.path_in_expression (
47 578 : {builder.get_path_start (), "cmp", trait (ordering), fn (ordering)}, true);
48 :
49 578 : return builder.call (ptrify (cmp_fn_path),
50 578 : vec (builder.ref (std::move (self_expr)),
51 867 : builder.ref (std::move (other_expr))));
52 289 : }
53 :
54 : std::unique_ptr<Item>
55 145 : DeriveOrd::cmp_impl (
56 : std::unique_ptr<BlockExpr> &&fn_block, Identifier type_name,
57 : const std::vector<std::unique_ptr<GenericParam>> &type_generics)
58 : {
59 145 : auto fn = cmp_fn (std::move (fn_block), type_name);
60 :
61 145 : auto trait = ordering == Ordering::Partial ? "PartialOrd" : "Ord";
62 324 : auto trait_path = [&, this] () {
63 716 : return builder.type_path ({builder.get_path_start (), "cmp", trait}, true);
64 145 : };
65 :
66 145 : auto trait_bound
67 34 : = [&, this] () { return builder.trait_bound (trait_path ()); };
68 :
69 145 : auto trait_items = vec (std::move (fn));
70 :
71 145 : auto cmp_generics
72 145 : = setup_impl_generics (type_name.as_string (), type_generics, trait_bound);
73 :
74 290 : return builder.trait_impl (trait_path (), std::move (cmp_generics.self_type),
75 : std::move (trait_items),
76 290 : std::move (cmp_generics.impl));
77 145 : }
78 :
79 : std::unique_ptr<AssociatedItem>
80 145 : DeriveOrd::cmp_fn (std::unique_ptr<BlockExpr> &&block, Identifier type_name)
81 : {
82 : // Ordering
83 145 : auto return_type
84 580 : = builder.type_path ({builder.get_path_start (), "cmp", "Ordering"}, true);
85 :
86 : // In the case of PartialOrd, we return an Option<Ordering>
87 145 : if (ordering == Ordering::Partial)
88 : {
89 73 : auto generic = GenericArg::create_type (ptrify (return_type));
90 :
91 73 : auto generic_seg = builder.type_path_segment_generic (
92 146 : "Option", GenericArgs ({}, {generic}, {}, loc));
93 73 : auto core = builder.type_path_segment (builder.get_path_start ());
94 73 : auto option = builder.type_path_segment ("option");
95 :
96 73 : return_type
97 146 : = builder.type_path (vec (std::move (core), std::move (option),
98 : std::move (generic_seg)),
99 73 : true);
100 73 : }
101 :
102 : // &self, other: &Self
103 : //
104 : // this must be Self for a generic type Wrapping<T> the bare struct name has
105 : // no type arguments
106 145 : auto params
107 290 : = vec (builder.self_ref_param (),
108 290 : builder.function_param (builder.identifier_pattern ("other"),
109 290 : builder.reference_type (
110 435 : ptrify (builder.type_path ("Self")))));
111 :
112 145 : auto function_name = fn (ordering);
113 :
114 580 : return builder.function (function_name, std::move (params),
115 435 : ptrify (return_type), std::move (block));
116 145 : }
117 :
118 : std::unique_ptr<Pattern>
119 98 : DeriveOrd::make_equal ()
120 : {
121 490 : std::unique_ptr<Pattern> equal = ptrify (builder.path_in_expression (
122 98 : {builder.get_path_start (), "cmp", "Ordering", "Equal"}, true));
123 :
124 : // We need to wrap the pattern in Option::Some if we are doing partial
125 : // ordering
126 98 : if (ordering == Ordering::Partial)
127 : {
128 61 : auto pattern_items = std::unique_ptr<TupleStructItems> (
129 61 : new TupleStructItemsNoRest (vec (std::move (equal))));
130 :
131 61 : equal
132 122 : = std::make_unique<TupleStructPattern> (builder.path_in_expression (
133 : LangItem::Kind::OPTION_SOME),
134 61 : std::move (pattern_items));
135 61 : }
136 :
137 98 : return equal;
138 : }
139 :
140 : std::pair<MatchArm, MatchArm>
141 98 : DeriveOrd::make_cmp_arms ()
142 : {
143 : // All comparison results other than Ordering::Equal
144 98 : auto non_equal = builder.identifier_pattern (DeriveOrd::not_equal);
145 98 : auto equal = make_equal ();
146 :
147 98 : return {builder.match_arm (std::move (equal)),
148 98 : builder.match_arm (std::move (non_equal))};
149 98 : }
150 :
151 : std::unique_ptr<Expr>
152 161 : DeriveOrd::recursive_match (std::vector<SelfOther> &&members)
153 : {
154 161 : if (members.empty ())
155 : {
156 48 : std::unique_ptr<Expr> value = ptrify (builder.path_in_expression (
157 8 : {builder.get_path_start (), "cmp", "Ordering", "Equal"}, true));
158 :
159 8 : if (ordering == Ordering::Partial)
160 12 : value = builder.call (ptrify (builder.path_in_expression (
161 : LangItem::Kind::OPTION_SOME)),
162 4 : std::move (value));
163 :
164 : return value;
165 : }
166 :
167 153 : std::unique_ptr<Expr> final_expr = nullptr;
168 :
169 404 : for (auto it = members.rbegin (); it != members.rend (); it++)
170 : {
171 251 : auto &member = *it;
172 :
173 251 : auto call = cmp_call (std::move (member.self_expr),
174 251 : std::move (member.other_expr));
175 :
176 : // For the last member (so the first iterator), we just create a call
177 : // expression
178 251 : if (it == members.rbegin ())
179 : {
180 153 : final_expr = std::move (call);
181 153 : continue;
182 : }
183 :
184 : // If we aren't dealing with the last member, then we need to wrap all of
185 : // that in a big match expression and keep going
186 98 : auto match_arms = make_cmp_arms ();
187 :
188 98 : auto match_cases
189 : = {builder.match_case (std::move (match_arms.first),
190 : std::move (final_expr)),
191 : builder.match_case (std::move (match_arms.second),
192 490 : builder.identifier (DeriveOrd::not_equal))};
193 :
194 98 : final_expr = builder.match (std::move (call), std::move (match_cases));
195 545 : }
196 :
197 153 : return final_expr;
198 153 : }
199 :
200 : // we need to do a recursive match expression for all of the fields used in a
201 : // struct so for something like struct Foo { a: i32, b: i32, c: i32 } we must
202 : // first compare each `a` field, then `b`, then `c`, like this:
203 : //
204 : // match cmp_fn(self.<field>, other.<field>) {
205 : // Ordering::Equal => <recurse>,
206 : // cmp => cmp,
207 : // }
208 : //
209 : // and the recurse will be the exact same expression, on the next field. so that
210 : // our result looks like this:
211 : //
212 : // match cmp_fn(self.a, other.a) {
213 : // Ordering::Equal => match cmp_fn(self.b, other.b) {
214 : // Ordering::Equal =>cmp_fn(self.c, other.c),
215 : // cmp => cmp,
216 : // }
217 : // cmp => cmp,
218 : // }
219 : //
220 : // the last field comparison needs not to be a match but just the function call.
221 : // this is going to be annoying lol
222 : void
223 79 : DeriveOrd::visit_struct (StructStruct &item)
224 : {
225 79 : auto fields = SelfOther::fields (builder, item.get_fields ());
226 :
227 79 : auto match_expr = recursive_match (std::move (fields));
228 :
229 158 : expanded = cmp_impl (builder.block (std::move (match_expr)),
230 158 : item.get_identifier (), item.get_generic_params ());
231 79 : }
232 :
233 : // same as structs, but for each field index instead of each field name -
234 : // straightforward once we have `visit_struct` working
235 : void
236 28 : DeriveOrd::visit_tuple (TupleStruct &item)
237 : {
238 28 : auto fields = SelfOther::indexes (builder, item.get_fields ());
239 :
240 28 : auto match_expr = recursive_match (std::move (fields));
241 :
242 56 : expanded = cmp_impl (builder.block (std::move (match_expr)),
243 56 : item.get_identifier (), item.get_generic_params ());
244 28 : }
245 :
246 : // for enums, we need to generate a match for each of the enum's variant that
247 : // contains data and then do the same thing as visit_struct or visit_enum. if
248 : // the two aren't the same variant, then compare the two discriminant values for
249 : // all the dataless enum variants and in the general case.
250 : //
251 : // so for enum Foo { A(i32, i32), B, C } we need to do the following
252 : //
253 : // match (self, other) {
254 : // (A(self_0, self_1), A(other_0, other_1)) => {
255 : // match cmp_fn(self_0, other_0) {
256 : // Ordering::Equal => cmp_fn(self_1, other_1),
257 : // cmp => cmp,
258 : // },
259 : // _ => cmp_fn(discr_value(self), discr_value(other))
260 : // }
261 : void
262 38 : DeriveOrd::visit_enum (Enum &item)
263 : {
264 : // NOTE: We can factor this even further with DerivePartialEq, but this is
265 : // getting out of scope for this PR surely
266 :
267 38 : auto cases = std::vector<MatchCase> ();
268 76 : auto type_name = item.get_identifier ().as_string ();
269 :
270 38 : auto let_sd = builder.discriminant_value (DeriveOrd::self_discr, "self");
271 38 : auto let_od = builder.discriminant_value (DeriveOrd::other_discr, "other");
272 :
273 76 : auto discr_cmp = cmp_call (builder.identifier (DeriveOrd::self_discr),
274 114 : builder.identifier (DeriveOrd::other_discr));
275 :
276 92 : auto recursive_match_fn = [this] (std::vector<SelfOther> &&fields) {
277 54 : return recursive_match (std::move (fields));
278 38 : };
279 :
280 122 : for (auto &variant : item.get_variants ())
281 : {
282 84 : auto enum_builder
283 168 : = EnumMatchBuilder (type_name, variant->get_identifier ().as_string (),
284 84 : recursive_match_fn, builder);
285 :
286 84 : switch (variant->get_enum_item_kind ())
287 : {
288 8 : case EnumItem::Kind::Struct:
289 8 : cases.emplace_back (enum_builder.strukt (*variant));
290 8 : break;
291 46 : case EnumItem::Kind::Tuple:
292 46 : cases.emplace_back (enum_builder.tuple (*variant));
293 46 : break;
294 : case EnumItem::Kind::Identifier:
295 : case EnumItem::Kind::Discriminant:
296 : // We don't need to do anything for these, as they are handled by the
297 : // discriminant value comparison
298 : break;
299 : }
300 84 : }
301 :
302 : // Add the last case which compares the discriminant values in case `self` and
303 : // `other` are actually different variants of the enum
304 38 : cases.emplace_back (
305 76 : builder.match_case (builder.wildcard (), std::move (discr_cmp)));
306 :
307 38 : auto match
308 76 : = builder.match (builder.tuple (vec (builder.identifier ("self"),
309 76 : builder.identifier ("other"))),
310 38 : std::move (cases));
311 :
312 38 : expanded
313 76 : = cmp_impl (builder.block (vec (std::move (let_sd), std::move (let_od)),
314 : std::move (match)),
315 190 : type_name, item.get_generic_params ());
316 38 : }
317 :
318 : void
319 0 : DeriveOrd::visit_union (Union &item)
320 : {
321 0 : auto trait_name = trait (ordering);
322 :
323 0 : rust_error_at (item.get_locus (), "derive(%s) cannot be used on unions",
324 : trait_name.c_str ());
325 0 : }
326 :
327 : } // namespace AST
328 : } // namespace Rust
|