Address review comments

This commit is contained in:
Duncan Ogilvie
2026-05-13 20:28:28 +02:00
parent 44840dc279
commit 5e5a6f4124
3 changed files with 181 additions and 13 deletions
+96 -13
View File
@@ -370,6 +370,41 @@ struct ValidityToken {
bool is_valid() const { return valid.load(); }
};
struct ModuleTokenRegistry {
std::mutex mutex;
std::unordered_map<LLVMModuleRef, std::weak_ptr<ValidityToken>> tokens;
static ModuleTokenRegistry &instance() {
static ModuleTokenRegistry registry;
return registry;
}
void register_module(LLVMModuleRef mod,
const std::shared_ptr<ValidityToken> &token) {
if (!mod || !token)
return;
std::lock_guard<std::mutex> lock(mutex);
tokens[mod] = token;
}
void unregister_module(LLVMModuleRef mod) {
if (!mod)
return;
std::lock_guard<std::mutex> lock(mutex);
tokens.erase(mod);
}
std::shared_ptr<ValidityToken> lookup(LLVMModuleRef mod) {
if (!mod)
return nullptr;
std::lock_guard<std::mutex> lock(mutex);
auto it = tokens.find(mod);
if (it == tokens.end())
return nullptr;
return it->second.lock();
}
};
// =============================================================================
// Base class to prevent copy
// =============================================================================
@@ -6297,13 +6332,16 @@ struct LLVMModuleWrapper : NoMoveCopy {
: m_context_token(std::move(context_token)),
m_token(std::make_shared<ValidityToken>()), m_ctx_ref(ctx) {
m_ref = LLVMModuleCreateWithNameInContext(name.c_str(), ctx);
ModuleTokenRegistry::instance().register_module(m_ref, m_token);
}
// Constructor for wrapping an existing module (from bitcode parsing)
LLVMModuleWrapper(LLVMModuleRef mod, LLVMContextRef ctx,
std::shared_ptr<ValidityToken> context_token)
: m_ref(mod), m_context_token(std::move(context_token)),
m_token(std::make_shared<ValidityToken>()), m_ctx_ref(ctx) {}
m_token(std::make_shared<ValidityToken>()), m_ctx_ref(ctx) {
ModuleTokenRegistry::instance().register_module(m_ref, m_token);
}
// Constructor for non-owning (borrowed) reference to an existing module.
// Used by function.module to return the function's parent module.
@@ -6324,6 +6362,7 @@ struct LLVMModuleWrapper : NoMoveCopy {
~LLVMModuleWrapper() {
if (m_ref && !m_borrowed) {
ModuleTokenRegistry::instance().unregister_module(m_ref);
// Only dispose the module if its context is still alive.
// If the context was already destroyed, the module memory is already
// freed and calling LLVMDisposeModule would cause a use-after-free crash.
@@ -6358,6 +6397,8 @@ struct LLVMModuleWrapper : NoMoveCopy {
void dispose() {
if (m_ref) {
if (!m_borrowed)
ModuleTokenRegistry::instance().unregister_module(m_ref);
LLVMDisposeModule(m_ref);
m_ref = nullptr;
}
@@ -9057,6 +9098,9 @@ struct LLVMJITWrapper : NoMoveCopy {
LLVMErrorRef err = LLVMOrcLLJITAddLLVMIRModule(
m_ref, LLVMOrcLLJITGetMainJITDylib(m_ref), tsm);
// The ORC C API documents ThreadSafeContext ownership as shared. The
// ThreadSafeModule/JIT retains the underlying context data; disposing this
// local reference after creating the ThreadSafeModule is required.
LLVMOrcDisposeThreadSafeContext(tsc);
if (err) {
throw LLVMError("Failed to add module to JIT: " +
@@ -10436,6 +10480,21 @@ struct LLVMMetadataWrapper {
}
};
static LLVMModuleRef module_for_metadata_value(LLVMValueRef value) {
if (LLVMIsAInstruction(value)) {
LLVMBasicBlockRef bb = LLVMGetInstructionParent(value);
if (!bb)
return nullptr;
LLVMValueRef fn = LLVMGetBasicBlockParent(bb);
if (!fn)
return nullptr;
return LLVMGetGlobalParent(fn);
}
if (LLVMIsAGlobalValue(value))
return LLVMGetGlobalParent(value);
return nullptr;
}
static LLVMContextRef context_for_metadata_value(LLVMValueRef value) {
if (!LLVMIsAInstruction(value) && !LLVMIsAGlobalValue(value)) {
throw LLVMAssertionError(
@@ -10468,6 +10527,10 @@ static LLVMContextRef context_for_metadata_value(LLVMValueRef value) {
static LLVMMetadataRef normalize_metadata_for_attachment(
LLVMContextRef ctx, const LLVMMetadataWrapper &md) {
md.check_valid();
if (md.m_ctx_ref && md.m_ctx_ref != ctx) {
throw LLVMAssertionError(
"metadata attachment requires metadata from the same context");
}
LLVMMetadataRef metadata_ref = md.m_ref;
LLVMValueRef md_value = LLVMMetadataAsValue(ctx, metadata_ref);
bool is_md_node = LLVMIsAMDNode(md_value) && !LLVMIsAValueAsMetadata(md_value);
@@ -10482,11 +10545,15 @@ struct LLVMMetadataMapWrapper {
LLVMValueRef m_value_ref = nullptr;
LLVMContextRef m_ctx_ref = nullptr;
std::shared_ptr<ValidityToken> m_context_token;
std::shared_ptr<ValidityToken> m_module_token;
LLVMMetadataMapWrapper() = default;
LLVMMetadataMapWrapper(LLVMValueRef value, LLVMContextRef ctx,
std::shared_ptr<ValidityToken> token)
: m_value_ref(value), m_ctx_ref(ctx), m_context_token(std::move(token)) {}
std::shared_ptr<ValidityToken> context_token,
std::shared_ptr<ValidityToken> module_token)
: m_value_ref(value), m_ctx_ref(ctx),
m_context_token(std::move(context_token)),
m_module_token(std::move(module_token)) {}
void check_valid() const {
if (!m_value_ref)
@@ -10495,6 +10562,8 @@ struct LLVMMetadataMapWrapper {
throw LLVMMemoryError("Metadata view has no context");
if (!m_context_token || !m_context_token->is_valid())
throw LLVMMemoryError("Metadata view used after context was destroyed");
if (m_module_token && !m_module_token->is_valid())
throw LLVMMemoryError("Metadata view used after module was disposed");
if (!LLVMIsAInstruction(m_value_ref) && !LLVMIsAGlobalValue(m_value_ref)) {
throw LLVMAssertionError(
"metadata view requires an instruction or global value");
@@ -10656,6 +10725,10 @@ struct LLVMNamedMetadataListWrapper {
void append(const LLVMMetadataWrapper &md) {
check_valid();
md.check_valid();
if (md.m_ctx_ref && md.m_ctx_ref != m_ctx_ref) {
throw LLVMAssertionError(
"named metadata requires metadata from the same context");
}
LLVMValueRef val = LLVMMetadataAsValue(m_ctx_ref, md.m_ref);
LLVMAddNamedMetadataOperand(m_module_ref, m_name.c_str(), val);
}
@@ -10768,6 +10841,10 @@ struct LLVMModuleFlagsWrapper {
const LLVMMetadataWrapper &value) {
check_valid();
value.check_valid();
if (value.m_ctx_ref && value.m_ctx_ref != m_ctx_ref) {
throw LLVMAssertionError(
"module flag requires metadata from the same context");
}
if (key.empty())
throw LLVMAssertionError("module flag key cannot be empty");
LLVMAddModuleFlag(m_module_ref, behavior, key.c_str(), key.size(), value.m_ref);
@@ -10885,7 +10962,10 @@ inline LLVMMetadataWrapper LLVMValueWrapper::as_metadata() const {
inline LLVMMetadataMapWrapper LLVMValueWrapper::metadata() const {
check_valid();
LLVMContextRef ctx_ref = context_for_metadata_value(m_ref);
return LLVMMetadataMapWrapper(m_ref, ctx_ref, m_context_token);
LLVMModuleRef mod_ref = module_for_metadata_value(m_ref);
std::shared_ptr<ValidityToken> module_token =
ModuleTokenRegistry::instance().lookup(mod_ref);
return LLVMMetadataMapWrapper(m_ref, ctx_ref, m_context_token, module_token);
}
// Implementation of replace_md_node_operand_with - needs LLVMMetadataWrapper
@@ -17043,14 +17123,15 @@ Valid when:
"__next__",
[](LLVMSectionIteratorWrapper &self) -> LLVMSectionIteratorWrapper * {
self.check_valid();
if (self.is_at_end())
throw nb::stop_iteration();
if (self.m_python_iter_started) {
self.move_next();
if (self.is_at_end())
throw nb::stop_iteration();
} else {
self.m_python_iter_started = true;
}
if (self.is_at_end()) {
throw nb::stop_iteration();
}
return &self;
},
nb::rv_policy::reference_internal);
@@ -17108,14 +17189,15 @@ TODO: needs safe LLVM-C API equivalent.
"__next__",
[](LLVMSymbolIteratorWrapper &self) -> LLVMSymbolIteratorWrapper * {
self.check_valid();
if (self.is_at_end())
throw nb::stop_iteration();
if (self.m_python_iter_started) {
self.move_next();
if (self.is_at_end())
throw nb::stop_iteration();
} else {
self.m_python_iter_started = true;
}
if (self.is_at_end()) {
throw nb::stop_iteration();
}
return &self;
},
nb::rv_policy::reference_internal);
@@ -17178,14 +17260,15 @@ Valid when:
[](LLVMRelocationIteratorWrapper &self)
-> LLVMRelocationIteratorWrapper * {
self.check_valid();
if (self.is_at_end())
throw nb::stop_iteration();
if (self.m_python_iter_started) {
self.move_next();
if (self.is_at_end())
throw nb::stop_iteration();
} else {
self.m_python_iter_started = true;
}
if (self.is_at_end()) {
throw nb::stop_iteration();
}
return &self;
},
nb::rv_policy::reference_internal);
@@ -29,6 +29,22 @@ def _manual_symbol_names(binary: llvm.Binary) -> list[str]:
return names
def _assert_exhausted_next_is_stable(iterator) -> None:
while True:
try:
next(iterator)
except StopIteration:
break
for _ in range(2):
try:
next(iterator)
except StopIteration:
pass
else:
raise AssertionError("expected exhausted iterator to stay exhausted")
def test_binary_iterators_match_manual_iteration():
obj = Path("llvm-c/llvm-c-test/inputs/simple.o")
if not obj.exists():
@@ -52,6 +68,23 @@ def test_binary_iterators_match_manual_iteration():
)
def test_binary_iterators_stay_stopped_after_exhaustion():
obj = Path("llvm-c/llvm-c-test/inputs/simple.o")
if not obj.exists():
print(f"SKIP: object file not found: {obj.resolve()}")
return
with llvm.BinaryManager.from_file(obj) as binary:
_assert_exhausted_next_is_stable(iter(binary.sections))
_assert_exhausted_next_is_stable(iter(binary.symbols))
sections = binary.sections
if not sections.is_at_end():
_assert_exhausted_next_is_stable(iter(sections.relocations))
if __name__ == "__main__":
test_binary_iterators_match_manual_iteration()
print("test_binary_iterators_match_manual_iteration: PASSED")
test_binary_iterators_stay_stopped_after_exhaustion()
print("test_binary_iterators_stay_stopped_after_exhaustion: PASSED")
@@ -108,6 +108,52 @@ def test_md_string_value_round_trips_to_python_string():
raise AssertionError("expected Metadata.value to be unimplemented")
def test_metadata_attachment_rejects_cross_context_metadata():
with llvm.create_context() as ctx1:
with llvm.create_context() as ctx2:
with ctx1.create_module("metadata_context_mismatch") as mod:
i32 = ctx1.types.i32
fn = mod.add_function("f", ctx1.types.function(i32, [i32]))
entry = fn.append_basic_block("entry")
with entry.create_builder() as builder:
inst = builder.add(fn.get_param(0), i32.constant(1), "sum")
foreign = ctx2.md_node([ctx2.md_string("foreign")])
try:
inst.metadata["example.foreign"] = foreign
except llvm.LLVMAssertionError as exc:
assert "same context" in str(exc)
else:
raise AssertionError("expected cross-context metadata error")
builder.ret(inst)
assert "example.foreign" not in str(mod)
assert mod.verify(), mod.verification_error
def test_metadata_map_rejects_use_after_module_disposed():
ctx_mgr = llvm.create_context()
ctx = ctx_mgr.__enter__()
try:
mod_mgr = ctx.create_module("metadata_map_lifetime")
mod = mod_mgr.__enter__()
fn = mod.add_function("f", ctx.types.function(ctx.types.void, []))
metadata_view = fn.metadata
mod_mgr.__exit__(None, None, None)
try:
metadata_view.get("example.kind")
except llvm.LLVMMemoryError as exc:
assert "module was disposed" in str(exc)
else:
raise AssertionError("expected disposed-module metadata view error")
finally:
ctx_mgr.__exit__(None, None, None)
def test_metadata_copy_to_replaces_raw_kind_id_copying():
with llvm.create_context() as ctx:
with ctx.create_module("metadata_copy") as mod:
@@ -280,6 +326,12 @@ if __name__ == "__main__":
test_md_string_value_round_trips_to_python_string()
print("test_md_string_value_round_trips_to_python_string: PASSED")
test_metadata_attachment_rejects_cross_context_metadata()
print("test_metadata_attachment_rejects_cross_context_metadata: PASSED")
test_metadata_map_rejects_use_after_module_disposed()
print("test_metadata_map_rejects_use_after_module_disposed: PASSED")
test_metadata_copy_to_replaces_raw_kind_id_copying()
print("test_metadata_copy_to_replaces_raw_kind_id_copying: PASSED")