1 //===-- lib/Semantics/check-case.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 "check-case.h"
10 #include "flang/Common/idioms.h"
11 #include "flang/Common/reference.h"
12 #include "flang/Common/template.h"
13 #include "flang/Evaluate/fold.h"
14 #include "flang/Evaluate/type.h"
15 #include "flang/Parser/parse-tree.h"
16 #include "flang/Semantics/semantics.h"
17 #include "flang/Semantics/tools.h"
18 #include <tuple>
19 
20 namespace Fortran::semantics {
21 
22 template <typename T> class CaseValues {
23 public:
24   CaseValues(SemanticsContext &c, const evaluate::DynamicType &t)
25       : context_{c}, caseExprType_{t} {}
26 
27   void Check(const std::list<parser::CaseConstruct::Case> &cases) {
28     for (const parser::CaseConstruct::Case &c : cases) {
29       AddCase(c);
30     }
31     if (!hasErrors_) {
32       cases_.sort(Comparator{});
33       if (!AreCasesDisjoint()) { // C1149
34         ReportConflictingCases();
35       }
36     }
37   }
38 
39 private:
40   using Value = evaluate::Scalar<T>;
41 
42   void AddCase(const parser::CaseConstruct::Case &c) {
43     const auto &stmt{std::get<parser::Statement<parser::CaseStmt>>(c.t)};
44     const parser::CaseStmt &caseStmt{stmt.statement};
45     const auto &selector{std::get<parser::CaseSelector>(caseStmt.t)};
46     common::visit(
47         common::visitors{
48             [&](const std::list<parser::CaseValueRange> &ranges) {
49               for (const auto &range : ranges) {
50                 auto pair{ComputeBounds(range)};
51                 if (pair.first && pair.second && *pair.first > *pair.second) {
52                   context_.Say(stmt.source,
53                       "CASE has lower bound greater than upper bound"_warn_en_US);
54                 } else {
55                   if constexpr (T::category == TypeCategory::Logical) { // C1148
56                     if ((pair.first || pair.second) &&
57                         (!pair.first || !pair.second ||
58                             *pair.first != *pair.second)) {
59                       context_.Say(stmt.source,
60                           "CASE range is not allowed for LOGICAL"_err_en_US);
61                     }
62                   }
63                   cases_.emplace_back(stmt);
64                   cases_.back().lower = std::move(pair.first);
65                   cases_.back().upper = std::move(pair.second);
66                 }
67               }
68             },
69             [&](const parser::Default &) { cases_.emplace_front(stmt); },
70         },
71         selector.u);
72   }
73 
74   std::optional<Value> GetValue(const parser::CaseValue &caseValue) {
75     const parser::Expr &expr{caseValue.thing.thing.value()};
76     auto *x{expr.typedExpr.get()};
77     if (x && x->v) { // C1147
78       auto type{x->v->GetType()};
79       if (type && type->category() == caseExprType_.category() &&
80           (type->category() != TypeCategory::Character ||
81               type->kind() == caseExprType_.kind())) {
82         x->v = evaluate::Fold(context_.foldingContext(),
83             evaluate::ConvertToType(T::GetType(), std::move(*x->v)));
84         if (x->v) {
85           if (auto value{evaluate::GetScalarConstantValue<T>(*x->v)}) {
86             return *value;
87           }
88         }
89         context_.Say(
90             expr.source, "CASE value must be a constant scalar"_err_en_US);
91       } else {
92         std::string typeStr{type ? type->AsFortran() : "typeless"s};
93         context_.Say(expr.source,
94             "CASE value has type '%s' which is not compatible with the SELECT CASE expression's type '%s'"_err_en_US,
95             typeStr, caseExprType_.AsFortran());
96       }
97       hasErrors_ = true;
98     }
99     return std::nullopt;
100   }
101 
102   using PairOfValues = std::pair<std::optional<Value>, std::optional<Value>>;
103   PairOfValues ComputeBounds(const parser::CaseValueRange &range) {
104     return common::visit(
105         common::visitors{
106             [&](const parser::CaseValue &x) {
107               auto value{GetValue(x)};
108               return PairOfValues{value, value};
109             },
110             [&](const parser::CaseValueRange::Range &x) {
111               std::optional<Value> lo, hi;
112               if (x.lower) {
113                 lo = GetValue(*x.lower);
114               }
115               if (x.upper) {
116                 hi = GetValue(*x.upper);
117               }
118               if ((x.lower && !lo) || (x.upper && !hi)) {
119                 return PairOfValues{}; // error case
120               }
121               return PairOfValues{std::move(lo), std::move(hi)};
122             },
123         },
124         range.u);
125   }
126 
127   struct Case {
128     explicit Case(const parser::Statement<parser::CaseStmt> &s) : stmt{s} {}
129     bool IsDefault() const { return !lower && !upper; }
130     std::string AsFortran() const {
131       std::string result;
132       {
133         llvm::raw_string_ostream bs{result};
134         if (lower) {
135           evaluate::Constant<T>{*lower}.AsFortran(bs << '(');
136           if (!upper) {
137             bs << ':';
138           } else if (*lower != *upper) {
139             evaluate::Constant<T>{*upper}.AsFortran(bs << ':');
140           }
141           bs << ')';
142         } else if (upper) {
143           evaluate::Constant<T>{*upper}.AsFortran(bs << "(:") << ')';
144         } else {
145           bs << "DEFAULT";
146         }
147       }
148       return result;
149     }
150 
151     const parser::Statement<parser::CaseStmt> &stmt;
152     std::optional<Value> lower, upper;
153   };
154 
155   // Defines a comparator for use with std::list<>::sort().
156   // Returns true if and only if the highest value in range x is less
157   // than the least value in range y.  The DEFAULT case is arbitrarily
158   // defined to be less than all others.  When two ranges overlap,
159   // neither is less than the other.
160   struct Comparator {
161     bool operator()(const Case &x, const Case &y) const {
162       if (x.IsDefault()) {
163         return !y.IsDefault();
164       } else {
165         return x.upper && y.lower && *x.upper < *y.lower;
166       }
167     }
168   };
169 
170   bool AreCasesDisjoint() const {
171     auto endIter{cases_.end()};
172     for (auto iter{cases_.begin()}; iter != endIter; ++iter) {
173       auto next{iter};
174       if (++next != endIter && !Comparator{}(*iter, *next)) {
175         return false;
176       }
177     }
178     return true;
179   }
180 
181   // This has quadratic time, but only runs in error cases
182   void ReportConflictingCases() {
183     for (auto iter{cases_.begin()}; iter != cases_.end(); ++iter) {
184       parser::Message *msg{nullptr};
185       for (auto p{cases_.begin()}; p != cases_.end(); ++p) {
186         if (p->stmt.source.begin() < iter->stmt.source.begin() &&
187             !Comparator{}(*p, *iter) && !Comparator{}(*iter, *p)) {
188           if (!msg) {
189             msg = &context_.Say(iter->stmt.source,
190                 "CASE %s conflicts with previous cases"_err_en_US,
191                 iter->AsFortran());
192           }
193           msg->Attach(
194               p->stmt.source, "Conflicting CASE %s"_en_US, p->AsFortran());
195         }
196       }
197     }
198   }
199 
200   SemanticsContext &context_;
201   const evaluate::DynamicType &caseExprType_;
202   std::list<Case> cases_;
203   bool hasErrors_{false};
204 };
205 
206 template <TypeCategory CAT> struct TypeVisitor {
207   using Result = bool;
208   using Types = evaluate::CategoryTypes<CAT>;
209   template <typename T> Result Test() {
210     if (T::kind == exprType.kind()) {
211       CaseValues<T>(context, exprType).Check(caseList);
212       return true;
213     } else {
214       return false;
215     }
216   }
217   SemanticsContext &context;
218   const evaluate::DynamicType &exprType;
219   const std::list<parser::CaseConstruct::Case> &caseList;
220 };
221 
222 void CaseChecker::Enter(const parser::CaseConstruct &construct) {
223   const auto &selectCaseStmt{
224       std::get<parser::Statement<parser::SelectCaseStmt>>(construct.t)};
225   const auto &selectCase{selectCaseStmt.statement};
226   const auto &selectExpr{
227       std::get<parser::Scalar<parser::Expr>>(selectCase.t).thing};
228   const auto *x{GetExpr(selectExpr)};
229   if (!x) {
230     return; // expression semantics failed
231   }
232   if (auto exprType{x->GetType()}) {
233     const auto &caseList{
234         std::get<std::list<parser::CaseConstruct::Case>>(construct.t)};
235     switch (exprType->category()) {
236     case TypeCategory::Integer:
237       common::SearchTypes(
238           TypeVisitor<TypeCategory::Integer>{context_, *exprType, caseList});
239       return;
240     case TypeCategory::Logical:
241       CaseValues<evaluate::Type<TypeCategory::Logical, 1>>{context_, *exprType}
242           .Check(caseList);
243       return;
244     case TypeCategory::Character:
245       common::SearchTypes(
246           TypeVisitor<TypeCategory::Character>{context_, *exprType, caseList});
247       return;
248     default:
249       break;
250     }
251   }
252   context_.Say(selectExpr.source,
253       "SELECT CASE expression must be integer, logical, or character"_err_en_US);
254 }
255 } // namespace Fortran::semantics
256