1use super::*;
19use derive_generic_visitor::*;
20use hax_lib_macros_types::AttrPayload;
21
22pub mod wrappers {
23 use std::ops::Deref;
30
31 use super::{infallible::AstVisitable as AstVisitableInfallible, *};
32 use diagnostics::*;
33
34 pub struct SpanWrapper<'a, V>(pub &'a mut V);
38
39 impl<'a, V: HasSpan> SpanWrapper<'a, V> {
40 fn spanned_action<T: Deref, U>(
44 &mut self,
45 ast_fragment: T,
46 action: impl Fn(&mut Self, T) -> U,
47 ) -> U
48 where
49 T::Target: HasSpan,
50 {
51 let span_before = self.0.span();
52 *self.0.span_mut() = ast_fragment.span();
53 let result = action(self, ast_fragment);
55 *self.0.span_mut() = span_before;
56 result
57 }
58 }
59
60 impl<'a, V: AstVisitorMut + HasSpan> AstVisitorMut for SpanWrapper<'a, V> {
61 fn visit_inner<T>(&mut self, x: &mut T)
62 where
63 T: AstVisitableInfallible,
64 T: for<'s> DriveMut<'s, AstVisitableInfallibleWrapper<Self>>,
65 {
66 x.drive_map(self.0)
67 }
68 fn visit_item(&mut self, x: &mut Item) {
69 self.spanned_action(x, Self::visit_inner)
70 }
71 fn visit_expr(&mut self, x: &mut Expr) {
72 self.spanned_action(x, Self::visit_inner)
73 }
74 fn visit_pat(&mut self, x: &mut Pat) {
75 self.spanned_action(x, Self::visit_inner)
76 }
77 fn visit_guard(&mut self, x: &mut Guard) {
78 self.spanned_action(x, Self::visit_inner)
79 }
80 fn visit_arm(&mut self, x: &mut Arm) {
81 self.spanned_action(x, Self::visit_inner)
82 }
83 fn visit_impl_item(&mut self, x: &mut ImplItem) {
84 self.spanned_action(x, Self::visit_inner)
85 }
86 fn visit_trait_item(&mut self, x: &mut TraitItem) {
87 self.spanned_action(x, Self::visit_inner)
88 }
89 fn visit_generic_param(&mut self, x: &mut GenericParam) {
90 self.spanned_action(x, Self::visit_inner)
91 }
92 fn visit_attribute(&mut self, x: &mut Attribute) {
93 self.spanned_action(x, Self::visit_inner)
94 }
95 fn visit_spanned_ty(&mut self, x: &mut SpannedTy) {
96 self.spanned_action(x, Self::visit_inner)
97 }
98 }
99
100 pub struct ErrorWrapper<'a, V>(pub &'a mut V);
105
106 #[derive(Default)]
109 pub struct ErrorVault(Vec<Diagnostic>);
110 impl ErrorVault {
111 fn add(&mut self, diagnostic: Diagnostic) {
112 self.0.push(diagnostic);
113 }
114 }
115
116 pub struct ErrorHandlingState(pub Span, pub ErrorVault);
119 impl Default for ErrorHandlingState {
120 fn default() -> Self {
121 Self(Span::dummy(), Default::default())
122 }
123 }
124
125 #[macro_export]
126 macro_rules! setup_error_handling_impl {
128 () => {
129 fn visit<T: $crate::ast::visitors::AstVisitableInfallible>(&mut self, x: &mut T) {
130 $crate::ast::visitors::wrappers::SpanWrapper(
131 &mut $crate::ast::visitors::wrappers::ErrorWrapper(self),
132 )
133 .visit(x)
134 }
135 };
136 }
137 pub use setup_error_handling_impl;
138
139 pub trait VisitorWithContext {
141 fn context(&self) -> Context;
143 }
144
145 impl<T: HasSpan> HasSpan for ErrorWrapper<'_, T> {
146 fn span(&self) -> Span {
147 self.0.span()
148 }
149
150 fn span_mut(&mut self) -> &mut Span {
151 self.0.span_mut()
152 }
153 }
154
155 pub trait VisitorWithErrors: HasSpan + VisitorWithContext {
161 fn error_vault(&mut self) -> &mut ErrorVault;
163 fn error(&mut self, node: impl Into<Fragment>, kind: DiagnosticInfoKind) {
165 let context = self.context();
166 let span = self.span();
167 self.error_vault().add(Diagnostic::new(
168 node,
169 DiagnosticInfo {
170 context,
171 span,
172 kind,
173 },
174 ));
175 }
176 }
177
178 impl<'a, V: VisitorWithErrors> ErrorWrapper<'a, V> {
179 fn error_handled_action<
180 T: FallibleAstNode + Clone + std::fmt::Debug + Into<Fragment>,
181 U,
182 >(
183 &mut self,
184 x: &mut T,
185 action: impl Fn(&mut Self, &mut T) -> U,
186 ) -> U {
187 let diagnostics_snapshot = self.0.error_vault().0.clone();
188 self.0.error_vault().0.clear();
189 let result = action(self, x);
190 let diagnostics: Vec<_> = self.0.error_vault().0.drain(..).collect();
191 if !diagnostics.is_empty() {
192 x.set_error(ErrorNode {
193 fragment: Box::new(x.clone().into()),
194 diagnostics,
195 });
196 }
197 self.0.error_vault().0 = diagnostics_snapshot;
198 result
199 }
200 }
201
202 impl<'a, V: AstVisitorMut + VisitorWithErrors> AstVisitorMut for ErrorWrapper<'a, V> {
203 fn visit_inner<T>(&mut self, x: &mut T)
204 where
205 T: AstVisitableInfallible,
206 T: for<'s> DriveMut<'s, AstVisitableInfallibleWrapper<Self>>,
207 {
208 x.drive_map(self.0)
209 }
210 fn visit_item(&mut self, x: &mut Item) {
211 self.error_handled_action(x, Self::visit_inner)
212 }
213 fn visit_pat(&mut self, x: &mut Pat) {
214 self.error_handled_action(x, Self::visit_inner)
215 }
216 fn visit_expr(&mut self, x: &mut Expr) {
217 self.error_handled_action(x, Self::visit_inner)
218 }
219 fn visit_ty(&mut self, x: &mut Ty) {
220 self.error_handled_action(x, Self::visit_inner)
221 }
222 }
223}
224
225#[hax_rust_engine_macros::replace(AstNodes => include(VisitableAstNodes))]
226mod replaced {
227 use super::*;
228 pub mod infallible {
229 use super::*;
230
231 #[visitable_group(
232 visitor(drive_map(
233 &mut AstVisitorMut
254 ), infallible),
255 visitor(drive(
256 &AstVisitor
258 ), infallible),
259 skip(
260 String, bool, char, hax_frontend_exporter::Span,
261 ),
262 drive(
263 for<T: AstVisitable> Box<T>, for<T: AstVisitable> Option<T>, for<T: AstVisitable> Vec<T>,
264 for<A: AstVisitable, B: AstVisitable> (A, B),
265 for<A: AstVisitable, B: AstVisitable, C: AstVisitable> (A, B, C),
266 usize
267 ),
268 override(AstNodes),
269 override_skip(
270 Span, Fragment, GlobalId, Diagnostic, AttrPayload,
271 ),
272 )]
273 pub trait AstVisitable {}
275 }
276
277 #[allow(missing_docs)]
278 pub mod fallible {
279 use super::*;
280
281 #[visitable_group(
282 visitor(drive(
283 &AstEarlyExitVisitor
285 )),
286 visitor(drive_mut(
287 &mut AstEarlyExitVisitorMut
289 )),
290 skip(
291 String, bool, char, hax_frontend_exporter::Span,
292 ),
293 drive(
294 for<T: AstVisitable> Box<T>, for<T: AstVisitable> Option<T>, for<T: AstVisitable> Vec<T>,
295 for<A: AstVisitable, B: AstVisitable> (A, B),
296 for<A: AstVisitable, B: AstVisitable, C: AstVisitable> (A, B, C),
297 usize
298 ),
299 override(AstNodes),
300 override_skip(
301 Span, Fragment, GlobalId, Diagnostic, AttrPayload,
302 ),
303 )]
304 pub trait AstVisitable {}
306 }
307
308 pub mod dyn_compatible {
310 use super::*;
311
312 macro_rules! derive_erased_ast_visitors {
313 ({$($attrs:tt)*}, $name: ident, $helper: ident, $($ty:ty),*) => {
314 $($attrs)*
315 pub trait $name<'a>: $($helper<'a, $ty> + )* {}
316 };
317 }
318
319 macro_rules! render_path {
320 ($head:ident) => {stringify!($head)};
321 ($head:ident $(::$tail:ident)*) => {
322 concat!(stringify!($head), "::", render_path!($($tail)::*))
323 };
324 }
325
326 macro_rules! make_dyn_compatible {
327 ($($visitable_trait:ident)::*, $($visitor_trait:ident)::*, $helper_name: ident, $name: ident, mut:{$($mut:tt)?}, super:{$($super:ident)::*}, $ret:ty) => {
328 #[doc = concat!("A [dyn-compatible](https://doc.rust-lang.org/reference/items/traits.html#dyn-compatibility) trait similar to [`", render_path!($($visitor_trait)::*),"`].")]
329 #[doc = concat!("This trait provides one `visit` method to visit a given type `T` with a given visitor.")]
330 pub trait $helper_name<'a, T: ?Sized>: $($super)::* {
331 fn visit(&mut self, _: &'a $($mut)? T) -> $ret;
333 }
334
335 impl<'a, T: $($visitable_trait)::*, V: $($visitor_trait)::*> $helper_name<'a, T> for V {
336 fn visit(&mut self, e: &'a $($mut)? T) -> $ret {
337 <Self as $($visitor_trait)::*>::visit(self, e)
338 }
339 }
340 derive_erased_ast_visitors!({
341 #[doc = concat!("A [dyn-compatible](https://doc.rust-lang.org/reference/items/traits.html#dyn-compatibility) trait similar to [`", render_path!($($visitor_trait)::*),"`].")]
342 #[doc = concat!("This trait is empty, but it implies a super bound for every type in the AST, so that you can use [`", stringify!($helper_name), "::visit", "`] with the entire AST.")]
343 }, $name, $helper_name, AstNodes);
344
345 impl<'a, V: $($visitor_trait)::*> $name<'a> for V {}
346 };
347 }
348
349 make_dyn_compatible!(
350 infallible::AstVisitable,
351 infallible::AstVisitorMut,
352 AstVisitableMut,
353 AstVisitorMut,
354 mut:{mut},
355 super:{},
356 ()
357 );
358 make_dyn_compatible!(
359 infallible::AstVisitable,
360 infallible::AstVisitor,
361 AstVisitable,
362 AstVisitor,
363 mut:{},
364 super:{},
365 ()
366 );
367
368 make_dyn_compatible!(
369 fallible::AstVisitable,
370 fallible::AstEarlyExitVisitorMut,
371 AstEarlyExitVisitableMut,
372 AstEarlyExitVisitorMut,
373 mut:{mut},
374 super:{Visitor},
375 ControlFlow<<Self as Visitor>::Break>
376 );
377 make_dyn_compatible!(
378 fallible::AstVisitable,
379 fallible::AstEarlyExitVisitor,
380 AstEarlyExitVisitable,
381 AstEarlyExitVisitor,
382 mut:{},
383 super:{Visitor},
384 ControlFlow<<Self as Visitor>::Break>
385 );
386 }
387}
388
389pub use replaced::dyn_compatible;
390use replaced::{fallible, infallible};
391
392pub use fallible::{
393 AstEarlyExitVisitor, AstEarlyExitVisitorMut, AstVisitable as AstVisitableFallible,
394 AstVisitableWrapper,
395};
396pub use hax_rust_engine_macros::setup_error_handling_struct;
397pub use infallible::{
398 AstVisitable as AstVisitableInfallible, AstVisitableInfallibleWrapper, AstVisitor,
399 AstVisitorMut,
400};
401pub use wrappers::{VisitorWithContext, VisitorWithErrors, setup_error_handling_impl};
402
403#[test]
404fn double_literals_in_ast() {
405 use crate::ast::diagnostics::*;
406
407 #[setup_error_handling_struct]
408 #[derive(Default)]
409 struct DoubleU8Literals;
410
411 impl VisitorWithContext for DoubleU8Literals {
412 fn context(&self) -> Context {
413 Context::Import
414 }
415 }
416
417 impl AstVisitorMut for DoubleU8Literals {
418 setup_error_handling_impl!();
419
420 fn visit_literal(&mut self, x: &mut Literal) {
421 let Literal::Int { value, .. } = x else {
422 return;
423 };
424 let Ok(n): Result<u8, _> = str::parse(value) else {
425 return self.error(
426 x.clone(),
427 DiagnosticInfoKind::AssertionFailure {
428 details: "Bad literal".into(),
429 },
430 );
431 };
432 let n = (n as u16) * 2;
433 if n >= u8::MAX as u16 {
434 return self.error(
435 x.clone(),
436 DiagnosticInfoKind::AssertionFailure {
437 details: "Literal too big".into(),
438 },
439 );
440 }
441 *value = Symbol::new(&format!("{}", n));
442 }
443 }
444
445 let int_kind = IntKind {
447 size: IntSize::S8,
448 signedness: Signedness::Signed,
449 };
450 let mk_lit = |n: isize| Literal::Int {
451 value: Symbol::new(&format!("{}", n)),
452 negative: false,
453 kind: int_kind.clone(),
454 };
455 let meta = Metadata {
456 span: Span::dummy(),
457 attributes: vec![],
458 };
459 let mk_lit_expr = |n| Expr {
460 kind: Box::new(ExprKind::Literal(mk_lit(n))),
461 ty: Ty(Box::new(TyKind::Primitive(PrimitiveTy::Int(
462 int_kind.clone(),
463 )))),
464 meta: meta.clone(),
465 };
466 let mk_array = |exprs| Expr {
467 kind: Box::new(ExprKind::Array(exprs)),
468 ty: Ty(Box::new(TyKind::RawPointer)), meta: meta.clone(),
470 };
471 let mut lit_expr_200 = mk_lit_expr(200);
472
473 let mut e = mk_array(vec![
475 mk_lit_expr(50),
476 mk_lit_expr(100),
477 lit_expr_200.clone(),
478 ]);
479
480 DoubleU8Literals::default().visit(&mut e);
482
483 lit_expr_200.set_error(ErrorNode {
485 fragment: Box::new(lit_expr_200.clone().into()),
486 diagnostics: vec![Diagnostic::new(
487 mk_lit(200),
488 DiagnosticInfo {
489 span: lit_expr_200.span(),
490 context: Context::Import,
491 kind: DiagnosticInfoKind::AssertionFailure {
492 details: "Literal too big".into(),
493 },
494 },
495 )],
496 });
497
498 assert_eq!(
500 e,
501 mk_array(vec![mk_lit_expr(100), mk_lit_expr(200), lit_expr_200])
502 );
503}