1 //===-- lib/Evaluate/fold.cpp ---------------------------------------------===//
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 #include "flang/Evaluate/fold.h"
10 #include "fold-implementation.h"
11 #include "flang/Evaluate/characteristics.h"
12 
13 namespace Fortran::evaluate {
14 
15 characteristics::TypeAndShape Fold(
16     FoldingContext &context, characteristics::TypeAndShape &&x) {
17   x.Rewrite(context);
18   return std::move(x);
19 }
20 
21 std::optional<Constant<SubscriptInteger>> GetConstantSubscript(
22     FoldingContext &context, Subscript &ss, const NamedEntity &base, int dim) {
23   ss = FoldOperation(context, std::move(ss));
24   return common::visit(
25       common::visitors{
26           [](IndirectSubscriptIntegerExpr &expr)
27               -> std::optional<Constant<SubscriptInteger>> {
28             if (const auto *constant{
29                     UnwrapConstantValue<SubscriptInteger>(expr.value())}) {
30               return *constant;
31             } else {
32               return std::nullopt;
33             }
34           },
35           [&](Triplet &triplet) -> std::optional<Constant<SubscriptInteger>> {
36             auto lower{triplet.lower()}, upper{triplet.upper()};
37             std::optional<ConstantSubscript> stride{ToInt64(triplet.stride())};
38             if (!lower) {
39               lower = GetLBOUND(context, base, dim);
40             }
41             if (!upper) {
42               if (auto lb{GetLBOUND(context, base, dim)}) {
43                 upper = ComputeUpperBound(
44                     context, std::move(*lb), GetExtent(context, base, dim));
45               }
46             }
47             auto lbi{ToInt64(lower)}, ubi{ToInt64(upper)};
48             if (lbi && ubi && stride && *stride != 0) {
49               std::vector<SubscriptInteger::Scalar> values;
50               while ((*stride > 0 && *lbi <= *ubi) ||
51                   (*stride < 0 && *lbi >= *ubi)) {
52                 values.emplace_back(*lbi);
53                 *lbi += *stride;
54               }
55               return Constant<SubscriptInteger>{std::move(values),
56                   ConstantSubscripts{
57                       static_cast<ConstantSubscript>(values.size())}};
58             } else {
59               return std::nullopt;
60             }
61           },
62       },
63       ss.u);
64 }
65 
66 Expr<SomeDerived> FoldOperation(
67     FoldingContext &context, StructureConstructor &&structure) {
68   StructureConstructor ctor{structure.derivedTypeSpec()};
69   bool isConstant{true};
70   auto restorer{context.WithPDTInstance(structure.derivedTypeSpec())};
71   for (auto &&[symbol, value] : std::move(structure)) {
72     auto expr{Fold(context, std::move(value.value()))};
73     if (IsPointer(symbol)) {
74       if (IsProcedure(symbol)) {
75         isConstant &= IsInitialProcedureTarget(expr);
76       } else {
77         isConstant &= IsInitialDataTarget(expr);
78       }
79     } else {
80       isConstant &= IsActuallyConstant(expr);
81       if (auto valueShape{GetConstantExtents(context, expr)}) {
82         if (auto componentShape{GetConstantExtents(context, symbol)}) {
83           if (GetRank(*componentShape) > 0 && GetRank(*valueShape) == 0) {
84             expr = ScalarConstantExpander{std::move(*componentShape)}.Expand(
85                 std::move(expr));
86             isConstant &= expr.Rank() > 0;
87           } else {
88             isConstant &= *valueShape == *componentShape;
89           }
90         }
91       }
92     }
93     ctor.Add(symbol, std::move(expr));
94   }
95   if (isConstant) {
96     return Expr<SomeDerived>{Constant<SomeDerived>{std::move(ctor)}};
97   } else {
98     return Expr<SomeDerived>{std::move(ctor)};
99   }
100 }
101 
102 Component FoldOperation(FoldingContext &context, Component &&component) {
103   return {FoldOperation(context, std::move(component.base())),
104       component.GetLastSymbol()};
105 }
106 
107 NamedEntity FoldOperation(FoldingContext &context, NamedEntity &&x) {
108   if (Component * c{x.UnwrapComponent()}) {
109     return NamedEntity{FoldOperation(context, std::move(*c))};
110   } else {
111     return std::move(x);
112   }
113 }
114 
115 Triplet FoldOperation(FoldingContext &context, Triplet &&triplet) {
116   MaybeExtentExpr lower{triplet.lower()};
117   MaybeExtentExpr upper{triplet.upper()};
118   return {Fold(context, std::move(lower)), Fold(context, std::move(upper)),
119       Fold(context, triplet.stride())};
120 }
121 
122 Subscript FoldOperation(FoldingContext &context, Subscript &&subscript) {
123   return common::visit(
124       common::visitors{
125           [&](IndirectSubscriptIntegerExpr &&expr) {
126             expr.value() = Fold(context, std::move(expr.value()));
127             return Subscript(std::move(expr));
128           },
129           [&](Triplet &&triplet) {
130             return Subscript(FoldOperation(context, std::move(triplet)));
131           },
132       },
133       std::move(subscript.u));
134 }
135 
136 ArrayRef FoldOperation(FoldingContext &context, ArrayRef &&arrayRef) {
137   NamedEntity base{FoldOperation(context, std::move(arrayRef.base()))};
138   for (Subscript &subscript : arrayRef.subscript()) {
139     subscript = FoldOperation(context, std::move(subscript));
140   }
141   return ArrayRef{std::move(base), std::move(arrayRef.subscript())};
142 }
143 
144 CoarrayRef FoldOperation(FoldingContext &context, CoarrayRef &&coarrayRef) {
145   std::vector<Subscript> subscript;
146   for (Subscript x : coarrayRef.subscript()) {
147     subscript.emplace_back(FoldOperation(context, std::move(x)));
148   }
149   std::vector<Expr<SubscriptInteger>> cosubscript;
150   for (Expr<SubscriptInteger> x : coarrayRef.cosubscript()) {
151     cosubscript.emplace_back(Fold(context, std::move(x)));
152   }
153   CoarrayRef folded{std::move(coarrayRef.base()), std::move(subscript),
154       std::move(cosubscript)};
155   if (std::optional<Expr<SomeInteger>> stat{coarrayRef.stat()}) {
156     folded.set_stat(Fold(context, std::move(*stat)));
157   }
158   if (std::optional<Expr<SomeInteger>> team{coarrayRef.team()}) {
159     folded.set_team(
160         Fold(context, std::move(*team)), coarrayRef.teamIsTeamNumber());
161   }
162   return folded;
163 }
164 
165 DataRef FoldOperation(FoldingContext &context, DataRef &&dataRef) {
166   return common::visit(common::visitors{
167                            [&](SymbolRef symbol) { return DataRef{*symbol}; },
168                            [&](auto &&x) {
169                              return DataRef{
170                                  FoldOperation(context, std::move(x))};
171                            },
172                        },
173       std::move(dataRef.u));
174 }
175 
176 Substring FoldOperation(FoldingContext &context, Substring &&substring) {
177   auto lower{Fold(context, substring.lower())};
178   auto upper{Fold(context, substring.upper())};
179   if (const DataRef * dataRef{substring.GetParentIf<DataRef>()}) {
180     return Substring{FoldOperation(context, DataRef{*dataRef}),
181         std::move(lower), std::move(upper)};
182   } else {
183     auto p{*substring.GetParentIf<StaticDataObject::Pointer>()};
184     return Substring{std::move(p), std::move(lower), std::move(upper)};
185   }
186 }
187 
188 ComplexPart FoldOperation(FoldingContext &context, ComplexPart &&complexPart) {
189   DataRef complex{complexPart.complex()};
190   return ComplexPart{
191       FoldOperation(context, std::move(complex)), complexPart.part()};
192 }
193 
194 std::optional<std::int64_t> GetInt64Arg(
195     const std::optional<ActualArgument> &arg) {
196   if (const auto *intExpr{UnwrapExpr<Expr<SomeInteger>>(arg)}) {
197     return ToInt64(*intExpr);
198   } else {
199     return std::nullopt;
200   }
201 }
202 
203 std::optional<std::int64_t> GetInt64ArgOr(
204     const std::optional<ActualArgument> &arg, std::int64_t defaultValue) {
205   if (!arg) {
206     return defaultValue;
207   } else if (const auto *intExpr{UnwrapExpr<Expr<SomeInteger>>(arg)}) {
208     return ToInt64(*intExpr);
209   } else {
210     return std::nullopt;
211   }
212 }
213 
214 Expr<ImpliedDoIndex::Result> FoldOperation(
215     FoldingContext &context, ImpliedDoIndex &&iDo) {
216   if (std::optional<ConstantSubscript> value{context.GetImpliedDo(iDo.name)}) {
217     return Expr<ImpliedDoIndex::Result>{*value};
218   } else {
219     return Expr<ImpliedDoIndex::Result>{std::move(iDo)};
220   }
221 }
222 
223 template class ExpressionBase<SomeDerived>;
224 template class ExpressionBase<SomeType>;
225 
226 } // namespace Fortran::evaluate
227