@@ -136,6 +136,9 @@ namespace clad {
136136 // expression to be of object type in the reverse mode as well.
137137 clang::Expr* m_ThisExprDerivative = nullptr ;
138138
139+ // / The currently visited statement. Useful for crash pretty-printing.
140+ const clang::Stmt* m_CurVisitedStmt = nullptr ;
141+
139142 // / A function used to wrap result of visiting E in a lambda. Returns a call
140143 // / to the built lambda. Func is a functor that will be invoked inside
141144 // / lambda scope and block. Statements inside lambda are expected to be
@@ -671,6 +674,23 @@ namespace clad {
671674 clang::TemplateDecl* m_CladConstructorPushforwardTag = nullptr ;
672675 clang::TemplateDecl* m_CladConstructorReverseForwTag = nullptr ;
673676 };
677+
678+ // / A class that generates prettier stack traces when we crash on generating
679+ // / a derivative.
680+ class PrettyStackTraceDerivative : public llvm ::PrettyStackTraceEntry {
681+ const DiffRequest& m_DiffReq;
682+ using Blocks = std::vector<llvm::SmallVector<clang::Stmt*, 16 >>;
683+ const Blocks& m_Blocks;
684+ const clang::Sema& m_Sema;
685+ const clang::Stmt** m_Stmt = nullptr ;
686+
687+ public:
688+ PrettyStackTraceDerivative (const DiffRequest& DiffReq, const Blocks& B,
689+ const clang::Sema& Sema, const clang::Stmt** S)
690+ : m_DiffReq(DiffReq), m_Blocks(B), m_Sema(Sema), m_Stmt(S) {}
691+ void print (llvm::raw_ostream& OS ) const override ;
692+ };
693+
674694} // end namespace clad
675695
676696#endif // CLAD_VISITOR_BASE_H
0 commit comments