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