diff --git a/nostrdb/NdbTxn.swift b/nostrdb/NdbTxn.swift index 6141f111..fb93d4d4 100644 --- a/nostrdb/NdbTxn.swift +++ b/nostrdb/NdbTxn.swift @@ -11,18 +11,26 @@ import Foundation fileprivate var txn_count: Int = 0 #endif +/// Standard timeout for nostrdb transaction operations. +/// +/// This timeout prevents indefinite hangs while allowing normal database operations to complete. +/// The value is chosen empirically: long enough for typical database operations on the main thread, +/// short enough to prevent user-visible freezes. +private extension DispatchTimeInterval { + static let ndbTransactionTimeout = DispatchTimeInterval.milliseconds(200) +} + // Would use struct and ~Copyable but generics aren't supported well class NdbTxn: RawNdbTxnAccessible { var txn: ndb_txn private var val: T! var moved: Bool - var inherited: Bool var ndb: Ndb var generation: Int var name: String static func pure(ndb: Ndb, val: T) -> NdbTxn { - .init(ndb: ndb, txn: ndb_txn(), val: val, generation: ndb.generation, inherited: true, name: "pure_txn") + .init(ndb: ndb, txn: ndb_txn(), val: val, generation: ndb.generation, name: "pure_txn") } /// Simple helper struct for the init function to avoid compiler errors encountered by using other techniques @@ -36,36 +44,21 @@ class NdbTxn: RawNdbTxnAccessible { self.name = name ?? "txn" self.ndb = ndb self.generation = ndb.generation - if let active_txn = Thread.current.threadDictionary["ndb_txn"] as? ndb_txn, - let txn_generation = Thread.current.threadDictionary["txn_generation"] as? Int, - txn_generation == ndb.generation - { - // some parent thread is active, use that instead - print("txn: inherited txn") - self.txn = active_txn - self.inherited = true - self.generation = Thread.current.threadDictionary["txn_generation"] as! Int - let ref_count = Thread.current.threadDictionary["ndb_txn_ref_count"] as! Int - let new_ref_count = ref_count + 1 - Thread.current.threadDictionary["ndb_txn_ref_count"] = new_ref_count - } else { - let result: R? = try? ndb.withNdb({ - var txn = ndb_txn() - #if TXNDEBUG - txn_count += 1 - #endif - let ok = ndb_begin_query(ndb.ndb.ndb, &txn) != 0 - guard ok else { return .none } - return .some(R(txn: txn, generation: ndb.generation)) - }, maxWaitTimeout: .milliseconds(200)) - guard let result else { return nil } - self.txn = result.txn - self.generation = result.generation - Thread.current.threadDictionary["ndb_txn"] = self.txn - Thread.current.threadDictionary["ndb_txn_ref_count"] = 1 - Thread.current.threadDictionary["txn_generation"] = ndb.generation - self.inherited = false - } + + // Always create fresh transaction + let result: R? = try? ndb.withNdb({ + var txn = ndb_txn() + let ok = ndb_begin_query(ndb.ndb.ndb, &txn) != 0 + guard ok else { return .none } + #if TXNDEBUG + txn_count += 1 + #endif + return .some(R(txn: txn, generation: ndb.generation)) + }, maxWaitTimeout: .ndbTransactionTimeout) + guard let result else { return nil } + self.txn = result.txn + self.generation = result.generation + #if TXNDEBUG print("txn: open gen\(self.generation) '\(self.name)' \(txn_count)") #endif @@ -73,11 +66,10 @@ class NdbTxn: RawNdbTxnAccessible { self.val = with(self) } - private init(ndb: Ndb, txn: ndb_txn, val: T, generation: Int, inherited: Bool, name: String) { + private init(ndb: Ndb, txn: ndb_txn, val: T, generation: Int, name: String) { self.txn = txn self.val = val self.moved = false - self.inherited = inherited self.ndb = ndb self.generation = generation self.name = name @@ -99,27 +91,15 @@ class NdbTxn: RawNdbTxnAccessible { print("txn: not closing. db closed") return } - if let ref_count = Thread.current.threadDictionary["ndb_txn_ref_count"] as? Int { - let new_ref_count = ref_count - 1 - Thread.current.threadDictionary["ndb_txn_ref_count"] = new_ref_count - assert(new_ref_count >= 0, "NdbTxn reference count should never be below zero") - if new_ref_count <= 0 { - _ = try? ndb.withNdb({ - ndb_end_query(&self.txn) - }, maxWaitTimeout: .milliseconds(200)) - Thread.current.threadDictionary.removeObject(forKey: "ndb_txn") - Thread.current.threadDictionary.removeObject(forKey: "ndb_txn_ref_count") - } - } - if inherited { - print("txn: not closing. inherited ") - return - } if moved { //print("txn: not closing. moved") return } + _ = try? ndb.withNdb({ + ndb_end_query(&self.txn) + }, maxWaitTimeout: .ndbTransactionTimeout) + #if TXNDEBUG txn_count -= 1; print("txn: close gen\(generation) '\(name)' \(txn_count)") @@ -129,14 +109,14 @@ class NdbTxn: RawNdbTxnAccessible { // functor func map(_ transform: (T) -> Y) -> NdbTxn { self.moved = true - return .init(ndb: self.ndb, txn: self.txn, val: transform(val), generation: generation, inherited: inherited, name: self.name) + return .init(ndb: self.ndb, txn: self.txn, val: transform(val), generation: generation, name: self.name) } // comonad!? // useful for moving ownership of a transaction to another value func extend(_ with: (NdbTxn) -> Y) -> NdbTxn { self.moved = true - return .init(ndb: self.ndb, txn: self.txn, val: with(self), generation: generation, inherited: inherited, name: self.name) + return .init(ndb: self.ndb, txn: self.txn, val: with(self), generation: generation, name: self.name) } } @@ -156,13 +136,12 @@ class SafeNdbTxn { var txn: ndb_txn var val: T! var moved: Bool - var inherited: Bool var ndb: Ndb var generation: Int var name: String static func pure(ndb: Ndb, val: consuming T) -> SafeNdbTxn { - .init(ndb: ndb, txn: ndb_txn(), val: val, generation: ndb.generation, inherited: true, name: "pure_txn") + .init(ndb: ndb, txn: ndb_txn(), val: val, generation: ndb.generation, name: "pure_txn") } /// Simple helper struct for the init function to avoid compiler errors encountered by using other techniques @@ -173,52 +152,42 @@ class SafeNdbTxn { static func new(on ndb: Ndb, with valueGetter: (PlaceholderNdbTxn) -> T? = { _ in () }, name: String = "txn") -> SafeNdbTxn? { guard !ndb.is_closed else { return nil } - let generation: Int - let txn: ndb_txn - let inherited: Bool - if let active_txn = Thread.current.threadDictionary["ndb_txn"] as? ndb_txn, - let txn_generation = Thread.current.threadDictionary["txn_generation"] as? Int, - txn_generation == ndb.generation - { - // some parent thread is active, use that instead - print("txn: inherited txn") - txn = active_txn - inherited = true - generation = Thread.current.threadDictionary["txn_generation"] as! Int - let ref_count = Thread.current.threadDictionary["ndb_txn_ref_count"] as! Int - let new_ref_count = ref_count + 1 - Thread.current.threadDictionary["ndb_txn_ref_count"] = new_ref_count - } else { - let result: R? = try? ndb.withNdb({ - var txn = ndb_txn() - #if TXNDEBUG - txn_count += 1 - #endif - let ok = ndb_begin_query(ndb.ndb.ndb, &txn) != 0 - guard ok else { return .none } - return .some(R(txn: txn, generation: ndb.generation)) - }, maxWaitTimeout: .milliseconds(200)) - guard let result else { return nil } - txn = result.txn - generation = result.generation - Thread.current.threadDictionary["ndb_txn"] = txn - Thread.current.threadDictionary["ndb_txn_ref_count"] = 1 - Thread.current.threadDictionary["txn_generation"] = ndb.generation - inherited = false - } + + // Always create fresh transaction + let result: R? = try? ndb.withNdb({ + var txn = ndb_txn() + let ok = ndb_begin_query(ndb.ndb.ndb, &txn) != 0 + guard ok else { return .none } + #if TXNDEBUG + txn_count += 1 + #endif + return .some(R(txn: txn, generation: ndb.generation)) + }, maxWaitTimeout: .ndbTransactionTimeout) + guard let result else { return nil } + let txn = result.txn + let generation = result.generation + #if TXNDEBUG print("txn: open gen\(generation) '\(name)' \(txn_count)") #endif let placeholderTxn = PlaceholderNdbTxn(txn: txn) - guard let val = valueGetter(placeholderTxn) else { return nil } - return SafeNdbTxn(ndb: ndb, txn: txn, val: val, generation: generation, inherited: inherited, name: name) + guard let val = valueGetter(placeholderTxn) else { + // Fix leak: Close transaction before returning nil + var mutableTxn = txn + _ = try? ndb.withNdb({ ndb_end_query(&mutableTxn) }, maxWaitTimeout: .ndbTransactionTimeout) + #if TXNDEBUG + txn_count -= 1 + print("txn: close (valueGetter nil) gen\(generation) '\(name)' \(txn_count)") + #endif + return nil + } + return SafeNdbTxn(ndb: ndb, txn: txn, val: val, generation: generation, name: name) } - private init(ndb: Ndb, txn: ndb_txn, val: consuming T, generation: Int, inherited: Bool, name: String) { + private init(ndb: Ndb, txn: ndb_txn, val: consuming T, generation: Int, name: String) { self.txn = txn self.val = consume val self.moved = false - self.inherited = inherited self.ndb = ndb self.generation = generation self.name = name @@ -233,27 +202,15 @@ class SafeNdbTxn { print("txn: not closing. db closed") return } - if let ref_count = Thread.current.threadDictionary["ndb_txn_ref_count"] as? Int { - let new_ref_count = ref_count - 1 - Thread.current.threadDictionary["ndb_txn_ref_count"] = new_ref_count - assert(new_ref_count >= 0, "NdbTxn reference count should never be below zero") - if new_ref_count <= 0 { - _ = try? ndb.withNdb({ - ndb_end_query(&self.txn) - }, maxWaitTimeout: .milliseconds(200)) - Thread.current.threadDictionary.removeObject(forKey: "ndb_txn") - Thread.current.threadDictionary.removeObject(forKey: "ndb_txn_ref_count") - } - } - if inherited { - print("txn: not closing. inherited ") - return - } if moved { //print("txn: not closing. moved") return } + _ = try? ndb.withNdb({ + ndb_end_query(&self.txn) + }, maxWaitTimeout: .ndbTransactionTimeout) + #if TXNDEBUG txn_count -= 1; print("txn: close gen\(generation) '\(name)' \(txn_count)") @@ -263,14 +220,14 @@ class SafeNdbTxn { // functor func map(_ transform: (borrowing T) -> Y) -> SafeNdbTxn { self.moved = true - return .init(ndb: self.ndb, txn: self.txn, val: transform(val), generation: generation, inherited: inherited, name: self.name) + return .init(ndb: self.ndb, txn: self.txn, val: transform(val), generation: generation, name: self.name) } // comonad!? // useful for moving ownership of a transaction to another value func extend(_ with: (SafeNdbTxn) -> Y) -> SafeNdbTxn { self.moved = true - return .init(ndb: self.ndb, txn: self.txn, val: with(self), generation: generation, inherited: inherited, name: self.name) + return .init(ndb: self.ndb, txn: self.txn, val: with(self), generation: generation, name: self.name) } consuming func maybeExtend(_ with: (consuming SafeNdbTxn) -> Y?) -> SafeNdbTxn? where Y: ~Copyable { @@ -278,10 +235,20 @@ class SafeNdbTxn { let ndb = self.ndb let txn = self.txn let generation = self.generation - let inherited = self.inherited let name = self.name - guard let newVal = with(consume self) else { return nil } - return .init(ndb: ndb, txn: txn, val: newVal, generation: generation, inherited: inherited, name: name) + + guard let newVal = with(consume self) else { + // Fix leak: Close transaction on nil path if we own it + var mutableTxn = txn + _ = try? ndb.withNdb({ ndb_end_query(&mutableTxn) }, maxWaitTimeout: .ndbTransactionTimeout) + #if TXNDEBUG + txn_count -= 1 + print("txn: close (maybeExtend nil) gen\(generation) '\(name)' \(txn_count)") + #endif + return nil + } + + return .init(ndb: ndb, txn: txn, val: newVal, generation: generation, name: name) } } @@ -304,7 +271,7 @@ extension NdbTxn where T: OptionalType { return nil } self.moved = true - return NdbTxn(ndb: self.ndb, txn: self.txn, val: unwrappedVal, generation: generation, inherited: inherited, name: name) + return NdbTxn(ndb: self.ndb, txn: self.txn, val: unwrappedVal, generation: generation, name: name) } } diff --git a/nostrdb/Test/NdbTests.swift b/nostrdb/Test/NdbTests.swift index 47426b00..43d71441 100644 --- a/nostrdb/Test/NdbTests.swift +++ b/nostrdb/Test/NdbTests.swift @@ -164,25 +164,6 @@ final class NdbTests: XCTestCase { let testNote = NdbNote.owned_from_json(json: testJSONWithEscapedSlashes)! XCTAssertEqual(testNote.content, "https://cdn.nostr.build/i/5c1d3296f66c2630131bf123106486aeaf051ed8466031c0e0532d70b33cddb2.jpg") } - - func test_inherited_transactions() throws { - let ndb = Ndb(path: db_dir)! - do { - guard let txn1 = NdbTxn(ndb: ndb) else { return XCTAssert(false) } - - let ntxn = (Thread.current.threadDictionary.value(forKey: "ndb_txn") as? ndb_txn)! - XCTAssertEqual(txn1.txn.lmdb, ntxn.lmdb) - XCTAssertEqual(txn1.txn.mdb_txn, ntxn.mdb_txn) - - guard let txn2 = NdbTxn(ndb: ndb) else { return XCTAssert(false) } - - XCTAssertEqual(txn1.inherited, false) - XCTAssertEqual(txn2.inherited, true) - } - - let ndb_txn = Thread.current.threadDictionary.value(forKey: "ndb_txn") - XCTAssertNil(ndb_txn) - } func test_decode_perf() throws { // This is an example of a performance test case.