1//===- Marshallers.h - Generic matcher function marshallers -----*- 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/// \file
10/// Functions templates and classes to wrap matcher construct functions.
11///
12/// A collection of template function and classes that provide a generic
13/// marshalling layer on top of matcher construct functions.
14/// These are used by the registry to export all marshaller constructors with
15/// the same generic interface.
16//
17//===----------------------------------------------------------------------===//
18
19#ifndef LLVM_CLANG_LIB_ASTMATCHERS_DYNAMIC_MARSHALLERS_H
20#define LLVM_CLANG_LIB_ASTMATCHERS_DYNAMIC_MARSHALLERS_H
21
22#include "clang/AST/ASTTypeTraits.h"
23#include "clang/AST/OperationKinds.h"
24#include "clang/ASTMatchers/ASTMatchersInternal.h"
25#include "clang/ASTMatchers/Dynamic/Diagnostics.h"
26#include "clang/ASTMatchers/Dynamic/VariantValue.h"
27#include "clang/Basic/AttrKinds.h"
28#include "clang/Basic/BuiltinTraits.h"
29#include "clang/Basic/LLVM.h"
30#include "clang/Basic/OpenMPKinds.h"
31#include "llvm/ADT/ArrayRef.h"
32#include "llvm/ADT/STLExtras.h"
33#include "llvm/ADT/StringRef.h"
34#include "llvm/ADT/StringSwitch.h"
35#include "llvm/ADT/Twine.h"
36#include "llvm/Support/Regex.h"
37#include <cassert>
38#include <cstddef>
39#include <iterator>
40#include <limits>
41#include <memory>
42#include <optional>
43#include <string>
44#include <utility>
45#include <vector>
46
47namespace clang {
48namespace ast_matchers {
49namespace dynamic {
50namespace internal {
51
52/// Helper template class to just from argument type to the right is/get
53/// functions in VariantValue.
54/// Used to verify and extract the matcher arguments below.
55template <class T> struct ArgTypeTraits;
56template <class T> struct ArgTypeTraits<const T &> : public ArgTypeTraits<T> {
57};
58
59template <> struct ArgTypeTraits<std::string> {
60 static bool hasCorrectType(const VariantValue &Value) {
61 return Value.isString();
62 }
63 static bool hasCorrectValue(const VariantValue &Value) { return true; }
64
65 static const std::string &get(const VariantValue &Value) {
66 return Value.getString();
67 }
68
69 static ArgKind getKind() {
70 return ArgKind(ArgKind::AK_String);
71 }
72
73 static std::optional<std::string> getBestGuess(const VariantValue &) {
74 return std::nullopt;
75 }
76};
77
78template <>
79struct ArgTypeTraits<StringRef> : public ArgTypeTraits<std::string> {
80};
81
82template <class T> struct ArgTypeTraits<ast_matchers::internal::Matcher<T>> {
83 static bool hasCorrectType(const VariantValue& Value) {
84 return Value.isMatcher();
85 }
86 static bool hasCorrectValue(const VariantValue &Value) {
87 return Value.getMatcher().hasTypedMatcher<T>();
88 }
89
90 static ast_matchers::internal::Matcher<T> get(const VariantValue &Value) {
91 return Value.getMatcher().getTypedMatcher<T>();
92 }
93
94 static ArgKind getKind() {
95 return ArgKind::MakeMatcherArg(MatcherKind: ASTNodeKind::getFromNodeKind<T>());
96 }
97
98 static std::optional<std::string> getBestGuess(const VariantValue &) {
99 return std::nullopt;
100 }
101};
102
103template <> struct ArgTypeTraits<bool> {
104 static bool hasCorrectType(const VariantValue &Value) {
105 return Value.isBoolean();
106 }
107 static bool hasCorrectValue(const VariantValue &Value) { return true; }
108
109 static bool get(const VariantValue &Value) {
110 return Value.getBoolean();
111 }
112
113 static ArgKind getKind() {
114 return ArgKind(ArgKind::AK_Boolean);
115 }
116
117 static std::optional<std::string> getBestGuess(const VariantValue &) {
118 return std::nullopt;
119 }
120};
121
122template <> struct ArgTypeTraits<double> {
123 static bool hasCorrectType(const VariantValue &Value) {
124 return Value.isDouble();
125 }
126 static bool hasCorrectValue(const VariantValue &Value) { return true; }
127
128 static double get(const VariantValue &Value) {
129 return Value.getDouble();
130 }
131
132 static ArgKind getKind() {
133 return ArgKind(ArgKind::AK_Double);
134 }
135
136 static std::optional<std::string> getBestGuess(const VariantValue &) {
137 return std::nullopt;
138 }
139};
140
141template <> struct ArgTypeTraits<unsigned> {
142 static bool hasCorrectType(const VariantValue &Value) {
143 return Value.isUnsigned();
144 }
145 static bool hasCorrectValue(const VariantValue &Value) { return true; }
146
147 static unsigned get(const VariantValue &Value) {
148 return Value.getUnsigned();
149 }
150
151 static ArgKind getKind() {
152 return ArgKind(ArgKind::AK_Unsigned);
153 }
154
155 static std::optional<std::string> getBestGuess(const VariantValue &) {
156 return std::nullopt;
157 }
158};
159
160template <> struct ArgTypeTraits<attr::Kind> {
161private:
162 static std::optional<attr::Kind> getAttrKind(llvm::StringRef AttrKind) {
163 if (!AttrKind.consume_front(Prefix: "attr::"))
164 return std::nullopt;
165 return llvm::StringSwitch<std::optional<attr::Kind>>(AttrKind)
166#define ATTR(X) .Case(#X, attr::X)
167#include "clang/Basic/AttrList.inc"
168 .Default(Value: std::nullopt);
169 }
170
171public:
172 static bool hasCorrectType(const VariantValue &Value) {
173 return Value.isString();
174 }
175 static bool hasCorrectValue(const VariantValue& Value) {
176 return getAttrKind(AttrKind: Value.getString()).has_value();
177 }
178
179 static attr::Kind get(const VariantValue &Value) {
180 return *getAttrKind(AttrKind: Value.getString());
181 }
182
183 static ArgKind getKind() {
184 return ArgKind(ArgKind::AK_String);
185 }
186
187 static std::optional<std::string> getBestGuess(const VariantValue &Value);
188};
189
190template <> struct ArgTypeTraits<CastKind> {
191private:
192 static std::optional<CastKind> getCastKind(llvm::StringRef AttrKind) {
193 if (!AttrKind.consume_front(Prefix: "CK_"))
194 return std::nullopt;
195 return llvm::StringSwitch<std::optional<CastKind>>(AttrKind)
196#define CAST_OPERATION(Name) .Case(#Name, CK_##Name)
197#include "clang/AST/OperationKinds.def"
198 .Default(Value: std::nullopt);
199 }
200
201public:
202 static bool hasCorrectType(const VariantValue &Value) {
203 return Value.isString();
204 }
205 static bool hasCorrectValue(const VariantValue& Value) {
206 return getCastKind(AttrKind: Value.getString()).has_value();
207 }
208
209 static CastKind get(const VariantValue &Value) {
210 return *getCastKind(AttrKind: Value.getString());
211 }
212
213 static ArgKind getKind() {
214 return ArgKind(ArgKind::AK_String);
215 }
216
217 static std::optional<std::string> getBestGuess(const VariantValue &Value);
218};
219
220template <> struct ArgTypeTraits<llvm::Regex::RegexFlags> {
221private:
222 static std::optional<llvm::Regex::RegexFlags> getFlags(llvm::StringRef Flags);
223
224public:
225 static bool hasCorrectType(const VariantValue &Value) {
226 return Value.isString();
227 }
228 static bool hasCorrectValue(const VariantValue& Value) {
229 return getFlags(Flags: Value.getString()).has_value();
230 }
231
232 static llvm::Regex::RegexFlags get(const VariantValue &Value) {
233 return *getFlags(Flags: Value.getString());
234 }
235
236 static ArgKind getKind() { return ArgKind(ArgKind::AK_String); }
237
238 static std::optional<std::string> getBestGuess(const VariantValue &Value);
239};
240
241template <> struct ArgTypeTraits<OpenMPClauseKind> {
242private:
243 static std::optional<OpenMPClauseKind>
244 getClauseKind(llvm::StringRef ClauseKind) {
245 return llvm::StringSwitch<std::optional<OpenMPClauseKind>>(ClauseKind)
246#define GEN_CLANG_CLAUSE_CLASS
247#define CLAUSE_CLASS(Enum, Str, Class) .Case(#Enum, llvm::omp::Clause::Enum)
248#include "llvm/Frontend/OpenMP/OMP.inc"
249 .Default(Value: std::nullopt);
250 }
251
252public:
253 static bool hasCorrectType(const VariantValue &Value) {
254 return Value.isString();
255 }
256 static bool hasCorrectValue(const VariantValue& Value) {
257 return getClauseKind(ClauseKind: Value.getString()).has_value();
258 }
259
260 static OpenMPClauseKind get(const VariantValue &Value) {
261 return *getClauseKind(ClauseKind: Value.getString());
262 }
263
264 static ArgKind getKind() { return ArgKind(ArgKind::AK_String); }
265
266 static std::optional<std::string> getBestGuess(const VariantValue &Value);
267};
268
269template <> struct ArgTypeTraits<UnaryExprOrTypeTrait> {
270private:
271 static std::optional<UnaryExprOrTypeTrait>
272 getUnaryOrTypeTraitKind(llvm::StringRef ClauseKind) {
273 if (!ClauseKind.consume_front(Prefix: "UETT_"))
274 return std::nullopt;
275 return llvm::StringSwitch<std::optional<UnaryExprOrTypeTrait>>(ClauseKind)
276#define UNARY_EXPR_OR_TYPE_TRAIT(Spelling, Name, Key) .Case(#Name, UETT_##Name)
277#define CXX11_UNARY_EXPR_OR_TYPE_TRAIT(Spelling, Name, Key) \
278 .Case(#Name, UETT_##Name)
279#include "clang/Basic/BuiltinTraits.inc"
280 .Default(Value: std::nullopt);
281 }
282
283public:
284 static bool hasCorrectType(const VariantValue &Value) {
285 return Value.isString();
286 }
287 static bool hasCorrectValue(const VariantValue& Value) {
288 return getUnaryOrTypeTraitKind(ClauseKind: Value.getString()).has_value();
289 }
290
291 static UnaryExprOrTypeTrait get(const VariantValue &Value) {
292 return *getUnaryOrTypeTraitKind(ClauseKind: Value.getString());
293 }
294
295 static ArgKind getKind() { return ArgKind(ArgKind::AK_String); }
296
297 static std::optional<std::string> getBestGuess(const VariantValue &Value);
298};
299
300/// Matcher descriptor interface.
301///
302/// Provides a \c create() method that constructs the matcher from the provided
303/// arguments, and various other methods for type introspection.
304class MatcherDescriptor {
305public:
306 virtual ~MatcherDescriptor() = default;
307
308 virtual VariantMatcher create(SourceRange NameRange,
309 ArrayRef<ParserValue> Args,
310 Diagnostics *Error) const = 0;
311
312 virtual ASTNodeKind nodeMatcherType() const { return ASTNodeKind(); }
313
314 virtual bool isBuilderMatcher() const { return false; }
315
316 virtual std::unique_ptr<MatcherDescriptor>
317 buildMatcherCtor(SourceRange NameRange, ArrayRef<ParserValue> Args,
318 Diagnostics *Error) const {
319 return {};
320 }
321
322 /// Returns whether the matcher is variadic. Variadic matchers can take any
323 /// number of arguments, but they must be of the same type.
324 virtual bool isVariadic() const = 0;
325
326 /// Returns the number of arguments accepted by the matcher if not variadic.
327 virtual unsigned getNumArgs() const = 0;
328
329 /// Given that the matcher is being converted to type \p ThisKind, append the
330 /// set of argument types accepted for argument \p ArgNo to \p ArgKinds.
331 // FIXME: We should provide the ability to constrain the output of this
332 // function based on the types of other matcher arguments.
333 virtual void getArgKinds(ASTNodeKind ThisKind, unsigned ArgNo,
334 std::vector<ArgKind> &ArgKinds) const = 0;
335
336 /// Returns whether this matcher is convertible to the given type. If it is
337 /// so convertible, store in *Specificity a value corresponding to the
338 /// "specificity" of the converted matcher to the given context, and in
339 /// *LeastDerivedKind the least derived matcher kind which would result in the
340 /// same matcher overload. Zero specificity indicates that this conversion
341 /// would produce a trivial matcher that will either always or never match.
342 /// Such matchers are excluded from code completion results.
343 virtual bool
344 isConvertibleTo(ASTNodeKind Kind, unsigned *Specificity = nullptr,
345 ASTNodeKind *LeastDerivedKind = nullptr) const = 0;
346
347 /// Returns whether the matcher will, given a matcher of any type T, yield a
348 /// matcher of type T.
349 virtual bool isPolymorphic() const { return false; }
350};
351
352inline bool isRetKindConvertibleTo(ArrayRef<ASTNodeKind> RetKinds,
353 ASTNodeKind Kind, unsigned *Specificity,
354 ASTNodeKind *LeastDerivedKind) {
355 for (const ASTNodeKind &NodeKind : RetKinds) {
356 if (ArgKind::MakeMatcherArg(MatcherKind: NodeKind).isConvertibleTo(
357 To: ArgKind::MakeMatcherArg(MatcherKind: Kind), Specificity)) {
358 if (LeastDerivedKind)
359 *LeastDerivedKind = NodeKind;
360 return true;
361 }
362 }
363 return false;
364}
365
366/// Simple callback implementation. Marshaller and function are provided.
367///
368/// This class wraps a function of arbitrary signature and a marshaller
369/// function into a MatcherDescriptor.
370/// The marshaller is in charge of taking the VariantValue arguments, checking
371/// their types, unpacking them and calling the underlying function.
372class FixedArgCountMatcherDescriptor : public MatcherDescriptor {
373public:
374 using MarshallerType = VariantMatcher (*)(void (*Func)(),
375 StringRef MatcherName,
376 SourceRange NameRange,
377 ArrayRef<ParserValue> Args,
378 Diagnostics *Error);
379
380 /// \param Marshaller Function to unpack the arguments and call \c Func
381 /// \param Func Matcher construct function. This is the function that
382 /// compile-time matcher expressions would use to create the matcher.
383 /// \param RetKinds The list of matcher types to which the matcher is
384 /// convertible.
385 /// \param ArgKinds The types of the arguments this matcher takes.
386 FixedArgCountMatcherDescriptor(MarshallerType Marshaller, void (*Func)(),
387 StringRef MatcherName,
388 ArrayRef<ASTNodeKind> RetKinds,
389 ArrayRef<ArgKind> ArgKinds)
390 : Marshaller(Marshaller), Func(Func), MatcherName(MatcherName),
391 RetKinds(RetKinds.begin(), RetKinds.end()),
392 ArgKinds(ArgKinds.begin(), ArgKinds.end()) {}
393
394 VariantMatcher create(SourceRange NameRange,
395 ArrayRef<ParserValue> Args,
396 Diagnostics *Error) const override {
397 return Marshaller(Func, MatcherName, NameRange, Args, Error);
398 }
399
400 bool isVariadic() const override { return false; }
401 unsigned getNumArgs() const override { return ArgKinds.size(); }
402
403 void getArgKinds(ASTNodeKind ThisKind, unsigned ArgNo,
404 std::vector<ArgKind> &Kinds) const override {
405 Kinds.push_back(x: ArgKinds[ArgNo]);
406 }
407
408 bool isConvertibleTo(ASTNodeKind Kind, unsigned *Specificity,
409 ASTNodeKind *LeastDerivedKind) const override {
410 return isRetKindConvertibleTo(RetKinds, Kind, Specificity,
411 LeastDerivedKind);
412 }
413
414private:
415 const MarshallerType Marshaller;
416 void (* const Func)();
417 const std::string MatcherName;
418 const std::vector<ASTNodeKind> RetKinds;
419 const std::vector<ArgKind> ArgKinds;
420};
421
422/// Helper methods to extract and merge all possible typed matchers
423/// out of the polymorphic object.
424template <class PolyMatcher>
425void mergePolyMatchers(const PolyMatcher &Poly,
426 std::vector<DynTypedMatcher> &Out,
427 ast_matchers::internal::EmptyTypeList) {}
428
429template <class PolyMatcher, class TypeList>
430void mergePolyMatchers(const PolyMatcher &Poly,
431 std::vector<DynTypedMatcher> &Out, TypeList) {
432 Out.push_back(ast_matchers::internal::Matcher<typename TypeList::head>(Poly));
433 mergePolyMatchers(Poly, Out, typename TypeList::tail());
434}
435
436/// Convert the return values of the functions into a VariantMatcher.
437///
438/// There are 2 cases right now: The return value is a Matcher<T> or is a
439/// polymorphic matcher. For the former, we just construct the VariantMatcher.
440/// For the latter, we instantiate all the possible Matcher<T> of the poly
441/// matcher.
442inline VariantMatcher outvalueToVariantMatcher(const DynTypedMatcher &Matcher) {
443 return VariantMatcher::SingleMatcher(Matcher);
444}
445
446template <typename T>
447VariantMatcher outvalueToVariantMatcher(const T &PolyMatcher,
448 typename T::ReturnTypes * = nullptr) {
449 std::vector<DynTypedMatcher> Matchers;
450 mergePolyMatchers(PolyMatcher, Matchers, typename T::ReturnTypes());
451 VariantMatcher Out = VariantMatcher::PolymorphicMatcher(Matchers: std::move(Matchers));
452 return Out;
453}
454
455template <typename T>
456inline void
457buildReturnTypeVectorFromTypeList(std::vector<ASTNodeKind> &RetTypes) {
458 RetTypes.push_back(ASTNodeKind::getFromNodeKind<typename T::head>());
459 buildReturnTypeVectorFromTypeList<typename T::tail>(RetTypes);
460}
461
462template <>
463inline void
464buildReturnTypeVectorFromTypeList<ast_matchers::internal::EmptyTypeList>(
465 std::vector<ASTNodeKind> &RetTypes) {}
466
467template <typename T>
468struct BuildReturnTypeVector {
469 static void build(std::vector<ASTNodeKind> &RetTypes) {
470 buildReturnTypeVectorFromTypeList<typename T::ReturnTypes>(RetTypes);
471 }
472};
473
474template <typename T>
475struct BuildReturnTypeVector<ast_matchers::internal::Matcher<T>> {
476 static void build(std::vector<ASTNodeKind> &RetTypes) {
477 RetTypes.push_back(ASTNodeKind::getFromNodeKind<T>());
478 }
479};
480
481template <typename T>
482struct BuildReturnTypeVector<ast_matchers::internal::BindableMatcher<T>> {
483 static void build(std::vector<ASTNodeKind> &RetTypes) {
484 RetTypes.push_back(ASTNodeKind::getFromNodeKind<T>());
485 }
486};
487
488/// Variadic marshaller function.
489template <typename ResultT, typename ArgT,
490 ResultT (*Func)(ArrayRef<const ArgT *>)>
491VariantMatcher
492variadicMatcherDescriptor(StringRef MatcherName, SourceRange NameRange,
493 ArrayRef<ParserValue> Args, Diagnostics *Error) {
494 SmallVector<ArgT *, 8> InnerArgsPtr;
495 InnerArgsPtr.resize_for_overwrite(Args.size());
496 SmallVector<ArgT, 8> InnerArgs;
497 InnerArgs.reserve(Args.size());
498
499 for (size_t i = 0, e = Args.size(); i != e; ++i) {
500 using ArgTraits = ArgTypeTraits<ArgT>;
501
502 const ParserValue &Arg = Args[i];
503 const VariantValue &Value = Arg.Value;
504 if (!ArgTraits::hasCorrectType(Value)) {
505 Error->addError(Range: Arg.Range, Error: Error->ET_RegistryWrongArgType)
506 << (i + 1) << ArgTraits::getKind().asString() << Value.getTypeAsString();
507 return {};
508 }
509 if (!ArgTraits::hasCorrectValue(Value)) {
510 if (std::optional<std::string> BestGuess =
511 ArgTraits::getBestGuess(Value)) {
512 Error->addError(Range: Arg.Range, Error: Error->ET_RegistryUnknownEnumWithReplace)
513 << i + 1 << Value.getString() << *BestGuess;
514 } else if (Value.isString()) {
515 Error->addError(Range: Arg.Range, Error: Error->ET_RegistryValueNotFound)
516 << Value.getString();
517 } else {
518 // This isn't ideal, but it's better than reporting an empty string as
519 // the error in this case.
520 Error->addError(Range: Arg.Range, Error: Error->ET_RegistryWrongArgType)
521 << (i + 1) << ArgTraits::getKind().asString()
522 << Value.getTypeAsString();
523 }
524 return {};
525 }
526 assert(InnerArgs.size() < InnerArgs.capacity());
527 InnerArgs.emplace_back(ArgTraits::get(Value));
528 InnerArgsPtr[i] = &InnerArgs[i];
529 }
530 return outvalueToVariantMatcher(Func(InnerArgsPtr));
531}
532
533/// Matcher descriptor for variadic functions.
534///
535/// This class simply wraps a VariadicFunction with the right signature to export
536/// it as a MatcherDescriptor.
537/// This allows us to have one implementation of the interface for as many free
538/// functions as we want, reducing the number of symbols and size of the
539/// object file.
540class VariadicFuncMatcherDescriptor : public MatcherDescriptor {
541public:
542 using RunFunc = VariantMatcher (*)(StringRef MatcherName,
543 SourceRange NameRange,
544 ArrayRef<ParserValue> Args,
545 Diagnostics *Error);
546
547 template <typename ResultT, typename ArgT,
548 ResultT (*F)(ArrayRef<const ArgT *>)>
549 VariadicFuncMatcherDescriptor(
550 ast_matchers::internal::VariadicFunction<ResultT, ArgT, F> Func,
551 StringRef MatcherName)
552 : Func(&variadicMatcherDescriptor<ResultT, ArgT, F>),
553 MatcherName(MatcherName.str()),
554 ArgsKind(ArgTypeTraits<ArgT>::getKind()) {
555 BuildReturnTypeVector<ResultT>::build(RetKinds);
556 }
557
558 VariantMatcher create(SourceRange NameRange,
559 ArrayRef<ParserValue> Args,
560 Diagnostics *Error) const override {
561 return Func(MatcherName, NameRange, Args, Error);
562 }
563
564 bool isVariadic() const override { return true; }
565 unsigned getNumArgs() const override { return 0; }
566
567 void getArgKinds(ASTNodeKind ThisKind, unsigned ArgNo,
568 std::vector<ArgKind> &Kinds) const override {
569 Kinds.push_back(x: ArgsKind);
570 }
571
572 bool isConvertibleTo(ASTNodeKind Kind, unsigned *Specificity,
573 ASTNodeKind *LeastDerivedKind) const override {
574 return isRetKindConvertibleTo(RetKinds, Kind, Specificity,
575 LeastDerivedKind);
576 }
577
578 ASTNodeKind nodeMatcherType() const override { return RetKinds[0]; }
579
580private:
581 const RunFunc Func;
582 const std::string MatcherName;
583 std::vector<ASTNodeKind> RetKinds;
584 const ArgKind ArgsKind;
585};
586
587/// Matcher descriptor for VariadicDynCastAllOfMatchers.
588class DynCastAllOfMatcherDescriptor : public MatcherDescriptor {
589public:
590 DynCastAllOfMatcherDescriptor(ASTNodeKind BaseKind, ASTNodeKind DerivedKind)
591 : BaseKind(BaseKind), DerivedKind(DerivedKind) {}
592
593 VariantMatcher create(SourceRange NameRange, ArrayRef<ParserValue> Args,
594 Diagnostics *Error) const override;
595
596 bool isVariadic() const override { return true; }
597 unsigned getNumArgs() const override { return 0; }
598
599 void getArgKinds(ASTNodeKind, unsigned,
600 std::vector<ArgKind> &Kinds) const override {
601 Kinds.push_back(x: ArgKind::MakeMatcherArg(MatcherKind: DerivedKind));
602 }
603
604 bool isConvertibleTo(ASTNodeKind Kind, unsigned *Specificity,
605 ASTNodeKind *LeastDerivedKind) const override;
606
607 ASTNodeKind nodeMatcherType() const override { return DerivedKind; }
608
609private:
610 const ASTNodeKind BaseKind;
611 const ASTNodeKind DerivedKind;
612};
613
614/// Helper macros to check the arguments on all marshaller functions.
615#define CHECK_ARG_COUNT(count) \
616 if (Args.size() != count) { \
617 Error->addError(NameRange, Error->ET_RegistryWrongArgCount) \
618 << count << Args.size(); \
619 return VariantMatcher(); \
620 }
621
622#define CHECK_ARG_TYPE(index, type) \
623 if (!ArgTypeTraits<type>::hasCorrectType(Args[index].Value)) { \
624 Error->addError(Args[index].Range, Error->ET_RegistryWrongArgType) \
625 << (index + 1) << ArgTypeTraits<type>::getKind().asString() \
626 << Args[index].Value.getTypeAsString(); \
627 return VariantMatcher(); \
628 } \
629 if (!ArgTypeTraits<type>::hasCorrectValue(Args[index].Value)) { \
630 if (std::optional<std::string> BestGuess = \
631 ArgTypeTraits<type>::getBestGuess(Args[index].Value)) { \
632 Error->addError(Args[index].Range, \
633 Error->ET_RegistryUnknownEnumWithReplace) \
634 << index + 1 << Args[index].Value.getString() << *BestGuess; \
635 } else if (Args[index].Value.isString()) { \
636 Error->addError(Args[index].Range, Error->ET_RegistryValueNotFound) \
637 << Args[index].Value.getString(); \
638 } \
639 return VariantMatcher(); \
640 }
641
642/// 0-arg marshaller function.
643template <typename ReturnType>
644VariantMatcher
645matcherMarshall0(void (*Func)(), StringRef MatcherName, SourceRange NameRange,
646 ArrayRef<ParserValue> Args, Diagnostics *Error) {
647 using FuncType = ReturnType (*)();
648 CHECK_ARG_COUNT(0);
649 return outvalueToVariantMatcher(reinterpret_cast<FuncType>(Func)());
650}
651
652/// 1-arg marshaller function.
653template <typename ReturnType, typename ArgType1>
654VariantMatcher
655matcherMarshall1(void (*Func)(), StringRef MatcherName, SourceRange NameRange,
656 ArrayRef<ParserValue> Args, Diagnostics *Error) {
657 using FuncType = ReturnType (*)(ArgType1);
658 CHECK_ARG_COUNT(1);
659 CHECK_ARG_TYPE(0, ArgType1);
660 return outvalueToVariantMatcher(reinterpret_cast<FuncType>(Func)(
661 ArgTypeTraits<ArgType1>::get(Args[0].Value)));
662}
663
664/// 2-arg marshaller function.
665template <typename ReturnType, typename ArgType1, typename ArgType2>
666VariantMatcher
667matcherMarshall2(void (*Func)(), StringRef MatcherName, SourceRange NameRange,
668 ArrayRef<ParserValue> Args, Diagnostics *Error) {
669 using FuncType = ReturnType (*)(ArgType1, ArgType2);
670 CHECK_ARG_COUNT(2);
671 CHECK_ARG_TYPE(0, ArgType1);
672 CHECK_ARG_TYPE(1, ArgType2);
673 return outvalueToVariantMatcher(reinterpret_cast<FuncType>(Func)(
674 ArgTypeTraits<ArgType1>::get(Args[0].Value),
675 ArgTypeTraits<ArgType2>::get(Args[1].Value)));
676}
677
678#undef CHECK_ARG_COUNT
679#undef CHECK_ARG_TYPE
680
681/// Helper class used to collect all the possible overloads of an
682/// argument adaptative matcher function.
683template <template <typename ToArg, typename FromArg> class ArgumentAdapterT,
684 typename FromTypes, typename ToTypes>
685class AdaptativeOverloadCollector {
686public:
687 AdaptativeOverloadCollector(
688 StringRef Name, std::vector<std::unique_ptr<MatcherDescriptor>> &Out)
689 : Name(Name), Out(Out) {
690 collect(FromTypes());
691 }
692
693private:
694 using AdaptativeFunc = ast_matchers::internal::ArgumentAdaptingMatcherFunc<
695 ArgumentAdapterT, FromTypes, ToTypes>;
696
697 /// End case for the recursion
698 static void collect(ast_matchers::internal::EmptyTypeList) {}
699
700 /// Recursive case. Get the overload for the head of the list, and
701 /// recurse to the tail.
702 template <typename FromTypeList>
703 inline void collect(FromTypeList);
704
705 StringRef Name;
706 std::vector<std::unique_ptr<MatcherDescriptor>> &Out;
707};
708
709/// MatcherDescriptor that wraps multiple "overloads" of the same
710/// matcher.
711///
712/// It will try every overload and generate appropriate errors for when none or
713/// more than one overloads match the arguments.
714class OverloadedMatcherDescriptor : public MatcherDescriptor {
715public:
716 OverloadedMatcherDescriptor(
717 MutableArrayRef<std::unique_ptr<MatcherDescriptor>> Callbacks)
718 : Overloads(std::make_move_iterator(i: Callbacks.begin()),
719 std::make_move_iterator(i: Callbacks.end())) {}
720
721 ~OverloadedMatcherDescriptor() override = default;
722
723 VariantMatcher create(SourceRange NameRange,
724 ArrayRef<ParserValue> Args,
725 Diagnostics *Error) const override {
726 std::vector<VariantMatcher> Constructed;
727 Diagnostics::OverloadContext Ctx(Error);
728 for (const auto &O : Overloads) {
729 VariantMatcher SubMatcher = O->create(NameRange, Args, Error);
730 if (!SubMatcher.isNull()) {
731 Constructed.push_back(x: SubMatcher);
732 }
733 }
734
735 if (Constructed.empty()) return VariantMatcher(); // No overload matched.
736 // We ignore the errors if any matcher succeeded.
737 Ctx.revertErrors();
738 if (Constructed.size() > 1) {
739 // More than one constructed. It is ambiguous.
740 Error->addError(Range: NameRange, Error: Error->ET_RegistryAmbiguousOverload);
741 return VariantMatcher();
742 }
743 return Constructed[0];
744 }
745
746 bool isVariadic() const override {
747 bool Overload0Variadic = Overloads[0]->isVariadic();
748#ifndef NDEBUG
749 for (const auto &O : Overloads) {
750 assert(Overload0Variadic == O->isVariadic());
751 }
752#endif
753 return Overload0Variadic;
754 }
755
756 unsigned getNumArgs() const override {
757 unsigned Overload0NumArgs = Overloads[0]->getNumArgs();
758#ifndef NDEBUG
759 for (const auto &O : Overloads) {
760 assert(Overload0NumArgs == O->getNumArgs());
761 }
762#endif
763 return Overload0NumArgs;
764 }
765
766 void getArgKinds(ASTNodeKind ThisKind, unsigned ArgNo,
767 std::vector<ArgKind> &Kinds) const override {
768 for (const auto &O : Overloads) {
769 if (O->isConvertibleTo(Kind: ThisKind))
770 O->getArgKinds(ThisKind, ArgNo, ArgKinds&: Kinds);
771 }
772 }
773
774 bool isConvertibleTo(ASTNodeKind Kind, unsigned *Specificity,
775 ASTNodeKind *LeastDerivedKind) const override {
776 for (const auto &O : Overloads) {
777 if (O->isConvertibleTo(Kind, Specificity, LeastDerivedKind))
778 return true;
779 }
780 return false;
781 }
782
783private:
784 std::vector<std::unique_ptr<MatcherDescriptor>> Overloads;
785};
786
787template <typename ReturnType>
788class RegexMatcherDescriptor : public MatcherDescriptor {
789public:
790 RegexMatcherDescriptor(ReturnType (*WithFlags)(StringRef,
791 llvm::Regex::RegexFlags),
792 ReturnType (*NoFlags)(StringRef),
793 ArrayRef<ASTNodeKind> RetKinds)
794 : WithFlags(WithFlags), NoFlags(NoFlags),
795 RetKinds(RetKinds.begin(), RetKinds.end()) {}
796 bool isVariadic() const override { return true; }
797 unsigned getNumArgs() const override { return 0; }
798
799 void getArgKinds(ASTNodeKind ThisKind, unsigned ArgNo,
800 std::vector<ArgKind> &Kinds) const override {
801 assert(ArgNo < 2);
802 Kinds.push_back(x: ArgKind::AK_String);
803 }
804
805 bool isConvertibleTo(ASTNodeKind Kind, unsigned *Specificity,
806 ASTNodeKind *LeastDerivedKind) const override {
807 return isRetKindConvertibleTo(RetKinds, Kind, Specificity,
808 LeastDerivedKind);
809 }
810
811 VariantMatcher create(SourceRange NameRange, ArrayRef<ParserValue> Args,
812 Diagnostics *Error) const override {
813 if (Args.size() < 1 || Args.size() > 2) {
814 Error->addError(Range: NameRange, Error: Diagnostics::ET_RegistryWrongArgCount)
815 << "1 or 2" << Args.size();
816 return VariantMatcher();
817 }
818 if (!ArgTypeTraits<StringRef>::hasCorrectType(Value: Args[0].Value)) {
819 Error->addError(Range: Args[0].Range, Error: Error->ET_RegistryWrongArgType)
820 << 1 << ArgTypeTraits<StringRef>::getKind().asString()
821 << Args[0].Value.getTypeAsString();
822 return VariantMatcher();
823 }
824 if (Args.size() == 1) {
825 return outvalueToVariantMatcher(
826 NoFlags(ArgTypeTraits<StringRef>::get(Value: Args[0].Value)));
827 }
828 if (!ArgTypeTraits<llvm::Regex::RegexFlags>::hasCorrectType(
829 Value: Args[1].Value)) {
830 Error->addError(Range: Args[1].Range, Error: Error->ET_RegistryWrongArgType)
831 << 2 << ArgTypeTraits<llvm::Regex::RegexFlags>::getKind().asString()
832 << Args[1].Value.getTypeAsString();
833 return VariantMatcher();
834 }
835 if (!ArgTypeTraits<llvm::Regex::RegexFlags>::hasCorrectValue(
836 Value: Args[1].Value)) {
837 if (std::optional<std::string> BestGuess =
838 ArgTypeTraits<llvm::Regex::RegexFlags>::getBestGuess(
839 Value: Args[1].Value)) {
840 Error->addError(Range: Args[1].Range, Error: Error->ET_RegistryUnknownEnumWithReplace)
841 << 2 << Args[1].Value.getString() << *BestGuess;
842 } else {
843 Error->addError(Range: Args[1].Range, Error: Error->ET_RegistryValueNotFound)
844 << Args[1].Value.getString();
845 }
846 return VariantMatcher();
847 }
848 return outvalueToVariantMatcher(
849 WithFlags(ArgTypeTraits<StringRef>::get(Value: Args[0].Value),
850 ArgTypeTraits<llvm::Regex::RegexFlags>::get(Value: Args[1].Value)));
851 }
852
853private:
854 ReturnType (*const WithFlags)(StringRef, llvm::Regex::RegexFlags);
855 ReturnType (*const NoFlags)(StringRef);
856 const std::vector<ASTNodeKind> RetKinds;
857};
858
859/// Variadic operator marshaller function.
860class VariadicOperatorMatcherDescriptor : public MatcherDescriptor {
861public:
862 using VarOp = DynTypedMatcher::VariadicOperator;
863
864 VariadicOperatorMatcherDescriptor(unsigned MinCount, unsigned MaxCount,
865 VarOp Op, StringRef MatcherName)
866 : MinCount(MinCount), MaxCount(MaxCount), Op(Op),
867 MatcherName(MatcherName) {}
868
869 VariantMatcher create(SourceRange NameRange,
870 ArrayRef<ParserValue> Args,
871 Diagnostics *Error) const override {
872 if (Args.size() < MinCount || MaxCount < Args.size()) {
873 const std::string MaxStr =
874 (MaxCount == std::numeric_limits<unsigned>::max() ? ""
875 : Twine(MaxCount))
876 .str();
877 Error->addError(Range: NameRange, Error: Error->ET_RegistryWrongArgCount)
878 << ("(" + Twine(MinCount) + ", " + MaxStr + ")") << Args.size();
879 return VariantMatcher();
880 }
881
882 std::vector<VariantMatcher> InnerArgs;
883 for (size_t i = 0, e = Args.size(); i != e; ++i) {
884 const ParserValue &Arg = Args[i];
885 const VariantValue &Value = Arg.Value;
886 if (!Value.isMatcher()) {
887 Error->addError(Range: Arg.Range, Error: Error->ET_RegistryWrongArgType)
888 << (i + 1) << "Matcher<>" << Value.getTypeAsString();
889 return VariantMatcher();
890 }
891 InnerArgs.push_back(x: Value.getMatcher());
892 }
893 return VariantMatcher::VariadicOperatorMatcher(Op, Args: std::move(InnerArgs));
894 }
895
896 bool isVariadic() const override { return true; }
897 unsigned getNumArgs() const override { return 0; }
898
899 void getArgKinds(ASTNodeKind ThisKind, unsigned ArgNo,
900 std::vector<ArgKind> &Kinds) const override {
901 Kinds.push_back(x: ArgKind::MakeMatcherArg(MatcherKind: ThisKind));
902 }
903
904 bool isConvertibleTo(ASTNodeKind Kind, unsigned *Specificity,
905 ASTNodeKind *LeastDerivedKind) const override {
906 if (Specificity)
907 *Specificity = 1;
908 if (LeastDerivedKind)
909 *LeastDerivedKind = Kind;
910 return true;
911 }
912
913 bool isPolymorphic() const override { return true; }
914
915private:
916 const unsigned MinCount;
917 const unsigned MaxCount;
918 const VarOp Op;
919 const StringRef MatcherName;
920};
921
922class MapAnyOfMatcherDescriptor : public MatcherDescriptor {
923 ASTNodeKind CladeNodeKind;
924 std::vector<ASTNodeKind> NodeKinds;
925
926public:
927 MapAnyOfMatcherDescriptor(ASTNodeKind CladeNodeKind,
928 std::vector<ASTNodeKind> NodeKinds)
929 : CladeNodeKind(CladeNodeKind), NodeKinds(std::move(NodeKinds)) {}
930
931 VariantMatcher create(SourceRange NameRange, ArrayRef<ParserValue> Args,
932 Diagnostics *Error) const override {
933
934 std::vector<DynTypedMatcher> NodeArgs;
935
936 for (auto NK : NodeKinds) {
937 std::vector<DynTypedMatcher> InnerArgs;
938
939 for (const auto &Arg : Args) {
940 if (!Arg.Value.isMatcher())
941 return {};
942 const VariantMatcher &VM = Arg.Value.getMatcher();
943 if (VM.hasTypedMatcher(NK)) {
944 auto DM = VM.getTypedMatcher(NK);
945 InnerArgs.push_back(x: DM);
946 }
947 }
948
949 if (InnerArgs.empty()) {
950 NodeArgs.push_back(
951 x: DynTypedMatcher::trueMatcher(NodeKind: NK).dynCastTo(Kind: CladeNodeKind));
952 } else {
953 NodeArgs.push_back(
954 x: DynTypedMatcher::constructVariadic(
955 Op: ast_matchers::internal::DynTypedMatcher::VO_AllOf, SupportedKind: NK,
956 InnerMatchers: InnerArgs)
957 .dynCastTo(Kind: CladeNodeKind));
958 }
959 }
960
961 auto Result = DynTypedMatcher::constructVariadic(
962 Op: ast_matchers::internal::DynTypedMatcher::VO_AnyOf, SupportedKind: CladeNodeKind,
963 InnerMatchers: NodeArgs);
964 Result.setAllowBind(true);
965 return VariantMatcher::SingleMatcher(Matcher: Result);
966 }
967
968 bool isVariadic() const override { return true; }
969 unsigned getNumArgs() const override { return 0; }
970
971 void getArgKinds(ASTNodeKind ThisKind, unsigned,
972 std::vector<ArgKind> &Kinds) const override {
973 Kinds.push_back(x: ArgKind::MakeMatcherArg(MatcherKind: ThisKind));
974 }
975
976 bool isConvertibleTo(ASTNodeKind Kind, unsigned *Specificity,
977 ASTNodeKind *LeastDerivedKind) const override {
978 if (Specificity)
979 *Specificity = 1;
980 if (LeastDerivedKind)
981 *LeastDerivedKind = CladeNodeKind;
982 return true;
983 }
984};
985
986class MapAnyOfBuilderDescriptor : public MatcherDescriptor {
987public:
988 VariantMatcher create(SourceRange, ArrayRef<ParserValue>,
989 Diagnostics *) const override {
990 return {};
991 }
992
993 bool isBuilderMatcher() const override { return true; }
994
995 std::unique_ptr<MatcherDescriptor>
996 buildMatcherCtor(SourceRange, ArrayRef<ParserValue> Args,
997 Diagnostics *) const override {
998
999 std::vector<ASTNodeKind> NodeKinds;
1000 for (const auto &Arg : Args) {
1001 if (!Arg.Value.isNodeKind())
1002 return {};
1003 NodeKinds.push_back(x: Arg.Value.getNodeKind());
1004 }
1005
1006 if (NodeKinds.empty())
1007 return {};
1008
1009 ASTNodeKind CladeNodeKind = NodeKinds.front().getCladeKind();
1010
1011 for (auto NK : NodeKinds)
1012 {
1013 if (!NK.getCladeKind().isSame(Other: CladeNodeKind))
1014 return {};
1015 }
1016
1017 return std::make_unique<MapAnyOfMatcherDescriptor>(args&: CladeNodeKind,
1018 args: std::move(NodeKinds));
1019 }
1020
1021 bool isVariadic() const override { return true; }
1022
1023 unsigned getNumArgs() const override { return 0; }
1024
1025 void getArgKinds(ASTNodeKind ThisKind, unsigned,
1026 std::vector<ArgKind> &ArgKinds) const override {
1027 ArgKinds.push_back(x: ArgKind::MakeNodeArg(MatcherKind: ThisKind));
1028 }
1029 bool isConvertibleTo(ASTNodeKind Kind, unsigned *Specificity = nullptr,
1030 ASTNodeKind *LeastDerivedKind = nullptr) const override {
1031 if (Specificity)
1032 *Specificity = 1;
1033 if (LeastDerivedKind)
1034 *LeastDerivedKind = Kind;
1035 return true;
1036 }
1037
1038 bool isPolymorphic() const override { return false; }
1039};
1040
1041/// Helper functions to select the appropriate marshaller functions.
1042/// They detect the number of arguments, arguments types and return type.
1043
1044/// 0-arg overload
1045template <typename ReturnType>
1046std::unique_ptr<MatcherDescriptor>
1047makeMatcherAutoMarshall(ReturnType (*Func)(), StringRef MatcherName) {
1048 std::vector<ASTNodeKind> RetTypes;
1049 BuildReturnTypeVector<ReturnType>::build(RetTypes);
1050 return std::make_unique<FixedArgCountMatcherDescriptor>(
1051 matcherMarshall0<ReturnType>, reinterpret_cast<void (*)()>(Func),
1052 MatcherName, RetTypes, ArrayRef<ArgKind>());
1053}
1054
1055/// 1-arg overload
1056template <typename ReturnType, typename ArgType1>
1057std::unique_ptr<MatcherDescriptor>
1058makeMatcherAutoMarshall(ReturnType (*Func)(ArgType1), StringRef MatcherName) {
1059 std::vector<ASTNodeKind> RetTypes;
1060 BuildReturnTypeVector<ReturnType>::build(RetTypes);
1061 ArgKind AK = ArgTypeTraits<ArgType1>::getKind();
1062 return std::make_unique<FixedArgCountMatcherDescriptor>(
1063 matcherMarshall1<ReturnType, ArgType1>,
1064 reinterpret_cast<void (*)()>(Func), MatcherName, RetTypes, AK);
1065}
1066
1067/// 2-arg overload
1068template <typename ReturnType, typename ArgType1, typename ArgType2>
1069std::unique_ptr<MatcherDescriptor>
1070makeMatcherAutoMarshall(ReturnType (*Func)(ArgType1, ArgType2),
1071 StringRef MatcherName) {
1072 std::vector<ASTNodeKind> RetTypes;
1073 BuildReturnTypeVector<ReturnType>::build(RetTypes);
1074 ArgKind AKs[] = { ArgTypeTraits<ArgType1>::getKind(),
1075 ArgTypeTraits<ArgType2>::getKind() };
1076 return std::make_unique<FixedArgCountMatcherDescriptor>(
1077 matcherMarshall2<ReturnType, ArgType1, ArgType2>,
1078 reinterpret_cast<void (*)()>(Func), MatcherName, RetTypes, AKs);
1079}
1080
1081template <typename ReturnType>
1082std::unique_ptr<MatcherDescriptor> makeMatcherRegexMarshall(
1083 ReturnType (*FuncFlags)(llvm::StringRef, llvm::Regex::RegexFlags),
1084 ReturnType (*Func)(llvm::StringRef)) {
1085 std::vector<ASTNodeKind> RetTypes;
1086 BuildReturnTypeVector<ReturnType>::build(RetTypes);
1087 return std::make_unique<RegexMatcherDescriptor<ReturnType>>(FuncFlags, Func,
1088 RetTypes);
1089}
1090
1091/// Variadic overload.
1092template <typename ResultT, typename ArgT,
1093 ResultT (*Func)(ArrayRef<const ArgT *>)>
1094std::unique_ptr<MatcherDescriptor> makeMatcherAutoMarshall(
1095 ast_matchers::internal::VariadicFunction<ResultT, ArgT, Func> VarFunc,
1096 StringRef MatcherName) {
1097 return std::make_unique<VariadicFuncMatcherDescriptor>(VarFunc, MatcherName);
1098}
1099
1100/// Overload for VariadicDynCastAllOfMatchers.
1101///
1102/// Not strictly necessary, but DynCastAllOfMatcherDescriptor gives us better
1103/// completion results for that type of matcher.
1104template <typename BaseT, typename DerivedT>
1105std::unique_ptr<MatcherDescriptor> makeMatcherAutoMarshall(
1106 ast_matchers::internal::VariadicDynCastAllOfMatcher<BaseT, DerivedT>,
1107 StringRef) {
1108 return std::make_unique<DynCastAllOfMatcherDescriptor>(
1109 ASTNodeKind::getFromNodeKind<BaseT>(),
1110 ASTNodeKind::getFromNodeKind<DerivedT>());
1111}
1112
1113/// Argument adaptative overload.
1114template <template <typename ToArg, typename FromArg> class ArgumentAdapterT,
1115 typename FromTypes, typename ToTypes>
1116std::unique_ptr<MatcherDescriptor> makeMatcherAutoMarshall(
1117 ast_matchers::internal::ArgumentAdaptingMatcherFunc<ArgumentAdapterT,
1118 FromTypes, ToTypes>,
1119 StringRef MatcherName) {
1120 std::vector<std::unique_ptr<MatcherDescriptor>> Overloads;
1121 AdaptativeOverloadCollector<ArgumentAdapterT, FromTypes, ToTypes>(MatcherName,
1122 Overloads);
1123 return std::make_unique<OverloadedMatcherDescriptor>(args&: Overloads);
1124}
1125
1126template <template <typename ToArg, typename FromArg> class ArgumentAdapterT,
1127 typename FromTypes, typename ToTypes>
1128template <typename FromTypeList>
1129inline void AdaptativeOverloadCollector<ArgumentAdapterT, FromTypes,
1130 ToTypes>::collect(FromTypeList) {
1131 Out.push_back(makeMatcherAutoMarshall(
1132 &AdaptativeFunc::template create<typename FromTypeList::head>, Name));
1133 collect(typename FromTypeList::tail());
1134}
1135
1136/// Variadic operator overload.
1137template <unsigned MinCount, unsigned MaxCount>
1138std::unique_ptr<MatcherDescriptor> makeMatcherAutoMarshall(
1139 ast_matchers::internal::VariadicOperatorMatcherFunc<MinCount, MaxCount>
1140 Func,
1141 StringRef MatcherName) {
1142 return std::make_unique<VariadicOperatorMatcherDescriptor>(
1143 MinCount, MaxCount, Func.Op, MatcherName);
1144}
1145
1146template <typename CladeType, typename... MatcherT>
1147std::unique_ptr<MatcherDescriptor> makeMatcherAutoMarshall(
1148 ast_matchers::internal::MapAnyOfMatcherImpl<CladeType, MatcherT...>,
1149 StringRef MatcherName) {
1150 return std::make_unique<MapAnyOfMatcherDescriptor>(
1151 ASTNodeKind::getFromNodeKind<CladeType>(),
1152 std::vector<ASTNodeKind>{ASTNodeKind::getFromNodeKind<MatcherT>()...});
1153}
1154
1155} // namespace internal
1156} // namespace dynamic
1157} // namespace ast_matchers
1158} // namespace clang
1159
1160#endif // LLVM_CLANG_LIB_ASTMATCHERS_DYNAMIC_MARSHALLERS_H
1161