Revert "Remove transaction inheritance"

This reverts commit c1fb8b592f.

It turns out we still need transaction inheritance.

Even though the current Ndb interface hides transaction objects from callers and prevents those transactions from being held open beyond the small sync window where the underlying query result is borrowed, there is nothing preventing that caller from making another nested ndb query call from within that closure.

Closes: https://github.com/damus-io/damus/issues/3688
Changelog-Fixed: Fixed issue where mentioned profile names would not render properly
Signed-off-by: Daniel D’Aquino <daniel@daquino.me>
This commit is contained in:
Daniel D’Aquino
2026-03-18 18:55:30 -07:00
parent 92cca89c06
commit bfad604e5b
2 changed files with 132 additions and 80 deletions
+113 -80
View File
@@ -11,26 +11,18 @@ import Foundation
fileprivate var txn_count: Int = 0 fileprivate var txn_count: Int = 0
#endif #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 // Would use struct and ~Copyable but generics aren't supported well
class NdbTxn<T>: RawNdbTxnAccessible { class NdbTxn<T>: RawNdbTxnAccessible {
var txn: ndb_txn var txn: ndb_txn
private var val: T! private var val: T!
var moved: Bool var moved: Bool
var inherited: Bool
var ndb: Ndb var ndb: Ndb
var generation: Int var generation: Int
var name: String var name: String
static func pure(ndb: Ndb, val: T) -> NdbTxn<T> { static func pure(ndb: Ndb, val: T) -> NdbTxn<T> {
.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 /// Simple helper struct for the init function to avoid compiler errors encountered by using other techniques
@@ -44,21 +36,36 @@ class NdbTxn<T>: RawNdbTxnAccessible {
self.name = name ?? "txn" self.name = name ?? "txn"
self.ndb = ndb self.ndb = ndb
self.generation = ndb.generation self.generation = ndb.generation
if let active_txn = Thread.current.threadDictionary["ndb_txn"] as? ndb_txn,
// Always create fresh transaction let txn_generation = Thread.current.threadDictionary["txn_generation"] as? Int,
let result: R? = try? ndb.withNdb({ txn_generation == ndb.generation
var txn = ndb_txn() {
let ok = ndb_begin_query(ndb.ndb.ndb, &txn) != 0 // some parent thread is active, use that instead
guard ok else { return .none } print("txn: inherited txn")
#if TXNDEBUG self.txn = active_txn
txn_count += 1 self.inherited = true
#endif self.generation = Thread.current.threadDictionary["txn_generation"] as! Int
return .some(R(txn: txn, generation: ndb.generation)) let ref_count = Thread.current.threadDictionary["ndb_txn_ref_count"] as! Int
}, maxWaitTimeout: .ndbTransactionTimeout) let new_ref_count = ref_count + 1
guard let result else { return nil } Thread.current.threadDictionary["ndb_txn_ref_count"] = new_ref_count
self.txn = result.txn } else {
self.generation = result.generation 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 #if TXNDEBUG
print("txn: open gen\(self.generation) '\(self.name)' \(txn_count)") print("txn: open gen\(self.generation) '\(self.name)' \(txn_count)")
#endif #endif
@@ -66,10 +73,11 @@ class NdbTxn<T>: RawNdbTxnAccessible {
self.val = with(self) 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.txn = txn
self.val = val self.val = val
self.moved = false self.moved = false
self.inherited = inherited
self.ndb = ndb self.ndb = ndb
self.generation = generation self.generation = generation
self.name = name self.name = name
@@ -91,15 +99,27 @@ class NdbTxn<T>: RawNdbTxnAccessible {
print("txn: not closing. db closed") print("txn: not closing. db closed")
return 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 { if moved {
//print("txn: not closing. moved") //print("txn: not closing. moved")
return return
} }
_ = try? ndb.withNdb({
ndb_end_query(&self.txn)
}, maxWaitTimeout: .ndbTransactionTimeout)
#if TXNDEBUG #if TXNDEBUG
txn_count -= 1; txn_count -= 1;
print("txn: close gen\(generation) '\(name)' \(txn_count)") print("txn: close gen\(generation) '\(name)' \(txn_count)")
@@ -109,14 +129,14 @@ class NdbTxn<T>: RawNdbTxnAccessible {
// functor // functor
func map<Y>(_ transform: (T) -> Y) -> NdbTxn<Y> { func map<Y>(_ transform: (T) -> Y) -> NdbTxn<Y> {
self.moved = true 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!? // comonad!?
// useful for moving ownership of a transaction to another value // useful for moving ownership of a transaction to another value
func extend<Y>(_ with: (NdbTxn<T>) -> Y) -> NdbTxn<Y> { func extend<Y>(_ with: (NdbTxn<T>) -> Y) -> NdbTxn<Y> {
self.moved = true 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<T: ~Copyable> {
var txn: ndb_txn var txn: ndb_txn
var val: T! var val: T!
var moved: Bool var moved: Bool
var inherited: Bool
var ndb: Ndb var ndb: Ndb
var generation: Int var generation: Int
var name: String var name: String
static func pure(ndb: Ndb, val: consuming T) -> SafeNdbTxn<T> { static func pure(ndb: Ndb, val: consuming T) -> SafeNdbTxn<T> {
.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 /// Simple helper struct for the init function to avoid compiler errors encountered by using other techniques
@@ -152,42 +173,52 @@ class SafeNdbTxn<T: ~Copyable> {
static func new(on ndb: Ndb, with valueGetter: (PlaceholderNdbTxn) -> T? = { _ in () }, name: String = "txn") -> SafeNdbTxn<T>? { static func new(on ndb: Ndb, with valueGetter: (PlaceholderNdbTxn) -> T? = { _ in () }, name: String = "txn") -> SafeNdbTxn<T>? {
guard !ndb.is_closed else { return nil } guard !ndb.is_closed else { return nil }
let generation: Int
// Always create fresh transaction let txn: ndb_txn
let result: R? = try? ndb.withNdb({ let inherited: Bool
var txn = ndb_txn() if let active_txn = Thread.current.threadDictionary["ndb_txn"] as? ndb_txn,
let ok = ndb_begin_query(ndb.ndb.ndb, &txn) != 0 let txn_generation = Thread.current.threadDictionary["txn_generation"] as? Int,
guard ok else { return .none } txn_generation == ndb.generation
#if TXNDEBUG {
txn_count += 1 // some parent thread is active, use that instead
#endif print("txn: inherited txn")
return .some(R(txn: txn, generation: ndb.generation)) txn = active_txn
}, maxWaitTimeout: .ndbTransactionTimeout) inherited = true
guard let result else { return nil } generation = Thread.current.threadDictionary["txn_generation"] as! Int
let txn = result.txn let ref_count = Thread.current.threadDictionary["ndb_txn_ref_count"] as! Int
let generation = result.generation 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 #if TXNDEBUG
print("txn: open gen\(generation) '\(name)' \(txn_count)") print("txn: open gen\(generation) '\(name)' \(txn_count)")
#endif #endif
let placeholderTxn = PlaceholderNdbTxn(txn: txn) let placeholderTxn = PlaceholderNdbTxn(txn: txn)
guard let val = valueGetter(placeholderTxn) else { guard let val = valueGetter(placeholderTxn) else { return nil }
// Fix leak: Close transaction before returning nil return SafeNdbTxn<T>(ndb: ndb, txn: txn, val: val, generation: generation, inherited: inherited, name: name)
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<T>(ndb: ndb, txn: txn, val: val, generation: generation, 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.txn = txn
self.val = consume val self.val = consume val
self.moved = false self.moved = false
self.inherited = inherited
self.ndb = ndb self.ndb = ndb
self.generation = generation self.generation = generation
self.name = name self.name = name
@@ -202,15 +233,27 @@ class SafeNdbTxn<T: ~Copyable> {
print("txn: not closing. db closed") print("txn: not closing. db closed")
return 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 { if moved {
//print("txn: not closing. moved") //print("txn: not closing. moved")
return return
} }
_ = try? ndb.withNdb({
ndb_end_query(&self.txn)
}, maxWaitTimeout: .ndbTransactionTimeout)
#if TXNDEBUG #if TXNDEBUG
txn_count -= 1; txn_count -= 1;
print("txn: close gen\(generation) '\(name)' \(txn_count)") print("txn: close gen\(generation) '\(name)' \(txn_count)")
@@ -220,14 +263,14 @@ class SafeNdbTxn<T: ~Copyable> {
// functor // functor
func map<Y>(_ transform: (borrowing T) -> Y) -> SafeNdbTxn<Y> { func map<Y>(_ transform: (borrowing T) -> Y) -> SafeNdbTxn<Y> {
self.moved = true 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!? // comonad!?
// useful for moving ownership of a transaction to another value // useful for moving ownership of a transaction to another value
func extend<Y>(_ with: (SafeNdbTxn<T>) -> Y) -> SafeNdbTxn<Y> { func extend<Y>(_ with: (SafeNdbTxn<T>) -> Y) -> SafeNdbTxn<Y> {
self.moved = true 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<Y>(_ with: (consuming SafeNdbTxn<T>) -> Y?) -> SafeNdbTxn<Y>? where Y: ~Copyable { consuming func maybeExtend<Y>(_ with: (consuming SafeNdbTxn<T>) -> Y?) -> SafeNdbTxn<Y>? where Y: ~Copyable {
@@ -235,20 +278,10 @@ class SafeNdbTxn<T: ~Copyable> {
let ndb = self.ndb let ndb = self.ndb
let txn = self.txn let txn = self.txn
let generation = self.generation let generation = self.generation
let inherited = self.inherited
let name = self.name let name = self.name
guard let newVal = with(consume self) else { return nil }
guard let newVal = with(consume self) else { return .init(ndb: ndb, txn: txn, val: newVal, generation: generation, inherited: inherited, name: name)
// 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)
} }
} }
@@ -271,7 +304,7 @@ extension NdbTxn where T: OptionalType {
return nil return nil
} }
self.moved = true self.moved = true
return NdbTxn<T.Wrapped>(ndb: self.ndb, txn: self.txn, val: unwrappedVal, generation: generation, name: name) return NdbTxn<T.Wrapped>(ndb: self.ndb, txn: self.txn, val: unwrappedVal, generation: generation, inherited: inherited, name: name)
} }
} }
+19
View File
@@ -164,6 +164,25 @@ final class NdbTests: XCTestCase {
let testNote = NdbNote.owned_from_json(json: testJSONWithEscapedSlashes)! let testNote = NdbNote.owned_from_json(json: testJSONWithEscapedSlashes)!
XCTAssertEqual(testNote.content, "https://cdn.nostr.build/i/5c1d3296f66c2630131bf123106486aeaf051ed8466031c0e0532d70b33cddb2.jpg") 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 { func test_decode_perf() throws {
// This is an example of a performance test case. // This is an example of a performance test case.