use std::fmt::Debug;
use std::marker::PhantomData;
#[cfg(debug_assertions)]
use ruff_db::files::File;
use ruff_db::parsed::ParsedModuleRef;
#[cfg(debug_assertions)]
use ruff_python_ast::PythonVersion;
use ruff_python_ast::{AnyNodeRef, NodeIndex};
use ruff_python_ast::{AnyRootNodeRef, HasNodeIndex};
use ruff_text_size::Ranged;
#[derive(Clone)]
pub struct AstNodeRef<T> {
index: NodeIndex,
#[cfg(debug_assertions)]
kind: ruff_python_ast::NodeKind,
#[cfg(debug_assertions)]
range: ruff_text_size::TextRange,
#[cfg(debug_assertions)]
file: File,
#[cfg(debug_assertions)]
python_version: PythonVersion,
_node: PhantomData<T>,
}
impl<T> AstNodeRef<T> {
pub fn index(&self) -> NodeIndex {
self.index
}
}
impl<T> AstNodeRef<T>
where
T: HasNodeIndex + Ranged + PartialEq + Debug,
for<'ast> AnyNodeRef<'ast>: From<&'ast T>,
for<'ast> &'ast T: TryFrom<AnyRootNodeRef<'ast>>,
{
pub(super) fn new(module_ref: &ParsedModuleRef, node: &T) -> Self {
let index = node.node_index().load();
debug_assert_eq!(module_ref.get_by_index(index).try_into().ok(), Some(node));
Self {
index,
#[cfg(debug_assertions)]
file: module_ref.module().file(),
#[cfg(debug_assertions)]
python_version: module_ref.module().python_version(),
#[cfg(debug_assertions)]
kind: AnyNodeRef::from(node).kind(),
#[cfg(debug_assertions)]
range: node.range(),
_node: PhantomData,
}
}
#[track_caller]
pub fn node<'ast>(&self, module_ref: &'ast ParsedModuleRef) -> &'ast T {
#[cfg(debug_assertions)]
assert_eq!(
(
module_ref.module().file(),
module_ref.module().python_version()
),
(self.file, self.python_version),
"an `AstNodeRef` cannot be used with a module parsed for a different file or Python version"
);
module_ref
.get_by_index(self.index)
.try_into()
.ok()
.expect("AST indices should never change within the same revision")
}
}
impl<T> get_size2::GetSize for AstNodeRef<T> {}
#[expect(clippy::missing_fields_in_debug)]
impl<T> Debug for AstNodeRef<T>
where
T: Debug,
for<'ast> &'ast T: TryFrom<AnyRootNodeRef<'ast>>,
{
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
cfg_select! {
debug_assertions => {
f.debug_struct("AstNodeRef")
.field("kind", &self.kind)
.field("range", &self.range)
.finish()
},
_ => {
f.debug_tuple("AstNodeRef").finish_non_exhaustive()
},
}
}
}
#[cfg(all(test, debug_assertions))]
mod tests {
use ruff_db::PythonFile;
use ruff_db::files::system_path_to_file;
use ruff_db::parsed::parsed_module;
use ruff_python_ast::PythonVersion;
use crate::ast_node_ref::AstNodeRef;
use crate::db::tests::TestDbBuilder;
#[test]
#[should_panic(
expected = "an `AstNodeRef` cannot be used with a module parsed for a different file or Python version"
)]
fn rejects_module_parsed_for_different_python_version() {
let db = TestDbBuilder::new()
.with_file("test.py", "x = 1")
.build()
.unwrap();
let file = system_path_to_file(&db, "test.py").unwrap();
let parsed_py311 =
parsed_module(&db, PythonFile::new(&db, file, PythonVersion::PY311)).load(&db);
let parsed_py312 =
parsed_module(&db, PythonFile::new(&db, file, PythonVersion::PY312)).load(&db);
let assignment = parsed_py311.syntax().body[0].as_assign_stmt().unwrap();
let node = AstNodeRef::new(&parsed_py311, assignment);
node.node(&parsed_py312);
}
}