ty_python_core/
ast_node_ref.rs1use std::fmt::Debug;
2use std::marker::PhantomData;
3
4#[cfg(debug_assertions)]
5use ruff_db::files::File;
6use ruff_db::parsed::ParsedModuleRef;
7#[cfg(debug_assertions)]
8use ruff_python_ast::PythonVersion;
9use ruff_python_ast::{AnyNodeRef, NodeIndex};
10use ruff_python_ast::{AnyRootNodeRef, HasNodeIndex};
11use ruff_text_size::Ranged;
12
13#[derive(Clone)]
38pub struct AstNodeRef<T> {
39 index: NodeIndex,
41
42 #[cfg(debug_assertions)]
44 kind: ruff_python_ast::NodeKind,
45 #[cfg(debug_assertions)]
46 range: ruff_text_size::TextRange,
47 #[cfg(debug_assertions)]
51 file: File,
52 #[cfg(debug_assertions)]
53 python_version: PythonVersion,
54
55 _node: PhantomData<T>,
56}
57
58impl<T> AstNodeRef<T> {
59 pub fn index(&self) -> NodeIndex {
60 self.index
61 }
62}
63
64impl<T> AstNodeRef<T>
65where
66 T: HasNodeIndex + Ranged + PartialEq + Debug,
67 for<'ast> AnyNodeRef<'ast>: From<&'ast T>,
68 for<'ast> &'ast T: TryFrom<AnyRootNodeRef<'ast>>,
69{
70 pub(super) fn new(module_ref: &ParsedModuleRef, node: &T) -> Self {
75 let index = node.node_index().load();
76 debug_assert_eq!(module_ref.get_by_index(index).try_into().ok(), Some(node));
77
78 Self {
79 index,
80 #[cfg(debug_assertions)]
81 file: module_ref.module().file(),
82 #[cfg(debug_assertions)]
83 python_version: module_ref.module().python_version(),
84 #[cfg(debug_assertions)]
85 kind: AnyNodeRef::from(node).kind(),
86 #[cfg(debug_assertions)]
87 range: node.range(),
88 _node: PhantomData,
89 }
90 }
91
92 #[track_caller]
97 pub fn node<'ast>(&self, module_ref: &'ast ParsedModuleRef) -> &'ast T {
98 #[cfg(debug_assertions)]
99 assert_eq!(
100 (
101 module_ref.module().file(),
102 module_ref.module().python_version()
103 ),
104 (self.file, self.python_version),
105 "an `AstNodeRef` cannot be used with a module parsed for a different file or Python version"
106 );
107 module_ref
110 .get_by_index(self.index)
111 .try_into()
112 .ok()
113 .expect("AST indices should never change within the same revision")
114 }
115}
116
117impl<T> get_size2::GetSize for AstNodeRef<T> {}
118
119#[expect(clippy::missing_fields_in_debug)]
120impl<T> Debug for AstNodeRef<T>
121where
122 T: Debug,
123 for<'ast> &'ast T: TryFrom<AnyRootNodeRef<'ast>>,
124{
125 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
126 cfg_select! {
127 debug_assertions => {
128 f.debug_struct("AstNodeRef")
129 .field("kind", &self.kind)
130 .field("range", &self.range)
131 .finish()
132 },
133 _ => {
134 f.debug_tuple("AstNodeRef").finish_non_exhaustive()
136 },
137 }
138 }
139}
140
141#[cfg(all(test, debug_assertions))]
142mod tests {
143 use ruff_db::PythonFile;
144 use ruff_db::files::system_path_to_file;
145 use ruff_db::parsed::parsed_module;
146 use ruff_python_ast::PythonVersion;
147
148 use crate::ast_node_ref::AstNodeRef;
149 use crate::db::tests::TestDbBuilder;
150
151 #[test]
152 #[should_panic(
153 expected = "an `AstNodeRef` cannot be used with a module parsed for a different file or Python version"
154 )]
155 fn rejects_module_parsed_for_different_python_version() {
156 let db = TestDbBuilder::new()
157 .with_file("test.py", "x = 1")
158 .build()
159 .unwrap();
160 let file = system_path_to_file(&db, "test.py").unwrap();
161
162 let parsed_py311 =
163 parsed_module(&db, PythonFile::new(&db, file, PythonVersion::PY311)).load(&db);
164 let parsed_py312 =
165 parsed_module(&db, PythonFile::new(&db, file, PythonVersion::PY312)).load(&db);
166 let assignment = parsed_py311.syntax().body[0].as_assign_stmt().unwrap();
167
168 let node = AstNodeRef::new(&parsed_py311, assignment);
169 node.node(&parsed_py312);
170 }
171}