#include #include #include #include #include #include #include #include #include #include #include #include #include #include import mean_field; import test_helpers; class CheckReporter : public Catch::StreamingReporterBase { // Accumulate failure messages for the current test case std::vector m_currentFailures; public: using StreamingReporterBase::StreamingReporterBase; static std::string getDescription() { return "Fixed-width table reporter with detailed assertion failures."; } void testRunStarting(Catch::TestRunInfo const& _testRunInfo) override { StreamingReporterBase::testRunStarting(_testRunInfo); std::cout << '\n'; std::cout << std::left << std::setw(55) << "Test Case Name" << "Status " << std::right << std::setw(8) << "Passed" << std::setw(8) << "Failed" << '\n'; std::cout << std::string(81, '-') << '\n'; } // 1. Hook into every assertion to catch failures void assertionEnded(Catch::AssertionStats const& assertionStats) override { StreamingReporterBase::assertionEnded(assertionStats); // If the assertion failed, build a detailed message if (!assertionStats.assertionResult.isOk()) { auto const& result = assertionStats.assertionResult; std::ostringstream oss; // Format: -> FAILED: [file:line] oss << " \033[31m-> FAILED:\033[0m " << result.getSourceInfo().file << ":" << result.getSourceInfo().line << '\n'; // Print the macro used (e.g., REQUIRE, CHECK) and the expression oss << " " << result.getTestMacroName() << "( " << result.getExpression() << " )\n"; // Print what it actually evaluated to (e.g., 1 == 2) if (result.hasExpandedExpression()) { oss << " with expansion:\n" << " " << result.getExpandedExpression() << '\n'; } // Capture any INFO() messages attached to this assertion for (auto const& msg : assertionStats.infoMessages) { oss << " info: " << msg.message << '\n'; } m_currentFailures.push_back(oss.str()); } } void testCaseEnded(Catch::TestCaseStats const& stats) override { StreamingReporterBase::testCaseEnded(stats); bool passed = stats.totals.assertions.allPassed(); std::string mark = passed ? "\033[32m✓\033[0m" : "\033[31m✗\033[0m"; std::string name = stats.testInfo->name; if (name.length() > 53) { name = name.substr(0, 50) + "..."; } // Print the table row std::cout << std::left << std::setw(55) << name << mark << " " << std::right << std::setw(8) << stats.totals.assertions.passed << std::setw(8) << stats.totals.assertions.failed << '\n'; // 2. Print all accumulated failures under the row if (!m_currentFailures.empty()) { std::cout << '\n'; for (auto const& failure : m_currentFailures) { std::cout << failure << '\n'; } // Add a separator so multiple failing tests don't blur together std::cout << std::string(81, '-') << '\n'; // Clear the buffer for the next test case m_currentFailures.clear(); } } void testRunEnded(Catch::TestRunStats const& _testRunStats) override { StreamingReporterBase::testRunEnded(_testRunStats); std::cout << std::string(81, '=') << '\n'; auto const& tc = _testRunStats.totals.testCases; auto const& as = _testRunStats.totals.assertions; std::string tc_passed_str = tc.passed > 0 ? "\033[32m" + std::to_string(tc.passed) + " passed\033[0m" : "0 passed"; std::string tc_failed_str = tc.failed > 0 ? "\033[31m" + std::to_string(tc.failed) + " failed\033[0m" : "0 failed"; std::string as_passed_str = as.passed > 0 ? "\033[32m" + std::to_string(as.passed) + " passed\033[0m" : "0 passed"; std::string as_failed_str = as.failed > 0 ? "\033[31m" + std::to_string(as.failed) + " failed\033[0m" : "0 failed"; std::cout << "Test Cases: " << tc_passed_str << ", " << tc_failed_str << ", " << tc.total() << " total\n"; std::cout << "Assertions: " << as_passed_str << ", " << as_failed_str << ", " << as.total() << " total\n\n"; } }; CATCH_REGISTER_REPORTER("check", CheckReporter) int main(int argc, char* argv[]) { fourdst::config::Config cfg; CLI::App app{"Mean Field Tests"}; app.allow_extras(); app.set_help_flag("--config-help", "Show mean-field configuration options"); fourdst::config::register_as_cli(cfg, app); std::vector config_arguments; std::vector forced_catch_arguments; config_arguments.emplace_back(argv[0]); bool parsing_catch_arguments = false; for (int i = 1; i < argc; ++i) { if (std::string_view(argv[i]) == "--catch2") { parsing_catch_arguments = true; continue; } if (parsing_catch_arguments) { forced_catch_arguments.emplace_back(argv[i]); } else { config_arguments.emplace_back(argv[i]); } } std::vector config_argv; config_argv.reserve(config_arguments.size()); for (const std::string& argument : config_arguments) { config_argv.push_back(argument.c_str()); } try { app.parse(static_cast(config_argv.size()), config_argv.data()); } catch (const CLI::ParseError& error) { return app.exit(error); } std::vector catch_arguments; catch_arguments.emplace_back(argv[0]); for (const std::string& argument : app.remaining()) { catch_arguments.push_back(argument); } for (const std::string& argument : forced_catch_arguments) { catch_arguments.push_back(argument); } const auto is_reporter_option = [](const std::string& argument) { return argument == "-r" || argument == "--reporter" || argument.starts_with("-r=") || argument.starts_with("--reporter="); }; if (const bool has_reporter = std::ranges::any_of(catch_arguments, is_reporter_option); !has_reporter) { catch_arguments.emplace_back("--reporter"); catch_arguments.emplace_back("check"); } std::vector catch_argv; catch_argv.reserve(catch_arguments.size()); for (const std::string& argument : catch_arguments) { catch_argv.push_back(argument.c_str()); } Catch::Session session; if (const int catch_parse_result = session.applyCommandLine(static_cast(catch_argv.size()), catch_argv.data()); catch_parse_result != 0) { return catch_parse_result; } mfem::Mpi::Init(argc, argv); constexpr std::string device_config = "cpu"; mfem::Device device(device_config); const int hdiv_max_q1d = mfem::DeviceDofQuadLimits::Get().HDIV_MAX_Q1D; std::cout << "H(div) maximum Q1D = " << hdiv_max_q1d << '\n'; std::cout << "Approximate maximum safe integration order = " << 2 * hdiv_max_q1d - 1 << '\n'; mean_field::utils::Args test_args = cfg.main(); if (app.count("--mesh_file") == 0) { test_args.mesh_file = "sandbox.smesh"; } if (app.count("--p.rtol") == 0) { test_args.p.rtol = 1.0e-12; } if (app.count("--p.atol") == 0) { test_args.p.atol = 1.0e-12; } test_utils::set_args(std::move(test_args)); return session.run(); }