Skip to main content

rustpython_vm/
import.rs

1//! Import mechanics
2
3use crate::{
4    AsObject, Py, PyObject, PyObjectRef, PyPayload, PyRef, PyResult,
5    builtins::{PyCode, PyStr, PyUtf8Str, PyUtf8StrRef, traceback::PyTraceback},
6    exceptions::types::PyBaseException,
7    scope::Scope,
8    vm::{VirtualMachine, resolve_frozen_alias, thread},
9};
10
11pub(crate) fn check_pyc_magic_number_bytes(buf: &[u8]) -> bool {
12    buf.starts_with(&crate::version::PYC_MAGIC_NUMBER_BYTES)
13}
14
15pub(crate) fn init_importlib_base(vm: &mut VirtualMachine) -> PyResult<PyObjectRef> {
16    flame_guard!("init importlib");
17
18    // importlib_bootstrap needs these and it inlines checks to sys.modules before calling into
19    // import machinery, so this should bring some speedup
20    #[cfg(all(feature = "threading", not(target_os = "wasi")))]
21    import_builtin(vm, "_thread")?;
22    import_builtin(vm, "_warnings")?;
23    import_builtin(vm, "_weakref")?;
24
25    let importlib = thread::enter_vm(vm, || {
26        let bootstrap = import_frozen(vm, "_frozen_importlib")?;
27        let install = bootstrap.get_attr("_install", vm)?;
28        let imp = import_builtin(vm, "_imp")?;
29        install.call((vm.sys_module.clone(), imp), vm)?;
30        Ok(bootstrap)
31    })?;
32    vm.import_func = importlib.get_attr(identifier!(vm, __import__), vm)?;
33    vm.importlib = importlib.clone();
34    Ok(importlib)
35}
36
37#[cfg(feature = "host_env")]
38pub(crate) fn init_importlib_package(vm: &VirtualMachine, importlib: &PyObject) -> PyResult<()> {
39    use crate::{TryFromObject, builtins::PyListRef};
40
41    thread::enter_vm(vm, || {
42        flame_guard!("install_external");
43
44        // same deal as imports above
45        import_builtin(vm, crate::stdlib::os::MODULE_NAME)?;
46        #[cfg(windows)]
47        import_builtin(vm, "winreg")?;
48        import_builtin(vm, "_io")?;
49        import_builtin(vm, "marshal")?;
50
51        let install_external = importlib.get_attr("_install_external_importers", vm)?;
52        install_external.call((), vm)?;
53        let zipimport_res = (|| -> PyResult<()> {
54            let zipimport = vm.import("zipimport", 0)?;
55            let zipimporter = zipimport.get_attr("zipimporter", vm)?;
56            let path_hooks = vm.sys_module.get_attr("path_hooks", vm)?;
57            let path_hooks = PyListRef::try_from_object(vm, path_hooks)?;
58            path_hooks.insert(0, zipimporter);
59            Ok(())
60        })();
61        if zipimport_res.is_err() {
62            warn!("couldn't init zipimport")
63        }
64        Ok(())
65    })
66}
67
68pub fn make_frozen(vm: &VirtualMachine, name: &str) -> PyResult<PyRef<PyCode>> {
69    let frozen = vm.state.frozen.get(name).ok_or_else(|| {
70        vm.new_import_error(
71            format!("No such frozen object named {name}"),
72            vm.ctx.new_utf8_str(name),
73        )
74    })?;
75    Ok(PyCode::new_ref_from_frozen(vm, frozen.code))
76}
77
78pub fn import_frozen(vm: &VirtualMachine, module_name: &str) -> PyResult {
79    let frozen = vm.state.frozen.get(module_name).ok_or_else(|| {
80        vm.new_import_error(
81            format!("No such frozen object named {module_name}"),
82            vm.ctx.new_utf8_str(module_name),
83        )
84    })?;
85    let module = import_code_obj(
86        vm,
87        module_name,
88        PyCode::new_ref_from_frozen(vm, frozen.code),
89        false,
90    )?;
91    debug_assert!(module.get_attr(identifier!(vm, __name__), vm).is_ok());
92    let origname = resolve_frozen_alias(module_name);
93    module.set_attr("__origname__", vm.ctx.new_utf8_str(origname), vm)?;
94    Ok(module)
95}
96
97pub fn import_builtin(vm: &VirtualMachine, module_name: &str) -> PyResult {
98    let sys_modules = vm.sys_module.get_attr("modules", vm)?;
99
100    // Check if already in sys.modules (handles recursive imports)
101    if let Ok(module) = sys_modules.get_item(module_name, vm) {
102        return Ok(module);
103    }
104
105    // Try multi-phase init first (preferred for modules that import other modules)
106    if let Some(&def) = vm.state.module_defs.get(module_name) {
107        // Phase 1: Create and initialize module
108        let module = def.create_module(vm)?;
109
110        // Add to sys.modules BEFORE exec (critical for circular import handling)
111        sys_modules.set_item(module_name, module.clone().into(), vm)?;
112
113        // Phase 2: Call exec slot (can safely import other modules now)
114        // If exec fails, remove the partially-initialized module from sys.modules
115        if let Err(e) = def.exec_module(vm, &module) {
116            let _ = sys_modules.del_item(module_name, vm);
117            return Err(e);
118        }
119
120        return Ok(module.into());
121    }
122
123    // Module not found in module_defs
124    Err(vm.new_import_error(
125        format!("Cannot import builtin module {module_name}"),
126        vm.ctx.new_utf8_str(module_name),
127    ))
128}
129
130#[cfg(feature = "rustpython-compiler")]
131pub fn import_file(
132    vm: &VirtualMachine,
133    module_name: &str,
134    file_path: &str,
135    content: &str,
136) -> PyResult {
137    let code = vm
138        .compile_with_opts(
139            content,
140            crate::compiler::Mode::Exec,
141            file_path,
142            vm.compile_opts(),
143        )
144        .map_err(|err| err.into_pyexception(vm, Some(content)))?;
145    import_code_obj(vm, module_name, code, true)
146}
147
148#[cfg(feature = "rustpython-compiler")]
149pub fn import_source(vm: &VirtualMachine, module_name: &str, content: &str) -> PyResult {
150    let code = vm
151        .compile_with_opts(
152            content,
153            crate::compiler::Mode::Exec,
154            "<source>",
155            vm.compile_opts(),
156        )
157        .map_err(|err| err.into_pyexception(vm, Some(content)))?;
158    import_code_obj(vm, module_name, code, false)
159}
160
161/// Check whether `module.__spec__._initializing` is true, i.e. the module
162/// is currently in the middle of being executed for the first time and is
163/// not yet safe to hand out as a finished result (used both by the slow
164/// import path below and by [`crate::VirtualMachine::import`]'s
165/// `sys.modules`-cache fast path).
166pub(crate) fn is_module_initializing(module: &PyObject, vm: &VirtualMachine) -> PyResult<bool> {
167    match vm.get_attribute_opt(module, vm.ctx.intern_str("__spec__"))? {
168        Some(spec) => match vm.get_attribute_opt(&spec, vm.ctx.intern_str("_initializing"))? {
169            Some(v) => v.try_to_bool(vm),
170            None => Ok(false),
171        },
172        None => Ok(false),
173    }
174}
175
176/// If `__spec__._initializing` is true, wait for the module to finish
177/// initializing by calling `_lock_unlock_module`.
178fn import_ensure_initialized(
179    module: &PyObject,
180    name: &Py<PyUtf8Str>,
181    vm: &VirtualMachine,
182) -> PyResult<()> {
183    if is_module_initializing(module, vm)? {
184        let lock_unlock = vm.importlib.get_attr("_lock_unlock_module", vm)?;
185        lock_unlock.call((name.to_owned(),), vm)?;
186    }
187    Ok(())
188}
189
190pub fn import_code_obj(
191    vm: &VirtualMachine,
192    module_name: &str,
193    code_obj: PyRef<PyCode>,
194    set_file_attr: bool,
195) -> PyResult {
196    let attrs = vm.ctx.new_dict();
197    attrs.set_item(
198        identifier!(vm, __name__),
199        vm.ctx.new_utf8_str(module_name).into(),
200        vm,
201    )?;
202    if set_file_attr {
203        attrs.set_item(
204            identifier!(vm, __file__),
205            code_obj.source_path().to_object(),
206            vm,
207        )?;
208    }
209    let module = vm.new_module(module_name, attrs.clone(), None);
210
211    // Store module in cache to prevent infinite loop with mutual importing libs:
212    let sys_modules = vm.sys_module.get_attr("modules", vm)?;
213    sys_modules.set_item(module_name, module.clone().into(), vm)?;
214
215    // Execute main code in module:
216    let scope = Scope::with_builtins(None, attrs, vm);
217    vm.run_code_obj(code_obj, scope)?;
218    Ok(module.into())
219}
220
221fn remove_importlib_frames_inner(
222    vm: &VirtualMachine,
223    tb: Option<PyRef<PyTraceback>>,
224    always_trim: bool,
225) -> (Option<PyRef<PyTraceback>>, bool) {
226    let traceback = if let Some(tb) = tb {
227        tb
228    } else {
229        return (None, false);
230    };
231
232    let file_name = traceback.frame.iframe().code().source_path().as_str();
233
234    let (inner_tb, mut now_in_importlib) =
235        remove_importlib_frames_inner(vm, traceback.next.lock().clone(), always_trim);
236    if file_name == "<frozen importlib._bootstrap>"
237        || file_name == "<frozen importlib._bootstrap_external>"
238        || file_name == "_frozen_importlib"
239        || file_name == "_frozen_importlib_external"
240    {
241        if traceback.frame.iframe().code().obj_name.as_str() == "_call_with_frames_removed" {
242            now_in_importlib = true;
243        }
244        if always_trim || now_in_importlib {
245            return (inner_tb, now_in_importlib);
246        }
247    } else {
248        now_in_importlib = false;
249    }
250
251    (
252        Some(
253            PyTraceback::new(
254                inner_tb,
255                traceback.frame.clone(),
256                traceback.lasti,
257                traceback.lineno,
258            )
259            .into_ref(&vm.ctx),
260        ),
261        now_in_importlib,
262    )
263}
264
265// TODO: This function should do nothing on verbose mode.
266// TODO: Fix this function after making PyTraceback.next mutable
267pub fn remove_importlib_frames(vm: &VirtualMachine, exc: &Py<PyBaseException>) {
268    if vm.state.config.settings.verbose != 0 {
269        return;
270    }
271
272    let always_trim = exc.fast_isinstance(vm.ctx.exceptions.import_error);
273
274    if let Some(tb) = exc.traceback() {
275        let trimmed_tb = remove_importlib_frames_inner(vm, Some(tb), always_trim).0;
276        exc.set_traceback(trimmed_tb);
277    }
278}
279
280/// Get origin path from a module spec, checking has_location first.
281pub(crate) fn get_spec_file_origin(spec: Option<&PyObject>, vm: &VirtualMachine) -> Option<String> {
282    let spec = spec?;
283
284    let has_location = spec
285        .get_attr("has_location", vm)
286        .ok()
287        .and_then(|v| v.try_to_bool(vm).ok())
288        .unwrap_or(false);
289    if !has_location {
290        return None;
291    }
292    spec.get_attr("origin", vm).ok().and_then(|origin| {
293        if vm.is_none(&origin) {
294            None
295        } else {
296            origin
297                .downcast_ref::<PyStr>()
298                .and_then(|s| s.to_str().map(|s| s.to_owned()))
299        }
300    })
301}
302
303/// Check if a module file possibly shadows another module of the same name.
304/// Compares the module's directory with the original sys.path[0] (derived from sys.argv[0]).
305pub(crate) fn is_possibly_shadowing_path(origin: &str, vm: &VirtualMachine) -> bool {
306    use std::path::Path;
307
308    if vm.state.config.settings.safe_path {
309        return false;
310    }
311
312    let origin_path = Path::new(origin);
313    let parent = match origin_path.parent() {
314        Some(p) => p,
315        None => return false,
316    };
317    // For packages (__init__.py), look one directory further up
318    let root = if origin_path.file_name() == Some("__init__.py".as_ref()) {
319        parent.parent().unwrap_or_else(|| Path::new(""))
320    } else {
321        parent
322    };
323
324    // Compute original sys.path[0] from sys.argv[0] (the script path).
325    // See: config->sys_path_0, which is set once
326    // at initialization and never changes even if sys.path is modified.
327    let sys_path_0 = (|| -> Option<String> {
328        let argv = vm.sys_module.get_attr("argv", vm).ok()?;
329        let argv0 = argv.get_item(&0usize, vm).ok()?;
330        let argv0_str = argv0.downcast_ref::<PyUtf8Str>()?;
331        let s = argv0_str.as_str();
332
333        // For -c and REPL, original sys.path[0] is ""
334        if s == "-c" || s.is_empty() {
335            return Some(String::new());
336        }
337        // For scripts, original sys.path[0] is dirname(argv[0])
338        Some(
339            Path::new(s)
340                .parent()
341                .and_then(|p| p.to_str())
342                .unwrap_or("")
343                .to_owned(),
344        )
345    })();
346
347    let sys_path_0 = match sys_path_0 {
348        Some(p) => p,
349        None => return false,
350    };
351
352    let cmp_path = if sys_path_0.is_empty() {
353        match crate::host_env::os::current_dir() {
354            Ok(d) => d.to_string_lossy().to_string(),
355            Err(_) => return false,
356        }
357    } else {
358        sys_path_0
359    };
360
361    root.to_str() == Some(cmp_path.as_str())
362}
363
364/// Check if a module name is in sys.stdlib_module_names.
365/// Takes the original __name__ object to preserve str subclass behavior.
366/// Propagates errors (e.g. TypeError for unhashable str subclass).
367pub(crate) fn is_stdlib_module_name(name: &PyObject, vm: &VirtualMachine) -> PyResult<bool> {
368    let stdlib_names = match vm.sys_module.get_attr("stdlib_module_names", vm) {
369        Ok(names) => names,
370        Err(_) => return Ok(false),
371    };
372    if !stdlib_names.class().fast_issubclass(vm.ctx.types.set_type)
373        && !stdlib_names
374            .class()
375            .fast_issubclass(vm.ctx.types.frozenset_type)
376    {
377        return Ok(false);
378    }
379    let result = vm.call_method(&stdlib_names, "__contains__", (name.to_owned(),))?;
380    result.try_to_bool(vm)
381}
382
383/// PyImport_ImportModuleLevelObject
384pub(crate) fn import_module_level(
385    name: &Py<PyStr>,
386    globals: Option<&PyObject>,
387    fromlist: Option<PyObjectRef>,
388    level: i32,
389    vm: &VirtualMachine,
390) -> PyResult {
391    if level < 0 {
392        return Err(vm.new_value_error("level must be >= 0"));
393    }
394
395    let name = match name.as_utf8() {
396        Some(name) => name,
397        None => {
398            // Name contains surrogates. Try sys.modules lookup with the
399            // Python string key directly.
400            if level == 0 {
401                let sys_modules = vm.sys_module.get_attr("modules", vm)?;
402                return sys_modules.get_item(name, vm).map_err(|_| {
403                    vm.new_import_error(format!("No module named '{name}'"), name.to_owned())
404                });
405            }
406            return Err(vm.new_import_error(format!("No module named '{name}'"), name.to_owned()));
407        }
408    };
409
410    // Resolve absolute name
411    let abs_name = if level > 0 {
412        // When globals is not provided (Rust None), raise KeyError
413        // matching resolve_name() where globals==NULL
414        let Some(globals_ref) = globals else {
415            return Err(vm.new_key_error(vm.ctx.new_str("'__name__' not in globals").into()));
416        };
417        // When globals is Python None, treat like empty mapping
418        let empty_dict_obj: PyObjectRef;
419        let globals_ref = if vm.is_none(globals_ref) {
420            empty_dict_obj = vm.ctx.new_dict().into();
421            &empty_dict_obj
422        } else {
423            globals_ref
424        };
425        let package = calc_package(Some(globals_ref), vm)?;
426        if package.is_empty() {
427            return Err(vm.new_import_error(
428                "attempted relative import with no known parent package",
429                vm.ctx.new_utf8_str(""),
430            ));
431        }
432        resolve_name(name, &package, level as usize, vm)?
433    } else {
434        if name.is_empty() {
435            return Err(vm.new_value_error("Empty module name"));
436        }
437        name.to_owned()
438    };
439
440    // import_get_module + import_find_and_load
441    let sys_modules = vm.sys_module.get_attr("modules", vm)?;
442    let module = match sys_modules.get_item(&*abs_name, vm) {
443        Ok(m) if !vm.is_none(&m) => {
444            import_ensure_initialized(&m, &abs_name, vm)?;
445            m
446        }
447        _ => {
448            let find_and_load = vm.importlib.get_attr("_find_and_load", vm)?;
449            find_and_load.call((abs_name.clone(), vm.import_func.clone()), vm)?
450        }
451    };
452
453    // Handle fromlist
454    let has_from = match fromlist.as_ref().filter(|fl| !vm.is_none(fl)) {
455        Some(fl) => fl.try_to_bool(vm)?,
456        None => false,
457    };
458
459    if has_from {
460        let fromlist = fromlist.unwrap();
461        // Only call _handle_fromlist if the module looks like a package
462        // (has __path__). Non-module objects without __name__/__path__ would
463        // crash inside _handle_fromlist; IMPORT_FROM handles per-attribute
464        // errors with proper ImportError conversion.
465        let has_path = vm
466            .get_attribute_opt(&module, vm.ctx.intern_str("__path__"))?
467            .is_some();
468        if has_path {
469            let handle_fromlist = vm.importlib.get_attr("_handle_fromlist", vm)?;
470            handle_fromlist.call((module, fromlist, vm.import_func.clone()), vm)
471        } else {
472            Ok(module)
473        }
474    } else if level == 0 || !name.is_empty() {
475        let name_str = name.as_str();
476        match name_str.find('.') {
477            None => Ok(module),
478            Some(dot) => {
479                let to_return = if level == 0 {
480                    vm.ctx.new_utf8_str(&name_str[..dot])
481                } else {
482                    let cut_off = name_str.len() - dot;
483                    let abs = abs_name.as_str();
484                    vm.ctx.new_utf8_str(&abs[..abs.len() - cut_off])
485                };
486                match sys_modules.get_item(&*to_return, vm) {
487                    Ok(m) => Ok(m),
488                    Err(_) if level == 0 => {
489                        // For absolute imports (level 0), try importing the
490                        // parent. Matches _bootstrap.__import__ behavior.
491                        let find_and_load = vm.importlib.get_attr("_find_and_load", vm)?;
492                        find_and_load.call((to_return, vm.import_func.clone()), vm)
493                    }
494                    Err(_) => {
495                        // For relative imports (level > 0), raise KeyError
496                        let to_return_obj: PyObjectRef = vm
497                            .ctx
498                            .new_utf8_str(format!("'{to_return}' not in sys.modules as expected"))
499                            .into();
500                        Err(vm.new_key_error(to_return_obj))
501                    }
502                }
503            }
504        }
505    } else {
506        Ok(module)
507    }
508}
509
510/// resolve_name in import.c - resolve relative import name
511fn resolve_name(
512    name: &Py<PyUtf8Str>,
513    package: &Py<PyUtf8Str>,
514    level: usize,
515    vm: &VirtualMachine,
516) -> PyResult<PyUtf8StrRef> {
517    // Python: bits = package.rsplit('.', level - 1)
518    // Rust: rsplitn(level, '.') gives maxsplit=level-1
519    let package_str = package.as_str();
520    let parts: Vec<&str> = package_str.rsplitn(level, '.').collect();
521    if parts.len() < level {
522        return Err(vm.new_import_error(
523            "attempted relative import beyond top-level package",
524            name.to_owned(),
525        ));
526    }
527    // rsplitn returns parts right-to-left, so last() is the leftmost (base)
528    let base = *parts.last().unwrap();
529    let abs_name = if name.is_empty() {
530        if base.len() == package_str.len() {
531            package.to_owned()
532        } else {
533            vm.ctx.new_utf8_str(base)
534        }
535    } else {
536        vm.ctx.new_utf8_str(format!("{base}.{}", name.as_str()))
537    };
538    Ok(abs_name)
539}
540
541/// _calc___package__ - calculate package from globals for relative imports
542fn calc_package(globals: Option<&PyObject>, vm: &VirtualMachine) -> PyResult<PyUtf8StrRef> {
543    let globals = globals.ok_or_else(|| {
544        vm.new_import_error(
545            "attempted relative import with no known parent package",
546            vm.ctx.new_utf8_str(""),
547        )
548    })?;
549
550    let package = globals.get_item("__package__", vm).ok();
551    let spec = globals.get_item("__spec__", vm).ok();
552
553    if let Some(ref pkg) = package
554        && !vm.is_none(pkg)
555    {
556        let pkg_str: PyUtf8StrRef = pkg
557            .clone()
558            .downcast()
559            .map_err(|_| vm.new_type_error("package must be a string"))?;
560        // Warn if __package__ != __spec__.parent
561        if let Some(ref spec) = spec
562            && !vm.is_none(spec)
563            && let Ok(parent) = spec.get_attr("parent", vm)
564            && !pkg_str.is(&parent)
565            && pkg_str
566                .as_object()
567                .rich_compare_bool(&parent, crate::types::PyComparisonOp::Ne, vm)
568                .unwrap_or(false)
569        {
570            let parent_repr = parent
571                .repr_utf8(vm)
572                .map(|s| s.as_str().to_owned())
573                .unwrap_or_default();
574            let msg = format!(
575                "__package__ != __spec__.parent ('{}' != {})",
576                pkg_str.as_str(),
577                parent_repr
578            );
579            let warn = vm
580                .import("_warnings", 0)
581                .and_then(|w| w.get_attr("warn", vm));
582            if let Ok(warn_fn) = warn {
583                let _ = warn_fn.call(
584                    (
585                        vm.ctx.new_str(msg),
586                        vm.ctx.exceptions.deprecation_warning.to_owned(),
587                    ),
588                    vm,
589                );
590            }
591        }
592        return Ok(pkg_str);
593    } else if let Some(ref spec) = spec
594        && !vm.is_none(spec)
595        && let Ok(parent) = spec.get_attr("parent", vm)
596        && !vm.is_none(&parent)
597    {
598        let parent_str: PyUtf8StrRef = parent
599            .downcast()
600            .map_err(|_| vm.new_type_error("package set to non-string"))?;
601        return Ok(parent_str);
602    }
603
604    // Fall back to __name__ and __path__
605    let warn = vm
606        .import("_warnings", 0)
607        .and_then(|w| w.get_attr("warn", vm));
608    if let Ok(warn_fn) = warn {
609        let _ = warn_fn.call(
610            (
611                vm.ctx.new_str("can't resolve package from __spec__ or __package__, falling back on __name__ and __path__"),
612                vm.ctx.exceptions.import_warning.to_owned(),
613            ),
614            vm,
615        );
616    }
617
618    let mod_name = globals.get_item("__name__", vm).map_err(|_| {
619        vm.new_import_error(
620            "attempted relative import with no known parent package",
621            vm.ctx.new_utf8_str(""),
622        )
623    })?;
624    let mod_name_str: PyUtf8StrRef = mod_name
625        .downcast()
626        .map_err(|_| vm.new_type_error("__name__ must be a string"))?;
627    // If not a package (no __path__), strip last component.
628    // Uses rpartition('.')[0] semantics: returns empty string when no dot.
629    if globals.get_item("__path__", vm).is_err() {
630        let s = mod_name_str.as_str();
631        Ok(match s.rfind('.') {
632            Some(dot) => vm.ctx.new_utf8_str(&s[..dot]),
633            None => vm.ctx.new_utf8_str(""),
634        })
635    } else {
636        Ok(mod_name_str)
637    }
638}