diff --git a/src/sm/database/database.cpp b/src/sm/database/database.cpp index 30d9f53d4..7f3de3a54 100644 --- a/src/sm/database/database.cpp +++ b/src/sm/database/database.cpp @@ -812,6 +812,57 @@ Error Database::RemoveTrafficMonitorData(const String& chain) return ErrorEnum::eNone; } +Error Database::BeginTransaction() +{ + std::lock_guard lock {mMutex}; + + LOG_DBG() << "Begin transaction"; + + try { + if (!mSession->isTransaction()) { + mSession->begin(); + } + } catch (const std::exception& e) { + return AOS_ERROR_WRAP(common::utils::ToAosError(e)); + } + + return ErrorEnum::eNone; +} + +Error Database::CommitTransaction() +{ + std::lock_guard lock {mMutex}; + + LOG_DBG() << "Commit transaction"; + + try { + if (mSession->isTransaction()) { + mSession->commit(); + } + } catch (const std::exception& e) { + return AOS_ERROR_WRAP(common::utils::ToAosError(e)); + } + + return ErrorEnum::eNone; +} + +Error Database::RollbackTransaction() +{ + std::lock_guard lock {mMutex}; + + LOG_DBG() << "Rollback transaction"; + + try { + if (mSession->isTransaction()) { + mSession->rollback(); + } + } catch (const std::exception& e) { + return AOS_ERROR_WRAP(common::utils::ToAosError(e)); + } + + return ErrorEnum::eNone; +} + Error Database::AddInstanceNetworkInfo(const sm::networkmanager::InstanceNetworkInfo& info) { std::lock_guard lock {mMutex}; diff --git a/src/sm/database/database.hpp b/src/sm/database/database.hpp index 8862b0a40..2c76d9ad6 100644 --- a/src/sm/database/database.hpp +++ b/src/sm/database/database.hpp @@ -208,6 +208,28 @@ class Database : public DatabaseItf, public sm::alerts::StorageItf { */ Error RemoveTrafficMonitorData(const String& chain) override; + /** + * Begins a SQLite transaction so subsequent writes commit together (one + * fsync) on CommitTransaction. + * + * @return Error. + */ + Error BeginTransaction() override; + + /** + * Commits the current SQLite transaction. + * + * @return Error. + */ + Error CommitTransaction() override; + + /** + * Rolls back the current SQLite transaction, discarding its writes. + * + * @return Error. + */ + Error RollbackTransaction() override; + // sm::alerts::StorageItf interface /** diff --git a/src/sm/database/tests/database.cpp b/src/sm/database/tests/database.cpp index 77728b8c2..a93033ba2 100644 --- a/src/sm/database/tests/database.cpp +++ b/src/sm/database/tests/database.cpp @@ -622,6 +622,102 @@ TEST_F(DatabaseTest, SetUpdateAndRemoveTrafficMonitorDataSucceeds) ASSERT_TRUE(mDB.GetTrafficMonitorData(chain, resTime, resValue).Is(aos::ErrorEnum::eNotFound)); } +TEST_F(DatabaseTest, TransactionCommitPersistsWrites) +{ + ASSERT_TRUE(mDB.Init(mWorkingDir.string(), mMigrationConfig).IsNone()); + + aos::sm::networkmanager::InstanceNetworkInfo info1; + info1.mInstanceID = "instance-1"; + info1.mNetworkID = "network-1"; + info1.mHostIfName = "veth-initial"; + + aos::sm::networkmanager::InstanceNetworkInfo info2; + info2.mInstanceID = "instance-2"; + info2.mNetworkID = "network-2"; + + ASSERT_TRUE(mDB.BeginTransaction().IsNone()); + + ASSERT_TRUE(mDB.AddInstanceNetworkInfo(info1).IsNone()); + + info1.mHostIfName = "veth-updated"; + + ASSERT_TRUE(mDB.UpdateInstanceNetworkInfo(info1).IsNone()); + ASSERT_TRUE(mDB.AddInstanceNetworkInfo(info2).IsNone()); + + ASSERT_TRUE(mDB.CommitTransaction().IsNone()); + + aos::StaticArray result; + + ASSERT_TRUE(mDB.GetInstanceNetworksInfo(result).IsNone()); + + ASSERT_EQ(result.Size(), 2); + + EXPECT_EQ(result[0].mInstanceID, info1.mInstanceID); + EXPECT_EQ(result[0].mHostIfName, info1.mHostIfName); + EXPECT_EQ(result[1].mInstanceID, info2.mInstanceID); +} + +TEST_F(DatabaseTest, TransactionRollbackDiscardsWrites) +{ + ASSERT_TRUE(mDB.Init(mWorkingDir.string(), mMigrationConfig).IsNone()); + + aos::sm::networkmanager::InstanceNetworkInfo committed; + committed.mInstanceID = "instance-0"; + committed.mNetworkID = "network-0"; + + ASSERT_TRUE(mDB.AddInstanceNetworkInfo(committed).IsNone()); + + aos::sm::networkmanager::InstanceNetworkInfo staged; + staged.mInstanceID = "instance-1"; + staged.mNetworkID = "network-1"; + + ASSERT_TRUE(mDB.BeginTransaction().IsNone()); + + ASSERT_TRUE(mDB.AddInstanceNetworkInfo(staged).IsNone()); + ASSERT_TRUE(mDB.SetTrafficMonitorData("chain", aos::Time::Now(), 100).IsNone()); + + ASSERT_TRUE(mDB.RollbackTransaction().IsNone()); + + aos::StaticArray result; + + ASSERT_TRUE(mDB.GetInstanceNetworksInfo(result).IsNone()); + + ASSERT_EQ(result.Size(), 1); + EXPECT_EQ(result[0].mInstanceID, committed.mInstanceID); + + aos::Time resTime; + uint64_t resValue = 0; + + EXPECT_TRUE(mDB.GetTrafficMonitorData("chain", resTime, resValue).Is(aos::ErrorEnum::eNotFound)); +} + +TEST_F(DatabaseTest, TransactionCommitAndRollbackWithoutBeginAreNoOp) +{ + ASSERT_TRUE(mDB.Init(mWorkingDir.string(), mMigrationConfig).IsNone()); + + EXPECT_TRUE(mDB.CommitTransaction().IsNone()); + EXPECT_TRUE(mDB.RollbackTransaction().IsNone()); + + aos::sm::networkmanager::InstanceNetworkInfo info; + info.mInstanceID = "instance-1"; + info.mNetworkID = "network-1"; + + ASSERT_TRUE(mDB.AddInstanceNetworkInfo(info).IsNone()); + + ASSERT_TRUE(mDB.BeginTransaction().IsNone()); + ASSERT_TRUE(mDB.CommitTransaction().IsNone()); + + EXPECT_TRUE(mDB.CommitTransaction().IsNone()); + EXPECT_TRUE(mDB.RollbackTransaction().IsNone()); + + aos::StaticArray result; + + ASSERT_TRUE(mDB.GetInstanceNetworksInfo(result).IsNone()); + + ASSERT_EQ(result.Size(), 1); + EXPECT_EQ(result[0].mInstanceID, info.mInstanceID); +} + /*********************************************************************************************************************** * Tests - alerts::StorageItf **********************************************************************************************************************/ diff --git a/src/sm/networkmanager/bandwidth.cpp b/src/sm/networkmanager/bandwidth.cpp index 5ba25a441..cb59d6e17 100644 --- a/src/sm/networkmanager/bandwidth.cpp +++ b/src/sm/networkmanager/bandwidth.cpp @@ -122,31 +122,13 @@ Error Bandwidth::Clear(const String& ifName) { LOG_DBG() << "Clear bandwidth" << Log::Field("ifName", ifName); - Error err; - - if (auto rootErr = mTC->DelRootTBFQDisc(ifName); !rootErr.IsNone()) { - LOG_ERR() << "Failed to delete root TBF qdisc" << Log::Field(rootErr); - - err = AOS_ERROR_WRAP(rootErr); - } - - if (auto ingErr = mTC->DelIngressQDisc(ifName); !ingErr.IsNone()) { - LOG_ERR() << "Failed to delete ingress qdisc" << Log::Field(ingErr); - - if (err.IsNone()) { - err = AOS_ERROR_WRAP(ingErr); - } - } - - if (auto ifbErr = mIfMgr->DeleteLink(IFBName(ifName)); !ifbErr.IsNone() && !ifbErr.Is(ErrorEnum::eNotFound)) { - LOG_ERR() << "Failed to delete IFB" << Log::Field(ifbErr); + if (auto err = mIfMgr->DeleteLink(IFBName(ifName)); !err.IsNone() && !err.Is(ErrorEnum::eNotFound)) { + LOG_ERR() << "Failed to delete IFB" << Log::Field(err); - if (err.IsNone()) { - err = AOS_ERROR_WRAP(ifbErr); - } + return AOS_ERROR_WRAP(err); } - return err; + return ErrorEnum::eNone; } /*********************************************************************************************************************** diff --git a/src/sm/networkmanager/firewall.cpp b/src/sm/networkmanager/firewall.cpp index 5a76c76fd..984f9d13d 100644 --- a/src/sm/networkmanager/firewall.cpp +++ b/src/sm/networkmanager/firewall.cpp @@ -226,6 +226,12 @@ Error Firewall::Start() } } + { + std::lock_guard lock {mBatchMutex}; + + mInstanceJumps.clear(); + } + mMasqueradeRules.clear(); return ErrorEnum::eNone; @@ -322,6 +328,12 @@ Error Firewall::Stop() // Keep the table and base chains (they outlive SM); drop only the // per-instance state we added. Nothing to do if the table is already gone. if (auto err = mBackend->ListChainRules(mTable, cForwardChain, forwardRules); !err.IsNone()) { + { + std::lock_guard lock {mBatchMutex}; + + mInstanceJumps.clear(); + } + mMasqueradeRules.clear(); return ErrorEnum::eNone; @@ -331,6 +343,12 @@ Error Firewall::Stop() return AOS_ERROR_WRAP(err); } + { + std::lock_guard lock {mBatchMutex}; + + mInstanceJumps.clear(); + } + mMasqueradeRules.clear(); return ErrorEnum::eNone; @@ -433,11 +451,47 @@ Error Firewall::AddInstance(const String& instanceID, const InstanceFirewallPara const auto chain = ChainName(instanceID); + { + std::lock_guard lock {mBatchMutex}; + + if (mBatchMode && mBatchTxn) { + if (auto err = AppendInstanceChain(*mBatchTxn, chain, params); !err.IsNone()) { + return AOS_ERROR_WRAP(err); + } + + mBatchChains.insert(chain); + + return ErrorEnum::eNone; + } + } + auto txn = mBackend->NewTxn(); - txn->AddChain({mTable, chain}); + if (auto err = AppendInstanceChain(*txn, chain, params); !err.IsNone()) { + return AOS_ERROR_WRAP(err); + } - if (auto err = AppendInstanceRules(*txn, mTable, chain, params); !err.IsNone()) { + std::vector handles; + + if (auto err = txn->Commit(handles); !err.IsNone()) { + return AOS_ERROR_WRAP(err); + } + + if (handles.size() >= 2) { + std::lock_guard lock {mBatchMutex}; + + mInstanceJumps[chain] = {handles[handles.size() - 2], handles[handles.size() - 1]}; + } + + return ErrorEnum::eNone; +} + +Error Firewall::AppendInstanceChain( + nftables::FWTxnItf& txn, const std::string& chain, const InstanceFirewallParams& params) +{ + txn.AddChain({mTable, chain}); + + if (auto err = AppendInstanceRules(txn, mTable, chain, params); !err.IsNone()) { return AOS_ERROR_WRAP(err); } @@ -446,7 +500,7 @@ Error Firewall::AddInstance(const String& instanceID, const InstanceFirewallPara jumpIn.mAction = nftables::FWActionEnum::eJump; jumpIn.mJumpTarget = chain; - if (auto err = txn->AddRule(mTable, cForwardChain, jumpIn); !err.IsNone()) { + if (auto err = txn.AddRule(mTable, cForwardChain, jumpIn); !err.IsNone()) { return AOS_ERROR_WRAP(err); } @@ -455,15 +509,22 @@ Error Firewall::AddInstance(const String& instanceID, const InstanceFirewallPara jumpOut.mAction = nftables::FWActionEnum::eJump; jumpOut.mJumpTarget = chain; - if (auto err = txn->AddRule(mTable, cForwardChain, jumpOut); !err.IsNone()) { + if (auto err = txn.AddRule(mTable, cForwardChain, jumpOut); !err.IsNone()) { return AOS_ERROR_WRAP(err); } - if (auto err = txn->Commit(); !err.IsNone()) { - return AOS_ERROR_WRAP(err); + return ErrorEnum::eNone; +} + +void Firewall::DeleteInstanceChain( + nftables::FWTxnItf& txn, const std::string& chain, const std::vector& jumpHandles) +{ + for (const auto handle : jumpHandles) { + txn.DeleteRuleByHandle(mTable, cForwardChain, handle); } - return ErrorEnum::eNone; + txn.FlushChain(mTable, chain); + txn.DeleteChain(mTable, chain); } Error Firewall::RemoveInstance(const String& instanceID) @@ -472,32 +533,183 @@ Error Firewall::RemoveInstance(const String& instanceID) const auto chain = ChainName(instanceID); - std::vector forwardRules; + std::vector jumpHandles; - if (auto err = mBackend->ListChainRules(mTable, cForwardChain, forwardRules); !err.IsNone()) { + { + std::lock_guard lock {mBatchMutex}; + + if (auto it = mInstanceJumps.find(chain); it != mInstanceJumps.end()) { + jumpHandles = {it->second.first, it->second.second}; + + mInstanceJumps.erase(it); + } + } + + if (jumpHandles.empty()) { + std::vector forwardRules; + + if (auto err = mBackend->ListChainRules(mTable, cForwardChain, forwardRules); !err.IsNone()) { + return AOS_ERROR_WRAP(err); + } + + for (const auto& r : forwardRules) { + if (r.mRule.mAction == nftables::FWActionEnum::eJump && r.mRule.mJumpTarget == chain) { + jumpHandles.push_back(r.mHandle); + } + } + + if (jumpHandles.empty()) { + return ErrorEnum::eNone; + } + } + + { + std::lock_guard lock {mBatchMutex}; + + if (mBatchMode && mBatchTxn) { + DeleteInstanceChain(*mBatchTxn, chain, jumpHandles); + + return ErrorEnum::eNone; + } + } + + auto txn = mBackend->NewTxn(); + + DeleteInstanceChain(*txn, chain, jumpHandles); + + if (auto err = txn->Commit(); !err.IsNone()) { return AOS_ERROR_WRAP(err); } - std::vector jumpHandles; + return ErrorEnum::eNone; +} - for (const auto& r : forwardRules) { - if (r.mRule.mAction == nftables::FWActionEnum::eJump && r.mRule.mJumpTarget == chain) { - jumpHandles.push_back(r.mHandle); +Error Firewall::BeginBatch() +{ + LOG_DBG() << "Begin firewall batch"; + + std::lock_guard lock {mBatchMutex}; + + mBatchTxn = mBackend->NewTxn(); + + mBatchChains.clear(); + mAppliedHandles.clear(); + + mBatchMode = true; + + return ErrorEnum::eNone; +} + +Error Firewall::FlushBatch() +{ + LOG_DBG() << "Flush firewall batch"; + + std::unique_ptr txn; + + { + std::lock_guard lock {mBatchMutex}; + + mBatchMode = false; + txn = std::move(mBatchTxn); + } + + if (!txn) { + return ErrorEnum::eNone; + } + + std::vector added; + + const auto err = txn->Commit(added); + + std::lock_guard lock {mBatchMutex}; + + if (!err.IsNone()) { + mBatchChains.clear(); + + return AOS_ERROR_WRAP(err); + } + + std::unordered_map> jumpsByChain; + + for (const auto& r : added) { + mAppliedHandles.insert(r.mHandle); + + if (r.mRule.mAction == nftables::FWActionEnum::eJump) { + jumpsByChain[r.mRule.mJumpTarget].push_back(r.mHandle); } } - if (jumpHandles.empty()) { + for (const auto& [chain, hs] : jumpsByChain) { + if (hs.size() >= 2) { + mInstanceJumps[chain] = {hs[hs.size() - 2], hs[hs.size() - 1]}; + } + } + + return ErrorEnum::eNone; +} + +Error Firewall::AbortBatch() +{ + LOG_DBG() << "Abort firewall batch"; + + std::lock_guard lock {mBatchMutex}; + + mBatchMode = false; + + mBatchTxn.reset(); + + mBatchChains.clear(); + mAppliedHandles.clear(); + + return ErrorEnum::eNone; +} + +Error Firewall::Revert() +{ + LOG_DBG() << "Revert firewall batch"; + + std::set chains; + std::set handles; + + { + std::lock_guard lock {mBatchMutex}; + + chains = std::move(mBatchChains); + handles = std::move(mAppliedHandles); + + mBatchChains.clear(); + mAppliedHandles.clear(); + + for (const auto& chain : chains) { + mInstanceJumps.erase(chain); + } + } + + if (chains.empty() && handles.empty()) { return ErrorEnum::eNone; } + std::vector forwardRules; + + if (auto err = mBackend->ListChainRules(mTable, cForwardChain, forwardRules); !err.IsNone()) { + return AOS_ERROR_WRAP(err); + } + auto txn = mBackend->NewTxn(); - for (const auto handle : jumpHandles) { - txn->DeleteRuleByHandle(mTable, cForwardChain, handle); + for (const auto& r : forwardRules) { + const bool batchJump + = r.mRule.mAction == nftables::FWActionEnum::eJump && chains.count(r.mRule.mJumpTarget) != 0; + + if (batchJump || handles.count(r.mHandle) != 0) { + txn->DeleteRuleByHandle(mTable, cForwardChain, r.mHandle); + } } - txn->FlushChain(mTable, chain); - txn->DeleteChain(mTable, chain); + for (const auto& chain : chains) { + txn->FlushChain(mTable, chain); + txn->DeleteChain(mTable, chain); + } if (auto err = txn->Commit(); !err.IsNone()) { return AOS_ERROR_WRAP(err); @@ -556,10 +768,18 @@ Error Firewall::UpdateInstance(const String& instanceID, const InstanceFirewallP return AOS_ERROR_WRAP(err); } - if (auto err = txn->Commit(); !err.IsNone()) { + std::vector handles; + + if (auto err = txn->Commit(handles); !err.IsNone()) { return AOS_ERROR_WRAP(err); } + if (handles.size() >= 2) { + std::lock_guard lock {mBatchMutex}; + + mInstanceJumps[chain] = {handles[handles.size() - 2], handles[handles.size() - 1]}; + } + return ErrorEnum::eNone; } diff --git a/src/sm/networkmanager/firewall.hpp b/src/sm/networkmanager/firewall.hpp index 8b9810163..92d10f268 100644 --- a/src/sm/networkmanager/firewall.hpp +++ b/src/sm/networkmanager/firewall.hpp @@ -7,8 +7,13 @@ #ifndef AOS_SM_NETWORKMANAGER_FIREWALL_HPP_ #define AOS_SM_NETWORKMANAGER_FIREWALL_HPP_ +#include +#include #include +#include +#include #include +#include #include #include @@ -98,6 +103,37 @@ class Firewall : public FirewallItf, private NonCopyable { */ Error RemoveMasquerade(const String& subnet, const String& outIf) override; + /** + * Begins batch mode: AddInstance/RemoveInstance stage their nft operations + * into a single shared transaction instead of committing per instance. + * + * @return Error. + */ + Error BeginBatch() override; + + /** + * Commits the staged batch in one nft transaction, records the handles it + * added and leaves batch mode. + * + * @return Error. + */ + Error FlushBatch() override; + + /** + * Discards the staged batch and leaves batch mode without applying anything. + * + * @return Error. + */ + Error AbortBatch() override; + + /** + * Deletes by handle everything the last flushed batch added, together with + * the instance chains it created. + * + * @return Error. + */ + Error Revert() override; + private: static constexpr auto cTableName = "aos"; static constexpr auto cForwardChain = "forward"; @@ -110,10 +146,21 @@ class Firewall : public FirewallItf, private NonCopyable { Error CreateSkeleton(); Error ReconcileArtifacts(const std::vector& forwardRules); + Error AppendInstanceChain(nftables::FWTxnItf& txn, const std::string& chain, const InstanceFirewallParams& params); + void DeleteInstanceChain( + nftables::FWTxnItf& txn, const std::string& chain, const std::vector& jumpHandles); const std::string mTable {cTableName}; nftables::FWBackendItf* mBackend {}; std::set> mMasqueradeRules; + + std::mutex mBatchMutex; + bool mBatchMode {false}; + std::unique_ptr mBatchTxn; + std::set mBatchChains; + std::set mAppliedHandles; + + std::unordered_map> mInstanceJumps; }; } // namespace aos::sm::networkmanager diff --git a/src/sm/networkmanager/tests/bandwidth.cpp b/src/sm/networkmanager/tests/bandwidth.cpp index 1f032362e..ca5505c5c 100644 --- a/src/sm/networkmanager/tests/bandwidth.cpp +++ b/src/sm/networkmanager/tests/bandwidth.cpp @@ -169,10 +169,8 @@ TEST_F(BandwidthTest, ApplyRollsBackOnTBFFailure) * Clear **********************************************************************************************************************/ -TEST_F(BandwidthTest, ClearRemovesEverything) +TEST_F(BandwidthTest, ClearRemovesIFB) { - EXPECT_CALL(mTC, DelRootTBFQDisc(String(cHostIfName))).WillOnce(Return(ErrorEnum::eNone)); - EXPECT_CALL(mTC, DelIngressQDisc(String(cHostIfName))).WillOnce(Return(ErrorEnum::eNone)); EXPECT_CALL(mIfMgr, DeleteLink(String(cIFBName))).WillOnce(Return(ErrorEnum::eNone)); EXPECT_TRUE(mBandwidth.Clear(cHostIfName).IsNone()); @@ -180,18 +178,14 @@ TEST_F(BandwidthTest, ClearRemovesEverything) TEST_F(BandwidthTest, ClearTreatsMissingIFBAsSuccess) { - EXPECT_CALL(mTC, DelRootTBFQDisc(String(cHostIfName))).WillOnce(Return(ErrorEnum::eNone)); - EXPECT_CALL(mTC, DelIngressQDisc(String(cHostIfName))).WillOnce(Return(ErrorEnum::eNone)); EXPECT_CALL(mIfMgr, DeleteLink(String(cIFBName))).WillOnce(Return(Error(ErrorEnum::eNotFound))); EXPECT_TRUE(mBandwidth.Clear(cHostIfName).IsNone()); } -TEST_F(BandwidthTest, ClearRunsEveryStepOnFailure) +TEST_F(BandwidthTest, ClearFailsWhenIFBDeleteFails) { - EXPECT_CALL(mTC, DelRootTBFQDisc(String(cHostIfName))).WillOnce(Return(Error(ErrorEnum::eFailed))); - EXPECT_CALL(mTC, DelIngressQDisc(String(cHostIfName))).WillOnce(Return(ErrorEnum::eNone)); - EXPECT_CALL(mIfMgr, DeleteLink(String(cIFBName))).WillOnce(Return(ErrorEnum::eNone)); + EXPECT_CALL(mIfMgr, DeleteLink(String(cIFBName))).WillOnce(Return(Error(ErrorEnum::eFailed))); EXPECT_FALSE(mBandwidth.Clear(cHostIfName).IsNone()); } @@ -225,8 +219,6 @@ TEST_F(BandwidthTest, ApplyAndClearShareIFBName) StaticString capturedOnClear; - EXPECT_CALL(mTC, DelRootTBFQDisc(_)).WillOnce(Return(ErrorEnum::eNone)); - EXPECT_CALL(mTC, DelIngressQDisc(_)).WillOnce(Return(ErrorEnum::eNone)); EXPECT_CALL(mIfMgr, DeleteLink(_)).WillOnce([&capturedOnClear](const String& name) { capturedOnClear = name; return ErrorEnum::eNone; diff --git a/src/sm/networkmanager/tests/firewall.cpp b/src/sm/networkmanager/tests/firewall.cpp index e9ba0de25..7012f880c 100644 --- a/src/sm/networkmanager/tests/firewall.cpp +++ b/src/sm/networkmanager/tests/firewall.cpp @@ -355,7 +355,7 @@ TEST_F(FirewallTest, AddInstanceInputRulesTranslated) EXPECT_CALL(*mTxnPtr, AddRule(_, std::string("instance_test"), TerminalOutRule(std::string("10.0.0.5"), FWActionEnum::eAccept))); EXPECT_CALL(*mTxnPtr, AddRule(_, std::string("forward"), _)).Times(2); - EXPECT_CALL(*mTxnPtr, Commit()).WillOnce(Return(ErrorEnum::eNone)); + EXPECT_CALL(*mTxnPtr, Commit(An&>())).WillOnce(Return(ErrorEnum::eNone)); EXPECT_TRUE(mFirewall.AddInstance("test", params).IsNone()); } @@ -385,7 +385,7 @@ TEST_F(FirewallTest, AddInstanceOutputRulesTranslated) EXPECT_CALL(*mTxnPtr, AddRule(_, std::string("instance_test"), TerminalOutRule(std::string("10.0.0.5"), FWActionEnum::eAccept))); EXPECT_CALL(*mTxnPtr, AddRule(_, std::string("forward"), _)).Times(2); - EXPECT_CALL(*mTxnPtr, Commit()).WillOnce(Return(ErrorEnum::eNone)); + EXPECT_CALL(*mTxnPtr, Commit(An&>())).WillOnce(Return(ErrorEnum::eNone)); EXPECT_TRUE(mFirewall.AddInstance("test", params).IsNone()); } @@ -404,7 +404,7 @@ TEST_F(FirewallTest, AddInstanceDenyPublicProducesDrop) EXPECT_CALL(*mTxnPtr, AddRule(_, std::string("instance_test"), TerminalOutRule(std::string("10.0.0.5"), FWActionEnum::eDrop))); EXPECT_CALL(*mTxnPtr, AddRule(_, std::string("forward"), _)).Times(2); - EXPECT_CALL(*mTxnPtr, Commit()).WillOnce(Return(ErrorEnum::eNone)); + EXPECT_CALL(*mTxnPtr, Commit(An&>())).WillOnce(Return(ErrorEnum::eNone)); EXPECT_TRUE(mFirewall.AddInstance("test", params).IsNone()); } @@ -433,7 +433,7 @@ TEST_F(FirewallTest, AddInstanceInstallsBothJumpsInForward) AllOf(Field(&FWRule::mSrcAddr, "10.0.0.5"), Field(&FWRule::mDstAddr, ""), Field(&FWRule::mAction, FWActionEnum::eJump), Field(&FWRule::mJumpTarget, "instance_test")))); - EXPECT_CALL(*mTxnPtr, Commit()).WillOnce(Return(ErrorEnum::eNone)); + EXPECT_CALL(*mTxnPtr, Commit(An&>())).WillOnce(Return(ErrorEnum::eNone)); EXPECT_TRUE(mFirewall.AddInstance("test", params).IsNone()); } @@ -463,7 +463,7 @@ TEST_F(FirewallTest, AddInstanceSameNetworkAcceptsIntraSubnetBeforeAccessRules) EXPECT_CALL(*mTxnPtr, AddRule(_, std::string("instance_test"), TerminalOutRule(std::string("10.0.0.5"), FWActionEnum::eAccept))); EXPECT_CALL(*mTxnPtr, AddRule(_, std::string("forward"), _)).Times(2); - EXPECT_CALL(*mTxnPtr, Commit()).WillOnce(Return(ErrorEnum::eNone)); + EXPECT_CALL(*mTxnPtr, Commit(An&>())).WillOnce(Return(ErrorEnum::eNone)); EXPECT_TRUE(mFirewall.AddInstance("test", params).IsNone()); } @@ -484,7 +484,7 @@ TEST_F(FirewallTest, AddInstanceNoSubnetSkipsSameNetworkAccepts) EXPECT_CALL(*mTxnPtr, AddRule(_, std::string("instance_test"), TerminalOutRule(std::string("10.0.0.5"), FWActionEnum::eAccept))); EXPECT_CALL(*mTxnPtr, AddRule(_, std::string("forward"), _)).Times(2); - EXPECT_CALL(*mTxnPtr, Commit()).WillOnce(Return(ErrorEnum::eNone)); + EXPECT_CALL(*mTxnPtr, Commit(An&>())).WillOnce(Return(ErrorEnum::eNone)); EXPECT_TRUE(mFirewall.AddInstance("test", params).IsNone()); } @@ -501,7 +501,7 @@ TEST_F(FirewallTest, AddInstanceSanitisesInstanceID) EXPECT_CALL(*mTxnPtr, AddRule(_, std::string("instance_abc_123_de"), _)).Times(2); EXPECT_CALL(*mTxnPtr, AddRule(_, std::string("forward"), Field(&FWRule::mJumpTarget, "instance_abc_123_de"))) .Times(2); - EXPECT_CALL(*mTxnPtr, Commit()).WillOnce(Return(ErrorEnum::eNone)); + EXPECT_CALL(*mTxnPtr, Commit(An&>())).WillOnce(Return(ErrorEnum::eNone)); EXPECT_TRUE(mFirewall.AddInstance("abc-123-de", params).IsNone()); } @@ -524,7 +524,7 @@ TEST_F(FirewallTest, AddInstanceDefaultsMissingProtocolToTcp) EXPECT_CALL(*mTxnPtr, AddRule(_, std::string("instance_test"), TerminalOutRule(std::string("10.0.0.5"), FWActionEnum::eAccept))); EXPECT_CALL(*mTxnPtr, AddRule(_, std::string("forward"), _)).Times(2); - EXPECT_CALL(*mTxnPtr, Commit()).WillOnce(Return(ErrorEnum::eNone)); + EXPECT_CALL(*mTxnPtr, Commit(An&>())).WillOnce(Return(ErrorEnum::eNone)); EXPECT_TRUE(mFirewall.AddInstance("test", params).IsNone()); } @@ -674,7 +674,7 @@ TEST_F(FirewallTest, UpdateInstanceFlushesRepopulatesAndRepointsJumps) EXPECT_CALL(*mTxnPtr, AddRule(_, std::string("forward"), AllOf(Field(&FWRule::mSrcAddr, "10.0.0.9"), Field(&FWRule::mJumpTarget, "instance_test")))); - EXPECT_CALL(*mTxnPtr, Commit()).WillOnce(Return(ErrorEnum::eNone)); + EXPECT_CALL(*mTxnPtr, Commit(An&>())).WillOnce(Return(ErrorEnum::eNone)); EXPECT_TRUE(mFirewall.UpdateInstance("test", params).IsNone()); } @@ -718,6 +718,156 @@ TEST_F(FirewallTest, RemoveInstanceNoMatchIsNoOp) EXPECT_TRUE(mFirewall.RemoveInstance("test").IsNone()); } +/*********************************************************************************************************************** + * Batch + **********************************************************************************************************************/ + +TEST_F(FirewallTest, BatchStagesInstancesIntoSingleCommit) +{ + auto tx = NewMockTx(); + + InSequence seq; + EXPECT_CALL(mBackend, NewTxn()).WillOnce(Return(ByMove(std::move(tx)))); + EXPECT_CALL(*mTxnPtr, AddChain(ChainNamed("instance_inst1"))); + EXPECT_CALL(*mTxnPtr, AddRule(_, std::string("instance_inst1"), _)).Times(2); + EXPECT_CALL(*mTxnPtr, AddRule(_, std::string("forward"), Field(&FWRule::mJumpTarget, "instance_inst1"))).Times(2); + EXPECT_CALL(*mTxnPtr, AddChain(ChainNamed("instance_inst2"))); + EXPECT_CALL(*mTxnPtr, AddRule(_, std::string("instance_inst2"), _)).Times(2); + EXPECT_CALL(*mTxnPtr, AddRule(_, std::string("forward"), Field(&FWRule::mJumpTarget, "instance_inst2"))).Times(2); + EXPECT_CALL(*mTxnPtr, Commit(An&>())).WillOnce(Return(ErrorEnum::eNone)); + + ASSERT_TRUE(mFirewall.BeginBatch().IsNone()); + ASSERT_TRUE(mFirewall.AddInstance("inst1", MakeParams("10.0.0.5", true)).IsNone()); + ASSERT_TRUE(mFirewall.AddInstance("inst2", MakeParams("10.0.0.6", true)).IsNone()); + + EXPECT_TRUE(mFirewall.FlushBatch().IsNone()); +} + +TEST_F(FirewallTest, BatchStagesRemoveInstanceDeletes) +{ + std::vector forwardRules; + forwardRules.push_back({{"10.0.0.5", "", "", 0, "", FWActionEnum::eJump, "instance_test"}, FWRuleHandle {11}}); + forwardRules.push_back({{"", "10.0.0.5", "", 0, "", FWActionEnum::eJump, "instance_test"}, FWRuleHandle {12}}); + + auto tx = NewMockTx(); + + InSequence seq; + EXPECT_CALL(mBackend, NewTxn()).WillOnce(Return(ByMove(std::move(tx)))); + EXPECT_CALL(mBackend, ListChainRules(_, std::string("forward"), _)) + .WillOnce(DoAll(SetArgReferee<2>(forwardRules), Return(ErrorEnum::eNone))); + EXPECT_CALL(*mTxnPtr, DeleteRuleByHandle(_, std::string("forward"), FWRuleHandle {11})); + EXPECT_CALL(*mTxnPtr, DeleteRuleByHandle(_, std::string("forward"), FWRuleHandle {12})); + EXPECT_CALL(*mTxnPtr, FlushChain(_, std::string("instance_test"))); + EXPECT_CALL(*mTxnPtr, DeleteChain(_, std::string("instance_test"))); + EXPECT_CALL(*mTxnPtr, Commit(An&>())).WillOnce(Return(ErrorEnum::eNone)); + + ASSERT_TRUE(mFirewall.BeginBatch().IsNone()); + ASSERT_TRUE(mFirewall.RemoveInstance("test").IsNone()); + + EXPECT_TRUE(mFirewall.FlushBatch().IsNone()); +} + +TEST_F(FirewallTest, AddInstanceCommitsImmediatelyAfterFlushBatch) +{ + auto batchTx = NewMockTx(); + auto* batchPtr = mTxnPtr; + + auto directTx = NewMockTx(); + auto* directPtr = mTxnPtr; + + EXPECT_CALL(mBackend, NewTxn()) + .WillOnce(Return(ByMove(std::move(batchTx)))) + .WillOnce(Return(ByMove(std::move(directTx)))); + + EXPECT_CALL(*batchPtr, AddChain(_)); + EXPECT_CALL(*batchPtr, AddRule(_, _, _)).Times(4).WillRepeatedly(Return(ErrorEnum::eNone)); + EXPECT_CALL(*batchPtr, Commit(An&>())).WillOnce(Return(ErrorEnum::eNone)); + + ASSERT_TRUE(mFirewall.BeginBatch().IsNone()); + ASSERT_TRUE(mFirewall.AddInstance("inst1", MakeParams("10.0.0.5", true)).IsNone()); + ASSERT_TRUE(mFirewall.FlushBatch().IsNone()); + + // Batch is over: the next instance gets its own transaction and commits now. + EXPECT_CALL(*directPtr, AddChain(ChainNamed("instance_inst2"))); + EXPECT_CALL(*directPtr, AddRule(_, _, _)).Times(4).WillRepeatedly(Return(ErrorEnum::eNone)); + EXPECT_CALL(*directPtr, Commit(An&>())).WillOnce(Return(ErrorEnum::eNone)); + + EXPECT_TRUE(mFirewall.AddInstance("inst2", MakeParams("10.0.0.6", true)).IsNone()); +} + +TEST_F(FirewallTest, RevertDeletesFlushedHandlesAndBatchChains) +{ + auto batchTx = NewMockTx(); + auto* batchPtr = mTxnPtr; + + auto revertTx = NewMockTx(); + auto* revertPtr = mTxnPtr; + + EXPECT_CALL(mBackend, NewTxn()) + .WillOnce(Return(ByMove(std::move(batchTx)))) + .WillOnce(Return(ByMove(std::move(revertTx)))); + + EXPECT_CALL(*batchPtr, AddChain(_)); + EXPECT_CALL(*batchPtr, AddRule(_, _, _)).Times(4).WillRepeatedly(Return(ErrorEnum::eNone)); + EXPECT_CALL(*batchPtr, Commit(An&>())).WillOnce([](std::vector& added) { + added = { + {{}, FWRuleHandle {100}}, + {{}, FWRuleHandle {101}}, + {{"10.0.0.5", "", "", 0, "", FWActionEnum::eJump, "instance_inst1"}, FWRuleHandle {102}}, + {{"", "10.0.0.5", "", 0, "", FWActionEnum::eJump, "instance_inst1"}, FWRuleHandle {103}}, + }; + + return Error(ErrorEnum::eNone); + }); + + ASSERT_TRUE(mFirewall.BeginBatch().IsNone()); + ASSERT_TRUE(mFirewall.AddInstance("inst1", MakeParams("10.0.0.5", true)).IsNone()); + ASSERT_TRUE(mFirewall.FlushBatch().IsNone()); + + std::vector forwardRules; + forwardRules.push_back({{"10.0.0.5", "", "", 0, "", FWActionEnum::eJump, "instance_inst1"}, FWRuleHandle {102}}); + forwardRules.push_back({{"", "10.0.0.5", "", 0, "", FWActionEnum::eJump, "instance_inst1"}, FWRuleHandle {103}}); + forwardRules.push_back({{"", "10.0.0.9", "", 0, "", FWActionEnum::eJump, "instance_other"}, FWRuleHandle {7}}); + + InSequence seq; + EXPECT_CALL(mBackend, ListChainRules(_, std::string("forward"), _)) + .WillOnce(DoAll(SetArgReferee<2>(forwardRules), Return(ErrorEnum::eNone))); + EXPECT_CALL(*revertPtr, DeleteRuleByHandle(_, std::string("forward"), FWRuleHandle {102})); + EXPECT_CALL(*revertPtr, DeleteRuleByHandle(_, std::string("forward"), FWRuleHandle {103})); + EXPECT_CALL(*revertPtr, FlushChain(_, std::string("instance_inst1"))); + EXPECT_CALL(*revertPtr, DeleteChain(_, std::string("instance_inst1"))); + EXPECT_CALL(*revertPtr, Commit()).WillOnce(Return(ErrorEnum::eNone)); + + EXPECT_TRUE(mFirewall.Revert().IsNone()); +} + +TEST_F(FirewallTest, RevertAfterFailedFlushIsNoOp) +{ + auto tx = NewMockTx(); + + EXPECT_CALL(mBackend, NewTxn()).WillOnce(Return(ByMove(std::move(tx)))); + EXPECT_CALL(*mTxnPtr, AddChain(_)); + EXPECT_CALL(*mTxnPtr, AddRule(_, _, _)).Times(4).WillRepeatedly(Return(ErrorEnum::eNone)); + EXPECT_CALL(*mTxnPtr, Commit(An&>())).WillOnce(Return(Error(ErrorEnum::eFailed))); + + ASSERT_TRUE(mFirewall.BeginBatch().IsNone()); + ASSERT_TRUE(mFirewall.AddInstance("inst1", MakeParams("10.0.0.5", true)).IsNone()); + + EXPECT_FALSE(mFirewall.FlushBatch().IsNone()); + + // The batch was atomic: nothing was applied, so there is nothing to undo. + EXPECT_TRUE(mFirewall.Revert().IsNone()); +} + +TEST_F(FirewallTest, FlushBatchAndRevertWithoutBeginAreNoOp) +{ + EXPECT_CALL(mBackend, NewTxn()).Times(0); + EXPECT_CALL(mBackend, ListChainRules(_, _, _)).Times(0); + + EXPECT_TRUE(mFirewall.FlushBatch().IsNone()); + EXPECT_TRUE(mFirewall.Revert().IsNone()); +} + /*********************************************************************************************************************** * Masquerade **********************************************************************************************************************/ diff --git a/src/sm/networkmanager/tests/trafficmonitor.cpp b/src/sm/networkmanager/tests/trafficmonitor.cpp index bc7eb960c..b0e98f4ba 100644 --- a/src/sm/networkmanager/tests/trafficmonitor.cpp +++ b/src/sm/networkmanager/tests/trafficmonitor.cpp @@ -190,6 +190,191 @@ TEST_F(TrafficMonitorTest, StopInstanceMonitoring) EXPECT_EQ(mMonitor->StopInstanceMonitoring("test-instance"), ErrorEnum::eNone); } +TEST_F(TrafficMonitorTest, BatchStagesInstancesIntoSingleCommit) +{ + ExpectInit(); + ASSERT_EQ(mMonitor->Init(*mStorage, *mBackend), ErrorEnum::eNone); + + auto txn = MakeTxn(); + EXPECT_CALL(*txn, AddChain(_)).Times(4); + EXPECT_CALL(*txn, AddRule(std::string(cTable), std::string("in_inst1"), _)).Times(AtLeast(1)); + EXPECT_CALL(*txn, AddRule(std::string(cTable), std::string("out_inst1"), _)).Times(AtLeast(1)); + EXPECT_CALL(*txn, AddRule(std::string(cTable), std::string("in_inst2"), _)).Times(AtLeast(1)); + EXPECT_CALL(*txn, AddRule(std::string(cTable), std::string("out_inst2"), _)).Times(AtLeast(1)); + EXPECT_CALL(*txn, AddRule(std::string(cTable), std::string(cForwardChain), JumpTo(std::string("in_inst1")))); + EXPECT_CALL(*txn, AddRule(std::string(cTable), std::string(cForwardChain), JumpTo(std::string("out_inst1")))); + EXPECT_CALL(*txn, AddRule(std::string(cTable), std::string(cForwardChain), JumpTo(std::string("in_inst2")))); + EXPECT_CALL(*txn, AddRule(std::string(cTable), std::string(cForwardChain), JumpTo(std::string("out_inst2")))); + EXPECT_CALL(*txn, Commit()).Times(0); + EXPECT_CALL(*txn, Commit(An&>())).WillOnce(Return(ErrorEnum::eNone)); + + EXPECT_CALL(*mBackend, NewTxn()).WillOnce(Return(ByMove(std::move(txn)))); + EXPECT_CALL(*mStorage, GetTrafficMonitorData(_, _, _)).WillRepeatedly(Return(ErrorEnum::eNotFound)); + + ASSERT_EQ(mMonitor->BeginBatch(), ErrorEnum::eNone); + ASSERT_EQ(mMonitor->StartInstanceMonitoring("inst1", "192.168.1.100", 0, 0), ErrorEnum::eNone); + ASSERT_EQ(mMonitor->StartInstanceMonitoring("inst2", "192.168.1.101", 0, 0), ErrorEnum::eNone); + + EXPECT_EQ(mMonitor->FlushBatch(), ErrorEnum::eNone); +} + +TEST_F(TrafficMonitorTest, BatchStagesStopInstanceDeletes) +{ + ExpectInit(); + ASSERT_EQ(mMonitor->Init(*mStorage, *mBackend), ErrorEnum::eNone); + + const std::string inChain = "in_test_instance"; + const std::string outChain = "out_test_instance"; + + auto startTxn = MakeTxn(); + auto batchTxn = MakeTxn(); + + auto* batchPtr = batchTxn.get(); + + EXPECT_CALL(*mBackend, NewTxn()) + .WillOnce(Return(ByMove(std::move(startTxn)))) + .WillOnce(Return(ByMove(std::move(batchTxn)))); + EXPECT_CALL(*mStorage, GetTrafficMonitorData(_, _, _)).WillRepeatedly(Return(ErrorEnum::eNotFound)); + + ASSERT_EQ(mMonitor->StartInstanceMonitoring("test-instance", "192.168.1.100", 0, 0), ErrorEnum::eNone); + + std::vector forwardRules; + forwardRules.push_back({{"", "", "", 0, "", FWActionEnum::eJump, inChain}, FWRuleHandle {10}}); + forwardRules.push_back({{"", "", "", 0, "", FWActionEnum::eJump, outChain}, FWRuleHandle {11}}); + + EXPECT_CALL(*mBackend, ListChainRules(std::string(cTable), std::string(cForwardChain), _)) + .WillOnce(DoAll(SetArgReferee<2>(forwardRules), Return(ErrorEnum::eNone))); + EXPECT_CALL(*batchPtr, DeleteRuleByHandle(std::string(cTable), std::string(cForwardChain), FWRuleHandle {10})); + EXPECT_CALL(*batchPtr, DeleteRuleByHandle(std::string(cTable), std::string(cForwardChain), FWRuleHandle {11})); + EXPECT_CALL(*batchPtr, FlushChain(std::string(cTable), inChain)); + EXPECT_CALL(*batchPtr, DeleteChain(std::string(cTable), inChain)); + EXPECT_CALL(*batchPtr, FlushChain(std::string(cTable), outChain)); + EXPECT_CALL(*batchPtr, DeleteChain(std::string(cTable), outChain)); + EXPECT_CALL(*batchPtr, Commit()).Times(0); + EXPECT_CALL(*batchPtr, Commit(An&>())).WillOnce(Return(ErrorEnum::eNone)); + + ASSERT_EQ(mMonitor->BeginBatch(), ErrorEnum::eNone); + ASSERT_EQ(mMonitor->StopInstanceMonitoring("test-instance"), ErrorEnum::eNone); + + EXPECT_EQ(mMonitor->FlushBatch(), ErrorEnum::eNone); +} + +TEST_F(TrafficMonitorTest, RevertDeletesFlushedHandlesAndClearsInstanceState) +{ + ExpectInit(); + ASSERT_EQ(mMonitor->Init(*mStorage, *mBackend), ErrorEnum::eNone); + + const std::string inChain = "in_inst1"; + const std::string outChain = "out_inst1"; + + auto batchTxn = MakeTxn(); + auto revertTxn = MakeTxn(); + auto retryTxn = MakeTxn(); + + auto* batchPtr = batchTxn.get(); + auto* revertPtr = revertTxn.get(); + auto* retryPtr = retryTxn.get(); + + EXPECT_CALL(*mBackend, NewTxn()) + .WillOnce(Return(ByMove(std::move(batchTxn)))) + .WillOnce(Return(ByMove(std::move(revertTxn)))) + .WillOnce(Return(ByMove(std::move(retryTxn)))); + EXPECT_CALL(*mStorage, GetTrafficMonitorData(_, _, _)).WillRepeatedly(Return(ErrorEnum::eNotFound)); + EXPECT_CALL(*mStorage, SetTrafficMonitorData(_, _, _)).Times(0); + + EXPECT_CALL(*batchPtr, Commit(An&>())) + .WillOnce([inChain, outChain](std::vector& added) { + added = { + {{"", "", "", 0, "", FWActionEnum::eJump, inChain}, FWRuleHandle {20}}, + {{"", "", "", 0, "", FWActionEnum::eJump, outChain}, FWRuleHandle {21}}, + }; + + return Error(ErrorEnum::eNone); + }); + + ASSERT_EQ(mMonitor->BeginBatch(), ErrorEnum::eNone); + ASSERT_EQ(mMonitor->StartInstanceMonitoring("inst1", "192.168.1.100", 0, 0), ErrorEnum::eNone); + ASSERT_EQ(mMonitor->FlushBatch(), ErrorEnum::eNone); + + std::vector forwardRules; + forwardRules.push_back({{"", "", "", 0, "", FWActionEnum::eJump, inChain}, FWRuleHandle {20}}); + forwardRules.push_back({{"", "", "", 0, "", FWActionEnum::eJump, outChain}, FWRuleHandle {21}}); + forwardRules.push_back({{"", "", "", 0, "", FWActionEnum::eJump, "in_other"}, FWRuleHandle {9}}); + + EXPECT_CALL(*mBackend, ListChainRules(std::string(cTable), std::string(cForwardChain), _)) + .WillOnce(DoAll(SetArgReferee<2>(forwardRules), Return(ErrorEnum::eNone))); + EXPECT_CALL(*revertPtr, DeleteRuleByHandle(std::string(cTable), std::string(cForwardChain), FWRuleHandle {20})); + EXPECT_CALL(*revertPtr, DeleteRuleByHandle(std::string(cTable), std::string(cForwardChain), FWRuleHandle {21})); + EXPECT_CALL(*revertPtr, FlushChain(std::string(cTable), inChain)); + EXPECT_CALL(*revertPtr, DeleteChain(std::string(cTable), inChain)); + EXPECT_CALL(*revertPtr, FlushChain(std::string(cTable), outChain)); + EXPECT_CALL(*revertPtr, DeleteChain(std::string(cTable), outChain)); + + ASSERT_EQ(mMonitor->Revert(), ErrorEnum::eNone); + + uint64_t inputTraffic = 0, outputTraffic = 0; + + EXPECT_EQ(mMonitor->GetInstanceTraffic("inst1", inputTraffic, outputTraffic), ErrorEnum::eNotFound); + + // The instance is unknown again, so the per-instance retry really re-applies. + EXPECT_CALL(*retryPtr, AddChain(_)).Times(2); + EXPECT_CALL(*retryPtr, AddRule(std::string(cTable), inChain, _)).Times(AtLeast(1)); + EXPECT_CALL(*retryPtr, AddRule(std::string(cTable), outChain, _)).Times(AtLeast(1)); + EXPECT_CALL(*retryPtr, AddRule(std::string(cTable), std::string(cForwardChain), JumpTo(inChain))); + EXPECT_CALL(*retryPtr, AddRule(std::string(cTable), std::string(cForwardChain), JumpTo(outChain))); + + EXPECT_EQ(mMonitor->StartInstanceMonitoring("inst1", "192.168.1.100", 0, 0), ErrorEnum::eNone); +} + +TEST_F(TrafficMonitorTest, FailedFlushClearsStagedInstanceState) +{ + ExpectInit(); + ASSERT_EQ(mMonitor->Init(*mStorage, *mBackend), ErrorEnum::eNone); + + const std::string inChain = "in_inst1"; + const std::string outChain = "out_inst1"; + + auto batchTxn = MakeTxn(); + auto retryTxn = MakeTxn(); + + auto* batchPtr = batchTxn.get(); + auto* retryPtr = retryTxn.get(); + + EXPECT_CALL(*mBackend, NewTxn()) + .WillOnce(Return(ByMove(std::move(batchTxn)))) + .WillOnce(Return(ByMove(std::move(retryTxn)))); + EXPECT_CALL(*mStorage, GetTrafficMonitorData(_, _, _)).WillRepeatedly(Return(ErrorEnum::eNotFound)); + + EXPECT_CALL(*batchPtr, Commit(An&>())).WillOnce(Return(Error(ErrorEnum::eFailed))); + + ASSERT_EQ(mMonitor->BeginBatch(), ErrorEnum::eNone); + ASSERT_EQ(mMonitor->StartInstanceMonitoring("inst1", "192.168.1.100", 0, 0), ErrorEnum::eNone); + + EXPECT_FALSE(mMonitor->FlushBatch().IsNone()); + + // Nothing was applied, so the per-instance retry must build and commit its + // own transaction instead of returning early for a known instance. + EXPECT_CALL(*retryPtr, AddChain(_)).Times(2); + EXPECT_CALL(*retryPtr, AddRule(std::string(cTable), inChain, _)).Times(AtLeast(1)); + EXPECT_CALL(*retryPtr, AddRule(std::string(cTable), outChain, _)).Times(AtLeast(1)); + EXPECT_CALL(*retryPtr, AddRule(std::string(cTable), std::string(cForwardChain), JumpTo(inChain))); + EXPECT_CALL(*retryPtr, AddRule(std::string(cTable), std::string(cForwardChain), JumpTo(outChain))); + EXPECT_CALL(*retryPtr, Commit()).WillOnce(Return(ErrorEnum::eNone)); + + EXPECT_EQ(mMonitor->StartInstanceMonitoring("inst1", "192.168.1.100", 0, 0), ErrorEnum::eNone); +} + +TEST_F(TrafficMonitorTest, FlushBatchAndRevertWithoutBeginAreNoOp) +{ + ExpectInit(); + ASSERT_EQ(mMonitor->Init(*mStorage, *mBackend), ErrorEnum::eNone); + + EXPECT_CALL(*mBackend, ListChainRules(_, _, _)).Times(0); + + EXPECT_EQ(mMonitor->FlushBatch(), ErrorEnum::eNone); + EXPECT_EQ(mMonitor->Revert(), ErrorEnum::eNone); +} + TEST_F(TrafficMonitorTest, GetSystemData) { ExpectInit(); diff --git a/src/sm/networkmanager/trafficmonitor.cpp b/src/sm/networkmanager/trafficmonitor.cpp index 0989b5061..baeca538c 100644 --- a/src/sm/networkmanager/trafficmonitor.cpp +++ b/src/sm/networkmanager/trafficmonitor.cpp @@ -152,22 +152,35 @@ Error TrafficMonitor::StartInstanceMonitoring( std::string {cOutChainPrefix} + safeID, }; - auto txn = mBackend->NewTxn(); - StagedTrafficData staged; - if (auto err = CreateInstanceChain(*txn, chains.mInChain, true, chains.mIP, cForwardChain, downloadLimit, staged); - !err.IsNone()) { - return AOS_ERROR_WRAP(err); - } + bool batched = false; - if (auto err = CreateInstanceChain(*txn, chains.mOutChain, false, chains.mIP, cForwardChain, uploadLimit, staged); - !err.IsNone()) { - return AOS_ERROR_WRAP(err); + { + std::lock_guard lock {mBatchMutex}; + + if (mBatchMode && mBatchTxn) { + if (auto err = BuildInstanceMonitoring(*mBatchTxn, chains, staged, downloadLimit, uploadLimit); + !err.IsNone()) { + return AOS_ERROR_WRAP(err); + } + + mBatchInstances.emplace_back(instanceID.CStr()); + + batched = true; + } } - if (auto err = txn->Commit(); !err.IsNone()) { - return AOS_ERROR_WRAP(err); + if (!batched) { + auto txn = mBackend->NewTxn(); + + if (auto err = BuildInstanceMonitoring(*txn, chains, staged, downloadLimit, uploadLimit); !err.IsNone()) { + return AOS_ERROR_WRAP(err); + } + + if (auto err = txn->Commit(); !err.IsNone()) { + return AOS_ERROR_WRAP(err); + } } PublishTrafficData(staged); @@ -202,28 +215,45 @@ Error TrafficMonitor::StopInstanceMonitoring(const String& instanceID) chains = it->second; } - std::vector forwardRules; + std::vector jumpHandles; - if (auto err = mBackend->ListChainRules(cTable, cForwardChain, forwardRules); !err.IsNone()) { - return AOS_ERROR_WRAP(err); + if (chains.mInHandle != 0 && chains.mOutHandle != 0) { + jumpHandles = {chains.mInHandle, chains.mOutHandle}; + } else { + std::vector forwardRules; + + if (auto err = mBackend->ListChainRules(cTable, cForwardChain, forwardRules); !err.IsNone()) { + return AOS_ERROR_WRAP(err); + } + + for (const auto& r : forwardRules) { + if (r.mRule.mAction == nftables::FWActionEnum::eJump + && (r.mRule.mJumpTarget == chains.mInChain || r.mRule.mJumpTarget == chains.mOutChain)) { + jumpHandles.push_back(r.mHandle); + } + } } - auto txn = mBackend->NewTxn(); + bool batched = false; - for (const auto& r : forwardRules) { - if (r.mRule.mAction == nftables::FWActionEnum::eJump - && (r.mRule.mJumpTarget == chains.mInChain || r.mRule.mJumpTarget == chains.mOutChain)) { - txn->DeleteRuleByHandle(cTable, cForwardChain, r.mHandle); + { + std::lock_guard lock {mBatchMutex}; + + if (mBatchMode && mBatchTxn) { + DeleteInstanceMonitoring(*mBatchTxn, chains, jumpHandles); + + batched = true; } } - txn->FlushChain(cTable, chains.mInChain); - txn->DeleteChain(cTable, chains.mInChain); - txn->FlushChain(cTable, chains.mOutChain); - txn->DeleteChain(cTable, chains.mOutChain); + if (!batched) { + auto txn = mBackend->NewTxn(); - if (auto err = txn->Commit(); !err.IsNone()) { - return AOS_ERROR_WRAP(err); + DeleteInstanceMonitoring(*txn, chains, jumpHandles); + + if (auto err = txn->Commit(); !err.IsNone()) { + return AOS_ERROR_WRAP(err); + } } { @@ -248,6 +278,193 @@ Error TrafficMonitor::StopInstanceMonitoring(const String& instanceID) return ErrorEnum::eNone; } +Error TrafficMonitor::BeginBatch() +{ + LOG_DBG() << "Begin traffic monitor batch"; + + std::lock_guard lock {mBatchMutex}; + + mBatchTxn = mBackend->NewTxn(); + + mBatchInstances.clear(); + mAppliedHandles.clear(); + + mBatchMode = true; + + return ErrorEnum::eNone; +} + +Error TrafficMonitor::FlushBatch() +{ + LOG_DBG() << "Flush traffic monitor batch"; + + std::unique_ptr txn; + + { + std::lock_guard lock {mBatchMutex}; + + mBatchMode = false; + txn = std::move(mBatchTxn); + } + + if (!txn) { + return ErrorEnum::eNone; + } + + std::vector added; + + const auto err = txn->Commit(added); + + std::unordered_map jumpByChain; + + for (const auto& r : added) { + if (r.mRule.mAction == nftables::FWActionEnum::eJump) { + jumpByChain[r.mRule.mJumpTarget] = r.mHandle; + } + } + + std::vector failedInstances; + std::vector committedInstances; + + { + std::lock_guard lock {mBatchMutex}; + + if (!err.IsNone()) { + failedInstances = std::move(mBatchInstances); + + mBatchInstances.clear(); + } else { + for (const auto& r : added) { + mAppliedHandles.insert(r.mHandle); + } + + committedInstances = mBatchInstances; + } + } + + if (!err.IsNone()) { + DropBatchInstanceState(failedInstances); + + return AOS_ERROR_WRAP(err); + } + + { + std::unique_lock lock {mMutex}; + + for (const auto& instanceID : committedInstances) { + auto it = mInstanceChains.find(instanceID); + if (it == mInstanceChains.end()) { + continue; + } + + if (auto j = jumpByChain.find(it->second.mInChain); j != jumpByChain.end()) { + it->second.mInHandle = j->second; + } + + if (auto j = jumpByChain.find(it->second.mOutChain); j != jumpByChain.end()) { + it->second.mOutHandle = j->second; + } + } + } + + return ErrorEnum::eNone; +} + +Error TrafficMonitor::AbortBatch() +{ + LOG_DBG() << "Abort traffic monitor batch"; + + std::vector instances; + + { + std::lock_guard lock {mBatchMutex}; + + mBatchMode = false; + + mBatchTxn.reset(); + + instances = std::move(mBatchInstances); + + mBatchInstances.clear(); + mAppliedHandles.clear(); + } + + DropBatchInstanceState(instances); + + return ErrorEnum::eNone; +} + +Error TrafficMonitor::Revert() +{ + LOG_DBG() << "Revert traffic monitor batch"; + + std::vector instances; + std::set handles; + + { + std::lock_guard lock {mBatchMutex}; + + instances = std::move(mBatchInstances); + handles = std::move(mAppliedHandles); + + mBatchInstances.clear(); + mAppliedHandles.clear(); + } + + std::vector> reverted; + + { + std::shared_lock lock {mMutex}; + + for (const auto& instanceID : instances) { + if (auto it = mInstanceChains.find(instanceID); it != mInstanceChains.end()) { + reverted.emplace_back(instanceID, it->second); + } + } + } + + if (reverted.empty()) { + return ErrorEnum::eNone; + } + + std::vector forwardRules; + + if (auto err = mBackend->ListChainRules(cTable, cForwardChain, forwardRules); !err.IsNone()) { + return AOS_ERROR_WRAP(err); + } + + std::set chainNames; + + for (const auto& [instanceID, chains] : reverted) { + chainNames.insert(chains.mInChain); + chainNames.insert(chains.mOutChain); + } + + auto txn = mBackend->NewTxn(); + + for (const auto& r : forwardRules) { + const bool batchJump + = r.mRule.mAction == nftables::FWActionEnum::eJump && chainNames.count(r.mRule.mJumpTarget) != 0; + + if (batchJump || handles.count(r.mHandle) != 0) { + txn->DeleteRuleByHandle(cTable, cForwardChain, r.mHandle); + } + } + + for (const auto& chain : chainNames) { + txn->FlushChain(cTable, chain); + txn->DeleteChain(cTable, chain); + } + + if (auto err = txn->Commit(); !err.IsNone()) { + return AOS_ERROR_WRAP(err); + } + + DropBatchInstanceState(instances); + + return ErrorEnum::eNone; +} + Error TrafficMonitor::GetSystemTraffic(uint64_t& inputTraffic, uint64_t& outputTraffic) const { std::shared_lock lock {mMutex}; @@ -360,6 +577,52 @@ Error TrafficMonitor::CreateInstanceChain(nftables::FWTxnItf& txn, const std::st return ErrorEnum::eNone; } +Error TrafficMonitor::BuildInstanceMonitoring(nftables::FWTxnItf& txn, const InstanceChains& chains, + StagedTrafficData& staged, uint64_t downloadLimit, uint64_t uploadLimit) +{ + if (auto err = CreateInstanceChain(txn, chains.mInChain, true, chains.mIP, cForwardChain, downloadLimit, staged); + !err.IsNone()) { + return AOS_ERROR_WRAP(err); + } + + if (auto err = CreateInstanceChain(txn, chains.mOutChain, false, chains.mIP, cForwardChain, uploadLimit, staged); + !err.IsNone()) { + return AOS_ERROR_WRAP(err); + } + + return ErrorEnum::eNone; +} + +void TrafficMonitor::DeleteInstanceMonitoring( + nftables::FWTxnItf& txn, const InstanceChains& chains, const std::vector& jumpHandles) +{ + for (const auto handle : jumpHandles) { + txn.DeleteRuleByHandle(cTable, cForwardChain, handle); + } + + txn.FlushChain(cTable, chains.mInChain); + txn.DeleteChain(cTable, chains.mInChain); + txn.FlushChain(cTable, chains.mOutChain); + txn.DeleteChain(cTable, chains.mOutChain); +} + +void TrafficMonitor::DropBatchInstanceState(const std::vector& instanceIDs) +{ + std::unique_lock lock {mMutex}; + + for (const auto& instanceID : instanceIDs) { + auto it = mInstanceChains.find(instanceID); + if (it == mInstanceChains.end()) { + continue; + } + + mTrafficData.erase(it->second.mInChain); + mTrafficData.erase(it->second.mOutChain); + + mInstanceChains.erase(it); + } +} + void TrafficMonitor::PublishTrafficData(StagedTrafficData& staged) { std::unique_lock lock {mMutex}; diff --git a/src/sm/networkmanager/trafficmonitor.hpp b/src/sm/networkmanager/trafficmonitor.hpp index 545e0b6ef..618ba1fe0 100644 --- a/src/sm/networkmanager/trafficmonitor.hpp +++ b/src/sm/networkmanager/trafficmonitor.hpp @@ -8,7 +8,11 @@ #define AOS_SM_NETWORKMANAGER_TRAFFICMONITOR_HPP_ #include +#include +#include +#include #include +#include #include #include #include @@ -87,6 +91,39 @@ class TrafficMonitor : public TrafficMonitorItf { */ Error GetInstanceTraffic(const String& instanceID, uint64_t& inputTraffic, uint64_t& outputTraffic) const override; + /** + * Begins batch mode: StartInstanceMonitoring/StopInstanceMonitoring stage + * their nft operations into a single shared transaction instead of + * committing per instance. + * + * @return Error. + */ + Error BeginBatch() override; + + /** + * Commits the staged batch in one nft transaction, records the handles it + * added and leaves batch mode. + * + * @return Error. + */ + Error FlushBatch() override; + + /** + * Discards the staged batch and leaves batch mode without applying anything, + * dropping the monitoring state the batch staged. + * + * @return Error. + */ + Error AbortBatch() override; + + /** + * Deletes by handle everything the last flushed batch added, together with + * the counter chains it created, and drops their monitoring state. + * + * @return Error. + */ + Error Revert() override; + private: static constexpr auto cTable = "aos-traffic"; static constexpr auto cInputChain = "input"; @@ -108,9 +145,11 @@ class TrafficMonitor : public TrafficMonitorItf { }; struct InstanceChains { - std::string mIP; - std::string mInChain; - std::string mOutChain; + std::string mIP; + std::string mInChain; + std::string mOutChain; + nftables::FWRuleHandle mInHandle {}; + nftables::FWRuleHandle mOutHandle {}; }; using StagedTrafficData = std::vector>; @@ -119,6 +158,11 @@ class TrafficMonitor : public TrafficMonitorItf { Error DeleteTrafficTable(); Error CreateInstanceChain(nftables::FWTxnItf& txn, const std::string& chain, bool isInChain, const std::string& address, const std::string& parentBaseChain, uint64_t limit, StagedTrafficData& staged); + Error BuildInstanceMonitoring(nftables::FWTxnItf& txn, const InstanceChains& chains, StagedTrafficData& staged, + uint64_t downloadLimit, uint64_t uploadLimit); + void DeleteInstanceMonitoring( + nftables::FWTxnItf& txn, const InstanceChains& chains, const std::vector& jumpHandles); + void DropBatchInstanceState(const std::vector& instanceIDs); void PublishTrafficData(StagedTrafficData& staged); Error AppendChainCounterRules( nftables::FWTxnItf& txn, const std::string& chain, bool isInChain, const std::string& address, bool disabled); @@ -136,10 +180,17 @@ class TrafficMonitor : public TrafficMonitorItf { std::unordered_map mTrafficData {}; std::unordered_map mInstanceChains {}; mutable std::shared_mutex mMutex {}; - aos::Timer mTimer {}; - TrafficPeriod mTrafficPeriod {}; - Duration mUpdatePeriod {}; - bool mStop {}; + + std::mutex mBatchMutex; + bool mBatchMode {false}; + std::unique_ptr mBatchTxn; + std::vector mBatchInstances; + std::set mAppliedHandles; + + aos::Timer mTimer {}; + TrafficPeriod mTrafficPeriod {}; + Duration mUpdatePeriod {}; + bool mStop {}; }; } // namespace aos::sm::networkmanager diff --git a/src/sm/nftables/itf/firewallbackend.hpp b/src/sm/nftables/itf/firewallbackend.hpp index 9af4b9f00..7a186280b 100644 --- a/src/sm/nftables/itf/firewallbackend.hpp +++ b/src/sm/nftables/itf/firewallbackend.hpp @@ -228,6 +228,26 @@ class FWTxnItf { * @return error. */ virtual Error Commit() = 0; + + /** + * Submits the queued batch and returns the handles of the rules it added, + * in the order they were added. Lets the caller record handles for O(1) + * deletion later without re-listing the chain. + * + * @param[out] addedHandles handles of the rules added by this batch. + * @return error. + */ + virtual Error Commit(std::vector& addedHandles) = 0; + + /** + * Submits the queued batch and returns the jump rules it added with their + * handles, so the caller can attribute added jumps to their target chains + * when many instances are committed together. + * + * @param[out] addedRules jump rules added by this batch. + * @return error. + */ + virtual Error Commit(std::vector& addedRules) = 0; }; /** diff --git a/src/sm/nftables/nftables.cpp b/src/sm/nftables/nftables.cpp index 9cf297a4d..8c1e5044b 100644 --- a/src/sm/nftables/nftables.cpp +++ b/src/sm/nftables/nftables.cpp @@ -133,33 +133,47 @@ void AppendRuleExpr(std::ostringstream& buf, const FWRule& rule) bool ParseRuleLine(const std::string& line, FWListedRule& out) { + static const std::regex ctStateRe(R"(ct\s+state\s+([a-z,]+))"); + static const std::regex saddrRe(R"(ip\s+saddr\s+(\S+))"); + static const std::regex daddrRe(R"(ip\s+daddr\s+(\S+))"); + static const std::regex dportRe(R"((tcp|udp)\s+dport\s+(\d+))"); + static const std::regex protoRe(R"(\b(tcp|udp)\b)"); + static const std::regex oifnameRe(R"rx(oifname\s+(!=\s+)?"([^"]+)")rx"); + static const std::regex counterRe(R"(counter\s+packets\s+(\d+)\s+bytes\s+(\d+))"); + static const std::regex jumpRe(R"(\bjump\s+(\S+))"); + static const std::regex masqueradeRe(R"(\bmasquerade\b)"); + static const std::regex acceptRe(R"(\baccept\b)"); + static const std::regex dropRe(R"(\bdrop\b)"); + static const std::regex returnRe(R"(\breturn\b)"); + static const std::regex handleRe(R"(#\s+handle\s+(\d+))"); + std::smatch m; - if (std::regex_search(line, m, std::regex(R"(ct\s+state\s+([a-z,]+))"))) { + if (std::regex_search(line, m, ctStateRe)) { out.mRule.mCtState = m[1]; } - if (std::regex_search(line, m, std::regex(R"(ip\s+saddr\s+(\S+))"))) { + if (std::regex_search(line, m, saddrRe)) { out.mRule.mSrcAddr = m[1]; } - if (std::regex_search(line, m, std::regex(R"(ip\s+daddr\s+(\S+))"))) { + if (std::regex_search(line, m, daddrRe)) { out.mRule.mDstAddr = m[1]; } - if (std::regex_search(line, m, std::regex(R"((tcp|udp)\s+dport\s+(\d+))"))) { + if (std::regex_search(line, m, dportRe)) { out.mRule.mProto = m[1]; out.mRule.mDstPort = static_cast(std::stoi(m[2])); - } else if (std::regex_search(line, m, std::regex(R"(\b(tcp|udp)\b)"))) { + } else if (std::regex_search(line, m, protoRe)) { out.mRule.mProto = m[1]; } - if (std::regex_search(line, m, std::regex(R"rx(oifname\s+(!=\s+)?"([^"]+)")rx"))) { + if (std::regex_search(line, m, oifnameRe)) { out.mRule.mOIFNeg = m[1].matched; out.mRule.mOIFName = m[2]; } - if (std::regex_search(line, m, std::regex(R"(counter\s+packets\s+(\d+)\s+bytes\s+(\d+))"))) { + if (std::regex_search(line, m, counterRe)) { out.mRule.mCounter = true; out.mPackets = std::stoull(m[1]); out.mBytes = std::stoull(m[2]); @@ -167,25 +181,25 @@ bool ParseRuleLine(const std::string& line, FWListedRule& out) bool actionFound = false; - if (std::regex_search(line, m, std::regex(R"(\bjump\s+(\S+))"))) { + if (std::regex_search(line, m, jumpRe)) { out.mRule.mAction = FWActionEnum::eJump; out.mRule.mJumpTarget = m[1]; actionFound = true; - } else if (std::regex_search(line, m, std::regex(R"(\bmasquerade\b)"))) { + } else if (std::regex_search(line, m, masqueradeRe)) { out.mRule.mAction = FWActionEnum::eMasquerade; actionFound = true; - } else if (std::regex_search(line, m, std::regex(R"(\baccept\b)"))) { + } else if (std::regex_search(line, m, acceptRe)) { out.mRule.mAction = FWActionEnum::eAccept; actionFound = true; - } else if (std::regex_search(line, m, std::regex(R"(\bdrop\b)"))) { + } else if (std::regex_search(line, m, dropRe)) { out.mRule.mAction = FWActionEnum::eDrop; actionFound = true; - } else if (std::regex_search(line, m, std::regex(R"(\breturn\b)"))) { + } else if (std::regex_search(line, m, returnRe)) { out.mRule.mAction = FWActionEnum::eReturn; actionFound = true; } - if (!std::regex_search(line, m, std::regex(R"(#\s+handle\s+(\d+))"))) { + if (!std::regex_search(line, m, handleRe)) { return false; } @@ -268,6 +282,34 @@ class NFTables::NFTxn : public FWTxnItf { return mParent.RunBuffer(cmd); } + Error Commit(std::vector& addedHandles) override + { + const auto cmd = mBuf.str(); + + mBuf.str(std::string {}); + mBuf.clear(); + + if (cmd.empty()) { + return ErrorEnum::eNone; + } + + return mParent.RunBufferEcho(cmd, addedHandles); + } + + Error Commit(std::vector& addedRules) override + { + const auto cmd = mBuf.str(); + + mBuf.str(std::string {}); + mBuf.clear(); + + if (cmd.empty()) { + return ErrorEnum::eNone; + } + + return mParent.RunBufferEchoRules(cmd, addedRules); + } + private: NFTables& mParent; std::string mFamily; @@ -368,4 +410,94 @@ Error NFTables::RunBufferWithOutput(const std::string& cmd, std::string& output) return ErrorEnum::eNone; } +Error NFTables::RunBufferEcho(const std::string& cmd, std::vector& handles) +{ + std::lock_guard lock {mMutex}; + + NFTCtxGuard ctx; + if (ctx.Get() == nullptr) { + return AOS_ERROR_WRAP(Error(ErrorEnum::eFailed, "nft_ctx_new failed")); + } + + nft_ctx_output_set_flags(ctx.Get(), nft_ctx_output_get_flags(ctx.Get()) | NFT_CTX_OUTPUT_ECHO); + + if (nft_run_cmd_from_buffer(ctx.Get(), cmd.c_str()) != 0) { + const auto errText = ctx.ErrorBuffer(); + + if (IsNotFoundError(errText)) { + return Error(ErrorEnum::eNotFound, errText.empty() ? "nftables object not found" : errText.c_str()); + } + + LOG_ERR() << "nftables command failed" << Log::Field("cmd", cmd.c_str()) << Log::Field("err", errText.c_str()); + + return AOS_ERROR_WRAP(Error(ErrorEnum::eFailed, errText.empty() ? "nftables command failed" : errText.c_str())); + } + + const auto output = ctx.OutputBuffer(); + std::istringstream iss(output); + std::string line; + const std::regex re(R"(#\s+handle\s+(\d+))"); + + while (std::getline(iss, line)) { + std::smatch m; + + if (std::regex_search(line, m, re)) { + handles.push_back(static_cast(std::stoull(m[1]))); + } + } + + return ErrorEnum::eNone; +} + +Error NFTables::RunBufferEchoRules(const std::string& cmd, std::vector& rules) +{ + std::lock_guard lock {mMutex}; + + NFTCtxGuard ctx; + if (ctx.Get() == nullptr) { + return AOS_ERROR_WRAP(Error(ErrorEnum::eFailed, "nft_ctx_new failed")); + } + + nft_ctx_output_set_flags(ctx.Get(), nft_ctx_output_get_flags(ctx.Get()) | NFT_CTX_OUTPUT_ECHO); + + if (nft_run_cmd_from_buffer(ctx.Get(), cmd.c_str()) != 0) { + const auto errText = ctx.ErrorBuffer(); + + if (IsNotFoundError(errText)) { + return Error(ErrorEnum::eNotFound, errText.empty() ? "nftables object not found" : errText.c_str()); + } + + LOG_ERR() << "nftables command failed" << Log::Field("cmd", cmd.c_str()) << Log::Field("err", errText.c_str()); + + return AOS_ERROR_WRAP(Error(ErrorEnum::eFailed, errText.empty() ? "nftables command failed" : errText.c_str())); + } + + const auto output = ctx.OutputBuffer(); + std::istringstream iss(output); + std::string line; + const std::regex jumpRe(R"(\bjump\s+(\S+))"); + const std::regex handleRe(R"(#\s+handle\s+(\d+))"); + + while (std::getline(iss, line)) { + std::smatch jm; + if (!std::regex_search(line, jm, jumpRe)) { + continue; + } + + std::smatch hm; + if (!std::regex_search(line, hm, handleRe)) { + continue; + } + + FWListedRule listed {}; + listed.mRule.mAction = FWActionEnum::eJump; + listed.mRule.mJumpTarget = jm[1]; + listed.mHandle = static_cast(std::stoull(hm[1])); + + rules.push_back(std::move(listed)); + } + + return ErrorEnum::eNone; +} + } // namespace aos::sm::nftables diff --git a/src/sm/nftables/nftables.hpp b/src/sm/nftables/nftables.hpp index d84297d6f..53633a73b 100644 --- a/src/sm/nftables/nftables.hpp +++ b/src/sm/nftables/nftables.hpp @@ -53,6 +53,8 @@ class NFTables : public FWBackendItf, private NonCopyable { Error RunBuffer(const std::string& cmd); Error RunBufferWithOutput(const std::string& cmd, std::string& output); + Error RunBufferEcho(const std::string& cmd, std::vector& handles); + Error RunBufferEchoRules(const std::string& cmd, std::vector& rules); std::string mFamily; std::mutex mMutex; diff --git a/src/sm/tests/mocks/firewallbackendmock.hpp b/src/sm/tests/mocks/firewallbackendmock.hpp index 28f34de04..40c647859 100644 --- a/src/sm/tests/mocks/firewallbackendmock.hpp +++ b/src/sm/tests/mocks/firewallbackendmock.hpp @@ -25,6 +25,8 @@ class MockFWTxn : public FWTxnItf { MOCK_METHOD(void, DeleteRuleByHandle, (const std::string& table, const std::string& chain, FWRuleHandle handle), (override)); MOCK_METHOD(Error, Commit, (), (override)); + MOCK_METHOD(Error, Commit, (std::vector & addedHandles), (override)); + MOCK_METHOD(Error, Commit, (std::vector & addedRules), (override)); }; class MockFWBackend : public FWBackendItf {