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
163 // Calls
164 Result operator()(const SpecificIntrinsic &) const {
165 return visitor_.Default();
166 }
167 Result operator()(const ProcedureDesignator &x) const {
168 if (const Component * component{x.GetComponent()}) {
169 return visitor_(*component);
170 } else if (const Symbol * symbol{x.GetSymbol()}) {
171 return visitor_(*symbol);
172 } else {
173 return visitor_(DEREF(x.GetSpecificIntrinsic()));
174 }
175 }
176 Result operator()(const ActualArgument &x) const {
177 if (const auto *symbol{x.GetAssumedTypeDummy()}) {
178 return visitor_(*symbol);
179 }
180 if (const auto *condArg{x.GetConditionalArg()}) {
181 return TraverseConditionalArg(*condArg);
182 }
183 return visitor_(x.UnwrapExpr());
184 }
185 Result TraverseConditionalArg(
186 const ActualArgument::ConditionalArg &ca) const {
187 Result result{visitor_.Default()};
188 result = visitor_.Combine(std::move(result), visitor_(ca.condition()));
189 if (ca.consequent()) {
190 result = visitor_.Combine(
191 std::move(result), visitor_(ca.consequent()->value()));
192 }
193 return ca.VisitTail(
194 [&](const ActualArgument::ConditionalArg &inner) {
195 return visitor_.Combine(
196 std::move(result), TraverseConditionalArg(inner));
197 },
198 [&](const ActualArgument::ConditionalArg::Consequent &cons) -> Result {
199 if (cons) {
200 return visitor_.Combine(std::move(result), visitor_(cons->value()));
201 }
202 return result;
203 });
204 }
205 Result operator()(const ProcedureRef &x) const {
206 return Combine(x.proc(), x.arguments());
207 }
208 template <typename T> Result operator()(const FunctionRef<T> &x) const {
209 return visitor_(static_cast<const ProcedureRef &>(x));
210 }
211
212 // Other primaries
213 template <typename T>
214 Result operator()(const ArrayConstructorValue<T> &x) const {
215 return visitor_(x.u);
216 }
217 template <typename T>
218 Result operator()(const ArrayConstructorValues<T> &x) const {
219 return CombineContents(x);
220 }
221 template <typename T> Result operator()(const ImpliedDo<T> &x) const {
222 return Combine(x.lower(), x.upper(), x.stride(), x.values());
223 }
224 Result operator()(const semantics::ParamValue &x) const {
225 return visitor_(x.GetExplicit());
226 }
227 Result operator()(
228 const semantics::DerivedTypeSpec::ParameterMapType::value_type &x) const {
229 return visitor_(x.second);
230 }
231 Result operator()(
232 const semantics::DerivedTypeSpec::ParameterMapType &x) const {
233 return CombineContents(x);
234 }
235 Result operator()(const semantics::DerivedTypeSpec &x) const {
236 return Combine(x.originalTypeSymbol(), x.parameters());
237 }
238 Result operator()(const StructureConstructorValues::value_type &x) const {
239 return visitor_(x.second);
240 }
241 Result operator()(const StructureConstructorValues &x) const {
242 return CombineContents(x);
243 }
244 Result operator()(const StructureConstructor &x) const {
245 return visitor_.Combine(visitor_(x.derivedTypeSpec()), CombineContents(x));
246 }
247 // Conditional expressions (Fortran 2023)
248 template <typename T> Result operator()(const ConditionalExpr<T> &x) const {
249 return Combine(x.condition(), x.thenValue(), x.elseValue());
250 }
251
252 // Operations and wrappers
253 // Have a single operator() for all Operations.
254 template <typename D, typename R, typename... Os>
255 Result operator()(const Operation<D, R, Os...> &op) const {
256 if constexpr (sizeof...(Os) == 1) {
257 return visitor_(op.left());
258 } else {
259 return CombineOperands(op, std::index_sequence_for<Os...>{});
260 }
261 }
262 Result operator()(const Relational<SomeType> &x) const {
263 return visitor_(x.u);
264 }
265 template <typename T> Result operator()(const Expr<T> &x) const {
266 return visitor_(x.u);
267 }
268 Result operator()(const Assignment &x) const {
269 return Combine(x.lhs, x.rhs, x.u);
270 }
271 Result operator()(const Assignment::Intrinsic &) const {
272 return visitor_.Default();
273 }
274 Result operator()(const GenericExprWrapper &x) const { return visitor_(x.v); }
275 Result operator()(const GenericAssignmentWrapper &x) const {
276 return visitor_(x.v);
277 }
278
279private:
280 template <typename ITER> Result CombineRange(ITER iter, ITER end) const {
281 if (iter == end) {
282 return visitor_.Default();
283 } else {
284 Result result{visitor_(*iter)};
285 for (++iter; iter != end; ++iter) {
286 result = visitor_.Combine(std::move(result), visitor_(*iter));
287 }
288 return result;
289 }
290 }
291
292 template <typename A> Result CombineContents(const A &x) const {
293 return CombineRange(x.begin(), x.end());
294 }
295
296 template <typename D, typename R, typename... Os, size_t... Is>
297 Result CombineOperands(
298 const Operation<D, R, Os...> &op, std::index_sequence<Is...>) const {
299 static_assert(sizeof...(Os) > 1 && "Expecting multiple operands");
300 return Combine(op.template operand<Is>()...);
301 }
302
303 template <typename A, typename... Bs>
304 Result Combine(const A &x, const Bs &...ys) const {
305 if constexpr (sizeof...(Bs) == 0) {
306 return visitor_(x);
307 } else {
308 return visitor_.Combine(visitor_(x), Combine(ys...));
309 }
310 }
311
312 Visitor &visitor_;
313};
314
315// For validity checks across an expression: if any operator() result is
316// false, so is the overall result.
317template <typename Visitor, bool DefaultValue,
318 bool TraverseAssocEntityDetails = true,
320struct AllTraverse : public Base {
321 explicit AllTraverse(Visitor &v) : Base{v} {}
322 using Base::operator();
323 static bool Default() { return DefaultValue; }
324 static bool Combine(bool x, bool y) { return x && y; }
325};
326
327// For searches over an expression: the first operator() result that
328// is truthful is the final result. Works for Booleans, pointers,
329// and std::optional<>.
330template <typename Visitor, typename Result = bool,
331 bool TraverseAssocEntityDetails = true,
333class AnyTraverse : public Base {
334public:
335 explicit AnyTraverse(Visitor &v) : Base{v} {}
336 using Base::operator();
337 Result Default() const { return default_; }
338 static Result Combine(Result &&x, Result &&y) {
339 if (x) {
340 return std::move(x);
341 } else {
342 return std::move(y);
343 }
344 }
345
346private:
347 Result default_{};
348};
349
350template <typename Visitor, typename Set,
351 bool TraverseAssocEntityDetails = true,
353struct SetTraverse : public Base {
354 explicit SetTraverse(Visitor &v) : Base{v} {}
355 using Base::operator();
356 static Set Default() { return {}; }
357 static Set Combine(Set &&x, Set &&y) {
358#if defined __GNUC__ && !defined __APPLE__ && !(CLANG_LIBRARIES)
359 x.merge(y);
360#else
361 // std::set::merge() not available (yet)
362 for (auto &value : y) {
363 x.insert(std::move(value));
364 }
365#endif
366 return std::move(x);
367 }
368};
369
370} // namespace Fortran::evaluate
371#endif
Definition indirection.h:127
Definition indirection.h:31
Definition variable.h:205
Definition expression.h:920
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:697
Definition static-data.h:29
Definition expression.h:781
Definition variable.h:304
Definition traverse.h:48
Definition variable.h:160
Definition variable.h:136
Definition type.h:95
Definition symbol.h:896
Definition call.h:34
Definition expression.h:472
Definition expression.h:925
Definition variable.h:50
Definition variable.h:288
Definition expression.h:938
Definition expression.h:436
Definition expression.h:869
Definition variable.h:191