use crate::builtins::{PyCode, PyStrInterned};
use crate::frozen::FrozenModule;
use crate::{Py, VirtualMachine, builtins::PyBaseExceptionRef};
use core::borrow::Borrow;
pub(crate) use _imp::module_def;
pub(super) use crate::vm::resolve_frozen_alias;
#[cfg(feature = "threading")]
#[pymodule(sub, name = "_imp")]
mod lock {
use crate::{PyResult, VirtualMachine, stdlib::_thread::RawRMutex};
use core::cell::Cell;
static IMP_LOCK: RawRMutex = RawRMutex::INIT;
thread_local! {
static IMP_LOCK_DEPTH: Cell<usize> = const { Cell::new(0) };
}
fn bump_depth() {
IMP_LOCK_DEPTH.with(|c| c.set(c.get() + 1));
}
fn drop_depth() {
IMP_LOCK_DEPTH.with(|c| c.set(c.get().saturating_sub(1)));
}
#[pyfunction]
fn acquire_lock(vm: &VirtualMachine) {
vm.allow_threads(acquire_lock_for_fork);
}
#[pyfunction]
fn release_lock(vm: &VirtualMachine) -> PyResult<()> {
if !IMP_LOCK.is_locked() || !IMP_LOCK.is_owned_by_current_thread() {
Err(vm.new_runtime_error("Global import lock not held"))
} else {
unsafe { IMP_LOCK.unlock() };
drop_depth();
Ok(())
}
}
#[pyfunction]
fn lock_held(_vm: &VirtualMachine) -> bool {
IMP_LOCK.is_locked()
}
pub(super) fn acquire_lock_for_fork() {
IMP_LOCK.lock();
bump_depth();
}
#[cfg(all(unix, feature = "host_env"))]
pub(super) fn release_lock_after_fork_parent() {
if IMP_LOCK.is_locked() && IMP_LOCK.is_owned_by_current_thread() {
unsafe { IMP_LOCK.unlock() };
drop_depth();
}
}
#[cfg(all(unix, feature = "host_env"))]
pub(crate) unsafe fn reinit_after_fork() {
let depth = IMP_LOCK_DEPTH.with(Cell::get);
unsafe { rustpython_common::lock::zero_reinit_after_fork(&IMP_LOCK) };
for _ in 0..depth {
IMP_LOCK.lock();
}
}
#[cfg(all(unix, feature = "host_env"))]
pub(super) unsafe fn after_fork_child_reinit_and_release() {
unsafe { reinit_after_fork() };
if IMP_LOCK.is_locked() && IMP_LOCK.is_owned_by_current_thread() {
unsafe { IMP_LOCK.unlock() };
drop_depth();
}
}
}
#[cfg(all(unix, feature = "threading", feature = "host_env"))]
pub(crate) fn acquire_imp_lock_for_fork(vm: &VirtualMachine) {
vm.allow_threads(lock::acquire_lock_for_fork);
}
#[cfg(all(unix, feature = "threading", feature = "host_env"))]
pub(crate) fn release_imp_lock_after_fork_parent() {
lock::release_lock_after_fork_parent();
}
#[cfg(all(unix, feature = "threading", feature = "host_env"))]
pub(crate) unsafe fn reinit_imp_lock_after_fork() {
unsafe { lock::reinit_after_fork() }
}
#[cfg(all(unix, feature = "threading", feature = "host_env"))]
pub(crate) unsafe fn after_fork_child_imp_lock_release() {
unsafe { lock::after_fork_child_reinit_and_release() }
}
#[cfg(not(feature = "threading"))]
#[pymodule(sub, name = "_imp")]
mod lock {
use crate::vm::VirtualMachine;
#[pyfunction]
pub(super) const fn acquire_lock(_vm: &VirtualMachine) {}
#[pyfunction]
pub(super) const fn release_lock(_vm: &VirtualMachine) {}
#[pyfunction]
pub(super) const fn lock_held(_vm: &VirtualMachine) -> bool {
false
}
}
#[allow(dead_code)]
enum FrozenError {
BadName, NotFound, Disabled, Excluded, Invalid, }
impl FrozenError {
fn to_pyexception(&self, mod_name: &str, vm: &VirtualMachine) -> PyBaseExceptionRef {
let msg = match self {
Self::BadName | Self::NotFound => format!("No such frozen object named {mod_name}"),
Self::Disabled => format!(
"Frozen modules are disabled and the frozen object named {mod_name} is not essential"
),
Self::Excluded => format!("Excluded frozen object named {mod_name}"),
Self::Invalid => format!("Frozen object named {mod_name} is invalid"),
};
vm.new_import_error(msg, vm.ctx.new_utf8_str(mod_name))
}
}
fn find_frozen(name: &str, vm: &VirtualMachine) -> Result<FrozenModule, FrozenError> {
let frozen = vm
.state
.frozen
.get(name)
.copied()
.ok_or(FrozenError::NotFound)?;
if matches!(
name,
"_frozen_importlib" | "_frozen_importlib_external" | "zipimport"
) {
return Ok(frozen);
}
let override_val = vm.state.override_frozen_modules.load();
if override_val < 0 {
return Err(FrozenError::NotFound);
}
Ok(frozen)
}
#[pymodule(with(lock))]
mod _imp {
use crate::{
PyObjectRef, PyPayload, PyRef, PyResult, VirtualMachine,
builtins::{PyBytesRef, PyCode, PyMemoryView, PyModule, PyStrRef, PyUtf8StrRef},
import, version,
};
use super::FrozenError;
#[pyattr]
fn check_hash_based_pycs(vm: &VirtualMachine) -> PyStrRef {
vm.ctx
.new_str(vm.state.config.settings.check_hash_pycs_mode.to_string())
}
#[pyattr(name = "pyc_magic_number_token")]
use version::PYC_MAGIC_NUMBER_TOKEN;
#[pyfunction]
const fn extension_suffixes() -> Vec<PyObjectRef> {
Vec::new()
}
#[pyfunction]
fn is_builtin(name: PyUtf8StrRef, vm: &VirtualMachine) -> bool {
vm.state.module_defs.contains_key(name.as_str())
}
#[pyfunction]
fn is_frozen(name: PyUtf8StrRef, vm: &VirtualMachine) -> bool {
super::find_frozen(name.as_str(), vm).is_ok()
}
#[pyfunction]
fn create_builtin(spec: PyObjectRef, vm: &VirtualMachine) -> PyResult {
let sys_modules = vm.sys_module.get_attr("modules", vm).unwrap();
let name: PyUtf8StrRef = spec.get_attr("name", vm)?.try_into_value(vm)?;
if let Ok(module) = sys_modules.get_item(&*name, vm) {
return Ok(module);
}
let name_str = name.as_str();
if let Some(&def) = vm.state.module_defs.get(name_str) {
let module = if let Some(create) = def.slots.create {
create(vm, &spec, def)?
} else {
PyModule::from_def(def).into_ref(&vm.ctx)
};
PyModule::__init_dict_from_def(vm, &module);
module.__init_methods(vm)?;
sys_modules.set_item(name.as_pystr(), module.clone().into(), vm)?;
if let Some(exec) = def.slots.exec {
exec(vm, &module)?;
}
return Ok(module.into());
}
Ok(vm.ctx.none())
}
#[derive(FromArgs)]
struct CreateDynamicArgs {
#[pyarg(positional)]
spec: PyObjectRef,
#[pyarg(positional, optional)]
_file: crate::function::OptionalArg<PyObjectRef>,
}
#[pyfunction]
fn create_dynamic(args: CreateDynamicArgs, vm: &VirtualMachine) -> PyResult {
let name_obj = args.spec.get_attr("name", vm)?;
let name: PyUtf8StrRef = name_obj.try_into_value(vm)?;
if name.as_str().contains('\0') {
return Err(vm.new_value_error("embedded null character".to_owned()));
}
let origin_obj = args.spec.get_attr("origin", vm)?;
if vm.is_none(&origin_obj) {
return Err(vm.new_value_error("origin must be set".to_owned()));
}
let origin: PyUtf8StrRef = origin_obj.try_into_value(vm)?;
if origin.as_str().contains('\0') {
return Err(vm.new_value_error("embedded null character".to_owned()));
}
let sys_modules = vm.sys_module.get_attr("modules", vm)?;
if let Ok(module) = sys_modules.get_item(&*name, vm) {
return Ok(module);
}
#[cfg(all(feature = "host_env", any(unix, windows)))]
{
let origin_str = origin.as_str();
let short_name = name.as_str().rsplit('.').next().unwrap_or(name.as_str());
let export_func_name = format!("PyModExport_{short_name}");
#[cfg(unix)]
let handle_res = {
let mode = rustpython_host_env::ctypes::dlopen_mode(None);
rustpython_host_env::ctypes::open_library_with_mode(origin_str, mode)
};
#[cfg(windows)]
let handle_res = rustpython_host_env::ctypes::open_library(origin_str);
if let Ok(handle) = handle_res
&& let Ok(export_fn_addr) = rustpython_host_env::ctypes::lookup_function_symbol_addr(
handle,
export_func_name.as_bytes(),
)
&& export_fn_addr != 0
{
type ModExportFn = unsafe extern "C" fn() -> *mut crate::PyObject;
let export_fn: ModExportFn =
unsafe { core::mem::transmute(export_fn_addr as *const ()) };
let mod_ptr = unsafe { export_fn() };
if let Some(mod_nonnull) = core::ptr::NonNull::new(mod_ptr) {
let py_obj = unsafe { crate::PyObjectRef::from_raw(mod_nonnull) };
return Ok(py_obj);
}
}
}
Err(vm.new_import_error(
format!(
"dynamic module does not define module export function (PyModExport_{})",
name.as_str()
),
name.into_wtf8(),
))
}
#[pyfunction]
fn exec_dynamic(_module: PyRef<PyModule>) -> i32 {
0
}
#[pyfunction]
fn exec_builtin(_mod: PyRef<PyModule>) -> i32 {
0
}
#[derive(FromArgs)]
struct FrozenObjectArgs {
#[pyarg(positional)]
name: PyUtf8StrRef,
#[pyarg(positional, optional)]
data: Option<PyObjectRef>,
}
#[pyfunction]
fn get_frozen_object(args: FrozenObjectArgs, vm: &VirtualMachine) -> PyResult<PyRef<PyCode>> {
let FrozenObjectArgs { name, data } = args;
if let Some(data) = data
&& !vm.is_none(&data)
{
let invalid_err = || {
vm.new_import_error(
format!("Frozen object named '{}' is invalid", name.as_str()),
name.clone().into_wtf8(),
)
};
crate::protocol::PyBuffer::from_object(
vm,
&data,
crate::protocol::BufferFlags::SIMPLE,
)?;
let loads = vm.import("marshal", 0)?.get_attr("loads", vm)?;
let code = loads.call((data,), vm).map_err(|_| invalid_err())?;
return code.downcast::<PyCode>().map_err(|_| invalid_err());
}
import::make_frozen(vm, name.as_str())
}
#[pyfunction]
fn init_frozen(name: PyUtf8StrRef, vm: &VirtualMachine) -> PyResult {
import::import_frozen(vm, name.as_str())
}
#[pyfunction]
fn is_frozen_package(name: PyUtf8StrRef, vm: &VirtualMachine) -> PyResult<bool> {
let name_str = name.as_str();
super::find_frozen(name_str, vm)
.map(|frozen| frozen.package)
.map_err(|e| e.to_pyexception(name_str, vm))
}
#[pyfunction]
fn _override_frozen_modules_for_tests(r#override: isize, vm: &VirtualMachine) {
vm.state.override_frozen_modules.store(r#override);
}
#[pyfunction]
fn _fix_co_filename(code: PyRef<PyCode>, path: PyStrRef, vm: &VirtualMachine) {
let old_name = code.source_path();
let new_name = vm.ctx.intern_str(path.as_wtf8());
super::update_code_filenames(&code, old_name, new_name);
}
#[pyfunction]
fn _frozen_module_names(vm: &VirtualMachine) -> Vec<PyObjectRef> {
vm.state
.frozen
.keys()
.map(|&name| vm.ctx.new_utf8_str(name).into())
.collect()
}
#[derive(FromArgs)]
struct FindFrozenArgs {
#[pyarg(positional)]
name: PyUtf8StrRef,
#[pyarg(named, default)]
withdata: bool,
}
#[allow(clippy::type_complexity)]
#[pyfunction]
fn find_frozen(
args: FindFrozenArgs,
vm: &VirtualMachine,
) -> PyResult<Option<(Option<PyRef<PyMemoryView>>, bool, Option<PyStrRef>)>> {
let FindFrozenArgs { name, withdata } = args;
let name_str = name.as_str();
let info = match super::find_frozen(name_str, vm) {
Ok(info) => info,
Err(FrozenError::NotFound | FrozenError::Disabled | FrozenError::BadName) => {
return Ok(None);
}
Err(e) => return Err(e.to_pyexception(name_str, vm)),
};
let data = if withdata {
let code = PyCode::new_ref_from_frozen(vm, info.code);
let dumps = vm.import("marshal", 0)?.get_attr("dumps", vm)?;
let bytes = dumps.call((code,), vm)?;
Some(PyMemoryView::from_object(&bytes, vm)?.into_ref(&vm.ctx))
} else {
None
};
let origname_str = super::resolve_frozen_alias(name_str);
let origname = if origname_str.is_empty() {
None
} else {
Some(vm.ctx.new_utf8_str(origname_str).into())
};
Ok(Some((data, info.package, origname)))
}
#[derive(FromArgs)]
struct SourceHashArgs {
#[pyarg(any)]
key: u64,
#[pyarg(any)]
source: PyBytesRef,
}
#[pyfunction]
fn source_hash(SourceHashArgs { key, source }: SourceHashArgs) -> Vec<u8> {
let hash: u64 = crate::common::hash::keyed_hash(key, source.as_bytes());
hash.to_le_bytes().to_vec()
}
}
fn update_code_filenames(
code: &Py<PyCode>,
old_name: &'static PyStrInterned,
new_name: &'static PyStrInterned,
) {
let current = code.source_path();
if !core::ptr::eq(current, old_name) && current.as_str() != old_name.as_str() {
return;
}
code.set_source_path(new_name);
for constant in code.code.constants.iter() {
let obj: &crate::PyObject = constant.borrow();
if let Some(inner_code) = obj.downcast_ref::<PyCode>() {
update_code_filenames(inner_code, old_name, new_name);
}
}
}