use std::fmt;
use std::num::NonZeroU32;
use std::ops::Deref;
use compact_str::{CompactString, ToCompactString};
use ruff_db::files::File;
use ruff_python_ast::{self as ast, PythonVersion};
use ruff_python_stdlib::identifiers::is_identifier;
use crate::db::Db;
use crate::resolve::file_to_module;
use crate::{ResolverEnvironment, ResolverFile};
#[derive(Clone, Debug, Eq, PartialEq, Hash, PartialOrd, Ord, get_size2::GetSize)]
pub struct ModuleName(compact_str::CompactString);
impl ModuleName {
#[inline]
#[must_use]
pub fn new(name: &str) -> Option<Self> {
Self::is_valid_name(name).then(|| Self(CompactString::from(name)))
}
#[inline]
#[must_use]
pub fn new_static(name: &'static str) -> Option<Self> {
Self::is_valid_name(name).then(|| Self(CompactString::const_new(name)))
}
#[must_use]
fn is_valid_name(name: &str) -> bool {
!name.is_empty() && name.split('.').all(is_identifier)
}
#[must_use]
pub fn components(&self) -> impl DoubleEndedIterator<Item = &str> {
self.0.split('.')
}
#[must_use]
pub fn first_component(&self) -> &str {
self.components()
.next()
.expect("at least one module component")
}
#[must_use]
pub fn last_component(&self) -> &str {
self.components()
.next_back()
.expect("at least one module component")
}
#[must_use]
pub fn parent(&self) -> Option<ModuleName> {
let (parent, _) = self.0.rsplit_once('.')?;
Some(Self(parent.to_compact_string()))
}
#[must_use]
pub fn starts_with(&self, other: &ModuleName) -> bool {
let mut self_components = self.components();
let other_components = other.components();
for other_component in other_components {
if self_components.next() != Some(other_component) {
return false;
}
}
true
}
#[must_use]
pub fn relative_to(&self, parent: &ModuleName) -> Option<ModuleName> {
let relative_name = self.0.strip_prefix(&*parent.0)?.strip_prefix('.')?;
assert!(!relative_name.is_empty());
debug_assert!(self.starts_with(parent));
Some(ModuleName(CompactString::from(relative_name)))
}
#[must_use]
#[inline]
pub fn as_str(&self) -> &str {
&self.0
}
#[must_use]
pub fn from_components<'a>(components: impl IntoIterator<Item = &'a str>) -> Option<Self> {
let mut components = components.into_iter();
let first_part = components.next()?;
if !is_identifier(first_part) {
return None;
}
let mut name = CompactString::from(first_part);
for part in components {
if !is_identifier(part) {
return None;
}
name.push('.');
name.push_str(part);
}
Some(Self(name))
}
pub fn extend(&mut self, other: &ModuleName) {
self.0.push('.');
self.0.push_str(other);
}
pub fn ancestors(&self) -> impl Iterator<Item = Self> {
std::iter::successors(Some(self.clone()), Self::parent)
}
pub fn from_import_statement<'db>(
db: &'db dyn Db,
importing_file: ImportingFile<'db>,
node: &ast::StmtImportFrom,
) -> Result<Self, ModuleNameResolutionError> {
let ast::StmtImportFrom {
module,
level,
names: _,
is_lazy: _,
range: _,
node_index: _,
} = node;
Self::from_identifier_parts(db, importing_file, module.as_deref(), *level)
}
pub fn from_identifier_parts<'db>(
db: &'db dyn Db,
importing_file: ImportingFile<'db>,
module: Option<&str>,
level: u32,
) -> Result<Self, ModuleNameResolutionError> {
if let Some(level) = NonZeroU32::new(level) {
relative_module_name(db, importing_file.resolver_file(db), module, level)
} else {
module
.and_then(Self::new)
.ok_or(ModuleNameResolutionError::InvalidSyntax)
}
}
pub fn package_for_file<'db>(
db: &'db dyn Db,
importing_file: ImportingFile<'db>,
) -> Result<Self, ModuleNameResolutionError> {
Self::from_identifier_parts(db, importing_file, None, 1)
}
pub fn is_test_module(&self) -> bool {
if self.last_component() == "conftest" {
return true;
}
self.components()
.skip(1)
.any(|c| c == "test" || c == "tests")
}
pub fn is_private(&self) -> bool {
self.components().any(|c| c.starts_with('_'))
}
}
impl Deref for ModuleName {
type Target = str;
#[inline]
fn deref(&self) -> &Self::Target {
self.as_str()
}
}
impl PartialEq<str> for ModuleName {
fn eq(&self, other: &str) -> bool {
self.as_str() == other
}
}
impl PartialEq<ModuleName> for str {
fn eq(&self, other: &ModuleName) -> bool {
self == other.as_str()
}
}
impl std::fmt::Display for ModuleName {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(&self.0)
}
}
#[derive(Clone, Copy)]
pub enum ImportingFile<'db> {
ResolverFile(ResolverFile<'db>),
File(File, ResolverEnvironment<'db>),
}
impl<'db> ImportingFile<'db> {
pub fn file(self, db: &dyn Db) -> File {
match self {
Self::ResolverFile(file) => file.file(db),
Self::File(file, _) => file,
}
}
pub fn resolver_environment(self, db: &'db dyn Db) -> ResolverEnvironment<'db> {
match self {
Self::ResolverFile(file) => file.environment(db),
Self::File(_, resolver_environment) => resolver_environment,
}
}
pub fn python_version(self, db: &'db dyn Db) -> PythonVersion {
self.resolver_environment(db).python_version(db)
}
pub fn resolver_file(self, db: &'db dyn Db) -> ResolverFile<'db> {
match self {
Self::ResolverFile(file) => file,
Self::File(file, resolver_environment) => {
ResolverFile::new(db, file, resolver_environment)
}
}
}
}
fn relative_module_name<'db>(
db: &'db dyn Db,
importing_file: ResolverFile<'db>,
tail: Option<&str>,
level: NonZeroU32,
) -> Result<ModuleName, ModuleNameResolutionError> {
let module = file_to_module(db, importing_file)
.ok_or(ModuleNameResolutionError::UnknownCurrentModule)?;
let mut level = level.get();
if module.kind(db).is_package() {
level = level.saturating_sub(1);
}
let mut module_name = module
.name(db)
.ancestors()
.nth(level as usize)
.ok_or(ModuleNameResolutionError::TooManyDots)?;
if let Some(tail) = tail {
let tail = ModuleName::new(tail).ok_or(ModuleNameResolutionError::InvalidSyntax)?;
module_name.extend(&tail);
}
Ok(module_name)
}
#[derive(Debug, Copy, Clone, PartialEq, Eq)]
pub enum ModuleNameResolutionError {
InvalidSyntax,
UnknownCurrentModule,
TooManyDots,
}