diff --git a/nostrdb/NdbTxn.swift b/nostrdb/NdbTxn.swift index fb93d4d4..6141f111 100644 --- a/nostrdb/NdbTxn.swift +++ b/nostrdb/NdbTxn.swift @@ -11,26 +11,18 @@ 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, name: "pure_txn") + .init(ndb: ndb, txn: ndb_txn(), val: val, generation: ndb.generation, inherited: true, name: "pure_txn") } /// Simple helper struct for the init function to avoid compiler errors encountered by using other techniques @@ -44,21 +36,36 @@ class NdbTxn: RawNdbTxnAccessible { self.name = name ?? "txn" self.ndb = ndb self.generation = ndb.generation - - // 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 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 + } #if TXNDEBUG print("txn: open gen\(self.generation) '\(self.name)' \(txn_count)") #endif @@ -66,10 +73,11 @@ class NdbTxn: RawNdbTxnAccessible { self.val = with(self) } - private init(ndb: Ndb, txn: ndb_txn, val: T, generation: Int, name: String) { + private init(ndb: Ndb, txn: ndb_txn, val: T, generation: Int, inherited: Bool, name: String) { self.txn = txn self.val = val self.moved = false + self.inherited = inherited self.ndb = ndb self.generation = generation self.name = name @@ -91,15 +99,27 @@ 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)") @@ -109,14 +129,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, name: self.name) + return .init(ndb: self.ndb, txn: self.txn, val: transform(val), generation: generation, inherited: inherited, 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, name: self.name) + return .init(ndb: self.ndb, txn: self.txn, val: with(self), generation: generation, inherited: inherited, name: self.name) } } @@ -136,12 +156,13 @@ 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, name: "pure_txn") + .init(ndb: ndb, txn: ndb_txn(), val: val, generation: ndb.generation, inherited: true, name: "pure_txn") } /// Simple helper struct for the init function to avoid compiler errors encountered by using other techniques @@ -152,42 +173,52 @@ class SafeNdbTxn { static func new(on ndb: Ndb, with valueGetter: (PlaceholderNdbTxn) -> T? = { _ in () }, name: String = "txn") -> SafeNdbTxn? { guard !ndb.is_closed else { return nil } - - // 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 - + 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 + } #if TXNDEBUG print("txn: open gen\(generation) '\(name)' \(txn_count)") #endif let placeholderTxn = PlaceholderNdbTxn(txn: txn) - 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) + guard let val = valueGetter(placeholderTxn) else { return nil } + return SafeNdbTxn(ndb: ndb, txn: txn, val: val, generation: generation, inherited: inherited, name: name) } - private init(ndb: Ndb, txn: ndb_txn, val: consuming T, generation: Int, name: String) { + private init(ndb: Ndb, txn: ndb_txn, val: consuming T, generation: Int, inherited: Bool, name: String) { self.txn = txn self.val = consume val self.moved = false + self.inherited = inherited self.ndb = ndb self.generation = generation self.name = name @@ -202,15 +233,27 @@ 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)") @@ -220,14 +263,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, name: self.name) + return .init(ndb: self.ndb, txn: self.txn, val: transform(val), generation: generation, inherited: inherited, 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, name: self.name) + return .init(ndb: self.ndb, txn: self.txn, val: with(self), generation: generation, inherited: inherited, name: self.name) } consuming func maybeExtend(_ with: (consuming SafeNdbTxn) -> Y?) -> SafeNdbTxn? where Y: ~Copyable { @@ -235,20 +278,10 @@ 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 { - // 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) + guard let newVal = with(consume self) else { return nil } + return .init(ndb: ndb, txn: txn, val: newVal, generation: generation, inherited: inherited, name: name) } } @@ -271,7 +304,7 @@ extension NdbTxn where T: OptionalType { return nil } self.moved = true - return NdbTxn(ndb: self.ndb, txn: self.txn, val: unwrappedVal, generation: generation, name: name) + return NdbTxn(ndb: self.ndb, txn: self.txn, val: unwrappedVal, generation: generation, inherited: inherited, name: name) } } diff --git a/nostrdb/Test/NdbTests.swift b/nostrdb/Test/NdbTests.swift index 43d71441..47426b00 100644 --- a/nostrdb/Test/NdbTests.swift +++ b/nostrdb/Test/NdbTests.swift @@ -164,6 +164,25 @@ 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.