FLANG
traverse.h
1//===-- include/flang/Evaluate/traverse.h -----------------------*- C++ -*-===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8
9#ifndef FORTRAN_EVALUATE_TRAVERSE_H_
10#define FORTRAN_EVALUATE_TRAVERSE_H_
11
12// A utility for scanning all of the constituent objects in an Expr<>
13// expression representation using a collection of mutually recursive
14// functions to compose a function object.
15//
16// The class template Traverse<> below implements a function object that
17// can handle every type that can appear in or around an Expr<>.
18// Each of its overloads for operator() should be viewed as a *default*
19// handler; some of these must be overridden by the client to accomplish
20// its particular task.
21//
22// The client (Visitor) of Traverse<Visitor,Result> must define:
23// - a member function "Result Default();"
24// - a member function "Result Combine(Result &&, Result &&)"
25// - overrides for "Result operator()"
26//
27// Boilerplate classes also appear below to ease construction of visitors.
28// See CheckSpecificationExpr() in check-expression.cpp for an example client.
29//
30// How this works:
31// - The operator() overloads in Traverse<> invoke the visitor's Default() for
32// expression leaf nodes. They invoke the visitor's operator() for the
33// subtrees of interior nodes, and the visitor's Combine() to merge their
34// results together.
35// - Overloads of operator() in each visitor handle the cases of interest.
36//
37// The default handler for semantics::Symbol will descend into the associated
38// expression of an ASSOCIATE (or related) construct entity.
39
40#include "expression.h"
41#include "flang/Common/indirection.h"
42#include "flang/Semantics/symbol.h"
43#include "flang/Semantics/type.h"
44
45namespace Fortran::evaluate {
46template <typename Visitor, typename Result,
47 bool TraverseAssocEntityDetails = true>
48class Traverse {
49public:
50 explicit Traverse(Visitor &v) : visitor_{v} {}
51
52 // Packaging
53 template <typename A, bool C>
54 Result operator()(const common::Indirection<A, C> &x) const {
55 return visitor_(x.value());
56 }
57 template <typename A>
58 Result operator()(const common::ForwardOwningPointer<A> &p) const {
59 return visitor_(p.get());
60 }
61 template <typename _> Result operator()(const SymbolRef x) const {
62 return visitor_(*x);
63 }
64 template <typename A> Result operator()(const std::unique_ptr<A> &x) const {
65 return visitor_(x.get());
66 }
67 template <typename A> Result operator()(const std::shared_ptr<A> &x) const {
68 return visitor_(x.get());
69 }
70 template <typename A> Result operator()(const A *x) const {
71 if (x) {
72 return visitor_(*x);
73 } else {
74 return visitor_.Default();
75 }
76 }
77 template <typename A> Result operator()(const std::optional<A> &x) const {
78 if (x) {
79 return visitor_(*x);
80 } else {
81 return visitor_.Default();
82 }
83 }
84 template <typename... As>
85 Result operator()(const std::variant<As...> &u) const {
86 return common::visit([=](const auto &y) { return visitor_(y); }, u);
87 }
88 template <typename A> Result operator()(const std::vector<A> &x) const {
89 return CombineContents(x);
90 }
91 template <typename A, typename B>
92 Result operator()(const std::pair<A, B> &x) const {
93 return Combine(x.first, x.second);
94 }
95
96 // Leaves
97 Result operator()(const BOZLiteralConstant &) const {
98 return visitor_.Default();
99 }
100 Result operator()(const NullPointer &) const { return visitor_.Default(); }
101 template <typename T> Result operator()(const Constant<T> &x) const {
102 if constexpr (T::category == TypeCategory::Derived) {
103 return visitor_.Combine(
104 visitor_(x.result().derivedTypeSpec()), CombineContents(x.values()));
105 } else {
106 return visitor_.Default();
107 }
108 }
109 Result operator()(const Symbol &symbol) const {
110 const Symbol &ultimate{symbol.GetUltimate()};
111 if constexpr (TraverseAssocEntityDetails) {
112 if (const auto *assoc{
113 ultimate.detailsIf<semantics::AssocEntityDetails>()}) {
114 return visitor_(assoc->expr());
115 }
116 }
117 return visitor_.Default();
118 }
119 Result operator()(const StaticDataObject &) const {
120 return visitor_.Default();
121 }
122 Result operator()(const ImpliedDoIndex &) const { return visitor_.Default(); }
123
124 // Variables
125 Result operator()(const BaseObject &x) const { return visitor_(x.u); }
126 Result operator()(const Component &x) const {
127 return Combine(x.base(), x.symbol());
128 }
129 Result operator()(const NamedEntity &x) const {
130 if (const Component * component{x.UnwrapComponent()}) {
131 return visitor_(*component);
132 } else {
133 return visitor_(DEREF(x.UnwrapSymbolRef()));
134 }
135 }
136 Result operator()(const TypeParamInquiry &x) const {
137 return visitor_(x.base());
138 }
139 Result operator()(const Triplet &x) const {
140 return Combine(x.GetLower(), x.GetUpper(), x.GetStride());
141 }
142 Result operator()(const Subscript &x) const { return visitor_(x.u); }
143 Result operator()(const ArrayRef &x) const {
144 return Combine(x.base(), x.subscript());
145 }
146 Result operator()(const CoarrayRef &x) const {
147 return Combine(x.base(), x.cosubscript(), x.notify(), x.stat(), x.team());
148 }
149 Result operator()(const DataRef &x) const { return visitor_(x.u); }
150 Result operator()(const Substring &x) const {
151 return Combine(x.parent(), x.GetLower(), x.GetUpper());
152 }
153 Result operator()(const ComplexPart &x) const {
154 return visitor_(x.complex());
155 }
156 template <typename T> Result operator()(const Designator<T> &x) const {
157 return visitor_(x.u);
158 }
159 Result operator()(const DescriptorInquiry &x) const {
160 return visitor_(x.base());
161 }
162 Result operator()(const RankOneBoundElement &x) const {
163 return visitor_(x.base());
164 }
165
166 // Calls
167 Result operator()(const SpecificIntrinsic &) const {
168 return visitor_.Default();
169 }
170 Result operator()(const ProcedureDesignator &x) const {
171 if (const Component * component{x.GetComponent()}) {
172 return visitor_(*component);
173 } else if (const Symbol * symbol{x.GetSymbol()}) {
174 return visitor_(*symbol);
175 } else {
176 return visitor_(DEREF(x.GetSpecificIntrinsic()));
177 }
178 }
179 Result operator()(const ActualArgument &x) const {
180 if (const auto *symbol{x.GetAssumedTypeDummy()}) {
181 return visitor_(*symbol);
182 }
183 if (const auto *condArg{x.GetConditionalArg()}) {
184 return TraverseConditionalArg(*condArg);
185 }
186 return visitor_(x.UnwrapExpr());
187 }
188 Result TraverseConditionalArg(
189 const ActualArgument::ConditionalArg &ca) const {
190 Result result{visitor_.Default()};
191 result = visitor_.Combine(std::move(result), visitor_(ca.condition()));
192 if (ca.consequent()) {
193 result = visitor_.Combine(
194 std::move(result), visitor_(ca.consequent()->value()));
195 }
196 return ca.VisitTail(
197 [&](const ActualArgument::ConditionalArg &inner) {
198 return visitor_.Combine(
199 std::move(result), TraverseConditionalArg(inner));
200 },
201 [&](const ActualArgument::ConditionalArg::Consequent &cons) -> Result {
202 if (cons) {
203 return visitor_.Combine(std::move(result), visitor_(cons->value()));
204 }
205 return result;
206 });
207 }
208 Result operator()(const ProcedureRef &x) const {
209 return Combine(x.proc(), x.arguments());
210 }
211 template <typename T> Result operator()(const FunctionRef<T> &x) const {
212 return visitor_(static_cast<const ProcedureRef &>(x));
213 }
214
215 // Other primaries
216 template <typename T>
217 Result operator()(const ArrayConstructorValue<T> &x) const {
218 return visitor_(x.u);
219 }
220 template <typename T>
221 Result operator()(const ArrayConstructorValues<T> &x) const {
222 return CombineContents(x);
223 }
224 template <typename T> Result operator()(const ImpliedDo<T> &x) const {
225 return Combine(x.lower(), x.upper(), x.stride(), x.values());
226 }
227 Result operator()(const semantics::ParamValue &x) const {
228 return visitor_(x.GetExplicit());
229 }
230 Result operator()(
231 const semantics::DerivedTypeSpec::ParameterMapType::value_type &x) const {
232 return visitor_(x.second);
233 }
234 Result operator()(
235 const semantics::DerivedTypeSpec::ParameterMapType &x) const {
236 return CombineContents(x);
237 }
238 Result operator()(const semantics::DerivedTypeSpec &x) const {
239 return Combine(x.originalTypeSymbol(), x.parameters());
240 }
241 Result operator()(const StructureConstructorValues::value_type &x) const {
242 return visitor_(x.second);
243 }
244 Result operator()(const StructureConstructorValues &x) const {
245 return CombineContents(x);
246 }
247 Result operator()(const StructureConstructor &x) const {
248 return visitor_.Combine(visitor_(x.derivedTypeSpec()), CombineContents(x));
249 }
250 // Conditional expressions (Fortran 2023)
251 template <typename T> Result operator()(const ConditionalExpr<T> &x) const {
252 return Combine(x.condition(), x.thenValue(), x.elseValue());
253 }
254
255 // Operations and wrappers
256 // Have a single operator() for all Operations.
257 template <typename D, typename R, typename... Os>
258 Result operator()(const Operation<D, R, Os...> &op) const {
259 if constexpr (sizeof...(Os) == 1) {
260 return visitor_(op.left());
261 } else {
262 return CombineOperands(op, std::index_sequence_for<Os...>{});
263 }
264 }
265 Result operator()(const Relational<SomeType> &x) const {
266 return visitor_(x.u);
267 }
268 template <typename T> Result operator()(const Expr<T> &x) const {
269 return visitor_(x.u);
270 }
271 Result operator()(const Assignment &x) const {
272 return Combine(x.lhs, x.rhs, x.u);
273 }
274 Result operator()(const Assignment::Intrinsic &) const {
275 return visitor_.Default();
276 }
277 Result operator()(const GenericExprWrapper &x) const { return visitor_(x.v); }
278 Result operator()(const GenericAssignmentWrapper &x) const {
279 return visitor_(x.v);
280 }
281
282private:
283 template <typename ITER> Result CombineRange(ITER iter, ITER end) const {
284 if (iter == end) {
285 return visitor_.Default();
286 } else {
287 Result result{visitor_(*iter)};
288 for (++iter; iter != end; ++iter) {
289 result = visitor_.Combine(std::move(result), visitor_(*iter));
290 }
291 return result;
292 }
293 }
294
295 template <typename A> Result CombineContents(const A &x) const {
296 return CombineRange(x.begin(), x.end());
297 }
298
299 template <typename D, typename R, typename... Os, size_t... Is>
300 Result CombineOperands(
301 const Operation<D, R, Os...> &op, std::index_sequence<Is...>) const {
302 static_assert(sizeof...(Os) > 1 && "Expecting multiple operands");
303 return Combine(op.template operand<Is>()...);
304 }
305
306 template <typename A, typename... Bs>
307 Result Combine(const A &x, const Bs &...ys) const {
308 if constexpr (sizeof...(Bs) == 0) {
309 return visitor_(x);
310 } else {
311 return visitor_.Combine(visitor_(x), Combine(ys...));
312 }
313 }
314
315 Visitor &visitor_;
316};
317
318// For validity checks across an expression: if any operator() result is
319// false, so is the overall result.
320template <typename Visitor, bool DefaultValue,
321 bool TraverseAssocEntityDetails = true,
323struct AllTraverse : public Base {
324 explicit AllTraverse(Visitor &v) : Base{v} {}
325 using Base::operator();
326 static bool Default() { return DefaultValue; }
327 static bool Combine(bool x, bool y) { return x && y; }
328};
329
330// For searches over an expression: the first operator() result that
331// is truthful is the final result. Works for Booleans, pointers,
332// and std::optional<>.
333template <typename Visitor, typename Result = bool,
334 bool TraverseAssocEntityDetails = true,
336class AnyTraverse : public Base {
337public:
338 explicit AnyTraverse(Visitor &v) : Base{v} {}
339 using Base::operator();
340 Result Default() const { return default_; }
341 static Result Combine(Result &&x, Result &&y) {
342 if (x) {
343 return std::move(x);
344 } else {
345 return std::move(y);
346 }
347 }
348
349private:
350 Result default_{};
351};
352
353template <typename Visitor, typename Set,
354 bool TraverseAssocEntityDetails = true,
356struct SetTraverse : public Base {
357 explicit SetTraverse(Visitor &v) : Base{v} {}
358 using Base::operator();
359 static Set Default() { return {}; }
360 static Set Combine(Set &&x, Set &&y) {
361#if defined __GNUC__ && !defined __APPLE__ && !(CLANG_LIBRARIES)
362 x.merge(y);
363#else
364 // std::set::merge() not available (yet)
365 for (auto &value : y) {
366 x.insert(std::move(value));
367 }
368#endif
369 return std::move(x);
370 }
371};
372
373} // namespace Fortran::evaluate
374#endif
Definition indirection.h:127
Definition indirection.h:31
Definition variable.h:205
Definition expression.h:923
Definition variable.h:243
Definition variable.h:357
Definition variable.h:73
Definition expression.h:394
Definition constant.h:147
Definition variable.h:413
Definition variable.h:381
Definition common.h:215
Definition call.h:394
Definition expression.h:444
Definition variable.h:101
Definition expression.h:113
Definition call.h:334
Definition expression.h:700
Definition static-data.h:29
Definition expression.h:784
Definition variable.h:304
Definition traverse.h:48
Definition variable.h:160
Definition variable.h:136
Definition type.h:95
Definition symbol.h:907
Definition call.h:34
Definition expression.h:472
Definition expression.h:928
Definition variable.h:50
Definition variable.h:288
Definition expression.h:941
Definition expression.h:436
Definition expression.h:872
Definition variable.h:191