diff --git a/proptest-regressions/domain/lexical_scope/tests/property.txt b/proptest-regressions/domain/lexical_scope/tests/property.txt new file mode 100644 index 00000000..0d621e24 --- /dev/null +++ b/proptest-regressions/domain/lexical_scope/tests/property.txt @@ -0,0 +1,7 @@ +# Seeds for failure cases proptest has generated in the past. It is +# automatically read and these particular cases re-run before any +# novel cases are generated. +# +# It is recommended to check this file in to source control so that +# everyone who runs the test benefits from these saved cases. +cc aa2c89921549b725fd6d19803ed345a5c69e25e36c3e5af9b185cc271dc96384 # shrinks to count = 1 diff --git a/proptest-regressions/domain/package/tests/add_export.txt b/proptest-regressions/domain/package/tests/add_export.txt new file mode 100644 index 00000000..3e6ad627 --- /dev/null +++ b/proptest-regressions/domain/package/tests/add_export.txt @@ -0,0 +1,7 @@ +# Seeds for failure cases proptest has generated in the past. It is +# automatically read and these particular cases re-run before any +# novel cases are generated. +# +# It is recommended to check this file in to source control so that +# everyone who runs the test benefits from these saved cases. +cc c30c07b2fcda86b095d61cd5831bb7cb620baf0207cc55f1d72ad0168cba8a47 # shrinks to package = "a", symbol = "a" diff --git a/proptest-regressions/domain/package/tests/merge_options.txt b/proptest-regressions/domain/package/tests/merge_options.txt new file mode 100644 index 00000000..318014b4 --- /dev/null +++ b/proptest-regressions/domain/package/tests/merge_options.txt @@ -0,0 +1,7 @@ +# Seeds for failure cases proptest has generated in the past. It is +# automatically read and these particular cases re-run before any +# novel cases are generated. +# +# It is recommended to check this file in to source control so that +# everyone who runs the test benefits from these saved cases. +cc fd8f5ca493deed6f17aa556dd701edf1c5e24c596dba9b2dc8d7a811d814c0e7 # shrinks to package = "a", mut symbols = ["a", "b"] diff --git a/proptest-regressions/domain/package/tests/rename.txt b/proptest-regressions/domain/package/tests/rename.txt new file mode 100644 index 00000000..2901c69d --- /dev/null +++ b/proptest-regressions/domain/package/tests/rename.txt @@ -0,0 +1,7 @@ +# Seeds for failure cases proptest has generated in the past. It is +# automatically read and these particular cases re-run before any +# novel cases are generated. +# +# It is recommended to check this file in to source control so that +# everyone who runs the test benefits from these saved cases. +cc 00e6776bf24d8eed932106bbf7fdd15d5d02f3a8bb6a107a6f4751e24a29a432 # shrinks to from = "a", to = "a.a", symbol = "a" diff --git a/proptest-regressions/domain/package/tests/sort_exports.txt b/proptest-regressions/domain/package/tests/sort_exports.txt new file mode 100644 index 00000000..110f07b7 --- /dev/null +++ b/proptest-regressions/domain/package/tests/sort_exports.txt @@ -0,0 +1,7 @@ +# Seeds for failure cases proptest has generated in the past. It is +# automatically read and these particular cases re-run before any +# novel cases are generated. +# +# It is recommended to check this file in to source control so that +# everyone who runs the test benefits from these saved cases. +cc b4e1c869e62416cc01a52c09f5deb90cc8ae69d618a90ad16f85d2f11aed3c77 # shrinks to package = "a", mut symbols = ["a", "aa"] diff --git a/proptest-regressions/domain/package/tests/sort_options.txt b/proptest-regressions/domain/package/tests/sort_options.txt new file mode 100644 index 00000000..c59a6b3b --- /dev/null +++ b/proptest-regressions/domain/package/tests/sort_options.txt @@ -0,0 +1,7 @@ +# Seeds for failure cases proptest has generated in the past. It is +# automatically read and these particular cases re-run before any +# novel cases are generated. +# +# It is recommended to check this file in to source control so that +# everyone who runs the test benefits from these saved cases. +cc 8b75b40d618a71bb30508a771d4569c96ac7b021df9e4c8f0614b0d4ec0b541c # shrinks to package = "a", mut option_indexes = [0, 1] diff --git a/src/application/usecase/conditional_sugar.rs b/src/application/usecase/conditional_sugar.rs index 2639565b..9d7bcea5 100644 --- a/src/application/usecase/conditional_sugar.rs +++ b/src/application/usecase/conditional_sugar.rs @@ -9,7 +9,8 @@ use crate::domain::sexpr::SyntaxTree; pub use domain::{ConditionalConversionPlan, ConditionalConversionRequest}; fn safe(request: &ConditionalConversionRequest<'_>) -> Result<()> { - let tree = SyntaxTree::parse(request.input)?; + domain::require_supported_dialect(request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect)?; Ok(reject_common_lisp_reader_conditionals( &tree, request.dialect, @@ -40,3 +41,58 @@ pub fn plan_convert_if_to_unless( safe(&request)?; domain::plan_convert_if_to_unless(request) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::domain::dialect::Dialect; + + const DIALECTS: [Dialect; 7] = [ + Dialect::CommonLisp, + Dialect::EmacsLisp, + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ]; + + fn request(input: &str, dialect: Dialect) -> ConditionalConversionRequest<'_> { + ConditionalConversionRequest { + input, + dialect, + path: "0".parse().expect("path"), + } + } + + #[test] + fn every_command_gates_all_dialects_before_parsing() { + let support_error = "conditional conversion supports only Common Lisp and Emacs Lisp"; + for dialect in DIALECTS { + let errors = [ + plan_convert_when_to_if(request(")", dialect)).unwrap_err(), + plan_convert_unless_to_if(request(")", dialect)).unwrap_err(), + plan_convert_if_to_when(request(")", dialect)).unwrap_err(), + plan_convert_if_to_unless(request(")", dialect)).unwrap_err(), + ]; + for error in errors { + if matches!(dialect, Dialect::CommonLisp | Dialect::EmacsLisp) { + assert_ne!(error.to_string(), support_error, "{dialect:?}: {error:#}"); + } else { + assert_eq!(error.to_string(), support_error, "{dialect:?}"); + } + } + } + } + + #[test] + fn supported_reader_collisions_use_the_requested_dialect() { + for (dialect, input) in [ + (Dialect::CommonLisp, r"(when ok one two) #\)"), + (Dialect::EmacsLisp, r"(when ok one two) ?\)"), + ] { + let plan = plan_convert_when_to_if(request(input, dialect)).expect("conversion"); + assert!(plan.changed); + } + } +} diff --git a/src/application/usecase/convert_cond_to_if.rs b/src/application/usecase/convert_cond_to_if.rs index d476e212..f88d2ed9 100644 --- a/src/application/usecase/convert_cond_to_if.rs +++ b/src/application/usecase/convert_cond_to_if.rs @@ -9,7 +9,59 @@ use crate::domain::sexpr::SyntaxTree; pub use domain::{ConvertCondToIfPlan, ConvertCondToIfRequest}; pub fn plan_convert_cond_to_if(request: ConvertCondToIfRequest<'_>) -> Result { - let tree = SyntaxTree::parse(request.input)?; + domain::require_supported_dialect(request.dialect, "convert-cond-to-if")?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect)?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; domain::plan_convert_cond_to_if(request) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::domain::dialect::Dialect; + + const DIALECTS: [Dialect; 7] = [ + Dialect::CommonLisp, + Dialect::EmacsLisp, + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ]; + + fn request<'a>(input: &'a str, dialect: Dialect, path: &str) -> ConvertCondToIfRequest<'a> { + ConvertCondToIfRequest { + input, + dialect, + path: path.parse().expect("path"), + } + } + + #[test] + fn all_dialects_are_gated_before_parsing() { + let support_error = "convert-cond-to-if currently supports only Common Lisp and Emacs Lisp"; + for dialect in DIALECTS { + let error = plan_convert_cond_to_if(request(")", dialect, "0")).unwrap_err(); + if matches!(dialect, Dialect::CommonLisp | Dialect::EmacsLisp) { + assert_ne!(error.to_string(), support_error, "{dialect:?}: {error:#}"); + } else { + assert_eq!(error.to_string(), support_error, "{dialect:?}"); + } + } + } + + #[test] + fn supported_reader_collisions_use_the_requested_dialect() { + for (dialect, input) in [ + ( + Dialect::CommonLisp, + r"#\) (cond (ready yes) ((quote t) no))", + ), + (Dialect::EmacsLisp, r"?\) (cond (ready yes) ((quote t) no))"), + ] { + let plan = plan_convert_cond_to_if(request(input, dialect, "1")).expect("conversion"); + assert!(plan.changed); + } + } +} diff --git a/src/application/usecase/convert_flet_to_labels.rs b/src/application/usecase/convert_flet_to_labels.rs index a9a7accf..6a55aada 100644 --- a/src/application/usecase/convert_flet_to_labels.rs +++ b/src/application/usecase/convert_flet_to_labels.rs @@ -3,7 +3,6 @@ use anyhow::Result; use crate::application::usecase::mutation_safety::reject_common_lisp_reader_conditionals; -use crate::domain::dialect::Dialect; use crate::domain::local_function_binding as domain; use crate::domain::sexpr::SyntaxTree; @@ -12,9 +11,54 @@ pub use domain::{ConvertFletToLabelsPlan, ConvertFletToLabelsRequest}; pub fn plan_convert_flet_to_labels( request: ConvertFletToLabelsRequest<'_>, ) -> Result { - let tree = SyntaxTree::parse(request.input)?; - if request.dialect == Dialect::CommonLisp { - reject_common_lisp_reader_conditionals(&tree, request.dialect)?; - } + domain::validate_convert_flet_to_labels_dialect(request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect)?; + reject_common_lisp_reader_conditionals(&tree, request.dialect)?; domain::plan_convert_flet_to_labels(request) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::domain::dialect::Dialect; + + #[test] + fn accepts_common_lisp_reader_literal() { + let input = r"#\) (flet ((helper (value) value)) (helper 1))"; + let plan = plan_convert_flet_to_labels(ConvertFletToLabelsRequest { + input, + dialect: Dialect::CommonLisp, + path: "1".parse().expect("path"), + }) + .expect("plan"); + + assert_eq!( + plan.rewritten, + r"#\) (labels ((helper (value) value)) (helper 1))" + ); + } + + #[test] + fn unsupported_dialect_gate_precedes_parsing() { + for dialect in [ + Dialect::EmacsLisp, + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ] { + let error = plan_convert_flet_to_labels(ConvertFletToLabelsRequest { + input: ")", + dialect, + path: "0".parse().expect("path"), + }) + .expect_err("unsupported dialect"); + + assert_eq!( + error.to_string(), + "convert-flet-to-labels supports only Common Lisp" + ); + } + } +} diff --git a/src/application/usecase/convert_if_to_cond.rs b/src/application/usecase/convert_if_to_cond.rs index 09de9d18..3771b9f6 100644 --- a/src/application/usecase/convert_if_to_cond.rs +++ b/src/application/usecase/convert_if_to_cond.rs @@ -9,7 +9,56 @@ use crate::domain::sexpr::SyntaxTree; pub use domain::{ConvertIfToCondPlan, ConvertIfToCondRequest}; pub fn plan_convert_if_to_cond(request: ConvertIfToCondRequest<'_>) -> Result { - let tree = SyntaxTree::parse(request.input)?; + domain::require_supported_dialect(request.dialect, "convert-if-to-cond")?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect)?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; domain::plan_convert_if_to_cond(request) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::domain::dialect::Dialect; + + const DIALECTS: [Dialect; 7] = [ + Dialect::CommonLisp, + Dialect::EmacsLisp, + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ]; + + fn request<'a>(input: &'a str, dialect: Dialect, path: &str) -> ConvertIfToCondRequest<'a> { + ConvertIfToCondRequest { + input, + dialect, + path: path.parse().expect("path"), + } + } + + #[test] + fn all_dialects_are_gated_before_parsing() { + let support_error = "convert-if-to-cond currently supports only Common Lisp and Emacs Lisp"; + for dialect in DIALECTS { + let error = plan_convert_if_to_cond(request(")", dialect, "0")).unwrap_err(); + if matches!(dialect, Dialect::CommonLisp | Dialect::EmacsLisp) { + assert_ne!(error.to_string(), support_error, "{dialect:?}: {error:#}"); + } else { + assert_eq!(error.to_string(), support_error, "{dialect:?}"); + } + } + } + + #[test] + fn supported_reader_collisions_use_the_requested_dialect() { + for (dialect, input) in [ + (Dialect::CommonLisp, r"#\) (if ready yes no)"), + (Dialect::EmacsLisp, r"?\) (if ready yes no)"), + ] { + let plan = plan_convert_if_to_cond(request(input, dialect, "1")).expect("conversion"); + assert!(plan.changed); + } + } +} diff --git a/src/application/usecase/convert_labels_to_flet.rs b/src/application/usecase/convert_labels_to_flet.rs index 335c2ebf..afd56e8d 100644 --- a/src/application/usecase/convert_labels_to_flet.rs +++ b/src/application/usecase/convert_labels_to_flet.rs @@ -3,7 +3,6 @@ use anyhow::Result; use crate::application::usecase::mutation_safety::reject_common_lisp_reader_conditionals; -use crate::domain::dialect::Dialect; use crate::domain::local_function_binding as domain; use crate::domain::sexpr::SyntaxTree; @@ -12,9 +11,54 @@ pub use domain::{ConvertLabelsToFletPlan, ConvertLabelsToFletRequest}; pub fn plan_convert_labels_to_flet( request: ConvertLabelsToFletRequest<'_>, ) -> Result { - let tree = SyntaxTree::parse(request.input)?; - if request.dialect == Dialect::CommonLisp { - reject_common_lisp_reader_conditionals(&tree, request.dialect)?; - } + domain::validate_convert_labels_to_flet_dialect(request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect)?; + reject_common_lisp_reader_conditionals(&tree, request.dialect)?; domain::plan_convert_labels_to_flet(request) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::domain::dialect::Dialect; + + #[test] + fn accepts_common_lisp_reader_literal() { + let input = r"#\) (labels ((helper (value) value)) (helper 1))"; + let plan = plan_convert_labels_to_flet(ConvertLabelsToFletRequest { + input, + dialect: Dialect::CommonLisp, + path: "1".parse().expect("path"), + }) + .expect("plan"); + + assert_eq!( + plan.rewritten, + r"#\) (flet ((helper (value) value)) (helper 1))" + ); + } + + #[test] + fn unsupported_dialect_gate_precedes_parsing() { + for dialect in [ + Dialect::EmacsLisp, + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ] { + let error = plan_convert_labels_to_flet(ConvertLabelsToFletRequest { + input: ")", + dialect, + path: "0".parse().expect("path"), + }) + .expect_err("unsupported dialect"); + + assert_eq!( + error.to_string(), + "convert-labels-to-flet supports only Common Lisp" + ); + } + } +} diff --git a/src/application/usecase/convert_let_star_to_let.rs b/src/application/usecase/convert_let_star_to_let.rs index 17a23a2e..2408f2e5 100644 --- a/src/application/usecase/convert_let_star_to_let.rs +++ b/src/application/usecase/convert_let_star_to_let.rs @@ -10,7 +10,51 @@ pub use domain::{ConvertLetStarToLetPlan, ConvertLetStarToLetRequest}; pub fn plan_convert_let_star_to_let( request: ConvertLetStarToLetRequest<'_>, ) -> Result { - let tree = SyntaxTree::parse(request.input)?; + domain::validate_convert_let_star_to_let_dialect(request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect)?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; domain::plan_convert_let_star_to_let(request) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::domain::dialect::Dialect; + + #[test] + fn accepts_common_lisp_reader_literal() { + let input = r"#\) (let* ((value 1)) value)"; + let plan = plan_convert_let_star_to_let(ConvertLetStarToLetRequest { + input, + dialect: Dialect::CommonLisp, + path: "1".parse().expect("path"), + }) + .expect("plan"); + + assert_eq!(plan.rewritten, r"#\) (let ((value 1)) value)"); + } + + #[test] + fn unsupported_dialect_gate_precedes_parsing() { + for dialect in [ + Dialect::EmacsLisp, + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ] { + let error = plan_convert_let_star_to_let(ConvertLetStarToLetRequest { + input: ")", + dialect, + path: "0".parse().expect("path"), + }) + .expect_err("unsupported dialect"); + + assert_eq!( + error.to_string(), + "convert-let-star-to-let currently supports only Common Lisp" + ); + } + } +} diff --git a/src/application/usecase/convert_let_to_let_star.rs b/src/application/usecase/convert_let_to_let_star.rs index c29fb4c5..1ff6c04b 100644 --- a/src/application/usecase/convert_let_to_let_star.rs +++ b/src/application/usecase/convert_let_to_let_star.rs @@ -10,7 +10,64 @@ pub use domain::{ConvertLetToLetStarPlan, ConvertLetToLetStarRequest}; pub fn plan_convert_let_to_let_star( request: ConvertLetToLetStarRequest<'_>, ) -> Result { - let tree = SyntaxTree::parse(request.input)?; + domain::validate_convert_let_to_let_star_dialect(request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect)?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; domain::plan_convert_let_to_let_star(request) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::domain::dialect::Dialect; + + #[test] + fn accepts_supported_dialect_reader_literals() { + let cases = [ + ( + Dialect::CommonLisp, + r"#\) (let ((value 1)) value)", + r"#\) (let* ((value 1)) value)", + ), + ( + Dialect::EmacsLisp, + r"?\) (let ((value 1)) value)", + r"?\) (let* ((value 1)) value)", + ), + ]; + + for (dialect, input, expected) in cases { + let plan = plan_convert_let_to_let_star(ConvertLetToLetStarRequest { + input, + dialect, + path: "1".parse().expect("path"), + }) + .unwrap_or_else(|error| panic!("{}: {error}", dialect.label())); + + assert_eq!(plan.rewritten, expected); + } + } + + #[test] + fn unsupported_dialect_gate_precedes_parsing() { + for dialect in [ + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ] { + let error = plan_convert_let_to_let_star(ConvertLetToLetStarRequest { + input: ")", + dialect, + path: "0".parse().expect("path"), + }) + .expect_err("unsupported dialect"); + + assert_eq!( + error.to_string(), + "convert-let-to-let-star supports only Common Lisp and Emacs Lisp" + ); + } + } +} diff --git a/src/application/usecase/convert_sequential_binding.rs b/src/application/usecase/convert_sequential_binding.rs index 441e507a..2f97a22d 100644 --- a/src/application/usecase/convert_sequential_binding.rs +++ b/src/application/usecase/convert_sequential_binding.rs @@ -8,8 +8,9 @@ use crate::domain::sexpr::SyntaxTree; pub use domain::{ConvertSequentialBindingPlan, ConvertSequentialBindingRequest}; -fn safe(request: &ConvertSequentialBindingRequest<'_>) -> Result<()> { - let tree = SyntaxTree::parse(request.input)?; +fn safe(request: &ConvertSequentialBindingRequest<'_>, command: &str) -> Result<()> { + domain::require_supported_dialect(request.dialect, command)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect)?; Ok(reject_common_lisp_reader_conditionals( &tree, request.dialect, @@ -19,12 +20,68 @@ fn safe(request: &ConvertSequentialBindingRequest<'_>) -> Result<()> { pub fn plan_convert_do_star_to_do( request: ConvertSequentialBindingRequest<'_>, ) -> Result { - safe(&request)?; + safe(&request, "convert-do-star-to-do")?; domain::plan_convert_do_star_to_do(request) } pub fn plan_convert_prog_star_to_prog( request: ConvertSequentialBindingRequest<'_>, ) -> Result { - safe(&request)?; + safe(&request, "convert-prog-star-to-prog")?; domain::plan_convert_prog_star_to_prog(request) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::domain::dialect::Dialect; + + const DIALECTS: [Dialect; 7] = [ + Dialect::CommonLisp, + Dialect::EmacsLisp, + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ]; + + fn request(input: &str, dialect: Dialect) -> ConvertSequentialBindingRequest<'_> { + ConvertSequentialBindingRequest { + input, + dialect, + path: "0".parse().expect("path"), + } + } + + #[test] + fn every_command_gates_all_dialects_before_parsing() { + for dialect in DIALECTS { + let cases = [ + ( + plan_convert_do_star_to_do(request(")", dialect)).unwrap_err(), + "convert-do-star-to-do currently supports only Common Lisp", + ), + ( + plan_convert_prog_star_to_prog(request(")", dialect)).unwrap_err(), + "convert-prog-star-to-prog currently supports only Common Lisp", + ), + ]; + for (error, support_error) in cases { + if dialect == Dialect::CommonLisp { + assert_ne!(error.to_string(), support_error, "{dialect:?}: {error:#}"); + } else { + assert_eq!(error.to_string(), support_error, "{dialect:?}"); + } + } + } + } + + #[test] + fn common_lisp_reader_collision_uses_the_requested_dialect() { + let input = + r"(do* ((x (first) (next-x)) (y (second) (next-y))) ((done-p x y) y) (work x)) #\)"; + let plan = + plan_convert_do_star_to_do(request(input, Dialect::CommonLisp)).expect("conversion"); + assert!(plan.changed); + } +} diff --git a/src/application/usecase/eliminate_empty_binding_form.rs b/src/application/usecase/eliminate_empty_binding_form.rs index ae71dcc1..717fb988 100644 --- a/src/application/usecase/eliminate_empty_binding_form.rs +++ b/src/application/usecase/eliminate_empty_binding_form.rs @@ -14,7 +14,9 @@ pub use domain::{EliminateEmptyBindingFormPlan, EliminateEmptyBindingFormRequest pub fn plan_eliminate_empty_binding_form( request: EliminateEmptyBindingFormRequest<'_>, ) -> Result { - let tree = SyntaxTree::parse(request.input).context("input is not valid")?; + domain::require_supported(request.dialect, "eliminate-empty-binding-form")?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("input is not valid")?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; require_known_expression_context(&tree, &request.path, request.dialect)?; domain::plan_eliminate_empty_binding_form(request) @@ -67,3 +69,56 @@ fn require_known_expression_context( bail!("eliminate-empty-binding-form requires a known expression position") } } + +#[cfg(test)] +mod tests { + use super::*; + + const DIALECTS: [Dialect; 7] = [ + Dialect::CommonLisp, + Dialect::EmacsLisp, + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ]; + + fn request<'a>( + input: &'a str, + dialect: Dialect, + path: &str, + ) -> EliminateEmptyBindingFormRequest<'a> { + EliminateEmptyBindingFormRequest { + input, + dialect, + path: path.parse().expect("path"), + } + } + + #[test] + fn all_dialects_are_gated_before_parsing() { + let support_error = "eliminate-empty-binding-form supports only Common Lisp and Emacs Lisp"; + for dialect in DIALECTS { + let error = + plan_eliminate_empty_binding_form(request(")", dialect, "0.1")).unwrap_err(); + if matches!(dialect, Dialect::CommonLisp | Dialect::EmacsLisp) { + assert_ne!(error.to_string(), support_error, "{dialect:?}: {error:#}"); + } else { + assert_eq!(error.to_string(), support_error, "{dialect:?}"); + } + } + } + + #[test] + fn supported_reader_collisions_use_the_requested_dialect() { + for (dialect, input) in [ + (Dialect::CommonLisp, r"(progn (let () a b)) #\)"), + (Dialect::EmacsLisp, r"(progn (let () a b)) ?\)"), + ] { + let plan = plan_eliminate_empty_binding_form(request(input, dialect, "0.1")) + .expect("elimination"); + assert!(plan.changed); + } + } +} diff --git a/src/application/usecase/flatten_progn.rs b/src/application/usecase/flatten_progn.rs index d17b0918..146c12af 100644 --- a/src/application/usecase/flatten_progn.rs +++ b/src/application/usecase/flatten_progn.rs @@ -11,10 +11,11 @@ use crate::domain::sexpr::{Path, SyntaxTree}; pub use domain::{FlattenPrognPlan, FlattenPrognRequest}; pub fn plan_flatten_progn(request: FlattenPrognRequest<'_>) -> Result { + domain::require_supported(request.dialect, "flatten-progn")?; if request.path.indexes().len() < 2 { bail!("flatten-progn refuses to rewrite a top-level progn"); } - let tree = SyntaxTree::parse(request.input)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect)?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; reject_unsafe_context(&tree, &request.path)?; domain::plan_flatten_progn(request) @@ -45,3 +46,51 @@ fn reject_unsafe_context(tree: &SyntaxTree, path: &Path) -> Result<()> { } Ok(()) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::domain::dialect::Dialect; + + const DIALECTS: [Dialect; 7] = [ + Dialect::CommonLisp, + Dialect::EmacsLisp, + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ]; + + fn request<'a>(input: &'a str, dialect: Dialect, path: &str) -> FlattenPrognRequest<'a> { + FlattenPrognRequest { + input, + dialect, + path: path.parse().expect("path"), + } + } + + #[test] + fn all_dialects_are_gated_before_parsing() { + let support_error = "flatten-progn supports only Common Lisp and Emacs Lisp"; + for dialect in DIALECTS { + let error = plan_flatten_progn(request(")", dialect, "0.1")).unwrap_err(); + if matches!(dialect, Dialect::CommonLisp | Dialect::EmacsLisp) { + assert_ne!(error.to_string(), support_error, "{dialect:?}: {error:#}"); + } else { + assert_eq!(error.to_string(), support_error, "{dialect:?}"); + } + } + } + + #[test] + fn supported_reader_collisions_use_the_requested_dialect() { + for (dialect, input) in [ + (Dialect::CommonLisp, r"(progn (progn a (progn b c))) #\)"), + (Dialect::EmacsLisp, r"(progn (progn a (progn b c))) ?\)"), + ] { + let plan = plan_flatten_progn(request(input, dialect, "0.1")).expect("flatten"); + assert!(plan.changed); + } + } +} diff --git a/src/application/usecase/inline_lambda.rs b/src/application/usecase/inline_lambda.rs index 58ec5671..8817f8b6 100644 --- a/src/application/usecase/inline_lambda.rs +++ b/src/application/usecase/inline_lambda.rs @@ -30,7 +30,8 @@ pub struct InlineLambdaPlan { } pub fn plan_inline_lambda(request: InlineLambdaRequest<'_>) -> Result { - let tree = SyntaxTree::parse(request.input)?; + inline_lambda::validate_dialect(request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect)?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; let plan = inline_lambda::plan(DomainRequest { input: request.input, @@ -55,3 +56,33 @@ pub fn plan_inline_lambda(request: InlineLambdaRequest<'_>) -> Result, ) -> Result { - let tree = SyntaxTree::parse(request.input)?; + inline_local_function::validate_dialect(request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect)?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; let plan = inline_local_function::plan(DomainRequest { input: request.input, @@ -63,3 +64,33 @@ pub fn plan_inline_local_function( changed: plan.changed, }) } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn validates_dialect_before_parsing_and_uses_dialect_parser() { + let plan = plan_inline_local_function(InlineLocalFunctionRequest { + input: r"#\) (flet ((identity (x) x)) (identity value))", + dialect: Dialect::CommonLisp, + path: "1".parse().expect("path"), + }) + .expect("Common Lisp"); + SyntaxTree::parse_with_dialect(&plan.rewritten, Dialect::CommonLisp) + .expect("rewritten input"); + + for dialect in [Dialect::EmacsLisp, Dialect::Unknown] { + let error = plan_inline_local_function(InlineLocalFunctionRequest { + input: ")", + dialect, + path: "0".parse().expect("path"), + }) + .expect_err("unsupported dialect"); + assert_eq!( + error.to_string(), + "inline-local-function currently supports only Common Lisp" + ); + } + } +} diff --git a/src/application/usecase/merge_nested_flet.rs b/src/application/usecase/merge_nested_flet.rs index 967f3b66..731dbc24 100644 --- a/src/application/usecase/merge_nested_flet.rs +++ b/src/application/usecase/merge_nested_flet.rs @@ -24,7 +24,8 @@ pub struct MergeNestedFletPlan { } pub fn plan_merge_nested_flet(request: MergeNestedFletRequest<'_>) -> Result { - let tree = SyntaxTree::parse(request.input)?; + flet_composition::validate_dialect(request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect)?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; let plan = flet_composition::plan(DomainRequest { input: request.input, @@ -41,3 +42,33 @@ pub fn plan_merge_nested_flet(request: MergeNestedFletRequest<'_>) -> Result) -> Result { - let tree = SyntaxTree::parse(request.input)?; + let_composition::validate_dialect(request.dialect, "merge-nested-let")?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect)?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; let plan = let_composition::plan_merge_nested_let(DomainRequest { input: request.input, @@ -41,3 +42,35 @@ pub fn plan_merge_nested_let(request: MergeNestedLetRequest<'_>) -> Result, ) -> Result { - let tree = SyntaxTree::parse(request.input)?; + let_composition::validate_dialect(request.dialect, "merge-nested-let-star")?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect)?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; let plan = let_composition::plan_merge_nested_let_star(DomainRequest { input: request.input, @@ -43,3 +44,35 @@ pub fn plan_merge_nested_let_star( changed: plan.changed, }) } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn validates_dialect_before_parsing_and_uses_dialect_parser() { + for (dialect, prefix) in [(Dialect::CommonLisp, r"#\)"), (Dialect::EmacsLisp, r"?\)")] { + let input = format!("{prefix} (let* ((x 1)) (let* ((y (+ x 1))) y))"); + let plan = plan_merge_nested_let_star(MergeNestedLetStarRequest { + input: &input, + dialect, + path: "1".parse().expect("path"), + }) + .expect("supported dialect"); + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect).expect("rewritten input"); + } + + for dialect in [Dialect::Scheme, Dialect::Unknown] { + let error = plan_merge_nested_let_star(MergeNestedLetStarRequest { + input: ")", + dialect, + path: "0".parse().expect("path"), + }) + .expect_err("unsupported dialect"); + assert_eq!( + error.to_string(), + "merge-nested-let-star supports only Common Lisp and Emacs Lisp" + ); + } + } +} diff --git a/src/application/usecase/split_let.rs b/src/application/usecase/split_let.rs index fc66b77a..5ba6ea70 100644 --- a/src/application/usecase/split_let.rs +++ b/src/application/usecase/split_let.rs @@ -27,7 +27,8 @@ pub struct SplitLetPlan { } pub fn plan_split_let(request: SplitLetRequest<'_>) -> Result { - let tree = SyntaxTree::parse(request.input)?; + let_composition::validate_dialect(request.dialect, "split-let")?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect)?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; let binding_index = BindingIndex::new(request.binding_index)?; let plan = let_composition::plan_split_let(DomainRequest { @@ -47,3 +48,37 @@ pub fn plan_split_let(request: SplitLetRequest<'_>) -> Result { changed: plan.changed, }) } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn validates_dialect_before_parsing_and_uses_dialect_parser() { + for (dialect, prefix) in [(Dialect::CommonLisp, r"#\)"), (Dialect::EmacsLisp, r"?\)")] { + let input = format!("{prefix} (let ((x 1) (y 2)) (+ x y))"); + let plan = plan_split_let(SplitLetRequest { + input: &input, + dialect, + path: "1".parse().expect("path"), + binding_index: 1, + }) + .expect("supported dialect"); + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect).expect("rewritten input"); + } + + for dialect in [Dialect::Scheme, Dialect::Unknown] { + let error = plan_split_let(SplitLetRequest { + input: ")", + dialect, + path: "0".parse().expect("path"), + binding_index: 0, + }) + .expect_err("unsupported dialect"); + assert_eq!( + error.to_string(), + "split-let supports only Common Lisp and Emacs Lisp" + ); + } + } +} diff --git a/src/application/usecase/split_let_star.rs b/src/application/usecase/split_let_star.rs index 2379c85e..4b751c3f 100644 --- a/src/application/usecase/split_let_star.rs +++ b/src/application/usecase/split_let_star.rs @@ -26,7 +26,8 @@ pub struct SplitLetStarPlan { pub changed: bool, } pub fn plan_split_let_star(request: SplitLetStarRequest<'_>) -> Result { - let tree = SyntaxTree::parse(request.input)?; + let_star_composition::validate_dialect(request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect)?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; let binding_index = BindingIndex::new(request.binding_index)?; let p = let_star_composition::plan(DomainRequest { @@ -46,3 +47,37 @@ pub fn plan_split_let_star(request: SplitLetStarRequest<'_>) -> Result) -> Result { - let tree = SyntaxTree::parse(request.input)?; + crate::domain::unwrap_call::validate_dialect(request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect)?; reject_overlapping_common_lisp_reader_time_forms( &tree, request.dialect, @@ -65,8 +66,8 @@ mod tests { use crate::domain::sexpr::Path; use proptest::prelude::*; - fn target(input: &str) -> ExpressionView { - let tree = SyntaxTree::parse(input).expect("parse"); + fn target(input: &str, dialect: Dialect) -> ExpressionView { + let tree = SyntaxTree::parse_with_dialect(input, dialect).expect("parse"); tree.select_path(&"0".parse::().expect("path")) .expect("select") .view() @@ -79,7 +80,7 @@ mod tests { input, dialect: Dialect::CommonLisp, path: Some("0".parse().expect("path")), - target: target(input), + target: target(input, Dialect::CommonLisp), expected_function: Some(SymbolName::new("with-cache").expect("symbol")), argument_index: 0, }) @@ -99,7 +100,7 @@ mod tests { input, dialect: Dialect::EmacsLisp, path: Some("0".parse().expect("path")), - target: target(input), + target: target(input, Dialect::EmacsLisp), expected_function: None, argument_index: 1, }) @@ -117,7 +118,7 @@ mod tests { input, dialect: Dialect::CommonLisp, path: Some("0".parse().expect("path")), - target: target(input), + target: target(input, Dialect::CommonLisp), expected_function: Some(SymbolName::new("with-transaction").expect("symbol")), argument_index: 0, }) @@ -126,6 +127,52 @@ mod tests { assert!(err.to_string().contains("expected function")); } + #[test] + fn accepts_reader_forms_for_all_known_dialects() { + let cases = [ + (Dialect::CommonLisp, r"(wrap #\))", r"#\)"), + (Dialect::EmacsLisp, r"(wrap ?\))", r"?\)"), + (Dialect::Scheme, "(wrap #u8(1 2))", "#u8(1 2)"), + ( + Dialect::Clojure, + r#"(wrap #inst "2020-01-01")"#, + r#"#inst "2020-01-01""#, + ), + (Dialect::Janet, "(wrap ;value)", ";value"), + (Dialect::Fennel, "(wrap #(value))", "#(value)"), + ]; + + for (dialect, input, expected_replacement) in cases { + let plan = plan_unwrap_call(UnwrapCallRequest { + input, + dialect, + path: Some("0".parse().expect("path")), + target: target(input, dialect), + expected_function: Some(SymbolName::new("wrap").expect("symbol")), + argument_index: 0, + }) + .unwrap_or_else(|error| panic!("{}: {error}", dialect.label())); + + assert_eq!(plan.replacement, expected_replacement); + assert_eq!(plan.rewritten, expected_replacement); + } + } + + #[test] + fn unknown_dialect_gate_precedes_parsing_and_span_safety() { + let err = plan_unwrap_call(UnwrapCallRequest { + input: ")", + dialect: Dialect::Unknown, + path: Some("0".parse().expect("path")), + target: target("(wrap value)", Dialect::CommonLisp), + expected_function: None, + argument_index: 0, + }) + .expect_err("unknown dialect"); + + assert_eq!(err.to_string(), "unwrap-call requires a known dialect"); + } + fn symbol_strategy() -> impl Strategy { "[a-z][a-z0-9-]{0,8}".prop_filter("reserved symbol", |name| { !matches!( @@ -152,7 +199,7 @@ mod tests { input: &input, dialect: Dialect::Clojure, path: Some("0".parse().expect("path")), - target: target(&input), + target: target(&input, Dialect::Clojure), expected_function: Some(SymbolName::new(&wrapper).expect("symbol")), argument_index: 0, }) @@ -160,7 +207,9 @@ mod tests { prop_assert_eq!(plan.replacement, format!("({callee} {value})")); prop_assert_eq!(&plan.rewritten, &format!("({callee} {value})")); - prop_assert!(SyntaxTree::parse(&plan.rewritten).is_ok()); + prop_assert!( + SyntaxTree::parse_with_dialect(&plan.rewritten, Dialect::Clojure).is_ok() + ); prop_assert!(plan.changed); } } diff --git a/src/domain/common_lisp/reader_condition/mod.rs b/src/domain/common_lisp/reader_condition/mod.rs index e742031d..49a15d64 100644 --- a/src/domain/common_lisp/reader_condition/mod.rs +++ b/src/domain/common_lisp/reader_condition/mod.rs @@ -1,9 +1,9 @@ //! Common Lisp reader-conditional dispatch detection. //! -//! The S-expression parser intentionally represents `#+` and `#-` as atom -//! siblings of their feature expression and guarded datum. This module owns -//! the Common Lisp meaning of those dispatch atoms without changing that -//! general-purpose parser representation. +//! Legacy S-expression trees represent `#+` and `#-` as atom siblings of their +//! feature expression and guarded datum. Dialect-aware Common Lisp trees keep +//! the complete conditional as one opaque atom. This module owns the semantic +//! query across both representations. mod dispatch; mod query; diff --git a/src/domain/common_lisp/reader_condition/query.rs b/src/domain/common_lisp/reader_condition/query.rs index 62a3c2c5..a3283e0f 100644 --- a/src/domain/common_lisp/reader_condition/query.rs +++ b/src/domain/common_lisp/reader_condition/query.rs @@ -1,6 +1,6 @@ #[cfg(test)] use crate::domain::sexpr::ExpressionPath; -use crate::domain::sexpr::{ByteSpan, ExpressionKind, ExpressionView, SyntaxTree}; +use crate::domain::sexpr::{ByteOffset, ByteSpan, ExpressionKind, ExpressionView, SyntaxTree}; #[cfg(test)] use super::CommonLispReaderConditionalDispatch; @@ -8,8 +8,9 @@ use super::{CommonLispReaderConditionalForm, CommonLispReaderConditionalKind}; /// Returns every Common Lisp `#+` or `#-` dispatch atom in source order. /// -/// The parser keeps the dispatch, feature expression, and guarded datum as -/// sibling expressions. A bare dispatch is still reported so callers can +/// Legacy trees keep the dispatch, feature expression, and guarded datum as +/// siblings. Dialect-aware Common Lisp trees keep the complete conditional as +/// one opaque atom. A bare legacy dispatch is still reported so callers can /// reject incomplete input safely before attempting a structural refactor. #[cfg(test)] pub fn common_lisp_reader_conditional_dispatches( @@ -26,9 +27,9 @@ pub fn common_lisp_reader_conditional_dispatches( /// Returns the complete source region consumed by every reader conditional. /// -/// The parser represents a reader conditional as three sibling expressions: -/// its dispatch atom, feature expression, and guarded datum. The returned span -/// protects all three, rather than only the dispatch token. +/// This supports both legacy trees, where the dispatch, feature expression, +/// and guarded datum are siblings, and dialect-aware Common Lisp trees, where +/// the complete conditional is one opaque atom. pub fn common_lisp_reader_conditional_forms( tree: &SyntaxTree, ) -> Vec { @@ -43,11 +44,11 @@ fn collect_dispatches( path: &ExpressionPath, dispatches: &mut Vec, ) { - if let Some(kind) = reader_conditional_kind(view) { + if let Some((kind, span, _)) = reader_conditional(view) { dispatches.push(CommonLispReaderConditionalDispatch { kind, path: path.clone(), - span: view.content_span, + span, }); } @@ -62,16 +63,22 @@ fn collect_forms(view: &ExpressionView, forms: &mut Vec child.span, + ReaderConditionalShape::LegacyDispatch => { + let end = view + .children + .get(index + 2) + .or_else(|| view.children.get(index + 1)) + .map_or(child.span.end(), |component| component.span.end()); + ByteSpan::new(child.span.start(), end) + } + }; forms.push(CommonLispReaderConditionalForm { kind, - dispatch_span: child.content_span, - span: ByteSpan::new(child.span.start(), end), + dispatch_span, + span, }); } @@ -83,17 +90,48 @@ fn collect_forms(view: &ExpressionView, forms: &mut Vec Option { + reader_conditional(view).map(|(kind, _, _)| kind) +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +enum ReaderConditionalShape { + LegacyDispatch, + OpaqueForm, +} + +fn reader_conditional( + view: &ExpressionView, +) -> Option<( + CommonLispReaderConditionalKind, + ByteSpan, + ReaderConditionalShape, +)> { if view.kind != ExpressionKind::Atom { return None; } - match view - .text - .as_deref() - .and_then(|text| text.get(view.symbol_offset..)) - { - Some("#+") => Some(CommonLispReaderConditionalKind::Include), - Some("#-") => Some(CommonLispReaderConditionalKind::Exclude), - _ => None, - } + let text = view.text.as_deref()?.get(view.symbol_offset..)?; + let (kind, shape) = match text { + "#+" => ( + CommonLispReaderConditionalKind::Include, + ReaderConditionalShape::LegacyDispatch, + ), + "#-" => ( + CommonLispReaderConditionalKind::Exclude, + ReaderConditionalShape::LegacyDispatch, + ), + text if text.starts_with("#+") => ( + CommonLispReaderConditionalKind::Include, + ReaderConditionalShape::OpaqueForm, + ), + text if text.starts_with("#-") => ( + CommonLispReaderConditionalKind::Exclude, + ReaderConditionalShape::OpaqueForm, + ), + _ => return None, + }; + let dispatch_start = view.content_span.start(); + let dispatch_end = ByteOffset::new(dispatch_start.get() + 2); + + Some((kind, ByteSpan::new(dispatch_start, dispatch_end), shape)) } diff --git a/src/domain/common_lisp/tests/reader_condition.rs b/src/domain/common_lisp/tests/reader_condition.rs index 89e8ffde..f9725dc5 100644 --- a/src/domain/common_lisp/tests/reader_condition.rs +++ b/src/domain/common_lisp/tests/reader_condition.rs @@ -1,6 +1,7 @@ use super::*; use crate::domain::common_lisp::{ CommonLispReaderConditionalKind, common_lisp_reader_conditional_dispatches, + common_lisp_reader_conditional_forms, }; #[test] @@ -110,3 +111,39 @@ fn does_not_confuse_clojure_conditionals_or_reader_comments_with_common_lisp_dis assert!(common_lisp_reader_conditional_dispatches(&tree).is_empty()); } } + +#[test] +fn collects_complete_forms_from_legacy_and_dialect_aware_trees() { + let input = "#+sbcl (compile-file source) #-(and sbcl x86-64) (load source)"; + let legacy = SyntaxTree::parse(input).expect("legacy parse succeeds"); + let dialect_aware = SyntaxTree::parse_with_dialect(input, Dialect::CommonLisp) + .expect("Common Lisp parse succeeds"); + + for tree in [&legacy, &dialect_aware] { + let forms = common_lisp_reader_conditional_forms(tree); + + assert_eq!(forms.len(), 2); + assert_eq!(forms[0].kind, CommonLispReaderConditionalKind::Include); + assert_eq!(forms[1].kind, CommonLispReaderConditionalKind::Exclude); + assert_eq!(forms[0].dispatch_span.slice(input), "#+"); + assert_eq!(forms[1].dispatch_span.slice(input), "#-"); + assert_eq!(forms[0].span.slice(input), "#+sbcl (compile-file source)"); + assert_eq!( + forms[1].span.slice(input), + "#-(and sbcl x86-64) (load source)" + ); + } +} + +#[test] +fn reports_dispatch_spans_from_dialect_aware_opaque_forms() { + let input = "#+sbcl selected #-sbcl rejected"; + let tree = SyntaxTree::parse_with_dialect(input, Dialect::CommonLisp) + .expect("Common Lisp parse succeeds"); + + let dispatches = common_lisp_reader_conditional_dispatches(&tree); + + assert_eq!(dispatches.len(), 2); + assert_eq!(dispatches[0].span.slice(input), "#+"); + assert_eq!(dispatches[1].span.slice(input), "#-"); +} diff --git a/src/domain/conditional_sugar.rs b/src/domain/conditional_sugar.rs index 28ee7f55..777d1d0f 100644 --- a/src/domain/conditional_sugar.rs +++ b/src/domain/conditional_sugar.rs @@ -32,14 +32,19 @@ fn replace_span(input: &str, span: ByteSpan, replacement: &str) -> String { output } +pub(crate) fn require_supported_dialect(dialect: Dialect) -> Result<()> { + if !matches!(dialect, Dialect::CommonLisp | Dialect::EmacsLisp) { + bail!("conditional conversion supports only Common Lisp and Emacs Lisp"); + } + Ok(()) +} + fn prepare<'a>( request: &ConditionalConversionRequest<'a>, head: &str, ) -> Result<(SyntaxTree, ExpressionView)> { - if !matches!(request.dialect, Dialect::CommonLisp | Dialect::EmacsLisp) { - bail!("conditional conversion supports only Common Lisp and Emacs Lisp"); - } - let tree = SyntaxTree::parse(request.input) + require_supported_dialect(request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) .context("conditional conversion input is not a valid S-expression document")?; let form = tree.select_path(&request.path)?.view(); if tree.has_comment_in(form.span) { @@ -71,7 +76,7 @@ fn finish( replacement: String, ) -> Result { let rewritten = replace_span(request.input, form.span, &replacement); - SyntaxTree::parse(&rewritten) + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) .context("conditional conversion output is not a valid S-expression document")?; Ok(ConditionalConversionPlan { dialect: request.dialect, @@ -181,20 +186,32 @@ mod tests { } #[test] - fn conversions_preserve_parseability_and_dialect_boundary() { - for dialect in [Dialect::CommonLisp, Dialect::EmacsLisp] { - let plan = plan_convert_when_to_if(request("(when ok one two)", dialect)).unwrap(); + fn supported_dialects_preserve_reader_forms_and_validate_with_the_same_dialect() { + for (dialect, input) in [ + (Dialect::CommonLisp, "(when ok one two) #\\)"), + (Dialect::EmacsLisp, "(when ok one two) ?\\)"), + ] { + let plan = plan_convert_when_to_if(request(input, dialect)).unwrap(); assert!(plan.changed); assert_eq!(plan.body_count, 2); - SyntaxTree::parse(&plan.rewritten).unwrap(); + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect).unwrap(); } + } + + #[test] + fn unsupported_dialects_fail_before_parsing_input() { for dialect in [ + Dialect::Scheme, Dialect::Clojure, Dialect::Janet, Dialect::Fennel, - Dialect::Scheme, + Dialect::Unknown, ] { - assert!(plan_convert_when_to_if(request("(when ok one)", dialect)).is_err()); + let error = plan_convert_when_to_if(request(")", dialect)).unwrap_err(); + assert_eq!( + error.to_string(), + "conditional conversion supports only Common Lisp and Emacs Lisp" + ); } } diff --git a/src/domain/convert_control.rs b/src/domain/convert_control.rs index 3b862567..880ff18e 100644 --- a/src/domain/convert_control.rs +++ b/src/domain/convert_control.rs @@ -26,7 +26,7 @@ pub struct ConvertIfToCondPlan { pub fn plan_convert_if_to_cond(request: ConvertIfToCondRequest<'_>) -> Result { require_supported_dialect(request.dialect, "convert-if-to-cond")?; - let tree = SyntaxTree::parse(request.input) + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) .context("convert-if-to-cond input is not a valid S-expression document")?; let form = tree.select_path(&request.path)?.view(); if tree.has_comment_in(form.span) { @@ -47,7 +47,7 @@ pub fn plan_convert_if_to_cond(request: ConvertIfToCondRequest<'_>) -> Result format!("(cond ({test} {then}))"), }; let rewritten = replace_span(request.input, form.span, &replacement); - parse_output(&rewritten, "convert-if-to-cond")?; + parse_output(&rewritten, request.dialect, "convert-if-to-cond")?; Ok(ConvertIfToCondPlan { dialect: request.dialect, @@ -78,7 +78,7 @@ pub struct ConvertCondToIfPlan { pub fn plan_convert_cond_to_if(request: ConvertCondToIfRequest<'_>) -> Result { require_supported_dialect(request.dialect, "convert-cond-to-if")?; - let tree = SyntaxTree::parse(request.input) + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) .context("convert-cond-to-if input is not a valid S-expression document")?; let form = tree.select_path(&request.path)?.view(); if tree.has_comment_in(form.span) { @@ -109,7 +109,7 @@ pub fn plan_convert_cond_to_if(request: ConvertCondToIfRequest<'_>) -> Result) -> Result Result<()> { +pub(crate) fn require_supported_dialect(dialect: Dialect, operation: &str) -> Result<()> { if !matches!(dialect, Dialect::CommonLisp | Dialect::EmacsLisp) { bail!("{operation} currently supports only Common Lisp and Emacs Lisp"); } @@ -161,8 +161,8 @@ fn replace_span(input: &str, span: ByteSpan, replacement: &str) -> String { rewritten } -fn parse_output(rewritten: &str, operation: &str) -> Result<()> { - SyntaxTree::parse(rewritten) +fn parse_output(rewritten: &str, dialect: Dialect, operation: &str) -> Result<()> { + SyntaxTree::parse_with_dialect(rewritten, dialect) .with_context(|| format!("{operation} output is not a valid S-expression document"))?; Ok(()) } @@ -244,4 +244,63 @@ mod tests { .is_err() ); } + + #[test] + fn dialect_support_matrix_is_enforced_before_parsing_and_reparses_output() { + for (dialect, prefix) in [(Dialect::CommonLisp, "#\\)"), (Dialect::EmacsLisp, "?\\)")] { + let if_input = format!("{prefix} (if ready yes no)"); + let if_plan = plan_convert_if_to_cond(ConvertIfToCondRequest { + input: &if_input, + dialect, + path: "1".parse().expect("path"), + }) + .expect("supported if conversion"); + SyntaxTree::parse_with_dialect(&if_plan.rewritten, dialect) + .expect("dialect-specific cond output"); + + let cond_input = format!("{prefix} (cond (ready yes) ((quote t) no))"); + let cond_plan = plan_convert_cond_to_if(ConvertCondToIfRequest { + input: &cond_input, + dialect, + path: "1".parse().expect("path"), + }) + .expect("supported cond conversion"); + SyntaxTree::parse_with_dialect(&cond_plan.rewritten, dialect) + .expect("dialect-specific if output"); + } + + for dialect in [ + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ] { + let if_error = plan_convert_if_to_cond(ConvertIfToCondRequest { + input: ")", + dialect, + path: "0".parse().expect("path"), + }) + .expect_err("unsupported if conversion"); + assert!( + if_error + .to_string() + .contains("currently supports only Common Lisp and Emacs Lisp"), + "{dialect:?}: {if_error:#}" + ); + + let cond_error = plan_convert_cond_to_if(ConvertCondToIfRequest { + input: ")", + dialect, + path: "0".parse().expect("path"), + }) + .expect_err("unsupported cond conversion"); + assert!( + cond_error + .to_string() + .contains("currently supports only Common Lisp and Emacs Lisp"), + "{dialect:?}: {cond_error:#}" + ); + } + } } diff --git a/src/domain/convert_sequential_binding.rs b/src/domain/convert_sequential_binding.rs index aebaac04..9deeb983 100644 --- a/src/domain/convert_sequential_binding.rs +++ b/src/domain/convert_sequential_binding.rs @@ -35,6 +35,13 @@ fn replace_span(input: &str, span: ByteSpan, replacement: &str) -> String { output } +pub(crate) fn require_supported_dialect(dialect: Dialect, command: &str) -> Result<()> { + if dialect != Dialect::CommonLisp { + bail!("{command} currently supports only Common Lisp"); + } + Ok(()) +} + #[derive(Clone, Copy)] enum Conversion { Do, @@ -85,10 +92,8 @@ fn plan_conversion( conversion: Conversion, ) -> Result { let command = conversion.command(); - if request.dialect != Dialect::CommonLisp { - bail!("{command} currently supports only Common Lisp"); - } - let tree = SyntaxTree::parse(request.input) + require_supported_dialect(request.dialect, command)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) .with_context(|| format!("{command} input is not a valid S-expression document"))?; let form = tree.select_path(&request.path)?.view(); if tree.has_comment_in(form.span) { @@ -147,7 +152,7 @@ fn plan_conversion( } let head = &form.children[0]; let rewritten = replace_span(request.input, head.span, conversion.target_head()); - SyntaxTree::parse(&rewritten) + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) .with_context(|| format!("{command} output is not a valid S-expression document"))?; Ok(ConvertSequentialBindingPlan { dialect: request.dialect, @@ -271,15 +276,34 @@ mod tests { } #[test] - fn independent_bindings_parse_after_conversion() { - let input = "(do* ((x (first) (next-x)) (y (second) (next-y))) ((done-p x y) y) (work x))"; + fn independent_bindings_preserve_reader_forms_and_validate_as_common_lisp() { + let input = + "(do* ((x (first) (next-x)) (y (second) (next-y))) ((done-p x y) y) (work x)) #\\)"; let plan = plan_convert_do_star_to_do(request(input, Dialect::CommonLisp)).unwrap(); assert_eq!(plan.binding_names.len(), 2); - SyntaxTree::parse(&plan.rewritten).unwrap(); + SyntaxTree::parse_with_dialect(&plan.rewritten, Dialect::CommonLisp).unwrap(); } #[test] - fn rejects_dependency_duplicates_malformed_forms_and_other_dialects() { + fn unsupported_dialects_fail_before_parsing_input() { + for dialect in [ + Dialect::EmacsLisp, + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ] { + let error = plan_convert_do_star_to_do(request(")", dialect)).unwrap_err(); + assert_eq!( + error.to_string(), + "convert-do-star-to-do currently supports only Common Lisp" + ); + } + } + + #[test] + fn rejects_dependency_duplicates_and_malformed_forms() { assert!( plan_convert_do_star_to_do(request( "(do* ((x 1) (y (+ x 1))) ((done-p)))", @@ -301,9 +325,5 @@ mod tests { )) .is_err() ); - assert!( - plan_convert_do_star_to_do(request("(do* ((x 1)) ((done-p)))", Dialect::Clojure)) - .is_err() - ); } } diff --git a/src/domain/definition/classify.rs b/src/domain/definition/classify.rs index 8e100dc4..c143efe1 100644 --- a/src/domain/definition/classify.rs +++ b/src/domain/definition/classify.rs @@ -4,43 +4,106 @@ use crate::domain::dialect::Dialect; use super::DefinitionCategory; pub(super) fn classify_definition_head(dialect: Dialect, head: &str) -> Option { - if matches!(dialect, Dialect::CommonLisp | Dialect::Unknown) { - if let Some(category) = - CommonLispOperator::from_head(head).and_then(CommonLispOperator::definition_category) - { - return Some(category); - } + if dialect == Dialect::CommonLisp { + return CommonLispOperator::from_head(head) + .and_then(CommonLispOperator::definition_category) + .or_else(|| { + let normalized_lower = + normalize_common_lisp_operator_head(head).to_ascii_lowercase(); + match normalized_lower.as_str() { + "deftest" => Some(DefinitionCategory::Test), + _ if normalized_lower.starts_with("define-") => { + Some(DefinitionCategory::UnknownMacro) + } + _ => None, + } + }); + } + + if dialect != Dialect::Unknown && !dialect.is_definition_head(head) { + return None; } let normalized = normalize_common_lisp_operator_head(head); let normalized_lower = normalized.to_ascii_lowercase(); - let category = match normalized_lower.as_str() { - "cl-defun" | "defsubst" | "definline" | "defn" | "defn-" => DefinitionCategory::Function, - "cl-defmacro" => DefinitionCategory::Macro, - "cl-defgeneric" => DefinitionCategory::GenericFunction, - "cl-defmethod" => DefinitionCategory::Method, - "cl-defclass" => DefinitionCategory::Class, - "cl-defstruct" | "defrecord" => DefinitionCategory::Struct, - "def" | "setq-default" => DefinitionCategory::Variable, - "defconst" => DefinitionCategory::Constant, - "defparameter" | "defcustom" => { - if normalized_lower == "defcustom" { - DefinitionCategory::Customization - } else { - DefinitionCategory::Parameter + let category = match dialect { + Dialect::EmacsLisp => match normalized_lower.as_str() { + "defun" | "defsubst" | "cl-defun" => DefinitionCategory::Function, + "defmacro" | "cl-defmacro" => DefinitionCategory::Macro, + "cl-defgeneric" => DefinitionCategory::GenericFunction, + "cl-defmethod" => DefinitionCategory::Method, + "defvar" => DefinitionCategory::Variable, + "defconst" => DefinitionCategory::Constant, + "defcustom" | "defgroup" => DefinitionCategory::Customization, + "define-minor-mode" | "define-derived-mode" => DefinitionCategory::Mode, + "provide" | "require" => DefinitionCategory::Package, + _ => return None, + }, + Dialect::Scheme => match normalized_lower.as_str() { + "define" | "lambda" => DefinitionCategory::Function, + "define-syntax" => DefinitionCategory::Macro, + "define-library" => DefinitionCategory::Package, + "let" | "let*" => DefinitionCategory::Other, + _ => return None, + }, + Dialect::Clojure => match normalized_lower.as_str() { + "ns" => DefinitionCategory::Package, + "def" => DefinitionCategory::Variable, + "defn" => DefinitionCategory::Function, + "defmacro" => DefinitionCategory::Macro, + "defrecord" => DefinitionCategory::Struct, + "deftype" | "defprotocol" => DefinitionCategory::Class, + "defmulti" => DefinitionCategory::GenericFunction, + "defmethod" => DefinitionCategory::Method, + _ => return None, + }, + Dialect::Janet => match normalized_lower.as_str() { + "def" | "def-" => DefinitionCategory::Variable, + "defn" | "defn-" => DefinitionCategory::Function, + "defmacro" => DefinitionCategory::Macro, + _ => return None, + }, + Dialect::Fennel => match normalized_lower.as_str() { + "fn" | "lambda" => DefinitionCategory::Function, + "macro" => DefinitionCategory::Macro, + "local" | "global" => DefinitionCategory::Variable, + _ => return None, + }, + Dialect::Unknown => { + if let Some(category) = CommonLispOperator::from_head(head) + .and_then(CommonLispOperator::definition_category) + { + return Some(category); + } + + match normalized_lower.as_str() { + "cl-defun" | "defsubst" | "definline" | "defn" | "defn-" => { + DefinitionCategory::Function + } + "cl-defmacro" => DefinitionCategory::Macro, + "cl-defgeneric" => DefinitionCategory::GenericFunction, + "cl-defmethod" => DefinitionCategory::Method, + "cl-defclass" => DefinitionCategory::Class, + "cl-defstruct" | "defrecord" => DefinitionCategory::Struct, + "def" | "setq-default" => DefinitionCategory::Variable, + "defconst" => DefinitionCategory::Constant, + "defparameter" => DefinitionCategory::Parameter, + "defcustom" => DefinitionCategory::Customization, + "deftest" | "define-test" | "ert-deftest" | "define-ert-test" => { + DefinitionCategory::Test + } + "provide" | "require" => DefinitionCategory::Package, + "defgroup" | "defface" => DefinitionCategory::Customization, + "define-minor-mode" | "define-derived-mode" | "define-globalized-minor-mode" => { + DefinitionCategory::Mode + } + _ if dialect.is_definition_head(head) => DefinitionCategory::Other, + _ if normalized_lower.starts_with("define-") => DefinitionCategory::UnknownMacro, + _ => return None, } } - "deftest" | "define-test" | "ert-deftest" | "define-ert-test" => DefinitionCategory::Test, - "provide" | "require" => DefinitionCategory::Package, - "defgroup" | "defface" => DefinitionCategory::Customization, - "define-minor-mode" | "define-derived-mode" | "define-globalized-minor-mode" => { - DefinitionCategory::Mode - } - _ if dialect.is_definition_head(head) => DefinitionCategory::Other, - _ if normalized_lower.starts_with("define-") => DefinitionCategory::UnknownMacro, - _ => return None, + Dialect::CommonLisp => unreachable!("Common Lisp is handled before dialect dispatch"), }; - Some(category) } diff --git a/src/domain/definition/mod.rs b/src/domain/definition/mod.rs index 0fe3c640..e80dc61c 100644 --- a/src/domain/definition/mod.rs +++ b/src/domain/definition/mod.rs @@ -315,38 +315,69 @@ mod tests { #[test] fn classifies_common_lisp_and_emacs_definition_heads() { - assert_eq!( - classify_definition_head(Dialect::CommonLisp, "defun"), - Some(DefinitionCategory::Function) - ); - assert_eq!( - classify_definition_head(Dialect::CommonLisp, "cl:defmacro"), - Some(DefinitionCategory::Macro) - ); - assert_eq!( - classify_definition_head(Dialect::CommonLisp, "cl:defgeneric"), - Some(DefinitionCategory::GenericFunction) - ); - assert_eq!( - classify_definition_head(Dialect::CommonLisp, "define-setf-expander"), - Some(DefinitionCategory::Macro) - ); - assert_eq!( - classify_definition_head(Dialect::CommonLisp, "define-symbol-macro"), - Some(DefinitionCategory::Variable) - ); - assert_eq!( - classify_definition_head(Dialect::CommonLisp, "asdf:defsystem"), - Some(DefinitionCategory::System) - ); - assert_eq!( - classify_definition_head(Dialect::EmacsLisp, "defcustom"), - Some(DefinitionCategory::Customization) - ); - assert_eq!( - classify_definition_head(Dialect::EmacsLisp, "define-minor-mode"), - Some(DefinitionCategory::Mode) - ); + for (head, category) in [ + ("defun", DefinitionCategory::Function), + ("DEFUN", DefinitionCategory::Function), + ("CL:DEFMACRO", DefinitionCategory::Macro), + ("cl:defgeneric", DefinitionCategory::GenericFunction), + ("define-setf-expander", DefinitionCategory::Macro), + ("define-symbol-macro", DefinitionCategory::Variable), + ("asdf:defsystem", DefinitionCategory::System), + ("deftest", DefinitionCategory::Test), + ] { + assert_eq!( + classify_definition_head(Dialect::CommonLisp, head), + Some(category), + "Common Lisp head {head}" + ); + } + + for (head, category) in [ + ("defun", DefinitionCategory::Function), + ("defsubst", DefinitionCategory::Function), + ("cl-defun", DefinitionCategory::Function), + ("defmacro", DefinitionCategory::Macro), + ("cl-defmacro", DefinitionCategory::Macro), + ("cl-defgeneric", DefinitionCategory::GenericFunction), + ("cl-defmethod", DefinitionCategory::Method), + ("defvar", DefinitionCategory::Variable), + ("defconst", DefinitionCategory::Constant), + ("defcustom", DefinitionCategory::Customization), + ("defgroup", DefinitionCategory::Customization), + ("define-minor-mode", DefinitionCategory::Mode), + ("define-derived-mode", DefinitionCategory::Mode), + ("provide", DefinitionCategory::Package), + ("require", DefinitionCategory::Package), + ] { + assert_eq!( + classify_definition_head(Dialect::EmacsLisp, head), + Some(category), + "Emacs Lisp head {head}" + ); + } + } + + #[test] + fn classifies_common_lisp_definition_extensions_without_accepting_arbitrary_heads() { + for (head, expected) in [ + ( + "define-trading-strategy", + Some(DefinitionCategory::UnknownMacro), + ), + ( + "CL:DEFINE-TRADING-STRATEGY", + Some(DefinitionCategory::UnknownMacro), + ), + ("deftest", Some(DefinitionCategory::Test)), + ("CL:DEFTEST", Some(DefinitionCategory::Test)), + ("trading-strategy", None), + ] { + assert_eq!( + classify_definition_head(Dialect::CommonLisp, head), + expected, + "Common Lisp extension head {head}" + ); + } } #[test] @@ -429,15 +460,66 @@ mod tests { } #[test] - fn classifies_clojure_and_custom_define_heads() { - assert_eq!( - classify_definition_head(Dialect::Clojure, "defn-"), - Some(DefinitionCategory::Function) - ); - assert_eq!( - classify_definition_head(Dialect::Clojure, "defrecord"), - Some(DefinitionCategory::Struct) - ); + fn classifies_scheme_clojure_janet_and_fennel_definition_heads() { + for (dialect, head, category) in [ + (Dialect::Scheme, "define", DefinitionCategory::Function), + (Dialect::Scheme, "define-syntax", DefinitionCategory::Macro), + ( + Dialect::Scheme, + "define-library", + DefinitionCategory::Package, + ), + (Dialect::Scheme, "lambda", DefinitionCategory::Function), + (Dialect::Scheme, "let", DefinitionCategory::Other), + (Dialect::Scheme, "let*", DefinitionCategory::Other), + (Dialect::Clojure, "ns", DefinitionCategory::Package), + (Dialect::Clojure, "def", DefinitionCategory::Variable), + (Dialect::Clojure, "defn", DefinitionCategory::Function), + (Dialect::Clojure, "defmacro", DefinitionCategory::Macro), + (Dialect::Clojure, "defrecord", DefinitionCategory::Struct), + (Dialect::Clojure, "deftype", DefinitionCategory::Class), + (Dialect::Clojure, "defprotocol", DefinitionCategory::Class), + ( + Dialect::Clojure, + "defmulti", + DefinitionCategory::GenericFunction, + ), + (Dialect::Clojure, "defmethod", DefinitionCategory::Method), + (Dialect::Janet, "def", DefinitionCategory::Variable), + (Dialect::Janet, "def-", DefinitionCategory::Variable), + (Dialect::Janet, "defn", DefinitionCategory::Function), + (Dialect::Janet, "defn-", DefinitionCategory::Function), + (Dialect::Janet, "defmacro", DefinitionCategory::Macro), + (Dialect::Fennel, "fn", DefinitionCategory::Function), + (Dialect::Fennel, "lambda", DefinitionCategory::Function), + (Dialect::Fennel, "macro", DefinitionCategory::Macro), + (Dialect::Fennel, "local", DefinitionCategory::Variable), + (Dialect::Fennel, "global", DefinitionCategory::Variable), + ] { + assert_eq!( + classify_definition_head(dialect, head), + Some(category), + "{dialect:?} head {head}" + ); + } + } + + #[test] + fn rejects_cross_dialect_definition_heads_but_keeps_unknown_compatibility() { + for (dialect, head) in [ + (Dialect::Scheme, "defn"), + (Dialect::Fennel, "defconst"), + (Dialect::EmacsLisp, "ns"), + (Dialect::Clojure, "defn-"), + (Dialect::Janet, "defrecord"), + ] { + assert_eq!( + classify_definition_head(dialect, head), + None, + "{dialect:?} must reject foreign head {head}" + ); + } + assert_eq!( classify_definition_head(Dialect::Unknown, "define-widget"), Some(DefinitionCategory::Other) diff --git a/src/domain/definition_report.rs b/src/domain/definition_report.rs index ad91d0f2..d60f8b1b 100644 --- a/src/domain/definition_report.rs +++ b/src/domain/definition_report.rs @@ -2,7 +2,7 @@ use std::collections::HashSet; use std::path::PathBuf; use std::thread; -use anyhow::{Result, anyhow}; +use anyhow::{Context, Result, anyhow}; use crate::domain::common_lisp::CommonLispPackageDeclarationForm; use crate::domain::definition::{DefinitionCategory, definition_shape}; @@ -161,14 +161,23 @@ pub fn build_parsed_definition_file( pub fn collect_unused_definition_candidates( files: &[ParsedDefinitionFile], ) -> Result> { - let views: Vec<_> = files + for file in files { + if file.dialect == Dialect::Unknown { + anyhow::bail!( + "unused-definition analysis does not support dialect unknown: {}", + file.path.display() + ); + } + } + + let views: Vec> = files .iter() .map(|file| { - SyntaxTree::parse(&file.text) - .ok() - .map(|tree| tree.root_view()) + SyntaxTree::parse_with_dialect(&file.text, file.dialect) + .with_context(|| format!("failed to parse {}", file.path.display())) + .map(|tree| Some(tree.root_view())) }) - .collect(); + .collect::>()?; let package_form_spans: Vec> = files .iter() @@ -404,6 +413,12 @@ pub fn evaluate_unused_definition_policy( mod tests { use super::*; + fn parsed_file(path: &str, dialect: Dialect, text: &str) -> ParsedDefinitionFile { + let tree = SyntaxTree::parse_with_dialect(text, dialect).expect("fixture must parse"); + build_parsed_definition_file(PathBuf::from(path), dialect, &tree, text) + .expect("fixture report must build") + } + #[test] fn validates_unused_definition_threshold() { assert!(UnusedDefinitionPolicyOptions::new(true, Some(1)).is_ok()); @@ -412,4 +427,121 @@ mod tests { "require-unused-definitions must be greater than zero" ); } + + #[test] + fn collects_unused_definitions_for_every_known_dialect() { + let fixtures = [ + ( + "common-lisp.lisp", + Dialect::CommonLisp, + "(defun cl-unused () 1)\n", + ), + ( + "emacs-lisp.el", + Dialect::EmacsLisp, + "(defun el-unused () 1)\n", + ), + ( + "scheme.scm", + Dialect::Scheme, + "(define scheme-unused (lambda () 1))\n", + ), + ("clojure.clj", Dialect::Clojure, "(defn clj-unused [] 1)\n"), + ("janet.janet", Dialect::Janet, "(defn janet-unused [] 1)\n"), + ("fennel.fnl", Dialect::Fennel, "(fn fennel-unused [] 1)\n"), + ]; + let files: Vec<_> = fixtures + .iter() + .map(|(path, dialect, text)| parsed_file(path, *dialect, text)) + .collect(); + + let reports = collect_unused_definition_candidates(&files).expect("report must build"); + + assert_eq!(reports.len(), fixtures.len()); + assert_eq!(unused_definition_candidate_count(&reports), fixtures.len()); + for ((_, dialect, _), report) in fixtures.iter().zip(&reports) { + assert_eq!(report.dialect, *dialect); + assert_eq!(report.definitions.len(), 1); + assert!(report.definitions[0].references.is_empty()); + } + } + + #[test] + fn rejects_unknown_dialect_before_parsing_any_file() { + let files = vec![ + ParsedDefinitionFile { + path: PathBuf::from("broken.lisp"), + dialect: Dialect::CommonLisp, + package: None, + definitions: Vec::new(), + atoms: Vec::new(), + text: "(defun broken ()".to_owned(), + }, + ParsedDefinitionFile { + path: PathBuf::from("unknown.lisp"), + dialect: Dialect::Unknown, + package: None, + definitions: Vec::new(), + atoms: Vec::new(), + text: "()".to_owned(), + }, + ]; + + let error = collect_unused_definition_candidates(&files).expect_err("input must fail"); + + assert_eq!( + error.to_string(), + "unused-definition analysis does not support dialect unknown: unknown.lisp" + ); + } + + #[test] + fn propagates_malformed_input_errors() { + let files = vec![ParsedDefinitionFile { + path: PathBuf::from("broken.lisp"), + dialect: Dialect::CommonLisp, + package: None, + definitions: Vec::new(), + atoms: Vec::new(), + text: "(defun broken ()".to_owned(), + }]; + + let error = collect_unused_definition_candidates(&files).expect_err("input must fail"); + + assert_eq!(error.to_string(), "failed to parse broken.lisp"); + } + + #[test] + fn parses_reader_syntax_with_the_input_dialect() { + let text = "(defun el-reader () [?\\)])\n"; + assert!(SyntaxTree::parse_with_dialect(text, Dialect::EmacsLisp).is_ok()); + assert!(SyntaxTree::parse_with_dialect(text, Dialect::CommonLisp).is_err()); + let files = vec![parsed_file("reader.el", Dialect::EmacsLisp, text)]; + + let reports = collect_unused_definition_candidates(&files).expect("report must build"); + + assert_eq!(unused_definition_candidate_count(&reports), 1); + } + + #[test] + fn common_lisp_symbol_matching_is_package_aware_but_scheme_is_exact() { + let common_lisp = vec![parsed_file( + "symbols.lisp", + Dialect::CommonLisp, + "(defun Foo () 1)\n(pkg:foo)\n", + )]; + let scheme = vec![parsed_file( + "symbols.scm", + Dialect::Scheme, + "(define Foo (lambda () 1))\n(pkg:foo)\n", + )]; + + let common_lisp_report = + collect_unused_definition_candidates(&common_lisp).expect("report must build"); + let scheme_report = + collect_unused_definition_candidates(&scheme).expect("report must build"); + + assert_eq!(unused_definition_candidate_count(&common_lisp_report), 0); + assert_eq!(unused_definition_candidate_count(&scheme_report), 1); + } } diff --git a/src/domain/dialect/mod.rs b/src/domain/dialect/mod.rs index 715fc187..eea70339 100644 --- a/src/domain/dialect/mod.rs +++ b/src/domain/dialect/mod.rs @@ -3,6 +3,13 @@ mod capability; mod parse; +mod semantic; + +pub use semantic::{ + BinderShape, BindingVisibility, BodyShape, DefinitionShape, ExtractFunctionOperation, + IntroduceLetOperation, ParameterShape, RelativeNodePath, RenameBindingOperation, ScopeShape, + SemanticOperation, UnsupportedSemanticOperation, VerifiedSemanticPolicy, +}; #[cfg(test)] mod tests; diff --git a/src/domain/dialect/semantic.rs b/src/domain/dialect/semantic.rs new file mode 100644 index 00000000..c4473d01 --- /dev/null +++ b/src/domain/dialect/semantic.rs @@ -0,0 +1,1156 @@ +use std::{fmt, marker::PhantomData}; + +use crate::domain::common_lisp::{common_lisp_operator_head_eq, common_lisp_symbol_identity_eq}; +use crate::domain::definition::DefinitionCategory; +use crate::domain::sexpr::{Delimiter, ExpressionKind, ExpressionView}; + +use super::Dialect; + +/// A refactoring operation whose semantic safety must be verified per dialect. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub enum SemanticOperation { + /// Introduces a dialect-appropriate lexical binding form. + IntroduceLet, + /// Renames a lexical binding and its references. + RenameBinding, + /// Extracts selected forms into a new function. + ExtractFunction, +} + +impl SemanticOperation { + /// Returns the stable CLI-facing operation name. + pub const fn label(self) -> &'static str { + match self { + Self::IntroduceLet => "introduce-let", + Self::RenameBinding => "rename-binding", + Self::ExtractFunction => "extract-function", + } + } +} + +mod sealed { + use super::SemanticOperation; + + pub(crate) trait SemanticOperationMarker { + const OPERATION: SemanticOperation; + } +} + +/// Type marker for an introduce-let semantic proof. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub struct IntroduceLetOperation; + +impl sealed::SemanticOperationMarker for IntroduceLetOperation { + const OPERATION: SemanticOperation = SemanticOperation::IntroduceLet; +} + +/// Type marker for a rename-binding semantic proof. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub struct RenameBindingOperation; + +impl sealed::SemanticOperationMarker for RenameBindingOperation { + const OPERATION: SemanticOperation = SemanticOperation::RenameBinding; +} + +/// Type marker for an extract-function semantic proof. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub struct ExtractFunctionOperation; + +impl sealed::SemanticOperationMarker for ExtractFunctionOperation { + const OPERATION: SemanticOperation = SemanticOperation::ExtractFunction; +} + +/// A path from a semantic form to one of its direct or nested children. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub enum RelativeNodePath { + /// A direct child of the form. + Child(usize), + /// A child of one of the form's direct children. + Grandchild { + /// The direct child index. + child: usize, + /// The nested child index. + grandchild: usize, + }, +} + +impl RelativeNodePath { + /// Returns the first child index in the path. + pub const fn child(self) -> usize { + match self { + Self::Child(child) | Self::Grandchild { child, .. } => child, + } + } + + /// Returns the nested child index when this is a two-level path. + pub const fn grandchild(self) -> Option { + match self { + Self::Child(_) => None, + Self::Grandchild { grandchild, .. } => Some(grandchild), + } + } +} + +/// Describes the parameter list of a callable form. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub struct ParameterShape { + container: RelativeNodePath, + first_parameter_index: usize, +} + +impl ParameterShape { + const fn new(container: RelativeNodePath, first_parameter_index: usize) -> Self { + Self { + container, + first_parameter_index, + } + } + + /// Returns the path to the parameter container. + pub const fn container(self) -> RelativeNodePath { + self.container + } + + /// Returns the first child in the container that denotes a parameter. + pub const fn first_parameter_index(self) -> usize { + self.first_parameter_index + } +} + +/// Describes where the executable body of a semantic form begins. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub enum BodyShape { + /// All direct children from this index are body forms. + ChildrenFrom(usize), + /// Body forms begin immediately after the node at this path. + ChildrenAfter(RelativeNodePath), + /// Each callable clause has body forms beginning at the given child index. + ClauseChildrenFrom { + /// Index of the first direct child that is an arity clause. + first_clause_index: usize, + /// Index of the first body form inside each arity clause. + body_child_index: usize, + }, +} + +/// A dialect-neutral definition layout. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct DefinitionShape { + category: DefinitionCategory, + name: Option, + parameters: Option, + body: BodyShape, +} + +impl DefinitionShape { + const fn new( + category: DefinitionCategory, + name: Option, + parameters: Option, + body: BodyShape, + ) -> Self { + Self { + category, + name, + parameters, + body, + } + } + + /// Returns the semantic category of this definition. + pub const fn category(self) -> DefinitionCategory { + self.category + } + + /// Returns the definition name path, if the form has a name. + pub const fn name(self) -> Option { + self.name + } + + /// Returns the callable parameter layout, if present. + pub const fn parameters(self) -> Option { + self.parameters + } + + /// Returns the body layout. + pub const fn body(self) -> BodyShape { + self.body + } +} + +/// Determines whether binding initializers can see earlier bindings. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub enum BindingVisibility { + /// Every initializer is evaluated in the enclosing scope. + Parallel, + /// Each initializer can reference preceding bindings. + Sequential, +} + +/// Describes where a scope obtains its lexical binders. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub enum BinderShape { + /// A container whose children are binding entries such as `(name value)`. + BindingList { + /// Path to the binding-entry container. + container: RelativeNodePath, + /// Path from each binding entry to its name. + name: RelativeNodePath, + /// Path from each binding entry to its initializer. + initializer: Option, + /// Visibility of earlier bindings from later initializers. + visibility: BindingVisibility, + }, + /// A named scope plus a container of binding entries, as in Scheme named let. + NamedBindingList { + /// Path to the name bound over the scope body. + scope_name: RelativeNodePath, + /// Path to the binding-entry container. + container: RelativeNodePath, + /// Path from each binding entry to its name. + name: RelativeNodePath, + /// Path from each binding entry to its initializer. + initializer: Option, + /// Visibility of earlier bindings from later initializers. + visibility: BindingVisibility, + }, + /// Alternating name and initializer nodes in one flat container. + FlatPairs { + /// Path to the flat binding container. + container: RelativeNodePath, + /// Index of the first binding name. + first_name_index: usize, + /// Number of children occupied by each binding pair. + stride: usize, + /// Visibility of earlier bindings from later initializers. + visibility: BindingVisibility, + }, + /// A callable parameter list. + Parameters(ParameterShape), + /// A callable name and parameter list that are both bound over its body. + NamedParameters { + /// Path to the callable's local name. + name: RelativeNodePath, + /// Parameter layout relative to the callable form. + parameters: ParameterShape, + }, + /// Parameter lists repeated in independently scoped callable clauses. + ParameterClauses { + /// Optional path to a callable name bound over every clause body. + name: Option, + /// Index of the first direct child that is an arity clause. + first_clause_index: usize, + /// Parameter layout relative to each arity clause. + parameters: ParameterShape, + }, +} + +/// A dialect-neutral lexical scope layout. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub struct ScopeShape { + binders: BinderShape, + body: BodyShape, +} + +impl ScopeShape { + const fn new(binders: BinderShape, body: BodyShape) -> Self { + Self { binders, body } + } + + /// Returns the lexical binder layout. + pub const fn binders(self) -> BinderShape { + self.binders + } + + /// Returns the executable body layout. + pub const fn body(self) -> BodyShape { + self.body + } +} + +/// Semantic metadata and verification rules used inside the domain layer. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) struct DialectSemanticPolicy { + dialect: Dialect, +} + +impl DialectSemanticPolicy { + pub(crate) const fn new(dialect: Dialect) -> Self { + Self { dialect } + } + + pub(crate) const fn dialect(self) -> Dialect { + self.dialect + } + + pub(crate) const fn supports(self, operation: SemanticOperation) -> bool { + matches!( + (self.dialect, operation), + ( + Dialect::CommonLisp + | Dialect::EmacsLisp + | Dialect::Scheme + | Dialect::Clojure + | Dialect::Janet + | Dialect::Fennel, + SemanticOperation::IntroduceLet + | SemanticOperation::RenameBinding + | SemanticOperation::ExtractFunction, + ) + ) + } + + fn verify( + self, + ) -> Result, UnsupportedSemanticOperation> { + if self.supports(O::OPERATION) { + Ok(VerifiedSemanticPolicy { + policy: self, + operation: PhantomData, + }) + } else { + Err(UnsupportedSemanticOperation { + dialect: self.dialect, + operation: O::OPERATION, + }) + } + } + + pub(crate) fn identifiers_equal(self, candidate: &str, expected: &str) -> bool { + match self.dialect { + Dialect::CommonLisp => common_lisp_symbol_identity_eq(candidate, expected), + Dialect::EmacsLisp + | Dialect::Scheme + | Dialect::Clojure + | Dialect::Janet + | Dialect::Fennel + | Dialect::Unknown => candidate == expected, + } + } + + pub(crate) fn definition_shape(self, form: &ExpressionView) -> Option { + definition_shape(self, form) + } + + pub(crate) fn scope_shape(self, form: &ExpressionView) -> Option { + scope_shape(self, form) + } +} + +impl Dialect { + /// Verifies that introduce-let has semantic support for this dialect. + pub fn verify_introduce_let( + self, + ) -> Result, UnsupportedSemanticOperation> { + DialectSemanticPolicy::new(self).verify() + } + + /// Verifies that rename-binding has semantic support for this dialect. + pub fn verify_rename_binding( + self, + ) -> Result, UnsupportedSemanticOperation> { + DialectSemanticPolicy::new(self).verify() + } + + /// Verifies that extract-function has semantic support for this dialect. + pub fn verify_extract_function( + self, + ) -> Result, UnsupportedSemanticOperation> + { + DialectSemanticPolicy::new(self).verify() + } +} + +/// Proof that semantic operation `O` is verified for a dialect. +/// +/// The operation marker is part of the token type, so a proof for one +/// operation cannot be passed to an API requiring another operation. +/// Raw policy construction is intentionally unavailable outside the crate. +/// +/// ```compile_fail +/// use paredit_cli::dialect::DialectSemanticPolicy; +/// ``` +/// +/// ```compile_fail +/// use paredit_cli::dialect::{ +/// IntroduceLetOperation, RenameBindingOperation, VerifiedSemanticPolicy, +/// }; +/// +/// fn requires_rename(_: Option>) {} +/// let introduce: Option> = None; +/// requires_rename(introduce); +/// ``` +/// +/// Its private fields also prevent safe callers from forging a proof. +/// +/// ```compile_fail +/// use paredit_cli::dialect::{RenameBindingOperation, VerifiedSemanticPolicy}; +/// +/// let _forged = VerifiedSemanticPolicy:: {}; +/// ``` +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct VerifiedSemanticPolicy { + policy: DialectSemanticPolicy, + operation: PhantomData O>, +} + +impl VerifiedSemanticPolicy { + /// Returns the verified dialect. + pub const fn dialect(self) -> Dialect { + self.policy.dialect() + } + + /// Compares identifiers using the verified dialect's identity rules. + pub fn identifiers_equal(self, candidate: &str, expected: &str) -> bool { + self.policy.identifiers_equal(candidate, expected) + } + + /// Resolves a definition layout after validating the actual form. + pub fn definition_shape(self, form: &ExpressionView) -> Option { + self.policy.definition_shape(form) + } + + /// Resolves a lexical scope layout after validating the actual form. + pub fn scope_shape(self, form: &ExpressionView) -> Option { + self.policy.scope_shape(form) + } +} + +impl VerifiedSemanticPolicy { + /// Returns the operation verified by this token type. + pub const fn operation(self) -> SemanticOperation { + SemanticOperation::IntroduceLet + } +} + +impl VerifiedSemanticPolicy { + /// Returns the operation verified by this token type. + pub const fn operation(self) -> SemanticOperation { + SemanticOperation::RenameBinding + } +} + +impl VerifiedSemanticPolicy { + /// Returns the operation verified by this token type. + pub const fn operation(self) -> SemanticOperation { + SemanticOperation::ExtractFunction + } +} + +/// Failure to verify a semantic operation for a dialect. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct UnsupportedSemanticOperation { + dialect: Dialect, + operation: SemanticOperation, +} + +impl UnsupportedSemanticOperation { + /// Returns the unsupported dialect. + pub const fn dialect(self) -> Dialect { + self.dialect + } + + /// Returns the unverified operation. + pub const fn operation(self) -> SemanticOperation { + self.operation + } +} + +impl fmt::Display for UnsupportedSemanticOperation { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + write!( + formatter, + "semantic operation {} is not verified for {:?}", + self.operation.label(), + self.dialect + ) + } +} + +impl std::error::Error for UnsupportedSemanticOperation {} + +const DIRECT_FUNCTION: DefinitionShape = DefinitionShape::new( + DefinitionCategory::Function, + Some(RelativeNodePath::Child(1)), + Some(ParameterShape::new(RelativeNodePath::Child(2), 0)), + BodyShape::ChildrenFrom(3), +); +const DIRECT_MACRO: DefinitionShape = DefinitionShape::new( + DefinitionCategory::Macro, + Some(RelativeNodePath::Child(1)), + Some(ParameterShape::new(RelativeNodePath::Child(2), 0)), + BodyShape::ChildrenFrom(3), +); +const DIRECT_VARIABLE: DefinitionShape = DefinitionShape::new( + DefinitionCategory::Variable, + Some(RelativeNodePath::Child(1)), + None, + BodyShape::ChildrenFrom(2), +); +const SCHEME_FUNCTION_DEFINE: DefinitionShape = DefinitionShape::new( + DefinitionCategory::Function, + Some(RelativeNodePath::Grandchild { + child: 1, + grandchild: 0, + }), + Some(ParameterShape::new(RelativeNodePath::Child(1), 1)), + BodyShape::ChildrenFrom(2), +); +const SCHEME_SYNTAX_DEFINE: DefinitionShape = DefinitionShape::new( + DefinitionCategory::Macro, + Some(RelativeNodePath::Child(1)), + None, + BodyShape::ChildrenFrom(2), +); + +fn definition_shape( + policy: DialectSemanticPolicy, + form: &ExpressionView, +) -> Option { + let head = form_head(form)?; + + match policy.dialect { + Dialect::CommonLisp if common_lisp_operator_head_eq(head, "defun") => { + direct_callable_shape(form, Delimiter::Paren, DIRECT_FUNCTION) + } + Dialect::CommonLisp if common_lisp_operator_head_eq(head, "defmacro") => { + direct_callable_shape(form, Delimiter::Paren, DIRECT_MACRO) + } + Dialect::CommonLisp + if common_lisp_operator_head_eq(head, "defvar") + || common_lisp_operator_head_eq(head, "defparameter") => + { + direct_variable_shape(form) + } + Dialect::EmacsLisp if head == "defun" => { + direct_callable_shape(form, Delimiter::Paren, DIRECT_FUNCTION) + } + Dialect::EmacsLisp if head == "defmacro" => { + direct_callable_shape(form, Delimiter::Paren, DIRECT_MACRO) + } + Dialect::EmacsLisp if matches!(head, "defvar" | "defconst" | "defcustom") => { + direct_variable_shape(form) + } + Dialect::Scheme if head == "define" => scheme_define_shape(form), + Dialect::Scheme if head == "define-syntax" => scheme_define_syntax_shape(form), + Dialect::Clojure if head == "defn" => { + direct_callable_shape(form, Delimiter::Bracket, DIRECT_FUNCTION) + } + Dialect::Clojure if head == "defmacro" => { + direct_callable_shape(form, Delimiter::Bracket, DIRECT_MACRO) + } + Dialect::Clojure if head == "def" => direct_variable_shape(form), + Dialect::Janet if matches!(head, "defn" | "defn-") => { + direct_callable_shape(form, Delimiter::Bracket, DIRECT_FUNCTION) + } + Dialect::Janet if head == "defmacro" => { + direct_callable_shape(form, Delimiter::Bracket, DIRECT_MACRO) + } + Dialect::Janet if matches!(head, "def" | "def-") => direct_variable_shape(form), + Dialect::Fennel if head == "fn" => { + direct_callable_shape(form, Delimiter::Bracket, DIRECT_FUNCTION) + } + Dialect::Fennel if head == "macro" => { + direct_callable_shape(form, Delimiter::Bracket, DIRECT_MACRO) + } + Dialect::Fennel if matches!(head, "local" | "global") => direct_variable_shape(form), + Dialect::Unknown + | Dialect::CommonLisp + | Dialect::EmacsLisp + | Dialect::Scheme + | Dialect::Clojure + | Dialect::Janet + | Dialect::Fennel => None, + } +} + +fn direct_callable_shape( + form: &ExpressionView, + parameter_delimiter: Delimiter, + shape: DefinitionShape, +) -> Option { + (form.children.len() >= 3 + && atom_text(form.children.get(1)?).is_some() + && is_plain_list(form.children.get(2)?, parameter_delimiter)) + .then_some(shape) +} + +fn direct_variable_shape(form: &ExpressionView) -> Option { + (form.children.len() >= 2 && atom_text(form.children.get(1)?).is_some()) + .then_some(DIRECT_VARIABLE) +} + +fn scheme_define_shape(form: &ExpressionView) -> Option { + if form.children.len() < 3 { + return None; + } + + let target = form.children.get(1)?; + if atom_text(target).is_some() { + return Some(DIRECT_VARIABLE); + } + + (is_plain_list(target, Delimiter::Paren) + && target.children.first().and_then(atom_text).is_some()) + .then_some(SCHEME_FUNCTION_DEFINE) +} + +fn scheme_define_syntax_shape(form: &ExpressionView) -> Option { + (form.children.len() == 3 && atom_text(form.children.get(1)?).is_some()) + .then_some(SCHEME_SYNTAX_DEFINE) +} + +const LIST_BINDINGS_PARALLEL: BinderShape = BinderShape::BindingList { + container: RelativeNodePath::Child(1), + name: RelativeNodePath::Child(0), + initializer: Some(RelativeNodePath::Child(1)), + visibility: BindingVisibility::Parallel, +}; +const LIST_BINDINGS_SEQUENTIAL: BinderShape = BinderShape::BindingList { + container: RelativeNodePath::Child(1), + name: RelativeNodePath::Child(0), + initializer: Some(RelativeNodePath::Child(1)), + visibility: BindingVisibility::Sequential, +}; +const FLAT_BINDINGS_SEQUENTIAL: BinderShape = BinderShape::FlatPairs { + container: RelativeNodePath::Child(1), + first_name_index: 0, + stride: 2, + visibility: BindingVisibility::Sequential, +}; +const PARAMETER_SCOPE: ScopeShape = ScopeShape::new( + BinderShape::Parameters(ParameterShape::new(RelativeNodePath::Child(1), 0)), + BodyShape::ChildrenFrom(2), +); +const LIST_LET_SCOPE: ScopeShape = + ScopeShape::new(LIST_BINDINGS_PARALLEL, BodyShape::ChildrenFrom(2)); +const LIST_LET_STAR_SCOPE: ScopeShape = + ScopeShape::new(LIST_BINDINGS_SEQUENTIAL, BodyShape::ChildrenFrom(2)); +const FLAT_LET_SCOPE: ScopeShape = + ScopeShape::new(FLAT_BINDINGS_SEQUENTIAL, BodyShape::ChildrenFrom(2)); +const SCHEME_NAMED_LET_SCOPE: ScopeShape = ScopeShape::new( + BinderShape::NamedBindingList { + scope_name: RelativeNodePath::Child(1), + container: RelativeNodePath::Child(2), + name: RelativeNodePath::Child(0), + initializer: Some(RelativeNodePath::Child(1)), + visibility: BindingVisibility::Parallel, + }, + BodyShape::ChildrenFrom(3), +); + +fn scope_shape(policy: DialectSemanticPolicy, form: &ExpressionView) -> Option { + let head = form_head(form)?; + + match policy.dialect { + Dialect::CommonLisp if common_lisp_operator_head_eq(head, "let") => { + list_scope(form, Delimiter::Paren, LIST_LET_SCOPE) + } + Dialect::CommonLisp if common_lisp_operator_head_eq(head, "let*") => { + list_scope(form, Delimiter::Paren, LIST_LET_STAR_SCOPE) + } + Dialect::CommonLisp if common_lisp_operator_head_eq(head, "lambda") => { + parameter_scope(form, Delimiter::Paren, PARAMETER_SCOPE) + } + Dialect::EmacsLisp if head == "let" => list_scope(form, Delimiter::Paren, LIST_LET_SCOPE), + Dialect::EmacsLisp if head == "let*" => { + list_scope(form, Delimiter::Paren, LIST_LET_STAR_SCOPE) + } + Dialect::EmacsLisp if head == "lambda" => { + parameter_scope(form, Delimiter::Paren, PARAMETER_SCOPE) + } + Dialect::Scheme if head == "let" => scheme_let_scope(form), + Dialect::Scheme if head == "let*" => { + list_scope(form, Delimiter::Paren, LIST_LET_STAR_SCOPE) + } + Dialect::Scheme if head == "lambda" => { + parameter_scope(form, Delimiter::Paren, PARAMETER_SCOPE) + } + Dialect::Clojure if head == "let" => flat_scope(form, FLAT_LET_SCOPE), + Dialect::Clojure if head == "fn" => clojure_fn_scope(form), + Dialect::Janet if head == "let" => flat_scope(form, FLAT_LET_SCOPE), + Dialect::Janet if head == "fn" => { + parameter_scope(form, Delimiter::Bracket, PARAMETER_SCOPE) + } + Dialect::Fennel if head == "let" => flat_scope(form, FLAT_LET_SCOPE), + Dialect::Fennel if head == "fn" => { + parameter_scope(form, Delimiter::Bracket, PARAMETER_SCOPE) + } + Dialect::Unknown + | Dialect::CommonLisp + | Dialect::EmacsLisp + | Dialect::Scheme + | Dialect::Clojure + | Dialect::Janet + | Dialect::Fennel => None, + } +} + +fn list_scope( + form: &ExpressionView, + binding_delimiter: Delimiter, + shape: ScopeShape, +) -> Option { + (form.children.len() >= 3 && is_plain_list(form.children.get(1)?, binding_delimiter)) + .then_some(shape) +} + +fn flat_scope(form: &ExpressionView, shape: ScopeShape) -> Option { + let bindings = form.children.get(1)?; + (form.children.len() >= 3 + && is_plain_list(bindings, Delimiter::Bracket) + && bindings.children.len() % 2 == 0) + .then_some(shape) +} + +fn parameter_scope( + form: &ExpressionView, + parameter_delimiter: Delimiter, + shape: ScopeShape, +) -> Option { + (form.children.len() >= 3 && is_plain_list(form.children.get(1)?, parameter_delimiter)) + .then_some(shape) +} + +fn scheme_let_scope(form: &ExpressionView) -> Option { + if form.children.len() >= 3 + && form + .children + .get(1) + .is_some_and(|bindings| is_plain_list(bindings, Delimiter::Paren)) + { + return Some(LIST_LET_SCOPE); + } + + (form.children.len() >= 4 + && form.children.get(1).and_then(atom_text).is_some() + && form + .children + .get(2) + .is_some_and(|bindings| is_plain_list(bindings, Delimiter::Paren))) + .then_some(SCHEME_NAMED_LET_SCOPE) +} + +fn clojure_fn_scope(form: &ExpressionView) -> Option { + let first = form.children.get(1)?; + let (name, first_shape_index) = if atom_text(first).is_some() { + (Some(RelativeNodePath::Child(1)), 2) + } else { + (None, 1) + }; + let first_shape = form.children.get(first_shape_index)?; + + if is_plain_list(first_shape, Delimiter::Bracket) { + if form.children.len() <= first_shape_index + 1 { + return None; + } + + let parameters = ParameterShape::new(RelativeNodePath::Child(first_shape_index), 0); + let binders = name.map_or(BinderShape::Parameters(parameters), |name| { + BinderShape::NamedParameters { name, parameters } + }); + return Some(ScopeShape::new( + binders, + BodyShape::ChildrenFrom(first_shape_index + 1), + )); + } + + let clauses = &form.children[first_shape_index..]; + if clauses.is_empty() || !clauses.iter().all(valid_clojure_arity_clause) { + return None; + } + + Some(ScopeShape::new( + BinderShape::ParameterClauses { + name, + first_clause_index: first_shape_index, + parameters: ParameterShape::new(RelativeNodePath::Child(0), 0), + }, + BodyShape::ClauseChildrenFrom { + first_clause_index: first_shape_index, + body_child_index: 1, + }, + )) +} + +fn valid_clojure_arity_clause(clause: &ExpressionView) -> bool { + is_plain_list(clause, Delimiter::Paren) + && clause.children.len() >= 2 + && clause + .children + .first() + .is_some_and(|parameters| is_plain_list(parameters, Delimiter::Bracket)) +} + +fn form_head(form: &ExpressionView) -> Option<&str> { + if !is_plain_list(form, Delimiter::Paren) { + return None; + } + form.children.first().and_then(atom_text) +} + +fn atom_text(view: &ExpressionView) -> Option<&str> { + (view.kind == ExpressionKind::Atom && view.reader_prefixes.is_empty()) + .then_some(view.text.as_deref()) + .flatten() +} + +fn is_plain_list(view: &ExpressionView, delimiter: Delimiter) -> bool { + view.kind == ExpressionKind::List + && view.delimiter == Some(delimiter) + && view.reader_prefixes.is_empty() +} + +#[cfg(test)] +mod tests { + use crate::domain::sexpr::SyntaxTree; + + use super::*; + + const OPERATIONS: [SemanticOperation; 3] = [ + SemanticOperation::IntroduceLet, + SemanticOperation::RenameBinding, + SemanticOperation::ExtractFunction, + ]; + + fn parsed_form(source: &str, dialect: Dialect) -> ExpressionView { + let root = SyntaxTree::parse_with_dialect(source, dialect) + .expect("fixture parses") + .root_view(); + + root.children + .first() + .cloned() + .expect("fixture has one form") + } + + fn verified_dialect( + dialect: Dialect, + operation: SemanticOperation, + ) -> Result { + match operation { + SemanticOperation::IntroduceLet => dialect + .verify_introduce_let() + .map(VerifiedSemanticPolicy::dialect), + SemanticOperation::RenameBinding => dialect + .verify_rename_binding() + .map(VerifiedSemanticPolicy::dialect), + SemanticOperation::ExtractFunction => dialect + .verify_extract_function() + .map(VerifiedSemanticPolicy::dialect), + } + } + + #[test] + fn semantic_support_matrix_covers_all_eighteen_dialect_operation_cells() { + let cases = [ + (Dialect::CommonLisp, true), + (Dialect::EmacsLisp, true), + (Dialect::Scheme, true), + (Dialect::Clojure, true), + (Dialect::Janet, true), + (Dialect::Fennel, true), + ]; + let mut checked_cells = 0; + + for (dialect, supported) in cases { + let policy = DialectSemanticPolicy::new(dialect); + for operation in OPERATIONS { + assert_eq!(policy.supports(operation), supported, "{dialect:?}"); + assert_eq!( + verified_dialect(dialect, operation).ok(), + supported.then_some(dialect), + "{dialect:?}: {operation:?}" + ); + checked_cells += 1; + } + } + + assert_eq!(checked_cells, 18); + } + + #[test] + fn unknown_dialect_fails_closed_for_every_verification_entry() { + let policy = DialectSemanticPolicy::new(Dialect::Unknown); + + for operation in OPERATIONS { + assert!(!policy.supports(operation)); + let error = verified_dialect(Dialect::Unknown, operation) + .expect_err("Unknown must fail every operation-specific factory"); + assert_eq!(error.dialect(), Dialect::Unknown); + assert_eq!(error.operation(), operation); + } + } + + #[test] + fn verified_token_type_is_bound_to_its_operation() { + fn accepts_rename(_: VerifiedSemanticPolicy) {} + + let verified = Dialect::CommonLisp + .verify_rename_binding() + .expect("Common Lisp rename-binding is verified"); + + accepts_rename(verified); + assert_eq!(verified.dialect(), Dialect::CommonLisp); + assert_eq!(verified.operation(), SemanticOperation::RenameBinding); + } + + #[test] + fn common_lisp_identifier_equality_is_package_aware_and_conservative() { + let policy = DialectSemanticPolicy::new(Dialect::CommonLisp); + + assert!(policy.identifiers_equal(":X", ":x")); + assert!(policy.identifiers_equal("A:X", "a::x")); + assert!(policy.identifiers_equal("CL:X", "COMMON-LISP:x")); + assert!(policy.identifiers_equal("A:|X|", "a:x")); + + assert!(!policy.identifiers_equal("A:X", "B:X")); + assert!(!policy.identifiers_equal("A:X", "X")); + assert!(!policy.identifiers_equal("X", "A:X")); + assert!(!policy.identifiers_equal("#:X", "X")); + assert!(!policy.identifiers_equal("#:X", "#:X")); + assert!(!policy.identifiers_equal("#:X", "#:x")); + assert!(!policy.identifiers_equal("A:|x|", "A:X")); + assert!(!policy.identifiers_equal("|a|:X", "A:X")); + } + + #[test] + fn non_common_lisp_identifier_equality_is_exact() { + for dialect in [ + Dialect::EmacsLisp, + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ] { + let policy = DialectSemanticPolicy::new(dialect); + assert!(policy.identifiers_equal("same", "same"), "{dialect:?}"); + assert!(!policy.identifiers_equal("Widget", "widget"), "{dialect:?}"); + } + } + + #[test] + fn definition_shape_matrix_covers_all_six_dialects() { + let cases = [ + ( + Dialect::CommonLisp, + "(defun f (x) x)", + DefinitionCategory::Function, + ), + ( + Dialect::EmacsLisp, + "(defun f (x) x)", + DefinitionCategory::Function, + ), + ( + Dialect::Scheme, + "(define (f x) x)", + DefinitionCategory::Function, + ), + ( + Dialect::Clojure, + "(defn f [x] x)", + DefinitionCategory::Function, + ), + ( + Dialect::Janet, + "(defn f [x] x)", + DefinitionCategory::Function, + ), + ( + Dialect::Fennel, + "(macro m [x] x)", + DefinitionCategory::Macro, + ), + ]; + + for (dialect, source, category) in cases { + let form = parsed_form(source, dialect); + let shape = DialectSemanticPolicy::new(dialect) + .definition_shape(&form) + .expect("known definition form"); + assert_eq!(shape.category(), category, "{dialect:?}"); + } + } + + #[test] + fn scheme_definition_resolver_discriminates_actual_form_shape() { + let policy = DialectSemanticPolicy::new(Dialect::Scheme); + let variable = parsed_form("(define answer 42)", Dialect::Scheme); + let function = parsed_form("(define (answer x) x)", Dialect::Scheme); + let syntax = parsed_form("(define-syntax when transformer)", Dialect::Scheme); + + assert_eq!( + policy + .definition_shape(&variable) + .map(DefinitionShape::category), + Some(DefinitionCategory::Variable) + ); + let function_shape = policy + .definition_shape(&function) + .expect("function define shape"); + assert_eq!(function_shape.category(), DefinitionCategory::Function); + assert_eq!( + function_shape.name(), + Some(RelativeNodePath::Grandchild { + child: 1, + grandchild: 0, + }) + ); + assert_eq!( + function_shape.parameters(), + Some(ParameterShape::new(RelativeNodePath::Child(1), 1)) + ); + + let syntax_shape = policy + .definition_shape(&syntax) + .expect("define-syntax shape"); + assert_eq!(syntax_shape.category(), DefinitionCategory::Macro); + assert_eq!(syntax_shape.name(), Some(RelativeNodePath::Child(1))); + assert_eq!(syntax_shape.parameters(), None); + assert_eq!(syntax_shape.body(), BodyShape::ChildrenFrom(2)); + } + + #[test] + fn definition_resolver_rejects_unverified_shapes() { + let cases = [ + (Dialect::Scheme, "(define)"), + (Dialect::Scheme, "(define (f))"), + (Dialect::Scheme, "(define-syntax x)"), + (Dialect::Scheme, "(define-syntax (x) transformer)"), + (Dialect::Clojure, "(defn f (not-a-parameter-vector) body)"), + (Dialect::Unknown, "(defun f (x) x)"), + ]; + + for (dialect, source) in cases { + let form = parsed_form(source, dialect); + assert_eq!( + DialectSemanticPolicy::new(dialect).definition_shape(&form), + None, + "{dialect:?}: {source}" + ); + } + } + + #[test] + fn scope_shape_matrix_covers_all_six_dialects() { + let cases = [ + ( + Dialect::CommonLisp, + "(let ((x 1)) x)", + LIST_BINDINGS_PARALLEL, + ), + ( + Dialect::EmacsLisp, + "(let ((x 1)) x)", + LIST_BINDINGS_PARALLEL, + ), + (Dialect::Scheme, "(let ((x 1)) x)", LIST_BINDINGS_PARALLEL), + (Dialect::Clojure, "(let [x 1] x)", FLAT_BINDINGS_SEQUENTIAL), + (Dialect::Janet, "(let [x 1] x)", FLAT_BINDINGS_SEQUENTIAL), + (Dialect::Fennel, "(let [x 1] x)", FLAT_BINDINGS_SEQUENTIAL), + ]; + + for (dialect, source, binders) in cases { + let form = parsed_form(source, dialect); + let shape = DialectSemanticPolicy::new(dialect) + .scope_shape(&form) + .expect("known let scope"); + assert_eq!(shape.binders(), binders, "{dialect:?}"); + assert_eq!(shape.body(), BodyShape::ChildrenFrom(2), "{dialect:?}"); + } + } + + #[test] + fn scheme_named_let_uses_shifted_binding_and_body_paths() { + let form = parsed_form("(let loop ((x 1)) (loop x))", Dialect::Scheme); + let shape = DialectSemanticPolicy::new(Dialect::Scheme) + .scope_shape(&form) + .expect("named let scope"); + + assert_eq!(shape, SCHEME_NAMED_LET_SCOPE); + assert_eq!(shape.body(), BodyShape::ChildrenFrom(3)); + + let malformed = parsed_form("(let loop body)", Dialect::Scheme); + assert_eq!( + DialectSemanticPolicy::new(Dialect::Scheme).scope_shape(&malformed), + None + ); + } + + #[test] + fn clojure_fn_resolver_handles_optional_name_and_multi_arity() { + let policy = DialectSemanticPolicy::new(Dialect::Clojure); + let anonymous = parsed_form("(fn [x] x)", Dialect::Clojure); + let named = parsed_form("(fn add [x] x)", Dialect::Clojure); + let multi = parsed_form("(fn ([x] x) ([x y] y))", Dialect::Clojure); + let named_multi = parsed_form("(fn add ([x] x) ([x y] y))", Dialect::Clojure); + + assert_eq!(policy.scope_shape(&anonymous), Some(PARAMETER_SCOPE)); + assert_eq!( + policy.scope_shape(&named), + Some(ScopeShape::new( + BinderShape::NamedParameters { + name: RelativeNodePath::Child(1), + parameters: ParameterShape::new(RelativeNodePath::Child(2), 0), + }, + BodyShape::ChildrenFrom(3), + )) + ); + assert_eq!( + policy.scope_shape(&multi), + Some(clojure_multi_arity_scope(None, 1)) + ); + assert_eq!( + policy.scope_shape(&named_multi), + Some(clojure_multi_arity_scope( + Some(RelativeNodePath::Child(1)), + 2, + )) + ); + } + + #[test] + fn clojure_fn_resolver_fails_closed_on_unverified_shapes() { + let policy = DialectSemanticPolicy::new(Dialect::Clojure); + + for source in [ + "(fn add)", + "(fn [x])", + "(fn ([x]))", + "(fn add (x))", + "(fn ([x] x) malformed)", + ] { + let form = parsed_form(source, Dialect::Clojure); + assert_eq!(policy.scope_shape(&form), None, "{source}"); + } + } + + fn clojure_multi_arity_scope( + name: Option, + first_clause_index: usize, + ) -> ScopeShape { + ScopeShape::new( + BinderShape::ParameterClauses { + name, + first_clause_index, + parameters: ParameterShape::new(RelativeNodePath::Child(0), 0), + }, + BodyShape::ClauseChildrenFrom { + first_clause_index, + body_child_index: 1, + }, + ) + } + + #[test] + fn unknown_dialect_has_no_semantic_shapes() { + let policy = DialectSemanticPolicy::new(Dialect::Unknown); + let definition = parsed_form("(defun f (x) x)", Dialect::Unknown); + let scope = parsed_form("(let ((x 1)) x)", Dialect::Unknown); + + assert_eq!(policy.definition_shape(&definition), None); + assert_eq!(policy.scope_shape(&scope), None); + } +} diff --git a/src/domain/extract_constant.rs b/src/domain/extract_constant.rs index 680e83bf..45e73d42 100644 --- a/src/domain/extract_constant.rs +++ b/src/domain/extract_constant.rs @@ -53,7 +53,7 @@ pub(crate) fn plan(request: Request<'_>) -> Result { let replacement = request.name.as_str().to_owned(); let definition = format!("({head} {} {selected})", request.name); let replaced = replace_span_checked(request.input, span, &replacement)?; - let replaced_tree = SyntaxTree::parse(&replaced) + let replaced_tree = SyntaxTree::parse_with_dialect(&replaced, request.dialect) .context("replacement output is not a valid S-expression document")?; let (rewritten, anchor_span) = insert_top_level_form( &replaced, @@ -63,7 +63,7 @@ pub(crate) fn plan(request: Request<'_>) -> Result { request.anchor_path.as_ref(), "extract-constant", )?; - SyntaxTree::parse(&rewritten) + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) .context("extracted output is not a valid S-expression document")?; Ok(Plan { @@ -162,3 +162,78 @@ fn find_path(view: &ExpressionView, target: ByteSpan, path: &mut Vec) -> } None } + +#[cfg(test)] +mod tests { + use std::cell::Cell; + + use super::*; + + fn extraction_plan(input: &str, dialect: Dialect) -> Result { + let tree = SyntaxTree::parse_with_dialect(input, dialect)?; + let path: Path = "0.1".parse()?; + let selection = tree.select_path(&path)?; + plan(Request { + input, + tree: &tree, + selection, + path, + dialect, + name: SymbolName::new("answer")?, + insert: TopLevelInsert::Append, + anchor_path: None, + }) + } + + #[test] + fn dialect_matrix_gates_unsupported_dialects_before_parsing() { + assert_eq!(dialect_head(Dialect::CommonLisp).unwrap(), "defconstant"); + assert_eq!(dialect_head(Dialect::EmacsLisp).unwrap(), "defconst"); + + for dialect in [ + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ] { + let parse_attempted = Cell::new(false); + let error = dialect_head(dialect) + .and_then(|_| { + parse_attempted.set(true); + SyntaxTree::parse_with_dialect(")", dialect)?; + Ok("") + }) + .expect_err("dialect must be rejected"); + assert!(!parse_attempted.get()); + assert_eq!( + error.to_string(), + "extract-constant supports only common-lisp and emacs-lisp" + ); + } + } + + #[test] + fn preserves_reader_character_literals_in_each_supported_dialect() { + for (dialect, input, definition) in [ + ( + Dialect::CommonLisp, + r"(print #\))", + r"(defconstant answer #\))", + ), + ( + Dialect::EmacsLisp, + r"(message ?\))", + r"(defconst answer ?\))", + ), + ] { + let plan = extraction_plan(input, dialect).expect("extraction plan"); + assert_eq!(plan.definition, definition); + assert!(plan.rewritten.contains(definition)); + SyntaxTree::parse_with_dialect(&plan.definition, dialect) + .expect("generated definition must use the request dialect"); + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect) + .expect("rewritten output must use the request dialect"); + } + } +} diff --git a/src/domain/extract_function.rs b/src/domain/extract_function.rs index fff1b829..a75bcead 100644 --- a/src/domain/extract_function.rs +++ b/src/domain/extract_function.rs @@ -18,8 +18,12 @@ use rewrite::{extracted_call, extracted_definition}; pub use types::{ExtractFunctionInsert, ExtractFunctionPlan, ExtractFunctionRequest}; pub fn plan_extract_function(request: ExtractFunctionRequest<'_>) -> Result { + let semantic = request + .dialect + .verify_extract_function() + .context("extract-function is not supported for this dialect")?; request.selection.validate_source(request.input)?; - let input_tree = SyntaxTree::parse(request.input) + let input_tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) .context("extract-function input is not a valid S-expression document")?; reject_common_lisp_reader_conditionals(&input_tree, request.dialect)?; @@ -27,14 +31,15 @@ pub fn plan_extract_function(request: ExtractFunctionRequest<'_>) -> Result) -> Result) -> Result Vec { - inference::infer_extract_function_params(dialect, selection, explicit_params) + let Ok(semantic) = dialect.verify_extract_function() else { + return Vec::new(); + }; + inference::infer_extract_function_params(semantic, selection, explicit_params) } diff --git a/src/domain/extract_function/inference.rs b/src/domain/extract_function/inference.rs index 20e163bc..0f071383 100644 --- a/src/domain/extract_function/inference.rs +++ b/src/domain/extract_function/inference.rs @@ -1,25 +1,29 @@ mod bindings; mod forms; mod patterns; +mod semantic; mod symbols; -use crate::domain::common_lisp::common_lisp_symbol_reference_eq; -use crate::domain::common_lisp::is_common_lisp_declaration_form; -use crate::domain::dialect::Dialect; +use crate::domain::common_lisp::{ + common_lisp_symbol_reference_eq, is_common_lisp_declaration_form, +}; +use crate::domain::dialect::{Dialect, ExtractFunctionOperation, VerifiedSemanticPolicy}; use crate::domain::sexpr::{Delimiter, ExpressionKind, ExpressionView, ReaderPrefix}; use super::syntax::{atom_text, list_head}; use forms::collect_inferred_extract_function_special_form; use symbols::is_extract_function_param_candidate; +pub(super) type ExtractFunctionSemantic = VerifiedSemanticPolicy; + pub(super) fn infer_extract_function_params( - dialect: Dialect, + semantic: ExtractFunctionSemantic, selection: &ExpressionView, explicit_params: &[String], ) -> Vec { let mut params = Vec::new(); collect_inferred_extract_function_params( - dialect, + semantic, selection, false, explicit_params, @@ -29,15 +33,19 @@ pub(super) fn infer_extract_function_params( params } -pub(super) fn extract_function_param_name_eq(dialect: Dialect, left: &str, right: &str) -> bool { - match dialect { - Dialect::CommonLisp | Dialect::Unknown => common_lisp_symbol_reference_eq(left, right), - _ => left == right, +pub(super) fn extract_function_param_name_eq( + semantic: ExtractFunctionSemantic, + left: &str, + right: &str, +) -> bool { + match semantic.dialect() { + Dialect::CommonLisp => common_lisp_symbol_reference_eq(left, right), + _ => semantic.identifiers_equal(left, right), } } fn collect_inferred_extract_function_params( - dialect: Dialect, + semantic: ExtractFunctionSemantic, view: &ExpressionView, is_call_head: bool, explicit_params: &[String], @@ -53,20 +61,23 @@ fn collect_inferred_extract_function_params( && is_extract_function_param_candidate(text) && !explicit_params .iter() - .any(|param| extract_function_param_name_eq(dialect, param, text)) + .any(|param| extract_function_param_name_eq(semantic, param, text)) && !bound_params .iter() - .any(|param| extract_function_param_name_eq(dialect, param, text)) + .any(|param| extract_function_param_name_eq(semantic, param, text)) && !params .iter() - .any(|param| extract_function_param_name_eq(dialect, param, text)) + .any(|param| extract_function_param_name_eq(semantic, param, text)) { params.push(text.to_owned()); } return; } - if view.kind == ExpressionKind::List && view.delimiter == Some(Delimiter::Paren) { + if semantic.dialect() == Dialect::CommonLisp + && view.kind == ExpressionKind::List + && view.delimiter == Some(Delimiter::Paren) + { if let Some(head) = list_head(view) { if is_common_lisp_declaration_form(head) { return; @@ -74,8 +85,18 @@ fn collect_inferred_extract_function_params( } } + if semantic::collect_inferred_extract_function_semantic_form( + semantic, + view, + explicit_params, + bound_params, + params, + ) { + return; + } + if collect_inferred_extract_function_special_form( - dialect, + semantic, view, explicit_params, bound_params, @@ -86,7 +107,7 @@ fn collect_inferred_extract_function_params( for (index, child) in view.children.iter().enumerate() { collect_inferred_extract_function_params( - dialect, + semantic, child, view.kind == ExpressionKind::List && view.delimiter == Some(Delimiter::Paren) diff --git a/src/domain/extract_function/inference/bindings.rs b/src/domain/extract_function/inference/bindings.rs index cca8b7e4..eff4744e 100644 --- a/src/domain/extract_function/inference/bindings.rs +++ b/src/domain/extract_function/inference/bindings.rs @@ -9,6 +9,7 @@ pub(super) struct ExtractFunctionBindingEntry { } pub(super) fn extract_function_binding_entries( + semantic: super::ExtractFunctionSemantic, binding_form: &ExpressionView, ) -> Option> { match binding_form.delimiter { @@ -21,7 +22,7 @@ pub(super) fn extract_function_binding_entries( .children .chunks_exact(2) .map(|pair| ExtractFunctionBindingEntry { - names: extract_function_pattern_names(&pair[0]), + names: extract_function_pattern_names(semantic, &pair[0]), value: Some(pair[1].clone()), }) .collect(), @@ -38,7 +39,7 @@ pub(super) fn extract_function_binding_entries( return None; } return Some(ExtractFunctionBindingEntry { - names: extract_function_pattern_names(pair), + names: extract_function_pattern_names(semantic, pair), value: None, }); } @@ -46,7 +47,7 @@ pub(super) fn extract_function_binding_entries( return None; } Some(ExtractFunctionBindingEntry { - names: extract_function_pattern_names(&pair.children[0]), + names: extract_function_pattern_names(semantic, &pair.children[0]), value: pair.children.get(1).cloned(), }) }) diff --git a/src/domain/extract_function/inference/forms/bindings.rs b/src/domain/extract_function/inference/forms/bindings.rs index e7cfcf20..95b75d7b 100644 --- a/src/domain/extract_function/inference/forms/bindings.rs +++ b/src/domain/extract_function/inference/forms/bindings.rs @@ -1,16 +1,15 @@ use std::iter; use crate::domain::common_lisp::CommonLispResourceBindingForm; -use crate::domain::dialect::Dialect; use crate::domain::sexpr::{Delimiter, ExpressionKind, ExpressionView}; use super::super::super::syntax::atom_text; use super::super::bindings::extract_function_binding_entries; use super::super::patterns::parameter_names; -use super::{extend_extract_function_bound_params, slot_spec_bound_name}; +use super::{ExtractFunctionSemantic, extend_extract_function_bound_params, slot_spec_bound_name}; pub(super) fn collect_inferred_extract_function_let( - dialect: Dialect, + semantic: ExtractFunctionSemantic, view: &ExpressionView, explicit_params: &[String], bound_params: &[String], @@ -21,21 +20,21 @@ pub(super) fn collect_inferred_extract_function_let( }; if binding_form.delimiter == Some(Delimiter::Bracket) { return collect_inferred_extract_function_let_star( - dialect, + semantic, view, explicit_params, bound_params, params, ); } - let Some(bindings) = extract_function_binding_entries(binding_form) else { + let Some(bindings) = extract_function_binding_entries(semantic, binding_form) else { return false; }; for binding in &bindings { if let Some(value) = &binding.value { super::super::collect_inferred_extract_function_params( - dialect, + semantic, value, false, explicit_params, @@ -46,14 +45,14 @@ pub(super) fn collect_inferred_extract_function_let( } let body_bound_params = extend_extract_function_bound_params( - dialect, + semantic, bound_params, bindings .iter() .flat_map(|binding| binding.names.iter().map(String::as_str)), ); collect_bodies( - dialect, + semantic, &view.children[2..], explicit_params, &body_bound_params, @@ -63,7 +62,7 @@ pub(super) fn collect_inferred_extract_function_let( } pub(super) fn collect_inferred_extract_function_let_star( - dialect: Dialect, + semantic: ExtractFunctionSemantic, view: &ExpressionView, explicit_params: &[String], bound_params: &[String], @@ -72,7 +71,7 @@ pub(super) fn collect_inferred_extract_function_let_star( let Some(binding_form) = view.children.get(1) else { return false; }; - let Some(bindings) = extract_function_binding_entries(binding_form) else { + let Some(bindings) = extract_function_binding_entries(semantic, binding_form) else { return false; }; @@ -80,7 +79,7 @@ pub(super) fn collect_inferred_extract_function_let_star( for binding in &bindings { if let Some(value) = &binding.value { super::super::collect_inferred_extract_function_params( - dialect, + semantic, value, false, explicit_params, @@ -89,12 +88,12 @@ pub(super) fn collect_inferred_extract_function_let_star( ); } for name in &binding.names { - super::push_extract_function_bound_param(dialect, &mut current_bound_params, name); + super::push_extract_function_bound_param(semantic, &mut current_bound_params, name); } } collect_bodies( - dialect, + semantic, &view.children[2..], explicit_params, ¤t_bound_params, @@ -104,7 +103,7 @@ pub(super) fn collect_inferred_extract_function_let_star( } pub(super) fn collect_inferred_extract_function_value_binding( - dialect: Dialect, + semantic: ExtractFunctionSemantic, view: &ExpressionView, explicit_params: &[String], bound_params: &[String], @@ -118,7 +117,7 @@ pub(super) fn collect_inferred_extract_function_value_binding( }; super::super::collect_inferred_extract_function_params( - dialect, + semantic, value_form, false, explicit_params, @@ -126,14 +125,14 @@ pub(super) fn collect_inferred_extract_function_value_binding( params, ); - let names = parameter_names(binding_form); + let names = parameter_names(semantic, binding_form); let body_bound_params = extend_extract_function_bound_params( - dialect, + semantic, bound_params, names.iter().map(String::as_str), ); collect_bodies( - dialect, + semantic, &view.children[3..], explicit_params, &body_bound_params, @@ -143,7 +142,7 @@ pub(super) fn collect_inferred_extract_function_value_binding( } pub(super) fn collect_inferred_extract_function_clause_form( - dialect: Dialect, + semantic: ExtractFunctionSemantic, view: &ExpressionView, explicit_params: &[String], bound_params: &[String], @@ -154,7 +153,7 @@ pub(super) fn collect_inferred_extract_function_clause_form( }; super::super::collect_inferred_extract_function_params( - dialect, + semantic, protected_form, false, explicit_params, @@ -165,7 +164,7 @@ pub(super) fn collect_inferred_extract_function_clause_form( for clause in &view.children[2..] { if clause.kind != ExpressionKind::List || clause.delimiter != Some(Delimiter::Paren) { super::super::collect_inferred_extract_function_params( - dialect, + semantic, clause, false, explicit_params, @@ -178,14 +177,14 @@ pub(super) fn collect_inferred_extract_function_clause_form( let Some(parameter_form) = clause.children.get(1) else { continue; }; - let names = parameter_names(parameter_form); + let names = parameter_names(semantic, parameter_form); let clause_bound_params = extend_extract_function_bound_params( - dialect, + semantic, bound_params, names.iter().map(String::as_str), ); collect_bodies( - dialect, + semantic, &clause.children[2..], explicit_params, &clause_bound_params, @@ -196,7 +195,7 @@ pub(super) fn collect_inferred_extract_function_clause_form( } pub(super) fn collect_inferred_extract_function_handler_bind( - dialect: Dialect, + semantic: ExtractFunctionSemantic, view: &ExpressionView, explicit_params: &[String], bound_params: &[String], @@ -214,7 +213,7 @@ pub(super) fn collect_inferred_extract_function_handler_bind( if let Some(function_form) = spec.children.get(1) { super::super::collect_inferred_extract_function_params( - dialect, + semantic, function_form, false, explicit_params, @@ -225,7 +224,7 @@ pub(super) fn collect_inferred_extract_function_handler_bind( if include_restart_options { collect_inferred_extract_function_restart_option_values( - dialect, + semantic, spec, explicit_params, bound_params, @@ -235,7 +234,7 @@ pub(super) fn collect_inferred_extract_function_handler_bind( } collect_bodies( - dialect, + semantic, &view.children[2..], explicit_params, bound_params, @@ -245,7 +244,7 @@ pub(super) fn collect_inferred_extract_function_handler_bind( } fn collect_inferred_extract_function_restart_option_values( - dialect: Dialect, + semantic: ExtractFunctionSemantic, spec: &ExpressionView, explicit_params: &[String], bound_params: &[String], @@ -254,7 +253,7 @@ fn collect_inferred_extract_function_restart_option_values( let mut index = 2; while index + 1 < spec.children.len() { super::super::collect_inferred_extract_function_params( - dialect, + semantic, &spec.children[index + 1], false, explicit_params, @@ -266,7 +265,7 @@ fn collect_inferred_extract_function_restart_option_values( } pub(super) fn collect_inferred_extract_function_iteration_binding( - dialect: Dialect, + semantic: ExtractFunctionSemantic, view: &ExpressionView, explicit_params: &[String], bound_params: &[String], @@ -278,7 +277,7 @@ pub(super) fn collect_inferred_extract_function_iteration_binding( if let Some(source_form) = binding_form.children.get(1) { super::super::collect_inferred_extract_function_params( - dialect, + semantic, source_form, false, explicit_params, @@ -288,7 +287,7 @@ pub(super) fn collect_inferred_extract_function_iteration_binding( } let body_bound_params = extend_extract_function_bound_params( - dialect, + semantic, bound_params, binding_form .children @@ -299,7 +298,7 @@ pub(super) fn collect_inferred_extract_function_iteration_binding( if let Some(result_form) = binding_form.children.get(2) { super::super::collect_inferred_extract_function_params( - dialect, + semantic, result_form, false, explicit_params, @@ -309,7 +308,7 @@ pub(super) fn collect_inferred_extract_function_iteration_binding( } collect_bodies( - dialect, + semantic, &view.children[2..], explicit_params, &body_bound_params, @@ -319,7 +318,7 @@ pub(super) fn collect_inferred_extract_function_iteration_binding( } pub(super) fn collect_inferred_extract_function_slot_binding( - dialect: Dialect, + semantic: ExtractFunctionSemantic, view: &ExpressionView, explicit_params: &[String], bound_params: &[String], @@ -333,7 +332,7 @@ pub(super) fn collect_inferred_extract_function_slot_binding( }; super::super::collect_inferred_extract_function_params( - dialect, + semantic, instance_form, false, explicit_params, @@ -342,12 +341,12 @@ pub(super) fn collect_inferred_extract_function_slot_binding( ); let body_bound_params = extend_extract_function_bound_params( - dialect, + semantic, bound_params, slot_specs.children.iter().filter_map(slot_spec_bound_name), ); collect_bodies( - dialect, + semantic, &view.children[3..], explicit_params, &body_bound_params, @@ -357,7 +356,7 @@ pub(super) fn collect_inferred_extract_function_slot_binding( } fn collect_bodies( - dialect: Dialect, + semantic: ExtractFunctionSemantic, bodies: &[ExpressionView], explicit_params: &[String], bound_params: &[String], @@ -365,7 +364,7 @@ fn collect_bodies( ) { for body in bodies { super::super::collect_inferred_extract_function_params( - dialect, + semantic, body, false, explicit_params, @@ -376,7 +375,7 @@ fn collect_bodies( } pub(super) fn collect_inferred_extract_function_resource_binding( - dialect: Dialect, + semantic: ExtractFunctionSemantic, view: &ExpressionView, explicit_params: &[String], bound_params: &[String], @@ -392,7 +391,7 @@ pub(super) fn collect_inferred_extract_function_resource_binding( for initializer in binding_spec.children.iter().skip(1) { super::super::collect_inferred_extract_function_params( - dialect, + semantic, initializer, false, explicit_params, @@ -402,10 +401,10 @@ pub(super) fn collect_inferred_extract_function_resource_binding( } let body_bound_params = - extend_extract_function_bound_params(dialect, bound_params, iter::once(binding_name)); + extend_extract_function_bound_params(semantic, bound_params, iter::once(binding_name)); for body_form in view.children.iter().skip(resource_form.body_start_index()) { super::super::collect_inferred_extract_function_params( - dialect, + semantic, body_form, false, explicit_params, diff --git a/src/domain/extract_function/inference/forms/callable.rs b/src/domain/extract_function/inference/forms/callable.rs index f5d87a07..048cf4fe 100644 --- a/src/domain/extract_function/inference/forms/callable.rs +++ b/src/domain/extract_function/inference/forms/callable.rs @@ -1,11 +1,10 @@ -use crate::domain::dialect::Dialect; use crate::domain::sexpr::{Delimiter, ExpressionKind, ExpressionView}; use super::super::patterns::parameter_names; -use super::extend_extract_function_bound_params; +use super::{ExtractFunctionSemantic, extend_extract_function_bound_params}; pub(super) fn collect_inferred_extract_function_lambda( - dialect: Dialect, + semantic: ExtractFunctionSemantic, view: &ExpressionView, parameter_index: usize, explicit_params: &[String], @@ -15,15 +14,15 @@ pub(super) fn collect_inferred_extract_function_lambda( let Some(parameter_form) = view.children.get(parameter_index) else { return false; }; - let names = parameter_names(parameter_form); + let names = parameter_names(semantic, parameter_form); let body_bound_params = extend_extract_function_bound_params( - dialect, + semantic, bound_params, names.iter().map(String::as_str), ); for body in &view.children[parameter_index + 1..] { super::super::collect_inferred_extract_function_params( - dialect, + semantic, body, false, explicit_params, @@ -35,6 +34,7 @@ pub(super) fn collect_inferred_extract_function_lambda( } pub(super) fn collect_inferred_extract_function_local_callable_form( + semantic: ExtractFunctionSemantic, view: &ExpressionView, explicit_params: &[String], bound_params: &[String], @@ -55,13 +55,15 @@ pub(super) fn collect_inferred_extract_function_local_callable_form( continue; }; let lambda_bound_params = extend_extract_function_bound_params( - Dialect::CommonLisp, + semantic, bound_params, - parameter_names(parameter_form).iter().map(String::as_str), + parameter_names(semantic, parameter_form) + .iter() + .map(String::as_str), ); for body in binding.children.iter().skip(2) { super::super::collect_inferred_extract_function_params( - Dialect::CommonLisp, + semantic, body, false, explicit_params, @@ -73,7 +75,7 @@ pub(super) fn collect_inferred_extract_function_local_callable_form( for body in &view.children[2..] { super::super::collect_inferred_extract_function_params( - Dialect::CommonLisp, + semantic, body, false, explicit_params, diff --git a/src/domain/extract_function/inference/forms/control.rs b/src/domain/extract_function/inference/forms/control.rs index 1219ca1c..86c437b9 100644 --- a/src/domain/extract_function/inference/forms/control.rs +++ b/src/domain/extract_function/inference/forms/control.rs @@ -1,13 +1,12 @@ -use crate::domain::dialect::Dialect; use crate::domain::sexpr::ExpressionView; use super::{ - extend_extract_function_bound_params, iteration_spec_bound_name, iteration_spec_init_form, - iteration_spec_step_form, push_extract_function_bound_param, + ExtractFunctionSemantic, extend_extract_function_bound_params, iteration_spec_bound_name, + iteration_spec_init_form, iteration_spec_step_form, push_extract_function_bound_param, }; pub(super) fn collect_inferred_extract_function_do( - dialect: Dialect, + semantic: ExtractFunctionSemantic, view: &ExpressionView, explicit_params: &[String], bound_params: &[String], @@ -19,15 +18,27 @@ pub(super) fn collect_inferred_extract_function_do( }; let body_bound_params = if sequential_scope { - collect_sequential_do_inits(dialect, binding_form, explicit_params, bound_params, params) + collect_sequential_do_inits( + semantic, + binding_form, + explicit_params, + bound_params, + params, + ) } else { - collect_parallel_do_inits(dialect, binding_form, explicit_params, bound_params, params) + collect_parallel_do_inits( + semantic, + binding_form, + explicit_params, + bound_params, + params, + ) }; for spec in &binding_form.children { if let Some(step_form) = iteration_spec_step_form(spec) { super::super::collect_inferred_extract_function_params( - dialect, + semantic, step_form, false, explicit_params, @@ -39,7 +50,7 @@ pub(super) fn collect_inferred_extract_function_do( for body in &view.children[2..] { super::super::collect_inferred_extract_function_params( - dialect, + semantic, body, false, explicit_params, @@ -51,7 +62,7 @@ pub(super) fn collect_inferred_extract_function_do( } pub(super) fn collect_inferred_extract_function_prog( - dialect: Dialect, + semantic: ExtractFunctionSemantic, view: &ExpressionView, explicit_params: &[String], bound_params: &[String], @@ -63,14 +74,26 @@ pub(super) fn collect_inferred_extract_function_prog( }; let body_bound_params = if sequential_scope { - collect_sequential_do_inits(dialect, binding_form, explicit_params, bound_params, params) + collect_sequential_do_inits( + semantic, + binding_form, + explicit_params, + bound_params, + params, + ) } else { - collect_parallel_do_inits(dialect, binding_form, explicit_params, bound_params, params) + collect_parallel_do_inits( + semantic, + binding_form, + explicit_params, + bound_params, + params, + ) }; for body in &view.children[2..] { super::super::collect_inferred_extract_function_params( - dialect, + semantic, body, false, explicit_params, @@ -82,7 +105,7 @@ pub(super) fn collect_inferred_extract_function_prog( } fn collect_sequential_do_inits( - dialect: Dialect, + semantic: ExtractFunctionSemantic, binding_form: &ExpressionView, explicit_params: &[String], bound_params: &[String], @@ -92,7 +115,7 @@ fn collect_sequential_do_inits( for spec in &binding_form.children { if let Some(init_form) = iteration_spec_init_form(spec) { super::super::collect_inferred_extract_function_params( - dialect, + semantic, init_form, false, explicit_params, @@ -101,14 +124,14 @@ fn collect_sequential_do_inits( ); } if let Some(name) = iteration_spec_bound_name(spec) { - push_extract_function_bound_param(dialect, &mut body_bound_params, name); + push_extract_function_bound_param(semantic, &mut body_bound_params, name); } } body_bound_params } fn collect_parallel_do_inits( - dialect: Dialect, + semantic: ExtractFunctionSemantic, binding_form: &ExpressionView, explicit_params: &[String], bound_params: &[String], @@ -117,7 +140,7 @@ fn collect_parallel_do_inits( for spec in &binding_form.children { if let Some(init_form) = iteration_spec_init_form(spec) { super::super::collect_inferred_extract_function_params( - dialect, + semantic, init_form, false, explicit_params, @@ -127,7 +150,7 @@ fn collect_parallel_do_inits( } } extend_extract_function_bound_params( - dialect, + semantic, bound_params, binding_form .children diff --git a/src/domain/extract_function/inference/forms/mod.rs b/src/domain/extract_function/inference/forms/mod.rs index 009e8529..8c17443c 100644 --- a/src/domain/extract_function/inference/forms/mod.rs +++ b/src/domain/extract_function/inference/forms/mod.rs @@ -3,18 +3,19 @@ mod callable; mod control; use crate::domain::common_lisp::{CommonLispLetBindingForm, CommonLispValueScopeForm}; -use crate::domain::dialect::Dialect; use crate::domain::sexpr::{Delimiter, ExpressionKind, ExpressionView}; use super::super::syntax::{atom_text, list_head}; +use super::ExtractFunctionSemantic; pub(super) fn collect_inferred_extract_function_special_form( - dialect: Dialect, + semantic: ExtractFunctionSemantic, view: &ExpressionView, explicit_params: &[String], bound_params: &[String], params: &mut Vec, ) -> bool { + let dialect = semantic.dialect(); if view.kind != ExpressionKind::List || view.delimiter != Some(Delimiter::Paren) { return false; } @@ -31,7 +32,7 @@ pub(super) fn collect_inferred_extract_function_special_form( CommonLispValueScopeForm::Let(CommonLispLetBindingForm::Parallel) | CommonLispValueScopeForm::Let(CommonLispLetBindingForm::SymbolMacro) => { bindings::collect_inferred_extract_function_let( - dialect, + semantic, view, explicit_params, bound_params, @@ -40,7 +41,7 @@ pub(super) fn collect_inferred_extract_function_special_form( } CommonLispValueScopeForm::Let(CommonLispLetBindingForm::Sequential) => { bindings::collect_inferred_extract_function_let_star( - dialect, + semantic, view, explicit_params, bound_params, @@ -49,7 +50,7 @@ pub(super) fn collect_inferred_extract_function_special_form( } CommonLispValueScopeForm::Value => { bindings::collect_inferred_extract_function_value_binding( - dialect, + semantic, view, explicit_params, bound_params, @@ -58,7 +59,7 @@ pub(super) fn collect_inferred_extract_function_special_form( } CommonLispValueScopeForm::Clause => { bindings::collect_inferred_extract_function_clause_form( - dialect, + semantic, view, explicit_params, bound_params, @@ -67,7 +68,7 @@ pub(super) fn collect_inferred_extract_function_special_form( } CommonLispValueScopeForm::Handler(form) => { bindings::collect_inferred_extract_function_handler_bind( - dialect, + semantic, view, explicit_params, bound_params, @@ -77,7 +78,7 @@ pub(super) fn collect_inferred_extract_function_special_form( } CommonLispValueScopeForm::Iteration => { bindings::collect_inferred_extract_function_iteration_binding( - dialect, + semantic, view, explicit_params, bound_params, @@ -88,7 +89,7 @@ pub(super) fn collect_inferred_extract_function_special_form( if dialect.common_lisp_variable_binding_has_step_forms_for_head(head) => { control::collect_inferred_extract_function_do( - dialect, + semantic, view, explicit_params, bound_params, @@ -98,7 +99,7 @@ pub(super) fn collect_inferred_extract_function_special_form( } CommonLispValueScopeForm::Variable(form) => { control::collect_inferred_extract_function_prog( - dialect, + semantic, view, explicit_params, bound_params, @@ -107,7 +108,7 @@ pub(super) fn collect_inferred_extract_function_special_form( ) } CommonLispValueScopeForm::Slot => bindings::collect_inferred_extract_function_slot_binding( - dialect, + semantic, view, explicit_params, bound_params, @@ -115,7 +116,7 @@ pub(super) fn collect_inferred_extract_function_special_form( ), CommonLispValueScopeForm::Resource(resource_form) => { bindings::collect_inferred_extract_function_resource_binding( - dialect, + semantic, view, explicit_params, bound_params, @@ -125,7 +126,7 @@ pub(super) fn collect_inferred_extract_function_special_form( } CommonLispValueScopeForm::Lambda | CommonLispValueScopeForm::FunctionLiteral => { callable::collect_inferred_extract_function_lambda( - dialect, + semantic, view, 1, explicit_params, @@ -134,7 +135,7 @@ pub(super) fn collect_inferred_extract_function_special_form( ) } CommonLispValueScopeForm::Definition => callable::collect_inferred_extract_function_lambda( - dialect, + semantic, view, 2, explicit_params, @@ -143,6 +144,7 @@ pub(super) fn collect_inferred_extract_function_special_form( ), CommonLispValueScopeForm::LocalCallable(_) => { callable::collect_inferred_extract_function_local_callable_form( + semantic, view, explicit_params, bound_params, @@ -153,26 +155,26 @@ pub(super) fn collect_inferred_extract_function_special_form( } pub(super) fn extend_extract_function_bound_params<'a>( - dialect: Dialect, + semantic: ExtractFunctionSemantic, bound_params: &[String], names: impl Iterator, ) -> Vec { let mut extended = bound_params.to_vec(); for name in names { - push_extract_function_bound_param(dialect, &mut extended, name); + push_extract_function_bound_param(semantic, &mut extended, name); } extended } pub(super) fn push_extract_function_bound_param( - dialect: Dialect, + semantic: ExtractFunctionSemantic, bound_params: &mut Vec, name: &str, ) { if super::symbols::is_extract_function_param_candidate(name) && !bound_params .iter() - .any(|param| super::extract_function_param_name_eq(dialect, param, name)) + .any(|param| super::extract_function_param_name_eq(semantic, param, name)) { bound_params.push(name.to_owned()); } diff --git a/src/domain/extract_function/inference/patterns.rs b/src/domain/extract_function/inference/patterns.rs index 3ab88e59..3168684e 100644 --- a/src/domain/extract_function/inference/patterns.rs +++ b/src/domain/extract_function/inference/patterns.rs @@ -1,29 +1,39 @@ -use crate::domain::dialect::Dialect; use crate::domain::sexpr::ExpressionView; use super::super::syntax::atom_text; +use super::ExtractFunctionSemantic; use super::symbols::is_extract_function_param_candidate; -pub(super) fn parameter_names(parameter_form: &ExpressionView) -> Vec { +pub(super) fn parameter_names( + semantic: ExtractFunctionSemantic, + parameter_form: &ExpressionView, +) -> Vec { let mut names = Vec::new(); - collect_lambda_list_parameter_names(parameter_form, &mut names); + collect_lambda_list_parameter_names(semantic, parameter_form, &mut names); names } -pub(super) fn extract_function_pattern_names(pattern: &ExpressionView) -> Vec { +pub(super) fn extract_function_pattern_names( + semantic: ExtractFunctionSemantic, + pattern: &ExpressionView, +) -> Vec { let mut names = Vec::new(); - collect_extract_function_pattern_names(pattern, &mut names); + collect_extract_function_pattern_names(semantic, pattern, &mut names); names } -fn collect_extract_function_pattern_names(pattern: &ExpressionView, names: &mut Vec) { +fn collect_extract_function_pattern_names( + semantic: ExtractFunctionSemantic, + pattern: &ExpressionView, + names: &mut Vec, +) { if let Some(text) = atom_text(pattern) { - push_extract_function_pattern_name(text, names); + push_extract_function_pattern_name(semantic, text, names); return; } for child in &pattern.children { - collect_extract_function_pattern_names(child, names); + collect_extract_function_pattern_names(semantic, child, names); } } @@ -35,7 +45,11 @@ enum LambdaListMode { Aux, } -fn collect_lambda_list_parameter_names(parameter_form: &ExpressionView, names: &mut Vec) { +fn collect_lambda_list_parameter_names( + semantic: ExtractFunctionSemantic, + parameter_form: &ExpressionView, + names: &mut Vec, +) { let mut mode = LambdaListMode::Required; let mut index = 0usize; @@ -60,7 +74,7 @@ fn collect_lambda_list_parameter_names(parameter_form: &ExpressionView, names: & } "&rest" | "&body" | "&whole" | "&environment" => { if let Some(next) = parameter_form.children.get(index + 1) { - collect_extract_function_pattern_names(next, names); + collect_extract_function_pattern_names(semantic, next, names); } index += 2; continue; @@ -77,18 +91,19 @@ fn collect_lambda_list_parameter_names(parameter_form: &ExpressionView, names: & } } - collect_lambda_list_parameter_spec_names(child, mode, names); + collect_lambda_list_parameter_spec_names(semantic, child, mode, names); index += 1; } } fn collect_lambda_list_parameter_spec_names( + semantic: ExtractFunctionSemantic, spec: &ExpressionView, mode: LambdaListMode, names: &mut Vec, ) { if atom_text(spec).is_some() || mode == LambdaListMode::Required { - collect_extract_function_pattern_names(spec, names); + collect_extract_function_pattern_names(semantic, spec, names); return; } @@ -97,44 +112,58 @@ fn collect_lambda_list_parameter_spec_names( } match mode { - LambdaListMode::Required => collect_extract_function_pattern_names(spec, names), + LambdaListMode::Required => collect_extract_function_pattern_names(semantic, spec, names), LambdaListMode::Optional => { - collect_extract_function_pattern_names(&spec.children[0], names); - collect_supplied_p_name(spec, names); + collect_extract_function_pattern_names(semantic, &spec.children[0], names); + collect_supplied_p_name(semantic, spec, names); } LambdaListMode::Key => { - collect_key_parameter_name(&spec.children[0], names); - collect_supplied_p_name(spec, names); + collect_key_parameter_name(semantic, &spec.children[0], names); + collect_supplied_p_name(semantic, spec, names); + } + LambdaListMode::Aux => { + collect_extract_function_pattern_names(semantic, &spec.children[0], names) } - LambdaListMode::Aux => collect_extract_function_pattern_names(&spec.children[0], names), } } -fn collect_key_parameter_name(spec_name: &ExpressionView, names: &mut Vec) { +fn collect_key_parameter_name( + semantic: ExtractFunctionSemantic, + spec_name: &ExpressionView, + names: &mut Vec, +) { if spec_name.children.len() >= 2 { if let Some(designator) = atom_text(&spec_name.children[0]) { if designator.starts_with(':') { - collect_extract_function_pattern_names(&spec_name.children[1], names); + collect_extract_function_pattern_names(semantic, &spec_name.children[1], names); return; } } } - collect_extract_function_pattern_names(spec_name, names); + collect_extract_function_pattern_names(semantic, spec_name, names); } -fn collect_supplied_p_name(spec: &ExpressionView, names: &mut Vec) { +fn collect_supplied_p_name( + semantic: ExtractFunctionSemantic, + spec: &ExpressionView, + names: &mut Vec, +) { if let Some(supplied_p) = spec.children.get(2) { - collect_extract_function_pattern_names(supplied_p, names); + collect_extract_function_pattern_names(semantic, supplied_p, names); } } -fn push_extract_function_pattern_name(text: &str, names: &mut Vec) { +fn push_extract_function_pattern_name( + semantic: ExtractFunctionSemantic, + text: &str, + names: &mut Vec, +) { if text != "_" && is_extract_function_param_candidate(text) && !names .iter() - .any(|name| super::extract_function_param_name_eq(Dialect::CommonLisp, name, text)) + .any(|name| super::extract_function_param_name_eq(semantic, name, text)) { names.push(text.to_owned()); } diff --git a/src/domain/extract_function/inference/semantic.rs b/src/domain/extract_function/inference/semantic.rs new file mode 100644 index 00000000..933f8b0d --- /dev/null +++ b/src/domain/extract_function/inference/semantic.rs @@ -0,0 +1,398 @@ +use crate::domain::dialect::{ + BinderShape, BindingVisibility, BodyShape, DefinitionShape, Dialect, ParameterShape, + RelativeNodePath, ScopeShape, +}; +use crate::domain::sexpr::ExpressionView; + +use super::bindings::{ExtractFunctionBindingEntry, extract_function_binding_entries}; +use super::forms::{extend_extract_function_bound_params, push_extract_function_bound_param}; +use super::patterns::{extract_function_pattern_names, parameter_names}; +use super::{ExtractFunctionSemantic, collect_inferred_extract_function_params}; + +pub(super) fn collect_inferred_extract_function_semantic_form( + semantic: ExtractFunctionSemantic, + view: &ExpressionView, + explicit_params: &[String], + bound_params: &[String], + params: &mut Vec, +) -> bool { + if semantic.dialect() == Dialect::CommonLisp { + return false; + } + + if let Some(scope) = semantic.scope_shape(view) { + collect_scope(semantic, view, scope, explicit_params, bound_params, params); + return true; + } + + if let Some(definition) = semantic.definition_shape(view) { + collect_definition( + semantic, + view, + definition, + explicit_params, + bound_params, + params, + ); + return true; + } + + false +} + +fn collect_scope( + semantic: ExtractFunctionSemantic, + view: &ExpressionView, + scope: ScopeShape, + explicit_params: &[String], + bound_params: &[String], + params: &mut Vec, +) { + match scope.binders() { + BinderShape::BindingList { + container, + visibility, + .. + } + | BinderShape::FlatPairs { + container, + visibility, + .. + } => collect_binding_scope( + semantic, + view, + container, + None, + visibility, + scope.body(), + explicit_params, + bound_params, + params, + ), + BinderShape::NamedBindingList { + scope_name, + container, + visibility, + .. + } => collect_binding_scope( + semantic, + view, + container, + Some(scope_name), + visibility, + scope.body(), + explicit_params, + bound_params, + params, + ), + BinderShape::Parameters(parameters) => collect_parameter_scope( + semantic, + view, + None, + parameters, + scope.body(), + explicit_params, + bound_params, + params, + ), + BinderShape::NamedParameters { name, parameters } => collect_parameter_scope( + semantic, + view, + Some(name), + parameters, + scope.body(), + explicit_params, + bound_params, + params, + ), + BinderShape::ParameterClauses { + name, + first_clause_index, + parameters, + } => collect_parameter_clauses( + semantic, + view, + name, + first_clause_index, + parameters, + scope.body(), + explicit_params, + bound_params, + params, + ), + } +} + +#[allow(clippy::too_many_arguments)] +fn collect_binding_scope( + semantic: ExtractFunctionSemantic, + view: &ExpressionView, + container_path: RelativeNodePath, + scope_name: Option, + visibility: BindingVisibility, + body: BodyShape, + explicit_params: &[String], + bound_params: &[String], + params: &mut Vec, +) { + let Some(container) = resolve_relative(view, container_path) else { + return; + }; + let Some(entries) = extract_function_binding_entries(semantic, container) else { + return; + }; + + let mut body_bound_params = bound_params.to_vec(); + match visibility { + BindingVisibility::Parallel => { + for entry in &entries { + collect_binding_initializer(semantic, entry, explicit_params, bound_params, params); + } + extend_with_binding_names(semantic, &mut body_bound_params, &entries); + } + BindingVisibility::Sequential => { + for entry in &entries { + collect_binding_initializer( + semantic, + entry, + explicit_params, + &body_bound_params, + params, + ); + extend_with_names(semantic, &mut body_bound_params, &entry.names); + } + } + } + + if let Some(scope_name) = scope_name { + let Some(name) = resolve_relative(view, scope_name) else { + return; + }; + extend_with_pattern(semantic, &mut body_bound_params, name); + } + + collect_body( + semantic, + view, + body, + explicit_params, + &body_bound_params, + params, + ); +} + +fn collect_binding_initializer( + semantic: ExtractFunctionSemantic, + entry: &ExtractFunctionBindingEntry, + explicit_params: &[String], + bound_params: &[String], + params: &mut Vec, +) { + if let Some(value) = &entry.value { + collect_inferred_extract_function_params( + semantic, + value, + false, + explicit_params, + bound_params, + params, + ); + } +} + +fn extend_with_binding_names( + semantic: ExtractFunctionSemantic, + bound_params: &mut Vec, + entries: &[ExtractFunctionBindingEntry], +) { + for entry in entries { + extend_with_names(semantic, bound_params, &entry.names); + } +} + +#[allow(clippy::too_many_arguments)] +fn collect_parameter_scope( + semantic: ExtractFunctionSemantic, + view: &ExpressionView, + name: Option, + parameters: ParameterShape, + body: BodyShape, + explicit_params: &[String], + bound_params: &[String], + params: &mut Vec, +) { + let Some(parameter_names) = parameter_names_at(semantic, view, parameters) else { + return; + }; + let mut body_bound_params = extend_extract_function_bound_params( + semantic, + bound_params, + parameter_names.iter().map(String::as_str), + ); + + if let Some(name) = name { + let Some(name) = resolve_relative(view, name) else { + return; + }; + extend_with_pattern(semantic, &mut body_bound_params, name); + } + + collect_body( + semantic, + view, + body, + explicit_params, + &body_bound_params, + params, + ); +} + +#[allow(clippy::too_many_arguments)] +fn collect_parameter_clauses( + semantic: ExtractFunctionSemantic, + view: &ExpressionView, + name: Option, + first_clause_index: usize, + parameters: ParameterShape, + body: BodyShape, + explicit_params: &[String], + bound_params: &[String], + params: &mut Vec, +) { + let BodyShape::ClauseChildrenFrom { + first_clause_index: body_first_clause_index, + body_child_index, + } = body + else { + return; + }; + if first_clause_index != body_first_clause_index { + return; + } + + let name = name.and_then(|path| resolve_relative(view, path)); + for clause in view.children.iter().skip(first_clause_index) { + let Some(parameter_names) = parameter_names_at(semantic, clause, parameters) else { + return; + }; + let mut clause_bound_params = extend_extract_function_bound_params( + semantic, + bound_params, + parameter_names.iter().map(String::as_str), + ); + if let Some(name) = name { + extend_with_pattern(semantic, &mut clause_bound_params, name); + } + + for child in clause.children.iter().skip(body_child_index) { + collect_inferred_extract_function_params( + semantic, + child, + false, + explicit_params, + &clause_bound_params, + params, + ); + } + } +} + +#[allow(clippy::too_many_arguments)] +fn collect_definition( + semantic: ExtractFunctionSemantic, + view: &ExpressionView, + definition: DefinitionShape, + explicit_params: &[String], + bound_params: &[String], + params: &mut Vec, +) { + let mut body_bound_params = bound_params.to_vec(); + if let Some(name) = definition.name() { + let Some(name) = resolve_relative(view, name) else { + return; + }; + extend_with_pattern(semantic, &mut body_bound_params, name); + } + if let Some(parameters) = definition.parameters() { + let Some(names) = parameter_names_at(semantic, view, parameters) else { + return; + }; + extend_with_names(semantic, &mut body_bound_params, &names); + } + + collect_body( + semantic, + view, + definition.body(), + explicit_params, + &body_bound_params, + params, + ); +} + +fn collect_body( + semantic: ExtractFunctionSemantic, + view: &ExpressionView, + body: BodyShape, + explicit_params: &[String], + bound_params: &[String], + params: &mut Vec, +) { + let first_body_index = match body { + BodyShape::ChildrenFrom(index) => index, + BodyShape::ChildrenAfter(path) => path.child() + 1, + BodyShape::ClauseChildrenFrom { .. } => return, + }; + + for child in view.children.iter().skip(first_body_index) { + collect_inferred_extract_function_params( + semantic, + child, + false, + explicit_params, + bound_params, + params, + ); + } +} + +fn parameter_names_at( + semantic: ExtractFunctionSemantic, + view: &ExpressionView, + parameters: ParameterShape, +) -> Option> { + let parameter_form = resolve_relative(view, parameters.container())?; + let first_parameter_index = parameters.first_parameter_index(); + if first_parameter_index > parameter_form.children.len() { + return None; + } + + let mut parameter_form = parameter_form.clone(); + parameter_form.children.drain(..first_parameter_index); + Some(parameter_names(semantic, ¶meter_form)) +} + +fn extend_with_pattern( + semantic: ExtractFunctionSemantic, + bound_params: &mut Vec, + pattern: &ExpressionView, +) { + let names = extract_function_pattern_names(semantic, pattern); + extend_with_names(semantic, bound_params, &names); +} + +fn extend_with_names( + semantic: ExtractFunctionSemantic, + bound_params: &mut Vec, + names: &[String], +) { + for name in names { + push_extract_function_bound_param(semantic, bound_params, name); + } +} + +fn resolve_relative(view: &ExpressionView, path: RelativeNodePath) -> Option<&ExpressionView> { + let child = view.children.get(path.child())?; + path.grandchild() + .map_or(Some(child), |grandchild| child.children.get(grandchild)) +} diff --git a/src/domain/extract_function/tests/inference.rs b/src/domain/extract_function/tests/inference.rs index 83da1046..746a30c1 100644 --- a/src/domain/extract_function/tests/inference.rs +++ b/src/domain/extract_function/tests/inference.rs @@ -374,6 +374,118 @@ fn treats_emacs_lisp_flet_lambda_list_as_local_to_function_body() { assert_eq!(params, vec!["outer", "input"]); } +#[test] +fn resolver_binding_scopes_preserve_dialect_visibility() { + let cases = [ + ( + Dialect::EmacsLisp, + "(let* ((first seed) (second first)) (list first second outer))", + &["seed", "outer"][..], + ), + ( + Dialect::Scheme, + "(let ((first seed) (second first)) (list first second outer))", + &["seed", "first", "outer"][..], + ), + ( + Dialect::Clojure, + "(let [first seed second first] [first second outer])", + &["seed", "outer"][..], + ), + ( + Dialect::Janet, + "(let [first seed second first] [first second outer])", + &["seed", "outer"][..], + ), + ( + Dialect::Fennel, + "(let [first seed second first] [first second outer])", + &["seed", "outer"][..], + ), + ]; + + for (dialect, input, expected) in cases { + assert_eq!( + infer_at_dialect(dialect, input, &[0], &[]), + expected, + "{dialect:?}" + ); + } +} + +#[test] +fn verified_dialects_exclude_callable_parameters() { + let cases = [ + (Dialect::CommonLisp, "(lambda (local) (list local outer))"), + (Dialect::EmacsLisp, "(lambda (local) (list local outer))"), + (Dialect::Scheme, "(lambda (local) (list local outer))"), + (Dialect::Clojure, "(fn [local] [local outer])"), + (Dialect::Janet, "(fn [local] [local outer])"), + (Dialect::Fennel, "(fn [local] [local outer])"), + ]; + + for (dialect, input) in cases { + assert_eq!( + infer_at_dialect(dialect, input, &[0], &[]), + ["outer"], + "{dialect:?}" + ); + } +} + +#[test] +fn verified_dialects_exclude_definition_parameters() { + let cases = [ + ( + Dialect::CommonLisp, + "(defun render (local) (list local outer))", + ), + ( + Dialect::EmacsLisp, + "(defun render (local) (list local outer))", + ), + ( + Dialect::Scheme, + "(define (render local) (list local outer))", + ), + (Dialect::Clojure, "(defn render [local] [local outer])"), + (Dialect::Janet, "(defn render [local] [local outer])"), + (Dialect::Fennel, "(fn render [local] [local outer])"), + ]; + + for (dialect, input) in cases { + assert_eq!( + infer_at_dialect(dialect, input, &[0], &[]), + ["outer"], + "{dialect:?}" + ); + } +} + +#[test] +fn treats_scheme_named_let_name_and_parameters_as_local() { + let params = infer_at_dialect( + Dialect::Scheme, + "(let loop ((local seed)) (list loop local outer))", + &[0], + &[], + ); + + assert_eq!(params, vec!["seed", "outer"]); +} + +#[test] +fn treats_clojure_named_multi_arity_fn_bindings_as_local() { + let params = infer_at_dialect( + Dialect::Clojure, + "(fn render ([local] [render local outer]) ([local other] [render local other extra]))", + &[0], + &[], + ); + + assert_eq!(params, vec!["outer", "extra"]); +} + #[test] fn treats_define_setf_expander_macro_lambda_list_as_local_to_expander_body() { let params = infer_at( @@ -416,8 +528,9 @@ fn keeps_flet_callable_name_as_free_value_reference() { #[test] fn excludes_destructured_lambda_parameters_and_explicit_params() { - let params = infer_at( - "(lambda [{:keys [inner]}] (+ inner outer ignored))", + let params = infer_at_dialect( + Dialect::CommonLisp, + "(lambda ((inner)) (+ inner outer ignored))", &[0], &["ignored"], ); diff --git a/src/domain/extract_function/tests/mod.rs b/src/domain/extract_function/tests/mod.rs index cee25e4b..0fcb79b8 100644 --- a/src/domain/extract_function/tests/mod.rs +++ b/src/domain/extract_function/tests/mod.rs @@ -16,7 +16,7 @@ fn infer_at_dialect( path: &[usize], explicit: &[&str], ) -> Vec { - let tree = SyntaxTree::parse(input).expect("parse fixture"); + let tree = SyntaxTree::parse_with_dialect(input, dialect).expect("parse fixture"); let selection = tree .select_path(&Path::from_indexes(path.to_vec())) .expect("select fixture"); @@ -35,7 +35,25 @@ fn plan_at( explicit: &[&str], infer_params: bool, ) -> ExtractFunctionPlan { - let tree = SyntaxTree::parse(input).expect("parse fixture"); + plan_at_dialect( + Dialect::CommonLisp, + input, + path, + name, + explicit, + infer_params, + ) +} + +fn plan_at_dialect( + dialect: Dialect, + input: &str, + path: &[usize], + name: &str, + explicit: &[&str], + infer_params: bool, +) -> ExtractFunctionPlan { + let tree = SyntaxTree::parse_with_dialect(input, dialect).expect("parse fixture"); let selection = tree .select_path(&Path::from_indexes(path.to_vec())) .expect("select fixture"); @@ -48,7 +66,7 @@ fn plan_at( input, selection, path: Some(Path::from_indexes(path.to_vec())), - dialect: Dialect::CommonLisp, + dialect, name: SymbolName::new(name).expect("symbol fixture"), explicit_params, infer_params, diff --git a/src/domain/extract_function/tests/planning.rs b/src/domain/extract_function/tests/planning.rs index 59b48ba7..4c6697e1 100644 --- a/src/domain/extract_function/tests/planning.rs +++ b/src/domain/extract_function/tests/planning.rs @@ -1,5 +1,34 @@ use super::*; +#[test] +fn rejects_unknown_dialect_before_extract_function_planning() { + let input = "(+ width height)"; + let tree = SyntaxTree::parse(input).expect("parse fixture"); + let path = Path::from_indexes(vec![0]); + let selection = tree.select_path(&path).expect("select fixture"); + + assert!(infer_extract_function_params(Dialect::Unknown, &selection.view(), &[]).is_empty()); + + let error = plan_extract_function(ExtractFunctionRequest { + input, + selection, + path: Some(path), + dialect: Dialect::Unknown, + name: SymbolName::new("area").expect("symbol fixture"), + explicit_params: Vec::new(), + infer_params: true, + insert: ExtractFunctionInsert::Append, + anchor_path: None, + }) + .expect_err("unknown dialect should be rejected"); + + assert!( + error + .to_string() + .contains("extract-function is not supported for this dialect") + ); +} + #[test] fn plans_extract_function_with_inferred_params() { let plan = plan_at( @@ -17,7 +46,28 @@ fn plans_extract_function_with_inferred_params() { ); assert_eq!(plan.inferred_params, vec!["width", "height"]); assert!(plan.changed); - SyntaxTree::parse(&plan.rewritten).expect("rewritten output remains parseable"); + SyntaxTree::parse_with_dialect(&plan.rewritten, Dialect::CommonLisp) + .expect("rewritten output remains parseable"); +} + +#[test] +fn extract_function_preserves_dialect_reader_collisions() { + let cases = [( + Dialect::Janet, + "(+ width height)\n# ignored ))", + "(area width height)", + "(defn area [width height] (+ width height))", + )]; + + for (dialect, input, expected_call, expected_definition) in cases { + let plan = plan_at_dialect(dialect, input, &[0], "area", &["width", "height"], false); + + assert_eq!(plan.call, expected_call); + assert_eq!(plan.definition, expected_definition); + assert!(plan.rewritten.contains("# ignored ))")); + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect) + .expect("rewritten output remains parseable"); + } } #[test] diff --git a/src/domain/extract_local_function.rs b/src/domain/extract_local_function.rs index 42ba222e..e3215519 100644 --- a/src/domain/extract_local_function.rs +++ b/src/domain/extract_local_function.rs @@ -46,15 +46,20 @@ pub struct ExtractLocalFunctionPlan { pub changed: bool, } +fn ensure_common_lisp_dialect(dialect: Dialect) -> Result<()> { + if dialect != Dialect::CommonLisp { + bail!("extract-local-function currently supports only Common Lisp"); + } + Ok(()) +} + pub fn plan_extract_local_function( request: ExtractLocalFunctionRequest<'_>, ) -> Result { request.selection.validate_source(request.input)?; request.enclosing.validate_source(request.input)?; - if request.dialect != Dialect::CommonLisp { - bail!("extract-local-function currently supports only Common Lisp"); - } - let tree = SyntaxTree::parse(request.input) + ensure_common_lisp_dialect(request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) .context("extract-local-function input is not a valid S-expression document")?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; let path = request @@ -114,7 +119,7 @@ pub fn plan_extract_local_function( enclosed ); let rewritten = replace_span(request.input, enclosing_span, &replacement); - SyntaxTree::parse(&rewritten) + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) .context("extracted local function output is not a valid S-expression document")?; Ok(ExtractLocalFunctionPlan { @@ -1142,6 +1147,63 @@ mod tests { }) } + #[test] + fn dialect_matrix_gates_unsupported_dialects_before_parsing() { + for (dialect, supported) in [ + (Dialect::CommonLisp, true), + (Dialect::EmacsLisp, false), + (Dialect::Scheme, false), + (Dialect::Clojure, false), + (Dialect::Janet, false), + (Dialect::Fennel, false), + (Dialect::Unknown, false), + ] { + let parse_attempted = std::cell::Cell::new(false); + match ensure_common_lisp_dialect(dialect) { + Ok(()) => { + parse_attempted.set(true); + assert!(SyntaxTree::parse_with_dialect(")", dialect).is_err()); + } + Err(error) => { + assert!(!supported); + assert_eq!( + error.to_string(), + "extract-local-function currently supports only Common Lisp" + ); + } + } + + assert_eq!(parse_attempted.get(), supported); + } + } + + #[test] + fn preserves_common_lisp_reader_character_literal() -> Result<()> { + let input = r"(progn (print #\)) (finish))"; + let tree = SyntaxTree::parse_with_dialect(input, Dialect::CommonLisp)?; + let path: Path = "0.1".parse()?; + let enclosing_path: Path = "0".parse()?; + let selection = tree.select_path(&path)?; + let enclosing = tree.select_path(&enclosing_path)?; + + let result = plan_extract_local_function(ExtractLocalFunctionRequest { + input, + selection, + path: Some(path), + enclosing, + enclosing_path, + dialect: Dialect::CommonLisp, + name: SymbolName::new("compute")?, + explicit_params: Vec::new(), + infer_params: true, + recursive: false, + })?; + + assert!(result.rewritten.contains(r"#\)")); + SyntaxTree::parse_with_dialect(&result.rewritten, Dialect::CommonLisp)?; + Ok(()) + } + #[test] fn extracts_into_flet_and_infers_free_values() { let result = plan( diff --git a/src/domain/flet_composition.rs b/src/domain/flet_composition.rs index 8114268c..6e1bb5d4 100644 --- a/src/domain/flet_composition.rs +++ b/src/domain/flet_composition.rs @@ -24,10 +24,9 @@ pub(crate) struct Plan { } pub(crate) fn plan(request: Request<'_>) -> Result { - if request.dialect != Dialect::CommonLisp { - bail!("merge-nested-flet supports only Common Lisp"); - } - let tree = SyntaxTree::parse(request.input).context("merge-nested-flet input is not valid")?; + validate_dialect(request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("merge-nested-flet input is not valid")?; let outer = tree.select_path(&request.path)?.view(); require_flet(&outer, "selected form")?; reject_unsafe(&tree, &outer)?; @@ -73,7 +72,8 @@ pub(crate) fn plan(request: Request<'_>) -> Result { let body = &request.input[inner_defs.span.end().get()..inner.span.end().get() - 1]; let replacement = format!("({head} ({left}{separator}{right}){body})"); let rewritten = replace_span(request.input, outer.span, &replacement); - SyntaxTree::parse(&rewritten).context("merge-nested-flet output is not valid")?; + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) + .context("merge-nested-flet output is not valid")?; Ok(Plan { dialect: request.dialect, path: request.path, @@ -84,6 +84,13 @@ pub(crate) fn plan(request: Request<'_>) -> Result { rewritten, }) } + +pub(crate) fn validate_dialect(dialect: Dialect) -> Result<()> { + if dialect != Dialect::CommonLisp { + bail!("merge-nested-flet supports only Common Lisp"); + } + Ok(()) +} fn require_flet(view: &ExpressionView, role: &str) -> Result<()> { if view.kind != ExpressionKind::List || !view.reader_prefixes.is_empty() @@ -196,6 +203,7 @@ fn replace_span(input: &str, span: ByteSpan, replacement: &str) -> String { #[cfg(test)] mod tests { use super::*; + fn request(input: &str, dialect: Dialect) -> Request<'_> { Request { input, @@ -203,20 +211,40 @@ mod tests { path: "0".parse().expect("path"), } } + #[test] - fn merges_independent_definitions() { - let plan = plan(request( - "(flet ((parse (x) (list x))) (flet ((emit (x) (print x))) (emit (parse value))))", - Dialect::CommonLisp, - )) - .expect("plan"); + fn merges_independent_definitions_with_common_lisp_reader_literal() { + let input = + r"(flet ((parse (x) (list x))) (flet ((emit (x) (print x))) (emit (parse value)))) #\)"; + let plan = plan(request(input, Dialect::CommonLisp)).expect("plan"); assert_eq!( plan.rewritten, - "(flet ((parse (x) (list x)) (emit (x) (print x))) (emit (parse value)))" + r"(flet ((parse (x) (list x)) (emit (x) (print x))) (emit (parse value))) #\)" ); + SyntaxTree::parse_with_dialect(&plan.rewritten, Dialect::CommonLisp) + .expect("rewritten input"); + } + + #[test] + fn rejects_unsupported_dialects_before_parsing() { + for dialect in [ + Dialect::EmacsLisp, + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ] { + let error = plan(request(")", dialect)).expect_err("dialect must be rejected"); + assert_eq!( + error.to_string(), + "merge-nested-flet supports only Common Lisp" + ); + } } + #[test] - fn rejects_scope_capture_duplicates_unsafe_syntax_and_dialect() { + fn rejects_scope_capture_duplicates_and_unsafe_syntax() { for input in [ "(flet ((parse (x) x)) (flet ((emit (x) (parse x))) (emit value)))", "(flet ((work () 1)) (flet ((work () 2)) (work)))", @@ -224,12 +252,5 @@ mod tests { ] { assert!(plan(request(input, Dialect::CommonLisp)).is_err()); } - assert!( - plan(request( - "(flet ((left () 1)) (flet ((right () 2)) (right)))", - Dialect::EmacsLisp - )) - .is_err() - ); } } diff --git a/src/domain/function_parameter/add/mod.rs b/src/domain/function_parameter/add/mod.rs index a78a2350..8ec90171 100644 --- a/src/domain/function_parameter/add/mod.rs +++ b/src/domain/function_parameter/add/mod.rs @@ -1,5 +1,6 @@ use anyhow::{Context, Result}; +use crate::dialect::Dialect; use crate::domain::mutation_safety::reject_common_lisp_reader_conditionals; use crate::domain::sexpr::SyntaxTree; @@ -19,8 +20,8 @@ use definition_insertion::{DefinitionInsertionPlan, resolve_definition_insertion pub fn plan_add_function_parameter( request: AddFunctionParameterRequest<'_>, ) -> Result { - let argument = validate_argument(&request.argument)?; - let tree = SyntaxTree::parse(request.input)?; + let argument = validate_argument(&request.argument, request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect)?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; let target = parse_add_function_parameter_definition( request.dialect, @@ -96,7 +97,7 @@ pub fn plan_add_function_parameter( sorted_call_spans.sort_by_key(|span| span.start()); ensure_non_overlapping_spans(sorted_call_spans)?; let rewritten = apply_byte_span_edits(request.input, edits)?; - SyntaxTree::parse(&rewritten) + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) .context("add-function-parameter output is not a valid S-expression document")?; let changed = rewritten != request.input; @@ -118,12 +119,12 @@ pub fn plan_add_function_parameter( }) } -fn validate_argument(argument: &str) -> Result { +fn validate_argument(argument: &str, dialect: Dialect) -> Result { let argument = argument.trim().to_owned(); if argument.is_empty() { anyhow::bail!("--argument must not be empty"); } - let argument_tree = SyntaxTree::parse(&argument) + let argument_tree = SyntaxTree::parse_with_dialect(&argument, dialect) .context("add-function-parameter argument is not a valid S-expression")?; if argument_tree.root_children().len() != 1 { anyhow::bail!("--argument must contain exactly one top-level S-expression"); diff --git a/src/domain/function_parameter/move_parameter.rs b/src/domain/function_parameter/move_parameter.rs index 82a9ecf0..f68b1fe9 100644 --- a/src/domain/function_parameter/move_parameter.rs +++ b/src/domain/function_parameter/move_parameter.rs @@ -21,7 +21,7 @@ use super::types::{MoveFunctionParameterPlan, MoveFunctionParameterRequest}; pub fn plan_move_function_parameter( request: MoveFunctionParameterRequest<'_>, ) -> Result { - let tree = SyntaxTree::parse(request.input)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect)?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; let target = parse_move_function_parameter_definition(request.dialect, &tree, &request.definition_path)?; @@ -121,7 +121,7 @@ pub fn plan_move_function_parameter( sorted_call_spans.sort_by_key(|span| span.start()); ensure_non_overlapping_spans(sorted_call_spans)?; let rewritten = apply_byte_span_edits(request.input, edits)?; - SyntaxTree::parse(&rewritten) + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) .context("move-function-parameter output is not a valid S-expression document")?; let changed = rewritten != request.input; diff --git a/src/domain/function_parameter/remove/mod.rs b/src/domain/function_parameter/remove/mod.rs index 7feb5899..dcc661a1 100644 --- a/src/domain/function_parameter/remove/mod.rs +++ b/src/domain/function_parameter/remove/mod.rs @@ -17,7 +17,7 @@ use metadata::resolve_remove_parameter_metadata; pub fn plan_remove_function_parameter( request: RemoveFunctionParameterRequest<'_>, ) -> Result { - let tree = SyntaxTree::parse(request.input)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect)?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; let target = parse_remove_function_parameter_definition( request.dialect, @@ -68,7 +68,7 @@ pub fn plan_remove_function_parameter( sorted_call_spans.sort_by_key(|span| span.start()); ensure_non_overlapping_spans(sorted_call_spans)?; let rewritten = apply_byte_span_edits(request.input, edits)?; - SyntaxTree::parse(&rewritten) + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) .context("remove-function-parameter output is not a valid S-expression document")?; let changed = rewritten != request.input; diff --git a/src/domain/function_parameter/reorder/mod.rs b/src/domain/function_parameter/reorder/mod.rs index 1d0e2b91..f56ef43a 100644 --- a/src/domain/function_parameter/reorder/mod.rs +++ b/src/domain/function_parameter/reorder/mod.rs @@ -32,7 +32,7 @@ pub fn plan_reorder_function_parameters( anyhow::bail!("reorder-function-parameters requires at least one --parameter"); } - let tree = SyntaxTree::parse(request.input)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect)?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; let target = parse_reorder_function_parameters_definition( request.dialect, @@ -106,7 +106,7 @@ pub fn plan_reorder_function_parameters( sorted_call_spans.sort_by_key(|span| span.start()); ensure_non_overlapping_spans(sorted_call_spans)?; let rewritten = apply_byte_span_edits(request.input, edits)?; - SyntaxTree::parse(&rewritten) + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) .context("reorder-function-parameters output is not a valid S-expression document")?; let changed = rewritten != request.input; diff --git a/src/domain/function_parameter/swap.rs b/src/domain/function_parameter/swap.rs index 26c14f22..7a9d2136 100644 --- a/src/domain/function_parameter/swap.rs +++ b/src/domain/function_parameter/swap.rs @@ -29,7 +29,7 @@ pub fn plan_swap_function_parameters( anyhow::bail!("swap-function-parameters requires two distinct parameter names"); } - let tree = SyntaxTree::parse(request.input)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect)?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; let target = parse_swap_function_parameters_definition( request.dialect, @@ -145,7 +145,7 @@ pub fn plan_swap_function_parameters( sorted_call_spans.sort_by_key(|span| span.start()); ensure_non_overlapping_spans(sorted_call_spans)?; let rewritten = apply_byte_span_edits(request.input, edits)?; - SyntaxTree::parse(&rewritten) + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) .context("swap-function-parameters output is not a valid S-expression document")?; let changed = rewritten != request.input; diff --git a/src/domain/function_parameter/tests/add/basic.rs b/src/domain/function_parameter/tests/add/basic.rs index 8e690685..a5881614 100644 --- a/src/domain/function_parameter/tests/add/basic.rs +++ b/src/domain/function_parameter/tests/add/basic.rs @@ -22,6 +22,35 @@ fn adds_parameter_to_definition_and_call() { assert!(plan.changed); } +#[test] +fn add_parameter_preserves_dialect_reader_collisions() { + let cases = [( + Dialect::Janet, + "(defn area [w] w)\n(area 3)\n# ignored ))", + "[value # ignored ))\n next]", + "(defn area [w h] w)\n(area 3 [value # ignored ))\n next])\n# ignored ))", + )]; + + for (dialect, input, argument, expected) in cases { + let plan = plan_add_function_parameter(AddFunctionParameterRequest { + input, + dialect, + definition_path: path("0"), + name: symbol("h"), + argument: argument.to_owned(), + call_paths: vec![path("1")], + all_calls: false, + insert: FunctionParameterInsert::End, + section: FunctionParameterSection::Auto, + }) + .expect("plan"); + + assert_eq!(plan.rewritten, expected); + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect) + .expect("rewritten output remains parseable"); + } +} + #[test] fn adds_parameter_to_package_qualified_common_lisp_definition() { let input = "(cl:defun area (w) w)\n(print (area 3))"; diff --git a/src/domain/inline_function.rs b/src/domain/inline_function.rs index 817791ff..0dada455 100644 --- a/src/domain/inline_function.rs +++ b/src/domain/inline_function.rs @@ -44,7 +44,8 @@ struct InlineFunctionParts { } pub fn plan_inline_function(request: InlineFunctionRequest<'_>) -> Result { - let tree = SyntaxTree::parse(request.input)?; + ensure_inline_function_dialect_supported(request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect)?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; let definition_selection = tree.select_path(&request.definition_path)?; let definition_span = definition_selection.span(); @@ -119,7 +120,7 @@ pub fn plan_inline_function(request: InlineFunctionRequest<'_>) -> Result) -> Result bool { + matches!(dialect, Dialect::CommonLisp | Dialect::EmacsLisp) +} + +fn ensure_inline_function_dialect_supported(dialect: Dialect) -> Result<()> { + if supports_inline_function_dialect(dialect) { + Ok(()) + } else { + anyhow::bail!( + "inline-function does not support dialect {}", + dialect.label() + ) + } +} + fn inline_function_parts( dialect: Dialect, input: &str, @@ -150,7 +166,7 @@ fn inline_function_parts( parse_inline_function_definition(dialect, input, definition_selection.clone())?; validate_macro_environment_parameters(dialect, input, &definition)?; let function_name = definition.name.clone(); - let call = parse_inline_function_call(call_selection.clone(), &function_name, input)?; + let call = parse_inline_function_call(dialect, call_selection.clone(), &function_name, input)?; let bindings = bind_inline_function_arguments( dialect, &definition.params, @@ -174,7 +190,8 @@ fn inline_function_parts( )?; let (argument_param_names, argument_args): (Vec<_>, Vec<_>) = bindings.argument_bindings.into_iter().unzip(); - let replacement_tree = SyntaxTree::parse(&intermediate_replacement)?; + let replacement_tree = + SyntaxTree::parse_with_dialect(&intermediate_replacement, dialect)?; let replacement_body = replacement_view_for_inline_function(&replacement_tree)?; let (replacement, argument_parameters) = substitute_inline_function_body( dialect, @@ -265,9 +282,10 @@ fn validate_macro_environment_parameters( )?; for (context, default_value) in &initialization_expressions { - let default_tree = SyntaxTree::parse(default_value).with_context(|| { - format!("inline-function could not parse {context}: {default_value}") - })?; + let default_tree = SyntaxTree::parse_with_dialect(default_value, dialect) + .with_context(|| { + format!("inline-function could not parse {context}: {default_value}") + })?; let default_expression = default_tree .select_path(&crate::domain::sexpr::Path::root_child(0))? .view(); diff --git a/src/domain/inline_function/calls/destructure.rs b/src/domain/inline_function/calls/destructure.rs index 8d434a5c..3e8e646a 100644 --- a/src/domain/inline_function/calls/destructure.rs +++ b/src/domain/inline_function/calls/destructure.rs @@ -18,7 +18,7 @@ pub(super) fn destructure_argument_entries( argument: &str, allow_drop_arguments: bool, ) -> Result> { - let tree = SyntaxTree::parse(argument).with_context(|| { + let tree = SyntaxTree::parse_with_dialect(argument, dialect).with_context(|| { format!("inline-function could not parse macro destructuring argument: {argument}") })?; let argument_expression = tree diff --git a/src/domain/inline_function/calls/discovery.rs b/src/domain/inline_function/calls/discovery.rs index 963d4ada..d341af5a 100644 --- a/src/domain/inline_function/calls/discovery.rs +++ b/src/domain/inline_function/calls/discovery.rs @@ -1,16 +1,16 @@ use anyhow::Result; use crate::domain::callable_scope::{ - common_lisp_local_callable_form, is_local_callable_bound, local_callable_binding_body_scope, - local_callable_body_scope, + common_lisp_local_callable_form, local_callable_binding_body_scope, local_callable_body_scope, }; -use crate::domain::common_lisp::{CommonLispLocalCallableForm, common_lisp_symbol_reference_eq}; +use crate::domain::common_lisp::CommonLispLocalCallableForm; use crate::domain::dialect::Dialect; use crate::domain::sexpr::{ ByteSpan, Delimiter, ExpressionKind, ExpressionView, Path, SymbolName, SyntaxTree, }; use super::super::syntax::{list_head, spans_overlap}; +use super::inline_function_symbol_reference_eq; struct InlineCallTraversal<'a> { dialect: Dialect, @@ -70,9 +70,17 @@ fn collect_function_call_paths( && view.delimiter == Some(Delimiter::Paren) && !spans_overlap(context.definition_span, view.span) && list_head(view).is_some_and(|head| { - common_lisp_symbol_reference_eq(head, context.function_name.as_str()) + inline_function_symbol_reference_eq( + context.dialect, + head, + context.function_name.as_str(), + ) }) - && !is_local_callable_bound(local_callables, context.function_name.as_str()) + && !is_inline_function_local_callable_bound( + context.dialect, + local_callables, + context.function_name.as_str(), + ) { output.push(path.clone()); } @@ -82,6 +90,12 @@ fn collect_function_call_paths( } } +fn is_inline_function_local_callable_bound(dialect: Dialect, scope: &[String], head: &str) -> bool { + scope + .iter() + .any(|name| inline_function_symbol_reference_eq(dialect, name, head)) +} + fn collect_local_callable_function_call_paths( view: &ExpressionView, path: Path, diff --git a/src/domain/inline_function/calls/mod.rs b/src/domain/inline_function/calls/mod.rs index dfa527ee..5d5395fc 100644 --- a/src/domain/inline_function/calls/mod.rs +++ b/src/domain/inline_function/calls/mod.rs @@ -55,6 +55,7 @@ pub(super) fn resolve_function_call_paths( } pub(super) fn parse_inline_function_call( + dialect: Dialect, view: ExpressionView, function_name: &SymbolName, input: &str, @@ -68,7 +69,7 @@ pub(super) fn parse_inline_function_call( .context("inline-function call must not be empty")?, ) .context("inline-function call must start with an atom")?; - if !common_lisp_symbol_reference_eq(head, function_name.as_str()) { + if !inline_function_symbol_reference_eq(dialect, head, function_name.as_str()) { anyhow::bail!( "inline-function call head '{}' does not match selected definition '{}'", head, @@ -105,7 +106,7 @@ fn validate_explicit_function_call_paths( let head = list_head(&view) .context("inline-function call must not be empty")? .to_owned(); - if !common_lisp_symbol_reference_eq(&head, function_name.as_str()) { + if !inline_function_symbol_reference_eq(dialect, &head, function_name.as_str()) { anyhow::bail!( "{command} --call-path {call_path} head '{}' does not match selected definition '{}'", head, @@ -123,6 +124,22 @@ fn validate_explicit_function_call_paths( Ok(()) } +pub(super) fn inline_function_symbol_reference_eq( + dialect: Dialect, + left: &str, + right: &str, +) -> bool { + match dialect { + Dialect::CommonLisp => common_lisp_symbol_reference_eq(left, right), + Dialect::EmacsLisp + | Dialect::Scheme + | Dialect::Clojure + | Dialect::Janet + | Dialect::Fennel => left == right, + Dialect::Unknown => false, + } +} + pub(super) fn validate_or_resolve_function_call_paths( tree: &SyntaxTree, dialect: Dialect, diff --git a/src/domain/inline_function/macro_expansion/expansion.rs b/src/domain/inline_function/macro_expansion/expansion.rs index bc9491da..0e3decbe 100644 --- a/src/domain/inline_function/macro_expansion/expansion.rs +++ b/src/domain/inline_function/macro_expansion/expansion.rs @@ -18,7 +18,8 @@ pub(super) fn expand_unquote_expression( ) -> Result { let literal_source = render_unquoted_source(view)?; let literal_tree = - crate::domain::sexpr::SyntaxTree::parse(&literal_source).context("invalid unquote form")?; + crate::domain::sexpr::SyntaxTree::parse_with_dialect(&literal_source, dialect) + .context("invalid unquote form")?; let expression = literal_tree .select_path(&crate::domain::sexpr::Path::root_child(0))? .view(); @@ -39,7 +40,7 @@ pub(super) fn expand_unquote_expression( )?; count_references_in_expanded_expression(dialect, &intermediate, reference_counts)?; - let intermediate_tree = parse_single_expression_tree(&intermediate)?; + let intermediate_tree = parse_single_expression_tree(dialect, &intermediate)?; let intermediate_expression = intermediate_tree .select_path(&crate::domain::sexpr::Path::root_child(0))? .view(); @@ -66,7 +67,7 @@ pub(super) fn count_references_in_expanded_expression( source: &str, reference_counts: &mut BTreeMap, ) -> Result<()> { - let expression_tree = parse_single_expression_tree(source)?; + let expression_tree = parse_single_expression_tree(dialect, source)?; let expression = expression_tree .select_path(&crate::domain::sexpr::Path::root_child(0))? .view(); @@ -82,9 +83,12 @@ pub(super) fn count_references_in_expanded_expression( } pub(super) fn parse_single_expression_tree( + dialect: Dialect, source: &str, ) -> Result { - Ok(crate::domain::sexpr::SyntaxTree::parse(source)?) + Ok(crate::domain::sexpr::SyntaxTree::parse_with_dialect( + source, dialect, + )?) } pub(super) fn expand_unquote_splicing( @@ -101,8 +105,8 @@ pub(super) fn expand_unquote_splicing( argument_bindings, reference_counts, )?; - let expanded_tree = - crate::domain::sexpr::SyntaxTree::parse(&expanded).context("invalid ,@ expansion")?; + let expanded_tree = crate::domain::sexpr::SyntaxTree::parse_with_dialect(&expanded, dialect) + .context("invalid ,@ expansion")?; let expression = expanded_tree .select_path(&crate::domain::sexpr::Path::root_child(0))? .view(); diff --git a/src/domain/inline_function/macro_expansion/mod.rs b/src/domain/inline_function/macro_expansion/mod.rs index 4ffeb008..22ce253f 100644 --- a/src/domain/inline_function/macro_expansion/mod.rs +++ b/src/domain/inline_function/macro_expansion/mod.rs @@ -101,7 +101,7 @@ fn expand_plain_macro_body( )?; count_references_in_expanded_expression(dialect, &intermediate, reference_counts)?; - let intermediate_tree = expansion::parse_single_expression_tree(&intermediate)?; + let intermediate_tree = expansion::parse_single_expression_tree(dialect, &intermediate)?; let intermediate_expression = intermediate_tree .select_path(&crate::domain::sexpr::Path::root_child(0))? .view(); diff --git a/src/domain/inline_function/substitution.rs b/src/domain/inline_function/substitution.rs index 8c711ed1..e8b6a092 100644 --- a/src/domain/inline_function/substitution.rs +++ b/src/domain/inline_function/substitution.rs @@ -33,7 +33,7 @@ pub(super) fn substitute_expression( params: &[String], args: &[String], ) -> Result { - let tree = SyntaxTree::parse(input)?; + let tree = SyntaxTree::parse_with_dialect(input, dialect)?; if tree.root_children().len() != 1 { anyhow::bail!("inline-function default value must be a single S-expression"); } diff --git a/src/domain/inline_function/tests/dialect.rs b/src/domain/inline_function/tests/dialect.rs new file mode 100644 index 00000000..43f46373 --- /dev/null +++ b/src/domain/inline_function/tests/dialect.rs @@ -0,0 +1,95 @@ +use super::super::calls::inline_function_symbol_reference_eq; +use super::*; + +#[test] +fn common_lisp_reader_collision_output_reparses_with_the_same_dialect() { + let input = "(defun identity (value) value)\n(print (identity #\\)))"; + let plan = inline_plan(input); + + assert_eq!( + plan.rewritten, + "(defun identity (value) value)\n(print #\\))" + ); + SyntaxTree::parse_with_dialect(&plan.rewritten, plan.dialect).expect("same-dialect reparse"); +} + +#[test] +fn unsupported_dialects_fail_before_malformed_input_is_parsed() { + for dialect in [ + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ] { + let error = plan_inline_function(InlineFunctionRequest { + input: "(", + dialect, + definition_path: path("0"), + call_paths: vec![path("1")], + all_calls: false, + remove_definition: false, + allow_duplicate_evaluation: false, + allow_drop_arguments: false, + }) + .expect_err("unsupported dialect"); + + assert_eq!( + error.to_string(), + format!( + "inline-function does not support dialect {}", + dialect.label() + ) + ); + } +} + +#[test] +fn emacs_lisp_macro_definitions_remain_unsupported() { + let error = plan_inline_function(InlineFunctionRequest { + input: "(defmacro helper (x) `(+ ,x 1))\n(helper 2)", + dialect: Dialect::EmacsLisp, + definition_path: path("0"), + call_paths: vec![path("1")], + all_calls: false, + remove_definition: false, + allow_duplicate_evaluation: false, + allow_drop_arguments: false, + }) + .expect_err("Emacs Lisp macro definitions stay outside inline-function support"); + + assert_eq!( + error.to_string(), + "inline-function does not support definition head: defmacro" + ); +} + +#[test] +fn function_reference_equality_is_dialect_aware_and_unknown_fails_closed() { + assert!(inline_function_symbol_reference_eq( + Dialect::CommonLisp, + "IDENTITY", + "identity" + )); + + for dialect in [ + Dialect::EmacsLisp, + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + ] { + assert!(inline_function_symbol_reference_eq( + dialect, "identity", "identity" + )); + assert!(!inline_function_symbol_reference_eq( + dialect, "IDENTITY", "identity" + )); + } + + assert!(!inline_function_symbol_reference_eq( + Dialect::Unknown, + "identity", + "identity" + )); +} diff --git a/src/domain/inline_function/tests/mod.rs b/src/domain/inline_function/tests/mod.rs index 0b5fa920..60845976 100644 --- a/src/domain/inline_function/tests/mod.rs +++ b/src/domain/inline_function/tests/mod.rs @@ -3,6 +3,7 @@ use crate::domain::sexpr::{Path, SyntaxTree}; use super::*; mod common_lisp; +mod dialect; mod property; fn path(value: &str) -> Path { diff --git a/src/domain/inline_function/tests/property.rs b/src/domain/inline_function/tests/property.rs index 23fc49b6..da84cd55 100644 --- a/src/domain/inline_function/tests/property.rs +++ b/src/domain/inline_function/tests/property.rs @@ -19,7 +19,7 @@ proptest! { &plan.rewritten, &format!("(defun {name} ({param}) (+ {param} {addend}))\n(print (+ {argument} {addend}))") ); - SyntaxTree::parse(&plan.rewritten) + SyntaxTree::parse_with_dialect(&plan.rewritten, plan.dialect) .map_err(|error| TestCaseError::fail(error.to_string()))?; } @@ -40,7 +40,7 @@ proptest! { &plan.rewritten, &format!("(defun {name} ({param}) (+ {param} {param}))\n(print (+ ({callee}) ({callee})))") ); - SyntaxTree::parse(&plan.rewritten) + SyntaxTree::parse_with_dialect(&plan.rewritten, plan.dialect) .map_err(|error| TestCaseError::fail(error.to_string()))?; } } diff --git a/src/domain/inline_lambda.rs b/src/domain/inline_lambda.rs index 29563be3..39fcf2d0 100644 --- a/src/domain/inline_lambda.rs +++ b/src/domain/inline_lambda.rs @@ -32,10 +32,9 @@ pub(crate) struct Plan { } pub(crate) fn plan(request: Request<'_>) -> Result { - if request.dialect != Dialect::CommonLisp { - bail!("inline-lambda currently supports only Common Lisp"); - } - let tree = SyntaxTree::parse(request.input).context("inline-lambda input is not valid")?; + validate_dialect(request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("inline-lambda input is not valid")?; let call = tree.select_path(&request.path)?.view(); if tree.has_comment_in(call.span) { bail!("inline-lambda cannot replace a call containing comments"); @@ -88,7 +87,8 @@ pub(crate) fn plan(request: Request<'_>) -> Result { .join(" "); let replacement = format!("(let ({rendered}) {})", body.span.slice(request.input)); let rewritten = replace_span(request.input, call.span, &replacement); - SyntaxTree::parse(&rewritten).context("inline-lambda output is not valid")?; + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) + .context("inline-lambda output is not valid")?; Ok(Plan { dialect: request.dialect, path: request.path, @@ -101,6 +101,13 @@ pub(crate) fn plan(request: Request<'_>) -> Result { }) } +pub(crate) fn validate_dialect(dialect: Dialect) -> Result<()> { + if dialect != Dialect::CommonLisp { + bail!("inline-lambda currently supports only Common Lisp"); + } + Ok(()) +} + fn plain_symbol(view: &ExpressionView, role: &str) -> Result { if view.kind != ExpressionKind::Atom || !view.reader_prefixes.is_empty() { bail!("inline-lambda requires a plain {role}"); @@ -187,4 +194,39 @@ mod tests { .is_err() ); } + + #[test] + fn dialect_support_matrix_is_enforced_before_parsing_and_reparses_output() { + let result = plan(Request { + input: "#\\) ((lambda (x) x) 1)", + dialect: Dialect::CommonLisp, + path: "1".parse().expect("path"), + }) + .expect("Common Lisp"); + assert!(result.rewritten.starts_with("#\\)")); + SyntaxTree::parse_with_dialect(&result.rewritten, Dialect::CommonLisp) + .expect("Common Lisp output"); + + for dialect in [ + Dialect::EmacsLisp, + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ] { + let error = plan(Request { + input: ")", + dialect, + path: "0".parse().expect("path"), + }) + .expect_err("unsupported dialect"); + assert!( + error + .to_string() + .contains("currently supports only Common Lisp"), + "{dialect:?}: {error:#}" + ); + } + } } diff --git a/src/domain/inline_let.rs b/src/domain/inline_let.rs index 022c4f9b..e8370dff 100644 --- a/src/domain/inline_let.rs +++ b/src/domain/inline_let.rs @@ -31,18 +31,42 @@ pub struct InlineLetPlan { pub changed: bool, } +pub(crate) const fn supports_inline_let_dialect(dialect: Dialect) -> bool { + matches!( + dialect, + Dialect::CommonLisp + | Dialect::EmacsLisp + | Dialect::Scheme + | Dialect::Clojure + | Dialect::Janet + | Dialect::Fennel + ) +} + +fn require_supported_dialect(dialect: Dialect) -> Result<()> { + if supports_inline_let_dialect(dialect) { + return Ok(()); + } + + anyhow::bail!( + "inline-let requires a known dialect because semantic safety cannot be verified for unknown input" + ) +} + pub fn plan_inline_let(request: InlineLetRequest<'_>) -> Result { - let input_tree = SyntaxTree::parse(request.input) + require_supported_dialect(request.dialect)?; + let input_tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) .context("inline-let input is not a valid S-expression document")?; crate::domain::mutation_safety::reject_common_lisp_reader_conditionals( &input_tree, request.dialect, )?; + let target = select_target_from_tree(&input_tree, request.path.as_ref(), &request.target)?; let plan = plan(CoreRequest { input: request.input, dialect: request.dialect, path: request.path, - target: request.target, + target, allow_duplicate_evaluation: request.allow_duplicate_evaluation, })?; Ok(InlineLetPlan { @@ -82,8 +106,28 @@ pub(crate) struct CorePlan { pub changed: bool, } +fn select_target_from_tree( + tree: &SyntaxTree, + path: Option<&Path>, + requested_target: &ExpressionView, +) -> Result { + let selected = match path { + Some(path) => tree.select_path(path)?, + None => tree.select_at(requested_target.span.start().get())?, + }; + let target = selected.view(); + if target.span != requested_target.span { + anyhow::bail!("inline-let target does not match the dialect-aware input tree"); + } + Ok(target) +} + pub(crate) fn plan(request: CoreRequest<'_>) -> Result { - let parts = parts(request.dialect, request.input, &request.target)?; + require_supported_dialect(request.dialect)?; + let input_tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("inline-let input is not a valid S-expression document")?; + let target = select_target_from_tree(&input_tree, request.path.as_ref(), &request.target)?; + let parts = parts(request.dialect, request.input, &target)?; let reference_count = parts.reference_spans.len(); if reference_count == 0 { anyhow::bail!("inline-let would drop an unused binding value"); @@ -101,7 +145,7 @@ pub(crate) fn plan(request: CoreRequest<'_>) -> Result { &parts.binding_value, ); let rewritten = replace_span(request.input, parts.let_span, &replacement); - SyntaxTree::parse(&rewritten) + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) .context("inline-let output is not a valid S-expression document")?; Ok(CorePlan { @@ -262,30 +306,214 @@ mod tests { use super::*; fn plan(input: &str, dialect: Dialect, allow_duplicate_evaluation: bool) -> Result { - let tree = SyntaxTree::parse(input)?; - let target = tree.select_path(&Path::from_indexes(vec![0]))?.view(); + let path = Path::from_indexes(vec![0]); + let tree = SyntaxTree::parse_with_dialect(input, dialect)?; + let target = tree.select_path(&path)?.view(); super::plan(CoreRequest { input, dialect, - path: Some(Path::from_indexes(vec![0])), + path: Some(path), target, allow_duplicate_evaluation, }) } #[test] - fn supports_list_and_vector_binding_dialects() { - let common_lisp = plan("(let ((x 1)) (+ x y))", Dialect::CommonLisp, false).unwrap(); - assert_eq!(common_lisp.rewritten, "(+ 1 y)"); - let clojure = plan("(let [x 1] (+ x y))", Dialect::Clojure, false).unwrap(); - assert_eq!(clojure.rewritten, "(+ 1 y)"); + fn all_known_dialects_rewrite_and_reparse_with_the_same_dialect() { + for (dialect, input, expected) in [ + ( + Dialect::CommonLisp, + "(let ((x 1)) (+ x y)) #\\)", + "(+ 1 y) #\\)", + ), + ( + Dialect::EmacsLisp, + r"(let ((x 1)) (+ x y)) ?\)", + r"(+ 1 y) ?\)", + ), + (Dialect::Scheme, "(let ((x 1)) (+ x y))", "(+ 1 y)"), + (Dialect::Clojure, "(let [x 1] (+ x y))", "(+ 1 y)"), + (Dialect::Janet, "(let [x 1] (+ x y))", "(+ 1 y)"), + (Dialect::Fennel, "(let [x 1] (+ x y))", "(+ 1 y)"), + ] { + assert!(supports_inline_let_dialect(dialect)); + let core = plan(input, dialect, false).unwrap(); + assert_eq!(core.rewritten, expected); + SyntaxTree::parse_with_dialect(&core.rewritten, dialect).unwrap(); + + let path = Path::from_indexes(vec![0]); + let tree = SyntaxTree::parse_with_dialect(input, dialect).unwrap(); + let target = tree.select_path(&path).unwrap().view(); + let public = plan_inline_let(InlineLetRequest { + input, + dialect, + path: Some(path), + target, + allow_duplicate_evaluation: false, + }) + .unwrap(); + assert_eq!(public.rewritten, core.rewritten); + SyntaxTree::parse_with_dialect(&public.rewritten, dialect).unwrap(); + } + } + + #[test] + fn support_predicate_rejects_only_unknown() { + assert!(!supports_inline_let_dialect(Dialect::Unknown)); + } + + #[test] + fn unknown_fails_closed_before_parsing_or_target_traversal() { + let valid_input = "(let ((x 1)) x)"; + let path = Path::from_indexes(vec![0]); + let tree = SyntaxTree::parse_with_dialect(valid_input, Dialect::CommonLisp).unwrap(); + let target = tree.select_path(&path).unwrap().view(); + + let expected = "inline-let requires a known dialect because semantic safety cannot be verified for unknown input"; + let core_error = super::plan(CoreRequest { + input: ")", + dialect: Dialect::Unknown, + path: Some(path.clone()), + target: target.clone(), + allow_duplicate_evaluation: false, + }) + .unwrap_err(); + assert_eq!(core_error.to_string(), expected); + + let public_error = plan_inline_let(InlineLetRequest { + input: ")", + dialect: Dialect::Unknown, + path: Some(path), + target, + allow_duplicate_evaluation: false, + }) + .unwrap_err(); + assert_eq!(public_error.to_string(), expected); + } + + #[test] + fn all_known_dialects_reject_unused_or_duplicated_evaluation() { + for (dialect, unused, duplicate) in [ + ( + Dialect::CommonLisp, + "(let ((x (effect))) y)", + "(let ((x (effect))) (+ x x))", + ), + ( + Dialect::EmacsLisp, + "(let ((x (effect))) y)", + "(let ((x (effect))) (+ x x))", + ), + ( + Dialect::Scheme, + "(let ((x (effect))) y)", + "(let ((x (effect))) (+ x x))", + ), + ( + Dialect::Clojure, + "(let [x (effect)] y)", + "(let [x (effect)] (+ x x))", + ), + ( + Dialect::Janet, + "(let [x (effect)] y)", + "(let [x (effect)] (+ x x))", + ), + ( + Dialect::Fennel, + "(let [x (effect)] y)", + "(let [x (effect)] (+ x x))", + ), + ] { + assert!(plan(unused, dialect, false).is_err()); + assert!(plan(duplicate, dialect, false).is_err()); + assert!(plan(duplicate, dialect, true).is_ok()); + } } #[test] - fn rejects_unused_duplicate_and_capture_cases() { - assert!(plan("(let ((x (effect))) y)", Dialect::CommonLisp, false).is_err()); - assert!(plan("(let ((x (effect))) (+ x x))", Dialect::CommonLisp, false).is_err()); - assert!(plan("(let ((x (effect))) (+ x x))", Dialect::CommonLisp, true).is_ok()); - assert!(plan("(let ((x y)) (let ((y 2)) x))", Dialect::CommonLisp, false).is_err()); + fn callable_bindings_are_not_rewritten_when_they_shadow_the_let_binding() { + for (dialect, input, expected) in [ + ( + Dialect::CommonLisp, + "(let ((x 1)) (list x (lambda (x) x)))", + "(list 1 (lambda (x) x))", + ), + ( + Dialect::EmacsLisp, + "(let ((x 1)) (list x (lambda (x) x)))", + "(list 1 (lambda (x) x))", + ), + ( + Dialect::Scheme, + "(let ((x 1)) (list x (lambda x x)))", + "(list 1 (lambda x x))", + ), + ( + Dialect::Clojure, + "(let [x 1] (list x (fn x [value] x)))", + "(list 1 (fn x [value] x))", + ), + ( + Dialect::Clojure, + "(let [x 1] (list x (fn ([x] x) ([x y] x))))", + "(list 1 (fn ([x] x) ([x y] x)))", + ), + ( + Dialect::Janet, + "(let [x 1] (list x (fn [x] x)))", + "(list 1 (fn [x] x))", + ), + ( + Dialect::Fennel, + "(let [x 1] (list x (fn [x] x)))", + "(list 1 (fn [x] x))", + ), + ] { + let shadowed = plan(input, dialect, false).unwrap(); + assert_eq!(shadowed.reference_count, 1); + assert_eq!(shadowed.rewritten, expected); + SyntaxTree::parse_with_dialect(&shadowed.rewritten, dialect).unwrap(); + } + } + + #[test] + fn rejects_inline_let_when_callable_bindings_would_capture_the_value() { + for (dialect, input) in [ + ( + Dialect::CommonLisp, + "(let ((target external)) (lambda (external) target))", + ), + ( + Dialect::EmacsLisp, + "(let ((target external)) (lambda (external) target))", + ), + ( + Dialect::Scheme, + "(let ((target external)) (lambda external target))", + ), + ( + Dialect::Clojure, + "(let [target recur-name] (fn recur-name ([value] target)))", + ), + ( + Dialect::Clojure, + "(let [target external] (fn ([external] target) ([other external] target)))", + ), + ( + Dialect::Janet, + "(let [target external] (fn [external] target))", + ), + ( + Dialect::Fennel, + "(let [target external] (fn [external] target))", + ), + ] { + let error = plan(input, dialect, false).unwrap_err(); + assert!( + error.to_string().contains("would capture variable"), + "unexpected {dialect:?} error: {error:#}" + ); + } } } diff --git a/src/domain/inline_literal_constant.rs b/src/domain/inline_literal_constant.rs index a59cecab..7ce1a231 100644 --- a/src/domain/inline_literal_constant.rs +++ b/src/domain/inline_literal_constant.rs @@ -36,7 +36,7 @@ pub fn plan_inline_literal_constant( if request.dialect != Dialect::CommonLisp { bail!("inline-literal-constant supports only Common Lisp"); } - let tree = SyntaxTree::parse(request.input) + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) .context("inline-literal-constant input is not a valid S-expression document")?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; let definition = tree.select_path(&request.path)?.view(); @@ -85,7 +85,7 @@ pub fn plan_inline_literal_constant( rewritten.replace_range(span.start().get()..span.end().get(), &replacement); } let rewritten = collapse_removed_definition_gap(&rewritten); - SyntaxTree::parse(&rewritten) + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) .context("inline-literal-constant output is not a valid S-expression document")?; Ok(InlineLiteralConstantPlan { @@ -327,4 +327,37 @@ mod tests { let result = plan("(defconstant +step+ 2) (incf value +step+)").unwrap(); assert_eq!(result.rewritten, " (incf value 2)"); } + + #[test] + fn dialect_support_matrix_is_enforced_before_parsing_and_reparses_output() { + let plan = plan_inline_literal_constant(InlineLiteralConstantRequest { + input: "#\\)\n(defconstant +x+ 1)\n(print +x+)", + dialect: Dialect::CommonLisp, + path: "1".parse().expect("path"), + }) + .expect("Common Lisp"); + assert!(plan.rewritten.starts_with("#\\)")); + SyntaxTree::parse_with_dialect(&plan.rewritten, Dialect::CommonLisp) + .expect("Common Lisp output"); + + for dialect in [ + Dialect::EmacsLisp, + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ] { + let error = plan_inline_literal_constant(InlineLiteralConstantRequest { + input: ")", + dialect, + path: "0".parse().expect("path"), + }) + .expect_err("unsupported dialect"); + assert!( + error.to_string().contains("supports only Common Lisp"), + "{dialect:?}: {error:#}" + ); + } + } } diff --git a/src/domain/inline_local_function.rs b/src/domain/inline_local_function.rs index c242221f..b05c868d 100644 --- a/src/domain/inline_local_function.rs +++ b/src/domain/inline_local_function.rs @@ -38,10 +38,8 @@ pub(crate) struct Plan { } pub(crate) fn plan(request: Request<'_>) -> Result { - if request.dialect != Dialect::CommonLisp { - bail!("inline-local-function currently supports only Common Lisp"); - } - let tree = SyntaxTree::parse(request.input) + validate_dialect(request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) .context("inline-local-function input is not a valid S-expression document")?; let form = tree.select_path(&request.path)?.view(); if tree.has_comment_in(form.span) { @@ -126,7 +124,7 @@ pub(crate) fn plan(request: Request<'_>) -> Result { .join(" "); let replacement = format!("(let ({bindings}) {body})"); let rewritten = replace_span(request.input, form.span, &replacement); - SyntaxTree::parse(&rewritten) + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) .context("inline-local-function output is not a valid S-expression document")?; Ok(Plan { dialect: request.dialect, @@ -141,6 +139,13 @@ pub(crate) fn plan(request: Request<'_>) -> Result { }) } +pub(crate) fn validate_dialect(dialect: Dialect) -> Result<()> { + if dialect != Dialect::CommonLisp { + bail!("inline-local-function currently supports only Common Lisp"); + } + Ok(()) +} + fn replace_span(input: &str, span: ByteSpan, replacement: &str) -> String { let mut output = String::with_capacity(input.len() + replacement.len()); output.push_str(&input[..span.start().get()]); @@ -234,4 +239,59 @@ mod tests { assert!(plan("(flet ((f (x) (go done))) (f 1))").is_err()); assert!(plan("(labels ((f (x) x)) (f 1))").is_err()); } + + #[test] + fn supports_only_common_lisp_before_input_validation() { + let cases = [ + (Dialect::CommonLisp, true), + (Dialect::EmacsLisp, false), + (Dialect::Scheme, false), + (Dialect::Clojure, false), + (Dialect::Janet, false), + (Dialect::Fennel, false), + (Dialect::Unknown, false), + ]; + + for (dialect, supported) in cases { + let input = if supported { + "(flet ((identity (x) x)) (identity value))" + } else { + ")" + }; + let result = super::plan(Request { + input, + dialect, + path: Path::from_indexes(vec![0]), + }); + + if supported { + let plan = result.expect("Common Lisp should be supported"); + assert_eq!(plan.dialect, dialect); + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect) + .expect("rewritten Common Lisp should parse"); + } else { + let error = match result { + Ok(_) => panic!("{dialect:?} should be rejected"), + Err(error) => error, + }; + assert_eq!( + error.to_string(), + "inline-local-function currently supports only Common Lisp" + ); + } + } + } + + #[test] + fn preserves_common_lisp_reader_atoms() { + let input = "(flet ((render (x) (list x #\\) #:done #x2a))) (render (next)))"; + let plan = plan(input).unwrap(); + + assert_eq!( + plan.rewritten, + "(let ((x (next))) (list x #\\) #:done #x2a))" + ); + SyntaxTree::parse_with_dialect(&plan.rewritten, Dialect::CommonLisp) + .expect("rewritten Common Lisp reader atoms should parse"); + } } diff --git a/src/domain/inline_symbol_macro.rs b/src/domain/inline_symbol_macro.rs index e069bba3..520bf97b 100644 --- a/src/domain/inline_symbol_macro.rs +++ b/src/domain/inline_symbol_macro.rs @@ -40,7 +40,7 @@ pub fn plan_inline_symbol_macro( if request.dialect != Dialect::CommonLisp { bail!("inline-symbol-macro currently supports only Common Lisp"); } - let tree = SyntaxTree::parse(request.input) + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) .context("inline-symbol-macro input is not a valid S-expression document")?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; let form = tree.select_path(&request.path)?.view(); @@ -221,4 +221,38 @@ mod tests { let error = plan("(symbol-macrolet ((x 1)) (list x 'x))").expect_err("reject quote"); assert!(error.to_string().contains("reader prefixes")); } + + #[test] + fn dialect_matrix_preserves_common_lisp_gate_precedence() { + let cases = [ + (Dialect::CommonLisp, r"(symbol-macrolet ((x #\))) (list x))"), + (Dialect::EmacsLisp, ")"), + (Dialect::Scheme, ")"), + (Dialect::Clojure, ")"), + (Dialect::Janet, ")"), + (Dialect::Fennel, ")"), + (Dialect::Unknown, ")"), + ]; + + for (dialect, input) in cases { + let result = plan_inline_symbol_macro(InlineSymbolMacroRequest { + input, + dialect, + path: "0".parse().unwrap(), + }); + + if dialect == Dialect::CommonLisp { + let plan = result.unwrap(); + assert_eq!(plan.rewritten, r"(list #\))"); + assert!( + SyntaxTree::parse_with_dialect(&plan.rewritten, Dialect::CommonLisp).is_ok() + ); + } else { + assert_eq!( + result.unwrap_err().to_string(), + "inline-symbol-macro currently supports only Common Lisp" + ); + } + } + } } diff --git a/src/domain/introduce_let.rs b/src/domain/introduce_let.rs index 6a3e9598..a20c99e3 100644 --- a/src/domain/introduce_let.rs +++ b/src/domain/introduce_let.rs @@ -10,6 +10,7 @@ mod types; use anyhow::{Context, Result, bail}; use super::mutation_safety::reject_common_lisp_reader_conditionals; +use crate::domain::dialect::{IntroduceLetOperation, VerifiedSemanticPolicy}; use crate::domain::sexpr::{ByteOffset, ByteSpan, Path, SyntaxTree}; use occurrences::{ @@ -20,15 +21,19 @@ use rewrite::{introduced_let, replace_span, replace_spans_within_span}; pub use types::{IntroduceLetPlan, IntroduceLetRequest}; pub fn plan_introduce_let(request: IntroduceLetRequest<'_>) -> Result { - let input_tree = SyntaxTree::parse(request.input) + let semantic = request + .dialect + .verify_introduce_let() + .context("introduce-let is not supported for this dialect")?; + let input_tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) .context("introduce-let input is not a valid S-expression document")?; reject_common_lisp_reader_conditionals(&input_tree, request.dialect)?; let selected_span = request.target.span; let binding_value = selected_span.slice(request.input).to_owned(); let enclosing = request.enclosing_span.slice(request.input); - let enclosing_tree = - SyntaxTree::parse(enclosing).context("failed to parse enclosing list for introduce-let")?; + let enclosing_tree = SyntaxTree::parse_with_dialect(enclosing, request.dialect) + .context("failed to parse enclosing list for introduce-let")?; let enclosing_view = enclosing_tree.select_path(&Path::root_child(0))?.view(); let selected_relative_span = ByteSpan::new( @@ -36,7 +41,7 @@ pub fn plan_introduce_let(request: IntroduceLetRequest<'_>) -> Result) -> Result) -> Result) -> Result) -> Result) -> Result { +fn selected_path_shadowed_by_binding( + semantic: VerifiedSemanticPolicy, + request: &IntroduceLetRequest<'_>, +) -> Result { let Some(path) = &request.path else { return Ok(false); }; @@ -125,14 +133,14 @@ fn selected_path_shadowed_by_binding(request: &IntroduceLetRequest<'_>) -> Resul return Ok(false); }; - let tree = - SyntaxTree::parse(request.input).context("failed to parse document for introduce-let")?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("failed to parse document for introduce-let")?; let top_level_view = tree .select_path(&Path::root_child(top_level_index.get()))? .view(); Ok(is_path_shadowed_by_binding( - request.dialect, + semantic, &top_level_view, relative_path, request.name.as_str(), diff --git a/src/domain/introduce_let/occurrences/iteration.rs b/src/domain/introduce_let/occurrences/iteration.rs index 2fdde518..90bba32f 100644 --- a/src/domain/introduce_let/occurrences/iteration.rs +++ b/src/domain/introduce_let/occurrences/iteration.rs @@ -1,5 +1,5 @@ use crate::domain::{ - dialect::Dialect, + dialect::{IntroduceLetOperation, VerifiedSemanticPolicy}, sexpr::{ByteSpan, ChildIndex, ExpressionView}, }; @@ -10,7 +10,7 @@ use super::{ use crate::domain::introduce_let::syntax::iteration_binding_child_shadowed; pub(super) fn collect_iteration_binding_spans( - dialect: Dialect, + semantic: VerifiedSemanticPolicy, binding_form: &ExpressionView, target: &ExpressionView, binding_name: &str, @@ -26,7 +26,7 @@ pub(super) fn collect_iteration_binding_spans( let child_shadowed = shadowed_by_binding || iteration_binding_child_shadowed(binding_form, binding_name, index); collect_equivalent_expression_spans( - dialect, + semantic, child, target, binding_name, @@ -37,7 +37,7 @@ pub(super) fn collect_iteration_binding_spans( } pub(super) fn is_span_shadowed_by_iteration_bindings( - dialect: Dialect, + semantic: VerifiedSemanticPolicy, binding_form: &ExpressionView, target_span: ByteSpan, binding_name: &str, @@ -54,12 +54,12 @@ pub(super) fn is_span_shadowed_by_iteration_bindings( .any(|(index, child)| { let child_shadowed = shadowed_by_binding || iteration_binding_child_shadowed(binding_form, binding_name, index); - is_span_shadowed_by_binding(dialect, child, target_span, binding_name, child_shadowed) + is_span_shadowed_by_binding(semantic, child, target_span, binding_name, child_shadowed) }) } pub(super) fn is_path_shadowed_by_iteration_bindings( - dialect: Dialect, + semantic: VerifiedSemanticPolicy, binding_form: &ExpressionView, target_path: &[ChildIndex], binding_name: &str, @@ -78,6 +78,6 @@ pub(super) fn is_path_shadowed_by_iteration_bindings( if rest.is_empty() { child_shadowed } else { - is_path_shadowed_by_binding(dialect, child, rest, binding_name, child_shadowed) + is_path_shadowed_by_binding(semantic, child, rest, binding_name, child_shadowed) } } diff --git a/src/domain/introduce_let/occurrences/let_star.rs b/src/domain/introduce_let/occurrences/let_star.rs index 0640841e..1a15fc0d 100644 --- a/src/domain/introduce_let/occurrences/let_star.rs +++ b/src/domain/introduce_let/occurrences/let_star.rs @@ -1,5 +1,5 @@ use crate::domain::{ - dialect::Dialect, + dialect::{IntroduceLetOperation, VerifiedSemanticPolicy}, sexpr::{ByteSpan, ChildIndex, ExpressionView}, }; @@ -10,7 +10,7 @@ use super::{ use crate::domain::introduce_let::syntax::binding_pair_binds_name; pub(super) fn collect_let_star_binding_spans( - dialect: Dialect, + semantic: VerifiedSemanticPolicy, binding_form: &ExpressionView, target: &ExpressionView, binding_name: &str, @@ -25,7 +25,7 @@ pub(super) fn collect_let_star_binding_spans( let mut sequential_shadowed = shadowed_by_binding; for binding in &binding_form.children { collect_let_star_binding_spec_spans( - dialect, + semantic, binding, target, binding_name, @@ -39,7 +39,7 @@ pub(super) fn collect_let_star_binding_spans( } fn collect_let_star_binding_spec_spans( - dialect: Dialect, + semantic: VerifiedSemanticPolicy, binding: &ExpressionView, target: &ExpressionView, binding_name: &str, @@ -53,7 +53,7 @@ fn collect_let_star_binding_spec_spans( for child in &binding.children { collect_equivalent_expression_spans( - dialect, + semantic, child, target, binding_name, @@ -64,7 +64,7 @@ fn collect_let_star_binding_spec_spans( } pub(super) fn is_span_shadowed_by_let_star_bindings( - dialect: Dialect, + semantic: VerifiedSemanticPolicy, binding_form: &ExpressionView, target_span: ByteSpan, binding_name: &str, @@ -78,7 +78,7 @@ pub(super) fn is_span_shadowed_by_let_star_bindings( for binding in &binding_form.children { if binding.span.contains_span(target_span) { return is_span_shadowed_by_binding( - dialect, + semantic, binding, target_span, binding_name, @@ -94,7 +94,7 @@ pub(super) fn is_span_shadowed_by_let_star_bindings( } pub(super) fn is_path_shadowed_by_let_star_bindings( - dialect: Dialect, + semantic: VerifiedSemanticPolicy, binding_form: &ExpressionView, target_path: &[ChildIndex], binding_name: &str, @@ -112,7 +112,7 @@ pub(super) fn is_path_shadowed_by_let_star_bindings( sequential_shadowed } else { is_path_shadowed_by_binding( - dialect, + semantic, binding, rest, binding_name, diff --git a/src/domain/introduce_let/occurrences/local_callable.rs b/src/domain/introduce_let/occurrences/local_callable.rs index 88a783f8..91b94780 100644 --- a/src/domain/introduce_let/occurrences/local_callable.rs +++ b/src/domain/introduce_let/occurrences/local_callable.rs @@ -1,5 +1,5 @@ use crate::domain::{ - dialect::Dialect, + dialect::{IntroduceLetOperation, VerifiedSemanticPolicy}, sexpr::{ByteSpan, ChildIndex, ExpressionView}, }; @@ -10,7 +10,7 @@ use super::{ use crate::domain::introduce_let::syntax::local_callable_binding_child_shadowed; pub(super) fn collect_local_callable_binding_spans( - dialect: Dialect, + semantic: VerifiedSemanticPolicy, binding_form: &ExpressionView, target: &ExpressionView, binding_name: &str, @@ -24,7 +24,7 @@ pub(super) fn collect_local_callable_binding_spans( for binding in &binding_form.children { collect_local_callable_binding_spec_spans( - dialect, + semantic, binding, target, binding_name, @@ -35,7 +35,7 @@ pub(super) fn collect_local_callable_binding_spans( } fn collect_local_callable_binding_spec_spans( - dialect: Dialect, + semantic: VerifiedSemanticPolicy, binding: &ExpressionView, target: &ExpressionView, binding_name: &str, @@ -51,7 +51,7 @@ fn collect_local_callable_binding_spec_spans( let child_shadowed = shadowed_by_binding || local_callable_binding_child_shadowed(binding, binding_name, index); collect_equivalent_expression_spans( - dialect, + semantic, child, target, binding_name, @@ -62,7 +62,7 @@ fn collect_local_callable_binding_spec_spans( } pub(super) fn is_span_shadowed_by_local_callable_binding( - dialect: Dialect, + semantic: VerifiedSemanticPolicy, binding_form: &ExpressionView, target_span: ByteSpan, binding_name: &str, @@ -74,7 +74,7 @@ pub(super) fn is_span_shadowed_by_local_callable_binding( binding_form.children.iter().any(|binding| { is_span_shadowed_by_local_callable_binding_spec( - dialect, + semantic, binding, target_span, binding_name, @@ -84,7 +84,7 @@ pub(super) fn is_span_shadowed_by_local_callable_binding( } pub(super) fn is_path_shadowed_by_local_callable_binding( - dialect: Dialect, + semantic: VerifiedSemanticPolicy, binding_form: &ExpressionView, target_path: &[ChildIndex], binding_name: &str, @@ -102,7 +102,7 @@ pub(super) fn is_path_shadowed_by_local_callable_binding( shadowed_by_binding } else { is_path_shadowed_by_local_callable_binding_spec( - dialect, + semantic, binding, rest, binding_name, @@ -112,7 +112,7 @@ pub(super) fn is_path_shadowed_by_local_callable_binding( } fn is_span_shadowed_by_local_callable_binding_spec( - dialect: Dialect, + semantic: VerifiedSemanticPolicy, binding: &ExpressionView, target_span: ByteSpan, binding_name: &str, @@ -125,12 +125,12 @@ fn is_span_shadowed_by_local_callable_binding_spec( binding.children.iter().enumerate().any(|(index, child)| { let child_shadowed = shadowed_by_binding || local_callable_binding_child_shadowed(binding, binding_name, index); - is_span_shadowed_by_binding(dialect, child, target_span, binding_name, child_shadowed) + is_span_shadowed_by_binding(semantic, child, target_span, binding_name, child_shadowed) }) } fn is_path_shadowed_by_local_callable_binding_spec( - dialect: Dialect, + semantic: VerifiedSemanticPolicy, binding: &ExpressionView, target_path: &[ChildIndex], binding_name: &str, @@ -149,6 +149,6 @@ fn is_path_shadowed_by_local_callable_binding_spec( if rest.is_empty() { child_shadowed } else { - is_path_shadowed_by_binding(dialect, child, rest, binding_name, child_shadowed) + is_path_shadowed_by_binding(semantic, child, rest, binding_name, child_shadowed) } } diff --git a/src/domain/introduce_let/occurrences/mod.rs b/src/domain/introduce_let/occurrences/mod.rs index 6e819e2f..68f75a19 100644 --- a/src/domain/introduce_let/occurrences/mod.rs +++ b/src/domain/introduce_let/occurrences/mod.rs @@ -4,7 +4,7 @@ mod local_callable; mod variable; use crate::domain::{ - dialect::Dialect, + dialect::{IntroduceLetOperation, VerifiedSemanticPolicy}, sexpr::{ByteOffset, ByteSpan, ChildIndex, ExpressionView}, }; @@ -37,13 +37,14 @@ pub(super) struct EquivalentExpressionSpans { } pub(super) fn collect_equivalent_expression_spans( - dialect: Dialect, + semantic: VerifiedSemanticPolicy, view: &ExpressionView, target: &ExpressionView, binding_name: &str, shadowed_by_binding: bool, output: &mut EquivalentExpressionSpans, ) { + let dialect = semantic.dialect(); if expressions_equivalent(view, target) { record_equivalent_span(output, view.span, shadowed_by_binding); return; @@ -52,7 +53,7 @@ pub(super) fn collect_equivalent_expression_spans( for (index, child) in view.children.iter().enumerate() { if let_star_bindings_child_index(dialect, view) == Some(index) { collect_let_star_binding_spans( - dialect, + semantic, child, target, binding_name, @@ -64,7 +65,7 @@ pub(super) fn collect_equivalent_expression_spans( if iteration_bindings_child_index(dialect, view) == Some(index) { collect_iteration_binding_spans( - dialect, + semantic, child, target, binding_name, @@ -76,7 +77,7 @@ pub(super) fn collect_equivalent_expression_spans( if variable_bindings_child_index(dialect, view) == Some(index) { let mut ctx = VariableBindingContext { - dialect, + semantic, target, binding_name, has_step_forms: variable_binding_form_has_step_forms(dialect, view), @@ -93,7 +94,7 @@ pub(super) fn collect_equivalent_expression_spans( if local_callable_bindings_child_index(dialect, view) == Some(index) { collect_local_callable_binding_spans( - dialect, + semantic, child, target, binding_name, @@ -104,9 +105,9 @@ pub(super) fn collect_equivalent_expression_spans( } let child_shadowed = - shadowed_by_binding || child_shadowed_by_binding(dialect, view, binding_name, index); + shadowed_by_binding || child_shadowed_by_binding(semantic, view, binding_name, index); collect_equivalent_expression_spans( - dialect, + semantic, child, target, binding_name, @@ -117,12 +118,13 @@ pub(super) fn collect_equivalent_expression_spans( } pub(super) fn is_span_shadowed_by_binding( - dialect: Dialect, + semantic: VerifiedSemanticPolicy, view: &ExpressionView, target_span: ByteSpan, binding_name: &str, shadowed_by_binding: bool, ) -> bool { + let dialect = semantic.dialect(); if view.span == target_span { return shadowed_by_binding; } @@ -130,7 +132,7 @@ pub(super) fn is_span_shadowed_by_binding( view.children.iter().enumerate().any(|(index, child)| { if let_star_bindings_child_index(dialect, view) == Some(index) { return is_span_shadowed_by_let_star_bindings( - dialect, + semantic, child, target_span, binding_name, @@ -140,7 +142,7 @@ pub(super) fn is_span_shadowed_by_binding( if iteration_bindings_child_index(dialect, view) == Some(index) { return is_span_shadowed_by_iteration_bindings( - dialect, + semantic, child, target_span, binding_name, @@ -150,7 +152,7 @@ pub(super) fn is_span_shadowed_by_binding( if variable_bindings_child_index(dialect, view) == Some(index) { return is_span_shadowed_by_variable_bindings( - dialect, + semantic, child, target_span, binding_name, @@ -162,7 +164,7 @@ pub(super) fn is_span_shadowed_by_binding( if local_callable_bindings_child_index(dialect, view) == Some(index) { return is_span_shadowed_by_local_callable_binding( - dialect, + semantic, child, target_span, binding_name, @@ -171,18 +173,19 @@ pub(super) fn is_span_shadowed_by_binding( } let child_shadowed = - shadowed_by_binding || child_shadowed_by_binding(dialect, view, binding_name, index); - is_span_shadowed_by_binding(dialect, child, target_span, binding_name, child_shadowed) + shadowed_by_binding || child_shadowed_by_binding(semantic, view, binding_name, index); + is_span_shadowed_by_binding(semantic, child, target_span, binding_name, child_shadowed) }) } pub(super) fn is_path_shadowed_by_binding( - dialect: Dialect, + semantic: VerifiedSemanticPolicy, view: &ExpressionView, target_path: &[ChildIndex], binding_name: &str, shadowed_by_binding: bool, ) -> bool { + let dialect = semantic.dialect(); let Some((index, rest)) = target_path.split_first() else { return shadowed_by_binding; }; @@ -193,7 +196,7 @@ pub(super) fn is_path_shadowed_by_binding( if let_star_bindings_child_index(dialect, view) == Some(index) { return is_path_shadowed_by_let_star_bindings( - dialect, + semantic, child, rest, binding_name, @@ -203,7 +206,7 @@ pub(super) fn is_path_shadowed_by_binding( if iteration_bindings_child_index(dialect, view) == Some(index) { return is_path_shadowed_by_iteration_bindings( - dialect, + semantic, child, rest, binding_name, @@ -213,7 +216,7 @@ pub(super) fn is_path_shadowed_by_binding( if variable_bindings_child_index(dialect, view) == Some(index) { return is_path_shadowed_by_variable_bindings( - dialect, + semantic, child, rest, binding_name, @@ -225,7 +228,7 @@ pub(super) fn is_path_shadowed_by_binding( if local_callable_bindings_child_index(dialect, view) == Some(index) { return is_path_shadowed_by_local_callable_binding( - dialect, + semantic, child, rest, binding_name, @@ -234,11 +237,11 @@ pub(super) fn is_path_shadowed_by_binding( } let child_shadowed = - shadowed_by_binding || child_shadowed_by_binding(dialect, view, binding_name, index); + shadowed_by_binding || child_shadowed_by_binding(semantic, view, binding_name, index); if rest.is_empty() { child_shadowed } else { - is_path_shadowed_by_binding(dialect, child, rest, binding_name, child_shadowed) + is_path_shadowed_by_binding(semantic, child, rest, binding_name, child_shadowed) } } diff --git a/src/domain/introduce_let/occurrences/variable.rs b/src/domain/introduce_let/occurrences/variable.rs index d8480c5a..8a2a96f1 100644 --- a/src/domain/introduce_let/occurrences/variable.rs +++ b/src/domain/introduce_let/occurrences/variable.rs @@ -1,5 +1,5 @@ use crate::domain::{ - dialect::Dialect, + dialect::{IntroduceLetOperation, VerifiedSemanticPolicy}, sexpr::{ByteSpan, ChildIndex, ExpressionView}, }; @@ -10,7 +10,7 @@ use super::{ use crate::domain::introduce_let::syntax::variable_spec_binds_name; pub(super) struct VariableBindingContext<'a> { - pub(super) dialect: Dialect, + pub(super) semantic: VerifiedSemanticPolicy, pub(super) target: &'a ExpressionView, pub(super) binding_name: &'a str, pub(super) has_step_forms: bool, @@ -66,12 +66,12 @@ fn collect_variable_binding_spec_spans( step_shadowed, ctx.has_step_forms, ); - let dialect = ctx.dialect; + let semantic = ctx.semantic; let target = ctx.target; let binding_name = ctx.binding_name; let output = &mut *ctx.output; collect_equivalent_expression_spans( - dialect, + semantic, child, target, binding_name, @@ -82,7 +82,7 @@ fn collect_variable_binding_spec_spans( } pub(super) fn is_span_shadowed_by_variable_bindings( - dialect: Dialect, + semantic: VerifiedSemanticPolicy, binding_form: &ExpressionView, target_span: ByteSpan, binding_name: &str, @@ -104,7 +104,7 @@ pub(super) fn is_span_shadowed_by_variable_bindings( for binding in &binding_form.children { if binding.span.contains_span(target_span) { return is_span_shadowed_by_variable_binding_spec( - dialect, + semantic, binding, target_span, binding_name, @@ -122,7 +122,7 @@ pub(super) fn is_span_shadowed_by_variable_bindings( } fn is_span_shadowed_by_variable_binding_spec( - dialect: Dialect, + semantic: VerifiedSemanticPolicy, binding: &ExpressionView, target_span: ByteSpan, binding_name: &str, @@ -141,12 +141,12 @@ fn is_span_shadowed_by_variable_binding_spec( step_shadowed, has_step_forms, ); - is_span_shadowed_by_binding(dialect, child, target_span, binding_name, child_shadowed) + is_span_shadowed_by_binding(semantic, child, target_span, binding_name, child_shadowed) }) } pub(super) fn is_path_shadowed_by_variable_bindings( - dialect: Dialect, + semantic: VerifiedSemanticPolicy, binding_form: &ExpressionView, target_path: &[ChildIndex], binding_name: &str, @@ -172,7 +172,7 @@ pub(super) fn is_path_shadowed_by_variable_bindings( sequential_shadowed } else { is_path_shadowed_by_variable_binding_spec( - dialect, + semantic, binding, rest, binding_name, @@ -191,7 +191,7 @@ pub(super) fn is_path_shadowed_by_variable_bindings( } fn is_path_shadowed_by_variable_binding_spec( - dialect: Dialect, + semantic: VerifiedSemanticPolicy, binding: &ExpressionView, target_path: &[ChildIndex], binding_name: &str, @@ -212,7 +212,7 @@ fn is_path_shadowed_by_variable_binding_spec( if rest.is_empty() { child_shadowed } else { - is_path_shadowed_by_binding(dialect, child, rest, binding_name, child_shadowed) + is_path_shadowed_by_binding(semantic, child, rest, binding_name, child_shadowed) } } diff --git a/src/domain/introduce_let/syntax.rs b/src/domain/introduce_let/syntax.rs index 4f5c1407..0c8e7e86 100644 --- a/src/domain/introduce_let/syntax.rs +++ b/src/domain/introduce_let/syntax.rs @@ -1,12 +1,24 @@ use crate::domain::common_lisp::{CommonLispValueScopeForm, common_lisp_symbol_reference_eq}; -use crate::domain::dialect::Dialect; +use crate::domain::dialect::{ + BinderShape, BodyShape, Dialect, IntroduceLetOperation, ParameterShape, RelativeNodePath, + ScopeShape, VerifiedSemanticPolicy, +}; use crate::domain::sexpr::{Delimiter, ExpressionView}; mod special; use special::special_declaration_shadows_child; -pub(super) fn binding_form_binds_name(dialect: Dialect, view: &ExpressionView, name: &str) -> bool { +pub(super) fn binding_form_binds_name( + semantic: VerifiedSemanticPolicy, + view: &ExpressionView, + name: &str, +) -> bool { + if semantic_scope_binds_name(semantic, view, name) { + return true; + } + + let dialect = semantic.dialect(); let Some(head) = list_head(view) else { return false; }; @@ -52,11 +64,12 @@ pub(super) fn binding_form_binds_name(dialect: Dialect, view: &ExpressionView, n } pub(super) fn child_shadowed_by_binding( - dialect: Dialect, + semantic: VerifiedSemanticPolicy, view: &ExpressionView, name: &str, child_index: usize, ) -> bool { + let dialect = semantic.dialect(); let Some(head) = list_head(view) else { return false; }; @@ -67,36 +80,213 @@ pub(super) fn child_shadowed_by_binding( return true; } + if semantic + .scope_shape(view) + .is_some_and(|shape| semantic_scope_shadows_child(semantic, view, shape, name, child_index)) + || semantic.definition_shape(view).is_some_and(|shape| { + body_contains_child(shape.body(), child_index) + && shape.parameters().is_some_and(|parameters| { + parameter_shape_contains_name(semantic, view, parameters, name) + }) + }) + { + return true; + } + match form { Some(CommonLispValueScopeForm::Let(_)) => { - child_index >= 2 && binding_form_binds_name(dialect, view, name) + child_index >= 2 && binding_form_binds_name(semantic, view, name) } Some(CommonLispValueScopeForm::Lambda | CommonLispValueScopeForm::FunctionLiteral) => { - child_index >= 2 && binding_form_binds_name(dialect, view, name) + child_index >= 2 && binding_form_binds_name(semantic, view, name) } Some(CommonLispValueScopeForm::Definition) => { - child_index >= 3 && binding_form_binds_name(dialect, view, name) + child_index >= 3 && binding_form_binds_name(semantic, view, name) } Some(CommonLispValueScopeForm::Value) => { - child_index >= 3 && binding_form_binds_name(dialect, view, name) + child_index >= 3 && binding_form_binds_name(semantic, view, name) } Some(CommonLispValueScopeForm::Clause) => view .children .get(child_index) .is_some_and(|clause| child_index >= 2 && clause_binds_name(clause, name)), Some(CommonLispValueScopeForm::Iteration) => { - child_index >= 2 && binding_form_binds_name(dialect, view, name) + child_index >= 2 && binding_form_binds_name(semantic, view, name) } Some(CommonLispValueScopeForm::Variable(_)) => { - child_index >= 2 && binding_form_binds_name(dialect, view, name) + child_index >= 2 && binding_form_binds_name(semantic, view, name) } Some(CommonLispValueScopeForm::Slot) => { - child_index >= 3 && binding_form_binds_name(dialect, view, name) + child_index >= 3 && binding_form_binds_name(semantic, view, name) } _ => false, } } +fn semantic_scope_binds_name( + semantic: VerifiedSemanticPolicy, + view: &ExpressionView, + name: &str, +) -> bool { + let Some(shape) = semantic.scope_shape(view) else { + return false; + }; + + binder_shape_contains_name(semantic, view, shape.binders(), name) +} + +fn semantic_scope_shadows_child( + semantic: VerifiedSemanticPolicy, + view: &ExpressionView, + shape: ScopeShape, + name: &str, + child_index: usize, +) -> bool { + match (shape.binders(), shape.body()) { + ( + BinderShape::ParameterClauses { + name: local_name, + first_clause_index, + parameters, + }, + BodyShape::ClauseChildrenFrom { .. }, + ) if child_index >= first_clause_index => { + let Some(clause) = view.children.get(child_index) else { + return false; + }; + local_name + .and_then(|path| resolve_relative(view, path)) + .is_some_and(|node| semantic_pattern_contains_name(semantic, node, name)) + || parameter_shape_contains_name(semantic, clause, parameters, name) + } + (binders, body) => { + body_contains_child(body, child_index) + && binder_shape_contains_name(semantic, view, binders, name) + } + } +} + +fn binder_shape_contains_name( + semantic: VerifiedSemanticPolicy, + view: &ExpressionView, + binders: BinderShape, + name: &str, +) -> bool { + match binders { + BinderShape::BindingList { + container, + name: path, + .. + } => binding_list_contains_name(semantic, view, container, path, name), + BinderShape::NamedBindingList { + scope_name, + container, + name: path, + .. + } => { + resolve_relative(view, scope_name) + .is_some_and(|node| semantic_pattern_contains_name(semantic, node, name)) + || binding_list_contains_name(semantic, view, container, path, name) + } + BinderShape::FlatPairs { + container, + first_name_index, + stride, + .. + } => resolve_relative(view, container).is_some_and(|bindings| { + bindings + .children + .iter() + .skip(first_name_index) + .step_by(stride) + .any(|node| semantic_pattern_contains_name(semantic, node, name)) + }), + BinderShape::Parameters(parameters) => { + parameter_shape_contains_name(semantic, view, parameters, name) + } + BinderShape::NamedParameters { + name: local_name, + parameters, + } => { + resolve_relative(view, local_name) + .is_some_and(|node| semantic_pattern_contains_name(semantic, node, name)) + || parameter_shape_contains_name(semantic, view, parameters, name) + } + BinderShape::ParameterClauses { + name: local_name, + first_clause_index, + parameters, + } => { + local_name + .and_then(|path| resolve_relative(view, path)) + .is_some_and(|node| semantic_pattern_contains_name(semantic, node, name)) + || view + .children + .iter() + .skip(first_clause_index) + .any(|clause| parameter_shape_contains_name(semantic, clause, parameters, name)) + } + } +} + +fn binding_list_contains_name( + semantic: VerifiedSemanticPolicy, + view: &ExpressionView, + container: RelativeNodePath, + name_path: RelativeNodePath, + name: &str, +) -> bool { + resolve_relative(view, container).is_some_and(|bindings| { + bindings.children.iter().any(|binding| { + resolve_relative(binding, name_path) + .is_some_and(|node| semantic_pattern_contains_name(semantic, node, name)) + }) + }) +} + +fn parameter_shape_contains_name( + semantic: VerifiedSemanticPolicy, + view: &ExpressionView, + parameters: ParameterShape, + name: &str, +) -> bool { + resolve_relative(view, parameters.container()).is_some_and(|container| { + container + .children + .iter() + .skip(parameters.first_parameter_index()) + .any(|parameter| semantic_pattern_contains_name(semantic, parameter, name)) + }) +} + +fn semantic_pattern_contains_name( + semantic: VerifiedSemanticPolicy, + view: &ExpressionView, + name: &str, +) -> bool { + atom_text(view).is_some_and(|text| semantic.identifiers_equal(text, name)) + || view + .children + .iter() + .any(|child| semantic_pattern_contains_name(semantic, child, name)) +} + +fn resolve_relative(view: &ExpressionView, path: RelativeNodePath) -> Option<&ExpressionView> { + let child = view.children.get(path.child())?; + path.grandchild() + .map_or(Some(child), |grandchild| child.children.get(grandchild)) +} + +fn body_contains_child(body: BodyShape, child_index: usize) -> bool { + match body { + BodyShape::ChildrenFrom(first) => child_index >= first, + BodyShape::ChildrenAfter(path) => child_index > path.child(), + BodyShape::ClauseChildrenFrom { + first_clause_index, .. + } => child_index >= first_clause_index, + } +} + pub(super) fn local_callable_bindings_child_index( dialect: Dialect, view: &ExpressionView, diff --git a/src/domain/introduce_let/tests/basic.rs b/src/domain/introduce_let/tests/basic.rs index 13d0daf7..f4327143 100644 --- a/src/domain/introduce_let/tests/basic.rs +++ b/src/domain/introduce_let/tests/basic.rs @@ -12,6 +12,62 @@ fn introduces_single_selected_occurrence_by_default() { ); } +#[test] +fn introduce_let_preserves_dialect_reader_collisions() { + let cases = [( + Dialect::Janet, + "(+ (* width height) margin)\n# ignored ))", + "0.1", + "(let [product (* width height)] (+ product margin))\n# ignored ))", + )]; + + for (dialect, input, path, expected) in cases { + assert_plan_with_dialect(input, dialect, path, false, 1, 0, expected); + } +} + +#[test] +fn introduces_dialect_appropriate_let_for_every_verified_dialect() { + let cases = [ + ( + Dialect::CommonLisp, + "(let ((product (* width height))) (list product outer))", + ), + ( + Dialect::EmacsLisp, + "(let ((product (* width height))) (list product outer))", + ), + ( + Dialect::Scheme, + "(let ((product (* width height))) (list product outer))", + ), + ( + Dialect::Clojure, + "(let [product (* width height)] (list product outer))", + ), + ( + Dialect::Janet, + "(let [product (* width height)] (list product outer))", + ), + ( + Dialect::Fennel, + "(let [product (* width height)] (list product outer))", + ), + ]; + + for (dialect, expected) in cases { + assert_plan_with_dialect( + "(list (* width height) outer)", + dialect, + "0.1", + false, + 1, + 0, + expected, + ); + } +} + #[test] fn introduces_all_structurally_equivalent_occurrences() { assert_plan( diff --git a/src/domain/introduce_let/tests/mod.rs b/src/domain/introduce_let/tests/mod.rs index c02a0581..39b6bb56 100644 --- a/src/domain/introduce_let/tests/mod.rs +++ b/src/domain/introduce_let/tests/mod.rs @@ -14,7 +14,7 @@ fn request_with_dialect<'a>( path: &str, all_occurrences: bool, ) -> IntroduceLetRequest<'a> { - let tree = SyntaxTree::parse(input).expect("parse"); + let tree = SyntaxTree::parse_with_dialect(input, dialect).expect("parse"); let path = path.parse::().expect("path"); let selection = tree.select_path(&path).expect("select"); IntroduceLetRequest { @@ -68,6 +68,8 @@ fn assert_plan_with_dialect( expected_skipped ); assert_eq!(plan.rewritten, expected_rewritten); + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect) + .expect("rewritten output remains parseable"); } fn assert_shadowed_error(input: &str, path: &str) { diff --git a/src/domain/introduce_let/tests/rejection.rs b/src/domain/introduce_let/tests/rejection.rs index b3c78518..3dad0237 100644 --- a/src/domain/introduce_let/tests/rejection.rs +++ b/src/domain/introduce_let/tests/rejection.rs @@ -1,5 +1,19 @@ use super::*; +#[test] +fn rejects_unknown_dialect_before_introduce_let_planning() { + let mut request = request("(+ width height)", "0.1", false); + request.dialect = Dialect::Unknown; + + let error = plan_introduce_let(request).expect_err("unknown dialect should be rejected"); + + assert!( + error + .to_string() + .contains("introduce-let is not supported for this dialect") + ); +} + #[test] fn rejects_selected_expression_inside_shadowing_binding_form() { assert_shadowed_error( diff --git a/src/domain/let_binding.rs b/src/domain/let_binding.rs index 945e25ae..89e6f51e 100644 --- a/src/domain/let_binding.rs +++ b/src/domain/let_binding.rs @@ -26,14 +26,26 @@ pub struct ConvertLetToLetStarPlan { pub changed: bool, } +pub(crate) fn validate_convert_let_to_let_star_dialect(dialect: Dialect) -> Result<()> { + if !matches!(dialect, Dialect::CommonLisp | Dialect::EmacsLisp) { + bail!("convert-let-to-let-star supports only Common Lisp and Emacs Lisp"); + } + Ok(()) +} + +pub(crate) fn validate_convert_let_star_to_let_dialect(dialect: Dialect) -> Result<()> { + if dialect != Dialect::CommonLisp { + bail!("convert-let-star-to-let currently supports only Common Lisp"); + } + Ok(()) +} + pub fn plan_convert_let_to_let_star( request: ConvertLetToLetStarRequest<'_>, ) -> Result { - if !matches!(request.dialect, Dialect::CommonLisp | Dialect::EmacsLisp) { - bail!("convert-let-to-let-star supports only Common Lisp and Emacs Lisp"); - } - let tree = - SyntaxTree::parse(request.input).context("convert-let-to-let-star input is not valid")?; + validate_convert_let_to_let_star_dialect(request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("convert-let-to-let-star input is not valid")?; let form = tree.select_path(&request.path)?.view(); validate_form( &form, @@ -46,7 +58,7 @@ pub fn plan_convert_let_to_let_star( analyze_bindings(&form, request.dialect, "convert-let-to-let-star")?; reject_dependencies(&names, &initializers, &request, "convert-let-to-let-star")?; let rewritten = replace_head(request.input, form.children[0].span, "let*"); - parse_output(&rewritten, "convert-let-to-let-star")?; + parse_output(&rewritten, request.dialect, "convert-let-to-let-star")?; Ok(ConvertLetToLetStarPlan { dialect: request.dialect, path: request.path, @@ -76,10 +88,8 @@ pub struct ConvertLetStarToLetPlan { pub fn plan_convert_let_star_to_let( request: ConvertLetStarToLetRequest<'_>, ) -> Result { - if request.dialect != Dialect::CommonLisp { - bail!("convert-let-star-to-let currently supports only Common Lisp"); - } - let tree = SyntaxTree::parse(request.input) + validate_convert_let_star_to_let_dialect(request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) .context("convert-let-star-to-let input is not a valid S-expression document")?; let form = tree.select_path(&request.path)?.view(); validate_form( @@ -93,7 +103,7 @@ pub fn plan_convert_let_star_to_let( analyze_bindings(&form, request.dialect, "convert-let-star-to-let")?; reject_dependencies(&names, &initializers, &request, "convert-let-star-to-let")?; let rewritten = replace_head(request.input, form.children[0].span, "let"); - parse_output(&rewritten, "convert-let-star-to-let")?; + parse_output(&rewritten, request.dialect, "convert-let-star-to-let")?; Ok(ConvertLetStarToLetPlan { dialect: request.dialect, path: request.path, @@ -260,8 +270,9 @@ fn replace_head(input: &str, span: ByteSpan, replacement: &str) -> String { output.push_str(&input[span.end().get()..]); output } -fn parse_output(output: &str, operation: &str) -> Result<()> { - SyntaxTree::parse(output).with_context(|| format!("{operation} output is not valid"))?; +fn parse_output(output: &str, dialect: Dialect, operation: &str) -> Result<()> { + SyntaxTree::parse_with_dialect(output, dialect) + .with_context(|| format!("{operation} output is not valid"))?; Ok(()) } @@ -333,4 +344,73 @@ mod tests { .is_err() ); } + + #[test] + fn dialect_support_matrix_is_enforced_before_parsing_and_reparses_output() { + for (dialect, input) in [ + (Dialect::CommonLisp, "#\\) (let ((x 1) (y 2)) (+ x y))"), + (Dialect::EmacsLisp, "?\\) (let ((x 1) (y 2)) (+ x y))"), + ] { + let plan = plan_convert_let_to_let_star(ConvertLetToLetStarRequest { + input, + dialect, + path: "1".parse().expect("path"), + }) + .expect("supported dialect"); + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect) + .expect("dialect-specific let output"); + } + + for dialect in [ + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ] { + let error = plan_convert_let_to_let_star(ConvertLetToLetStarRequest { + input: ")", + dialect, + path: "0".parse().expect("path"), + }) + .expect_err("unsupported dialect"); + assert!( + error + .to_string() + .contains("supports only Common Lisp and Emacs Lisp"), + "{dialect:?}: {error:#}" + ); + } + + let plan = plan_convert_let_star_to_let(ConvertLetStarToLetRequest { + input: "#\\) (let* ((x 1) (y 2)) (+ x y))", + dialect: Dialect::CommonLisp, + path: "1".parse().expect("path"), + }) + .expect("Common Lisp"); + SyntaxTree::parse_with_dialect(&plan.rewritten, Dialect::CommonLisp) + .expect("Common Lisp let output"); + + for dialect in [ + Dialect::EmacsLisp, + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ] { + let error = plan_convert_let_star_to_let(ConvertLetStarToLetRequest { + input: ")", + dialect, + path: "0".parse().expect("path"), + }) + .expect_err("unsupported dialect"); + assert!( + error + .to_string() + .contains("currently supports only Common Lisp"), + "{dialect:?}: {error:#}" + ); + } + } } diff --git a/src/domain/let_composition.rs b/src/domain/let_composition.rs index 834134d3..30c2268a 100644 --- a/src/domain/let_composition.rs +++ b/src/domain/let_composition.rs @@ -239,12 +239,18 @@ fn select( path: &Path, operation: &str, ) -> Result<(SyntaxTree, ExpressionView)> { + validate_dialect(dialect, operation)?; + let tree = SyntaxTree::parse_with_dialect(input, dialect) + .context("input is not a valid S-expression document")?; + let view = tree.select_path(path)?.view(); + Ok((tree, view)) +} + +pub(crate) fn validate_dialect(dialect: Dialect, operation: &str) -> Result<()> { if !matches!(dialect, Dialect::CommonLisp | Dialect::EmacsLisp) { bail!("{operation} supports only Common Lisp and Emacs Lisp"); } - let tree = SyntaxTree::parse(input).context("input is not a valid S-expression document")?; - let view = tree.select_path(path)?.view(); - Ok((tree, view)) + Ok(()) } fn require_form( dialect: Dialect, @@ -381,7 +387,8 @@ fn finish( operation: &str, ) -> Result { let rewritten = replace_span(input, span, &replacement); - SyntaxTree::parse(&rewritten).with_context(|| format!("{operation} output is not valid"))?; + SyntaxTree::parse_with_dialect(&rewritten, dialect) + .with_context(|| format!("{operation} output is not valid"))?; Ok(LetCompositionPlan { dialect, path, @@ -451,7 +458,7 @@ mod tests { path: path(), }) .expect("merge"); - SyntaxTree::parse(&merged.rewritten).expect("merged output"); + SyntaxTree::parse_with_dialect(&merged.rewritten, dialect).expect("merged output"); let split = plan_split_let(SplitLetRequest { input: "(let ((x 1) (y 2)) (+ x y))", dialect, @@ -462,7 +469,7 @@ mod tests { assert_eq!(split.binding_index.get(), 1); assert_eq!(split.outer_binding_count, 1); assert_eq!(split.inner_binding_count, 1); - SyntaxTree::parse(&split.rewritten).expect("split output"); + SyntaxTree::parse_with_dialect(&split.rewritten, dialect).expect("split output"); } } @@ -526,4 +533,85 @@ mod tests { .expect("let* merge"); assert_eq!(plan.rewritten, "(let* ((x 1) (y (+ x 1))) y)"); } + + #[test] + fn supported_dialects_handle_reader_collisions_and_reparse_output() { + for (dialect, prefix) in [(Dialect::CommonLisp, "#\\)"), (Dialect::EmacsLisp, "?\\)")] { + let path: Path = "1".parse().expect("path"); + + let merged = plan_merge_nested_let(MergeNestedLetRequest { + input: &format!("{prefix} (let ((x 1)) (let ((y 2)) (+ x y)))"), + dialect, + path: path.clone(), + }) + .expect("merge"); + SyntaxTree::parse_with_dialect(&merged.rewritten, dialect).expect("merged output"); + + let merged_star = plan_merge_nested_let_star(MergeNestedLetStarRequest { + input: &format!("{prefix} (let* ((x 1)) (let* ((y (+ x 1))) y))"), + dialect, + path: path.clone(), + }) + .expect("let* merge"); + SyntaxTree::parse_with_dialect(&merged_star.rewritten, dialect) + .expect("let* merged output"); + + let split = plan_split_let(SplitLetRequest { + input: &format!("{prefix} (let ((x 1) (y 2)) (+ x y))"), + dialect, + path, + binding_index: BindingIndex::new(1).expect("binding index"), + }) + .expect("split"); + SyntaxTree::parse_with_dialect(&split.rewritten, dialect).expect("split output"); + } + } + + #[test] + fn unsupported_dialects_are_rejected_before_parsing() { + for dialect in [ + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ] { + let merge_error = plan_merge_nested_let(MergeNestedLetRequest { + input: ")", + dialect, + path: path(), + }) + .expect_err("unsupported dialect"); + assert!( + merge_error + .to_string() + .contains("supports only Common Lisp and Emacs Lisp") + ); + + let merge_star_error = plan_merge_nested_let_star(MergeNestedLetStarRequest { + input: ")", + dialect, + path: path(), + }) + .expect_err("unsupported dialect"); + assert!( + merge_star_error + .to_string() + .contains("supports only Common Lisp and Emacs Lisp") + ); + + let split_error = plan_split_let(SplitLetRequest { + input: ")", + dialect, + path: path(), + binding_index: BindingIndex::new(1).expect("binding index"), + }) + .expect_err("unsupported dialect"); + assert!( + split_error + .to_string() + .contains("supports only Common Lisp and Emacs Lisp") + ); + } + } } diff --git a/src/domain/let_star_composition.rs b/src/domain/let_star_composition.rs index 91264520..75cda7f9 100644 --- a/src/domain/let_star_composition.rs +++ b/src/domain/let_star_composition.rs @@ -27,10 +27,8 @@ pub(crate) struct Plan { } pub(crate) fn plan(request: Request<'_>) -> Result { - if !matches!(request.dialect, Dialect::CommonLisp | Dialect::EmacsLisp) { - bail!("split-let-star supports only Common Lisp and Emacs Lisp"); - } - let tree = SyntaxTree::parse(request.input) + validate_dialect(request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) .context("split-let-star input is not a valid S-expression document")?; let form = tree.select_path(&request.path)?.view(); require_let_star(request.dialect, &form)?; @@ -71,7 +69,8 @@ pub(crate) fn plan(request: Request<'_>) -> Result { let body = &request.input[bindings.span.end().get()..form.span.end().get() - 1]; let replacement = format!("({head} ({outer}) ({head} ({inner}){body}))"); let rewritten = replace_span(request.input, form.span, &replacement); - SyntaxTree::parse(&rewritten).context("split-let-star output is not valid")?; + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) + .context("split-let-star output is not valid")?; Ok(Plan { dialect: request.dialect, path: request.path, @@ -84,6 +83,13 @@ pub(crate) fn plan(request: Request<'_>) -> Result { }) } +pub(crate) fn validate_dialect(dialect: Dialect) -> Result<()> { + if !matches!(dialect, Dialect::CommonLisp | Dialect::EmacsLisp) { + bail!("split-let-star supports only Common Lisp and Emacs Lisp"); + } + Ok(()) +} + fn replace_span(input: &str, span: ByteSpan, replacement: &str) -> String { format!( "{}{}{}", @@ -136,22 +142,54 @@ fn contains_headed_form(dialect: Dialect, view: &ExpressionView, expected: &str) #[cfg(test)] mod tests { use super::*; + #[test] - fn splits_in_both_dialects() { - for dialect in [Dialect::CommonLisp, Dialect::EmacsLisp] { - let p = plan(Request { - input: "(let* ((a 1) (b (+ a 1)) (c (+ b 1))) (+ a b c))", + fn splits_in_both_dialects_with_reader_literals_outside_target() { + for (dialect, reader_literal) in + [(Dialect::CommonLisp, r"#\)"), (Dialect::EmacsLisp, r"?\)")] + { + let input = + format!("(let* ((a 1) (b (+ a 1)) (c (+ b 1))) (+ a b c)) {reader_literal}"); + let plan = plan(Request { + input: &input, dialect, - path: "0".parse().unwrap(), + path: "0".parse().expect("path"), binding_index: BindingIndex::new(1).expect("binding index"), }) - .unwrap(); + .expect("plan"); assert_eq!( - p.rewritten, - "(let* ((a 1)) (let* ((b (+ a 1)) (c (+ b 1))) (+ a b c)))" + plan.rewritten, + format!( + "(let* ((a 1)) (let* ((b (+ a 1)) (c (+ b 1))) (+ a b c))) {reader_literal}" + ) ); + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect).expect("rewritten input"); } } + + #[test] + fn rejects_unsupported_dialects_before_parsing() { + for dialect in [ + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ] { + let error = plan(Request { + input: ")", + dialect, + path: "0".parse().expect("path"), + binding_index: BindingIndex::new(1).expect("binding index"), + }) + .expect_err("dialect must be rejected"); + assert_eq!( + error.to_string(), + "split-let-star supports only Common Lisp and Emacs Lisp" + ); + } + } + #[test] fn rejects_invalid_boundaries_and_declarations() { assert!( diff --git a/src/domain/lexical_scope/capture.rs b/src/domain/lexical_scope/capture.rs index 5ad1af75..025f1e60 100644 --- a/src/domain/lexical_scope/capture.rs +++ b/src/domain/lexical_scope/capture.rs @@ -41,6 +41,9 @@ pub fn value_capture( if free_vars.is_empty() { return Vec::new(); } + if dialect == Dialect::Unknown { + return free_vars; + } let base = scope_span.slice(input); let scope_start = scope_span.start().get(); @@ -55,25 +58,26 @@ pub fn value_capture( .collect(); sites.sort_unstable(); - let Ok(original) = SyntaxTree::parse(base) else { - return Vec::new(); + let Ok(original) = SyntaxTree::parse_with_dialect(base, dialect) else { + return free_vars; }; let original_root = original.root_view(); let Some(original_form) = original_root.children.first() else { - return Vec::new(); + return free_vars; }; let mut captured = Vec::new(); for symbol in &free_vars { let original_free = count_unshadowed(dialect, original_form, symbol, base); let probed = splice_all(base, &sites, symbol.as_str()); - let Ok(tree) = SyntaxTree::parse(&probed) else { + let Ok(tree) = SyntaxTree::parse_with_dialect(&probed, dialect) else { // A value that cannot be safely re-inserted is treated as unsafe. captured.push(symbol.clone()); continue; }; let root = tree.root_view(); let Some(form) = root.children.first() else { + captured.push(symbol.clone()); continue; }; let probed_free = count_unshadowed(dialect, form, symbol, &probed); diff --git a/src/domain/lexical_scope/tests/capture.rs b/src/domain/lexical_scope/tests/capture.rs new file mode 100644 index 00000000..b47786db --- /dev/null +++ b/src/domain/lexical_scope/tests/capture.rs @@ -0,0 +1,118 @@ +use super::*; + +fn parse_path(path: &str) -> Path { + path.parse().expect("path") +} + +fn captured_names( + input: &str, + parsed_dialect: Dialect, + capture_dialect: Dialect, + value_path: &str, + reference_path: &str, +) -> Vec { + let tree = SyntaxTree::parse_with_dialect(input, parsed_dialect).expect("dialect parse"); + let scope = tree.select_path(&parse_path("0")).expect("scope"); + let value = tree.select_path(&parse_path(value_path)).expect("value"); + let reference = tree + .select_path(&parse_path(reference_path)) + .expect("reference"); + + value_capture( + capture_dialect, + input, + scope.span(), + &SymbolName::new("target").expect("binding"), + &value.view(), + &[reference.span()], + ) + .into_iter() + .map(|symbol| symbol.as_str().to_owned()) + .collect() +} + +#[test] +fn common_lisp_dispatch_form_uses_the_dialect_reader_shape() { + let input = r"(let ((target external)) #S(holder :value external) target)"; + let dialect_tree = + SyntaxTree::parse_with_dialect(input, Dialect::CommonLisp).expect("Common Lisp parse"); + let generic_tree = SyntaxTree::parse(input).expect("legacy parse"); + assert_eq!(dialect_tree.root_view().children[0].children.len(), 4); + assert_eq!(generic_tree.root_view().children[0].children.len(), 5); + + assert!( + captured_names( + input, + Dialect::CommonLisp, + Dialect::CommonLisp, + "0.1.0.1", + "0.3", + ) + .is_empty() + ); +} + +#[test] +fn clojure_discard_and_tagged_literal_use_the_dialect_reader_shape() { + let input = r#"(let [target external] #_ignored #inst "2020-01-01" target)"#; + let dialect_tree = + SyntaxTree::parse_with_dialect(input, Dialect::Clojure).expect("Clojure parse"); + let generic_tree = SyntaxTree::parse(input).expect("legacy parse"); + assert_eq!(dialect_tree.root_view().children[0].children.len(), 4); + assert_eq!(generic_tree.root_view().children[0].children.len(), 5); + + assert!(captured_names(input, Dialect::Clojure, Dialect::Clojure, "0.1.1", "0.3",).is_empty()); +} + +#[test] +fn dialect_valid_scope_that_the_generic_reader_rejects_is_not_false_safe() { + let input = "(let [target external]\n # unmatched )\n (let [external 1] target))"; + assert!(SyntaxTree::parse(input).is_err()); + SyntaxTree::parse_with_dialect(input, Dialect::Janet).expect("Janet parse"); + + assert_eq!( + captured_names(input, Dialect::Janet, Dialect::Janet, "0.1.1", "0.2.2",), + ["external"] + ); +} + +#[test] +fn unknown_dialect_returns_all_free_variables_as_unsafe() { + let input = "(let ((target external)) target)"; + assert_eq!( + captured_names( + input, + Dialect::CommonLisp, + Dialect::Unknown, + "0.1.0.1", + "0.2", + ), + ["external"] + ); +} + +#[test] +fn scheme_variadic_lambda_parameter_captures_spliced_value() { + let input = "(let ((target external)) (lambda external target))"; + + assert_eq!( + captured_names(input, Dialect::Scheme, Dialect::Scheme, "0.1.0.1", "0.2.2",), + ["external"] + ); +} + +#[test] +fn clojure_named_multi_arity_fn_captures_spliced_value() { + let input = "(let [target recur-name] (fn recur-name ([value] target)))"; + + assert_eq!( + captured_names( + input, + Dialect::Clojure, + Dialect::Clojure, + "0.1.1", + "0.2.2.1", + ), + ["recur-name"] + ); +} diff --git a/src/domain/lexical_scope/tests/mod.rs b/src/domain/lexical_scope/tests/mod.rs index 25efcef3..c937360e 100644 --- a/src/domain/lexical_scope/tests/mod.rs +++ b/src/domain/lexical_scope/tests/mod.rs @@ -9,6 +9,16 @@ fn selected_form(input: &str) -> crate::domain::sexpr::ExpressionView { .view() } +fn selected_form_with_dialect( + input: &str, + dialect: Dialect, +) -> crate::domain::sexpr::ExpressionView { + let tree = SyntaxTree::parse_with_dialect(input, dialect).expect("parse"); + tree.select_path(&"0".parse::().expect("path")) + .expect("select") + .view() +} + fn reference_texts(input: &str, symbol: &str) -> Vec { let view = selected_form(input); let symbol = SymbolName::new(symbol).expect("symbol"); @@ -20,7 +30,19 @@ fn reference_texts(input: &str, symbol: &str) -> Vec { .collect() } +fn reference_texts_with_dialect(input: &str, dialect: Dialect, symbol: &str) -> Vec { + let view = selected_form_with_dialect(input, dialect); + let symbol = SymbolName::new(symbol).expect("symbol"); + let mut spans = Vec::new(); + collect_unshadowed_symbol_references(dialect, &view, &symbol, input, &mut spans); + spans + .into_iter() + .map(|span| span.slice(input).to_owned()) + .collect() +} + mod binding_forms; mod boundaries; +mod capture; mod property; mod shadowing; diff --git a/src/domain/lexical_scope/tests/shadowing.rs b/src/domain/lexical_scope/tests/shadowing.rs index 366167ef..900cb414 100644 --- a/src/domain/lexical_scope/tests/shadowing.rs +++ b/src/domain/lexical_scope/tests/shadowing.rs @@ -158,3 +158,107 @@ fn package_qualified_compiler_macrolet_expander_bodies_remain_outer_references() assert_eq!(reference_texts(input, "m"), vec!["m", "m", "m"]); } + +#[test] +fn common_lisp_lambda_parameters_preserve_shadowing_contract() { + assert!( + reference_texts_with_dialect("(lambda (value) value)", Dialect::CommonLisp, "value",) + .is_empty() + ); + assert_eq!( + reference_texts_with_dialect("(lambda (other) value)", Dialect::CommonLisp, "value",), + ["value"] + ); +} + +#[test] +fn emacs_lisp_lambda_parameters_preserve_shadowing_contract() { + assert!( + reference_texts_with_dialect("(lambda (value) value)", Dialect::EmacsLisp, "value",) + .is_empty() + ); + assert_eq!( + reference_texts_with_dialect("(lambda (other) value)", Dialect::EmacsLisp, "value",), + ["value"] + ); +} + +#[test] +fn scheme_variadic_lambda_atom_shadows_its_body() { + assert!( + reference_texts_with_dialect("(lambda value value)", Dialect::Scheme, "value",).is_empty() + ); + assert_eq!( + reference_texts_with_dialect("(lambda other value)", Dialect::Scheme, "value"), + ["value"] + ); +} + +#[test] +fn clojure_anonymous_fn_parameters_shadow_their_body() { + assert!( + reference_texts_with_dialect("(fn [value] value)", Dialect::Clojure, "value",).is_empty() + ); + assert_eq!( + reference_texts_with_dialect("(fn [other] value)", Dialect::Clojure, "value"), + ["value"] + ); +} + +#[test] +fn clojure_named_fn_name_and_parameters_shadow_their_body() { + assert!( + reference_texts_with_dialect( + "(fn recur-name [value] (list recur-name value))", + Dialect::Clojure, + "recur-name", + ) + .is_empty() + ); + assert_eq!( + reference_texts_with_dialect("(fn recur-name [other] value)", Dialect::Clojure, "value",), + ["value"] + ); +} + +#[test] +fn clojure_multi_arity_fn_scopes_each_parameter_clause() { + assert_eq!( + reference_texts_with_dialect( + "(fn ([value] value) ([other] value))", + Dialect::Clojure, + "value", + ), + ["value"] + ); + assert!( + reference_texts_with_dialect( + "(fn recur-name ([value] recur-name) ([other] recur-name))", + Dialect::Clojure, + "recur-name", + ) + .is_empty() + ); +} + +#[test] +fn janet_fn_vector_parameters_shadow_their_body() { + assert!( + reference_texts_with_dialect("(fn [value] value)", Dialect::Janet, "value",).is_empty() + ); + assert_eq!( + reference_texts_with_dialect("(fn [other] value)", Dialect::Janet, "value"), + ["value"] + ); +} + +#[test] +fn fennel_fn_vector_parameters_shadow_their_body() { + assert!( + reference_texts_with_dialect("(fn [value] value)", Dialect::Fennel, "value",).is_empty() + ); + assert_eq!( + reference_texts_with_dialect("(fn [other] value)", Dialect::Fennel, "value"), + ["value"] + ); +} diff --git a/src/domain/lexical_scope/traversal/binding_forms/mod.rs b/src/domain/lexical_scope/traversal/binding_forms/mod.rs index 4669b427..200fda5b 100644 --- a/src/domain/lexical_scope/traversal/binding_forms/mod.rs +++ b/src/domain/lexical_scope/traversal/binding_forms/mod.rs @@ -1,11 +1,11 @@ use crate::domain::common_lisp::CommonLispOperator; use crate::domain::definition::definition_shape; -use crate::domain::dialect::Dialect; +use crate::domain::dialect::{BinderShape, BodyShape, Dialect, ParameterShape, RelativeNodePath}; use crate::domain::sexpr::{ByteSpan, ExpressionKind, ExpressionView, SymbolName}; use super::body::collect_body_forms; -use super::collect_unshadowed_symbol_references_in_context; use super::lambda_lists::collect_lambda_list_references; +use super::{collect_unshadowed_symbol_references_in_context, symbol_name_matches}; use crate::domain::lexical_scope::bindings::{binding_binds, generic_binding_groups}; mod clause_bindings; @@ -38,6 +38,10 @@ pub(super) fn collect_shadow_aware_special_form( return false; }; + if is_dialect_callable_head(dialect, head) { + return collect_dialect_callable_references(dialect, view, symbol, input, output); + } + if dialect == Dialect::Scheme && is_scheme_named_let(view, head) { collect_named_let_references(dialect, view, symbol, input, output); return true; @@ -156,6 +160,132 @@ pub(super) fn collect_shadow_aware_special_form( } } +fn is_dialect_callable_head(dialect: Dialect, head: &str) -> bool { + matches!( + (dialect, head), + (Dialect::Scheme, "lambda") | (Dialect::Clojure | Dialect::Janet | Dialect::Fennel, "fn") + ) +} + +fn collect_dialect_callable_references( + dialect: Dialect, + view: &ExpressionView, + symbol: &SymbolName, + input: &str, + output: &mut Vec, +) -> bool { + if dialect == Dialect::Scheme { + let Some(parameter_form) = view.children.get(1) else { + return false; + }; + if parameter_form.kind == ExpressionKind::Atom { + let Some(parameter_name) = super::super::syntax::atom_text(parameter_form) else { + return false; + }; + if !symbol_name_matches(dialect, parameter_name, symbol.as_str()) { + collect_body_forms(dialect, &view.children[2..], symbol, input, output); + } + return true; + } + } + + let Some(scope) = dialect + .verify_rename_binding() + .ok() + .and_then(|policy| policy.scope_shape(view)) + else { + return false; + }; + + match (scope.binders(), scope.body()) { + (BinderShape::Parameters(parameters), BodyShape::ChildrenFrom(body_index)) => { + collect_parameter_scope_references( + dialect, + view, + parameters, + &view.children[body_index..], + symbol, + input, + output, + ) + } + ( + BinderShape::NamedParameters { name, parameters }, + BodyShape::ChildrenFrom(body_index), + ) => { + let Some(name) = resolve_relative(view, name) else { + return false; + }; + if super::super::syntax::atom_text(name) + .is_some_and(|name| symbol_name_matches(dialect, name, symbol.as_str())) + { + return true; + } + collect_parameter_scope_references( + dialect, + view, + parameters, + &view.children[body_index..], + symbol, + input, + output, + ) + } + ( + BinderShape::ParameterClauses { + name, + first_clause_index, + parameters, + }, + BodyShape::ClauseChildrenFrom { + first_clause_index: body_first_clause_index, + body_child_index, + }, + ) if first_clause_index == body_first_clause_index => { + if name + .and_then(|path| resolve_relative(view, path)) + .and_then(super::super::syntax::atom_text) + .is_some_and(|name| symbol_name_matches(dialect, name, symbol.as_str())) + { + return true; + } + + view.children.iter().skip(first_clause_index).all(|clause| { + clause.children.get(body_child_index..).is_some_and(|body| { + collect_parameter_scope_references( + dialect, clause, parameters, body, symbol, input, output, + ) + }) + }) + } + _ => false, + } +} + +fn collect_parameter_scope_references( + dialect: Dialect, + scope_root: &ExpressionView, + parameters: ParameterShape, + body_forms: &[ExpressionView], + symbol: &SymbolName, + input: &str, + output: &mut Vec, +) -> bool { + if parameters.first_parameter_index() != 0 { + return false; + } + let Some(parameter_form) = resolve_relative(scope_root, parameters.container()) else { + return false; + }; + collect_lambda_list_references(dialect, parameter_form, body_forms, symbol, input, output) +} + +fn resolve_relative(view: &ExpressionView, path: RelativeNodePath) -> Option<&ExpressionView> { + let child = view.children.get(path.child())?; + path.grandchild() + .map_or(Some(child), |grandchild| child.children.get(grandchild)) +} + fn is_scheme_named_let(view: &ExpressionView, head: &str) -> bool { (head == "let" || head == "let*") && view diff --git a/src/domain/lexical_scope/traversal/lambda_lists.rs b/src/domain/lexical_scope/traversal/lambda_lists.rs index fa71c529..761b9a91 100644 --- a/src/domain/lexical_scope/traversal/lambda_lists.rs +++ b/src/domain/lexical_scope/traversal/lambda_lists.rs @@ -3,7 +3,7 @@ use crate::domain::dialect::Dialect; use crate::domain::sexpr::{ByteSpan, ExpressionKind, ExpressionView, SymbolName}; use super::body::collect_body_forms; -use super::collect_unshadowed_symbol_references_in_context; +use super::{collect_unshadowed_symbol_references_in_context, symbol_name_matches}; use crate::domain::lexical_scope::patterns::binding_pattern_names; #[derive(Clone, Copy, Eq, PartialEq)] @@ -26,6 +26,20 @@ pub(super) fn collect_lambda_list_references( return false; } + if matches!( + dialect, + Dialect::Scheme | Dialect::Clojure | Dialect::Janet | Dialect::Fennel + ) { + return collect_simple_parameter_list_references( + dialect, + parameter_form, + body_forms, + symbol, + input, + output, + ); + } + let mut mode = LambdaListMode::Required; let mut index = 0usize; @@ -58,6 +72,26 @@ pub(super) fn collect_lambda_list_references( true } +fn collect_simple_parameter_list_references( + dialect: Dialect, + parameter_form: &ExpressionView, + body_forms: &[ExpressionView], + symbol: &SymbolName, + input: &str, + output: &mut Vec, +) -> bool { + let is_shadowed = parameter_form.children.iter().any(|parameter| { + lambda_list_binding_names(parameter, LambdaListMode::Required) + .iter() + .any(|name| symbol_name_matches(dialect, name, symbol.as_str())) + }); + + if !is_shadowed { + collect_body_forms(dialect, body_forms, symbol, input, output); + } + true +} + fn collect_lambda_list_marker( parameter_form: &ExpressionView, child: &ExpressionView, diff --git a/src/domain/local_function_binding.rs b/src/domain/local_function_binding.rs index 8c27db6c..5fdd7c12 100644 --- a/src/domain/local_function_binding.rs +++ b/src/domain/local_function_binding.rs @@ -24,6 +24,21 @@ pub struct ConvertFletToLabelsPlan { pub changed: bool, } +pub(crate) fn validate_convert_flet_to_labels_dialect(dialect: Dialect) -> Result<()> { + validate_common_lisp_dialect(dialect, "convert-flet-to-labels") +} + +pub(crate) fn validate_convert_labels_to_flet_dialect(dialect: Dialect) -> Result<()> { + validate_common_lisp_dialect(dialect, "convert-labels-to-flet") +} + +fn validate_common_lisp_dialect(dialect: Dialect, operation: &str) -> Result<()> { + if dialect != Dialect::CommonLisp { + bail!("{operation} supports only Common Lisp"); + } + Ok(()) +} + pub fn plan_convert_flet_to_labels( request: ConvertFletToLabelsRequest<'_>, ) -> Result { @@ -39,7 +54,7 @@ pub fn plan_convert_flet_to_labels( } } let rewritten = replace_head(request.input, &form, replace_flet_name(&head)); - parse_output(&rewritten, "convert-flet-to-labels")?; + parse_output(&rewritten, request.dialect, "convert-flet-to-labels")?; Ok(ConvertFletToLabelsPlan { dialect: request.dialect, path: request.path, @@ -82,7 +97,7 @@ pub fn plan_convert_labels_to_flet( } } let rewritten = replace_head(request.input, &form, replace_labels_name(&head)); - parse_output(&rewritten, "convert-labels-to-flet")?; + parse_output(&rewritten, request.dialect, "convert-labels-to-flet")?; Ok(ConvertLabelsToFletPlan { dialect: request.dialect, path: request.path, @@ -107,10 +122,8 @@ fn analyze_bindings<'a, R>( where R: BindingRequest<'a> + ?Sized, { - if request.dialect() != Dialect::CommonLisp { - bail!("{operation} supports only Common Lisp"); - } - let tree = SyntaxTree::parse(request.input()) + validate_common_lisp_dialect(request.dialect(), operation)?; + let tree = SyntaxTree::parse_with_dialect(request.input(), request.dialect()) .with_context(|| format!("{operation} input is not a valid S-expression document"))?; let form = tree.select_path(request.path())?.view(); if tree.has_comment_in(form.span) { @@ -245,8 +258,8 @@ fn replace_head(input: &str, form: &ExpressionView, replacement: String) -> Stri output } -fn parse_output(rewritten: &str, operation: &str) -> Result<()> { - SyntaxTree::parse(rewritten) +fn parse_output(rewritten: &str, dialect: Dialect, operation: &str) -> Result<()> { + SyntaxTree::parse_with_dialect(rewritten, dialect) .with_context(|| format!("{operation} output is not a valid S-expression document"))?; Ok(()) } @@ -255,26 +268,59 @@ fn parse_output(rewritten: &str, operation: &str) -> Result<()> { mod tests { use super::*; + const DIALECTS: [Dialect; 7] = [ + Dialect::CommonLisp, + Dialect::EmacsLisp, + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ]; + + fn flet_request(input: &str, dialect: Dialect) -> ConvertFletToLabelsRequest<'_> { + ConvertFletToLabelsRequest { + input, + dialect, + path: "0".parse().expect("path"), + } + } + + fn labels_request(input: &str, dialect: Dialect) -> ConvertLabelsToFletRequest<'_> { + ConvertLabelsToFletRequest { + input, + dialect, + path: "0".parse().expect("path"), + } + } + + fn assert_support_error(result: Result, operation: &str) { + let error = result.err().expect("unsupported dialect must fail"); + assert_eq!( + error.to_string(), + format!("{operation} supports only Common Lisp") + ); + } + #[test] fn converts_capture_free_and_non_recursive_forms() { - let flet = plan_convert_flet_to_labels(ConvertFletToLabelsRequest { - input: "(flet ((work () 1)) (work))", - dialect: Dialect::CommonLisp, - path: "0".parse().expect("path"), - }) + let flet = plan_convert_flet_to_labels(flet_request( + "(flet ((work () 1)) (work))", + Dialect::CommonLisp, + )) .expect("flet plan"); assert_eq!(flet.rewritten, "(labels ((work () 1)) (work))"); - let labels = plan_convert_labels_to_flet(ConvertLabelsToFletRequest { - input: "(labels ((work () 1)) (work))", - dialect: Dialect::CommonLisp, - path: "0".parse().expect("path"), - }) + + let labels = plan_convert_labels_to_flet(labels_request( + "(labels ((work () 1)) (work))", + Dialect::CommonLisp, + )) .expect("labels plan"); assert_eq!(labels.rewritten, "(flet ((work () 1)) (work))"); } #[test] - fn rejects_recursion_duplicates_malformed_forms_and_other_dialects() { + fn rejects_recursion_duplicates_and_malformed_forms() { for input in [ "(labels ((walk () (walk))) (walk))", "(labels ((walk () (function walk))) (walk))", @@ -284,27 +330,68 @@ mod tests { "(flet ((work () ; keep\n 1)) (work))", ] { assert!( - plan_convert_labels_to_flet(ConvertLabelsToFletRequest { - input, - dialect: Dialect::CommonLisp, - path: "0".parse().expect("path"), - }) - .is_err() - || plan_convert_flet_to_labels(ConvertFletToLabelsRequest { - input, - dialect: Dialect::CommonLisp, - path: "0".parse().expect("path"), - }) - .is_err() + plan_convert_labels_to_flet(labels_request(input, Dialect::CommonLisp)).is_err() + || plan_convert_flet_to_labels(flet_request(input, Dialect::CommonLisp)) + .is_err() ); } - assert!( - plan_convert_flet_to_labels(ConvertFletToLabelsRequest { - input: "(flet ((work () 1)) (work))", - dialect: Dialect::EmacsLisp, - path: "0".parse().expect("path"), - }) - .is_err() - ); + } + + #[test] + fn support_matrix_is_common_lisp_only_for_both_conversions() { + for dialect in DIALECTS { + let flet = + plan_convert_flet_to_labels(flet_request("(flet ((work () 1)) (work))", dialect)); + let labels = plan_convert_labels_to_flet(labels_request( + "(labels ((work () 1)) (work))", + dialect, + )); + + if dialect == Dialect::CommonLisp { + assert!(flet.is_ok(), "Common Lisp flet conversion must succeed"); + assert!(labels.is_ok(), "Common Lisp labels conversion must succeed"); + } else { + assert_support_error(flet, "convert-flet-to-labels"); + assert_support_error(labels, "convert-labels-to-flet"); + } + } + } + + #[test] + fn unsupported_dialect_gate_precedes_parsing_for_both_conversions() { + for dialect in DIALECTS + .into_iter() + .filter(|dialect| *dialect != Dialect::CommonLisp) + { + assert_support_error( + plan_convert_flet_to_labels(flet_request(")", dialect)), + "convert-flet-to-labels", + ); + assert_support_error( + plan_convert_labels_to_flet(labels_request(")", dialect)), + "convert-labels-to-flet", + ); + } + } + + #[test] + fn preserves_common_lisp_delimiter_character_literals() { + let flet = plan_convert_flet_to_labels(flet_request( + "(flet ((work () #\\))) (work))", + Dialect::CommonLisp, + )) + .expect("flet character literal plan"); + assert_eq!(flet.rewritten, "(labels ((work () #\\))) (work))"); + SyntaxTree::parse_with_dialect(&flet.rewritten, Dialect::CommonLisp) + .expect("flet output must reparse as Common Lisp"); + + let labels = plan_convert_labels_to_flet(labels_request( + "(labels ((work () #\\))) (work))", + Dialect::CommonLisp, + )) + .expect("labels character literal plan"); + assert_eq!(labels.rewritten, "(flet ((work () #\\))) (work))"); + SyntaxTree::parse_with_dialect(&labels.rewritten, Dialect::CommonLisp) + .expect("labels output must reparse as Common Lisp"); } } diff --git a/src/domain/package.rs b/src/domain/package.rs index 5ddb794c..8d7bc98d 100644 --- a/src/domain/package.rs +++ b/src/domain/package.rs @@ -2,6 +2,7 @@ use anyhow::{Context, Result}; +use crate::domain::dialect::Dialect; use crate::domain::mutation_safety::reject_common_lisp_reader_conditionals; use crate::domain::sexpr::SyntaxTree; @@ -28,8 +29,18 @@ use sort_options::defpackage_option_sort_edits; pub use sort_options::PackageOptionSortOrder; pub use types::*; +fn ensure_common_lisp_package_refactoring(dialect: Dialect) -> Result<()> { + anyhow::ensure!( + dialect == Dialect::CommonLisp, + "package refactoring currently supports only Common Lisp" + ); + Ok(()) +} + pub fn plan_add_export(request: AddExportRequest<'_>) -> Result { - let tree = SyntaxTree::parse(request.input).context("failed to parse input")?; + ensure_common_lisp_package_refactoring(request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("failed to parse input")?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; let edit = find_defpackage_export_edit(&tree, request.dialect, request.package, request.symbol)?; @@ -38,7 +49,7 @@ pub fn plan_add_export(request: AddExportRequest<'_>) -> Result { } else { replace_span(request.input, edit.insertion_span, &edit.replacement) }; - SyntaxTree::parse(&rewritten) + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) .context("add-export output is not a valid S-expression document")?; Ok(AddExportPlan { @@ -55,11 +66,13 @@ pub fn plan_add_export(request: AddExportRequest<'_>) -> Result { } pub fn plan_rename_package(request: RenamePackageRequest<'_>) -> Result { - let tree = SyntaxTree::parse(request.input).context("failed to parse input")?; + ensure_common_lisp_package_refactoring(request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("failed to parse input")?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; let occurrences = package_rename_occurrences(&tree, request.dialect, request.from, request.to)?; let rewritten = rewrite_package_occurrences(request.input, &occurrences); - SyntaxTree::parse(&rewritten) + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) .context("rename-package output is not a valid S-expression document")?; Ok(RenamePackagePlan { @@ -72,7 +85,9 @@ pub fn plan_rename_package(request: RenamePackageRequest<'_>) -> Result, ) -> Result { - let tree = SyntaxTree::parse(request.input).context("failed to parse input")?; + ensure_common_lisp_package_refactoring(request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("failed to parse input")?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; let edits = defpackage_export_sort_edits(request.input, &tree, request.dialect, request.package)?; @@ -86,7 +101,7 @@ pub fn plan_sort_package_exports( }) .collect::>(); let rewritten = rewrite_spans(request.input, &replacements); - SyntaxTree::parse(&rewritten) + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) .context("sort-package-exports output is not a valid S-expression document")?; let exports = edits @@ -113,7 +128,9 @@ pub fn plan_sort_package_exports( pub fn plan_sort_package_options( request: SortPackageOptionsRequest<'_>, ) -> Result { - let tree = SyntaxTree::parse(request.input).context("failed to parse input")?; + ensure_common_lisp_package_refactoring(request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("failed to parse input")?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; let edits = defpackage_option_sort_edits( request.input, @@ -132,7 +149,7 @@ pub fn plan_sort_package_options( }) .collect::>(); let rewritten = rewrite_spans(request.input, &replacements); - SyntaxTree::parse(&rewritten) + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) .context("sort-package-options output is not a valid S-expression document")?; let packages = edits @@ -157,7 +174,9 @@ pub fn plan_sort_package_options( pub fn plan_merge_package_options( request: MergePackageOptionsRequest<'_>, ) -> Result { - let tree = SyntaxTree::parse(request.input).context("failed to parse input")?; + ensure_common_lisp_package_refactoring(request.dialect)?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("failed to parse input")?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; let edits = defpackage_option_merge_edits(request.input, &tree, request.dialect, request.package)?; @@ -178,7 +197,7 @@ pub fn plan_merge_package_options( }) .collect::>(); let rewritten = rewrite_spans(request.input, &replacements); - SyntaxTree::parse(&rewritten) + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) .context("merge-package-options output is not a valid S-expression document")?; let merges = edits @@ -213,3 +232,152 @@ pub fn plan_merge_package_options( #[cfg(test)] mod tests; + +#[cfg(test)] +mod dialect_tests { + use anyhow::Result; + + use super::*; + use crate::domain::sexpr::SymbolName; + + const VALID_INPUT: &str = + "(defpackage demo (:use cl) (:export z) (:export a))\n(in-package demo)\n"; + const DIALECT_MATRIX: [(Dialect, bool); 7] = [ + (Dialect::CommonLisp, true), + (Dialect::EmacsLisp, false), + (Dialect::Scheme, false), + (Dialect::Clojure, false), + (Dialect::Janet, false), + (Dialect::Fennel, false), + (Dialect::Unknown, false), + ]; + + #[derive(Clone, Copy, Debug)] + enum PackageOperation { + AddExport, + RenamePackage, + SortExports, + SortOptions, + MergeOptions, + } + + impl PackageOperation { + fn run(self, input: &str, dialect: Dialect) -> Result { + let package = SymbolName::new("demo").expect("valid package name"); + + match self { + Self::AddExport => { + let symbol = SymbolName::new("z").expect("valid export name"); + plan_add_export(AddExportRequest { + input, + dialect, + package: Some(&package), + symbol: &symbol, + }) + .map(|plan| plan.rewritten) + } + Self::RenamePackage => { + let renamed = SymbolName::new("renamed").expect("valid package name"); + plan_rename_package(RenamePackageRequest { + input, + dialect, + from: &package, + to: &renamed, + }) + .map(|plan| plan.rewritten) + } + Self::SortExports => plan_sort_package_exports(SortPackageExportsRequest { + input, + dialect, + package: Some(&package), + }) + .map(|plan| plan.rewritten), + Self::SortOptions => plan_sort_package_options(SortPackageOptionsRequest { + input, + dialect, + package: Some(&package), + order: PackageOptionSortOrder::Canonical, + }) + .map(|plan| plan.rewritten), + Self::MergeOptions => plan_merge_package_options(MergePackageOptionsRequest { + input, + dialect, + package: Some(&package), + }) + .map(|plan| plan.rewritten), + } + } + } + + const OPERATIONS: [PackageOperation; 5] = [ + PackageOperation::AddExport, + PackageOperation::RenamePackage, + PackageOperation::SortExports, + PackageOperation::SortOptions, + PackageOperation::MergeOptions, + ]; + + fn assert_common_lisp_support_error( + operation: PackageOperation, + dialect: Dialect, + error: anyhow::Error, + ) { + assert!( + error.to_string().contains("supports only Common Lisp"), + "{operation:?} returned the wrong error for {dialect:?}: {error:#}" + ); + } + + #[test] + fn package_operations_follow_the_dialect_support_matrix() { + for operation in OPERATIONS { + for (dialect, supported) in DIALECT_MATRIX { + let result = operation.run(VALID_INPUT, dialect); + + if supported { + let rewritten = result.unwrap_or_else(|error| { + panic!("{operation:?} should support {dialect:?}: {error:#}") + }); + SyntaxTree::parse_with_dialect(&rewritten, dialect).unwrap_or_else(|error| { + panic!("{operation:?} output should reparse with {dialect:?}: {error:#}") + }); + } else { + assert_common_lisp_support_error( + operation, + dialect, + result.expect_err("unsupported dialect should fail"), + ); + } + } + } + } + + #[test] + fn unsupported_package_operations_reject_before_parsing() { + for operation in OPERATIONS { + for (dialect, supported) in DIALECT_MATRIX { + if supported { + continue; + } + + let error = operation + .run(")", dialect) + .expect_err("unsupported dialect should fail before parsing"); + assert_common_lisp_support_error(operation, dialect, error); + } + } + } + + #[test] + fn rename_package_preserves_common_lisp_closing_paren_character_literal() { + let input = "#\\)\n(defpackage demo (:export old))\n(in-package demo)\n"; + let rewritten = PackageOperation::RenamePackage + .run(input, Dialect::CommonLisp) + .expect("Common Lisp character literal should parse"); + + assert!(rewritten.starts_with("#\\)\n")); + assert!(rewritten.contains("(defpackage renamed")); + SyntaxTree::parse_with_dialect(&rewritten, Dialect::CommonLisp) + .expect("rewritten output should reparse as Common Lisp"); + } +} diff --git a/src/domain/progn.rs b/src/domain/progn.rs index b8100d44..59541183 100644 --- a/src/domain/progn.rs +++ b/src/domain/progn.rs @@ -27,7 +27,8 @@ pub struct FlattenPrognPlan { pub fn plan_flatten_progn(request: FlattenPrognRequest<'_>) -> Result { require_supported(request.dialect, "flatten-progn")?; - let tree = SyntaxTree::parse(request.input).context("flatten-progn input is not valid")?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("flatten-progn input is not valid")?; let form = tree.select_path(&request.path)?.view(); require_head(request.dialect, &form, "progn", "flatten-progn")?; if tree.has_comment_in(form.span) || contains_prefix(&form) { @@ -65,7 +66,8 @@ pub fn plan_flatten_progn(request: FlattenPrognRequest<'_>) -> Result, ) -> Result { require_supported(request.dialect, "eliminate-empty-binding-form")?; - let tree = SyntaxTree::parse(request.input) + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) .context("eliminate-empty-binding-form input is not valid")?; let form = tree.select_path(&request.path)?.view(); if form.kind != ExpressionKind::List @@ -138,7 +140,8 @@ pub fn plan_eliminate_empty_binding_form( ), }; let rewritten = replace_span(request.input, form.span, &replacement); - SyntaxTree::parse(&rewritten).context("eliminate-empty-binding-form output is not valid")?; + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) + .context("eliminate-empty-binding-form output is not valid")?; Ok(EliminateEmptyBindingFormPlan { dialect: request.dialect, path: request.path, @@ -150,7 +153,7 @@ pub fn plan_eliminate_empty_binding_form( }) } -fn require_supported(dialect: Dialect, operation: &str) -> Result<()> { +pub(crate) fn require_supported(dialect: Dialect, operation: &str) -> Result<()> { if matches!(dialect, Dialect::CommonLisp | Dialect::EmacsLisp) { Ok(()) } else { @@ -206,31 +209,78 @@ fn replace_span(input: &str, span: ByteSpan, replacement: &str) -> String { #[cfg(test)] mod tests { use super::*; + #[test] - fn flatten_round_trips_across_dialects() { - for dialect in [Dialect::CommonLisp, Dialect::EmacsLisp] { + fn flatten_round_trips_across_dialects_with_reader_literals_outside_target() { + for (dialect, reader_literal) in + [(Dialect::CommonLisp, r"#\)"), (Dialect::EmacsLisp, r"?\)")] + { + let input = format!("(progn a (progn b c)) {reader_literal}"); let plan = plan_flatten_progn(FlattenPrognRequest { - input: "(list (progn a (progn b c)))", + input: &input, dialect, - path: "0.1".parse().unwrap(), + path: "0".parse().expect("path"), }) - .unwrap(); + .expect("plan"); assert_eq!(plan.result_form_count, 3); - assert!(SyntaxTree::parse(&plan.rewritten).is_ok()); + assert_eq!(plan.rewritten, format!("(progn a b c) {reader_literal}")); + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect).expect("rewritten input"); } } + #[test] - fn eliminate_preserves_body_cardinality() { - let plan = plan_eliminate_empty_binding_form(EliminateEmptyBindingFormRequest { - input: "(if ok (let () a b) nil)", - dialect: Dialect::CommonLisp, - path: "0.2".parse().unwrap(), - }) - .unwrap(); - assert_eq!(plan.body_form_count, 2); - assert!(plan.introduced_progn); - assert!(SyntaxTree::parse(&plan.rewritten).is_ok()); + fn eliminate_round_trips_across_dialects_with_reader_literals_outside_target() { + for (dialect, reader_literal) in + [(Dialect::CommonLisp, r"#\)"), (Dialect::EmacsLisp, r"?\)")] + { + let input = format!("(let () a b) {reader_literal}"); + let plan = plan_eliminate_empty_binding_form(EliminateEmptyBindingFormRequest { + input: &input, + dialect, + path: "0".parse().expect("path"), + }) + .expect("plan"); + assert_eq!(plan.body_form_count, 2); + assert!(plan.introduced_progn); + assert_eq!(plan.rewritten, format!("(progn a b) {reader_literal}")); + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect).expect("rewritten input"); + } + } + + #[test] + fn rejects_unsupported_dialects_before_parsing_for_both_operations() { + for dialect in [ + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ] { + let flatten_error = plan_flatten_progn(FlattenPrognRequest { + input: ")", + dialect, + path: "0".parse().expect("path"), + }) + .expect_err("dialect must be rejected"); + assert_eq!( + flatten_error.to_string(), + "flatten-progn supports only Common Lisp and Emacs Lisp" + ); + + let eliminate_error = + plan_eliminate_empty_binding_form(EliminateEmptyBindingFormRequest { + input: ")", + dialect, + path: "0".parse().expect("path"), + }) + .expect_err("dialect must be rejected"); + assert_eq!( + eliminate_error.to_string(), + "eliminate-empty-binding-form supports only Common Lisp and Emacs Lisp" + ); + } } + #[test] fn rejects_unsupported_forms() { assert!( diff --git a/src/domain/remove_unused_binding/mod.rs b/src/domain/remove_unused_binding/mod.rs index c30ad760..0ee6e7f9 100644 --- a/src/domain/remove_unused_binding/mod.rs +++ b/src/domain/remove_unused_binding/mod.rs @@ -34,8 +34,11 @@ pub fn plan_remove_unused_binding( if request.name.is_none() && !request.all_bindings { anyhow::bail!("remove-unused-binding requires --name or --all-bindings"); } + if request.dialect == Dialect::Unknown { + anyhow::bail!("remove-unused-binding does not support dialect unknown"); + } - let input_tree = SyntaxTree::parse(request.input) + let input_tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) .context("remove-unused-binding input is not a valid S-expression document")?; reject_common_lisp_reader_conditionals(&input_tree, request.dialect)?; let parsed_target = input_tree @@ -54,7 +57,7 @@ pub fn plan_remove_unused_binding( request.all_bindings, )?; let rewritten = replace_span(request.input, parts.form_span, &parts.replacement); - SyntaxTree::parse(&rewritten) + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) .context("remove-unused-binding output is not a valid S-expression document")?; let bindings = parts @@ -130,7 +133,7 @@ fn remove_unused_binding_parts( let binding_form = &target.children[1]; let candidates = binding_removal_candidates(dialect, refactor_form, binding_form)?; - let input_tree = SyntaxTree::parse(input) + let input_tree = SyntaxTree::parse_with_dialect(input, dialect) .context("remove-unused-binding input is not a valid S-expression document")?; let selected = if all_bindings { let mut unused = Vec::new(); @@ -191,7 +194,7 @@ fn remove_unused_binding_parts( })?; let candidate = candidates .iter() - .find(|candidate| common_lisp_symbol_reference_eq(&candidate.name, name.as_str())) + .find(|candidate| binding_name_matches(dialect, &candidate.name, name.as_str())) .with_context(|| { format!( "binding {} was not found in selected binding form", @@ -241,7 +244,7 @@ fn remove_unused_binding_parts( .map(|binding| (binding.binding_span, String::new())) .collect(), )?; - format_single_replacement_form(&replacement)? + format_single_replacement_form(&replacement, dialect)? }; Ok(RemoveUnusedBindingParts { @@ -271,8 +274,20 @@ fn ensure_variable_binding_form_consistency( Ok(()) } -fn format_single_replacement_form(input: &str) -> Result { - let tree = SyntaxTree::parse(input) +fn binding_name_matches(dialect: Dialect, candidate: &str, expected: &str) -> bool { + match dialect { + Dialect::CommonLisp => common_lisp_symbol_reference_eq(candidate, expected), + Dialect::EmacsLisp + | Dialect::Scheme + | Dialect::Clojure + | Dialect::Janet + | Dialect::Fennel => candidate == expected, + Dialect::Unknown => false, + } +} + +fn format_single_replacement_form(input: &str, dialect: Dialect) -> Result { + let tree = SyntaxTree::parse_with_dialect(input, dialect) .context("remove-unused-binding replacement is not a valid S-expression form")?; if tree.root_children().len() != 1 { anyhow::bail!("remove-unused-binding replacement must contain exactly one form"); diff --git a/src/domain/remove_unused_binding/tests/mod.rs b/src/domain/remove_unused_binding/tests/mod.rs index 96824b19..900de8a8 100644 --- a/src/domain/remove_unused_binding/tests/mod.rs +++ b/src/domain/remove_unused_binding/tests/mod.rs @@ -24,6 +24,13 @@ fn target_at(input: &str, path: &str) -> ExpressionView { .view() } +fn target_at_with_dialect(input: &str, path: &str, dialect: Dialect) -> ExpressionView { + let tree = SyntaxTree::parse_with_dialect(input, dialect).expect("parse"); + tree.select_path(&path.parse::().expect("path")) + .expect("select") + .view() +} + fn plan_remove_unused_binding_for( input: &str, dialect: Dialect, @@ -38,7 +45,7 @@ fn plan_remove_unused_binding_for( input, dialect, path: parsed_path, - target: target_at(input, path.unwrap_or("0")), + target: target_at_with_dialect(input, path.unwrap_or("0"), dialect), name: symbol.as_ref(), all_bindings, allow_drop_value, @@ -58,7 +65,7 @@ fn remove_unused_binding_error( input, dialect, path: None, - target: target(input), + target: target_at_with_dialect(input, "0", dialect), name: symbol.as_ref(), all_bindings, allow_drop_value, @@ -81,7 +88,7 @@ fn remove_unused_binding_error_for( input, dialect, path: parsed_path, - target: target_at(input, path.unwrap_or("0")), + target: target_at_with_dialect(input, path.unwrap_or("0"), dialect), name: symbol.as_ref(), all_bindings, allow_drop_value, @@ -109,6 +116,110 @@ fn rejects_target_that_does_not_match_input() { assert!(error.to_string().contains("does not match the input")); } +#[test] +fn supports_known_dialects_and_rejects_unknown() { + let fixtures = [ + (Dialect::CommonLisp, "(let ((unused 1) (used 2)) used)"), + (Dialect::EmacsLisp, "(let ((unused 1) (used 2)) used)"), + (Dialect::Scheme, "(let ((unused 1) (used 2)) used)"), + (Dialect::Clojure, "(let [unused 1 used 2] used)"), + (Dialect::Janet, "(let [unused 1 used 2] used)"), + (Dialect::Fennel, "(let [unused 1 used 2] used)"), + ]; + + for (dialect, input) in fixtures { + let plan = + plan_remove_unused_binding_for(input, dialect, None, Some("unused"), false, true); + + assert_eq!(plan.binding_name.as_deref(), Some("unused")); + assert!(!plan.rewritten.contains("unused")); + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect).expect("reparse rewritten output"); + } + + let error = remove_unused_binding_error( + "(let ((unused 1) (used 2)) used)", + Dialect::Unknown, + Some("unused"), + false, + true, + ); + assert!(error.contains("does not support dialect unknown")); +} + +#[test] +fn unknown_dialect_fails_before_parsing_malformed_input() { + let symbol = symbol("unused"); + let error = plan_remove_unused_binding(RemoveUnusedBindingRequest { + input: "(", + dialect: Dialect::Unknown, + path: None, + target: target("(let ((unused 1)) 2)"), + name: Some(&symbol), + all_bindings: false, + allow_drop_value: true, + }) + .expect_err("unknown dialect must fail") + .to_string(); + + assert!(error.contains("does not support dialect unknown")); +} + +#[test] +fn non_common_lisp_binding_names_are_case_sensitive() { + let fixtures = [ + (Dialect::EmacsLisp, "(let ((foo 1) (used 2)) used)"), + (Dialect::Scheme, "(let ((foo 1) (used 2)) used)"), + (Dialect::Clojure, "(let [foo 1 used 2] used)"), + (Dialect::Janet, "(let [foo 1 used 2] used)"), + (Dialect::Fennel, "(let [foo 1 used 2] used)"), + ]; + + for (dialect, input) in fixtures { + let error = remove_unused_binding_error(input, dialect, Some("FOO"), false, true); + assert!( + error.contains("binding FOO was not found"), + "{dialect:?}: {error}" + ); + } + + let common_lisp = common_lisp_plan("(let ((foo 1) (used 2)) used)", Some("FOO"), false, true); + assert_eq!(common_lisp.binding_name.as_deref(), Some("foo")); +} + +#[test] +fn preserves_dialect_reader_forms_during_rewrite() { + let fixtures = [ + ( + Dialect::CommonLisp, + "(let ((unused 1) (used (list #\\) #:done #x2a))) used)", + &["#\\)", "#:done", "#x2a"][..], + ), + ( + Dialect::EmacsLisp, + "(let ((unused 1) (used (list ?\\)))) used)", + &["?\\)"][..], + ), + ( + Dialect::Clojure, + "(let [unused 1 used (list #foo/bar #:person{:x 1})] used)", + &["#foo/bar", "#:person{:x 1}"][..], + ), + ]; + + for (dialect, input, preserved) in fixtures { + let plan = + plan_remove_unused_binding_for(input, dialect, None, Some("unused"), false, true); + + for reader_form in preserved { + assert!( + plan.rewritten.contains(reader_form), + "{dialect:?}: {reader_form}" + ); + } + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect).expect("reparse rewritten output"); + } +} + fn common_lisp_plan( input: &str, name: Option<&str>, diff --git a/src/domain/remove_unused_control.rs b/src/domain/remove_unused_control.rs index dbc3bc5f..e78ed563 100644 --- a/src/domain/remove_unused_control.rs +++ b/src/domain/remove_unused_control.rs @@ -31,7 +31,8 @@ pub fn plan_remove_unused_block( request: RemoveUnusedControlRequest<'_>, ) -> Result { prepare(&request, "remove-unused-block")?; - let tree = SyntaxTree::parse(request.input).context("input is not valid")?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("input is not valid")?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; require_known_expression_context(&tree, &request.path)?; let form = tree.select_path(&request.path)?.view(); @@ -61,7 +62,8 @@ pub fn plan_remove_unused_tag( request: RemoveUnusedControlRequest<'_>, ) -> Result { prepare(&request, "remove-unused-tag")?; - let tree = SyntaxTree::parse(request.input).context("input is not valid")?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("input is not valid")?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; let form = tree.select_path(&request.path)?.view(); reject_unsafe(&tree, &form, "remove-unused-tag")?; @@ -106,7 +108,8 @@ fn finish( replacement: &str, ) -> Result { let rewritten = replace_span(request.input, span, replacement); - SyntaxTree::parse(&rewritten).context("rewritten output is not valid")?; + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) + .context("rewritten output is not valid")?; Ok(RemoveUnusedControlPlan { dialect: request.dialect, path: request.path, @@ -315,24 +318,38 @@ fn contains_head(view: &ExpressionView, expected: &str) -> bool { #[cfg(test)] mod tests { use super::*; + fn request<'a>(input: &'a str, path: &str, name: &str) -> RemoveUnusedControlRequest<'a> { + request_with_dialect(input, Dialect::CommonLisp, path, name) + } + + fn request_with_dialect<'a>( + input: &'a str, + dialect: Dialect, + path: &str, + name: &str, + ) -> RemoveUnusedControlRequest<'a> { RemoveUnusedControlRequest { input, - dialect: Dialect::CommonLisp, + dialect, path: path.parse().expect("path"), name: name.to_owned(), } } + #[test] - fn removes_unused_block_with_multiple_body_forms() { + fn removes_unused_block_with_common_lisp_reader_literal_outside_target() { let plan = plan_remove_unused_block(request( - "(if ok (block out (first) (second)) nil)", + r"(if ok (block out (first) (second)) nil) #\)", "0.2", "out", )) - .unwrap(); - assert_eq!(plan.rewritten, "(if ok (progn (first) (second)) nil)"); + .expect("plan"); + assert_eq!(plan.rewritten, r"(if ok (progn (first) (second)) nil) #\)"); + SyntaxTree::parse_with_dialect(&plan.rewritten, Dialect::CommonLisp) + .expect("rewritten input"); } + #[test] fn ignores_shadowed_return_from() { plan_remove_unused_block(request( @@ -342,6 +359,7 @@ mod tests { )) .unwrap(); } + #[test] fn rejects_referenced_block_and_unknown_context() { assert!( @@ -354,6 +372,7 @@ mod tests { ); assert!(plan_remove_unused_block(request("(list (block out 1))", "0.1", "out")).is_err()); } + #[test] fn removes_symbol_and_integer_tags() { assert_eq!( @@ -369,9 +388,47 @@ mod tests { "(tagbody (print 1))" ); } + + #[test] + fn removes_unused_tag_with_common_lisp_reader_literal_outside_target() { + let plan = plan_remove_unused_tag(request(r"(tagbody start (print 1)) #\)", "0", "start")) + .expect("plan"); + assert_eq!(plan.rewritten, r"(tagbody (print 1)) #\)"); + SyntaxTree::parse_with_dialect(&plan.rewritten, Dialect::CommonLisp) + .expect("rewritten input"); + } + #[test] fn rejects_referenced_tag_but_ignores_nested_shadow() { assert!(plan_remove_unused_tag(request("(tagbody x (go x))", "0", "x")).is_err()); plan_remove_unused_tag(request("(tagbody x (tagbody x (go x)))", "0", "x")).unwrap(); } + + #[test] + fn rejects_unsupported_dialects_before_parsing_for_both_operations() { + for dialect in [ + Dialect::EmacsLisp, + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ] { + let block_error = + plan_remove_unused_block(request_with_dialect(")", dialect, "0", "out")) + .expect_err("dialect must be rejected"); + assert_eq!( + block_error.to_string(), + "remove-unused-block supports only Common Lisp" + ); + + let tag_error = + plan_remove_unused_tag(request_with_dialect(")", dialect, "0", "start")) + .expect_err("dialect must be rejected"); + assert_eq!( + tag_error.to_string(), + "remove-unused-tag supports only Common Lisp" + ); + } + } } diff --git a/src/domain/remove_unused_definition.rs b/src/domain/remove_unused_definition.rs index 06d51295..689ed7c3 100644 --- a/src/domain/remove_unused_definition.rs +++ b/src/domain/remove_unused_definition.rs @@ -1,5 +1,6 @@ use anyhow::{Context, Result}; +use crate::domain::dialect::Dialect; use crate::domain::mutation_safety::reject_overlapping_common_lisp_reader_time_forms; use crate::domain::sexpr::SyntaxTree; @@ -23,6 +24,15 @@ pub use types::{ pub fn plan_remove_unused_definitions( request: RemoveUnusedDefinitionsRequest, ) -> Result { + for file in &request.files { + if file.dialect == Dialect::Unknown { + anyhow::bail!( + "remove-unused-definition does not support dialect unknown: {}", + file.path.display() + ); + } + } + let exported_symbols = collect_exported_symbol_index(&request.package_definitions); let unused_reports = collect_unused_definition_candidates(&request.files)?; let mut files = Vec::with_capacity(request.files.len()); @@ -62,6 +72,12 @@ fn plan_file_removals( ) -> Result { let mut removals = Vec::new(); let mut skipped = Vec::new(); + let empty_exported_symbols = std::collections::HashMap::new(); + let exported_symbols = if file.dialect == Dialect::CommonLisp { + exported_symbols + } else { + &empty_exported_symbols + }; for item in report .definitions @@ -137,7 +153,7 @@ fn rewrite_file_without_unused_definitions( for removal in removals { let expanded = expand_definition_removal(&rewritten, removal.definition.span); removal.removal_span = expanded; - let tree = SyntaxTree::parse(&rewritten).with_context(|| { + let tree = SyntaxTree::parse_with_dialect(&rewritten, file.dialect).with_context(|| { format!( "file would become invalid before removing unused definitions: {}", file.path.display() @@ -147,7 +163,7 @@ fn rewrite_file_without_unused_definitions( rewritten = replace_span(&rewritten, expanded, ""); } - SyntaxTree::parse(&rewritten).with_context(|| { + SyntaxTree::parse_with_dialect(&rewritten, file.dialect).with_context(|| { format!( "file would become invalid after removing unused definitions: {}", file.path.display() diff --git a/src/domain/remove_unused_definition/candidates.rs b/src/domain/remove_unused_definition/candidates.rs index b8fc2d8d..0f8ba61e 100644 --- a/src/domain/remove_unused_definition/candidates.rs +++ b/src/domain/remove_unused_definition/candidates.rs @@ -4,6 +4,7 @@ use crate::domain::common_lisp::common_lisp_symbol_reference_needle; use crate::domain::definition_reference::{ collect_package_form_spans, collect_reference_needles, collect_symbol_references, }; +use crate::domain::dialect::Dialect; use crate::domain::remove_unused_definition::types::{ RemoveUnusedDefinitionInputFile, UnusedDefinitionDefinition, }; @@ -26,10 +27,19 @@ pub(super) struct DefinitionReference; pub(super) fn collect_unused_definition_candidates( files: &[RemoveUnusedDefinitionInputFile], ) -> Result> { + for file in files { + if file.dialect == Dialect::Unknown { + anyhow::bail!( + "remove-unused-definition does not support dialect unknown: {}", + file.path.display() + ); + } + } + let parsed_files = files .iter() .map(|file| -> Result<_> { - let tree = SyntaxTree::parse(&file.text) + let tree = SyntaxTree::parse_with_dialect(&file.text, file.dialect) .with_context(|| format!("failed to parse {}", file.path.display()))?; Ok((file, tree.root_view())) }) diff --git a/src/domain/remove_unused_definition/tests/basic.rs b/src/domain/remove_unused_definition/tests/basic.rs index f6ef0192..5a5c1a9b 100644 --- a/src/domain/remove_unused_definition/tests/basic.rs +++ b/src/domain/remove_unused_definition/tests/basic.rs @@ -23,7 +23,8 @@ fn plans_private_unused_definition_removal() { assert_eq!(plan.removal_count, 1); assert_eq!(plan.skipped_count, 0); assert!(!plan.files[0].rewritten.contains(stale_form)); - SyntaxTree::parse(&plan.files[0].rewritten).expect("rewrite must stay parseable"); + SyntaxTree::parse_with_dialect(&plan.files[0].rewritten, plan.files[0].dialect) + .expect("rewrite must stay parseable"); } #[test] @@ -215,7 +216,8 @@ fn ignores_vector_literals_that_do_not_overlap_the_removal_span() { assert_eq!(plan.removal_count, 1); assert!(!plan.files[0].rewritten.contains(stale_form)); assert!(plan.files[0].rewritten.contains(vector_form)); - SyntaxTree::parse(&plan.files[0].rewritten).expect("rewrite must stay parseable"); + SyntaxTree::parse_with_dialect(&plan.files[0].rewritten, plan.files[0].dialect) + .expect("rewrite must stay parseable"); } #[test] @@ -250,3 +252,164 @@ fn keeps_definitions_referenced_from_another_input_file() { assert!(plan.files[0].rewritten.contains(shared_form)); assert_eq!(plan.files[1].rewritten, consumer_text); } + +#[test] +fn removes_unused_definitions_for_every_known_dialect_and_reparses_outputs() { + let fixtures = [ + ( + "common-lisp.lisp", + Dialect::CommonLisp, + "(defun cl-unused () 1)\n", + "(defun cl-unused () 1)", + "defun", + "cl-unused", + ), + ( + "emacs-lisp.el", + Dialect::EmacsLisp, + "(defun el-unused () 1)\n", + "(defun el-unused () 1)", + "defun", + "el-unused", + ), + ( + "scheme.scm", + Dialect::Scheme, + "(define scheme-unused (lambda () 1))\n", + "(define scheme-unused (lambda () 1))", + "define", + "scheme-unused", + ), + ( + "clojure.clj", + Dialect::Clojure, + "(defn clj-unused [] 1)\n", + "(defn clj-unused [] 1)", + "defn", + "clj-unused", + ), + ( + "janet.janet", + Dialect::Janet, + "(defn janet-unused [] 1)\n", + "(defn janet-unused [] 1)", + "defn", + "janet-unused", + ), + ( + "fennel.fnl", + Dialect::Fennel, + "(fn fennel-unused [] 1)\n", + "(fn fennel-unused [] 1)", + "fn", + "fennel-unused", + ), + ]; + let files = fixtures + .iter() + .map(|(path, dialect, text, form, head, name)| { + let mut item = definition(text, form, name, DefinitionCategory::Function); + item.head = (*head).to_owned(); + item.package = None; + file_with_dialect(PathBuf::from(path), *dialect, None, text, vec![item]) + }) + .collect(); + let request = RemoveUnusedDefinitionsRequest { + files, + package_definitions: Vec::new(), + include_protected: false, + include_exported: false, + }; + + let plan = plan_remove_unused_definitions(request).expect("plan should build"); + + assert_eq!(plan.candidate_count, fixtures.len()); + assert_eq!(plan.removal_count, fixtures.len()); + assert_eq!(plan.skipped_count, 0); + for ((_, dialect, _, form, _, _), file) in fixtures.iter().zip(&plan.files) { + assert_eq!(file.dialect, *dialect); + assert!(!file.rewritten.contains(form)); + SyntaxTree::parse_with_dialect(&file.rewritten, file.dialect) + .expect("rewritten fixture must parse with its dialect"); + } +} + +#[test] +fn uses_the_input_dialect_for_reader_syntax_during_removal() { + let text = "(defun el-reader () [?\\)])\n"; + let form = "(defun el-reader () [?\\)])"; + assert!(SyntaxTree::parse_with_dialect(text, Dialect::EmacsLisp).is_ok()); + assert!(SyntaxTree::parse_with_dialect(text, Dialect::CommonLisp).is_err()); + let mut item = definition(text, form, "el-reader", DefinitionCategory::Function); + item.package = None; + let request = RemoveUnusedDefinitionsRequest { + files: vec![file_with_dialect( + PathBuf::from("reader.el"), + Dialect::EmacsLisp, + None, + text, + vec![item], + )], + package_definitions: Vec::new(), + include_protected: false, + include_exported: false, + }; + + let plan = plan_remove_unused_definitions(request).expect("plan should build"); + + assert_eq!(plan.removal_count, 1); + assert!(!plan.files[0].rewritten.contains(form)); +} + +#[test] +fn common_lisp_symbol_matching_is_package_aware_but_scheme_is_exact() { + let common_lisp_text = "(defun Foo () 1)\n(pkg:foo)\n"; + let common_lisp_form = "(defun Foo () 1)"; + let common_lisp_request = RemoveUnusedDefinitionsRequest { + files: vec![file_with_dialect( + PathBuf::from("symbols.lisp"), + Dialect::CommonLisp, + Some("app"), + common_lisp_text, + vec![definition( + common_lisp_text, + common_lisp_form, + "Foo", + DefinitionCategory::Function, + )], + )], + package_definitions: Vec::new(), + include_protected: false, + include_exported: false, + }; + let scheme_text = "(define Foo (lambda () 1))\n(pkg:foo)\n"; + let scheme_form = "(define Foo (lambda () 1))"; + let mut scheme_definition = definition( + scheme_text, + scheme_form, + "Foo", + DefinitionCategory::Function, + ); + scheme_definition.head = "define".to_owned(); + scheme_definition.package = None; + let scheme_request = RemoveUnusedDefinitionsRequest { + files: vec![file_with_dialect( + PathBuf::from("symbols.scm"), + Dialect::Scheme, + None, + scheme_text, + vec![scheme_definition], + )], + package_definitions: Vec::new(), + include_protected: false, + include_exported: false, + }; + + let common_lisp_plan = + plan_remove_unused_definitions(common_lisp_request).expect("plan should build"); + let scheme_plan = plan_remove_unused_definitions(scheme_request).expect("plan should build"); + + assert_eq!(common_lisp_plan.candidate_count, 0); + assert_eq!(scheme_plan.candidate_count, 1); + assert_eq!(scheme_plan.removal_count, 1); +} diff --git a/src/domain/remove_unused_definition/tests/mod.rs b/src/domain/remove_unused_definition/tests/mod.rs index 0ad887d9..945b1e7c 100644 --- a/src/domain/remove_unused_definition/tests/mod.rs +++ b/src/domain/remove_unused_definition/tests/mod.rs @@ -26,11 +26,21 @@ fn file_with_text( text: &str, definitions: Vec, ) -> RemoveUnusedDefinitionInputFile { - let tree = SyntaxTree::parse(text).expect("fixture must parse"); + file_with_dialect(path, Dialect::CommonLisp, Some("app"), text, definitions) +} + +fn file_with_dialect( + path: PathBuf, + dialect: Dialect, + package: Option<&str>, + text: &str, + definitions: Vec, +) -> RemoveUnusedDefinitionInputFile { + let tree = SyntaxTree::parse_with_dialect(text, dialect).expect("fixture must parse"); RemoveUnusedDefinitionInputFile { path, - dialect: Dialect::CommonLisp, - package: Some("app".to_owned()), + dialect, + package: package.map(str::to_owned), definitions, atoms: tree.atom_occurrences(), text: text.to_owned(), diff --git a/src/domain/remove_unused_definition/tests/policy.rs b/src/domain/remove_unused_definition/tests/policy.rs index d9ac6ddd..7ab4c71d 100644 --- a/src/domain/remove_unused_definition/tests/policy.rs +++ b/src/domain/remove_unused_definition/tests/policy.rs @@ -169,3 +169,38 @@ fn skips_exported_definition_when_file_uses_package_nickname() { SkippedDefinitionRemovalReason::ExportedDefinition ); } + +#[test] +fn common_lisp_package_exports_do_not_protect_non_common_lisp_definitions() { + let text = "(define public-entry (lambda () 1))\n"; + let form = "(define public-entry (lambda () 1))"; + let mut public_entry = definition(text, form, "public-entry", DefinitionCategory::Function); + public_entry.head = "define".to_owned(); + let request = RemoveUnusedDefinitionsRequest { + files: vec![file_with_dialect( + PathBuf::from("public.scm"), + Dialect::Scheme, + Some("app"), + text, + vec![public_entry], + )], + package_definitions: vec![PackageDefinitionReport { + path: "0".to_owned(), + span: ByteSpan::new(ByteOffset::new(0), ByteOffset::new(0)), + name: "#:app".to_owned(), + nicknames: Vec::new(), + uses: Vec::new(), + exports: vec!["#:public-entry".to_owned()], + imports: Vec::new(), + option_count: 1, + }], + include_protected: false, + include_exported: false, + }; + + let plan = plan_remove_unused_definitions(request).expect("plan should build"); + + assert_eq!(plan.candidate_count, 1); + assert_eq!(plan.removal_count, 1); + assert_eq!(plan.skipped_count, 0); +} diff --git a/src/domain/remove_unused_definition/tests/validation.rs b/src/domain/remove_unused_definition/tests/validation.rs index cc86595e..598048aa 100644 --- a/src/domain/remove_unused_definition/tests/validation.rs +++ b/src/domain/remove_unused_definition/tests/validation.rs @@ -24,6 +24,40 @@ fn rejects_unparseable_input_files_instead_of_panicking() { ); } +#[test] +fn rejects_unknown_dialect_before_parsing_any_file() { + let request = RemoveUnusedDefinitionsRequest { + files: vec![ + RemoveUnusedDefinitionInputFile { + path: PathBuf::from("broken.lisp"), + dialect: Dialect::CommonLisp, + package: Some("app".to_owned()), + definitions: Vec::new(), + atoms: Vec::new(), + text: "(defun broken ()".to_owned(), + }, + RemoveUnusedDefinitionInputFile { + path: PathBuf::from("unknown.lisp"), + dialect: Dialect::Unknown, + package: None, + definitions: Vec::new(), + atoms: Vec::new(), + text: "()".to_owned(), + }, + ], + package_definitions: Vec::new(), + include_protected: false, + include_exported: false, + }; + + let error = plan_remove_unused_definitions(request).expect_err("invalid input must fail"); + + assert_eq!( + error.to_string(), + "remove-unused-definition does not support dialect unknown: unknown.lisp" + ); +} + #[test] fn rejects_invalid_definition_symbols_instead_of_panicking() { let text = "(in-package #:app)\n(defun still-valid () 1)\n"; @@ -42,7 +76,7 @@ fn rejects_invalid_definition_symbols_instead_of_panicking() { body_form_count: Some(1), package: Some("app".to_owned()), }], - atoms: SyntaxTree::parse(text) + atoms: SyntaxTree::parse_with_dialect(text, Dialect::CommonLisp) .expect("fixture must parse") .atom_occurrences(), text: text.to_owned(), diff --git a/src/domain/rename.rs b/src/domain/rename.rs index 69ef0941..545f32eb 100644 --- a/src/domain/rename.rs +++ b/src/domain/rename.rs @@ -2,6 +2,7 @@ mod at; mod binding; +mod call_identity; mod function; mod macrolet; mod reader; @@ -28,6 +29,7 @@ pub use crate::domain::rename_types::{ FunctionCallScope, ReplaceFunctionCallsScope, UnwrapFunctionCallsScope, WrapFunctionCallsScope, }; +pub(crate) use at::supports_rename_at_dialect; pub use at::{RenameAtError, RenameAtNamespace, RenameAtPlan, RenameAtRequest, plan_rename_at}; pub use function::{collect_callable_definition_renames, collect_function_call_head_renames}; pub use macrolet::{ @@ -56,7 +58,8 @@ pub use wrap::{ }; pub fn plan_rename_function(request: RenameFunctionRequest<'_>) -> Result { - let tree = SyntaxTree::parse(request.input).context("failed to parse input")?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("failed to parse input")?; let definitions = collect_callable_definition_renames(&tree, request.dialect, &request.from, &request.to)?; let calls = @@ -67,7 +70,8 @@ pub fn plan_rename_function(request: RenameFunctionRequest<'_>) -> Result>(); let rewritten = apply_byte_span_edits(request.input, edits)?; - SyntaxTree::parse(&rewritten).context("renamed output is not a valid S-expression document")?; + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) + .context("renamed output is not a valid S-expression document")?; Ok(RenameFunctionPlan { dialect: request.dialect, @@ -79,7 +83,8 @@ pub fn plan_rename_function(request: RenameFunctionRequest<'_>) -> Result) -> Result { - let tree = SyntaxTree::parse(request.input).context("failed to parse input")?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("failed to parse input")?; let definitions = collect_macrolet_binding_renames(&tree, request.dialect, &request.from, &request.to)?; let calls = @@ -90,7 +95,8 @@ pub fn plan_rename_macrolet(request: RenameMacroletRequest<'_>) -> Result>(); let rewritten = apply_byte_span_edits(request.input, edits)?; - SyntaxTree::parse(&rewritten).context("renamed output is not a valid S-expression document")?; + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) + .context("renamed output is not a valid S-expression document")?; Ok(RenameMacroletPlan { dialect: request.dialect, @@ -104,7 +110,8 @@ pub fn plan_rename_macrolet(request: RenameMacroletRequest<'_>) -> Result, ) -> Result { - let tree = SyntaxTree::parse(request.input).context("failed to parse input")?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("failed to parse input")?; let definitions = collect_define_symbol_macro_definition_renames( &tree, request.dialect, @@ -123,7 +130,8 @@ pub fn plan_rename_symbol_macro( .map(|occurrence| (occurrence.span, occurrence.replacement.clone())) .collect::>(); let rewritten = apply_byte_span_edits(request.input, edits)?; - SyntaxTree::parse(&rewritten).context("renamed output is not a valid S-expression document")?; + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) + .context("renamed output is not a valid S-expression document")?; Ok(RenameSymbolMacroPlan { dialect: request.dialect, @@ -137,7 +145,8 @@ pub fn plan_rename_symbol_macro( pub fn plan_rename_local_function( request: RenameLocalFunctionRequest<'_>, ) -> Result { - let tree = SyntaxTree::parse(request.input).context("failed to parse input")?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("failed to parse input")?; let definitions = collect_local_function_binding_renames(&tree, request.dialect, &request.from, &request.to)?; let calls = collect_local_function_call_head_renames( @@ -152,7 +161,8 @@ pub fn plan_rename_local_function( .map(|occurrence| (occurrence.span, occurrence.replacement.clone())) .collect::>(); let rewritten = apply_byte_span_edits(request.input, edits)?; - SyntaxTree::parse(&rewritten).context("renamed output is not a valid S-expression document")?; + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) + .context("renamed output is not a valid S-expression document")?; Ok(RenameLocalFunctionPlan { dialect: request.dialect, @@ -164,7 +174,8 @@ pub fn plan_rename_local_function( } pub fn plan_rename_in_form(request: RenameInFormRequest<'_>) -> Result { - let tree = SyntaxTree::parse(request.input).context("failed to parse input")?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("failed to parse input")?; let path = match &request.target { RenameTarget::Path(path) => Some(path.clone()), RenameTarget::Offset(_) => None, @@ -179,7 +190,8 @@ pub fn plan_rename_in_form(request: RenameInFormRequest<'_>) -> Result>(); let rewritten = apply_byte_span_edits(request.input, edits)?; - SyntaxTree::parse(&rewritten).context("renamed output is not a valid S-expression document")?; + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) + .context("renamed output is not a valid S-expression document")?; Ok(RenameInFormPlan { dialect: request.dialect, @@ -194,13 +206,18 @@ pub fn plan_rename_in_form(request: RenameInFormRequest<'_>) -> Result) -> Result { - let tree = SyntaxTree::parse(request.input).context("failed to parse input")?; + let semantic = request + .dialect + .verify_rename_binding() + .context("rename-binding is not supported for this dialect")?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("failed to parse input")?; let path = match &request.target { RenameTarget::Path(path) => Some(path.clone()), RenameTarget::Offset(_) => None, }; let view = select_rename_target(&tree, &request.target)?.view(); - let parts = binding_rename_parts(request.dialect, &view, &request.from, request.input)?; + let parts = binding_rename_parts(semantic, &view, &request.from, request.input)?; let mut edits = Vec::with_capacity(parts.reference_spans.len() + 1); edits.push(( @@ -214,7 +231,8 @@ pub fn plan_rename_binding(request: RenameBindingRequest<'_>) -> Result Result> { + let semantic = Dialect::CommonLisp + .verify_rename_binding() + .expect("Common Lisp rename-binding semantics are verified"); let selected_span = tree.select_path(path)?.span(); let mut candidates = Vec::new(); for view in ancestor_views(root_view, path)?.into_iter().rev() { - let Ok(parts) = binding_rename_parts(Dialect::CommonLisp, view, from, input) else { + let Ok(parts) = binding_rename_parts(semantic, view, from, input) else { continue; }; let reference_spans: Vec<_> = parts diff --git a/src/domain/rename/at/mod.rs b/src/domain/rename/at/mod.rs index 19f4a003..571bac73 100644 --- a/src/domain/rename/at/mod.rs +++ b/src/domain/rename/at/mod.rs @@ -16,8 +16,12 @@ pub use error::RenameAtError; use selection::AtomPathIndex; pub use types::{RenameAtNamespace, RenameAtPlan, RenameAtRequest}; +pub(crate) const fn supports_rename_at_dialect(dialect: Dialect) -> bool { + matches!(dialect, Dialect::CommonLisp) +} + pub fn plan_rename_at(request: RenameAtRequest<'_>) -> Result { - if request.dialect != Dialect::CommonLisp { + if !supports_rename_at_dialect(request.dialect) { return Err(RenameAtError::UnsupportedDialect.into()); } if request.at.get() >= request.input.len() || !request.input.is_char_boundary(request.at.get()) @@ -25,7 +29,8 @@ pub fn plan_rename_at(request: RenameAtRequest<'_>) -> Result { return Err(RenameAtError::InvalidSelection.into()); } - let tree = SyntaxTree::parse(request.input).context("failed to parse input")?; + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("failed to parse input")?; reject_common_lisp_reader_conditionals(&tree, request.dialect).map_err(RenameAtError::from)?; let atom_occurrences = tree.atom_occurrence_index(); let atom_paths = AtomPathIndex::new(&atom_occurrences); @@ -77,7 +82,7 @@ pub fn plan_rename_at(request: RenameAtRequest<'_>) -> Result { .ok_or_else(|| anyhow::anyhow!("one candidate"))?, _ => return Err(RenameAtError::Ambiguous.into()), }; - SyntaxTree::parse(&candidate.rewritten) + SyntaxTree::parse_with_dialect(&candidate.rewritten, request.dialect) .context("renamed output is not a valid S-expression document")?; Ok(RenameAtPlan { dialect: request.dialect, diff --git a/src/domain/rename/at/safety.rs b/src/domain/rename/at/safety.rs index c12f9b32..eedfed33 100644 --- a/src/domain/rename/at/safety.rs +++ b/src/domain/rename/at/safety.rs @@ -12,7 +12,10 @@ pub(super) fn ensure_binding_target_is_available( binding_span: ByteSpan, input: &str, ) -> Result<()> { - let Ok(existing) = binding_rename_parts(Dialect::CommonLisp, view, to, input) else { + let semantic = Dialect::CommonLisp + .verify_rename_binding() + .expect("Common Lisp rename-binding semantics are verified"); + let Ok(existing) = binding_rename_parts(semantic, view, to, input) else { return Ok(()); }; if existing.binding_span != binding_span && from != to { diff --git a/src/domain/rename/at/tests/global.rs b/src/domain/rename/at/tests/global.rs index 3e43f295..b309c25f 100644 --- a/src/domain/rename/at/tests/global.rs +++ b/src/domain/rename/at/tests/global.rs @@ -8,7 +8,7 @@ fn renames_global_function_from_call_head() { #[test] fn renames_global_definition_calls_and_callable_designators() { - let input = "(defun render (x) (render x)) (list #'render (function render))"; + let input = "(defun render (x) (render x)) (list #'render (function render) #:render)"; let plan = plan_rename_at(RenameAtRequest { input, dialect: Dialect::CommonLisp, @@ -19,7 +19,7 @@ fn renames_global_definition_calls_and_callable_designators() { assert_eq!(plan.namespace, RenameAtNamespace::Function); assert_eq!( plan.rewritten, - "(defun draw (x) (draw x)) (list #'draw (function draw))" + "(defun draw (x) (draw x)) (list #'draw (function draw) #:render)" ); } diff --git a/src/domain/rename/at/tests/validation.rs b/src/domain/rename/at/tests/validation.rs index 4bb47c36..69946d6b 100644 --- a/src/domain/rename/at/tests/validation.rs +++ b/src/domain/rename/at/tests/validation.rs @@ -130,4 +130,47 @@ fn rejects_package_syntax_in_replacement_symbol() { Some(&RenameAtError::UnsupportedPackageSyntax) ); } + +#[test] +fn support_predicate_accepts_only_common_lisp() { + assert!(super::super::supports_rename_at_dialect( + Dialect::CommonLisp + )); + + for dialect in [ + Dialect::EmacsLisp, + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ] { + assert!(!super::super::supports_rename_at_dialect(dialect)); + } +} + +#[test] +fn rejects_unsupported_dialects_before_parsing_malformed_input() { + for dialect in [ + Dialect::EmacsLisp, + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ] { + let error = plan_rename_at(RenameAtRequest { + input: "(", + dialect, + at: ByteOffset::new(0), + to: SymbolName::new("bar").unwrap(), + }) + .unwrap_err(); + + assert_eq!( + error.downcast_ref::(), + Some(&RenameAtError::UnsupportedDialect) + ); + } +} use super::*; diff --git a/src/domain/rename/binding/mod.rs b/src/domain/rename/binding/mod.rs index 05526414..e718b333 100644 --- a/src/domain/rename/binding/mod.rs +++ b/src/domain/rename/binding/mod.rs @@ -5,13 +5,14 @@ mod lambda_like; mod rewrite; mod scope; mod scoped; +mod semantic; mod types; mod value_like; use anyhow::{Context, Result}; use crate::domain::common_lisp::CommonLispBindingRefactorForm; -use crate::domain::dialect::Dialect; +use crate::domain::dialect::{RenameBindingOperation, VerifiedSemanticPolicy}; use crate::domain::sexpr::{ByteSpan, ExpressionView, SymbolName}; use lambda_like::{ @@ -42,14 +43,22 @@ pub(super) fn collect_shadow_aware_special_form( } pub(super) fn binding_rename_parts( - dialect: Dialect, + semantic: VerifiedSemanticPolicy, view: &ExpressionView, from: &SymbolName, input: &str, ) -> Result { + let dialect = semantic.dialect(); let form = super::selection::list_head(view) .context("selected form is not a supported binding form")? .to_owned(); + + if dialect != crate::domain::dialect::Dialect::CommonLisp + && (semantic.scope_shape(view).is_some() || semantic.definition_shape(view).is_some()) + { + return semantic::semantic_binding_rename_parts(semantic, view, from, form, input); + } + let Some(refactor_form) = dialect.common_lisp_binding_refactor_form_for_head(&form) else { anyhow::bail!("selected form is not a supported binding form"); }; diff --git a/src/domain/rename/binding/semantic.rs b/src/domain/rename/binding/semantic.rs new file mode 100644 index 00000000..c4d0309b --- /dev/null +++ b/src/domain/rename/binding/semantic.rs @@ -0,0 +1,831 @@ +use anyhow::{Context, Result}; + +use crate::domain::dialect::{ + BinderShape, BindingVisibility, BodyShape, DefinitionShape, ParameterShape, RelativeNodePath, + RenameBindingOperation, ScopeShape, VerifiedSemanticPolicy, +}; +use crate::domain::sexpr::{ByteSpan, ExpressionKind, ExpressionView, SymbolName}; + +use super::build_binding_rename_parts; +use super::destructure::binding_pattern_name_spans; +use super::forms::{binding_groups, parameter_name_spans}; +use super::types::{BindingGroup, BindingRenameParts, ParameterNameSpan}; + +pub(super) fn semantic_binding_rename_parts( + semantic: VerifiedSemanticPolicy, + view: &ExpressionView, + from: &SymbolName, + form: String, + input: &str, +) -> Result { + let mut reference_spans = Vec::new(); + let mut shadowed_scope_count = 0; + + let binding = if let Some(scope) = semantic.scope_shape(view) { + select_scope_binding_and_collect( + semantic, + view, + scope, + from, + input, + &mut reference_spans, + &mut shadowed_scope_count, + )? + } else if let Some(definition) = semantic.definition_shape(view) { + select_definition_binding_and_collect( + semantic, + view, + definition, + from, + input, + &mut reference_spans, + &mut shadowed_scope_count, + )? + } else { + anyhow::bail!("selected form has no verified semantic binding shape"); + }; + + Ok(build_binding_rename_parts( + form, + view.span, + binding.name_span, + binding.binding_edit, + reference_spans, + shadowed_scope_count, + )) +} + +fn select_scope_binding_and_collect( + semantic: VerifiedSemanticPolicy, + view: &ExpressionView, + scope: ScopeShape, + from: &SymbolName, + input: &str, + output: &mut Vec, + shadowed_scope_count: &mut usize, +) -> Result { + match scope.binders() { + BinderShape::BindingList { + container, + visibility, + .. + } + | BinderShape::FlatPairs { + container, + visibility, + .. + } => { + let groups = binding_groups( + semantic.dialect(), + resolve_relative(view, container).context("binding container is missing")?, + input, + )?; + let (binding, group_index) = select_group_binding(semantic, &groups, from)?; + collect_selected_binding_scope( + semantic, + view, + &groups, + group_index, + visibility, + scope.body(), + from, + input, + output, + shadowed_scope_count, + ); + Ok(binding) + } + BinderShape::NamedBindingList { + scope_name, + container, + visibility, + .. + } => { + let name = single_pattern_binding( + resolve_relative(view, scope_name).context("named scope name is missing")?, + input, + )?; + let groups = binding_groups( + semantic.dialect(), + resolve_relative(view, container).context("binding container is missing")?, + input, + )?; + let mut candidates = matching_group_bindings(semantic, &groups, from); + if identifiers_equal(semantic, &name.name, from) { + candidates.push((name.clone(), None)); + } + let (binding, group_index) = select_unique_indexed(candidates)?; + + if let Some(group_index) = group_index { + collect_selected_binding_scope( + semantic, + view, + &groups, + group_index, + visibility, + scope.body(), + from, + input, + output, + shadowed_scope_count, + ); + } else { + collect_body_references( + semantic, + view, + scope.body(), + from, + input, + output, + shadowed_scope_count, + ); + } + Ok(binding) + } + BinderShape::Parameters(parameters) => { + let bindings = parameter_bindings(view, parameters, input)?; + let binding = select_unique_binding(semantic, bindings, from)?; + collect_body_references( + semantic, + view, + scope.body(), + from, + input, + output, + shadowed_scope_count, + ); + Ok(binding) + } + BinderShape::NamedParameters { name, parameters } => { + let mut bindings = parameter_bindings(view, parameters, input)?; + bindings.push(single_pattern_binding( + resolve_relative(view, name).context("named callable name is missing")?, + input, + )?); + let binding = select_unique_binding(semantic, bindings, from)?; + collect_body_references( + semantic, + view, + scope.body(), + from, + input, + output, + shadowed_scope_count, + ); + Ok(binding) + } + BinderShape::ParameterClauses { + name, + first_clause_index, + parameters, + } => { + let local_name = name + .map(|path| { + resolve_relative(view, path) + .context("named callable name is missing") + .and_then(|name| single_pattern_binding(name, input)) + }) + .transpose()?; + let mut candidates = Vec::new(); + if let Some(name) = &local_name { + if identifiers_equal(semantic, &name.name, from) { + candidates.push((name.clone(), None)); + } + } + for (clause_index, clause) in view.children.iter().enumerate().skip(first_clause_index) + { + for binding in parameter_bindings(clause, parameters, input)? { + if identifiers_equal(semantic, &binding.name, from) { + candidates.push((binding, Some(clause_index))); + } + } + } + let (binding, clause_index) = select_unique_indexed(candidates)?; + match clause_index { + None => collect_body_references( + semantic, + view, + scope.body(), + from, + input, + output, + shadowed_scope_count, + ), + Some(clause_index) => collect_clause_body_references( + semantic, + view, + scope.body(), + clause_index, + from, + input, + output, + shadowed_scope_count, + )?, + } + Ok(binding) + } + } +} + +fn select_definition_binding_and_collect( + semantic: VerifiedSemanticPolicy, + view: &ExpressionView, + definition: DefinitionShape, + from: &SymbolName, + input: &str, + output: &mut Vec, + shadowed_scope_count: &mut usize, +) -> Result { + let parameters = definition + .parameters() + .context("selected definition has no lexical parameters")?; + let binding = + select_unique_binding(semantic, parameter_bindings(view, parameters, input)?, from)?; + collect_body_references( + semantic, + view, + definition.body(), + from, + input, + output, + shadowed_scope_count, + ); + Ok(binding) +} + +#[allow(clippy::too_many_arguments)] +fn collect_selected_binding_scope( + semantic: VerifiedSemanticPolicy, + view: &ExpressionView, + groups: &[BindingGroup], + selected_group: usize, + visibility: BindingVisibility, + body: BodyShape, + from: &SymbolName, + input: &str, + output: &mut Vec, + shadowed_scope_count: &mut usize, +) { + if visibility == BindingVisibility::Sequential { + for group in groups.iter().skip(selected_group + 1) { + if let Some(value) = &group.value { + collect_references( + semantic, + value, + from, + input, + output, + shadowed_scope_count, + false, + ); + } + } + } + collect_body_references( + semantic, + view, + body, + from, + input, + output, + shadowed_scope_count, + ); +} + +fn select_group_binding( + semantic: VerifiedSemanticPolicy, + groups: &[BindingGroup], + from: &SymbolName, +) -> Result<(ParameterNameSpan, usize)> { + let candidates = matching_group_bindings(semantic, groups, from) + .into_iter() + .map(|(binding, index)| Ok((binding, index.context("binding group index is missing")?))) + .collect::>>()?; + let mut candidates = candidates.into_iter(); + let candidate = candidates.next().context("binding name was not found")?; + if candidates.next().is_some() { + anyhow::bail!("binding name is ambiguous in the selected form"); + } + Ok(candidate) +} + +fn matching_group_bindings( + semantic: VerifiedSemanticPolicy, + groups: &[BindingGroup], + from: &SymbolName, +) -> Vec<(ParameterNameSpan, Option)> { + groups + .iter() + .enumerate() + .flat_map(|(index, group)| { + group + .names + .iter() + .filter(move |binding| identifiers_equal(semantic, &binding.name, from)) + .cloned() + .map(move |binding| (binding, Some(index))) + }) + .collect() +} + +fn select_unique_binding( + semantic: VerifiedSemanticPolicy, + bindings: Vec, + from: &SymbolName, +) -> Result { + let mut matches = bindings + .into_iter() + .filter(|binding| identifiers_equal(semantic, &binding.name, from)); + let binding = matches.next().context("binding name was not found")?; + if matches.next().is_some() { + anyhow::bail!("binding name is ambiguous in the selected form"); + } + Ok(binding) +} + +fn select_unique_indexed( + candidates: Vec<(ParameterNameSpan, Option)>, +) -> Result<(ParameterNameSpan, Option)> { + let mut candidates = candidates.into_iter(); + let candidate = candidates.next().context("binding name was not found")?; + if candidates.next().is_some() { + anyhow::bail!("binding name is ambiguous in the selected form"); + } + Ok(candidate) +} + +fn parameter_bindings( + view: &ExpressionView, + parameters: ParameterShape, + input: &str, +) -> Result> { + let mut container = resolve_relative(view, parameters.container()) + .context("parameter container is missing")? + .clone(); + let first = parameters.first_parameter_index(); + if first > container.children.len() { + anyhow::bail!("parameter layout starts outside its container"); + } + container.children.drain(..first); + parameter_name_spans(&container, input) +} + +fn single_pattern_binding(view: &ExpressionView, input: &str) -> Result { + let mut bindings = binding_pattern_name_spans(view, input).into_iter(); + let binding = bindings.next().context("binding pattern has no name")?; + if bindings.next().is_some() { + anyhow::bail!("binding pattern does not identify one name"); + } + Ok(binding) +} + +fn collect_references( + semantic: VerifiedSemanticPolicy, + view: &ExpressionView, + from: &SymbolName, + input: &str, + output: &mut Vec, + shadowed_scope_count: &mut usize, + is_call_head: bool, +) { + if view.kind == ExpressionKind::Atom { + if !is_lisp2_call_head(semantic, is_call_head) + && view + .text + .as_deref() + .is_some_and(|name| semantic.identifiers_equal(name, from.as_str())) + { + output.push(view.span); + } + return; + } + + if let Some(scope) = semantic.scope_shape(view) { + collect_nested_scope_references( + semantic, + view, + scope, + from, + input, + output, + shadowed_scope_count, + ); + return; + } + if let Some(definition) = semantic.definition_shape(view) { + collect_nested_definition_references( + semantic, + view, + definition, + from, + input, + output, + shadowed_scope_count, + ); + return; + } + if semantic.dialect() == crate::domain::dialect::Dialect::EmacsLisp + && super::collect_shadow_aware_special_form(view, from, output, shadowed_scope_count, input) + { + return; + } + + for (index, child) in view.children.iter().enumerate() { + collect_references( + semantic, + child, + from, + input, + output, + shadowed_scope_count, + index == 0, + ); + } +} + +#[allow(clippy::too_many_arguments)] +fn collect_nested_scope_references( + semantic: VerifiedSemanticPolicy, + view: &ExpressionView, + scope: ScopeShape, + from: &SymbolName, + input: &str, + output: &mut Vec, + shadowed_scope_count: &mut usize, +) { + match scope.binders() { + BinderShape::BindingList { + container, + visibility, + .. + } + | BinderShape::FlatPairs { + container, + visibility, + .. + } => { + let Some(container) = resolve_relative(view, container) else { + return; + }; + let Ok(groups) = binding_groups(semantic.dialect(), container, input) else { + return; + }; + collect_nested_binding_groups( + semantic, + view, + &groups, + visibility, + scope.body(), + false, + from, + input, + output, + shadowed_scope_count, + ); + } + BinderShape::NamedBindingList { + scope_name, + container, + visibility, + .. + } => { + let Some(container) = resolve_relative(view, container) else { + return; + }; + let Ok(groups) = binding_groups(semantic.dialect(), container, input) else { + return; + }; + let shadows_from = resolve_relative(view, scope_name) + .is_some_and(|name| pattern_binds(semantic, name, from, input)); + collect_nested_binding_groups( + semantic, + view, + &groups, + visibility, + scope.body(), + shadows_from, + from, + input, + output, + shadowed_scope_count, + ); + } + BinderShape::Parameters(parameters) => { + let shadows_from = parameter_bindings(view, parameters, input) + .is_ok_and(|bindings| bindings_bind(semantic, &bindings, from)); + collect_nested_parameter_body( + semantic, + view, + scope.body(), + shadows_from, + from, + input, + output, + shadowed_scope_count, + ); + } + BinderShape::NamedParameters { name, parameters } => { + let shadows_from = resolve_relative(view, name) + .is_some_and(|name| pattern_binds(semantic, name, from, input)) + || parameter_bindings(view, parameters, input) + .is_ok_and(|bindings| bindings_bind(semantic, &bindings, from)); + collect_nested_parameter_body( + semantic, + view, + scope.body(), + shadows_from, + from, + input, + output, + shadowed_scope_count, + ); + } + BinderShape::ParameterClauses { + name, + first_clause_index, + parameters, + } => { + let name_shadows = name + .and_then(|name| resolve_relative(view, name)) + .is_some_and(|name| pattern_binds(semantic, name, from, input)); + let BodyShape::ClauseChildrenFrom { + body_child_index, .. + } = scope.body() + else { + return; + }; + let mut counted_name_shadow = false; + for clause in view.children.iter().skip(first_clause_index) { + let parameter_shadows = parameter_bindings(clause, parameters, input) + .is_ok_and(|bindings| bindings_bind(semantic, &bindings, from)); + if name_shadows || parameter_shadows { + if parameter_shadows || !counted_name_shadow { + *shadowed_scope_count += 1; + } + counted_name_shadow |= name_shadows; + continue; + } + for child in clause.children.iter().skip(body_child_index) { + collect_references( + semantic, + child, + from, + input, + output, + shadowed_scope_count, + false, + ); + } + } + } + } +} + +#[allow(clippy::too_many_arguments)] +fn collect_nested_binding_groups( + semantic: VerifiedSemanticPolicy, + view: &ExpressionView, + groups: &[BindingGroup], + visibility: BindingVisibility, + body: BodyShape, + name_shadows: bool, + from: &SymbolName, + input: &str, + output: &mut Vec, + shadowed_scope_count: &mut usize, +) { + let mut binding_shadows = false; + for group in groups { + if visibility == BindingVisibility::Parallel || !binding_shadows { + if let Some(value) = &group.value { + collect_references( + semantic, + value, + from, + input, + output, + shadowed_scope_count, + false, + ); + } + } + binding_shadows |= bindings_bind(semantic, &group.names, from); + } + + if name_shadows || binding_shadows { + *shadowed_scope_count += 1; + } else { + collect_body_references( + semantic, + view, + body, + from, + input, + output, + shadowed_scope_count, + ); + } +} + +#[allow(clippy::too_many_arguments)] +fn collect_nested_parameter_body( + semantic: VerifiedSemanticPolicy, + view: &ExpressionView, + body: BodyShape, + shadows_from: bool, + from: &SymbolName, + input: &str, + output: &mut Vec, + shadowed_scope_count: &mut usize, +) { + if shadows_from { + *shadowed_scope_count += 1; + } else { + collect_body_references( + semantic, + view, + body, + from, + input, + output, + shadowed_scope_count, + ); + } +} + +#[allow(clippy::too_many_arguments)] +fn collect_nested_definition_references( + semantic: VerifiedSemanticPolicy, + view: &ExpressionView, + definition: DefinitionShape, + from: &SymbolName, + input: &str, + output: &mut Vec, + shadowed_scope_count: &mut usize, +) { + let name_shadows = definition + .name() + .and_then(|name| resolve_relative(view, name)) + .is_some_and(|name| pattern_binds(semantic, name, from, input)); + let parameter_shadows = definition.parameters().is_some_and(|parameters| { + parameter_bindings(view, parameters, input) + .is_ok_and(|bindings| bindings_bind(semantic, &bindings, from)) + }); + if name_shadows || parameter_shadows { + *shadowed_scope_count += 1; + } else { + collect_body_references( + semantic, + view, + definition.body(), + from, + input, + output, + shadowed_scope_count, + ); + } +} + +#[allow(clippy::too_many_arguments)] +fn collect_body_references( + semantic: VerifiedSemanticPolicy, + view: &ExpressionView, + body: BodyShape, + from: &SymbolName, + input: &str, + output: &mut Vec, + shadowed_scope_count: &mut usize, +) { + match body { + BodyShape::ChildrenFrom(first) => { + for child in view.children.iter().skip(first) { + collect_references( + semantic, + child, + from, + input, + output, + shadowed_scope_count, + false, + ); + } + } + BodyShape::ChildrenAfter(path) => { + for child in view.children.iter().skip(path.child() + 1) { + collect_references( + semantic, + child, + from, + input, + output, + shadowed_scope_count, + false, + ); + } + } + BodyShape::ClauseChildrenFrom { + first_clause_index, + body_child_index, + } => { + for clause in view.children.iter().skip(first_clause_index) { + for child in clause.children.iter().skip(body_child_index) { + collect_references( + semantic, + child, + from, + input, + output, + shadowed_scope_count, + false, + ); + } + } + } + } +} + +#[allow(clippy::too_many_arguments)] +fn collect_clause_body_references( + semantic: VerifiedSemanticPolicy, + view: &ExpressionView, + body: BodyShape, + clause_index: usize, + from: &SymbolName, + input: &str, + output: &mut Vec, + shadowed_scope_count: &mut usize, +) -> Result<()> { + let BodyShape::ClauseChildrenFrom { + first_clause_index, + body_child_index, + } = body + else { + anyhow::bail!("clause parameters require clause body metadata"); + }; + if clause_index < first_clause_index { + anyhow::bail!("selected parameter is outside callable clauses"); + } + let clause = view + .children + .get(clause_index) + .context("selected callable clause is missing")?; + for child in clause.children.iter().skip(body_child_index) { + collect_references( + semantic, + child, + from, + input, + output, + shadowed_scope_count, + false, + ); + } + Ok(()) +} + +fn resolve_relative(view: &ExpressionView, path: RelativeNodePath) -> Option<&ExpressionView> { + let child = view.children.get(path.child())?; + path.grandchild() + .map_or(Some(child), |grandchild| child.children.get(grandchild)) +} + +fn pattern_binds( + semantic: VerifiedSemanticPolicy, + pattern: &ExpressionView, + from: &SymbolName, + input: &str, +) -> bool { + binding_pattern_name_spans(pattern, input) + .iter() + .any(|binding| identifiers_equal(semantic, &binding.name, from)) +} + +fn bindings_bind( + semantic: VerifiedSemanticPolicy, + bindings: &[ParameterNameSpan], + from: &SymbolName, +) -> bool { + bindings + .iter() + .any(|binding| identifiers_equal(semantic, &binding.name, from)) +} + +fn identifiers_equal( + semantic: VerifiedSemanticPolicy, + candidate: &str, + from: &SymbolName, +) -> bool { + semantic.identifiers_equal(candidate, from.as_str()) +} + +fn is_lisp2_call_head( + semantic: VerifiedSemanticPolicy, + is_call_head: bool, +) -> bool { + is_call_head + && matches!( + semantic.dialect(), + crate::domain::dialect::Dialect::CommonLisp + | crate::domain::dialect::Dialect::EmacsLisp + ) +} diff --git a/src/domain/rename/call_identity.rs b/src/domain/rename/call_identity.rs new file mode 100644 index 00000000..a08cea2f --- /dev/null +++ b/src/domain/rename/call_identity.rs @@ -0,0 +1,25 @@ +use crate::domain::common_lisp::common_lisp_symbol_reference_eq; +use crate::domain::dialect::Dialect; + +pub(super) fn call_reference_eq(dialect: Dialect, candidate: &str, expected: &str) -> bool { + match dialect { + Dialect::CommonLisp => common_lisp_symbol_reference_eq(candidate, expected), + Dialect::EmacsLisp + | Dialect::Scheme + | Dialect::Clojure + | Dialect::Janet + | Dialect::Fennel => candidate == expected, + Dialect::Unknown => false, + } +} + +pub(super) fn is_local_call_bound( + dialect: Dialect, + local_callables: &[String], + expected: &str, +) -> bool { + local_callables + .iter() + .rev() + .any(|candidate| call_reference_eq(dialect, candidate, expected)) +} diff --git a/src/domain/rename/replace_call.rs b/src/domain/rename/replace_call.rs index 222d0fcd..ca7997e5 100644 --- a/src/domain/rename/replace_call.rs +++ b/src/domain/rename/replace_call.rs @@ -40,7 +40,20 @@ pub struct ReplaceFunctionCallsPlan { pub fn plan_replace_function_calls( request: ReplaceFunctionCallsRequest<'_>, ) -> Result { - let tree = SyntaxTree::parse(request.input).context("failed to parse input")?; + match request.dialect { + Dialect::CommonLisp + | Dialect::EmacsLisp + | Dialect::Scheme + | Dialect::Clojure + | Dialect::Janet + | Dialect::Fennel => {} + Dialect::Unknown => { + anyhow::bail!("replace-function-calls requires a known dialect"); + } + } + + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("failed to parse input")?; let calls = match &request.scope { ReplaceFunctionCallsScope::AllCalls => collect_all_replace_call_sites( &tree, @@ -63,7 +76,7 @@ pub fn plan_replace_function_calls( .map(|site| (site.head_span, site.replacement.clone())) .collect::>(); let rewritten = apply_byte_span_edits(request.input, edits)?; - SyntaxTree::parse(&rewritten) + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) .context("replace-function-calls output is not a valid S-expression document")?; Ok(ReplaceFunctionCallsPlan { diff --git a/src/domain/rename/replace_call/call_site.rs b/src/domain/rename/replace_call/call_site.rs index 51273f80..b98d5ae5 100644 --- a/src/domain/rename/replace_call/call_site.rs +++ b/src/domain/rename/replace_call/call_site.rs @@ -1,6 +1,6 @@ -use crate::domain::common_lisp::common_lisp_symbol_reference_eq; use crate::domain::definition::definition_shape; use crate::domain::dialect::Dialect; +use crate::domain::rename::call_identity::call_reference_eq; use crate::domain::rename::selection::list_head; use crate::domain::sexpr::{ExpressionView, SymbolName}; @@ -15,7 +15,7 @@ pub(super) fn replace_call_site_from_view( to: &SymbolName, ) -> Option { let head = list_head(view)?; - if !common_lisp_symbol_reference_eq(head, from.as_str()) + if !call_reference_eq(dialect, head, from.as_str()) || definition_shape(dialect, view, head).is_some() { return None; diff --git a/src/domain/rename/replace_call/collect.rs b/src/domain/rename/replace_call/collect.rs index b4853d13..071e8bf2 100644 --- a/src/domain/rename/replace_call/collect.rs +++ b/src/domain/rename/replace_call/collect.rs @@ -1,12 +1,13 @@ use anyhow::{Context, Result}; use crate::domain::callable_scope::{ - common_lisp_local_callable_form, is_local_callable_bound, local_callable_binding_body_scope, - local_callable_body_scope, local_callable_scope_at_path, + common_lisp_local_callable_form, local_callable_binding_body_scope, local_callable_body_scope, + local_callable_scope_at_path, }; use crate::domain::common_lisp::CommonLispLocalCallableForm; use crate::domain::definition::macro_expander_body_range; use crate::domain::dialect::Dialect; +use crate::domain::rename::call_identity::is_local_call_bound; use crate::domain::rename::reader::{ apply_reader_prefix_context, executable_reader_context_at_path, }; @@ -54,7 +55,7 @@ pub(super) fn collect_explicit_replace_call_sites( anyhow::bail!("call-path {path} is not in an executable reader context"); } let local_callables = local_callable_scope_at_path(tree, dialect, path)?; - if is_local_callable_bound(&local_callables, from.as_str()) { + if is_local_call_bound(dialect, &local_callables, from.as_str()) { anyhow::bail!("call-path {path} is shadowed by a local callable named {from}"); } let site = replace_call_site_from_view(&view, dialect, input, path.to_string(), from, to) @@ -119,7 +120,7 @@ fn collect_replace_call_sites_from_view( let macro_expander_body = head.and_then(|head| macro_expander_body_range(ctx.dialect, view, head)); - if !is_local_callable_bound(local_callables, ctx.from.as_str()) { + if !is_local_call_bound(ctx.dialect, local_callables, ctx.from.as_str()) { if let Some(site) = replace_call_site_from_view( view, ctx.dialect, diff --git a/src/domain/rename/tests/binding/basic/lexical.rs b/src/domain/rename/tests/binding/basic/lexical.rs index 6f9f8e4c..6f95077f 100644 --- a/src/domain/rename/tests/binding/basic/lexical.rs +++ b/src/domain/rename/tests/binding/basic/lexical.rs @@ -1,5 +1,24 @@ use super::*; +#[test] +fn rejects_unknown_dialect_before_binding_rename_planning() { + let input = "(let ((value 1)) value)"; + let error = plan_rename_binding(RenameBindingRequest { + input, + dialect: Dialect::Unknown, + target: RenameTarget::Path(Path::from_indexes(vec![0])), + from: SymbolName::new("value").unwrap(), + to: SymbolName::new("product").unwrap(), + }) + .expect_err("unknown dialect should be rejected"); + + assert!( + error + .to_string() + .contains("rename-binding is not supported for this dialect") + ); +} + #[test] fn plans_binding_rename_without_shadowed_inner_binding() { let input = "(let ((value 1)) (+ value (let ((value 2)) value) value))"; @@ -330,3 +349,21 @@ fn plans_janet_vector_let_binding_rename_through_later_binding_values() { assert_eq!(plan.rewritten, "(let [seed 1 next (+ seed 1)] [seed next])"); SyntaxTree::parse(&plan.rewritten).unwrap(); } + +#[test] +fn plans_fennel_vector_let_binding_rename_through_later_binding_values() { + let input = "(let [value 1 next (+ value 1)] [value next])"; + let plan = plan_rename_binding(RenameBindingRequest { + input, + dialect: Dialect::Fennel, + target: RenameTarget::Path(Path::from_indexes(vec![0])), + from: SymbolName::new("value").unwrap(), + to: SymbolName::new("seed").unwrap(), + }) + .unwrap(); + + assert_eq!(plan.form, "let"); + assert_eq!(plan.references.len(), 2); + assert_eq!(plan.rewritten, "(let [seed 1 next (+ seed 1)] [seed next])"); + SyntaxTree::parse(&plan.rewritten).unwrap(); +} diff --git a/src/domain/rename/tests/replace_call/dialect_contract.rs b/src/domain/rename/tests/replace_call/dialect_contract.rs new file mode 100644 index 00000000..130c4fa3 --- /dev/null +++ b/src/domain/rename/tests/replace_call/dialect_contract.rs @@ -0,0 +1,74 @@ +use super::*; + +fn plan(input: &str, dialect: Dialect, from: &str, to: &str) -> ReplaceFunctionCallsPlan { + plan_replace_function_calls(ReplaceFunctionCallsRequest { + input, + dialect, + from: SymbolName::new(from).unwrap(), + to: SymbolName::new(to).unwrap(), + scope: ReplaceFunctionCallsScope::AllCalls, + }) + .unwrap() +} + +#[test] +fn supports_known_dialects_with_their_reader_syntax() { + let cases = [ + (Dialect::CommonLisp, r"(foo #\) #:done #x2a)"), + (Dialect::EmacsLisp, r"(foo ?\))"), + (Dialect::Scheme, "(foo value)"), + (Dialect::Clojure, r#"(foo #inst "2020-01-01")"#), + (Dialect::Janet, "(foo value)"), + (Dialect::Fennel, "(foo value)"), + ]; + + for (dialect, input) in cases { + let plan = plan(input, dialect, "foo", "bar"); + assert_eq!(plan.calls.len(), 1, "{}", dialect.label()); + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect).unwrap(); + } +} + +#[test] +fn rejects_unknown_before_parsing_malformed_input() { + let error = plan_replace_function_calls(ReplaceFunctionCallsRequest { + input: ")", + dialect: Dialect::Unknown, + from: SymbolName::new("foo").unwrap(), + to: SymbolName::new("bar").unwrap(), + scope: ReplaceFunctionCallsScope::AllCalls, + }) + .unwrap_err(); + + assert_eq!( + error.to_string(), + "replace-function-calls requires a known dialect" + ); +} + +#[test] +fn common_lisp_matches_case_and_package_qualified_references() { + let plan = plan("(FOO x) (app::foo y)", Dialect::CommonLisp, "foo", "bar"); + + assert_eq!(plan.calls.len(), 2); + assert_eq!(plan.rewritten, "(bar x) (bar y)"); +} + +#[test] +fn non_common_lisp_matching_is_case_sensitive() { + let plan = plan("(FOO value)", Dialect::EmacsLisp, "foo", "bar"); + + assert!(plan.calls.is_empty()); + assert_eq!(plan.rewritten, "(FOO value)"); +} + +#[test] +fn local_callable_shadowing_uses_the_dialect_identity_rule() { + let input = "(flet ((FOO (x) x)) (foo value))"; + let common_lisp = plan(input, Dialect::CommonLisp, "foo", "bar"); + let emacs_lisp = plan(input, Dialect::EmacsLisp, "foo", "bar"); + + assert!(common_lisp.calls.is_empty()); + assert_eq!(emacs_lisp.calls.len(), 1); + assert_eq!(emacs_lisp.rewritten, "(flet ((FOO (x) x)) (bar value))"); +} diff --git a/src/domain/rename/tests/replace_call/mod.rs b/src/domain/rename/tests/replace_call/mod.rs index 6081d3f4..f88d2732 100644 --- a/src/domain/rename/tests/replace_call/mod.rs +++ b/src/domain/rename/tests/replace_call/mod.rs @@ -68,6 +68,7 @@ macro_rules! assert_shadowed_explicit_path { } mod basic_forms; +mod dialect_contract; mod local_callables; mod macro_forms; mod property; diff --git a/src/domain/rename/tests/scoped_form.rs b/src/domain/rename/tests/scoped_form.rs index efe65e48..faa363c9 100644 --- a/src/domain/rename/tests/scoped_form.rs +++ b/src/domain/rename/tests/scoped_form.rs @@ -19,7 +19,31 @@ fn plans_rename_in_form_only_inside_selected_form() { plan.rewritten, "(list product (list product other))\n(list value)" ); - SyntaxTree::parse(&plan.rewritten).unwrap(); + SyntaxTree::parse_with_dialect(&plan.rewritten, Dialect::CommonLisp).unwrap(); +} + +#[test] +fn rename_in_form_preserves_dialect_reader_collisions() { + let cases = [( + Dialect::Janet, + "(list value)\n# ignored ))", + "(list product)\n# ignored ))", + )]; + + for (dialect, input, expected) in cases { + let plan = plan_rename_in_form(RenameInFormRequest { + input, + dialect, + target: RenameTarget::Path(Path::from_indexes(vec![0])), + from: SymbolName::new("value").unwrap(), + to: SymbolName::new("product").unwrap(), + }) + .expect("plan"); + + assert_eq!(plan.rewritten, expected); + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect) + .expect("rewritten output remains parseable"); + } } proptest! { diff --git a/src/domain/rename/tests/unwrap/dialect_contract.rs b/src/domain/rename/tests/unwrap/dialect_contract.rs new file mode 100644 index 00000000..c2f9db38 --- /dev/null +++ b/src/domain/rename/tests/unwrap/dialect_contract.rs @@ -0,0 +1,84 @@ +use super::*; + +fn plan(input: &str, dialect: Dialect, function: &str, wrapper: &str) -> UnwrapFunctionCallsPlan { + plan_unwrap_function_calls(UnwrapFunctionCallsRequest { + input, + dialect, + function: SymbolName::new(function).unwrap(), + wrapper: SymbolName::new(wrapper).unwrap(), + scope: UnwrapFunctionCallsScope::AllCalls, + }) + .unwrap() +} + +#[test] +fn supports_known_dialects_with_their_reader_syntax() { + let cases = [ + (Dialect::CommonLisp, r"(trace (foo #\) #:done #x2a))"), + (Dialect::EmacsLisp, r"(trace (foo ?\)))"), + (Dialect::Scheme, "(trace (foo value))"), + (Dialect::Clojure, r#"(trace (foo #inst "2020-01-01"))"#), + (Dialect::Janet, "(trace (foo value))"), + (Dialect::Fennel, "(trace (foo value))"), + ]; + + for (dialect, input) in cases { + let plan = plan(input, dialect, "foo", "trace"); + assert_eq!(plan.calls.len(), 1, "{}", dialect.label()); + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect).unwrap(); + } +} + +#[test] +fn rejects_unknown_before_parsing_malformed_input() { + let error = plan_unwrap_function_calls(UnwrapFunctionCallsRequest { + input: ")", + dialect: Dialect::Unknown, + function: SymbolName::new("foo").unwrap(), + wrapper: SymbolName::new("trace").unwrap(), + scope: UnwrapFunctionCallsScope::AllCalls, + }) + .unwrap_err(); + + assert_eq!( + error.to_string(), + "unwrap-function-calls requires a known dialect" + ); +} + +#[test] +fn common_lisp_matches_case_and_package_qualified_references() { + let plan = plan( + "(TRACE (FOO x)) (app:trace (pkg::foo y))", + Dialect::CommonLisp, + "foo", + "trace", + ); + + assert_eq!(plan.calls.len(), 2); + assert_eq!(plan.rewritten, "(FOO x) (pkg::foo y)"); +} + +#[test] +fn non_common_lisp_wrapper_and_inner_matching_are_case_sensitive() { + let plan = plan( + "(TRACE (foo x)) (trace (FOO y))", + Dialect::EmacsLisp, + "foo", + "trace", + ); + + assert!(plan.calls.is_empty()); + assert_eq!(plan.rewritten, "(TRACE (foo x)) (trace (FOO y))"); +} + +#[test] +fn local_callable_shadowing_uses_the_dialect_identity_rule() { + let input = "(flet ((FOO (x) x)) (trace (foo value)))"; + let common_lisp = plan(input, Dialect::CommonLisp, "foo", "trace"); + let emacs_lisp = plan(input, Dialect::EmacsLisp, "foo", "trace"); + + assert!(common_lisp.calls.is_empty()); + assert_eq!(emacs_lisp.calls.len(), 1); + assert_eq!(emacs_lisp.rewritten, "(flet ((FOO (x) x)) (foo value))"); +} diff --git a/src/domain/rename/tests/unwrap/mod.rs b/src/domain/rename/tests/unwrap/mod.rs index 919412f9..eb02a6f0 100644 --- a/src/domain/rename/tests/unwrap/mod.rs +++ b/src/domain/rename/tests/unwrap/mod.rs @@ -70,6 +70,7 @@ macro_rules! assert_shadowed_unwrap_explicit_path { } mod basic_forms; +mod dialect_contract; mod local_callables; mod macro_forms; mod property; diff --git a/src/domain/rename/tests/wrap/dialect_contract.rs b/src/domain/rename/tests/wrap/dialect_contract.rs new file mode 100644 index 00000000..bc960d57 --- /dev/null +++ b/src/domain/rename/tests/wrap/dialect_contract.rs @@ -0,0 +1,119 @@ +use super::*; + +fn plan( + input: &str, + dialect: Dialect, + function: &str, + wrapper: &str, + wrapper_template: Option, +) -> WrapFunctionCallsPlan { + plan_wrap_function_calls(WrapFunctionCallsRequest { + input, + dialect, + function: SymbolName::new(function).unwrap(), + wrapper: SymbolName::new(wrapper).unwrap(), + wrapper_template, + scope: WrapFunctionCallsScope::AllCalls, + }) + .unwrap() +} + +#[test] +fn supports_known_dialects_with_their_reader_syntax() { + let cases = [ + (Dialect::CommonLisp, r"(foo #\) #:done #x2a)"), + (Dialect::EmacsLisp, r"(foo ?\))"), + (Dialect::Scheme, "(foo value)"), + (Dialect::Clojure, r#"(foo #inst "2020-01-01")"#), + (Dialect::Janet, "(foo value)"), + (Dialect::Fennel, "(foo value)"), + ]; + + for (dialect, input) in cases { + let plan = plan(input, dialect, "foo", "trace", None); + assert_eq!(plan.calls.len(), 1, "{}", dialect.label()); + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect).unwrap(); + } +} + +#[test] +fn rejects_unknown_before_parsing_input_or_template() { + let error = plan_wrap_function_calls(WrapFunctionCallsRequest { + input: ")", + dialect: Dialect::Unknown, + function: SymbolName::new("foo").unwrap(), + wrapper: SymbolName::new("trace").unwrap(), + wrapper_template: Some("(".to_owned()), + scope: WrapFunctionCallsScope::AllCalls, + }) + .unwrap_err(); + + assert_eq!( + error.to_string(), + "wrap-function-calls requires a known dialect" + ); +} + +#[test] +fn common_lisp_matches_case_and_package_qualified_references() { + let plan = plan( + "(FOO x) (app::foo y)", + Dialect::CommonLisp, + "foo", + "trace", + Some("(TRACE _)".to_owned()), + ); + + assert_eq!(plan.calls.len(), 2); + assert_eq!(plan.rewritten, "(TRACE (FOO x)) (TRACE (app::foo y))"); +} + +#[test] +fn non_common_lisp_matching_and_template_heads_are_case_sensitive() { + let plan = plan("(FOO value)", Dialect::EmacsLisp, "foo", "trace", None); + assert!(plan.calls.is_empty()); + + let error = plan_wrap_function_calls(WrapFunctionCallsRequest { + input: "(foo value)", + dialect: Dialect::EmacsLisp, + function: SymbolName::new("foo").unwrap(), + wrapper: SymbolName::new("trace").unwrap(), + wrapper_template: Some("(TRACE _)".to_owned()), + scope: WrapFunctionCallsScope::AllCalls, + }) + .unwrap_err(); + assert!( + error + .to_string() + .contains("wrapper template head must match --wrapper (trace)") + ); +} + +#[test] +fn already_wrapped_detection_uses_exact_identity_outside_common_lisp() { + let plan = plan( + "(TRACE (foo value))", + Dialect::EmacsLisp, + "foo", + "trace", + None, + ); + + assert_eq!(plan.calls.len(), 1); + assert!(plan.skipped_already_wrapped.is_empty()); + assert_eq!(plan.rewritten, "(TRACE (trace (foo value)))"); +} + +#[test] +fn local_callable_shadowing_uses_the_dialect_identity_rule() { + let input = "(flet ((FOO (x) x)) (foo value))"; + let common_lisp = plan(input, Dialect::CommonLisp, "foo", "trace", None); + let emacs_lisp = plan(input, Dialect::EmacsLisp, "foo", "trace", None); + + assert!(common_lisp.calls.is_empty()); + assert_eq!(emacs_lisp.calls.len(), 1); + assert_eq!( + emacs_lisp.rewritten, + "(flet ((FOO (x) x)) (trace (foo value)))" + ); +} diff --git a/src/domain/rename/tests/wrap/mod.rs b/src/domain/rename/tests/wrap/mod.rs index c9bee632..75f7e0ed 100644 --- a/src/domain/rename/tests/wrap/mod.rs +++ b/src/domain/rename/tests/wrap/mod.rs @@ -74,6 +74,7 @@ macro_rules! assert_shadowed_wrap_explicit_path { } mod basic_forms; +mod dialect_contract; mod local_callables; mod macro_forms; mod property; diff --git a/src/domain/rename/unwrap.rs b/src/domain/rename/unwrap.rs index 3234f681..7e7ebb14 100644 --- a/src/domain/rename/unwrap.rs +++ b/src/domain/rename/unwrap.rs @@ -48,7 +48,18 @@ pub struct UnwrapFunctionCallsPlan { pub fn plan_unwrap_function_calls( request: UnwrapFunctionCallsRequest<'_>, ) -> Result { - let tree = SyntaxTree::parse(request.input).context("failed to parse input")?; + match request.dialect { + Dialect::CommonLisp + | Dialect::EmacsLisp + | Dialect::Scheme + | Dialect::Clojure + | Dialect::Janet + | Dialect::Fennel => {} + Dialect::Unknown => anyhow::bail!("unwrap-function-calls requires a known dialect"), + } + + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("failed to parse input")?; let (calls, skipped_non_unary_wrapper, skipped_nested) = match &request.scope { UnwrapFunctionCallsScope::AllCalls => collect_unwrap_all_call_sites( &tree, @@ -71,7 +82,7 @@ pub fn plan_unwrap_function_calls( .map(|site| (site.span, site.replacement.clone())) .collect::>(); let rewritten = apply_byte_span_edits(request.input, edits)?; - SyntaxTree::parse(&rewritten) + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) .context("unwrapped output is not a valid S-expression document")?; Ok(UnwrapFunctionCallsPlan { diff --git a/src/domain/rename/unwrap/call_site.rs b/src/domain/rename/unwrap/call_site.rs index 6fbe1ddf..3ab13b3d 100644 --- a/src/domain/rename/unwrap/call_site.rs +++ b/src/domain/rename/unwrap/call_site.rs @@ -1,6 +1,6 @@ -use crate::domain::common_lisp::common_lisp_symbol_reference_eq; use crate::domain::definition::definition_shape; use crate::domain::dialect::Dialect; +use crate::domain::rename::call_identity::call_reference_eq; use crate::domain::rename::selection::list_head; use crate::domain::sexpr::{ExpressionView, SymbolName}; @@ -23,15 +23,14 @@ pub(super) fn unwrap_call_site_from_view( let Some(head) = list_head(view) else { return UnwrapCandidate::NotMatched; }; - if !common_lisp_symbol_reference_eq(head, wrapper.as_str()) + if !call_reference_eq(dialect, head, wrapper.as_str()) || definition_shape(dialect, view, head).is_some() { return UnwrapCandidate::NotMatched; } let matching_inner_call = view.children.iter().skip(1).find(|child| { - list_head(child) - .is_some_and(|head| common_lisp_symbol_reference_eq(head, function.as_str())) + list_head(child).is_some_and(|head| call_reference_eq(dialect, head, function.as_str())) }); let Some(inner_call) = matching_inner_call else { return UnwrapCandidate::NotMatched; diff --git a/src/domain/rename/unwrap/collect.rs b/src/domain/rename/unwrap/collect.rs index 17c5c82f..eed8d36d 100644 --- a/src/domain/rename/unwrap/collect.rs +++ b/src/domain/rename/unwrap/collect.rs @@ -1,12 +1,13 @@ use anyhow::Result; use crate::domain::callable_scope::{ - common_lisp_local_callable_form, is_local_callable_bound, local_callable_binding_body_scope, - local_callable_body_scope, local_callable_names, local_callable_scope_at_path, + common_lisp_local_callable_form, local_callable_binding_body_scope, local_callable_body_scope, + local_callable_names, local_callable_scope_at_path, }; use crate::domain::common_lisp::CommonLispLocalCallableForm; use crate::domain::definition::macro_expander_body_range; use crate::domain::dialect::Dialect; +use crate::domain::rename::call_identity::is_local_call_bound; use crate::domain::rename::reader::{ apply_reader_prefix_context, executable_reader_context_at_path, }; @@ -61,7 +62,7 @@ pub(super) fn collect_unwrap_explicit_call_sites( anyhow::bail!("call-path {path} is not in an executable reader context"); } let local_callables = local_callable_scope_at_path(tree, dialect, path)?; - if is_local_callable_bound(&local_callables, function.as_str()) { + if is_local_call_bound(dialect, &local_callables, function.as_str()) { anyhow::bail!("call-path {path} is shadowed by a local callable named {function}"); } match unwrap_call_site_from_view( @@ -216,7 +217,8 @@ impl<'a> UnwrapCollection<'a> { continue; } - if !is_local_callable_bound(&local_callables, self.function.as_str()) { + if !is_local_call_bound(self.dialect, &local_callables, self.function.as_str()) + { let materialized_paths = &mut self.traversal_stats.materialized_paths; match unwrap_call_site_from_view( view, @@ -386,7 +388,8 @@ mod tests { let mut input = "(".repeat(DEPTH); input.push_str("leaf"); input.push_str(&")".repeat(DEPTH)); - let tree = SyntaxTree::parse(&input).expect("deep expression parses"); + let tree = SyntaxTree::parse_with_dialect(&input, Dialect::CommonLisp) + .expect("deep expression parses"); let view = tree .select_path(&Path::root_child(0)) .expect("root expression") diff --git a/src/domain/rename/wrap.rs b/src/domain/rename/wrap.rs index 14523aa8..5e116ecc 100644 --- a/src/domain/rename/wrap.rs +++ b/src/domain/rename/wrap.rs @@ -1,8 +1,8 @@ use anyhow::{Context, Result, ensure}; -use crate::domain::common_lisp::common_lisp_symbol_reference_eq; use crate::domain::dialect::Dialect; pub use crate::domain::rename::WrapFunctionCallsScope; +use crate::domain::rename::call_identity::call_reference_eq; use crate::domain::sexpr::{ ByteSpan, ExpressionKind, ExpressionView, Path, SymbolName, SyntaxTree, }; @@ -56,8 +56,9 @@ pub(super) struct WrapFunctionCallTemplate { } impl WrapFunctionCallTemplate { - fn parse(source: String, wrapper: &SymbolName) -> Result { - let tree = SyntaxTree::parse(&source).context("failed to parse wrapper template")?; + fn parse(source: String, dialect: Dialect, wrapper: &SymbolName) -> Result { + let tree = SyntaxTree::parse_with_dialect(&source, dialect) + .context("failed to parse wrapper template")?; ensure!( tree.root_children().len() == 1, "wrapper template must contain exactly one root form" @@ -67,7 +68,7 @@ impl WrapFunctionCallTemplate { let head = crate::domain::rename::selection::list_head(&root) .context("wrapper template root form must be a parenthesized list")?; ensure!( - common_lisp_symbol_reference_eq(head, wrapper.as_str()), + call_reference_eq(dialect, head, wrapper.as_str()), "wrapper template head must match --wrapper ({})", wrapper.as_str() ); @@ -106,10 +107,21 @@ fn collect_template_placeholders(view: &ExpressionView, output: &mut Vec, ) -> Result { - let tree = SyntaxTree::parse(request.input).context("failed to parse input")?; + match request.dialect { + Dialect::CommonLisp + | Dialect::EmacsLisp + | Dialect::Scheme + | Dialect::Clojure + | Dialect::Janet + | Dialect::Fennel => {} + Dialect::Unknown => anyhow::bail!("wrap-function-calls requires a known dialect"), + } + + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("failed to parse input")?; let template = request .wrapper_template - .map(|source| WrapFunctionCallTemplate::parse(source, &request.wrapper)) + .map(|source| WrapFunctionCallTemplate::parse(source, request.dialect, &request.wrapper)) .transpose()?; let (calls, skipped_already_wrapped, skipped_nested) = match &request.scope { WrapFunctionCallsScope::AllCalls => collect_wrap_all_call_sites( @@ -135,7 +147,8 @@ pub fn plan_wrap_function_calls( .map(|site| (site.span, site.replacement.clone())) .collect::>(); let rewritten = apply_byte_span_edits(request.input, edits)?; - SyntaxTree::parse(&rewritten).context("wrapped output is not a valid S-expression document")?; + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) + .context("wrapped output is not a valid S-expression document")?; Ok(WrapFunctionCallsPlan { dialect: request.dialect, diff --git a/src/domain/rename/wrap/call_site.rs b/src/domain/rename/wrap/call_site.rs index 23682fef..f098e4b9 100644 --- a/src/domain/rename/wrap/call_site.rs +++ b/src/domain/rename/wrap/call_site.rs @@ -1,4 +1,5 @@ -use crate::domain::common_lisp::common_lisp_symbol_reference_eq; +use crate::domain::dialect::Dialect; +use crate::domain::rename::call_identity::call_reference_eq; use crate::domain::sexpr::{ExpressionView, SymbolName}; use super::{WrapFunctionCallSite, WrapFunctionCallTemplate}; @@ -6,6 +7,7 @@ use crate::domain::rename::selection::list_head; pub(super) fn wrap_call_site_from_view( view: &ExpressionView, + dialect: Dialect, input: &str, path: String, function: &SymbolName, @@ -13,7 +15,7 @@ pub(super) fn wrap_call_site_from_view( template: Option<&WrapFunctionCallTemplate>, ) -> Option { let head = list_head(view)?; - if !common_lisp_symbol_reference_eq(head, function.as_str()) { + if !call_reference_eq(dialect, head, function.as_str()) { return None; } let text = view.content_span.slice(input).to_owned(); diff --git a/src/domain/rename/wrap/collect.rs b/src/domain/rename/wrap/collect.rs index 1071ac50..bda6f2f7 100644 --- a/src/domain/rename/wrap/collect.rs +++ b/src/domain/rename/wrap/collect.rs @@ -1,12 +1,13 @@ use anyhow::{Context, Result}; use crate::domain::callable_scope::{ - common_lisp_local_callable_form, is_local_callable_bound, local_callable_binding_body_scope, - local_callable_body_scope, local_callable_scope_at_path, + common_lisp_local_callable_form, local_callable_binding_body_scope, local_callable_body_scope, + local_callable_scope_at_path, }; -use crate::domain::common_lisp::{CommonLispLocalCallableForm, common_lisp_symbol_reference_eq}; +use crate::domain::common_lisp::CommonLispLocalCallableForm; use crate::domain::definition::{definition_shape, macro_expander_body_range}; use crate::domain::dialect::Dialect; +use crate::domain::rename::call_identity::{call_reference_eq, is_local_call_bound}; use crate::domain::rename::reader::{ apply_reader_prefix_context, executable_reader_context_at_path, }; @@ -73,12 +74,19 @@ pub(super) fn collect_wrap_explicit_call_sites( anyhow::bail!("call-path {path} is not in an executable reader context"); } let local_callables = local_callable_scope_at_path(tree, dialect, path)?; - if is_local_callable_bound(&local_callables, function.as_str()) { + if is_local_call_bound(dialect, &local_callables, function.as_str()) { anyhow::bail!("call-path {path} is shadowed by a local callable named {function}"); } - let site = - wrap_call_site_from_view(&view, input, path.to_string(), function, wrapper, template) - .with_context(|| format!("call-path {path} is not a call to {function}"))?; + let site = wrap_call_site_from_view( + &view, + dialect, + input, + path.to_string(), + function, + wrapper, + template, + ) + .with_context(|| format!("call-path {path} is not a call to {function}"))?; if call_site_is_already_wrapped(tree, dialect, path, wrapper)? { skipped_already_wrapped.push(site); } else { @@ -142,9 +150,14 @@ fn collect_wrap_call_sites_from_view( } let current_head = list_head(view); - if !is_local_callable_bound(local_callables, collection.function.as_str()) { + if !is_local_call_bound( + collection.dialect, + local_callables, + collection.function.as_str(), + ) { if let Some(site) = wrap_call_site_from_view( view, + collection.dialect, collection.input, path.to_string(), collection.function, @@ -152,7 +165,7 @@ fn collect_wrap_call_sites_from_view( collection.template, ) { if parent_head.is_some_and(|head| { - common_lisp_symbol_reference_eq(head, collection.wrapper.as_str()) + call_reference_eq(collection.dialect, head, collection.wrapper.as_str()) }) { collection.skipped_already_wrapped.push(site); } else if current_head @@ -241,6 +254,6 @@ fn call_site_is_already_wrapped( let Some(head) = list_head(&parent) else { return Ok(false); }; - Ok(common_lisp_symbol_reference_eq(head, wrapper.as_str()) + Ok(call_reference_eq(dialect, head, wrapper.as_str()) && definition_shape(dialect, &parent, head).is_none()) } diff --git a/src/domain/rename_control.rs b/src/domain/rename_control.rs index 46e0557d..4945ab2b 100644 --- a/src/domain/rename_control.rs +++ b/src/domain/rename_control.rs @@ -54,8 +54,8 @@ fn plan(request: RenameControlRequest<'_>, kind: ControlKind) -> Result, kind: ControlKind) -> Result bool { #[cfg(test)] mod tests { use super::*; - fn req<'a>(input: &'a str, from: &str, to: &str) -> RenameControlRequest<'a> { + + const DIALECTS: [Dialect; 7] = [ + Dialect::CommonLisp, + Dialect::EmacsLisp, + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + Dialect::Unknown, + ]; + + fn req_for_dialect<'a>( + input: &'a str, + dialect: Dialect, + from: &str, + to: &str, + ) -> RenameControlRequest<'a> { RenameControlRequest { input, - dialect: Dialect::CommonLisp, + dialect, path: "0".parse().unwrap(), from: from.parse().unwrap(), to: to.parse().unwrap(), } } + + fn req<'a>(input: &'a str, from: &str, to: &str) -> RenameControlRequest<'a> { + req_for_dialect(input, Dialect::CommonLisp, from, to) + } + + fn assert_support_error(result: Result, operation: &str) { + let error = result.expect_err("unsupported dialect must fail"); + assert_eq!( + error.to_string(), + format!("{operation} supports only Common Lisp") + ); + } + #[test] fn renames_block_references_but_not_shadowed_ones() { let p = plan_rename_block(req( @@ -331,10 +361,12 @@ mod tests { "(block done (return-from done 1) (block out (return-from out 2)))" ); } + #[test] fn rejects_block_capture() { assert!(plan_rename_block(req("(block out (return-from done 1))", "out", "done")).is_err()); } + #[test] fn renames_tag_and_go_but_not_shadowed_go() { let p = plan_rename_tag(req( @@ -348,8 +380,64 @@ mod tests { "(tagbody next (go next) (tagbody start (go start)))" ); } + #[test] fn rejects_duplicate_tags() { assert!(plan_rename_tag(req("(tagbody x x (go x))", "x", "y")).is_err()); } + + #[test] + fn support_matrix_is_common_lisp_only_for_both_operations() { + for dialect in DIALECTS { + let block = plan_rename_block(req_for_dialect( + "(block out (return-from out 1))", + dialect, + "out", + "done", + )); + let tag = plan_rename_tag(req_for_dialect( + "(tagbody start (go start))", + dialect, + "start", + "next", + )); + if dialect == Dialect::CommonLisp { + assert!(block.is_ok(), "rename-block must support {dialect:?}"); + assert!(tag.is_ok(), "rename-tag must support {dialect:?}"); + } else { + assert_support_error(block, "rename-block"); + assert_support_error(tag, "rename-tag"); + } + } + } + + #[test] + fn unsupported_dialect_gate_precedes_parsing_for_both_operations() { + for dialect in DIALECTS + .into_iter() + .filter(|dialect| *dialect != Dialect::CommonLisp) + { + assert_support_error( + plan_rename_block(req_for_dialect(")", dialect, "out", "done")), + "rename-block", + ); + assert_support_error( + plan_rename_tag(req_for_dialect(")", dialect, "start", "next")), + "rename-tag", + ); + } + } + + #[test] + fn preserves_common_lisp_delimiter_character_literals() { + let plan = plan_rename_block(req( + "(block out #\\) (return-from out #\\)))", + "out", + "done", + )) + .expect("rename block containing character literals"); + assert_eq!(plan.rewritten, "(block done #\\) (return-from done #\\)))"); + SyntaxTree::parse_with_dialect(&plan.rewritten, Dialect::CommonLisp) + .expect("rewritten output must parse with the request dialect"); + } } diff --git a/src/domain/replace_forms.rs b/src/domain/replace_forms.rs index dbe73aef..e79ac65e 100644 --- a/src/domain/replace_forms.rs +++ b/src/domain/replace_forms.rs @@ -1,5 +1,6 @@ -use anyhow::{Context, Result}; +use anyhow::{Context, Result, bail}; +use crate::domain::dialect::Dialect; use crate::domain::form_shape::duplicate_shape; use crate::domain::mutation_safety::{ reject_common_lisp_reader_conditionals, reject_overlapping_common_lisp_reader_time_forms, @@ -19,14 +20,16 @@ use validation::{ }; pub fn plan_replace_forms(request: ReplaceFormsRequest<'_>) -> Result { - let input_tree = SyntaxTree::parse(request.input) + ensure_supported_dialect(request.dialect)?; + + let input_tree = SyntaxTree::parse_with_dialect(request.input, request.dialect) .context("replace-forms input is not a valid S-expression document")?; anyhow::ensure!( &input_tree == request.tree, "replace-forms input does not match the source used to build the syntax tree" ); - let replacement_tree = SyntaxTree::parse(request.replacement) + let replacement_tree = SyntaxTree::parse_with_dialect(request.replacement, request.dialect) .context("--with must be a valid S-expression document")?; // The replacement becomes source code in the rewritten document, so it // must satisfy the same reader-time safety contract as the input tree. @@ -52,7 +55,7 @@ pub fn plan_replace_forms(request: ReplaceFormsRequest<'_>) -> Result) -> Result Result<()> { + match dialect { + Dialect::CommonLisp + | Dialect::EmacsLisp + | Dialect::Scheme + | Dialect::Clojure + | Dialect::Janet + | Dialect::Fennel => Ok(()), + Dialect::Unknown => bail!("replace-forms requires a known dialect"), + } +} diff --git a/src/domain/replace_forms/tests.rs b/src/domain/replace_forms/tests.rs index ae09dc1a..fc99ad5b 100644 --- a/src/domain/replace_forms/tests.rs +++ b/src/domain/replace_forms/tests.rs @@ -7,7 +7,7 @@ use crate::domain::form_shape::FormShape; #[test] fn plans_multiple_replacements_in_reverse_span_order() { let input = "(foo 1)\n(foo 2)\n"; - let tree = SyntaxTree::parse(input).unwrap(); + let tree = SyntaxTree::parse_with_dialect(input, Dialect::CommonLisp).unwrap(); let plan = plan_replace_forms(ReplaceFormsRequest { input, @@ -32,7 +32,7 @@ fn plans_multiple_replacements_in_reverse_span_order() { #[test] fn rejects_duplicate_paths() { let input = "(foo 1)\n"; - let tree = SyntaxTree::parse(input).unwrap(); + let tree = SyntaxTree::parse_with_dialect(input, Dialect::CommonLisp).unwrap(); let error = plan_replace_forms(ReplaceFormsRequest { input, @@ -50,7 +50,7 @@ fn rejects_duplicate_paths() { #[test] fn rejects_shape_mismatch_when_required() { let input = "(foo 1)\n(foo 1 2)\n"; - let tree = SyntaxTree::parse(input).unwrap(); + let tree = SyntaxTree::parse_with_dialect(input, Dialect::CommonLisp).unwrap(); let error = plan_replace_forms(ReplaceFormsRequest { input, @@ -71,7 +71,7 @@ fn rejects_shape_mismatch_when_required() { #[test] fn rejects_input_that_does_not_match_tree_source() { - let tree = SyntaxTree::parse("(a x)").unwrap(); + let tree = SyntaxTree::parse_with_dialect("(a x)", Dialect::CommonLisp).unwrap(); let error = plan_replace_forms(ReplaceFormsRequest { input: "(é x)", @@ -86,6 +86,70 @@ fn rejects_input_that_does_not_match_tree_source() { assert!(error.to_string().contains("does not match")); } +#[test] +fn supports_each_known_dialect_with_its_reader_semantics() { + let cases = [ + (Dialect::CommonLisp, r"(list #\))", r"#\(", r"(list #\()"), + (Dialect::EmacsLisp, r"(list ?\))", r"?\(", r"(list ?\()"), + (Dialect::Scheme, "(list old)", "new", "(list new)"), + ( + Dialect::Clojure, + r#"(vector #inst "2020-01-01")"#, + r#"#uuid "00000000-0000-0000-0000-000000000000""#, + r#"(vector #uuid "00000000-0000-0000-0000-000000000000")"#, + ), + (Dialect::Janet, "(list old)", "new", "(list new)"), + (Dialect::Fennel, "(list old)", "new", "(list new)"), + ]; + + for (dialect, input, replacement, expected) in cases { + let tree = SyntaxTree::parse_with_dialect(input, dialect).unwrap(); + let plan = plan_replace_forms(ReplaceFormsRequest { + input, + tree: &tree, + dialect, + paths: vec![Path::from_indexes(vec![0, 1])], + replacement, + require_same_shape: false, + }) + .unwrap(); + + assert_eq!(plan.rewritten, expected, "dialect: {dialect:?}"); + assert!( + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect).is_ok(), + "dialect: {dialect:?}" + ); + } +} + +#[test] +fn dialect_specific_reader_semantics_do_not_leak() { + assert!(SyntaxTree::parse_with_dialect(r"(list #\))", Dialect::Clojure).is_err()); + assert!(SyntaxTree::parse_with_dialect(r"[?\)]", Dialect::EmacsLisp).is_ok()); + assert!(SyntaxTree::parse_with_dialect(r"[?\)]", Dialect::CommonLisp).is_err()); + assert!( + SyntaxTree::parse_with_dialect(r#"(vector #inst "2020-01-01")"#, Dialect::CommonLisp) + .is_err() + ); +} + +#[test] +fn rejects_unknown_dialect_before_parsing_malformed_documents() { + let tree = SyntaxTree::parse_with_dialect("placeholder", Dialect::CommonLisp).unwrap(); + + let error = plan_replace_forms(ReplaceFormsRequest { + input: ")", + tree: &tree, + dialect: Dialect::Unknown, + paths: vec![Path::from_indexes(vec![0])], + replacement: "(", + require_same_shape: false, + }) + .unwrap_err(); + + assert_eq!(error.to_string(), "replace-forms requires a known dialect"); +} + fn lisp_symbol_strategy() -> impl Strategy { "[a-z][a-z0-9-]{0,8}".prop_map(|name| name) } @@ -98,7 +162,7 @@ proptest! { replacement in lisp_symbol_strategy(), ) { let input = format!("({left} 1)\n({right} 2)\n"); - let tree = SyntaxTree::parse(&input).unwrap(); + let tree = SyntaxTree::parse_with_dialect(&input, Dialect::CommonLisp).unwrap(); let replacement_text = format!("({replacement} 0)"); let plan = plan_replace_forms(ReplaceFormsRequest { @@ -114,6 +178,8 @@ proptest! { prop_assert!(plan.changed); prop_assert_eq!(plan.targets.len(), 2); prop_assert_eq!(&plan.rewritten, &format!("{replacement_text}\n{replacement_text}\n")); - prop_assert!(SyntaxTree::parse(&plan.rewritten).is_ok()); + prop_assert!( + SyntaxTree::parse_with_dialect(&plan.rewritten, Dialect::CommonLisp).is_ok() + ); } } diff --git a/src/domain/sexpr.rs b/src/domain/sexpr.rs index 76d48fee..3408441d 100644 --- a/src/domain/sexpr.rs +++ b/src/domain/sexpr.rs @@ -5,6 +5,7 @@ mod edit; mod formatter; mod parser; pub mod reader; +mod reader_policy; #[cfg(test)] #[allow(clippy::unwrap_used)] mod tests; diff --git a/src/domain/sexpr/edit.rs b/src/domain/sexpr/edit.rs index e09a273e..19c9ce3e 100644 --- a/src/domain/sexpr/edit.rs +++ b/src/domain/sexpr/edit.rs @@ -7,12 +7,16 @@ use super::types::{ByteOffset, ByteSpan, NodeId}; pub struct Edit; impl Edit { - pub fn normalize_changed_line_trivia(input: &str, rewritten: String) -> Result { + pub fn normalize_changed_line_trivia( + input: &str, + rewritten: String, + dialect: crate::domain::dialect::Dialect, + ) -> Result { if input == rewritten { return Ok(rewritten); } - let tree = SyntaxTree::parse(&rewritten)?; + let tree = SyntaxTree::parse_with_dialect(&rewritten, dialect)?; let prefix = common_prefix_len(input, &rewritten); let suffix = common_suffix_len(input, &rewritten, prefix); let changed_end = rewritten.len().saturating_sub(suffix); diff --git a/src/domain/sexpr/formatter/core.rs b/src/domain/sexpr/formatter/core.rs index 37d650ca..e03636cb 100644 --- a/src/domain/sexpr/formatter/core.rs +++ b/src/domain/sexpr/formatter/core.rs @@ -1,6 +1,6 @@ use super::styles::ListStyle; use super::{Formatter, MAX_INLINE_WIDTH}; -use crate::domain::sexpr::tree::{Node, NodeKind, ReaderPrefix, SyntaxTree}; +use crate::domain::sexpr::tree::{Node, NodeKind, SyntaxTree}; use crate::domain::sexpr::types::Delimiter; use crate::domain::sexpr::types::NodeId; @@ -11,6 +11,9 @@ const MAX_RECURSIVE_FORMAT_DEPTH: usize = 256; enum TopLevelItem { Form { node_id: NodeId, + /// The last root node that belongs to this logical form. Metadata + /// descriptors and their target are separate parser nodes. + end_node_id: NodeId, /// Own-line comments emitted immediately above the form. leading: Vec, /// A comment trailing the form on the same source line, if any. @@ -58,21 +61,35 @@ impl Formatter { let mut order: Vec = (0..comments.len()).collect(); order.sort_by_key(|&index| comments[index].span.start().get()); + let root_children = tree.root_children(); + let mut node_index = 0usize; let mut cursor = 0usize; let mut items: Vec = Vec::new(); - for &node_id in tree.root_children() { - let node = tree.node(node_id); - let start = node.span.start().get(); - let end = node.span.end().get(); + while node_index < root_children.len() { + let node_id = root_children[node_index]; + let start = tree.node(node_id).span.start().get(); + let mut end_node_id = node_id; + + while tree + .node(end_node_id) + .reader_prefixes + .iter() + .any(|prefix| matches!(prefix, crate::domain::sexpr::tree::ReaderPrefix::Metadata)) + && node_index + 1 < root_children.len() + { + node_index += 1; + end_node_id = root_children[node_index]; + } + let end = tree.node(end_node_id).span.end().get(); let mut leading = Vec::new(); while cursor < order.len() && comments[order[cursor]].span.start().get() < start { leading.push(order[cursor]); cursor += 1; } - let mut verbatim = false; + let mut verbatim = node_id != end_node_id; while cursor < order.len() && comments[order[cursor]].span.start().get() < end { verbatim = true; cursor += 1; @@ -81,6 +98,7 @@ impl Formatter { let item_index = items.len(); items.push(TopLevelItem::Form { node_id, + end_node_id, leading, trailing: None, verbatim, @@ -98,6 +116,8 @@ impl Formatter { cursor += 1; } } + + node_index += 1; } if cursor < order.len() { @@ -112,6 +132,7 @@ impl Formatter { match item { TopLevelItem::Form { node_id, + end_node_id, leading, trailing, verbatim, @@ -121,7 +142,9 @@ impl Formatter { output.push('\n'); } if *verbatim { - output.push_str(&tree.source[tree.node(*node_id).span.as_range()]); + let start = tree.node(*node_id).span.start().get(); + let end = tree.node(*end_node_id).span.end().get(); + output.push_str(&tree.source[start..end]); } else { self.format_node(tree, *node_id, 0, output); } @@ -163,7 +186,7 @@ impl Formatter { output.push_str(node.span.slice(&tree.source)); return; } - self.write_reader_prefixes(node, output); + self.write_reader_prefixes(tree, node, output); let delimiter = self.list_delimiter(node); output.push(delimiter.open()); output.push(delimiter.close()); @@ -177,7 +200,7 @@ impl Formatter { output.push_str(node.span.slice(&tree.source)); return; } - self.write_reader_prefixes(node, output); + self.write_reader_prefixes(tree, node, output); if let Some(head) = self.head_text(tree, node_id) { match self.style_for_head(head) { ListStyle::Definition => { @@ -309,8 +332,8 @@ impl Formatter { return None; } - for prefix in &node.reader_prefixes { - push_bounded(&mut output, prefix.as_source())?; + for span in &node.reader_prefix_spans { + push_bounded(&mut output, span.slice(&tree.source))?; } let delimiter = self.list_delimiter(node); push_char_bounded(&mut output, delimiter.open())?; @@ -359,16 +382,23 @@ impl Formatter { }) } - pub(super) fn write_reader_prefixes(&self, node: &Node, output: &mut String) { - for prefix in &node.reader_prefixes { - output.push_str(prefix.as_source()); + pub(super) fn write_reader_prefixes( + &self, + tree: &SyntaxTree, + node: &Node, + output: &mut String, + ) { + for span in &node.reader_prefix_spans { + output.push_str(span.slice(&tree.source)); } } pub(super) fn is_opaque_reader_form(&self, node: &Node) -> bool { - node.reader_prefixes - .iter() - .any(|prefix| matches!(prefix, ReaderPrefix::ReadEval)) + node.opaque_reader_form + || node + .reader_prefixes + .iter() + .any(|prefix| prefix.is_opaque_reader_form()) } pub(super) fn indent(&self, depth: usize) -> String { diff --git a/src/domain/sexpr/formatter/lists/definitions.rs b/src/domain/sexpr/formatter/lists/definitions.rs index bcbade93..8f9fce78 100644 --- a/src/domain/sexpr/formatter/lists/definitions.rs +++ b/src/domain/sexpr/formatter/lists/definitions.rs @@ -54,7 +54,11 @@ impl Formatter { } let delimiter = self.list_delimiter(node); let mut output = String::new(); - self.write_reader_prefixes(node, &mut output); + let reader_prefix_len = node + .reader_prefix_spans + .iter() + .map(|span| span.slice(&tree.source).len()) + .sum::(); output.push(delimiter.open()); for (position, child) in node.children.iter().enumerate() { if position > 0 { @@ -63,7 +67,7 @@ impl Formatter { output.push_str(&self.compact_node(tree, *child)?); } output.push(delimiter.close()); - (output.len() <= MAX_INLINE_WIDTH).then_some(output) + (output.len().saturating_add(reader_prefix_len) <= MAX_INLINE_WIDTH).then_some(output) } pub(in crate::domain::sexpr::formatter) fn format_definition( diff --git a/src/domain/sexpr/parser.rs b/src/domain/sexpr/parser.rs index 6abe9c70..d94f00d9 100644 --- a/src/domain/sexpr/parser.rs +++ b/src/domain/sexpr/parser.rs @@ -1,12 +1,17 @@ use thiserror::Error; +use crate::domain::dialect::Dialect; + +use super::reader_policy::{DialectReaderPolicy, ReaderMacro}; use super::tree::{Comment, Node, NodeKind, ReaderPrefix, SyntaxTree}; -use super::types::{ByteOffset, ByteSpan, Delimiter, NodeId, is_symbol_boundary}; +use super::types::{ByteOffset, ByteSpan, Delimiter, NodeId}; #[derive(Debug, Error, PartialEq, Eq)] pub enum ParseError { #[error("unexpected closing delimiter '{delimiter}' at byte {position}")] UnexpectedClose { delimiter: char, position: usize }, + #[error("unsupported reader dispatch '{dispatch}' at byte {position}")] + UnsupportedReaderDispatch { dispatch: String, position: usize }, #[error("mismatched closing delimiter '{found}' at byte {position}; expected '{expected}'")] MismatchedClose { found: char, @@ -40,6 +45,7 @@ pub(in crate::domain::sexpr) struct Parser<'a> { nodes: Vec, stack: Vec, comments: Vec, + policy: DialectReaderPolicy, /// Nesting depth of `#;` datum comments currently being skipped. While /// positive, inner trivia is folded into the enclosing datum comment span /// instead of being recorded as separate comments. @@ -49,7 +55,7 @@ pub(in crate::domain::sexpr) struct Parser<'a> { #[derive(Debug, Clone, Copy)] struct PrefixToken { kind: ReaderPrefix, - start: ByteOffset, + span: ByteSpan, } #[derive(Debug, Clone, Copy)] @@ -65,16 +71,22 @@ enum SkipFrame { impl<'a> Parser<'a> { pub(in crate::domain::sexpr) fn new(input: &'a str) -> Self { + Self::with_dialect(input, Dialect::Unknown) + } + + pub(in crate::domain::sexpr) fn with_dialect(input: &'a str, dialect: Dialect) -> Self { let root = Node { kind: NodeKind::Root, delimiter: None, reader_prefixes: Vec::new(), + reader_prefix_spans: Vec::new(), parent: None, children: Vec::new(), span: ByteSpan::new(ByteOffset::new(0), ByteOffset::new(input.len())), open: None, close: None, symbol_offset: 0, + opaque_reader_form: false, }; Self { input, @@ -83,6 +95,7 @@ impl<'a> Parser<'a> { nodes: vec![root], stack: vec![NodeId::ROOT], comments: Vec::new(), + policy: DialectReaderPolicy::new(dialect), suppress_depth: 0, } } @@ -136,7 +149,9 @@ impl<'a> Parser<'a> { fn skip_trivia(&mut self) -> std::result::Result<(), ParseError> { loop { - while self.pos.get() < self.bytes.len() && self.current_byte().is_ascii_whitespace() { + while self.pos.get() < self.bytes.len() + && self.policy.is_whitespace(self.current_byte()) + { self.advance(); } if self.pos.get() < self.bytes.len() && self.current_byte_is_block_comment() { @@ -145,13 +160,16 @@ impl<'a> Parser<'a> { self.record_comment(start, self.pos); continue; } - if self.pos.get() < self.bytes.len() && self.current_byte() == b';' { - let start = self.pos; - while self.pos.get() < self.bytes.len() && self.current_byte() != b'\n' { - self.advance(); + if self.pos.get() < self.bytes.len() { + if let Some(width) = self.policy.line_comment_width(self.bytes, self.pos.get()) { + let start = self.pos; + self.advance_by(width); + while self.pos.get() < self.bytes.len() && self.current_byte() != b'\n' { + self.advance(); + } + self.record_comment(start, self.pos); + continue; } - self.record_comment(start, self.pos); - continue; } break; } @@ -177,10 +195,14 @@ impl<'a> Parser<'a> { fn is_line_start(&self, start: usize) -> bool { let mut index = start; while index > 0 { - match self.bytes[index - 1] { - b'\n' => return true, - b' ' | b'\t' | b'\r' | b'\x0c' => index -= 1, - _ => return false, + let byte = self.bytes[index - 1]; + if byte == b'\n' { + return true; + } + if self.policy.is_whitespace(byte) { + index -= 1; + } else { + return false; } } true @@ -192,7 +214,7 @@ impl<'a> Parser<'a> { return; }; let id = NodeId::new(self.nodes.len()); - let Some(delimiter) = Delimiter::from_open(self.current_byte()) else { + let Some(delimiter) = self.policy.delimiter_from_open(self.current_byte()) else { debug_assert!( false, "open_list_with_prefixes called on non-opening delimiter" @@ -201,18 +223,22 @@ impl<'a> Parser<'a> { }; let start = prefixes .first() - .map(|prefix| prefix.start) + .map(|prefix| prefix.span.start()) .unwrap_or(self.pos); + let reader_prefixes = prefixes.iter().map(|prefix| prefix.kind).collect(); + let reader_prefix_spans = prefixes.iter().map(|prefix| prefix.span).collect(); self.nodes.push(Node { kind: NodeKind::List, delimiter: Some(delimiter), - reader_prefixes: prefixes.into_iter().map(|prefix| prefix.kind).collect(), + reader_prefixes, + reader_prefix_spans, parent: Some(parent), children: Vec::new(), span: ByteSpan::new(start, ByteOffset::new(self.pos.get() + 1)), open: Some(self.pos), close: None, symbol_offset: 0, + opaque_reader_form: false, }); self.nodes[parent.get()].children.push(id); self.stack.push(id); @@ -220,7 +246,7 @@ impl<'a> Parser<'a> { } fn close_list(&mut self) -> std::result::Result<(), ParseError> { - let Some(delimiter) = Delimiter::from_close(self.current_byte()) else { + let Some(delimiter) = self.policy.delimiter_from_close(self.current_byte()) else { debug_assert!(false, "close_list called on non-closing delimiter"); return Ok(()); }; @@ -294,31 +320,56 @@ impl<'a> Parser<'a> { prefixes: Vec, ) -> std::result::Result<(), ParseError> { let start = self.pos; - if self.current_byte_is_feature_dispatch() { - self.advance(); - self.advance(); + if self.policy.is_legacy() + && matches!( + self.policy + .classify_reader_macro(self.bytes, self.pos.get()), + Some(ReaderMacro::MultiDatum { .. }) + ) + { + self.advance_by(2); self.push_atom(prefixes, start, self.pos); return Ok(()); } + self.consume_atom_body()?; + self.push_atom(prefixes, start, self.pos); + Ok(()) + } + + fn consume_atom_body(&mut self) -> std::result::Result<(), ParseError> { + self.consume_character_literal(); while self.pos.get() < self.bytes.len() { let byte = self.current_byte(); - if byte == b'\\' { + if self.policy.supports_symbol_escapes() && byte == b'\\' { self.consume_single_escape()?; continue; } - if byte == b'|' { + if self.policy.supports_symbol_escapes() && byte == b'|' { self.consume_multiple_escape()?; continue; } - if is_symbol_boundary(byte) { + if self.policy.is_atom_boundary(self.bytes, self.pos.get()) { break; } self.advance(); } - self.push_atom(prefixes, start, self.pos); Ok(()) } + fn consume_character_literal(&mut self) { + let Some(prefix_width) = self + .policy + .character_literal_prefix_width(self.bytes, self.pos.get()) + else { + return; + }; + self.advance_by(prefix_width); + let Some(character) = self.input[self.pos.get()..].chars().next() else { + return; + }; + self.advance_by(character.len_utf8()); + } + /// Consumes a Lisp single-escape (`\`) and the following character literally. /// /// This keeps character literals such as `#\[`, `#\)`, and `#\Space`, as well @@ -365,30 +416,65 @@ impl<'a> Parser<'a> { return; }; let id = NodeId::new(self.nodes.len()); - let span_start = prefixes.first().map(|prefix| prefix.start).unwrap_or(start); + let span_start = prefixes + .first() + .map(|prefix| prefix.span.start()) + .unwrap_or(start); // `start` is the position after `consume_reader_prefixes` already ran // `skip_trivia()`, so this is the true start of the atom's own // content even when whitespace or a comment separates a reader // prefix from what it prefixes (`#' foo` is valid, if unusual, CL // syntax). let symbol_offset = start.get() - span_start.get(); + let reader_prefixes = prefixes.iter().map(|prefix| prefix.kind).collect(); + let reader_prefix_spans = prefixes.iter().map(|prefix| prefix.span).collect(); self.nodes.push(Node { kind: NodeKind::Atom, delimiter: None, - reader_prefixes: prefixes.into_iter().map(|prefix| prefix.kind).collect(), + reader_prefixes, + reader_prefix_spans, parent: Some(parent), children: Vec::new(), span: ByteSpan::new(span_start, end), open: None, close: None, symbol_offset, + opaque_reader_form: false, + }); + self.nodes[parent.get()].children.push(id); + } + + fn push_opaque_reader_form(&mut self, start: ByteOffset, end: ByteOffset) { + let Some(&parent) = self.stack.last() else { + debug_assert!( + false, + "parser stack unexpectedly empty when pushing reader form" + ); + return; + }; + let id = NodeId::new(self.nodes.len()); + self.nodes.push(Node { + kind: NodeKind::Atom, + delimiter: None, + reader_prefixes: Vec::new(), + reader_prefix_spans: Vec::new(), + parent: Some(parent), + children: Vec::new(), + span: ByteSpan::new(start, end), + open: None, + close: None, + symbol_offset: 0, + opaque_reader_form: true, }); self.nodes[parent.get()].children.push(id); } fn form(&mut self) -> std::result::Result<(), ParseError> { - if self.current_byte_is_reader_comment() { - self.skip_reader_comment()?; + if let Some(ReaderMacro::Discard { width }) = self + .policy + .classify_reader_macro(self.bytes, self.pos.get()) + { + self.skip_reader_comment(width)?; return Ok(()); } let mut prefixes = Vec::new(); @@ -397,22 +483,48 @@ impl<'a> Parser<'a> { if self.pos.get() >= self.bytes.len() { break; } - if !self.current_byte_is_reader_comment() { - break; + match self + .policy + .classify_reader_macro(self.bytes, self.pos.get()) + { + Some(ReaderMacro::Discard { width }) => { + self.skip_reader_comment(width)?; + self.skip_trivia()?; + } + _ => break, } - self.skip_reader_comment()?; - self.skip_trivia()?; } if self.pos.get() >= self.bytes.len() { let missing_at = prefixes .first() - .map(|prefix| prefix.start.get()) + .map(|prefix| prefix.span.start().get()) .unwrap_or(self.pos.get()); return Err(ParseError::MissingReaderForm(missing_at)); } + match self + .policy + .classify_reader_macro(self.bytes, self.pos.get()) + { + Some(ReaderMacro::MultiDatum { + width, + payload_forms, + }) if !self.policy.is_legacy() => { + self.opaque_reader_form_with_prefixes(prefixes, width, payload_forms)?; + return Ok(()); + } + Some(ReaderMacro::UnsupportedDispatch { width }) => { + return Err(self.unsupported_reader_error(width)); + } + _ => {} + } match self.current_byte() { - byte if Delimiter::from_open(byte).is_some() => self.open_list_with_prefixes(prefixes), - byte if Delimiter::from_close(byte).is_some() => self.close_list()?, + byte if self.policy.delimiter_from_open(byte).is_some() => { + self.open_list_with_prefixes(prefixes) + } + byte if self.policy.delimiter_from_close(byte).is_some() => self.close_list()?, + byte if DialectReaderPolicy::is_raw_delimiter(byte) => { + return Err(self.raw_delimiter_error()); + } b'"' => self.atom_string_with_prefixes(prefixes)?, _ => self.atom_with_prefixes(prefixes)?, } @@ -421,10 +533,19 @@ impl<'a> Parser<'a> { fn consume_reader_prefixes(&mut self) -> std::result::Result, ParseError> { let mut prefixes = Vec::new(); - while let Some(kind) = self.current_reader_prefix() { + while let Some(ReaderMacro::Prefix { + semantic: kind, + width, + }) = self + .policy + .classify_reader_macro(self.bytes, self.pos.get()) + { let start = self.pos; - self.advance_reader_prefix(kind); - prefixes.push(PrefixToken { kind, start }); + self.advance_by(width); + prefixes.push(PrefixToken { + kind, + span: ByteSpan::new(start, self.pos), + }); self.skip_trivia()?; if self.pos.get() >= self.bytes.len() { break; @@ -433,10 +554,9 @@ impl<'a> Parser<'a> { Ok(prefixes) } - fn skip_reader_comment(&mut self) -> std::result::Result<(), ParseError> { + fn skip_reader_comment(&mut self, width: usize) -> std::result::Result<(), ParseError> { let start = self.pos; - self.advance(); - self.advance(); + self.advance_by(width); self.suppress_depth += 1; let result = self.skip_form(start.get()); self.suppress_depth -= 1; @@ -446,6 +566,26 @@ impl<'a> Parser<'a> { Ok(()) } + fn opaque_reader_form_with_prefixes( + &mut self, + prefixes: Vec, + width: usize, + payload_forms: usize, + ) -> std::result::Result<(), ParseError> { + let start = prefixes + .first() + .map(|prefix| prefix.span.start()) + .unwrap_or(self.pos); + let reader_start = self.pos.get(); + self.advance_by(width); + self.suppress_depth += 1; + let result = (0..payload_forms).try_for_each(|_| self.skip_form(reader_start)); + self.suppress_depth -= 1; + result?; + self.push_opaque_reader_form(start, self.pos); + Ok(()) + } + fn skip_form(&mut self, missing_at: usize) -> std::result::Result<(), ParseError> { let mut frames = Vec::new(); Self::push_skip_frame(&mut frames, SkipFrame::Form { missing_at }, self.pos.get())?; @@ -458,9 +598,15 @@ impl<'a> Parser<'a> { } let mut prefix_start = None; - while let Some(prefix) = self.current_reader_prefix() { + let mut additional_discarded_forms = 0; + while let Some(ReaderMacro::Prefix { semantic, width }) = self + .policy + .classify_reader_macro(self.bytes, self.pos.get()) + { prefix_start.get_or_insert(self.pos.get()); - self.advance_reader_prefix(prefix); + additional_discarded_forms += + self.policy.additional_discarded_forms_for_prefix(semantic); + self.advance_by(width); self.skip_trivia()?; if self.pos.get() >= self.bytes.len() { return Err(ParseError::MissingReaderForm( @@ -469,55 +615,68 @@ impl<'a> Parser<'a> { } } - if self.current_byte_is_reader_comment() { - let comment_start = self.pos.get(); - self.advance(); - self.advance(); + for _ in 0..additional_discarded_forms { Self::push_skip_frame( &mut frames, SkipFrame::Form { missing_at: prefix_start.unwrap_or(missing_at), }, - comment_start, + self.pos.get(), )?; - Self::push_skip_frame( - &mut frames, - SkipFrame::Form { - missing_at: comment_start, - }, - comment_start, - )?; - continue; } - if self.current_byte_is_feature_dispatch() { - let dispatch_start = self.pos.get(); - self.advance(); - self.advance(); - // A feature conditional is one reader form even though the - // normal syntax tree exposes its three parts as siblings. - // The guarded datum is pushed first because frames are LIFO. - Self::push_skip_frame( - &mut frames, - SkipFrame::Form { - missing_at: dispatch_start, - }, - dispatch_start, - )?; - Self::push_skip_frame( - &mut frames, - SkipFrame::Form { - missing_at: dispatch_start, - }, - dispatch_start, - )?; - continue; + match self + .policy + .classify_reader_macro(self.bytes, self.pos.get()) + { + Some(ReaderMacro::Discard { width }) => { + let comment_start = self.pos.get(); + self.advance_by(width); + Self::push_skip_frame( + &mut frames, + SkipFrame::Form { + missing_at: prefix_start.unwrap_or(missing_at), + }, + comment_start, + )?; + Self::push_skip_frame( + &mut frames, + SkipFrame::Form { + missing_at: comment_start, + }, + comment_start, + )?; + continue; + } + Some(ReaderMacro::MultiDatum { + width, + payload_forms, + }) => { + let dispatch_start = self.pos.get(); + self.advance_by(width); + for _ in 0..payload_forms { + Self::push_skip_frame( + &mut frames, + SkipFrame::Form { + missing_at: dispatch_start, + }, + dispatch_start, + )?; + } + continue; + } + Some(ReaderMacro::UnsupportedDispatch { width }) => { + return Err(self.unsupported_reader_error(width)); + } + Some(ReaderMacro::Prefix { .. }) | None => {} } match self.current_byte() { - byte if Delimiter::from_open(byte).is_some() => { + byte if self.policy.delimiter_from_open(byte).is_some() => { let open_pos = self.pos; - let expected_close = Delimiter::from_open(byte) + let expected_close = self + .policy + .delimiter_from_open(byte) .expect("opening delimiter checked above") .close() as u8; self.advance(); @@ -530,14 +689,19 @@ impl<'a> Parser<'a> { open_pos.get(), )?; } - byte if Delimiter::from_close(byte).is_some() => { - let delimiter = Delimiter::from_close(byte) + byte if self.policy.delimiter_from_close(byte).is_some() => { + let delimiter = self + .policy + .delimiter_from_close(byte) .expect("closing delimiter checked above"); return Err(ParseError::UnexpectedClose { delimiter: delimiter.close(), position: self.pos.get(), }); } + byte if DialectReaderPolicy::is_raw_delimiter(byte) => { + return Err(self.raw_delimiter_error()); + } b'"' => self.skip_string()?, _ => self.skip_atom()?, } @@ -641,86 +805,13 @@ impl<'a> Parser<'a> { } fn skip_atom(&mut self) -> std::result::Result<(), ParseError> { - while self.pos.get() < self.bytes.len() { - let byte = self.current_byte(); - if byte == b'\\' { - self.consume_single_escape()?; - continue; - } - if byte == b'|' { - self.consume_multiple_escape()?; - continue; - } - if is_symbol_boundary(byte) { - break; - } - self.advance(); - } - Ok(()) - } - - /// `#+`/`#-` (CLHS 2.4.8/2.4.9 feature-conditional dispatch) must scan as - /// their own fixed two-byte token, distinct from the feature expression - /// that follows. Without this, `#+sbcl` glues into one opaque atom while - /// `#+(and sbcl x86-64)` splits at the list delimiter, so equivalent - /// feature conditionals produce inconsistent tree shapes: the guarded - /// feature symbol is findable/renameable in one spelling but hidden - /// inside an opaque token in the other. - fn current_byte_is_feature_dispatch(&self) -> bool { - self.current_byte() == b'#' && matches!(self.peek_byte(), Some(b'+') | Some(b'-')) - } - - /// Matches Scheme/Common Lisp `#;` datum comments and Clojure `#_` - /// discard forms. Both are two-byte dispatch macros that read and - /// discard exactly one following form, so they share the same skip - /// path and are recorded as comments rather than tree nodes. - fn current_byte_is_reader_comment(&self) -> bool { - self.current_byte() == b'#' && matches!(self.peek_byte(), Some(b';') | Some(b'_')) + self.consume_atom_body() } fn current_byte_is_block_comment(&self) -> bool { - self.current_byte() == b'#' && self.peek_byte() == Some(b'|') - } - - fn current_reader_prefix(&self) -> Option { - if self.pos.get() >= self.bytes.len() { - return None; - } - match self.current_byte() { - b'\'' => Some(ReaderPrefix::Quote), - b'`' => Some(ReaderPrefix::Quasiquote), - b',' if self.peek_byte() == Some(b'@') => Some(ReaderPrefix::UnquoteSplicing), - b',' => Some(ReaderPrefix::Unquote), - b'^' => Some(ReaderPrefix::Metadata), - b'#' if self.peek_byte() == Some(b'.') => Some(ReaderPrefix::ReadEval), - b'#' if self.peek_byte() == Some(b'\'') => Some(ReaderPrefix::Function), - b'#' if self.peek_byte() == Some(b'?') && self.peek_byte_at(2) == Some(b'@') => { - Some(ReaderPrefix::ReaderConditionalSplicing) - } - b'#' if self.peek_byte() == Some(b'?') => Some(ReaderPrefix::ReaderConditional), - b'#' if Delimiter::from_open(self.peek_byte().unwrap_or(0)).is_some() => { - Some(ReaderPrefix::HashLiteral) - } - _ => None, - } - } - - fn advance_reader_prefix(&mut self, prefix: ReaderPrefix) { - let width = match prefix { - ReaderPrefix::ReaderConditionalSplicing => 3, - ReaderPrefix::UnquoteSplicing - | ReaderPrefix::Function - | ReaderPrefix::ReadEval - | ReaderPrefix::ReaderConditional => 2, - ReaderPrefix::Quote - | ReaderPrefix::Quasiquote - | ReaderPrefix::Unquote - | ReaderPrefix::HashLiteral - | ReaderPrefix::Metadata => 1, - }; - for _ in 0..width { - self.advance(); - } + self.policy.supports_block_comments() + && self.current_byte() == b'#' + && self.peek_byte() == Some(b'|') } fn current_byte(&self) -> u8 { @@ -731,11 +822,31 @@ impl<'a> Parser<'a> { self.bytes.get(self.pos.get() + 1).copied() } - fn peek_byte_at(&self, offset: usize) -> Option { - self.bytes.get(self.pos.get() + offset).copied() - } - fn advance(&mut self) { self.pos = ByteOffset::new(self.pos.get() + 1); } + + fn advance_by(&mut self, width: usize) { + self.pos = ByteOffset::new(self.pos.get() + width); + } + + fn unsupported_reader_error(&self, width: usize) -> ParseError { + let start = self.pos.get(); + let end = start.saturating_add(width).min(self.input.len()); + ParseError::UnsupportedReaderDispatch { + dispatch: self.input[start..end].to_owned(), + position: start, + } + } + + fn raw_delimiter_error(&self) -> ParseError { + let start = self.pos.get(); + ParseError::UnexpectedClose { + delimiter: self.input[start..] + .chars() + .next() + .unwrap_or(char::REPLACEMENT_CHARACTER), + position: start, + } + } } diff --git a/src/domain/sexpr/reader_policy.rs b/src/domain/sexpr/reader_policy.rs new file mode 100644 index 00000000..beedee42 --- /dev/null +++ b/src/domain/sexpr/reader_policy.rs @@ -0,0 +1,409 @@ +use crate::domain::dialect::Dialect; + +use super::tree::ReaderPrefix; +use super::types::Delimiter; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(super) enum ReaderMacro { + Prefix { + semantic: ReaderPrefix, + width: usize, + }, + Discard { + width: usize, + }, + MultiDatum { + width: usize, + payload_forms: usize, + }, + UnsupportedDispatch { + width: usize, + }, +} + +/// Dialect-specific lexical decisions shared by normal parsing and discarded +/// form scanning. Keeping these decisions in one place prevents the two paths +/// from disagreeing about the extent of a reader form. +#[derive(Debug, Clone, Copy)] +pub(super) struct DialectReaderPolicy { + dialect: Dialect, +} + +impl DialectReaderPolicy { + pub(super) const fn new(dialect: Dialect) -> Self { + Self { dialect } + } + + pub(super) const fn is_legacy(self) -> bool { + matches!(self.dialect, Dialect::Unknown) + } + + pub(super) const fn additional_discarded_forms_for_prefix(self, prefix: ReaderPrefix) -> usize { + if matches!( + (self.dialect, prefix), + (Dialect::Clojure, ReaderPrefix::Metadata) + ) { + 1 + } else { + 0 + } + } + + pub(super) fn is_whitespace(self, byte: u8) -> bool { + byte.is_ascii_whitespace() || matches!(self.dialect, Dialect::Clojure) && byte == b',' + } + + pub(super) fn line_comment_width(self, bytes: &[u8], pos: usize) -> Option { + let byte = *bytes.get(pos)?; + match self.dialect { + Dialect::Janet if byte == b'#' => Some(1), + Dialect::Janet => None, + _ if byte == b';' => Some(1), + _ => None, + } + } + + pub(super) fn supports_block_comments(self) -> bool { + matches!( + self.dialect, + Dialect::CommonLisp | Dialect::Scheme | Dialect::Unknown + ) + } + + pub(super) fn supports_symbol_escapes(self) -> bool { + matches!(self.dialect, Dialect::CommonLisp | Dialect::Unknown) + } + + pub(super) fn delimiter_from_open(self, byte: u8) -> Option { + let delimiter = Delimiter::from_open(byte)?; + self.allows_delimiter(delimiter).then_some(delimiter) + } + + pub(super) fn delimiter_from_close(self, byte: u8) -> Option { + let delimiter = Delimiter::from_close(byte)?; + self.allows_delimiter(delimiter).then_some(delimiter) + } + + pub(super) fn is_raw_delimiter(byte: u8) -> bool { + Delimiter::from_open(byte).is_some() || Delimiter::from_close(byte).is_some() + } + + pub(super) fn is_atom_boundary(self, bytes: &[u8], pos: usize) -> bool { + bytes.get(pos).is_none_or(|byte| { + self.is_whitespace(*byte) + || Self::is_raw_delimiter(*byte) + || self.line_comment_width(bytes, pos).is_some() + }) + } + + pub(super) fn character_literal_prefix_width(self, bytes: &[u8], pos: usize) -> Option { + let byte = *bytes.get(pos)?; + let next = bytes.get(pos + 1).copied(); + match self.dialect { + Dialect::Scheme if byte == b'#' && next == Some(b'\\') => Some(2), + Dialect::Clojure if byte == b'\\' => Some(1), + Dialect::EmacsLisp if byte == b'?' && next == Some(b'\\') => Some(2), + Dialect::EmacsLisp if byte == b'?' => Some(1), + _ => None, + } + } + + pub(super) fn classify_reader_macro(self, bytes: &[u8], pos: usize) -> Option { + let byte = *bytes.get(pos)?; + let next = bytes.get(pos + 1).copied(); + let third = bytes.get(pos + 2).copied(); + + match self.dialect { + Dialect::Unknown => self.classify_legacy(byte, next, third), + Dialect::CommonLisp => self.classify_common_lisp(bytes, pos), + Dialect::EmacsLisp => self.classify_emacs_lisp(byte, next), + Dialect::Scheme => self.classify_scheme(bytes, pos), + Dialect::Clojure => self.classify_clojure(bytes, pos), + Dialect::Janet => self.classify_janet(byte, next), + Dialect::Fennel => self.classify_fennel(byte, next), + } + } + + fn allows_delimiter(self, delimiter: Delimiter) -> bool { + match self.dialect { + Dialect::CommonLisp | Dialect::Scheme => matches!(delimiter, Delimiter::Paren), + Dialect::EmacsLisp => matches!(delimiter, Delimiter::Paren | Delimiter::Bracket), + Dialect::Clojure | Dialect::Janet | Dialect::Fennel | Dialect::Unknown => true, + } + } + + fn classify_legacy(self, byte: u8, next: Option, third: Option) -> Option { + if byte == b'#' && matches!(next, Some(b';' | b'_')) { + return Some(ReaderMacro::Discard { width: 2 }); + } + if byte == b'#' && matches!(next, Some(b'+' | b'-')) { + return Some(ReaderMacro::MultiDatum { + width: 2, + payload_forms: 2, + }); + } + classify_shared_prefix(byte, next, third) + } + + fn classify_common_lisp(self, bytes: &[u8], pos: usize) -> Option { + let byte = *bytes.get(pos)?; + let next = bytes.get(pos + 1).copied(); + + if byte == b'#' && matches!(next, Some(b'+' | b'-')) { + return Some(ReaderMacro::MultiDatum { + width: 2, + payload_forms: 2, + }); + } + if let Some(prefix) = classify_quote_prefix(byte, next) { + return Some(prefix); + } + if byte != b'#' { + return None; + } + if let Some(dispatch) = classify_numeric_dispatch(bytes, pos, true) { + return Some(dispatch); + } + if is_numeric_radix_dispatch(bytes, pos) { + return None; + } + match next { + Some(b':' | b'\\' | b'*' | b'b' | b'B' | b'o' | b'O' | b'd' | b'D' | b'x' | b'X') => { + None + } + Some(b'p' | b'P' | b's' | b'S') => Some(ReaderMacro::MultiDatum { + width: 2, + payload_forms: 1, + }), + Some(b'\'') => prefix(ReaderPrefix::Function, 2), + Some(b'.') => prefix(ReaderPrefix::ReadEval, 2), + Some(b'(') => prefix(ReaderPrefix::HashLiteral, 1), + _ => Some(ReaderMacro::UnsupportedDispatch { width: 1 }), + } + } + + fn classify_emacs_lisp(self, byte: u8, next: Option) -> Option { + if let Some(prefix) = classify_quote_prefix(byte, next) { + return Some(prefix); + } + if byte != b'#' { + return None; + } + match next { + Some(b'\'') => prefix(ReaderPrefix::Function, 2), + _ => Some(ReaderMacro::UnsupportedDispatch { width: 1 }), + } + } + + fn classify_scheme(self, bytes: &[u8], pos: usize) -> Option { + let byte = *bytes.get(pos)?; + let next = bytes.get(pos + 1).copied(); + if byte == b'#' && next == Some(b';') { + return Some(ReaderMacro::Discard { width: 2 }); + } + if let Some(prefix) = classify_quote_prefix(byte, next) { + return Some(prefix); + } + if byte != b'#' { + return None; + } + if let Some(dispatch) = classify_numeric_dispatch(bytes, pos, false) { + return Some(dispatch); + } + if matches!(next, Some(b'u' | b'U')) + && bytes.get(pos + 2) == Some(&b'8') + && bytes.get(pos + 3) == Some(&b'(') + { + return Some(ReaderMacro::MultiDatum { + width: 3, + payload_forms: 1, + }); + } + match next { + Some(b'(') => prefix(ReaderPrefix::HashLiteral, 1), + Some( + b'\\' | b't' | b'T' | b'f' | b'F' | b'b' | b'B' | b'o' | b'O' | b'd' | b'D' | b'x' + | b'X' | b'e' | b'E' | b'i' | b'I', + ) => None, + _ => Some(ReaderMacro::UnsupportedDispatch { width: 1 }), + } + } + + fn classify_clojure(self, bytes: &[u8], pos: usize) -> Option { + let byte = *bytes.get(pos)?; + let next = bytes.get(pos + 1).copied(); + let third = bytes.get(pos + 2).copied(); + let fourth = bytes.get(pos + 3).copied(); + match byte { + b'\'' => prefix(ReaderPrefix::Quote, 1), + b'`' => prefix(ReaderPrefix::Quasiquote, 1), + b'~' if next == Some(b'@') => prefix(ReaderPrefix::UnquoteSplicing, 2), + b'~' => prefix(ReaderPrefix::Unquote, 1), + b'@' => prefix(ReaderPrefix::Function, 1), + b'^' => prefix(ReaderPrefix::Metadata, 1), + b'#' if next == Some(b'_') => Some(ReaderMacro::Discard { width: 2 }), + b'#' if next == Some(b'?') && third == Some(b'@') && fourth == Some(b'(') => { + prefix(ReaderPrefix::ReaderConditionalSplicing, 3) + } + b'#' if next == Some(b'?') && third == Some(b'(') => { + prefix(ReaderPrefix::ReaderConditional, 2) + } + b'#' if next == Some(b'?') => Some(ReaderMacro::UnsupportedDispatch { + width: usize::from(third == Some(b'@')) + 2, + }), + b'#' if next == Some(b'\'') => prefix(ReaderPrefix::Function, 2), + b'#' if matches!(next, Some(b'(' | b'{')) => prefix(ReaderPrefix::HashLiteral, 1), + b'#' if next == Some(b'"') => Some(ReaderMacro::MultiDatum { + width: 1, + payload_forms: 1, + }), + b'#' if next == Some(b':') => self + .clojure_namespaced_map_width(bytes, pos) + .map(|width| ReaderMacro::MultiDatum { + width, + payload_forms: 1, + }) + .or(Some(ReaderMacro::UnsupportedDispatch { width: 1 })), + b'#' if next == Some(b'#') => None, + b'#' => self + .clojure_tagged_literal_width(bytes, pos) + .map(|width| ReaderMacro::MultiDatum { + width, + payload_forms: 1, + }) + .or(Some(ReaderMacro::UnsupportedDispatch { width: 1 })), + _ => None, + } + } + + fn clojure_namespaced_map_width(self, bytes: &[u8], pos: usize) -> Option { + let mut cursor = pos + 2; + let auto_resolved = bytes.get(cursor) == Some(&b':'); + if auto_resolved { + cursor += 1; + } + let namespace_start = cursor; + while let Some(&byte) = bytes.get(cursor) { + if byte == b'{' { + return (auto_resolved || cursor > namespace_start).then_some(cursor - pos); + } + if self.is_atom_boundary(bytes, cursor) { + return None; + } + cursor += 1; + } + None + } + + fn clojure_tagged_literal_width(self, bytes: &[u8], pos: usize) -> Option { + let first = *bytes.get(pos + 1)?; + if !(first.is_ascii_alphabetic() + || matches!( + first, + b'*' | b'+' | b'!' | b'-' | b'_' | b'\'' | b'?' | b'<' | b'>' | b'=' + )) + { + return None; + } + + let mut cursor = pos + 2; + while !self.is_atom_boundary(bytes, cursor) { + cursor += 1; + } + + let tag = &bytes[pos + 1..cursor]; + if tag.last() == Some(&b'/') || tag.iter().filter(|byte| **byte == b'/').count() > 1 { + return None; + } + Some(cursor - pos) + } + + fn classify_janet(self, byte: u8, next: Option) -> Option { + match byte { + b';' => prefix(ReaderPrefix::UnquoteSplicing, 1), + b'~' => prefix(ReaderPrefix::Quasiquote, 1), + b',' => prefix(ReaderPrefix::Unquote, 1), + b'|' => prefix(ReaderPrefix::Function, 1), + b'@' => prefix(ReaderPrefix::HashLiteral, 1), + // `#` is consumed as a line comment before reader classification. + b'#' if next.is_some() => None, + _ => None, + } + } + + fn classify_fennel(self, byte: u8, next: Option) -> Option { + match byte { + b'\'' => prefix(ReaderPrefix::Quote, 1), + b'`' => prefix(ReaderPrefix::Quasiquote, 1), + b',' if next == Some(b'@') => prefix(ReaderPrefix::UnquoteSplicing, 2), + b',' => prefix(ReaderPrefix::Unquote, 1), + b'#' => prefix(ReaderPrefix::Function, 1), + _ => None, + } + } +} + +fn classify_shared_prefix(byte: u8, next: Option, third: Option) -> Option { + if let Some(prefix) = classify_quote_prefix(byte, next) { + return Some(prefix); + } + match (byte, next, third) { + (b'^', _, _) => prefix(ReaderPrefix::Metadata, 1), + (b'#', Some(b'.'), _) => prefix(ReaderPrefix::ReadEval, 2), + (b'#', Some(b'\''), _) => prefix(ReaderPrefix::Function, 2), + (b'#', Some(b'?'), Some(b'@')) => prefix(ReaderPrefix::ReaderConditionalSplicing, 3), + (b'#', Some(b'?'), _) => prefix(ReaderPrefix::ReaderConditional, 2), + (b'#', Some(b'(' | b'[' | b'{'), _) => prefix(ReaderPrefix::HashLiteral, 1), + _ => None, + } +} + +fn classify_numeric_dispatch(bytes: &[u8], pos: usize, allow_array: bool) -> Option { + if bytes.get(pos) != Some(&b'#') { + return None; + } + + let mut marker_pos = pos + 1; + while matches!(bytes.get(marker_pos), Some(byte) if byte.is_ascii_digit()) { + marker_pos += 1; + } + + let has_numeric_argument = marker_pos > pos + 1; + let payload_forms = match bytes.get(marker_pos).copied() { + Some(b'=') if has_numeric_argument => 1, + Some(b'#') if has_numeric_argument => 0, + Some(b'a' | b'A') if allow_array => 1, + _ => return None, + }; + Some(ReaderMacro::MultiDatum { + width: marker_pos - pos + 1, + payload_forms, + }) +} + +fn is_numeric_radix_dispatch(bytes: &[u8], pos: usize) -> bool { + if bytes.get(pos) != Some(&b'#') { + return false; + } + + let mut marker_pos = pos + 1; + while matches!(bytes.get(marker_pos), Some(byte) if byte.is_ascii_digit()) { + marker_pos += 1; + } + + marker_pos > pos + 1 && matches!(bytes.get(marker_pos), Some(b'r' | b'R')) +} + +fn classify_quote_prefix(byte: u8, next: Option) -> Option { + match (byte, next) { + (b'\'', _) => prefix(ReaderPrefix::Quote, 1), + (b'`', _) => prefix(ReaderPrefix::Quasiquote, 1), + (b',', Some(b'@')) => prefix(ReaderPrefix::UnquoteSplicing, 2), + (b',', _) => prefix(ReaderPrefix::Unquote, 1), + _ => None, + } +} + +const fn prefix(semantic: ReaderPrefix, width: usize) -> Option { + Some(ReaderMacro::Prefix { semantic, width }) +} diff --git a/src/domain/sexpr/tests/edit.rs b/src/domain/sexpr/tests/edit.rs index 72c9a9b7..794fc937 100644 --- a/src/domain/sexpr/tests/edit.rs +++ b/src/domain/sexpr/tests/edit.rs @@ -1,4 +1,5 @@ use super::*; +use crate::domain::dialect::Dialect; #[test] fn replaces_expression() { @@ -233,7 +234,7 @@ fn normalizes_trailing_trivia_only_on_changed_lines() { let rewritten = "(alpha gamma) \n(unchanged) \n".to_owned(); assert_eq!( - Edit::normalize_changed_line_trivia(input, rewritten).unwrap(), + Edit::normalize_changed_line_trivia(input, rewritten, Dialect::CommonLisp).unwrap(), "(alpha gamma)\n(unchanged) \n" ); } @@ -244,7 +245,60 @@ fn preserves_trailing_spaces_inside_multiline_atoms_and_block_comments() { let rewritten = "(print \"new \nvalue\")\n#| new \ncomment |#\n".to_owned(); assert_eq!( - Edit::normalize_changed_line_trivia(input, rewritten.clone()).unwrap(), + Edit::normalize_changed_line_trivia(input, rewritten.clone(), Dialect::CommonLisp).unwrap(), rewritten ); } + +#[test] +fn preserves_changed_line_trivia_inside_clojure_discard() { + let input = "#_(old \n \\)) kept\n"; + let rewritten = "#_(new \n \\)) kept\n".to_owned(); + + assert_eq!( + Edit::normalize_changed_line_trivia(input, rewritten.clone(), Dialect::Clojure).unwrap(), + rewritten + ); +} + +#[test] +fn preserves_changed_line_trivia_inside_clojure_tagged_literal() { + let input = "#inst \n\"1985-04-12T23:20:50.52-00:00\"\n"; + let rewritten = "#inst \n\"1986-04-12T23:20:50.52-00:00\"\n".to_owned(); + + assert_eq!( + Edit::normalize_changed_line_trivia(input, rewritten.clone(), Dialect::Clojure).unwrap(), + rewritten + ); +} + +#[test] +fn preserves_changed_line_trivia_inside_common_lisp_dispatch() { + let input = "#S \n(point :x 1)\n"; + let rewritten = "#S \n(point :x 2)\n".to_owned(); + + assert_eq!( + Edit::normalize_changed_line_trivia(input, rewritten.clone(), Dialect::CommonLisp).unwrap(), + rewritten + ); +} + +#[test] +fn unknown_dialect_keeps_generic_compatibility_and_rejects_malformed_changes() { + let normalized = Edit::normalize_changed_line_trivia( + "(alpha beta) \n", + "(alpha gamma) \n".to_owned(), + Dialect::Unknown, + ) + .expect("generic changed document remains supported"); + assert_eq!(normalized, "(alpha gamma)\n"); + + assert!( + Edit::normalize_changed_line_trivia( + "(alpha beta)", + "(alpha gamma".to_owned(), + Dialect::Unknown, + ) + .is_err() + ); +} diff --git a/src/domain/sexpr/tests/formatter.rs b/src/domain/sexpr/tests/formatter.rs index c86f05b7..b43173ea 100644 --- a/src/domain/sexpr/tests/formatter.rs +++ b/src/domain/sexpr/tests/formatter.rs @@ -1,4 +1,5 @@ use super::*; +use crate::domain::dialect::Dialect; #[test] fn formats_short_atom_lists_inline() { @@ -61,6 +62,47 @@ fn preserves_common_lisp_reader_prefixes() { ); } +#[test] +fn preserves_dialect_reader_prefix_spellings() { + let cases = [ + (Dialect::Janet, ";(value)", ";(value)\n"), + (Dialect::Fennel, "#(value)", "#(value)\n"), + ]; + + for (dialect, input, expected) in cases { + let tree = SyntaxTree::parse_with_dialect(input, dialect).expect("valid reader form"); + assert_eq!( + Formatter::new(2).format(&tree), + expected, + "{}", + dialect.label() + ); + } +} + +#[test] +fn preserves_multi_datum_reader_forms_verbatim() { + let cases = [ + (Dialect::CommonLisp, "#+feature (guarded value)"), + (Dialect::Clojure, "^:private target"), + (Dialect::Clojure, r#"^{:doc "x"} target"#), + (Dialect::Scheme, "#u8(1 2 3)"), + (Dialect::Clojure, r##"#"foo.*""##), + (Dialect::Clojure, r#"#:person{:first "Ada"}"#), + (Dialect::Clojure, r#"#inst "2020-01-01""#), + ]; + + for (dialect, input) in cases { + let tree = SyntaxTree::parse_with_dialect(input, dialect).expect("valid reader form"); + assert_eq!( + Formatter::new(2).format(&tree), + format!("{input}\n"), + "{}", + dialect.label() + ); + } +} + #[test] fn preserves_common_lisp_reader_eval_forms_verbatim() { let input = "#.(foo (bar baz))\n#.(list 1 2 3)"; @@ -252,6 +294,16 @@ fn keeps_short_defsystem_forms_on_one_line() { ); } +#[test] +fn preserves_reader_prefix_on_short_defsystem_idempotently() { + let tree = SyntaxTree::parse("'(defsystem x)").expect("valid"); + let formatted = Formatter::new(2).format(&tree); + assert_eq!(formatted, "'(defsystem x)\n"); + + let reparsed = SyntaxTree::parse(&formatted).expect("formatted output is valid"); + assert_eq!(Formatter::new(2).format(&reparsed), formatted); +} + #[test] fn breaks_long_defsystem_forms_keeping_option_pairs_together() { let input = "(defsystem \"my-really-quite-long-system-name\" :description \"a considerably longer description string here\" :version \"0.1.0\" :depends-on (:alexandria :bordeaux-threads))"; diff --git a/src/domain/sexpr/tests/parser.rs b/src/domain/sexpr/tests/parser.rs index 7b69a415..e57b1180 100644 --- a/src/domain/sexpr/tests/parser.rs +++ b/src/domain/sexpr/tests/parser.rs @@ -1,4 +1,5 @@ use super::*; +use crate::domain::dialect::Dialect; use crate::domain::sexpr::parser::MAX_DISCARDED_FORM_STACK_FRAMES; #[test] @@ -7,6 +8,311 @@ fn parses_balanced_document() { assert_eq!(tree.root_children().len(), 1); } +#[test] +fn applies_dialect_reader_collisions_without_splitting_reader_forms() { + struct Case { + dialect: Dialect, + input: &'static str, + delimiter: Delimiter, + children: &'static [&'static str], + } + + let cases = [ + Case { + dialect: Dialect::CommonLisp, + input: "(#+feature guarded tail)", + delimiter: Delimiter::Paren, + children: &["#+feature guarded", "tail"], + }, + Case { + dialect: Dialect::EmacsLisp, + input: "[#'f tail]", + delimiter: Delimiter::Bracket, + children: &["#'f", "tail"], + }, + Case { + dialect: Dialect::Scheme, + input: "(#;discard kept)", + delimiter: Delimiter::Paren, + children: &["kept"], + }, + Case { + dialect: Dialect::Clojure, + input: "{left,right}", + delimiter: Delimiter::Brace, + children: &["left", "right"], + }, + Case { + dialect: Dialect::Janet, + input: "[;value # ignored\n next]", + delimiter: Delimiter::Bracket, + children: &[";value", "next"], + }, + Case { + dialect: Dialect::Fennel, + input: "{#(value) tail}", + delimiter: Delimiter::Brace, + children: &["#(value)", "tail"], + }, + ]; + + for case in cases { + let tree = SyntaxTree::parse_with_dialect(case.input, case.dialect) + .unwrap_or_else(|error| panic!("{}: {error}", case.dialect.label())); + let root = tree.root_view(); + assert_eq!(root.children.len(), 1, "{}", case.dialect.label()); + let form = &root.children[0]; + assert_eq!( + form.delimiter, + Some(case.delimiter), + "{}", + case.dialect.label() + ); + let children = form + .children + .iter() + .map(|child| child.span.slice(case.input)) + .collect::>(); + assert_eq!(children, case.children, "{}", case.dialect.label()); + } +} + +#[test] +fn multi_datum_reader_forms_are_single_siblings() { + let cases = [ + ( + Dialect::CommonLisp, + "(#+feature (guarded value) tail)", + &["#+feature (guarded value)", "tail"] as &[&str], + ), + ( + Dialect::Clojure, + "(^:private target tail)", + &["^:private", "target", "tail"] as &[&str], + ), + ]; + + for (dialect, input, expected) in cases { + let tree = SyntaxTree::parse_with_dialect(input, dialect).expect("valid reader form"); + let form = &tree.root_view().children[0]; + let children = form + .children + .iter() + .map(|child| child.span.slice(input)) + .collect::>(); + assert_eq!(children, expected, "{}", dialect.label()); + } +} + +#[test] +fn unsupported_dispatch_fails_closed_in_live_and_discarded_forms() { + let cases = [ + (Dialect::CommonLisp, "#?value"), + (Dialect::CommonLisp, "#12Q"), + (Dialect::CommonLisp, "#12Qvalue"), + (Dialect::EmacsLisp, "#(value)"), + (Dialect::Scheme, "#_value"), + (Dialect::Scheme, "#12Qvalue"), + (Dialect::Clojure, "#;value"), + (Dialect::Clojure, "#?value"), + (Dialect::Clojure, "#12Qvalue"), + (Dialect::CommonLisp, "#+feature #?value"), + (Dialect::CommonLisp, "#+feature #12Q"), + (Dialect::CommonLisp, "#+feature #12Qvalue"), + (Dialect::Scheme, "#;#?value"), + (Dialect::Scheme, "#;#12Qvalue"), + (Dialect::Clojure, "#_#;value"), + (Dialect::Clojure, "#_#?value"), + (Dialect::Clojure, "#_#12Qvalue"), + ]; + + for (dialect, input) in cases { + let error = SyntaxTree::parse_with_dialect(input, dialect).unwrap_err(); + assert!( + matches!(error, ParseError::UnsupportedReaderDispatch { .. }), + "{} returned the wrong error for {input}: {error}", + dialect.label(), + ); + assert!(error.to_string().contains("unsupported reader dispatch")); + } + + assert_eq!( + SyntaxTree::parse_with_dialect("#_value", Dialect::Scheme).unwrap_err(), + ParseError::UnsupportedReaderDispatch { + dispatch: "#".to_owned(), + position: 0, + } + ); +} + +#[test] +fn common_lisp_atom_like_dispatches_round_trip_losslessly() { + let cases = [ + "#:done", "#36rz", "#36RZ", "#37r10", "#16ra.", "#b1010", "#o17", "#d10", "#xFF", + ]; + + for input in cases { + let tree = SyntaxTree::parse_with_dialect(input, Dialect::CommonLisp) + .expect("valid atom dispatch"); + let root = tree.root_view(); + assert_eq!(root.children.len(), 1, "{input}"); + assert_eq!(root.children[0].span.slice(input), input); + assert_eq!(root.children[0].text.as_deref(), Some(input)); + } +} + +#[test] +fn standard_dialect_dispatch_forms_are_single_opaque_spans() { + let cases = [ + ( + Dialect::CommonLisp, + "#P\"/tmp/example.lisp\" tail", + "#P\"/tmp/example.lisp\"", + ), + ( + Dialect::CommonLisp, + "#S(point :x 1 :y 2) tail", + "#S(point :x 1 :y 2)", + ), + (Dialect::CommonLisp, "#A(1 2) tail", "#A(1 2)"), + ( + Dialect::CommonLisp, + "#2a((1 2) (3 4)) tail", + "#2a((1 2) (3 4))", + ), + ( + Dialect::CommonLisp, + "#1=(node . #1#) tail", + "#1=(node . #1#)", + ), + (Dialect::CommonLisp, "#1# tail", "#1#"), + (Dialect::Scheme, "#1=(node . #1#) tail", "#1=(node . #1#)"), + (Dialect::Scheme, "#1# tail", "#1#"), + (Dialect::Scheme, "#u8(1 2 3) tail", "#u8(1 2 3)"), + (Dialect::Clojure, r##"#"foo.*" tail"##, r##"#"foo.*""##), + ( + Dialect::Clojure, + r#"#:person{:first "Ada"} tail"#, + r#"#:person{:first "Ada"}"#, + ), + ( + Dialect::Clojure, + r#"#inst "1985-04-12T23:20:50.52-00:00" tail"#, + r#"#inst "1985-04-12T23:20:50.52-00:00""#, + ), + (Dialect::Clojure, "#+/foo 1 tail", "#+/foo 1"), + ]; + + for (dialect, input, expected_span) in cases { + let tree = SyntaxTree::parse_with_dialect(input, dialect).expect("valid dispatch form"); + let root = tree.root_view(); + assert_eq!(root.children.len(), 2, "{}", dialect.label()); + assert_eq!( + root.children[0].span.slice(input), + expected_span, + "{}", + dialect.label() + ); + assert_eq!(root.children[1].text.as_deref(), Some("tail")); + } +} + +#[test] +fn standard_dispatch_forms_require_their_payload_datum() { + let cases = [ + (Dialect::CommonLisp, "#P"), + (Dialect::CommonLisp, "#S"), + (Dialect::CommonLisp, "#A"), + (Dialect::CommonLisp, "#2A"), + (Dialect::CommonLisp, "#1="), + (Dialect::Scheme, "#1="), + ]; + + for (dialect, input) in cases { + assert_eq!( + SyntaxTree::parse_with_dialect(input, dialect), + Err(ParseError::MissingReaderForm(0)), + "{}: {input}", + dialect.label() + ); + } +} + +#[test] +fn standard_dispatch_forms_are_consumed_inside_skipped_datums() { + let cases = [ + ( + Dialect::CommonLisp, + "#+feature #2A((1 2) (3 4)) tail", + &["#+feature #2A((1 2) (3 4))", "tail"] as &[&str], + ), + ( + Dialect::CommonLisp, + "#+feature #1=(node . #1#) tail", + &["#+feature #1=(node . #1#)", "tail"] as &[&str], + ), + ( + Dialect::CommonLisp, + "#+feature #:done tail", + &["#+feature #:done", "tail"] as &[&str], + ), + ( + Dialect::CommonLisp, + "#+feature #36rz tail", + &["#+feature #36rz", "tail"] as &[&str], + ), + (Dialect::Scheme, "#;#1=(node . #1#) tail", &["tail"]), + (Dialect::Scheme, "#;#1# tail", &["tail"]), + (Dialect::Clojure, "#_#+/foo 1 tail", &["tail"]), + ]; + + for (dialect, input, expected) in cases { + let tree = SyntaxTree::parse_with_dialect(input, dialect).expect("valid skipped form"); + let spans = tree + .root_view() + .children + .iter() + .map(|child| child.span.slice(input)) + .collect::>(); + assert_eq!(spans, expected, "{}", dialect.label()); + } +} + +#[test] +fn opaque_dialect_dispatch_forms_are_not_traversed_by_rename() { + let cases = [ + ( + Dialect::Clojure, + "#:foo{:key foo} foo", + "#:foo{:key foo} bar", + ), + ( + Dialect::CommonLisp, + "#S(node :value foo) foo", + "#S(node :value foo) bar", + ), + ( + Dialect::CommonLisp, + "#1=(foo . #1#) foo", + "#1=(foo . #1#) bar", + ), + (Dialect::Scheme, "#1=(foo . #1#) foo", "#1=(foo . #1#) bar"), + ]; + + for (dialect, input, expected) in cases { + let tree = SyntaxTree::parse_with_dialect(input, dialect).expect("valid reader form"); + assert_eq!( + tree.rename_symbol( + &SymbolName::new("foo").expect("source symbol"), + &SymbolName::new("bar").expect("target symbol"), + ), + expected, + "{}", + dialect.label() + ); + } +} + #[test] fn parses_reader_delimiters() { let tree = SyntaxTree::parse("(mapv inc [1 2 {:x 3}])").expect("valid"); @@ -107,6 +413,47 @@ fn parses_clojure_metadata_prefix_on_map_and_atom() { assert_eq!(metadata_keyword.text.as_deref(), Some("^:private")); } +#[test] +fn clojure_metadata_keeps_target_live_and_discard_skips_target() { + let input = "^:private (defn foo [] (foo)) (foo)"; + let tree = SyntaxTree::parse_with_dialect(input, Dialect::Clojure).expect("valid metadata"); + + let root = tree.root_view(); + assert_eq!(root.children.len(), 3); + assert_eq!( + root.children[0].reader_prefixes, + vec![ReaderPrefix::Metadata] + ); + assert_eq!(root.children[0].text.as_deref(), Some("^:private")); + assert_eq!(root.children[1].span.slice(input), "(defn foo [] (foo))"); + assert_eq!(root.children[2].span.slice(input), "(foo)"); + + let foo_paths = tree + .atom_occurrences() + .into_iter() + .filter(|occurrence| occurrence.text == "foo") + .map(|occurrence| occurrence.path.to_string()) + .collect::>(); + assert_eq!(foo_paths, vec!["1.1", "1.3.0", "2.0"]); + + let outline = tree.outline(|head| Dialect::Clojure.is_definition_head(head)); + assert_eq!(outline.len(), 2); + assert_eq!(outline[0].path.to_string(), "1"); + assert_eq!(outline[0].head.as_deref(), Some("defn")); + assert!(outline[0].definition_like); + + let skipped_input = "#_^:private (defn foo [] (foo)) tail"; + let skipped = SyntaxTree::parse_with_dialect(skipped_input, Dialect::Clojure) + .expect("valid discarded metadata"); + let spans = skipped + .root_view() + .children + .iter() + .map(|child| child.span.slice(skipped_input)) + .collect::>(); + assert_eq!(spans, vec!["tail"]); +} + #[test] fn parses_clojure_reader_conditionals_as_one_node() { let tree = @@ -432,6 +779,42 @@ fn parses_named_and_whitespace_character_literals() { assert_eq!(form.children[1].text.as_deref(), Some("#\\ ")); } +#[test] +fn parses_dialect_character_literals_with_closing_delimiters() { + let cases = [ + (Dialect::Scheme, "(#\\))", "#\\)"), + (Dialect::Clojure, "(\\))", "\\)"), + (Dialect::EmacsLisp, "(?\\))", "?\\)"), + ]; + + for (dialect, input, expected) in cases { + let tree = SyntaxTree::parse_with_dialect(input, dialect).expect("valid character literal"); + let form = &tree.root_view().children[0]; + assert_eq!(form.children.len(), 1, "{}", dialect.label()); + assert_eq!( + form.children[0].span.slice(input), + expected, + "{}", + dialect.label() + ); + } +} + +#[test] +fn discarded_forms_use_the_same_dialect_character_literal_scanner() { + let cases = [ + (Dialect::Scheme, "#;(#\\)) kept"), + (Dialect::Clojure, "#_(\\)) kept"), + ]; + + for (dialect, input) in cases { + let tree = SyntaxTree::parse_with_dialect(input, dialect).expect("valid discarded form"); + let root = tree.root_view(); + assert_eq!(root.children.len(), 1, "{}", dialect.label()); + assert_eq!(root.children[0].text.as_deref(), Some("kept")); + } +} + #[test] fn character_literal_does_not_break_rename() { let input = "(defun f () (write-char #\\[ out) (foo))"; diff --git a/src/domain/sexpr/tree.rs b/src/domain/sexpr/tree.rs index 8206a04d..5a95288d 100644 --- a/src/domain/sexpr/tree.rs +++ b/src/domain/sexpr/tree.rs @@ -3,6 +3,7 @@ use std::fmt; use anyhow::{Result, anyhow}; use crate::domain::common_lisp::common_lisp_symbol_reference_eq; +use crate::domain::dialect::Dialect; use super::parser::{ParseError, Parser}; use super::types::{ByteOffset, ByteSpan, Delimiter, ExpressionPath, NodeId, SymbolName}; @@ -53,6 +54,9 @@ pub(in crate::domain::sexpr) struct Node { pub(in crate::domain::sexpr) kind: NodeKind, pub(in crate::domain::sexpr) delimiter: Option, pub(in crate::domain::sexpr) reader_prefixes: Vec, + /// Exact source ranges for `reader_prefixes`, kept separately from their + /// normalized semantics so dialect-specific spellings round-trip. + pub(in crate::domain::sexpr) reader_prefix_spans: Vec, pub(in crate::domain::sexpr) parent: Option, pub(in crate::domain::sexpr) children: Vec, pub(in crate::domain::sexpr) span: ByteSpan, @@ -66,6 +70,9 @@ pub(in crate::domain::sexpr) struct Node { /// prefix's fixed source length — it must be recorded while parsing. /// Meaningless (`0`) for non-atom nodes. pub(in crate::domain::sexpr) symbol_offset: usize, + /// Reader forms that consume multiple datums are represented by one + /// verbatim atom node so their payload cannot become editable siblings. + pub(in crate::domain::sexpr) opaque_reader_form: bool, } #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -349,6 +356,18 @@ impl SyntaxTree { parser.parse() } + /// Parses source using the lexical and reader-macro rules of `dialect`. + /// + /// [`Self::parse`] intentionally retains the historical permissive reader; + /// callers that know the file dialect should use this entry point. + pub fn parse_with_dialect( + input: &str, + dialect: Dialect, + ) -> std::result::Result { + let mut parser = Parser::with_dialect(input, dialect); + parser.parse() + } + /// Returns the direct children of the virtual root document node. pub fn root_children(&self) -> &[NodeId] { &self.node(NodeId::ROOT).children @@ -410,10 +429,11 @@ impl SyntaxTree { let mut pending = self.node(NodeId::ROOT).children.clone(); while let Some(node_id) = pending.pop() { let node = self.node(node_id); - if node - .reader_prefixes - .iter() - .any(|prefix| prefix.is_opaque_reader_form()) + if node.opaque_reader_form + || node + .reader_prefixes + .iter() + .any(|prefix| prefix.is_opaque_reader_form()) { continue; } @@ -451,10 +471,11 @@ impl SyntaxTree { while let Some((node_id, parent_id, index)) = pending.pop() { parent_steps[node_id.get()] = Some((parent_id, index)); let node = self.node(node_id); - if node - .reader_prefixes - .iter() - .any(|prefix| prefix.is_opaque_reader_form()) + if node.opaque_reader_form + || node + .reader_prefixes + .iter() + .any(|prefix| prefix.is_opaque_reader_form()) { continue; } @@ -646,10 +667,11 @@ impl SyntaxTree { pending.push(Frame::Leave); let node = self.node(node_id); - if node - .reader_prefixes - .iter() - .any(|prefix| prefix.is_opaque_reader_form()) + if node.opaque_reader_form + || node + .reader_prefixes + .iter() + .any(|prefix| prefix.is_opaque_reader_form()) { continue; } @@ -686,7 +708,10 @@ impl SyntaxTree { fn atom_text(&self, node_id: NodeId) -> Option<&str> { let node = self.node(node_id); - if node.kind != NodeKind::Atom || !node.reader_prefixes.is_empty() { + if node.kind != NodeKind::Atom + || node.opaque_reader_form + || !node.reader_prefixes.is_empty() + { return None; } Some(node.span.slice(&self.source)) diff --git a/src/domain/sort_definitions.rs b/src/domain/sort_definitions.rs index 6f6e9f33..e4032776 100644 --- a/src/domain/sort_definitions.rs +++ b/src/domain/sort_definitions.rs @@ -28,7 +28,11 @@ pub use types::{ const DEFAULT_ENTRY_SEPARATOR: &str = "\n\n"; pub fn plan_sort_definitions(request: SortDefinitionsRequest<'_>) -> Result { - let tree = SyntaxTree::parse(request.input)?; + if request.dialect == crate::domain::dialect::Dialect::Unknown { + anyhow::bail!("sort-definitions does not support the unknown dialect"); + } + + let tree = SyntaxTree::parse_with_dialect(request.input, request.dialect)?; reject_common_lisp_reader_conditionals(&tree, request.dialect)?; let blocks = collect_sortable_blocks(request.input, &tree, request.dialect)?; let mut replacements = Vec::new(); @@ -76,7 +80,7 @@ pub fn plan_sort_definitions(request: SortDefinitionsRequest<'_>) -> Result SortDefinitionsRequest<'_> { + request_for_dialect(input, Dialect::CommonLisp, strategy) +} + +fn request_for_dialect( + input: &str, + dialect: Dialect, + strategy: SortDefinitionsStrategy, +) -> SortDefinitionsRequest<'_> { SortDefinitionsRequest { file: PathBuf::from("core.lisp"), input, - dialect: Dialect::CommonLisp, + dialect, strategy, write: false, } @@ -37,6 +45,81 @@ fn sorts_contiguous_definitions_by_name() { assert!(SyntaxTree::parse(&plan.rewritten).is_ok()); } +#[test] +fn sorts_definitions_for_every_known_dialect_and_reparses_with_the_same_dialect() { + let fixtures = [ + ( + Dialect::CommonLisp, + "(in-package #:demo)\n\n(defun zeta () :zeta)\n(defun alpha () :alpha)\n", + "(defun alpha", + "(defun zeta", + ), + ( + Dialect::EmacsLisp, + "(defun zeta () 'zeta)\n(defun alpha () 'alpha)\n", + "(defun alpha", + "(defun zeta", + ), + ( + Dialect::Scheme, + "(define zeta (lambda () 'zeta))\n(define alpha (lambda () 'alpha))\n", + "(define alpha", + "(define zeta", + ), + ( + Dialect::Clojure, + "(defn zeta [] #?(:clj :zeta :cljs :zeta))\n(defn alpha [] :alpha)\n", + "(defn alpha", + "(defn zeta", + ), + ( + Dialect::Janet, + "(defn zeta [] :zeta)\n(defn alpha [] :alpha)\n", + "(defn alpha", + "(defn zeta", + ), + ( + Dialect::Fennel, + "(fn zeta [] :zeta)\n(fn alpha [] :alpha)\n", + "(fn alpha", + "(fn zeta", + ), + ]; + + for (dialect, input, alpha_marker, zeta_marker) in fixtures { + let plan = plan_sort_definitions(request_for_dialect( + input, + dialect, + SortDefinitionsStrategy::Name, + )) + .unwrap_or_else(|error| panic!("{dialect:?} sort failed: {error:#}")); + + assert!(plan.changed, "{dialect:?} fixture was not sorted"); + assert!( + plan.rewritten.find(alpha_marker).unwrap() < plan.rewritten.find(zeta_marker).unwrap(), + "{dialect:?} output was not sorted: {}", + plan.rewritten + ); + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect) + .unwrap_or_else(|error| panic!("{dialect:?} output did not reparse: {error:#}")); + } +} + +#[test] +fn rejects_unknown_dialect_before_parsing_malformed_input() { + let error = plan_sort_definitions(request_for_dialect( + "(", + Dialect::Unknown, + SortDefinitionsStrategy::Name, + )) + .unwrap_err(); + + assert!( + error.to_string().contains("unknown dialect"), + "unexpected gate error: {error:#}" + ); +} + #[test] fn does_not_cross_non_definition_barriers() { let input = "(defun zeta () :z)\n\ diff --git a/src/domain/split_file.rs b/src/domain/split_file.rs index 7eb478f7..cb7ec31c 100644 --- a/src/domain/split_file.rs +++ b/src/domain/split_file.rs @@ -23,14 +23,15 @@ pub fn plan_split_file(request: SplitFileRequest<'_>) -> Result { anyhow::bail!("split-file requires at least one --path, --name, or --kind selector"); } - let from_tree = SyntaxTree::parse(request.from_input) + let from_tree = SyntaxTree::parse_with_dialect(request.from_input, request.from_dialect) .with_context(|| format!("failed to parse {}", request.from_file.display()))?; - let to_tree = SyntaxTree::parse(request.to_input).with_context(|| { - format!( - "destination file is not a valid S-expression document: {}", - request.to_file.display() - ) - })?; + let to_tree = SyntaxTree::parse_with_dialect(request.to_input, request.to_dialect) + .with_context(|| { + format!( + "destination file is not a valid S-expression document: {}", + request.to_file.display() + ) + })?; let mut seen_paths = std::collections::BTreeSet::new(); let mut selected_paths = std::collections::BTreeMap::new(); @@ -125,15 +126,24 @@ pub fn plan_split_file(request: SplitFileRequest<'_>) -> Result { items.iter().map(|item| item.removal_span), )?; - let mut running_package = package_context_before_top_level( - &to_tree, - request.to_dialect, - to_tree.root_children().len(), - )?; + let destination_is_common_lisp = + request.to_dialect == crate::domain::dialect::Dialect::CommonLisp; + let mut running_package = if destination_is_common_lisp { + package_context_before_top_level( + &to_tree, + request.to_dialect, + to_tree.root_children().len(), + )? + } else { + None + }; let definition_texts = items .iter() .map(|item| match &item.definition.package { - Some(package) if running_package.as_deref() != Some(package.as_str()) => { + Some(package) + if destination_is_common_lisp + && running_package.as_deref() != Some(package.as_str()) => + { running_package = Some(package.clone()); format!("(in-package {package})\n\n{}", item.definition_text) } @@ -152,13 +162,13 @@ pub fn plan_split_file(request: SplitFileRequest<'_>) -> Result { } let to_rewritten = append_top_level_definitions(request.to_input, &definition_texts); - SyntaxTree::parse(&from_rewritten).with_context(|| { + SyntaxTree::parse_with_dialect(&from_rewritten, request.from_dialect).with_context(|| { format!( "source file would become invalid after splitting definitions: {}", request.from_file.display() ) })?; - SyntaxTree::parse(&to_rewritten).with_context(|| { + SyntaxTree::parse_with_dialect(&to_rewritten, request.to_dialect).with_context(|| { format!( "destination file would become invalid after receiving definitions: {}", request.to_file.display() diff --git a/src/domain/split_file/tests.rs b/src/domain/split_file/tests.rs index 9708cbde..fe14e70e 100644 --- a/src/domain/split_file/tests.rs +++ b/src/domain/split_file/tests.rs @@ -282,3 +282,71 @@ fn plan_split_file_preserves_unselected_common_lisp_vector_literals() { assert!(SyntaxTree::parse(&plan.from_rewritten).is_ok()); assert!(SyntaxTree::parse(&plan.to_rewritten).is_ok()); } + +#[test] +fn plan_split_file_uses_each_file_dialect_for_character_literal_collisions() { + for (case_name, dialect, character_literal) in [ + ("Common Lisp", Dialect::CommonLisp, "#\\)"), + ("Emacs Lisp", Dialect::EmacsLisp, "?\\)"), + ] { + let from_input = + format!("(defun keep () {character_literal})\n\n(defun moved () :moved)\n"); + let to_input = format!("(defun existing () {character_literal})\n"); + let mut request = split_request(&from_input, &to_input); + request.from_dialect = dialect; + request.to_dialect = dialect; + request.paths = vec![Path::from_indexes(vec![1])]; + + let plan = plan_split_file(request) + .unwrap_or_else(|error| panic!("{case_name} split should succeed: {error:#}")); + + assert!(plan.from_rewritten.contains(character_literal)); + assert!(plan.to_rewritten.contains(character_literal)); + assert!(plan.to_rewritten.contains("(defun moved () :moved)")); + SyntaxTree::parse_with_dialect(&plan.from_rewritten, dialect) + .unwrap_or_else(|error| panic!("{case_name} source output should parse: {error:#}")); + SyntaxTree::parse_with_dialect(&plan.to_rewritten, dialect).unwrap_or_else(|error| { + panic!("{case_name} destination output should parse: {error:#}") + }); + } +} + +#[test] +fn plan_split_file_preserves_unknown_generic_syntax_compatibility() { + let mut request = split_request( + "(defn keep [] :kept)\n\n(defn moved [] :moved)\n", + "(defn existing [] :existing)\n", + ); + request.from_dialect = Dialect::Unknown; + request.to_dialect = Dialect::Unknown; + request.paths = vec![Path::from_indexes(vec![1])]; + + let plan = plan_split_file(request).expect("Unknown dialect should retain generic parsing"); + + assert!(plan.from_rewritten.contains("(defn keep [] :kept)")); + assert!(plan.to_rewritten.contains("(defn moved [] :moved)")); + SyntaxTree::parse_with_dialect(&plan.from_rewritten, Dialect::Unknown) + .expect("Unknown source output should parse"); + SyntaxTree::parse_with_dialect(&plan.to_rewritten, Dialect::Unknown) + .expect("Unknown destination output should parse"); +} + +#[test] +fn plan_split_file_does_not_inject_common_lisp_package_into_non_common_lisp_destination() { + let mut request = split_request( + "(in-package #:demo)\n\n(defun moved () :moved)\n", + "(setq destination-ready t)\n", + ); + request.from_dialect = Dialect::CommonLisp; + request.to_dialect = Dialect::EmacsLisp; + request.paths = vec![Path::from_indexes(vec![1])]; + + let plan = plan_split_file(request).expect("Common Lisp to Emacs Lisp split should succeed"); + + assert!(!plan.to_rewritten.contains("(in-package")); + assert!(plan.to_rewritten.contains("(defun moved () :moved)")); + SyntaxTree::parse_with_dialect(&plan.from_rewritten, Dialect::CommonLisp) + .expect("Common Lisp source output should parse"); + SyntaxTree::parse_with_dialect(&plan.to_rewritten, Dialect::EmacsLisp) + .expect("Emacs Lisp destination output should parse"); +} diff --git a/src/domain/thread_expression.rs b/src/domain/thread_expression.rs index 6bc58146..74b48ec9 100644 --- a/src/domain/thread_expression.rs +++ b/src/domain/thread_expression.rs @@ -11,6 +11,7 @@ mod types; pub use types::{ThreadExpressionPlan, ThreadExpressionRequest, ThreadExpressionStep, ThreadStyle}; use crate::domain::common_lisp::common_lisp_symbol_reference_eq; +use crate::domain::dialect::Dialect; use crate::domain::mutation_safety::reject_common_lisp_reader_conditionals; use crate::domain::sexpr::SyntaxTree; use anyhow::{Context, Result}; @@ -21,11 +22,30 @@ use syntax::list_head; pub fn plan_thread_expression( request: ThreadExpressionRequest<'_>, ) -> Result { + match request.dialect { + Dialect::CommonLisp + | Dialect::EmacsLisp + | Dialect::Scheme + | Dialect::Clojure + | Dialect::Janet + | Dialect::Fennel => {} + Dialect::Unknown => { + anyhow::bail!("thread-expression does not support dialect unknown"); + } + } + reject_common_lisp_reader_conditionals(request.tree, request.dialect)?; - if list_head(&request.target) - .is_some_and(|head| common_lisp_symbol_reference_eq(head, request.operator.as_str())) - { + let already_threaded = list_head(&request.target).is_some_and(|head| match request.dialect { + Dialect::CommonLisp => common_lisp_symbol_reference_eq(head, request.operator.as_str()), + Dialect::EmacsLisp + | Dialect::Scheme + | Dialect::Clojure + | Dialect::Janet + | Dialect::Fennel => head == request.operator.as_str(), + Dialect::Unknown => false, + }); + if already_threaded { anyhow::bail!( "thread-expression selection is already threaded with {}", request.operator @@ -46,9 +66,11 @@ pub fn plan_thread_expression( anyhow::bail!("thread-expression target did not produce any pipeline steps"); } let replacement = thread_expression_replacement(&request.operator, &parts.base, &parts.steps); - SyntaxTree::parse(&replacement).context("thread-expression replacement does not parse")?; + SyntaxTree::parse_with_dialect(&replacement, request.dialect) + .context("thread-expression replacement does not parse")?; let rewritten = replace_span(request.input, request.target.span, &replacement); - SyntaxTree::parse(&rewritten).context("thread-expression rewritten output does not parse")?; + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) + .context("thread-expression rewritten output does not parse")?; let changed = rewritten != request.input; Ok(ThreadExpressionPlan { diff --git a/src/domain/thread_expression/tests.rs b/src/domain/thread_expression/tests.rs index 05450765..2bd31000 100644 --- a/src/domain/thread_expression/tests.rs +++ b/src/domain/thread_expression/tests.rs @@ -3,8 +3,8 @@ use crate::domain::dialect::Dialect; use crate::domain::sexpr::{ExpressionView, Path, SymbolName}; use proptest::prelude::*; -fn parsed(input: &str) -> (SyntaxTree, ExpressionView) { - let tree = SyntaxTree::parse(input).expect("parse fixture"); +fn parsed(input: &str, dialect: Dialect) -> (SyntaxTree, ExpressionView) { + let tree = SyntaxTree::parse_with_dialect(input, dialect).expect("parse fixture"); let target = tree .select_path(&"0".parse::().expect("path")) .expect("select fixture") @@ -24,7 +24,7 @@ fn symbol_strategy() -> impl Strategy { #[test] fn plans_thread_first_pipeline_from_nested_calls() { let input = "(render (normalize value mode))"; - let (tree, target) = parsed(input); + let (tree, target) = parsed(input, Dialect::Clojure); let plan = plan_thread_expression(ThreadExpressionRequest { input, tree: &tree, @@ -45,7 +45,7 @@ fn plans_thread_first_pipeline_from_nested_calls() { #[test] fn plans_thread_last_pipeline_from_nested_calls() { let input = "(map f (filter pred rows))"; - let (tree, target) = parsed(input); + let (tree, target) = parsed(input, Dialect::Clojure); let plan = plan_thread_expression(ThreadExpressionRequest { input, tree: &tree, @@ -66,7 +66,7 @@ fn plans_thread_last_pipeline_from_nested_calls() { #[test] fn rejects_package_qualified_already_threaded_expression() { let input = "(cl:-> x f)"; - let (tree, target) = parsed(input); + let (tree, target) = parsed(input, Dialect::CommonLisp); let err = plan_thread_expression(ThreadExpressionRequest { input, tree: &tree, @@ -81,10 +81,128 @@ fn rejects_package_qualified_already_threaded_expression() { assert!(err.to_string().contains("already threaded")); } +#[test] +fn supports_known_dialects_and_rejects_unknown() { + let cases = [ + (Dialect::CommonLisp, true), + (Dialect::EmacsLisp, true), + (Dialect::Scheme, true), + (Dialect::Clojure, true), + (Dialect::Janet, true), + (Dialect::Fennel, true), + (Dialect::Unknown, false), + ]; + + for (dialect, supported) in cases { + let input = "(render (normalize value mode))"; + let (tree, target) = parsed(input, dialect); + let result = plan_thread_expression(ThreadExpressionRequest { + input, + tree: &tree, + dialect, + path: Some("0".parse().expect("path")), + target, + style: ThreadStyle::First, + operator: SymbolName::new("->").expect("symbol"), + }); + + if supported { + let plan = result + .unwrap_or_else(|error| panic!("{} should be supported: {error}", dialect.label())); + assert!( + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect).is_ok(), + "{} output should parse in the same dialect", + dialect.label() + ); + } else { + let error = result.expect_err("unknown dialect should be rejected"); + assert!( + error + .to_string() + .contains("does not support dialect unknown") + ); + } + } +} + +#[test] +fn rejects_unknown_before_reader_and_target_validation() { + let (tree, target) = parsed("atom", Dialect::CommonLisp); + let error = plan_thread_expression(ThreadExpressionRequest { + input: ")", + tree: &tree, + dialect: Dialect::Unknown, + path: Some("0".parse().expect("path")), + target, + style: ThreadStyle::First, + operator: SymbolName::new("->").expect("symbol"), + }) + .expect_err("unknown dialect should fail before malformed input and target checks"); + + assert!( + error + .to_string() + .contains("does not support dialect unknown") + ); +} + +#[test] +fn clojure_namespace_does_not_match_an_unqualified_thread_operator() { + let input = "(ns/-> value f)"; + let (tree, target) = parsed(input, Dialect::Clojure); + let plan = plan_thread_expression(ThreadExpressionRequest { + input, + tree: &tree, + dialect: Dialect::Clojure, + path: Some("0".parse().expect("path")), + target, + style: ThreadStyle::First, + operator: SymbolName::new("->").expect("symbol"), + }) + .expect("namespace-qualified Clojure symbol is not the unqualified operator"); + + assert_eq!(plan.replacement, "(-> value (ns/-> f))"); +} + +#[test] +fn preserves_dialect_reader_atoms_and_forms() { + let cases = [ + ( + Dialect::CommonLisp, + "(render (normalize #\\) mode))", + "#\\)", + ), + (Dialect::EmacsLisp, "(render (normalize ?\\) mode))", "?\\)"), + ( + Dialect::Clojure, + "(render (normalize #foo/bar {:x 1} mode))", + "#foo/bar {:x 1}", + ), + ]; + + for (dialect, input, preserved) in cases { + let (tree, target) = parsed(input, dialect); + let plan = plan_thread_expression(ThreadExpressionRequest { + input, + tree: &tree, + dialect, + path: Some("0".parse().expect("path")), + target, + style: ThreadStyle::First, + operator: SymbolName::new("->").expect("symbol"), + }) + .unwrap_or_else(|error| panic!("{} reader case failed: {error}", dialect.label())); + + assert!(plan.rewritten.contains(preserved)); + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect) + .unwrap_or_else(|error| panic!("{} output did not reparse: {error}", dialect.label())); + } +} + #[test] fn rejects_nested_calls_with_an_interior_comment() { let input = "(render\n ;; note\n (normalize value mode))"; - let (tree, target) = parsed(input); + let (tree, target) = parsed(input, Dialect::Clojure); let err = plan_thread_expression(ThreadExpressionRequest { input, tree: &tree, @@ -114,7 +232,7 @@ proptest! { prop_assume!(inner != outer); let input = format!("({outer} ({inner} {base} {arg}))"); - let (tree, target) = parsed(&input); + let (tree, target) = parsed(&input, Dialect::Clojure); let plan = plan_thread_expression(ThreadExpressionRequest { input: &input, tree: &tree, @@ -130,7 +248,7 @@ proptest! { plan.replacement, format!("(-> {base} ({inner} {arg}) {outer})") ); - prop_assert!(SyntaxTree::parse(&plan.rewritten).is_ok()); + prop_assert!(SyntaxTree::parse_with_dialect(&plan.rewritten, Dialect::Clojure).is_ok()); prop_assert!(plan.changed); } } diff --git a/src/domain/unthread_expression.rs b/src/domain/unthread_expression.rs index d90ae6fd..fa83c3ca 100644 --- a/src/domain/unthread_expression.rs +++ b/src/domain/unthread_expression.rs @@ -12,6 +12,7 @@ pub use types::{ UnthreadExpressionPlan, UnthreadExpressionRequest, UnthreadExpressionStep, UnthreadStyle, }; +use crate::domain::dialect::Dialect; use crate::domain::mutation_safety::reject_common_lisp_reader_conditionals; use crate::domain::sexpr::{Delimiter, ExpressionKind, SymbolName, SyntaxTree}; use anyhow::{Context, Result}; @@ -22,6 +23,18 @@ use syntax::{atom_child, expression_source}; pub fn plan_unthread_expression( request: UnthreadExpressionRequest<'_>, ) -> Result { + match request.dialect { + Dialect::CommonLisp + | Dialect::EmacsLisp + | Dialect::Scheme + | Dialect::Clojure + | Dialect::Janet + | Dialect::Fennel => {} + Dialect::Unknown => { + anyhow::bail!("unthread-expression does not support dialect unknown"); + } + } + reject_common_lisp_reader_conditionals(request.tree, request.dialect)?; if request.target.kind != ExpressionKind::List @@ -87,9 +100,11 @@ pub fn plan_unthread_expression( .map(|view| pipeline_step(request.input, view)) .collect::>>()?; let (replacement, steps) = unthread_replacement(style, &base, pipeline_steps); - SyntaxTree::parse(&replacement).context("unthread-expression replacement does not parse")?; + SyntaxTree::parse_with_dialect(&replacement, request.dialect) + .context("unthread-expression replacement does not parse")?; let rewritten = replace_span(request.input, request.target.span, &replacement); - SyntaxTree::parse(&rewritten).context("unthread-expression rewritten output does not parse")?; + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) + .context("unthread-expression rewritten output does not parse")?; let changed = rewritten != request.input; Ok(UnthreadExpressionPlan { diff --git a/src/domain/unthread_expression/tests.rs b/src/domain/unthread_expression/tests.rs index 71b92e77..4600964f 100644 --- a/src/domain/unthread_expression/tests.rs +++ b/src/domain/unthread_expression/tests.rs @@ -3,8 +3,8 @@ use crate::domain::dialect::Dialect; use crate::domain::sexpr::{ExpressionView, Path, SymbolName}; use proptest::prelude::*; -fn parsed(input: &str) -> (SyntaxTree, ExpressionView) { - let tree = SyntaxTree::parse(input).expect("parse fixture"); +fn parsed(input: &str, dialect: Dialect) -> (SyntaxTree, ExpressionView) { + let tree = SyntaxTree::parse_with_dialect(input, dialect).expect("parse fixture"); let target = tree .select_path(&"0".parse::().expect("path")) .expect("select fixture") @@ -24,7 +24,7 @@ fn symbol_strategy() -> impl Strategy { #[test] fn plans_unthread_first_pipeline_into_nested_call() { let input = "(-> value (normalize mode) render)"; - let (tree, target) = parsed(input); + let (tree, target) = parsed(input, Dialect::Clojure); let plan = plan_unthread_expression(UnthreadExpressionRequest { input, tree: &tree, @@ -46,7 +46,7 @@ fn plans_unthread_first_pipeline_into_nested_call() { #[test] fn plans_unthread_last_pipeline_into_nested_call() { let input = "(->> rows (filter pred) (map f))"; - let (tree, target) = parsed(input); + let (tree, target) = parsed(input, Dialect::Clojure); let plan = plan_unthread_expression(UnthreadExpressionRequest { input, tree: &tree, @@ -68,7 +68,7 @@ fn plans_unthread_last_pipeline_into_nested_call() { #[test] fn rejects_style_alone_on_an_unrecognized_operator() { let input = "(+ a b)"; - let (tree, target) = parsed(input); + let (tree, target) = parsed(input, Dialect::CommonLisp); let err = plan_unthread_expression(UnthreadExpressionRequest { input, tree: &tree, @@ -86,7 +86,7 @@ fn rejects_style_alone_on_an_unrecognized_operator() { #[test] fn accepts_style_with_explicit_operator_confirming_a_custom_pipeline() { let input = "(my-pipe value (normalize mode) render)"; - let (tree, target) = parsed(input); + let (tree, target) = parsed(input, Dialect::CommonLisp); let plan = plan_unthread_expression(UnthreadExpressionRequest { input, tree: &tree, @@ -101,10 +101,114 @@ fn accepts_style_with_explicit_operator_confirming_a_custom_pipeline() { assert_eq!(plan.replacement, "(render (normalize value mode))"); } +#[test] +fn supports_known_dialects_and_rejects_unknown() { + let cases = [ + (Dialect::CommonLisp, true), + (Dialect::EmacsLisp, true), + (Dialect::Scheme, true), + (Dialect::Clojure, true), + (Dialect::Janet, true), + (Dialect::Fennel, true), + (Dialect::Unknown, false), + ]; + + for (dialect, supported) in cases { + let input = "(-> value (normalize mode) render)"; + let (tree, target) = parsed(input, dialect); + let result = plan_unthread_expression(UnthreadExpressionRequest { + input, + tree: &tree, + dialect, + path: Some("0".parse().expect("path")), + target, + style: None, + operator: None, + }); + + if supported { + let plan = result + .unwrap_or_else(|error| panic!("{} should be supported: {error}", dialect.label())); + assert!( + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect).is_ok(), + "{} output should parse in the same dialect", + dialect.label() + ); + } else { + let error = result.expect_err("unknown dialect should be rejected"); + assert!( + error + .to_string() + .contains("does not support dialect unknown") + ); + } + } +} + +#[test] +fn rejects_unknown_before_reader_and_target_validation() { + let (tree, target) = parsed("atom", Dialect::CommonLisp); + let error = plan_unthread_expression(UnthreadExpressionRequest { + input: ")", + tree: &tree, + dialect: Dialect::Unknown, + path: Some("0".parse().expect("path")), + target, + style: None, + operator: None, + }) + .expect_err("unknown dialect should fail before malformed input and target checks"); + + assert!( + error + .to_string() + .contains("does not support dialect unknown") + ); +} + +#[test] +fn preserves_dialect_reader_atoms_and_forms() { + let cases = [ + ( + Dialect::CommonLisp, + "(-> #\\) (normalize mode) render)", + "#\\)", + ), + ( + Dialect::EmacsLisp, + "(-> ?\\) (normalize mode) render)", + "?\\)", + ), + ( + Dialect::Clojure, + "(-> #foo/bar {:x 1} (normalize mode) render)", + "#foo/bar {:x 1}", + ), + ]; + + for (dialect, input, preserved) in cases { + let (tree, target) = parsed(input, dialect); + let plan = plan_unthread_expression(UnthreadExpressionRequest { + input, + tree: &tree, + dialect, + path: Some("0".parse().expect("path")), + target, + style: None, + operator: None, + }) + .unwrap_or_else(|error| panic!("{} reader case failed: {error}", dialect.label())); + + assert!(plan.rewritten.contains(preserved)); + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect) + .unwrap_or_else(|error| panic!("{} output did not reparse: {error}", dialect.label())); + } +} + #[test] fn rejects_pipeline_with_an_interior_comment() { let input = "(-> value\n ;; note\n (normalize mode)\n render)"; - let (tree, target) = parsed(input); + let (tree, target) = parsed(input, Dialect::Clojure); let err = plan_unthread_expression(UnthreadExpressionRequest { input, tree: &tree, @@ -134,7 +238,7 @@ proptest! { prop_assume!(inner != outer); let input = format!("(-> {base} ({inner} {arg}) {outer})"); - let (tree, target) = parsed(&input); + let (tree, target) = parsed(&input, Dialect::Clojure); let plan = plan_unthread_expression(UnthreadExpressionRequest { input: &input, tree: &tree, @@ -150,7 +254,7 @@ proptest! { plan.replacement, format!("({outer} ({inner} {base} {arg}))") ); - prop_assert!(SyntaxTree::parse(&plan.rewritten).is_ok()); + prop_assert!(SyntaxTree::parse_with_dialect(&plan.rewritten, Dialect::Clojure).is_ok()); prop_assert!(plan.changed); } } diff --git a/src/domain/unwrap_call.rs b/src/domain/unwrap_call.rs index 95c9d875..f782d952 100644 --- a/src/domain/unwrap_call.rs +++ b/src/domain/unwrap_call.rs @@ -2,6 +2,7 @@ use anyhow::{Context, Result}; +use crate::domain::common_lisp::common_lisp_symbol_reference_eq; use crate::domain::dialect::Dialect; use crate::domain::sexpr::{ ByteSpan, Delimiter, ExpressionKind, ExpressionView, Path, SymbolName, SyntaxTree, @@ -31,8 +32,23 @@ pub(crate) struct Plan { pub changed: bool, } +pub(crate) fn validate_dialect(dialect: Dialect) -> Result<()> { + match dialect { + Dialect::CommonLisp + | Dialect::EmacsLisp + | Dialect::Scheme + | Dialect::Clojure + | Dialect::Janet + | Dialect::Fennel => Ok(()), + Dialect::Unknown => anyhow::bail!("unwrap-call requires a known dialect"), + } +} + pub(crate) fn plan(request: Request<'_>) -> Result { - SyntaxTree::parse(request.input).context("unwrap-call input does not parse")?; + validate_dialect(request.dialect)?; + + SyntaxTree::parse_with_dialect(request.input, request.dialect) + .context("unwrap-call input does not parse")?; if request.target.kind != ExpressionKind::List || request.target.delimiter != Some(Delimiter::Paren) @@ -49,7 +65,18 @@ pub(crate) fn plan(request: Request<'_>) -> Result { let function = SymbolName::new(head)?; if let Some(expected) = &request.expected_function { - if expected.as_str() != function.as_str() { + let matches = match request.dialect { + Dialect::CommonLisp => { + common_lisp_symbol_reference_eq(expected.as_str(), function.as_str()) + } + Dialect::EmacsLisp + | Dialect::Scheme + | Dialect::Clojure + | Dialect::Janet + | Dialect::Fennel => expected.as_str() == function.as_str(), + Dialect::Unknown => unreachable!("dialect was validated before parsing"), + }; + if !matches { anyhow::bail!( "unwrap-call expected function {}, found {}", expected.as_str(), @@ -70,11 +97,13 @@ pub(crate) fn plan(request: Request<'_>) -> Result { ) })?; let replacement = argument.span.slice(request.input).to_owned(); - SyntaxTree::parse(&replacement).context("unwrap-call replacement is not parseable")?; + SyntaxTree::parse_with_dialect(&replacement, request.dialect) + .context("unwrap-call replacement is not parseable")?; let mut rewritten = request.input.to_owned(); rewritten.replace_range(request.target.span.as_range(), &replacement); - SyntaxTree::parse(&rewritten).context("unwrap-call rewritten output is not parseable")?; + SyntaxTree::parse_with_dialect(&rewritten, request.dialect) + .context("unwrap-call rewritten output is not parseable")?; Ok(Plan { dialect: request.dialect, @@ -89,3 +118,136 @@ pub(crate) fn plan(request: Request<'_>) -> Result { rewritten, }) } + +#[cfg(test)] +mod tests { + use super::*; + + fn target(input: &str, dialect: Dialect, path: &str) -> ExpressionView { + let tree = SyntaxTree::parse_with_dialect(input, dialect) + .unwrap_or_else(|error| panic!("{}: {error}", dialect.label())); + tree.select_path(&path.parse::().expect("valid test path")) + .expect("test target exists") + .view() + } + + fn request<'a>( + input: &'a str, + dialect: Dialect, + path: &str, + expected_function: Option<&str>, + argument_index: usize, + ) -> Request<'a> { + Request { + input, + dialect, + path: Some(path.parse::().expect("valid test path")), + target: target(input, dialect, path), + expected_function: expected_function + .map(SymbolName::new) + .transpose() + .expect("valid expected function"), + argument_index, + } + } + + #[test] + fn supports_all_known_dialects_with_their_reader_forms() { + let cases = [ + (Dialect::CommonLisp, r"(wrap #\))", r"#\)"), + (Dialect::EmacsLisp, r"(wrap ?\))", r"?\)"), + (Dialect::Scheme, "(wrap #u8(1 2))", "#u8(1 2)"), + ( + Dialect::Clojure, + r#"(wrap #inst "2020-01-01")"#, + r#"#inst "2020-01-01""#, + ), + (Dialect::Janet, "(wrap ;value)", ";value"), + (Dialect::Fennel, "(wrap #(value))", "#(value)"), + ]; + + for (dialect, input, expected_replacement) in cases { + let plan = plan(request(input, dialect, "0", Some("wrap"), 0)) + .unwrap_or_else(|error| panic!("{}: {error}", dialect.label())); + + assert_eq!(plan.dialect, dialect); + assert_eq!(plan.function.as_str(), "wrap"); + assert_eq!(plan.argument_index, 0); + assert_eq!(plan.call_argument_count, 1); + assert_eq!(plan.replacement, expected_replacement); + assert_eq!(plan.rewritten, expected_replacement); + assert!(plan.changed); + SyntaxTree::parse_with_dialect(&plan.rewritten, dialect) + .unwrap_or_else(|error| panic!("{} output: {error}", dialect.label())); + } + } + + #[test] + fn unknown_dialect_fails_before_malformed_input_is_parsed() { + let error = plan(Request { + input: ")", + dialect: Dialect::Unknown, + path: None, + target: target("(wrap value)", Dialect::CommonLisp, "0"), + expected_function: None, + argument_index: 0, + }) + .expect_err("unknown dialect must fail closed"); + + assert_eq!(error.to_string(), "unwrap-call requires a known dialect"); + } + + #[test] + fn common_lisp_expected_function_ignores_case_and_package_qualifiers() { + let input = r"(PKG:WRAP #\))"; + let plan = plan(request( + input, + Dialect::CommonLisp, + "0", + Some("other-package:wrap"), + 0, + )) + .expect("Common Lisp symbol references should match"); + + assert_eq!(plan.function.as_str(), "PKG:WRAP"); + assert_eq!(plan.replacement, r"#\)"); + assert_eq!(plan.rewritten, r"#\)"); + } + + #[test] + fn non_common_lisp_expected_function_comparison_is_exact() { + for dialect in [ + Dialect::EmacsLisp, + Dialect::Scheme, + Dialect::Clojure, + Dialect::Janet, + Dialect::Fennel, + ] { + for input in ["(WRAP value)", "(pkg:wrap value)"] { + let error = plan(request(input, dialect, "0", Some("wrap"), 0)) + .expect_err("non-Common-Lisp head comparison must be exact"); + assert!( + error + .to_string() + .starts_with("unwrap-call expected function wrap, found "), + "{}: {error}", + dialect.label() + ); + } + } + } + + #[test] + fn preserves_selected_call_and_argument_spans() { + let input = "(outer (wrap first second) tail)"; + let plan = plan(request(input, Dialect::Scheme, "0.1", Some("wrap"), 1)) + .expect("nested selected call should unwrap"); + + assert_eq!(plan.path, Some("0.1".parse::().expect("valid path"))); + assert_eq!(plan.span.slice(input), "(wrap first second)"); + assert_eq!(plan.argument_span.slice(input), "second"); + assert_eq!(plan.call_argument_count, 2); + assert_eq!(plan.replacement, "second"); + assert_eq!(plan.rewritten, "(outer second tail)"); + } +} diff --git a/src/presentation/cli.rs b/src/presentation/cli.rs index f45674e9..d3553ade 100644 --- a/src/presentation/cli.rs +++ b/src/presentation/cli.rs @@ -12,6 +12,7 @@ mod call_report; mod capabilities; mod command; mod conditional_conversion; +mod contract; mod convert_cond_to_if; mod convert_flet_to_labels; mod convert_if_to_cond; diff --git a/src/presentation/cli/analysis_report/workflow.rs b/src/presentation/cli/analysis_report/workflow.rs index 3c7a7c4b..22ebb1a7 100644 --- a/src/presentation/cli/analysis_report/workflow.rs +++ b/src/presentation/cli/analysis_report/workflow.rs @@ -17,7 +17,7 @@ pub(in crate::presentation::cli) fn check(args: AnalyzeArgs) -> Result<()> { OutputFormat::Json => { let (input, dialect) = read_input_and_dialect(args.file, args.dialect)?; let file = input.file.as_deref().map(|path| path.display().to_string()); - let parse_error = SyntaxTree::parse(&input.text).err(); + let parse_error = SyntaxTree::parse_with_dialect(&input.text, dialect).err(); let report = json!({ "schema_version": 1, "status": if parse_error.is_none() { "ok" } else { "error" }, diff --git a/src/presentation/cli/args.rs b/src/presentation/cli/args.rs index df866973..adb79279 100644 --- a/src/presentation/cli/args.rs +++ b/src/presentation/cli/args.rs @@ -44,6 +44,9 @@ pub(super) struct RepairArgs { /// Input file. Reads stdin when omitted. #[arg(short, long)] pub(super) file: Option, + /// Override extension-based dialect detection. + #[arg(long)] + pub(super) dialect: Option, /// Write the repaired document back to --file instead of stdout. #[arg(long)] pub(super) write: bool, @@ -57,6 +60,9 @@ pub(crate) struct TargetArgs { /// Input file. Reads stdin when omitted. #[arg(short, long)] pub(super) file: Option, + /// Override extension-based dialect detection. + #[arg(long)] + pub(super) dialect: Option, /// Select by child index path, for example 0.2.1. #[arg(long, conflicts_with = "at")] pub(super) path: Option, @@ -84,6 +90,9 @@ pub(super) struct ReplaceArgs { /// Input file. Reads stdin when omitted. #[arg(short, long)] pub(super) file: Option, + /// Override extension-based dialect detection. + #[arg(long)] + pub(super) dialect: Option, /// Select by child index path, for example 0.2.1. #[arg(long, conflicts_with = "at")] pub(super) path: Option, diff --git a/src/presentation/cli/basic_edit/workflow.rs b/src/presentation/cli/basic_edit/workflow.rs index e01ce224..0a10fa56 100644 --- a/src/presentation/cli/basic_edit/workflow.rs +++ b/src/presentation/cli/basic_edit/workflow.rs @@ -5,39 +5,40 @@ use crate::presentation::cli::args::{ EditTargetArgs, FormatArgs, RepairArgs, ReplaceArgs, TargetArgs, }; use crate::presentation::cli::shared::{ - edit_target, emit_document, read_input, read_input_dialect_and_tree, resolve_target, + edit_target, emit_document, read_input_and_dialect, read_input_dialect_and_tree, resolve_target, }; pub(in crate::presentation::cli) fn format(args: FormatArgs) -> Result<()> { - let (input, _, tree) = read_input_dialect_and_tree(args.file, args.dialect)?; + let (input, dialect, tree) = read_input_dialect_and_tree(args.file, args.dialect)?; let rendered = Formatter::new(args.indent).format(&tree); - emit_document(&input, args.write, args.diff, rendered) + emit_document(&input, dialect, args.write, args.diff, rendered) } pub(in crate::presentation::cli) fn repair_unclosed_lists(args: RepairArgs) -> Result<()> { - let input = read_input(args.file)?; + let (input, dialect) = read_input_and_dialect(args.file, args.dialect)?; let repaired = SyntaxTree::repair_unclosed_lists(&input.text) .context("repair-unclosed-lists only repairs unclosed lists")?; if repaired == input.text { bail!("input is already balanced"); } - emit_document(&input, args.write, args.diff, repaired) + emit_document(&input, dialect, args.write, args.diff, repaired) } pub(in crate::presentation::cli) fn select(args: TargetArgs) -> Result<()> { - let (_, _, tree) = read_input_dialect_and_tree(args.file, None)?; + let (_, _, tree) = read_input_dialect_and_tree(args.file, args.dialect)?; let selection = resolve_target(&tree, args.path.as_ref(), args.at)?; print!("{}", selection.text()); Ok(()) } pub(in crate::presentation::cli) fn replace(args: ReplaceArgs) -> Result<()> { - let (input, _, tree) = read_input_dialect_and_tree(args.file, None)?; - SyntaxTree::parse(&args.with).context("replacement is not a valid S-expression document")?; + let (input, dialect, tree) = read_input_dialect_and_tree(args.file, args.dialect)?; + SyntaxTree::parse_with_dialect(&args.with, dialect) + .context("replacement is not a valid S-expression document")?; let selection = resolve_target(&tree, args.path.as_ref(), args.at)?; let rewritten = Edit::replace(&input.text, selection, &args.with)?; - let rewritten = Edit::normalize_changed_line_trivia(&input.text, rewritten)?; - emit_document(&input, args.write, args.diff, rewritten) + let rewritten = Edit::normalize_changed_line_trivia(&input.text, rewritten, dialect)?; + emit_document(&input, dialect, args.write, args.diff, rewritten) } pub(in crate::presentation::cli) fn kill(args: EditTargetArgs) -> Result<()> { diff --git a/src/presentation/cli/capabilities.rs b/src/presentation/cli/capabilities.rs index 1aab9073..83813180 100644 --- a/src/presentation/cli/capabilities.rs +++ b/src/presentation/cli/capabilities.rs @@ -3,29 +3,53 @@ use anyhow::Result; use clap::builder::Command as ClapCommand; -use clap::{Arg, ArgAction, Args, CommandFactory}; +use clap::{Arg, ArgAction, Args, CommandFactory, ValueEnum}; use serde_json::{Value, json}; use super::args::OutputFormat; +#[derive(Clone, Copy, Debug, ValueEnum)] +enum CapabilitiesSchemaVersion { + #[value(name = "1")] + V1, + #[value(name = "2")] + V2, +} + +impl CapabilitiesSchemaVersion { + const fn number(self) -> u8 { + match self { + Self::V1 => 1, + Self::V2 => 2, + } + } +} + #[derive(Debug, Args)] pub(super) struct CapabilitiesArgs { /// Output format. Defaults to json: this command exists for agent discovery. #[arg(long, value_enum, default_value_t = OutputFormat::Json)] pub(super) output: OutputFormat, + + /// Machine-readable schema version. + #[arg(long, value_enum, default_value = "1")] + schema_version: CapabilitiesSchemaVersion, } pub(super) fn capabilities(args: CapabilitiesArgs) -> Result<()> { let root = super::Cli::command(); match args.output { OutputFormat::Json => { - let report = json!({ - "schema_version": 1, + let mut report = json!({ + "schema_version": args.schema_version.number(), "name": root.get_name(), "version": root.get_version(), "about": about_text(&root), - "commands": subcommand_reports(&root), + "commands": subcommand_reports(&root, args.schema_version), }); + if matches!(args.schema_version, CapabilitiesSchemaVersion::V2) { + report["dialect_contract"] = super::contract::dialect_contract_report(); + } println!("{}", serde_json::to_string_pretty(&report)?); } OutputFormat::Text => { @@ -37,16 +61,19 @@ pub(super) fn capabilities(args: CapabilitiesArgs) -> Result<()> { Ok(()) } -fn subcommand_reports(command: &ClapCommand) -> Vec { +fn subcommand_reports( + command: &ClapCommand, + schema_version: CapabilitiesSchemaVersion, +) -> Vec { command .get_subcommands() .filter(|subcommand| subcommand.get_name() != "help") .map(|subcommand| { - let nested = subcommand_reports(subcommand); + let nested = subcommand_reports(subcommand, schema_version); let mut report = json!({ "name": subcommand.get_name(), "about": about_text(subcommand), - "args": argument_reports(subcommand), + "args": argument_reports(subcommand, schema_version), }); if !nested.is_empty() { report["commands"] = Value::Array(nested); @@ -56,10 +83,18 @@ fn subcommand_reports(command: &ClapCommand) -> Vec { .collect() } -fn argument_reports(command: &ClapCommand) -> Vec { +fn argument_reports( + command: &ClapCommand, + schema_version: CapabilitiesSchemaVersion, +) -> Vec { command .get_arguments() - .filter(|arg| !matches!(arg.get_id().as_str(), "help" | "version")) + .filter(|arg| { + !matches!( + (arg.get_id().as_str(), schema_version), + ("help" | "version", _) | ("schema_version", CapabilitiesSchemaVersion::V1) + ) + }) .map(|arg| { json!({ "id": arg.get_id().as_str(), diff --git a/src/presentation/cli/contract.rs b/src/presentation/cli/contract.rs new file mode 100644 index 00000000..5e05aea0 --- /dev/null +++ b/src/presentation/cli/contract.rs @@ -0,0 +1,355 @@ +use serde_json::{Map, Value, json}; + +use crate::domain::dialect::Dialect; +use crate::domain::inline_function::supports_inline_function_dialect; +use crate::domain::inline_let::supports_inline_let_dialect; +use crate::domain::rename::supports_rename_at_dialect; + +pub(super) const DIALECTS: [&str; 6] = [ + "common-lisp", + "emacs-lisp", + "scheme", + "clojure", + "janet", + "fennel", +]; + +#[derive(Clone, Copy)] +enum CommandCategory { + Introspection, + Format, + Structural, + Semantic, +} + +impl CommandCategory { + const fn as_str(self) -> &'static str { + match self { + Self::Introspection => "introspection", + Self::Format => "format", + Self::Structural => "structural", + Self::Semantic => "semantic", + } + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub(super) enum SupportStatus { + Supported, + Unsupported, + Unknown, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub(super) enum DispatchDenial { + Unsupported, + Unknown, +} + +impl SupportStatus { + const fn as_str(self) -> &'static str { + match self { + Self::Supported => "supported", + Self::Unsupported => "unsupported", + Self::Unknown => "unknown", + } + } + + const fn dispatch_decision(self) -> Result<(), DispatchDenial> { + match self { + Self::Supported => Ok(()), + Self::Unsupported => Err(DispatchDenial::Unsupported), + Self::Unknown => Err(DispatchDenial::Unknown), + } + } +} + +const INTROSPECTION_COMMANDS: [&str; 21] = [ + "inspect check", + "inspect dialect", + "inspect stats", + "inspect agent-report", + "inspect capabilities", + "inspect outline", + "inspect form", + "inspect find-symbol", + "inspect symbols", + "inspect calls", + "inspect signature", + "inspect call-graph", + "inspect impact", + "inspect workspace", + "inspect dependencies", + "inspect packages", + "inspect definitions", + "inspect unused-definitions", + "inspect duplicates", + "inspect similarity", + "inspect lets", +]; + +const FORMAT_COMMANDS: [&str; 2] = ["edit format", "edit repair-unclosed-lists"]; + +const STRUCTURAL_COMMANDS: [&str; 12] = [ + "edit select", + "edit replace", + "edit kill", + "edit wrap", + "edit splice", + "edit raise", + "edit transpose-forward", + "edit transpose-backward", + "edit slurp-forward", + "edit slurp-backward", + "edit barf-forward", + "edit barf-backward", +]; + +const SEMANTIC_COMMANDS: [&str; 78] = [ + "refactor plan", + "refactor verify", + "refactor preview", + "refactor check", + "refactor status", + "refactor apply", + "refactor diff", + "refactor workspace-plan", + "refactor workspace-preview", + "refactor workspace-execute", + "refactor remove-definition", + "refactor remove-unused-definitions", + "refactor move-definition", + "refactor split-file", + "refactor sort-definitions", + "refactor move-form", + "refactor insert-top-level", + "refactor replacement-plan", + "refactor replace-forms", + "refactor add-export", + "refactor sort-package-exports", + "refactor sort-package-options", + "refactor merge-package-options", + "refactor rename-package", + "refactor rename-at", + "refactor rename-symbol", + "refactor rename-in-form", + "refactor rename-binding", + "refactor rename-block", + "refactor rename-tag", + "refactor remove-unused-block", + "refactor remove-unused-tag", + "refactor rename-symbols", + "refactor rename-function", + "refactor rename-macrolet", + "refactor rename-symbol-macro", + "refactor rename-local-function", + "refactor replace-function-calls", + "refactor wrap-function-calls", + "refactor unwrap-function-calls", + "refactor unwrap-call", + "refactor thread-expression", + "refactor unthread-expression", + "refactor extract-function", + "refactor extract-local-function", + "refactor extract-constant", + "refactor inline-function", + "refactor inline-lambda", + "refactor inline-local-function", + "refactor inline-symbol-macro", + "refactor inline-literal-constant", + "refactor add-function-parameter", + "refactor move-function-parameter", + "refactor swap-function-parameters", + "refactor reorder-function-parameters", + "refactor remove-function-parameter", + "refactor introduce-let", + "refactor inline-let", + "refactor convert-let-to-let-star", + "refactor convert-let-star-to-let", + "refactor convert-do-star-to-do", + "refactor convert-prog-star-to-prog", + "refactor merge-nested-let-star", + "refactor merge-nested-let", + "refactor merge-nested-flet", + "refactor split-let-star", + "refactor split-let", + "refactor eliminate-empty-binding-form", + "refactor flatten-progn", + "refactor convert-if-to-cond", + "refactor convert-cond-to-if", + "refactor convert-when-to-if", + "refactor convert-unless-to-if", + "refactor convert-if-to-when", + "refactor convert-if-to-unless", + "refactor convert-labels-to-flet", + "refactor convert-flet-to-labels", + "refactor remove-unused-binding", +]; + +const COMMAND_GROUPS: [(&[&str], CommandCategory); 4] = [ + (&INTROSPECTION_COMMANDS, CommandCategory::Introspection), + (&FORMAT_COMMANDS, CommandCategory::Format), + (&STRUCTURAL_COMMANDS, CommandCategory::Structural), + (&SEMANTIC_COMMANDS, CommandCategory::Semantic), +]; + +const STATUS_VALUES: [SupportStatus; 3] = [ + SupportStatus::Supported, + SupportStatus::Unsupported, + SupportStatus::Unknown, +]; + +pub(super) fn support_status(command_path: &str, dialect: &str) -> SupportStatus { + if !DIALECTS.contains(&dialect) || !contains_command(command_path) { + return SupportStatus::Unknown; + } + + let Ok(dialect) = dialect.parse::() else { + return SupportStatus::Unknown; + }; + + let supported = match command_path { + "refactor rename-at" => supports_rename_at_dialect(dialect), + "refactor inline-function" => supports_inline_function_dialect(dialect), + "refactor inline-let" => supports_inline_let_dialect(dialect), + _ => return SupportStatus::Unknown, + }; + + if supported { + SupportStatus::Supported + } else { + SupportStatus::Unsupported + } +} + +#[allow(dead_code)] +pub(super) fn dispatch_decision(command_path: &str, dialect: &str) -> Result<(), DispatchDenial> { + support_status(command_path, dialect).dispatch_decision() +} + +#[allow(dead_code)] +pub(super) fn dispatch_allowed(command_path: &str, dialect: &str) -> bool { + dispatch_decision(command_path, dialect).is_ok() +} + +pub(super) fn dialect_contract_report() -> Value { + let commands = COMMAND_GROUPS + .iter() + .flat_map(|(paths, category)| { + paths.iter().map(move |path| { + let path = *path; + let support = DIALECTS + .iter() + .map(|dialect| { + ( + (*dialect).to_owned(), + Value::String(support_status(path, dialect).as_str().to_owned()), + ) + }) + .collect::>(); + + json!({ + "path": path, + "category": category.as_str(), + "support": support, + }) + }) + }) + .collect::>(); + + json!({ + "command_count": commands.len(), + "dialect_count": DIALECTS.len(), + "cell_count": commands.len() * DIALECTS.len(), + "categories": ["introspection", "format", "structural", "semantic"], + "statuses": STATUS_VALUES.map(SupportStatus::as_str), + "dialects": DIALECTS, + "commands": commands, + }) +} + +fn contains_command(command_path: &str) -> bool { + COMMAND_GROUPS + .iter() + .any(|(paths, _)| paths.contains(&command_path)) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn support_status_decision_is_fail_closed() { + assert_eq!(SupportStatus::Supported.dispatch_decision(), Ok(())); + assert_eq!( + SupportStatus::Unsupported.dispatch_decision(), + Err(DispatchDenial::Unsupported) + ); + assert_eq!( + SupportStatus::Unknown.dispatch_decision(), + Err(DispatchDenial::Unknown) + ); + } + + #[test] + fn dispatch_adapter_allows_only_verified_supported_cells() { + let supported = [ + ("refactor rename-at", "common-lisp"), + ("refactor inline-function", "common-lisp"), + ("refactor inline-function", "emacs-lisp"), + ("refactor inline-let", "common-lisp"), + ("refactor inline-let", "emacs-lisp"), + ("refactor inline-let", "scheme"), + ("refactor inline-let", "clojure"), + ("refactor inline-let", "janet"), + ("refactor inline-let", "fennel"), + ]; + for (command_path, dialect) in supported { + assert_eq!(dispatch_decision(command_path, dialect), Ok(())); + assert!(dispatch_allowed(command_path, dialect)); + } + + let unsupported = [ + ("refactor rename-at", "emacs-lisp"), + ("refactor rename-at", "scheme"), + ("refactor rename-at", "clojure"), + ("refactor rename-at", "janet"), + ("refactor rename-at", "fennel"), + ("refactor inline-function", "scheme"), + ("refactor inline-function", "clojure"), + ("refactor inline-function", "janet"), + ("refactor inline-function", "fennel"), + ]; + for (command_path, dialect) in unsupported { + assert_eq!( + dispatch_decision(command_path, dialect), + Err(DispatchDenial::Unsupported) + ); + assert!(!dispatch_allowed(command_path, dialect)); + } + + let unknown = [ + ("refactor rename-at", "unknown"), + ("refactor inline-function", "unknown"), + ("refactor inline-let", "unknown"), + ("refactor rename-symbol", "scheme"), + ("refactor inline-function extra", "common-lisp"), + ("refactor inline", "common-lisp"), + ("future command", "common-lisp"), + ]; + for (command_path, dialect) in unknown { + assert_eq!( + dispatch_decision(command_path, dialect), + Err(DispatchDenial::Unknown) + ); + assert!(!dispatch_allowed(command_path, dialect)); + } + + for status in STATUS_VALUES { + assert_eq!( + status.dispatch_decision().is_ok(), + status == SupportStatus::Supported + ); + } + } +} diff --git a/src/presentation/cli/definition_movement/insert_top_level.rs b/src/presentation/cli/definition_movement/insert_top_level.rs index 149bae44..efcee5d2 100644 --- a/src/presentation/cli/definition_movement/insert_top_level.rs +++ b/src/presentation/cli/definition_movement/insert_top_level.rs @@ -1,6 +1,7 @@ use anyhow::{Context, Result, bail}; use serde_json::json; +use crate::domain::dialect::Dialect; use crate::domain::sexpr::SyntaxTree; use super::super::shared::{read_input_dialect_and_tree, write_file_with_rollback}; @@ -16,7 +17,8 @@ pub(in crate::presentation::cli) fn insert_top_level(args: InsertTopLevelArgs) - bail!("--insert before/after requires --anchor-path"); } - let replacement_tree = SyntaxTree::parse(&args.with) + let dialect = Dialect::detect(Some(&args.file), args.dialect.map(Into::into)); + let replacement_tree = SyntaxTree::parse_with_dialect(&args.with, dialect) .context("--with must contain a valid, complete top-level S-expression")?; if replacement_tree.root_children().len() != 1 { bail!("--with must contain exactly one top-level S-expression"); @@ -33,7 +35,8 @@ pub(in crate::presentation::cli) fn insert_top_level(args: InsertTopLevelArgs) - "insert-top-level", )?; - SyntaxTree::parse(&rewritten).context("insertion produced invalid Lisp syntax")?; + SyntaxTree::parse_with_dialect(&rewritten, dialect) + .context("insertion produced invalid Lisp syntax")?; let changed = input.text != rewritten; let written = args.write && changed; diff --git a/src/presentation/cli/definition_movement/move_definition.rs b/src/presentation/cli/definition_movement/move_definition.rs index 7b5ae39f..d4bed4e1 100644 --- a/src/presentation/cli/definition_movement/move_definition.rs +++ b/src/presentation/cli/definition_movement/move_definition.rs @@ -5,6 +5,7 @@ use anyhow::{Context, Result}; use crate::application::usecase::definition_report::DefinitionReportItem; use crate::application::usecase::leading_trivia::first_newline_or; use crate::domain::definition::definition_shape; +use crate::domain::dialect::Dialect; use crate::domain::sexpr::{ByteOffset, ByteSpan, Path, SyntaxTree}; use super::super::shared::{ @@ -32,12 +33,13 @@ pub(in crate::presentation::cli) fn move_definition(args: MoveDefinitionArgs) -> read_input_dialect_and_tree(Some(args.from_file.clone()), args.dialect)?; let (to_input, to_file_existed) = read_file_or_empty(&args.to_file)?; let to_dialect = detect_dialect(&to_input, args.dialect); - SyntaxTree::parse(&to_input.text).with_context(|| { - format!( - "destination file is not a valid S-expression document: {}", - args.to_file.display() - ) - })?; + let to_tree = + SyntaxTree::parse_with_dialect(&to_input.text, to_dialect).with_context(|| { + format!( + "destination file is not a valid S-expression document: {}", + args.to_file.display() + ) + })?; let target_index = match args.path.indexes() { [index] => index.get(), @@ -80,7 +82,7 @@ pub(in crate::presentation::cli) fn move_definition(args: MoveDefinitionArgs) -> .slice(&from_input.text) .trim_start_matches('\n') .to_owned(); - let source_package = package_context_before_top_level(&from_tree, target_index)?; + let source_package = package_context_before_top_level(&from_tree, from_dialect, target_index)?; let definition = DefinitionReportItem { path: args.path.to_string(), span, @@ -102,23 +104,26 @@ pub(in crate::presentation::cli) fn move_definition(args: MoveDefinitionArgs) -> &from_input.text[..move_span.start().get()], &from_input.text[move_span.end().get()..] ); - let to_tree = SyntaxTree::parse(&to_input.text)?; - let dest_package = package_context_before_top_level(&to_tree, to_tree.root_children().len())?; + let dest_package = + package_context_before_top_level(&to_tree, to_dialect, to_tree.root_children().len())?; let appended = match &source_package { - Some(package) if dest_package.as_deref() != Some(package.as_str()) => { + Some(package) + if to_dialect == Dialect::CommonLisp + && dest_package.as_deref() != Some(package.as_str()) => + { format!("(in-package {package})\n\n{definition_text}") } _ => definition_text.clone(), }; let to_rewritten = append_top_level_form(&to_input.text, &appended); - SyntaxTree::parse(&from_rewritten).with_context(|| { + SyntaxTree::parse_with_dialect(&from_rewritten, from_dialect).with_context(|| { format!( "source file would become invalid after moving definition: {}", args.from_file.display() ) })?; - SyntaxTree::parse(&to_rewritten).with_context(|| { + SyntaxTree::parse_with_dialect(&to_rewritten, to_dialect).with_context(|| { format!( "destination file would become invalid after receiving definition: {}", args.to_file.display() diff --git a/src/presentation/cli/definition_movement/move_form.rs b/src/presentation/cli/definition_movement/move_form.rs index 4a3da579..c557d74b 100644 --- a/src/presentation/cli/definition_movement/move_form.rs +++ b/src/presentation/cli/definition_movement/move_form.rs @@ -28,12 +28,13 @@ pub(in crate::presentation::cli) fn move_form(args: MoveFormArgs) -> Result<()> read_input_dialect_and_tree(Some(args.from_file.clone()), args.dialect)?; let (to_input, to_file_existed) = read_file_or_empty(&args.to_file)?; let to_dialect = detect_dialect(&to_input, args.dialect); - let to_tree = SyntaxTree::parse(&to_input.text).with_context(|| { - format!( - "destination file is not a valid S-expression document: {}", - args.to_file.display() - ) - })?; + let to_tree = + SyntaxTree::parse_with_dialect(&to_input.text, to_dialect).with_context(|| { + format!( + "destination file is not a valid S-expression document: {}", + args.to_file.display() + ) + })?; let target_index = top_level_path_index(&args.path, "move-form")?; if target_index >= from_tree.root_children().len() { @@ -56,13 +57,13 @@ pub(in crate::presentation::cli) fn move_form(args: MoveFormArgs) -> Result<()> "move-form", )?; - SyntaxTree::parse(&from_rewritten).with_context(|| { + SyntaxTree::parse_with_dialect(&from_rewritten, from_dialect).with_context(|| { format!( "source file would become invalid after moving form: {}", args.from_file.display() ) })?; - SyntaxTree::parse(&to_rewritten).with_context(|| { + SyntaxTree::parse_with_dialect(&to_rewritten, to_dialect).with_context(|| { format!( "destination file would become invalid after receiving form: {}", args.to_file.display() diff --git a/src/presentation/cli/definition_removal/remove_definition.rs b/src/presentation/cli/definition_removal/remove_definition.rs index 23980c5f..83c0fa10 100644 --- a/src/presentation/cli/definition_removal/remove_definition.rs +++ b/src/presentation/cli/definition_removal/remove_definition.rs @@ -44,11 +44,11 @@ pub(in crate::presentation::cli) fn remove_definition(args: RemoveDefinitionArgs category: shape.category, parameter_count: shape.lambda_parameter_count(&view), body_form_count: Some(shape.body_form_count(&view)), - package: package_context_before_top_level(&tree, target_index)?, + package: package_context_before_top_level(&tree, dialect, target_index)?, }; let rewritten = Edit::kill(&input.text, &tree, selection)?; - SyntaxTree::parse(&rewritten).with_context(|| { + SyntaxTree::parse_with_dialect(&rewritten, dialect).with_context(|| { format!( "file would become invalid after removing definition: {}", args.file.display() diff --git a/src/presentation/cli/io.rs b/src/presentation/cli/io.rs index a03af7c9..b8d557e2 100644 --- a/src/presentation/cli/io.rs +++ b/src/presentation/cli/io.rs @@ -373,15 +373,15 @@ pub(crate) fn read_input_dialect_and_tree( explicit: Option, ) -> Result<(SourceInput, Dialect, SyntaxTree)> { let (input, dialect) = read_input_and_dialect(file, explicit)?; - let tree = parse_document(&input)?; + let tree = parse_document(&input, dialect)?; Ok((input, dialect, tree)) } -/// Parses a source document, naming the input and the error's line/column in -/// the context. The underlying [`ParseError`] keeps the raw byte offset, -/// which feeds directly into `--at`. -pub(crate) fn parse_document(input: &SourceInput) -> Result { - SyntaxTree::parse(&input.text).map_err(|error| { +/// Parses a source document with its resolved dialect, naming the input and +/// the error's line/column in the context. The underlying [`ParseError`] keeps +/// the raw byte offset, which feeds directly into `--at`. +pub(crate) fn parse_document(input: &SourceInput, dialect: Dialect) -> Result { + SyntaxTree::parse_with_dialect(&input.text, dialect).map_err(|error| { let location = parse_error_line_column(&input.text, &error); let source = match input.file.as_deref() { Some(path) => path.display().to_string(), @@ -395,7 +395,8 @@ fn parse_error_line_column(text: &str, error: &ParseError) -> String { let position = match error { ParseError::UnexpectedClose { position, .. } | ParseError::MismatchedClose { position, .. } - | ParseError::ResourceLimitExceeded { position, .. } => *position, + | ParseError::ResourceLimitExceeded { position, .. } + | ParseError::UnsupportedReaderDispatch { position, .. } => *position, ParseError::UnclosedList(position) | ParseError::UnterminatedString(position) | ParseError::UnterminatedBlockComment(position) diff --git a/src/presentation/cli/refactor/manifest/check.rs b/src/presentation/cli/refactor/manifest/check.rs index a44eff65..b6c974ad 100644 --- a/src/presentation/cli/refactor/manifest/check.rs +++ b/src/presentation/cli/refactor/manifest/check.rs @@ -45,7 +45,7 @@ pub(in crate::presentation::cli) fn build_refactor_check_result( let rewritten = apply_byte_span_edits(&input, edits)?; let output_hash = stable_text_hash(&rewritten); let output_hash_matches = output_hash == file.output_hash; - let output_parse_ok = SyntaxTree::parse(&rewritten).is_ok(); + let output_parse_ok = SyntaxTree::parse_with_dialect(&rewritten, file.dialect).is_ok(); let changed = rewritten != input; let manifest_flags_match = changed == file.changed && output_parse_ok == file.output_parse_ok; diff --git a/src/presentation/cli/refactor/manifest/parse.rs b/src/presentation/cli/refactor/manifest/parse.rs index e4a74995..1d18cc0e 100644 --- a/src/presentation/cli/refactor/manifest/parse.rs +++ b/src/presentation/cli/refactor/manifest/parse.rs @@ -38,12 +38,18 @@ fn parse_refactor_apply_manifest_file( .as_object() .with_context(|| format!("files[{index}] must be a JSON object"))?; let edits = required_array(object.get("edits"), &format!("files[{index}].edits"))?; + let dialect_field = format!("files[{index}].dialect"); + let dialect_label = required_string(object.get("dialect"), &dialect_field)?; + let dialect = dialect_label.parse::().with_context(|| { + format!("manifest field {dialect_field} has invalid dialect {dialect_label:?}") + })?; Ok(RefactorApplyManifestFile { path: PathBuf::from(required_string( object.get("path"), &format!("files[{index}].path"), )?), + dialect, changed: required_bool(object.get("changed"), &format!("files[{index}].changed"))?, output_parse_ok: required_bool( object.get("output_parse_ok"), diff --git a/src/presentation/cli/refactor/types/manifest.rs b/src/presentation/cli/refactor/types/manifest.rs index 83dad329..90a9e3b6 100644 --- a/src/presentation/cli/refactor/types/manifest.rs +++ b/src/presentation/cli/refactor/types/manifest.rs @@ -13,6 +13,7 @@ pub(in crate::presentation::cli) struct RefactorApplyManifest { #[derive(Debug)] pub(in crate::presentation::cli) struct RefactorApplyManifestFile { pub(in crate::presentation::cli) path: PathBuf, + pub(in crate::presentation::cli) dialect: Dialect, pub(in crate::presentation::cli) changed: bool, pub(in crate::presentation::cli) output_parse_ok: bool, pub(in crate::presentation::cli) input_hash: String, diff --git a/src/presentation/cli/refactor/workflow/manifest/apply.rs b/src/presentation/cli/refactor/workflow/manifest/apply.rs index 581a1d33..781393ad 100644 --- a/src/presentation/cli/refactor/workflow/manifest/apply.rs +++ b/src/presentation/cli/refactor/workflow/manifest/apply.rs @@ -74,7 +74,7 @@ pub(in crate::presentation::cli) fn refactor_apply(args: RefactorApplyArgs) -> R let rewritten = apply_byte_span_edits(&input, edits)?; let output_hash = stable_text_hash(&rewritten); let output_hash_matches = output_hash == file.output_hash; - let output_parse_ok = SyntaxTree::parse(&rewritten).is_ok(); + let output_parse_ok = SyntaxTree::parse_with_dialect(&rewritten, file.dialect).is_ok(); let changed = rewritten != input; let manifest_flags_match = changed == file.changed && output_parse_ok == file.output_parse_ok; @@ -239,6 +239,7 @@ mod tests { "policy": { "passed": true }, "files": [{ "path": source.display().to_string(), + "dialect": "common-lisp", "changed": true, "output_parse_ok": true, "input_hash": stable_text_hash(original), diff --git a/src/presentation/cli/refactor/workflow/manifest/diff.rs b/src/presentation/cli/refactor/workflow/manifest/diff.rs index 4f16759d..bf9d7c73 100644 --- a/src/presentation/cli/refactor/workflow/manifest/diff.rs +++ b/src/presentation/cli/refactor/workflow/manifest/diff.rs @@ -48,7 +48,7 @@ pub(in crate::presentation::cli) fn refactor_diff(args: RefactorDiffArgs) -> Res let rewritten = apply_byte_span_edits(&input, edits)?; let output_hash = stable_text_hash(&rewritten); let output_hash_matches = output_hash == file.output_hash; - let output_parse_ok = SyntaxTree::parse(&rewritten).is_ok(); + let output_parse_ok = SyntaxTree::parse_with_dialect(&rewritten, file.dialect).is_ok(); let changed = rewritten != input; let manifest_flags_match = changed == file.changed && output_parse_ok == file.output_parse_ok; diff --git a/src/presentation/cli/refactor/workflow/preview/build.rs b/src/presentation/cli/refactor/workflow/preview/build.rs index b48388d9..8fc1759b 100644 --- a/src/presentation/cli/refactor/workflow/preview/build.rs +++ b/src/presentation/cli/refactor/workflow/preview/build.rs @@ -67,7 +67,8 @@ pub(in crate::presentation::cli::refactor::workflow) fn build_refactor_preview( total_definitions += definition_count; let changed = rewritten != input.text; - let output_parse_ok = !changed || SyntaxTree::parse(&rewritten).is_ok(); + let output_parse_ok = + !changed || SyntaxTree::parse_with_dialect(&rewritten, dialect).is_ok(); let edit_count = edits.len(); let preview = bounded_preview(&rewritten, request.max_preview_bytes); files.push(RefactorPreviewFile { diff --git a/src/presentation/cli/shared.rs b/src/presentation/cli/shared.rs index a491d60a..9b5126df 100644 --- a/src/presentation/cli/shared.rs +++ b/src/presentation/cli/shared.rs @@ -21,7 +21,7 @@ mod macos_acl; pub(crate) use diff::unified_diff; pub(crate) use io::{AnchoredExpectedWrite, write_files_with_rollback_expected_anchored}; pub(crate) use io::{ - ExpectedWriteTarget, MAX_SOURCE_INPUT_BYTES, parse_document, read_file_or_empty, read_input, + ExpectedWriteTarget, MAX_SOURCE_INPUT_BYTES, parse_document, read_file_or_empty, read_input_and_dialect, read_input_dialect_and_tree, read_text_file_with_expected_target, read_text_file_with_limit, read_text_with_limit, write_artifact_with_rollback, write_file_with_rollback, write_files_with_rollback, write_files_with_rollback_expected, @@ -135,8 +135,13 @@ fn ensure_non_overlapping_spans(spans: impl IntoIterator) -> Re pub(crate) fn package_context_before_top_level( tree: &SyntaxTree, + dialect: Dialect, target_index: usize, ) -> Result> { + if dialect != Dialect::CommonLisp { + return Ok(None); + } + let mut current_package = None; for index in 0..target_index { let path = Path::from_indexes(vec![index]); @@ -189,27 +194,27 @@ pub(crate) fn edit_target( f: fn(&str, &SyntaxTree, Selection<'_>) -> Result, ) -> Result<()> { let target = args.target; - let input = read_input(target.file)?; - let tree = parse_document(&input)?; + let (input, dialect) = read_input_and_dialect(target.file, target.dialect)?; + let tree = parse_document(&input, dialect)?; let selection = resolve_target(&tree, target.path.as_ref(), target.at)?; let rewritten = f(&input.text, &tree, selection)?; - let rewritten = Edit::normalize_changed_line_trivia(&input.text, rewritten)?; - emit_document(&input, args.write, args.diff, rewritten) + let rewritten = Edit::normalize_changed_line_trivia(&input.text, rewritten, dialect)?; + emit_document(&input, dialect, args.write, args.diff, rewritten) } /// Print the rewritten document to stdout, or with `write` persist it back to -/// the source file after confirming the result still parses as a balanced -/// document. With `diff`, stdout carries a unified diff against the input -/// instead of the whole rewritten document. +/// the source file after confirming the result reparses with the input dialect. +/// With `diff`, stdout carries a unified diff instead of the whole document. pub(crate) fn emit_document( input: &SourceInput, + dialect: Dialect, write: bool, diff: bool, rewritten: String, ) -> Result<()> { if write { let path = require_output_file(input.file.as_ref())?.clone(); - SyntaxTree::parse(&rewritten) + SyntaxTree::parse_with_dialect(&rewritten, dialect) .context("refusing to write: rewritten source does not reparse")?; if diff { print!("{}", unified_diff(&path, &input.text, &rewritten)); diff --git a/src/presentation/cli/similarity_report/workflow.rs b/src/presentation/cli/similarity_report/workflow.rs index aac79072..43c081a7 100644 --- a/src/presentation/cli/similarity_report/workflow.rs +++ b/src/presentation/cli/similarity_report/workflow.rs @@ -203,11 +203,12 @@ fn process_file( source: source.into(), })?; let dialect = Dialect::detect(Some(file), dialect.map(Into::into)); - let tree = SyntaxTree::parse(&text).map_err(|source| ProcessingError { - path: file.to_path_buf(), - stage: "parse", - source: source.into(), - })?; + let tree = + SyntaxTree::parse_with_dialect(&text, dialect).map_err(|source| ProcessingError { + path: file.to_path_buf(), + stage: "parse", + source: source.into(), + })?; let mut candidates = Vec::new(); let omitted_candidates = collect_similarity_candidates(&tree, &text, file, dialect, options, &mut candidates) diff --git a/src/presentation/cli/workspace_report/workflow.rs b/src/presentation/cli/workspace_report/workflow.rs index 49bb0be4..33f3e7f8 100644 --- a/src/presentation/cli/workspace_report/workflow.rs +++ b/src/presentation/cli/workspace_report/workflow.rs @@ -59,7 +59,7 @@ pub(in crate::presentation::cli) fn workspace_report(args: WorkspaceReportArgs) } }; - match SyntaxTree::parse(&text) { + match SyntaxTree::parse_with_dialect(&text, dialect) { Ok(tree) => { let (package, definitions) = collect_definition_forms(&tree, dialect) .with_context(|| format!("failed to analyze {}", file.display()))?; diff --git a/tests/cli.rs b/tests/cli.rs index e8975bea..aef97b94 100644 --- a/tests/cli.rs +++ b/tests/cli.rs @@ -47,6 +47,8 @@ mod definition_removal; mod definition_report; #[path = "cli/dependency_report.rs"] mod dependency_report; +#[path = "cli/dialect_contract.rs"] +mod dialect_contract; #[path = "cli/duplicate_report.rs"] mod duplicate_report; #[path = "cli/edit_transpose.rs"] diff --git a/tests/cli/analysis_report.rs b/tests/cli/analysis_report.rs index d9ac5162..9395c775 100644 --- a/tests/cli/analysis_report.rs +++ b/tests/cli/analysis_report.rs @@ -120,6 +120,49 @@ fn cli_check_reports_ok_as_json() { .stdout(predicate::str::contains("\"error\": null")); } +#[test] +fn cli_check_json_applies_reader_policy_and_preserves_unknown_parsing() { + for (dialect, input) in [ + ("common-lisp", "(list #\\))"), + ("emacs-lisp", "(list ?\\))"), + ] { + let mut cmd = paredit(); + cmd.args(["inspect", "check", "--dialect", dialect, "--output", "json"]) + .write_stdin(input) + .assert() + .success() + .stdout(predicate::str::contains("\"status\": \"ok\"")) + .stdout(predicate::str::contains(format!( + "\"dialect\": \"{dialect}\"" + ))); + } + + let mut known = paredit(); + known + .args([ + "inspect", + "check", + "--dialect", + "common-lisp", + "--output", + "json", + ]) + .write_stdin("#?value") + .assert() + .failure() + .stdout(predicate::str::contains("\"status\": \"error\"")) + .stdout(predicate::str::contains("unsupported reader dispatch")); + + let mut unknown = paredit(); + unknown + .args(["inspect", "check", "--output", "json"]) + .write_stdin("#?value") + .assert() + .success() + .stdout(predicate::str::contains("\"status\": \"ok\"")) + .stdout(predicate::str::contains("\"dialect\": \"unknown\"")); +} + #[test] fn cli_check_reports_parse_error_as_json_and_exits_nonzero() { let mut cmd = paredit(); diff --git a/tests/cli/basic_edit_write.rs b/tests/cli/basic_edit_write.rs index 8079f553..dec02d4c 100644 --- a/tests/cli/basic_edit_write.rs +++ b/tests/cli/basic_edit_write.rs @@ -151,6 +151,56 @@ fn assert_edit_output(subcommand: &str, path: &str, expected: &str) { .stdout(predicate::str::contains(expected)); } +#[test] +fn edit_replace_keeps_generic_stdin_compatibility_without_dialect_flag() { + paredit() + .args(["edit", "replace", "--path", "1", "--with", "(new)"]) + .write_stdin("(old) (keep)\n") + .assert() + .success() + .stdout(predicate::eq("(old) (new)\n")); +} + +#[test] +fn edit_replace_uses_explicit_clojure_dialect_for_stdin_reader_forms() { + paredit() + .args([ + "edit", + "replace", + "--dialect", + "clojure", + "--path", + "1", + "--with", + "(new)", + ]) + .write_stdin("#inst \"1985-04-12T23:20:50.52-00:00\" (old)\n") + .assert() + .success() + .stdout(predicate::eq( + "#inst \"1985-04-12T23:20:50.52-00:00\" (new)\n", + )); +} + +#[test] +fn edit_replace_uses_explicit_common_lisp_dialect_for_stdin_reader_forms() { + paredit() + .args([ + "edit", + "replace", + "--dialect", + "common-lisp", + "--path", + "1", + "--with", + "(new)", + ]) + .write_stdin("#S(point :x 1) (old)\n") + .assert() + .success() + .stdout(predicate::eq("#S(point :x 1) (new)\n")); +} + #[test] fn edit_splice_removes_one_list_pair() { assert_edit_output("splice", "0.1", "(a b c d e)"); diff --git a/tests/cli/capabilities_contract.rs b/tests/cli/capabilities_contract.rs index 8d37dccc..d2a5900b 100644 --- a/tests/cli/capabilities_contract.rs +++ b/tests/cli/capabilities_contract.rs @@ -1,8 +1,9 @@ use super::*; -fn capabilities_json() -> serde_json::Value { +fn capabilities_json_with_args(args: &[&str]) -> serde_json::Value { let output = paredit() .args(["inspect", "capabilities"]) + .args(args) .assert() .success() .get_output() @@ -11,6 +12,52 @@ fn capabilities_json() -> serde_json::Value { serde_json::from_slice(&output).expect("capabilities emits valid JSON") } +fn capabilities_json() -> serde_json::Value { + capabilities_json_with_args(&[]) +} + +#[test] +fn capabilities_defaults_to_unchanged_schema_v1_shape() { + let report = capabilities_json(); + let explicit_v1 = capabilities_json_with_args(&["--schema-version", "1"]); + + assert_eq!(report["schema_version"], 1); + assert_eq!(report, explicit_v1); + assert!(report.get("dialect_contract").is_none()); + + let mut keys = report + .as_object() + .expect("capabilities report is an object") + .keys() + .map(String::as_str) + .collect::>(); + keys.sort_unstable(); + assert_eq!( + keys, + ["about", "commands", "name", "schema_version", "version"] + ); + + let capabilities = report["commands"] + .as_array() + .expect("commands") + .iter() + .find(|command| command["name"] == "inspect") + .and_then(|inspect| inspect["commands"].as_array()) + .and_then(|commands| { + commands + .iter() + .find(|command| command["name"] == "capabilities") + }) + .expect("inspect capabilities command"); + assert!( + capabilities["args"] + .as_array() + .expect("inspect capabilities args") + .iter() + .all(|arg| arg["id"] != "schema_version") + ); +} + #[test] fn capabilities_reports_the_canonical_namespaces_and_meta_commands() { let report = capabilities_json(); diff --git a/tests/cli/definition_movement.rs b/tests/cli/definition_movement.rs index ced415c3..a8d91251 100644 --- a/tests/cli/definition_movement.rs +++ b/tests/cli/definition_movement.rs @@ -432,3 +432,129 @@ fn cli_writes_top_level_form_move_before_anchor() { .expect("anchor form should exist"); assert!(moved_index < anchor_index); } + +#[test] +fn cli_moves_common_lisp_character_literal_with_its_reader_policy() { + let dir = fresh_temp_dir("move-form-common-lisp-character"); + let from_file = dir.join("source.lisp"); + let to_file = dir.join("destination.lisp"); + fs::write( + &from_file, + "(defparameter *keep* :ok)\n(defparameter *close* #\\))\n", + ) + .expect("write Common Lisp source fixture"); + fs::write(&to_file, "(defparameter *existing* :ok)\n") + .expect("write Common Lisp destination fixture"); + + paredit() + .args([ + "refactor", + "move-form", + "--from-file", + from_file.to_str().unwrap(), + "--to-file", + to_file.to_str().unwrap(), + "--path", + "1", + "--write", + ]) + .assert() + .success(); + + assert!(!fs::read_to_string(&from_file).unwrap().contains("*close*")); + assert!( + fs::read_to_string(&to_file) + .unwrap() + .contains("(defparameter *close* #\\))") + ); +} + +#[test] +fn cli_inserts_emacs_lisp_character_literal_with_its_reader_policy() { + let dir = fresh_temp_dir("insert-top-level-emacs-lisp-character"); + let file = dir.join("source.el"); + fs::write(&file, "(defvar existing nil)\n").expect("write Emacs Lisp fixture"); + + paredit() + .args([ + "refactor", + "insert-top-level", + "--file", + file.to_str().unwrap(), + "--with", + "(defconst close ?\\))", + "--write", + ]) + .assert() + .success(); + + assert!( + fs::read_to_string(&file) + .unwrap() + .contains("(defconst close ?\\))") + ); +} + +#[test] +fn cli_moves_unknown_dialect_form_with_generic_reader_compatibility() { + let dir = fresh_temp_dir("move-form-unknown-dialect"); + let from_file = dir.join("source.txt"); + let to_file = dir.join("destination.txt"); + fs::write(&from_file, "(keep #?value)\n(move #?other)\n") + .expect("write unknown-dialect source fixture"); + fs::write(&to_file, "(existing #?value)\n").expect("write unknown-dialect destination fixture"); + + paredit() + .args([ + "refactor", + "move-form", + "--from-file", + from_file.to_str().unwrap(), + "--to-file", + to_file.to_str().unwrap(), + "--path", + "1", + "--write", + ]) + .assert() + .success(); + + assert!(!fs::read_to_string(&from_file).unwrap().contains("(move")); + assert!( + fs::read_to_string(&to_file) + .unwrap() + .contains("(move #?other)") + ); +} + +#[test] +fn cli_does_not_inject_common_lisp_package_when_moving_definition_to_emacs_lisp() { + let dir = fresh_temp_dir("move-definition-common-lisp-to-emacs-lisp"); + let from_file = dir.join("source.lisp"); + let to_file = dir.join("destination.el"); + fs::write( + &from_file, + "(in-package #:demo)\n(defun keep () :ok)\n(defun moved () :moved)\n", + ) + .expect("write Common Lisp source fixture"); + fs::write(&to_file, "(defvar close ?\\))\n").expect("write Emacs Lisp destination fixture"); + + paredit() + .args([ + "refactor", + "move-definition", + "--from-file", + from_file.to_str().unwrap(), + "--to-file", + to_file.to_str().unwrap(), + "--path", + "2", + "--write", + ]) + .assert() + .success(); + + let destination = fs::read_to_string(&to_file).expect("read Emacs Lisp destination"); + assert!(destination.contains("(defun moved () :moved)")); + assert!(!destination.contains("in-package")); +} diff --git a/tests/cli/definition_removal.rs b/tests/cli/definition_removal.rs index fdfbb3d9..70868ae7 100644 --- a/tests/cli/definition_removal.rs +++ b/tests/cli/definition_removal.rs @@ -59,6 +59,37 @@ fn cli_writes_definition_removal() { assert!(!rewritten.contains("stale-helper")); } +#[test] +fn cli_reparses_definition_removal_with_the_source_dialect() { + let dir = fresh_temp_dir("remove-definition-reader-policy"); + let cases = [ + ("common.lisp", "(defun keep () #\\))\n"), + ("emacs.el", "(defun keep () ?\\))\n"), + ]; + + for (file_name, keep_definition) in cases { + let file = dir.join(file_name); + fs::write(&file, format!("{keep_definition}(defun stale () :stale)\n")) + .expect("write fixture"); + + let mut cmd = paredit(); + cmd.arg("refactor") + .arg("remove-definition") + .arg("--file") + .arg(&file) + .arg("--path") + .arg("1") + .arg("--write") + .assert() + .success() + .stdout(predicate::str::contains("\"written\": true")); + + let rewritten = fs::read_to_string(&file).expect("read rewritten fixture"); + assert!(rewritten.contains(keep_definition.trim())); + assert!(!rewritten.contains("stale")); + } +} + #[test] fn cli_plans_unused_definition_removal_without_writing() { let dir = fresh_temp_dir("remove-unused-definitions-plan"); diff --git a/tests/cli/dialect_contract.rs b/tests/cli/dialect_contract.rs new file mode 100644 index 00000000..d5803bad --- /dev/null +++ b/tests/cli/dialect_contract.rs @@ -0,0 +1,208 @@ +use super::*; +use std::collections::{BTreeMap, BTreeSet}; + +fn capabilities_json(schema_version: &str) -> serde_json::Value { + let output = paredit() + .args([ + "inspect", + "capabilities", + "--schema-version", + schema_version, + ]) + .assert() + .success() + .get_output() + .stdout + .clone(); + serde_json::from_slice(&output).expect("capabilities emits valid JSON") +} + +fn collect_leaf_paths( + commands: &[serde_json::Value], + prefix: &mut Vec, + leaves: &mut BTreeSet, +) { + for command in commands { + let name = command["name"].as_str().expect("command name"); + prefix.push(name.to_owned()); + + match command + .get("commands") + .and_then(serde_json::Value::as_array) + { + Some(children) if !children.is_empty() => { + collect_leaf_paths(children, prefix, leaves); + } + _ => { + assert!(leaves.insert(prefix.join(" ")), "duplicate Clap leaf path"); + } + } + + prefix.pop(); + } +} + +fn clap_contract_leaf_paths(report: &serde_json::Value) -> BTreeSet { + let mut leaves = BTreeSet::new(); + for namespace in report["commands"].as_array().expect("root commands") { + let name = namespace["name"].as_str().expect("namespace name"); + if !matches!(name, "inspect" | "edit" | "refactor") { + continue; + } + + let mut prefix = vec![name.to_owned()]; + collect_leaf_paths( + namespace["commands"] + .as_array() + .expect("namespace commands"), + &mut prefix, + &mut leaves, + ); + } + leaves +} + +#[test] +fn schema_v2_registry_is_an_exact_bijection_with_clap_leaves() { + let v1 = capabilities_json("1"); + let v2 = capabilities_json("2"); + let commands = v2["dialect_contract"]["commands"] + .as_array() + .expect("dialect contract commands"); + let registry_paths = commands + .iter() + .map(|command| command["path"].as_str().expect("registry command path")) + .collect::>(); + let unique_registry_paths = registry_paths.iter().copied().collect::>(); + + assert_eq!(registry_paths.len(), unique_registry_paths.len()); + assert_eq!(registry_paths.len(), 113); + assert_eq!( + clap_contract_leaf_paths(&v1), + unique_registry_paths + .into_iter() + .map(str::to_owned) + .collect::>() + ); +} + +#[test] +fn schema_v2_reports_the_complete_dialect_matrix() { + let report = capabilities_json("2"); + assert_eq!(report["schema_version"], 2); + + let contract = &report["dialect_contract"]; + assert_eq!(contract["command_count"], 113); + assert_eq!(contract["dialect_count"], 6); + assert_eq!(contract["cell_count"], 678); + assert_eq!( + contract["dialects"], + serde_json::json!([ + "common-lisp", + "emacs-lisp", + "scheme", + "clojure", + "janet", + "fennel" + ]) + ); + assert_eq!( + contract["statuses"], + serde_json::json!(["supported", "unsupported", "unknown"]) + ); + + let commands = contract["commands"].as_array().expect("contract commands"); + let mut category_counts = BTreeMap::new(); + let mut cell_count = 0; + let mut supported_cells = BTreeSet::new(); + let mut unsupported_cells = BTreeSet::new(); + let expected_dialects = [ + "common-lisp", + "emacs-lisp", + "scheme", + "clojure", + "janet", + "fennel", + ] + .into_iter() + .collect::>(); + let valid_statuses = ["supported", "unsupported", "unknown"] + .into_iter() + .collect::>(); + + for command in commands { + let path = command["path"].as_str().expect("command path"); + let category = command["category"].as_str().expect("command category"); + *category_counts.entry(category).or_insert(0) += 1; + + let support = command["support"].as_object().expect("support map"); + assert_eq!( + support.keys().map(String::as_str).collect::>(), + expected_dialects, + "dialect columns for {path}" + ); + cell_count += support.len(); + + for (dialect, status) in support { + let status = status.as_str().expect("support status"); + assert!( + valid_statuses.contains(status), + "status for {path}/{dialect}" + ); + match status { + "supported" => { + supported_cells.insert(format!("{path}|{dialect}")); + } + "unsupported" => { + unsupported_cells.insert(format!("{path}|{dialect}")); + } + _ => {} + } + } + } + + assert_eq!( + category_counts, + BTreeMap::from([ + ("format", 2), + ("introspection", 21), + ("semantic", 78), + ("structural", 12), + ]) + ); + assert_eq!(cell_count, 678); + assert_eq!( + supported_cells, + [ + "refactor inline-function|common-lisp", + "refactor inline-function|emacs-lisp", + "refactor inline-let|clojure", + "refactor inline-let|common-lisp", + "refactor inline-let|emacs-lisp", + "refactor inline-let|fennel", + "refactor inline-let|janet", + "refactor inline-let|scheme", + "refactor rename-at|common-lisp", + ] + .into_iter() + .map(str::to_owned) + .collect::>() + ); + assert_eq!( + unsupported_cells, + [ + "refactor inline-function|clojure", + "refactor inline-function|fennel", + "refactor inline-function|janet", + "refactor inline-function|scheme", + "refactor rename-at|clojure", + "refactor rename-at|emacs-lisp", + "refactor rename-at|fennel", + "refactor rename-at|janet", + "refactor rename-at|scheme", + ] + .into_iter() + .map(str::to_owned) + .collect::>() + ); +} diff --git a/tests/cli/extract_function/inference/basic.rs b/tests/cli/extract_function/inference/basic.rs index 63e78515..f0c7da25 100644 --- a/tests/cli/extract_function/inference/basic.rs +++ b/tests/cli/extract_function/inference/basic.rs @@ -4,6 +4,8 @@ use super::assert_extract_function_inference; fn cli_infers_extract_function_params() { assert_extract_function_inference( &[ + "--dialect", + "common-lisp", "--path", "0.3", "--name", @@ -24,6 +26,8 @@ fn cli_infers_extract_function_params() { fn cli_infers_extract_function_params_without_call_heads_or_literals() { assert_extract_function_inference( &[ + "--dialect", + "common-lisp", "--path", "0.3", "--name", @@ -44,6 +48,8 @@ fn cli_infers_extract_function_params_without_call_heads_or_literals() { fn cli_does_not_infer_common_lisp_function_literals() { assert_extract_function_inference( &[ + "--dialect", + "common-lisp", "--path", "0.3", "--name", @@ -64,6 +70,8 @@ fn cli_does_not_infer_common_lisp_function_literals() { fn cli_infers_extract_function_params_without_local_let_bindings() { assert_extract_function_inference( &[ + "--dialect", + "common-lisp", "--path", "0.3", "--name", @@ -106,6 +114,8 @@ fn cli_infers_extract_function_params_without_emacs_lisp_local_let_bindings() { fn cli_infers_extract_function_params_without_sequential_let_bindings() { assert_extract_function_inference( &[ + "--dialect", + "common-lisp", "--path", "0.3", "--name", diff --git a/tests/cli/extract_function/inference/bindings.rs b/tests/cli/extract_function/inference/bindings.rs index e901d733..57f0728c 100644 --- a/tests/cli/extract_function/inference/bindings.rs +++ b/tests/cli/extract_function/inference/bindings.rs @@ -4,6 +4,8 @@ use super::assert_extract_function_inference; fn cli_infers_extract_function_params_without_symbol_macrolet_bindings() { assert_extract_function_inference( &[ + "--dialect", + "common-lisp", "--path", "0.3", "--name", @@ -26,6 +28,8 @@ fn cli_infers_extract_function_params_without_symbol_macrolet_bindings() { fn cli_infers_extract_function_params_without_package_qualified_symbol_macrolet_bindings() { assert_extract_function_inference( &[ + "--dialect", + "common-lisp", "--path", "0.3", "--name", @@ -114,6 +118,8 @@ fn cli_infers_extract_function_params_without_clojure_destructuring_lambda_bindi fn cli_infers_extract_function_params_without_bare_symbol_let_bindings() { assert_extract_function_inference( &[ + "--dialect", + "common-lisp", "--path", "0.3", "--name", diff --git a/tests/cli/extract_function/inference/local_callables.rs b/tests/cli/extract_function/inference/local_callables.rs index a5a03285..09fe7956 100644 --- a/tests/cli/extract_function/inference/local_callables.rs +++ b/tests/cli/extract_function/inference/local_callables.rs @@ -4,6 +4,8 @@ use super::assert_extract_function_inference; fn cli_infers_extract_function_params_without_common_lisp_flet_shadowing() { assert_extract_function_inference( &[ + "--dialect", + "common-lisp", "--path", "0.3", "--name", @@ -48,6 +50,8 @@ fn cli_infers_extract_function_params_without_emacs_lisp_flet_shadowing() { fn cli_infers_extract_function_params_without_common_lisp_labels_shadowing() { assert_extract_function_inference( &[ + "--dialect", + "common-lisp", "--path", "0.3", "--name", @@ -70,6 +74,8 @@ fn cli_infers_extract_function_params_without_common_lisp_labels_shadowing() { fn cli_infers_extract_function_params_without_common_lisp_macrolet_shadowing() { assert_extract_function_inference( &[ + "--dialect", + "common-lisp", "--path", "0.3", "--name", @@ -92,6 +98,8 @@ fn cli_infers_extract_function_params_without_common_lisp_macrolet_shadowing() { fn cli_infers_extract_function_params_without_cl_user_macrolet_shadowing() { assert_extract_function_inference( &[ + "--dialect", + "common-lisp", "--path", "0.3", "--name", @@ -114,6 +122,8 @@ fn cli_infers_extract_function_params_without_cl_user_macrolet_shadowing() { fn cli_infers_extract_function_params_without_common_lisp_package_macrolet_shadowing() { assert_extract_function_inference( &[ + "--dialect", + "common-lisp", "--path", "0.3", "--name", @@ -136,6 +146,8 @@ fn cli_infers_extract_function_params_without_common_lisp_package_macrolet_shado fn cli_infers_extract_function_params_without_common_lisp_compiler_macrolet_shadowing() { assert_extract_function_inference( &[ + "--dialect", + "common-lisp", "--path", "0.3", "--name", @@ -158,6 +170,8 @@ fn cli_infers_extract_function_params_without_common_lisp_compiler_macrolet_shad fn cli_infers_extract_function_params_without_common_lisp_package_compiler_macrolet_shadowing() { assert_extract_function_inference( &[ + "--dialect", + "common-lisp", "--path", "0.3", "--name", @@ -180,6 +194,8 @@ fn cli_infers_extract_function_params_without_common_lisp_package_compiler_macro fn cli_infers_extract_function_params_without_cl_user_compiler_macrolet_shadowing() { assert_extract_function_inference( &[ + "--dialect", + "common-lisp", "--path", "0.3", "--name", diff --git a/tests/cli/extract_function/inference/macro_lambda_lists.rs b/tests/cli/extract_function/inference/macro_lambda_lists.rs index 0ad03dec..a3b9e857 100644 --- a/tests/cli/extract_function/inference/macro_lambda_lists.rs +++ b/tests/cli/extract_function/inference/macro_lambda_lists.rs @@ -4,6 +4,8 @@ use super::assert_extract_function_inference; fn cli_infers_extract_function_params_without_common_lisp_lambda_list_init_bindings() { assert_extract_function_inference( &[ + "--dialect", + "common-lisp", "--path", "0.3", "--name", @@ -26,6 +28,8 @@ fn cli_infers_extract_function_params_without_common_lisp_lambda_list_init_bindi fn cli_infers_extract_function_params_without_define_setf_expander_macro_lambda_list_bindings() { assert_extract_function_inference( &[ + "--dialect", + "common-lisp", "--path", "0", "--name", @@ -48,6 +52,8 @@ fn cli_infers_extract_function_params_without_define_setf_expander_macro_lambda_ fn cli_infers_extract_function_params_without_define_compiler_macro_lambda_list_bindings() { assert_extract_function_inference( &[ + "--dialect", + "common-lisp", "--path", "0", "--name", @@ -71,6 +77,8 @@ fn cli_infers_extract_function_params_without_package_qualified_compiler_macro_l { assert_extract_function_inference( &[ + "--dialect", + "common-lisp", "--path", "0", "--name", diff --git a/tests/cli/extract_function/params.rs b/tests/cli/extract_function/params.rs index 03f81e92..67305ab0 100644 --- a/tests/cli/extract_function/params.rs +++ b/tests/cli/extract_function/params.rs @@ -4,6 +4,8 @@ use super::*; fn cli_plans_parameterized_extract_function() { let mut cmd = paredit(); cmd.args(["refactor", "extract-function", + "--dialect", + "common-lisp", "--path", "0.3", "--name", @@ -41,6 +43,8 @@ fn cli_merges_explicit_and_inferred_extract_function_params() { cmd.args([ "refactor", "extract-function", + "--dialect", + "common-lisp", "--path", "0.3", "--name", diff --git a/tests/cli/extract_function/planning.rs b/tests/cli/extract_function/planning.rs index 39a40a11..3b7d8e83 100644 --- a/tests/cli/extract_function/planning.rs +++ b/tests/cli/extract_function/planning.rs @@ -6,6 +6,8 @@ fn cli_plans_extract_function_for_common_lisp() { cmd.args([ "refactor", "extract-function", + "--dialect", + "common-lisp", "--path", "0.3", "--name", @@ -16,7 +18,7 @@ fn cli_plans_extract_function_for_common_lisp() { .write_stdin("(defun render () (+ 1 2))") .assert() .success() - .stdout(predicate::str::contains("\"dialect\": \"unknown\"")) + .stdout(predicate::str::contains("\"dialect\": \"common-lisp\"")) .stdout(predicate::str::contains("\"call\": \"(compute-sum)\"")) .stdout(predicate::str::contains( "\"definition\": \"(defun compute-sum () (+ 1 2))\"", @@ -30,6 +32,8 @@ fn cli_plans_extract_function_for_common_lisp() { fn cli_plans_extract_function_for_common_lisp_macrolet_body() { let mut cmd = paredit(); cmd.args(["refactor", "extract-function", + "--dialect", + "common-lisp", "--path", "0.3", "--name", @@ -56,6 +60,8 @@ fn cli_plans_extract_function_for_common_lisp_macrolet_body() { fn cli_plans_extract_function_for_common_lisp_symbol_macrolet_body() { let mut cmd = paredit(); cmd.args(["refactor", "extract-function", + "--dialect", + "common-lisp", "--path", "0.3", "--name", diff --git a/tests/cli/format/mod.rs b/tests/cli/format/mod.rs index cc0b7cf4..5efbbb56 100644 --- a/tests/cli/format/mod.rs +++ b/tests/cli/format/mod.rs @@ -20,3 +20,22 @@ fn assert_format_output(fixture_name: &str, file_name: &str, input: &str, expect .success() .stdout(predicate::str::contains(expected)); } + +#[test] +fn cli_formats_janet_hash_comment_without_changing_output() { + let input = "# keep this comment\n(foo)\n"; + let dir = fresh_temp_dir("format-janet-hash-comment"); + let file = dir.join(Path::new("source.janet")); + fs::write(&file, input).expect("write source fixture"); + + let mut cmd = paredit(); + cmd.arg("edit") + .arg("format") + .arg("--dialect") + .arg("janet") + .arg("--file") + .arg(&file) + .assert() + .success() + .stdout(predicate::eq(input)); +} diff --git a/tests/cli/help_contract.rs b/tests/cli/help_contract.rs index f5e8cdd1..a7df46ab 100644 --- a/tests/cli/help_contract.rs +++ b/tests/cli/help_contract.rs @@ -77,3 +77,14 @@ fn rename_local_function_help_surfaces_flet_and_labels_boundary() { "preserving the difference between non-recursive flet bodies and recursive labels bodies", )); } + +#[test] +fn basic_edit_help_surfaces_optional_dialect_override() { + for subcommand in ["repair-unclosed-lists", "select", "replace", "kill"] { + paredit() + .args(["edit", subcommand, "--help"]) + .assert() + .success() + .stdout(predicate::str::contains("--dialect ")); + } +} diff --git a/tests/cli/inline_function/all_calls.rs b/tests/cli/inline_function/all_calls.rs index 22b11a66..24b4dda0 100644 --- a/tests/cli/inline_function/all_calls.rs +++ b/tests/cli/inline_function/all_calls.rs @@ -3,7 +3,15 @@ use super::*; #[test] fn cli_plans_inline_function_with_all_calls() { assert_inline_success( - &["--definition-path", "0", "--all-calls", "--output", "json"], + &[ + "--dialect", + "common-lisp", + "--definition-path", + "0", + "--all-calls", + "--output", + "json", + ], "(defun area (width height) (* width height))\n\ (defun render () (area 10 20))\n\ (defun summarize () (+ (area 3 4) 1))", diff --git a/tests/cli/inline_function/basic.rs b/tests/cli/inline_function/basic.rs index 9b074112..66eeb023 100644 --- a/tests/cli/inline_function/basic.rs +++ b/tests/cli/inline_function/basic.rs @@ -13,6 +13,8 @@ fn cli_requires_file_for_inline_function_writes() { fn cli_plans_inline_function_with_parameters() { assert_inline_success( &[ + "--dialect", + "common-lisp", "--definition-path", "0", "--call-path", @@ -36,7 +38,14 @@ fn cli_plans_inline_function_with_parameters() { #[test] fn cli_rejects_inline_function_duplicate_evaluation_without_flag() { assert_inline_failure( - &["--definition-path", "0", "--call-path", "1.3"], + &[ + "--dialect", + "common-lisp", + "--definition-path", + "0", + "--call-path", + "1.3", + ], Some("(defun twice (x) (+ x x))\n(defun render () (twice (expensive)))"), &["inline-function would duplicate argument"], ); @@ -46,6 +55,8 @@ fn cli_rejects_inline_function_duplicate_evaluation_without_flag() { fn cli_rejects_inline_function_all_calls_with_explicit_call_path() { assert_inline_failure( &[ + "--dialect", + "common-lisp", "--definition-path", "0", "--call-path", diff --git a/tests/cli/let_refactor/inline.rs b/tests/cli/let_refactor/inline.rs index ee3a7ce7..18b8d499 100644 --- a/tests/cli/let_refactor/inline.rs +++ b/tests/cli/let_refactor/inline.rs @@ -17,13 +17,15 @@ fn cli_plans_inline_let_for_common_lisp() { "inline-let", "--path", "0.3", + "--dialect", + "common-lisp", "--output", "json", ]) .write_stdin("(defun render () (let ((product (* width height))) (+ product margin)))") .assert() .success() - .stdout(predicate::str::contains("\"dialect\": \"unknown\"")) + .stdout(predicate::str::contains("\"dialect\": \"common-lisp\"")) .stdout(predicate::str::contains("\"binding_name\": \"product\"")) .stdout(predicate::str::contains( "\"binding_value\": \"(* width height)\"", @@ -42,6 +44,8 @@ fn cli_plans_inline_let_for_common_lisp_symbol_macrolet() { "inline-let", "--path", "0.3", + "--dialect", + "common-lisp", "--output", "json", ]) @@ -50,7 +54,7 @@ fn cli_plans_inline_let_for_common_lisp_symbol_macrolet() { ) .assert() .success() - .stdout(predicate::str::contains("\"dialect\": \"unknown\"")) + .stdout(predicate::str::contains("\"dialect\": \"common-lisp\"")) .stdout(predicate::str::contains("\"binding_name\": \"product\"")) .stdout(predicate::str::contains( "\"binding_value\": \"(* width height)\"", @@ -70,6 +74,8 @@ fn cli_plans_inline_let_with_multiple_body_expressions() { "--path", "0.3", "--allow-duplicate-evaluation", + "--dialect", + "common-lisp", "--output", "json", ]) @@ -146,11 +152,18 @@ fn cli_writes_inline_let_for_emacs_lisp_cl_symbol_macrolet_file() { #[test] fn cli_rejects_inline_let_duplicate_evaluation_by_default() { let mut cmd = paredit(); - cmd.args(["refactor", "inline-let", "--path", "0"]) - .write_stdin("(let ((x (compute))) (+ x x))") - .assert() - .failure() - .stderr(predicate::str::contains("would duplicate")); + cmd.args([ + "refactor", + "inline-let", + "--path", + "0", + "--dialect", + "common-lisp", + ]) + .write_stdin("(let ((x (compute))) (+ x x))") + .assert() + .failure() + .stderr(predicate::str::contains("would duplicate")); } #[test] @@ -162,6 +175,8 @@ fn cli_allows_inline_let_duplicate_evaluation_when_explicit() { "--path", "0", "--allow-duplicate-evaluation", + "--dialect", + "common-lisp", "--output", "json", ]) @@ -173,34 +188,72 @@ fn cli_allows_inline_let_duplicate_evaluation_when_explicit() { } #[test] -fn cli_plans_inline_let_for_clojure_vector_binding() { +fn cli_plans_inline_let_for_all_known_stdin_dialects() { + for (dialect, input) in [ + ( + "common-lisp", + "(let ((product (* width height))) (+ product margin))", + ), + ( + "emacs-lisp", + "(let ((product (* width height))) (+ product margin))", + ), + ( + "scheme", + "(let ((product (* width height))) (+ product margin))", + ), + ( + "clojure", + "(let [product (* width height)] (+ product margin))", + ), + ( + "janet", + "(let [product (* width height)] (+ product margin))", + ), + ( + "fennel", + "(let [product (* width height)] (+ product margin))", + ), + ] { + let mut cmd = paredit(); + cmd.args([ + "refactor", + "inline-let", + "--dialect", + dialect, + "--path", + "0", + "--output", + "json", + ]) + .write_stdin(input) + .assert() + .success() + .stdout(predicate::str::contains(format!( + "\"dialect\": \"{dialect}\"" + ))) + .stdout(predicate::str::contains("\"binding_name\": \"product\"")) + .stdout(predicate::str::contains("(+ (* width height) margin)")); + } +} + +#[test] +fn cli_plans_inline_let_without_touching_shadowed_lambda_parameter() { let mut cmd = paredit(); cmd.args([ "refactor", "inline-let", - "--dialect", - "clojure", "--path", "0", + "--dialect", + "common-lisp", "--output", "json", ]) - .write_stdin("(let [product (* width height)] (+ product margin))") + .write_stdin("(let ((x 1)) (list x (lambda (x) x)))") .assert() .success() - .stdout(predicate::str::contains("\"dialect\": \"clojure\"")) - .stdout(predicate::str::contains("\"binding_name\": \"product\"")) - .stdout(predicate::str::contains("(+ (* width height) margin)")); -} - -#[test] -fn cli_plans_inline_let_without_touching_shadowed_lambda_parameter() { - let mut cmd = paredit(); - cmd.args(["refactor", "inline-let", "--path", "0", "--output", "json"]) - .write_stdin("(let ((x 1)) (list x (lambda (x) x)))") - .assert() - .success() - .stdout(predicate::str::contains("\"binding_name\": \"x\"")) - .stdout(predicate::str::contains("\"reference_count\": 1")) - .stdout(predicate::str::contains("(list 1 (lambda (x) x))")); + .stdout(predicate::str::contains("\"binding_name\": \"x\"")) + .stdout(predicate::str::contains("\"reference_count\": 1")) + .stdout(predicate::str::contains("(list 1 (lambda (x) x))")); } diff --git a/tests/cli/let_refactor/introduce/mod.rs b/tests/cli/let_refactor/introduce/mod.rs index e25f7dc0..a97256d1 100644 --- a/tests/cli/let_refactor/introduce/mod.rs +++ b/tests/cli/let_refactor/introduce/mod.rs @@ -9,6 +9,7 @@ fn assert_plan_output(args: &[&str], input: &str, checks: &[&str]) { let mut assert = cmd .arg("refactor") .args(args) + .args(["--dialect", "common-lisp"]) .write_stdin(input) .assert() .success(); diff --git a/tests/cli/let_refactor/introduce/plan.rs b/tests/cli/let_refactor/introduce/plan.rs index 0043bd9c..fc613b4b 100644 --- a/tests/cli/let_refactor/introduce/plan.rs +++ b/tests/cli/let_refactor/introduce/plan.rs @@ -14,7 +14,7 @@ fn cli_plans_introduce_let_for_common_lisp() { ], "(defun render () (+ (* width height) margin))", &[ - "\"dialect\": \"unknown\"", + "\"dialect\": \"common-lisp\"", "\"binding_value\": \"(* width height)\"", "\"replacement\": \"(let ((product (* width height))) (+ product margin))\"", "(defun render () (let ((product (* width height))) (+ product margin)))", diff --git a/tests/cli/let_refactor/report.rs b/tests/cli/let_refactor/report.rs index ef91e614..2aa7544e 100644 --- a/tests/cli/let_refactor/report.rs +++ b/tests/cli/let_refactor/report.rs @@ -3,18 +3,25 @@ use super::*; #[test] fn cli_reports_let_inline_safety_for_common_lisp() { let mut cmd = paredit(); - cmd.args(["inspect", "lets", "--output", "json"]) - .write_stdin("(defun render () (let ((product (* width height))) (+ product margin)))") - .assert() - .success() - .stdout(predicate::str::contains("\"let_form_count\": 1")) - .stdout(predicate::str::contains("\"path\": \"0.3\"")) - .stdout(predicate::str::contains("\"binding_style\": \"list-pair\"")) - .stdout(predicate::str::contains("\"name\": \"product\"")) - .stdout(predicate::str::contains("\"reference_count\": 1")) - .stdout(predicate::str::contains( - "\"can_inline_without_duplication\": true", - )); + cmd.args([ + "inspect", + "lets", + "--dialect", + "common-lisp", + "--output", + "json", + ]) + .write_stdin("(defun render () (let ((product (* width height))) (+ product margin)))") + .assert() + .success() + .stdout(predicate::str::contains("\"let_form_count\": 1")) + .stdout(predicate::str::contains("\"path\": \"0.3\"")) + .stdout(predicate::str::contains("\"binding_style\": \"list-pair\"")) + .stdout(predicate::str::contains("\"name\": \"product\"")) + .stdout(predicate::str::contains("\"reference_count\": 1")) + .stdout(predicate::str::contains( + "\"can_inline_without_duplication\": true", + )); } #[test] @@ -83,19 +90,26 @@ fn cli_reports_symbol_macrolet_bindings_without_counting_expansion_reference() { #[test] fn cli_reports_single_binding_symbol_macrolet_supported_by_inline_let() { let mut cmd = paredit(); - cmd.args(["inspect", "lets", "--output", "json"]) - .write_stdin("(symbol-macrolet ((used other)) (list used))") - .assert() - .success() - .stdout(predicate::str::contains("\"form\": \"symbol-macrolet\"")) - .stdout(predicate::str::contains( - "\"inline_supported_by_inline_let\": true", - )) - .stdout(predicate::str::contains("\"name\": \"used\"")) - .stdout(predicate::str::contains("\"reference_count\": 1")) - .stdout(predicate::str::contains( - "\"can_inline_without_duplication\": true", - )); + cmd.args([ + "inspect", + "lets", + "--dialect", + "common-lisp", + "--output", + "json", + ]) + .write_stdin("(symbol-macrolet ((used other)) (list used))") + .assert() + .success() + .stdout(predicate::str::contains("\"form\": \"symbol-macrolet\"")) + .stdout(predicate::str::contains( + "\"inline_supported_by_inline_let\": true", + )) + .stdout(predicate::str::contains("\"name\": \"used\"")) + .stdout(predicate::str::contains("\"reference_count\": 1")) + .stdout(predicate::str::contains( + "\"can_inline_without_duplication\": true", + )); } #[test] diff --git a/tests/cli/refactor_manifest/check.rs b/tests/cli/refactor_manifest/check.rs index 0e974808..0705b31b 100644 --- a/tests/cli/refactor_manifest/check.rs +++ b/tests/cli/refactor_manifest/check.rs @@ -83,6 +83,153 @@ fn cli_checks_refactor_manifest_without_writing_or_diffing() { ); } +#[test] +fn cli_refactor_manifest_round_trips_file_dialect() { + let cases = [ + ( + "common-lisp-character-literal", + "lisp", + "common-lisp", + "(old-name #\\))\n", + ), + ( + "emacs-lisp-character-literal", + "el", + "emacs-lisp", + "(old-name ?\\))\n", + ), + ]; + + for (case_name, extension, expected_dialect, original) in cases { + let dir = fresh_temp_dir(case_name); + let source = dir.join(format!("source.{extension}")); + let manifest_file = dir.join("rename.preview.json"); + fs::write(&source, original).expect("write dialect fixture"); + + let preview_output = paredit() + .args([ + "refactor", + "preview", + "--from", + "old-name", + "--to", + "new-name", + "--mode", + "symbol", + "--fail-on-parse-error", + "--output", + "json", + ]) + .arg(&source) + .assert() + .success() + .get_output() + .stdout + .clone(); + let manifest: serde_json::Value = + serde_json::from_slice(&preview_output).expect("preview output parses"); + assert_eq!( + manifest["files"][0]["dialect"].as_str(), + Some(expected_dialect) + ); + fs::write(&manifest_file, preview_output).expect("write refactor manifest"); + + for command in ["check", "diff"] { + paredit() + .args(["refactor", command]) + .arg("--manifest") + .arg(&manifest_file) + .arg("--root") + .arg(&dir) + .arg("--output") + .arg("json") + .assert() + .success() + .stdout(predicate::str::contains("\"can_apply\": true")); + } + + paredit() + .args(["refactor", "apply"]) + .arg("--manifest") + .arg(&manifest_file) + .arg("--root") + .arg(&dir) + .arg("--write") + .arg("--output") + .arg("json") + .assert() + .success(); + + assert_eq!( + fs::read_to_string(&source).expect("read rewritten dialect fixture"), + original.replace("old-name", "new-name") + ); + } +} + +#[test] +fn cli_refactor_manifest_requires_valid_file_dialect() { + let dir = fresh_temp_dir("refactor manifest dialect validation"); + let source = dir.join("core.lisp"); + fs::write(&source, "(old-name value)\n").expect("write dialect validation fixture"); + let preview_output = paredit() + .args([ + "refactor", "preview", "--from", "old-name", "--to", "new-name", "--mode", "symbol", + "--output", "json", + ]) + .arg(&source) + .assert() + .success() + .get_output() + .stdout + .clone(); + let manifest: serde_json::Value = + serde_json::from_slice(&preview_output).expect("preview output parses"); + + for (case_name, dialect, expected_error) in [ + ( + "missing", + None, + "missing required manifest field files[0].dialect", + ), + ( + "invalid", + Some("not-a-dialect"), + "manifest field files[0].dialect has invalid dialect", + ), + ] { + let mut malformed = manifest.clone(); + let file = malformed["files"][0] + .as_object_mut() + .expect("manifest file entry"); + match dialect { + Some(label) => { + file.insert("dialect".into(), serde_json::Value::String(label.into())); + } + None => { + file.remove("dialect"); + } + } + + let manifest_file = dir.join(format!("{case_name}.preview.json")); + fs::write( + &manifest_file, + serde_json::to_vec_pretty(&malformed).expect("serialize malformed manifest"), + ) + .expect("write malformed manifest"); + + paredit() + .args(["refactor", "check"]) + .arg("--manifest") + .arg(&manifest_file) + .arg("--root") + .arg(&dir) + .assert() + .failure() + .stderr(predicate::str::contains(expected_error)); + } +} + #[test] fn cli_refactor_check_rejects_unexpected_manifest_hash_without_writing() { let dir = fresh_temp_dir("refactor check-manifest-hash"); diff --git a/tests/cli/refactor_manifest/mod.rs b/tests/cli/refactor_manifest/mod.rs index 35ad3bee..90efbe70 100644 --- a/tests/cli/refactor_manifest/mod.rs +++ b/tests/cli/refactor_manifest/mod.rs @@ -53,6 +53,35 @@ fn preview_manifest_out_writes_manifest_and_reports_matching_hash() { assert!(rewritten.contains("(defun new-name (x) x)"), "{rewritten}"); } +#[test] +fn refactor_manifest_round_trip_has_no_dialect_agnostic_parse_calls() { + let sources = [ + ( + "preview", + include_str!("../../../src/presentation/cli/refactor/workflow/preview/build.rs"), + ), + ( + "check", + include_str!("../../../src/presentation/cli/refactor/manifest/check.rs"), + ), + ( + "diff", + include_str!("../../../src/presentation/cli/refactor/workflow/manifest/diff.rs"), + ), + ( + "apply", + include_str!("../../../src/presentation/cli/refactor/workflow/manifest/apply.rs"), + ), + ]; + + for (workflow, source) in sources { + assert!( + !source.contains("SyntaxTree::parse("), + "{workflow} must parse rewritten source with its manifest dialect" + ); + } +} + #[cfg(unix)] #[test] fn preview_manifest_out_refuses_symlink_without_modifying_its_target() { diff --git a/tests/cli/remove_unused_binding/planning.rs b/tests/cli/remove_unused_binding/planning.rs index 354a5444..3c717079 100644 --- a/tests/cli/remove_unused_binding/planning.rs +++ b/tests/cli/remove_unused_binding/planning.rs @@ -6,6 +6,8 @@ fn cli_plans_remove_unused_binding_without_writing() { cmd.args([ "refactor", "remove-unused-binding", + "--dialect", + "common-lisp", "--path", "0.3", "--name", @@ -32,6 +34,8 @@ fn cli_plans_remove_unused_binding_alongside_a_bare_symbol_sibling_binding() { cmd.args([ "refactor", "remove-unused-binding", + "--dialect", + "common-lisp", "--path", "0.3", "--name", @@ -55,6 +59,8 @@ fn cli_plans_remove_all_unused_bindings_without_writing() { cmd.args([ "refactor", "remove-unused-binding", + "--dialect", + "common-lisp", "--path", "0.3", "--all-bindings", @@ -79,6 +85,8 @@ fn cli_plans_remove_all_unused_bindings_with_multiple_body_expressions() { cmd.args([ "refactor", "remove-unused-binding", + "--dialect", + "common-lisp", "--path", "0", "--all-bindings", diff --git a/tests/cli/remove_unused_binding/scope/callables/compiler_macrolet.rs b/tests/cli/remove_unused_binding/scope/callables/compiler_macrolet.rs index a66aa9dd..efc5823f 100644 --- a/tests/cli/remove_unused_binding/scope/callables/compiler_macrolet.rs +++ b/tests/cli/remove_unused_binding/scope/callables/compiler_macrolet.rs @@ -6,6 +6,8 @@ fn cli_plans_remove_unused_compiler_macrolet_without_counting_expander_body_refe cmd.args([ "refactor", "remove-unused-binding", + "--dialect", + "common-lisp", "--path", "0", "--name", @@ -35,6 +37,8 @@ fn cli_plans_remove_unused_cl_compiler_macrolet_without_counting_expander_body_r cmd.args([ "refactor", "remove-unused-binding", + "--dialect", + "common-lisp", "--path", "0", "--name", @@ -63,7 +67,11 @@ fn cli_plans_remove_unused_cl_compiler_macrolet_without_counting_expander_body_r #[test] fn cli_plans_remove_unused_cl_user_compiler_macrolet_without_counting_expander_body_reference() { let mut cmd = paredit(); - cmd.args(["refactor", "remove-unused-binding", + cmd.args([ + "refactor", + "remove-unused-binding", + "--dialect", + "common-lisp", "--path", "0", "--name", @@ -93,6 +101,8 @@ fn cli_rejects_referenced_compiler_macrolet_binding() { cmd.args([ "refactor", "remove-unused-binding", + "--dialect", + "common-lisp", "--path", "0", "--name", diff --git a/tests/cli/remove_unused_binding/scope/callables/local_functions.rs b/tests/cli/remove_unused_binding/scope/callables/local_functions.rs index a7f81859..191e78bf 100644 --- a/tests/cli/remove_unused_binding/scope/callables/local_functions.rs +++ b/tests/cli/remove_unused_binding/scope/callables/local_functions.rs @@ -6,6 +6,8 @@ fn cli_plans_remove_unused_flet_binding_ignoring_definition_body_reference() { cmd.args([ "refactor", "remove-unused-binding", + "--dialect", + "common-lisp", "--path", "0", "--name", @@ -30,6 +32,8 @@ fn cli_rejects_recursive_labels_binding() { cmd.args([ "refactor", "remove-unused-binding", + "--dialect", + "common-lisp", "--path", "0", "--name", diff --git a/tests/cli/remove_unused_binding/scope/callables/macrolet.rs b/tests/cli/remove_unused_binding/scope/callables/macrolet.rs index 0637b536..e6a7f8de 100644 --- a/tests/cli/remove_unused_binding/scope/callables/macrolet.rs +++ b/tests/cli/remove_unused_binding/scope/callables/macrolet.rs @@ -6,6 +6,8 @@ fn cli_plans_remove_unused_macrolet_without_counting_expander_body_reference() { cmd.args([ "refactor", "remove-unused-binding", + "--dialect", + "common-lisp", "--path", "0", "--name", @@ -33,6 +35,8 @@ fn cli_plans_remove_unused_cl_macrolet_without_counting_expander_body_reference( cmd.args([ "refactor", "remove-unused-binding", + "--dialect", + "common-lisp", "--path", "0", "--name", @@ -60,6 +64,8 @@ fn cli_plans_remove_unused_cl_user_macrolet_without_counting_expander_body_refer cmd.args([ "refactor", "remove-unused-binding", + "--dialect", + "common-lisp", "--path", "0", "--name", @@ -89,6 +95,8 @@ fn cli_rejects_referenced_macrolet_binding() { cmd.args([ "refactor", "remove-unused-binding", + "--dialect", + "common-lisp", "--path", "0", "--name", diff --git a/tests/cli/remove_unused_binding/scope/slots.rs b/tests/cli/remove_unused_binding/scope/slots.rs index 22928e2a..d0f88562 100644 --- a/tests/cli/remove_unused_binding/scope/slots.rs +++ b/tests/cli/remove_unused_binding/scope/slots.rs @@ -6,6 +6,8 @@ fn cli_plans_remove_unused_with_slots_without_counting_instance_expression() { cmd.args([ "refactor", "remove-unused-binding", + "--dialect", + "common-lisp", "--path", "0", "--name", @@ -31,6 +33,8 @@ fn cli_plans_remove_unused_with_accessors_without_counting_instance_expression() cmd.args([ "refactor", "remove-unused-binding", + "--dialect", + "common-lisp", "--path", "0", "--name", diff --git a/tests/cli/remove_unused_binding/scope/symbol_macros.rs b/tests/cli/remove_unused_binding/scope/symbol_macros.rs index f50707fb..3ab7de40 100644 --- a/tests/cli/remove_unused_binding/scope/symbol_macros.rs +++ b/tests/cli/remove_unused_binding/scope/symbol_macros.rs @@ -6,6 +6,8 @@ fn cli_plans_remove_unused_symbol_macrolet_without_counting_expansion_reference( cmd.args([ "refactor", "remove-unused-binding", + "--dialect", + "common-lisp", "--path", "0", "--name", @@ -66,6 +68,8 @@ fn cli_plans_remove_unused_cl_user_symbol_macrolet_without_counting_expansion_re cmd.args([ "refactor", "remove-unused-binding", + "--dialect", + "common-lisp", "--path", "0", "--name", @@ -95,6 +99,8 @@ fn cli_rejects_referenced_symbol_macrolet_binding() { cmd.args([ "refactor", "remove-unused-binding", + "--dialect", + "common-lisp", "--path", "0", "--name", diff --git a/tests/cli/remove_unused_binding/scope/variables.rs b/tests/cli/remove_unused_binding/scope/variables.rs index 166cdb3a..45c58561 100644 --- a/tests/cli/remove_unused_binding/scope/variables.rs +++ b/tests/cli/remove_unused_binding/scope/variables.rs @@ -6,6 +6,8 @@ fn cli_plans_remove_unused_binding_ignoring_shadowed_lambda_parameter() { cmd.args([ "refactor", "remove-unused-binding", + "--dialect", + "common-lisp", "--path", "0", "--name", @@ -29,6 +31,8 @@ fn cli_rejects_remove_unused_let_star_binding_used_later() { cmd.args([ "refactor", "remove-unused-binding", + "--dialect", + "common-lisp", "--path", "0", "--name", @@ -48,6 +52,8 @@ fn cli_keeps_let_star_binding_used_by_later_binding_in_all_bindings() { cmd.args([ "refactor", "remove-unused-binding", + "--dialect", + "common-lisp", "--path", "0", "--all-bindings", diff --git a/tests/cli/remove_unused_binding/validation.rs b/tests/cli/remove_unused_binding/validation.rs index 9a5394a3..e2f0e2eb 100644 --- a/tests/cli/remove_unused_binding/validation.rs +++ b/tests/cli/remove_unused_binding/validation.rs @@ -17,6 +17,25 @@ fn cli_requires_file_for_remove_unused_binding_writes() { .stderr(predicate::str::contains("--write requires --file")); } +#[test] +fn cli_rejects_remove_unused_binding_for_unknown_stdin_dialect() { + let mut cmd = paredit(); + cmd.args([ + "refactor", + "remove-unused-binding", + "--path", + "0", + "--name", + "unused", + ]) + .write_stdin("(let ((unused 1)) :ok)") + .assert() + .failure() + .stderr(predicate::str::contains( + "remove-unused-binding does not support dialect unknown", + )); +} + #[test] fn cli_requires_drop_value_permission_for_remove_unused_binding_writes() { let dir = fresh_temp_dir("remove-unused-binding-permission"); @@ -48,6 +67,8 @@ fn cli_rejects_remove_unused_binding_with_references() { cmd.args([ "refactor", "remove-unused-binding", + "--dialect", + "common-lisp", "--path", "0", "--name", @@ -87,6 +108,8 @@ fn cli_rejects_remove_all_unused_bindings_when_none_unused() { cmd.args([ "refactor", "remove-unused-binding", + "--dialect", + "common-lisp", "--path", "0", "--all-bindings", diff --git a/tests/cli/rename/binding/lambda_list/functions.rs b/tests/cli/rename/binding/lambda_list/functions.rs index f8304ce9..013ab353 100644 --- a/tests/cli/rename/binding/lambda_list/functions.rs +++ b/tests/cli/rename/binding/lambda_list/functions.rs @@ -6,6 +6,8 @@ fn cli_plans_lambda_parameter_rename_without_shadow_capture() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", @@ -32,6 +34,8 @@ fn cli_plans_defun_parameter_rename() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", @@ -112,6 +116,8 @@ fn cli_plans_defmethod_specialized_parameter_rename_without_touching_specializer cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", @@ -135,6 +141,8 @@ fn cli_plans_defmethod_specialized_parameter_rename_without_touching_specializer fn cli_plans_defmethod_qualifier_parameter_rename() { let mut cmd = paredit(); cmd.args(["refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", @@ -160,6 +168,8 @@ fn cli_plans_defmethod_qualifier_parameter_rename() { fn cli_plans_cl_defmethod_optional_parameter_rename_without_touching_default_form() { let mut cmd = paredit(); cmd.args(["refactor", "rename-binding", + "--dialect", + "emacs-lisp", "--path", "0", "--from", diff --git a/tests/cli/rename/binding/lambda_list/macros.rs b/tests/cli/rename/binding/lambda_list/macros.rs index 2afaa5bc..29551257 100644 --- a/tests/cli/rename/binding/lambda_list/macros.rs +++ b/tests/cli/rename/binding/lambda_list/macros.rs @@ -85,6 +85,8 @@ fn cli_plans_defmacro_optional_parameter_rename_without_touching_default_form() cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", @@ -110,6 +112,8 @@ fn cli_plans_defmacro_optional_parameter_rename_without_touching_default_form() fn cli_plans_define_setf_expander_environment_parameter_rename() { let mut cmd = paredit(); cmd.args(["refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", @@ -137,6 +141,8 @@ fn cli_plans_define_setf_expander_environment_parameter_rename() { fn cli_plans_define_compiler_macro_environment_parameter_rename() { let mut cmd = paredit(); cmd.args(["refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", diff --git a/tests/cli/rename/binding/lexical.rs b/tests/cli/rename/binding/lexical.rs index 0d303a94..32ac8566 100644 --- a/tests/cli/rename/binding/lexical.rs +++ b/tests/cli/rename/binding/lexical.rs @@ -6,6 +6,8 @@ fn cli_plans_binding_rename_without_shadowed_inner_binding() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0.3", "--from", @@ -33,6 +35,8 @@ fn cli_plans_let_star_binding_rename_through_later_binding_values() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", @@ -114,6 +118,8 @@ fn cli_plans_common_lisp_bare_let_binding_rename_without_touching_later_init() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", @@ -139,6 +145,8 @@ fn cli_plans_common_lisp_bare_let_star_binding_rename_through_later_init() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", @@ -164,6 +172,8 @@ fn cli_plans_outer_let_binding_rename_without_touching_inner_bare_binding() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", diff --git a/tests/cli/rename/binding/macro_semantics/expander_visibility.rs b/tests/cli/rename/binding/macro_semantics/expander_visibility.rs index 8bb9c85c..ded46528 100644 --- a/tests/cli/rename/binding/macro_semantics/expander_visibility.rs +++ b/tests/cli/rename/binding/macro_semantics/expander_visibility.rs @@ -6,6 +6,8 @@ fn cli_plans_outer_binding_rename_through_macrolet_expander_only() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", @@ -32,6 +34,8 @@ fn cli_plans_outer_binding_rename_through_cl_user_macrolet_expander_only() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", @@ -60,6 +64,8 @@ fn cli_plans_outer_binding_rename_through_compiler_macrolet_expander_only() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", @@ -86,6 +92,8 @@ fn cli_plans_outer_binding_rename_through_compiler_macrolet_expander_only() { fn cli_plans_outer_binding_rename_through_cl_user_compiler_macrolet_expander_only() { let mut cmd = paredit(); cmd.args(["refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", diff --git a/tests/cli/rename/binding/macro_semantics/quasiquote.rs b/tests/cli/rename/binding/macro_semantics/quasiquote.rs index 3b29be4d..3beb96c6 100644 --- a/tests/cli/rename/binding/macro_semantics/quasiquote.rs +++ b/tests/cli/rename/binding/macro_semantics/quasiquote.rs @@ -6,6 +6,8 @@ fn cli_plans_binding_rename_inside_quasiquote_preserving_unquote_prefixes() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", @@ -31,6 +33,8 @@ fn cli_plans_binding_rename_only_after_matching_nested_unquote_depth() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", diff --git a/tests/cli/rename/binding/macro_semantics/symbol_macrolet.rs b/tests/cli/rename/binding/macro_semantics/symbol_macrolet.rs index e3d1ee30..b9eaf7e8 100644 --- a/tests/cli/rename/binding/macro_semantics/symbol_macrolet.rs +++ b/tests/cli/rename/binding/macro_semantics/symbol_macrolet.rs @@ -6,6 +6,8 @@ fn cli_plans_symbol_macrolet_binding_rename_without_touching_expansion_reference cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", @@ -58,6 +60,8 @@ fn cli_plans_outer_binding_rename_through_symbol_macrolet_expansion_only() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", diff --git a/tests/cli/rename/binding/special_forms/do_family.rs b/tests/cli/rename/binding/special_forms/do_family.rs index 51d1b945..c29ff07e 100644 --- a/tests/cli/rename/binding/special_forms/do_family.rs +++ b/tests/cli/rename/binding/special_forms/do_family.rs @@ -6,6 +6,8 @@ fn cli_plans_do_binding_rename_across_steps_end_clause_and_body() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", @@ -59,6 +61,8 @@ fn cli_plans_outer_binding_rename_without_touching_do_scope() { fn cli_plans_do_star_binding_rename_across_later_inits_steps_end_clause_and_body() { let mut cmd = paredit(); cmd.args(["refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", diff --git a/tests/cli/rename/binding/special_forms/iteration.rs b/tests/cli/rename/binding/special_forms/iteration.rs index bfe9c23f..bd12da1f 100644 --- a/tests/cli/rename/binding/special_forms/iteration.rs +++ b/tests/cli/rename/binding/special_forms/iteration.rs @@ -6,6 +6,8 @@ fn cli_plans_dolist_iteration_binding_rename_without_touching_source() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", @@ -31,6 +33,8 @@ fn cli_plans_dotimes_iteration_binding_rename_without_touching_count() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", diff --git a/tests/cli/rename/binding/special_forms/loop_clauses.rs b/tests/cli/rename/binding/special_forms/loop_clauses.rs index 7d1fad70..bc43beb9 100644 --- a/tests/cli/rename/binding/special_forms/loop_clauses.rs +++ b/tests/cli/rename/binding/special_forms/loop_clauses.rs @@ -6,6 +6,8 @@ fn cli_plans_loop_for_in_binding_rename_without_touching_source() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", @@ -31,6 +33,8 @@ fn cli_plans_loop_with_binding_rename_without_touching_init() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", @@ -56,6 +60,8 @@ fn cli_plans_loop_destructuring_binding_rename_without_touching_source() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", @@ -81,6 +87,8 @@ fn cli_plans_outer_binding_rename_without_touching_loop_shadow() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", @@ -106,6 +114,8 @@ fn cli_plans_outer_binding_rename_without_touching_loop_destructuring_shadow() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", diff --git a/tests/cli/rename/binding/special_forms/prog_family.rs b/tests/cli/rename/binding/special_forms/prog_family.rs index 84733862..fee2c547 100644 --- a/tests/cli/rename/binding/special_forms/prog_family.rs +++ b/tests/cli/rename/binding/special_forms/prog_family.rs @@ -6,6 +6,8 @@ fn cli_plans_prog_star_binding_rename_across_later_inits_and_body() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", diff --git a/tests/cli/rename/binding/special_forms/slot_binding_forms.rs b/tests/cli/rename/binding/special_forms/slot_binding_forms.rs index b403b739..784e8d63 100644 --- a/tests/cli/rename/binding/special_forms/slot_binding_forms.rs +++ b/tests/cli/rename/binding/special_forms/slot_binding_forms.rs @@ -6,6 +6,8 @@ fn cli_plans_with_slots_binding_rename_preserving_slot_name() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", @@ -59,6 +61,8 @@ fn cli_plans_outer_binding_rename_without_touching_with_slots_shadow() { fn cli_plans_with_accessors_binding_rename_preserving_accessor_name() { let mut cmd = paredit(); cmd.args(["refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", @@ -84,6 +88,8 @@ fn cli_plans_with_accessors_binding_rename_preserving_accessor_name() { fn cli_plans_outer_binding_rename_without_touching_with_accessors_shadow() { let mut cmd = paredit(); cmd.args(["refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", @@ -112,6 +118,8 @@ fn cli_rejects_ambiguous_with_slots_binding_rename() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", diff --git a/tests/cli/rename/binding/write.rs b/tests/cli/rename/binding/write.rs index 349d0d45..67604dfd 100644 --- a/tests/cli/rename/binding/write.rs +++ b/tests/cli/rename/binding/write.rs @@ -109,6 +109,8 @@ fn cli_rejects_rename_binding_write_without_file() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0.3", "--from", @@ -129,6 +131,8 @@ fn cli_rejects_missing_binding_rename_target() { cmd.args([ "refactor", "rename-binding", + "--dialect", + "common-lisp", "--path", "0", "--from", diff --git a/tests/cli/rename_at/failure.rs b/tests/cli/rename_at/failure.rs index be24e6fe..4e63762c 100644 --- a/tests/cli/rename_at/failure.rs +++ b/tests/cli/rename_at/failure.rs @@ -1,29 +1,42 @@ #[test] fn cli_rename_at_rejects_non_common_lisp_dialects_without_writing() { - let input = "(let ((value (lambda () 1))) (list value (value)))\n"; - let dir = fresh_temp_dir("rename-at-scheme-rejected"); - let file = dir.join("input.scm"); - fs::write(&file, input).expect("write non-Common-Lisp rename-at fixture"); + let input = "("; + for (dialect, extension) in [ + ("emacs-lisp", "el"), + ("scheme", "scm"), + ("clojure", "clj"), + ("janet", "janet"), + ("fennel", "fnl"), + ] { + let dir = fresh_temp_dir(&format!("rename-at-{dialect}-rejected")); + let file = dir.join(format!("input.{extension}")); + fs::write(&file, input).expect("write malformed non-Common-Lisp rename-at fixture"); - let mut cmd = paredit(); - cmd.arg("refactor") - .arg("rename-at") - .arg("--file") - .arg(&file) - .arg("--at") - .arg(byte_offset(input, "value (lambda").to_string()) - .arg("--to") - .arg("thunk") - .arg("--dialect") - .arg("scheme") - .arg("--write") - .assert() - .failure(); + let mut cmd = paredit(); + cmd.arg("refactor") + .arg("rename-at") + .arg("--file") + .arg(&file) + .arg("--at") + .arg("0") + .arg("--to") + .arg("thunk") + .arg("--dialect") + .arg(dialect) + .arg("--write") + .assert() + .failure() + .stderr( + predicate::str::contains("rename-at currently supports Common Lisp only") + .and(predicate::str::contains("failed to parse input").not()), + ); - assert_eq!( - fs::read_to_string(file).expect("read unchanged non-Common-Lisp fixture"), - input - ); + assert_eq!( + fs::read_to_string(&file).expect("read unchanged non-Common-Lisp fixture"), + input, + "dialect: {dialect}" + ); + } } #[test] diff --git a/tests/cli/repair_unclosed_lists.rs b/tests/cli/repair_unclosed_lists.rs index c609ed8e..a7131b85 100644 --- a/tests/cli/repair_unclosed_lists.rs +++ b/tests/cli/repair_unclosed_lists.rs @@ -3,7 +3,7 @@ use super::*; #[test] fn repair_unclosed_lists_writes_only_required_closers() { let dir = fresh_temp_dir("repair-unclosed-lists"); - let file = dir.join("source.lisp"); + let file = dir.join("source.clj"); fs::write(&file, "(outer [inner {leaf}").expect("write fixture"); paredit() diff --git a/tests/cli/replace_forms.rs b/tests/cli/replace_forms.rs index f3d52ae4..85104813 100644 --- a/tests/cli/replace_forms.rs +++ b/tests/cli/replace_forms.rs @@ -99,6 +99,8 @@ fn cli_rejects_replace_forms_overlapping_paths() { cmd.args([ "refactor", "replace-forms", + "--dialect", + "common-lisp", "--path", "0", "--path", @@ -118,6 +120,8 @@ fn cli_rejects_replace_forms_shape_mismatch_when_required() { cmd.args([ "refactor", "replace-forms", + "--dialect", + "common-lisp", "--path", "0", "--path", diff --git a/tests/cli/similarity_report.rs b/tests/cli/similarity_report.rs index 91e6c17a..51fd66e9 100644 --- a/tests/cli/similarity_report.rs +++ b/tests/cli/similarity_report.rs @@ -255,6 +255,45 @@ fn cli_dialect_override_includes_unknown_extensions() { })); } +#[test] +fn cli_applies_detected_reader_policy_to_similarity_inputs() { + let dir = fresh_temp_dir("similarity-reader-policy"); + let common_lisp_file = dir.join("common.lisp"); + let emacs_lisp_file = dir.join("emacs.el"); + let invalid_common_lisp_file = dir.join("invalid.lisp"); + fs::write(&common_lisp_file, "(list #\\))\n(list #\\))\n").unwrap(); + fs::write(&emacs_lisp_file, "(list ?\\))\n(list ?\\))\n").unwrap(); + fs::write(&invalid_common_lisp_file, "#?value\n").unwrap(); + + let output = paredit() + .args(["inspect", "similarity"]) + .arg("--threshold=1") + .arg("--min-node-count=2") + .arg("--error-policy=skip") + .arg(&dir) + .output() + .unwrap(); + + assert!(output.status.success()); + let report: Value = serde_json::from_slice(&output.stdout).unwrap(); + assert_eq!(report["summary"]["scanned_files"], 3); + assert_eq!(report["summary"]["processed_files"], 2); + assert_eq!(report["summary"]["skipped_error_files"], 1); + let errors = report["errors"].as_array().unwrap(); + assert_eq!(errors.len(), 1); + assert_eq!( + errors[0]["path"], + invalid_common_lisp_file.display().to_string() + ); + assert_eq!(errors[0]["stage"], "parse"); + assert!( + errors[0]["message"] + .as_str() + .unwrap() + .contains("unsupported reader dispatch") + ); +} + #[test] fn cli_json_contract_reports_options_summary_and_pair_count() { let dir = fresh_temp_dir("similarity-json-contract"); diff --git a/tests/cli/thread_expression/pbt.rs b/tests/cli/thread_expression/pbt.rs index 6c223e0d..f23f9f8b 100644 --- a/tests/cli/thread_expression/pbt.rs +++ b/tests/cli/thread_expression/pbt.rs @@ -49,6 +49,8 @@ fn assert_thread_expression_property( .args([ "refactor", "thread-expression", + "--dialect", + "clojure", "--path", "0", "--style", @@ -92,6 +94,8 @@ fn assert_unthread_expression_property(input: String) -> Result<(), TestCaseErro .args([ "refactor", "unthread-expression", + "--dialect", + "clojure", "--path", "0", "--output", diff --git a/tests/cli/thread_expression/thread.rs b/tests/cli/thread_expression/thread.rs index 32deeb3f..c9c47114 100644 --- a/tests/cli/thread_expression/thread.rs +++ b/tests/cli/thread_expression/thread.rs @@ -6,6 +6,8 @@ fn cli_plans_thread_first_expression_without_writing() { cmd.args([ "refactor", "thread-expression", + "--dialect", + "clojure", "--path", "0", "--style", @@ -59,6 +61,8 @@ fn cli_rejects_thread_expression_write_without_file() { cmd.args([ "refactor", "thread-expression", + "--dialect", + "clojure", "--path", "0", "--style", @@ -77,6 +81,8 @@ fn cli_rejects_already_threaded_expression() { cmd.args([ "refactor", "thread-expression", + "--dialect", + "clojure", "--path", "0", "--style", diff --git a/tests/cli/thread_expression/unthread.rs b/tests/cli/thread_expression/unthread.rs index 55ecbf5d..04c13f53 100644 --- a/tests/cli/thread_expression/unthread.rs +++ b/tests/cli/thread_expression/unthread.rs @@ -6,6 +6,8 @@ fn cli_plans_unthread_first_expression_without_writing() { cmd.args([ "refactor", "unthread-expression", + "--dialect", + "clojure", "--path", "0", "--output", @@ -58,14 +60,21 @@ fn cli_writes_unthread_last_expression_for_clojure_file() { #[test] fn cli_rejects_unthread_unrecognized_operator_without_explicit_confirmation() { let mut cmd = paredit(); - cmd.args(["refactor", "unthread-expression", "--path", "0"]) - .write_stdin("(my-> value step)") - .assert() - .failure() - .stderr(predicate::str::contains( - "not a recognized threading operator", - )) - .stderr(predicate::str::contains("--operator")); + cmd.args([ + "refactor", + "unthread-expression", + "--dialect", + "clojure", + "--path", + "0", + ]) + .write_stdin("(my-> value step)") + .assert() + .failure() + .stderr(predicate::str::contains( + "not a recognized threading operator", + )) + .stderr(predicate::str::contains("--operator")); } #[test] @@ -74,6 +83,8 @@ fn cli_rejects_unthread_custom_operator_without_style() { cmd.args([ "refactor", "unthread-expression", + "--dialect", + "clojure", "--path", "0", "--operator", diff --git a/tests/cli/unwrap_call.rs b/tests/cli/unwrap_call.rs index 2d01c6cc..56d0379a 100644 --- a/tests/cli/unwrap_call.rs +++ b/tests/cli/unwrap_call.rs @@ -6,6 +6,8 @@ fn cli_plans_unwrap_call_without_writing() { cmd.args([ "refactor", "unwrap-call", + "--dialect", + "common-lisp", "--path", "0", "--function", @@ -68,6 +70,8 @@ fn cli_rejects_unwrap_call_function_mismatch() { cmd.args([ "refactor", "unwrap-call", + "--dialect", + "common-lisp", "--path", "0", "--function", @@ -147,6 +151,8 @@ fn assert_unwrap_call_property(input: String) -> Result<(), TestCaseError> { .args([ "refactor", "unwrap-call", + "--dialect", + "common-lisp", "--path", "0", "--function", diff --git a/tests/cli/workspace_report.rs b/tests/cli/workspace_report.rs index 4431531b..3819ffc6 100644 --- a/tests/cli/workspace_report.rs +++ b/tests/cli/workspace_report.rs @@ -78,6 +78,37 @@ fn cli_reports_workspace_inventory_with_include_flags() { .stdout(predicate::str::contains("\"dialect\": \"unknown\"")); } +#[test] +fn cli_applies_detected_reader_policy_and_preserves_unknown_parsing() { + let dir = fresh_temp_dir("workspace report-reader-policy"); + let common_lisp_file = dir.join("common.lisp"); + let emacs_lisp_file = dir.join("emacs.el"); + let invalid_common_lisp_file = dir.join("invalid.lisp"); + let unknown_file = dir.join("generic.txt"); + fs::write(&common_lisp_file, "(defun cl-char () #\\))\n").expect("write common lisp fixture"); + fs::write(&emacs_lisp_file, "(defun el-char () ?\\))\n").expect("write emacs lisp fixture"); + fs::write(&invalid_common_lisp_file, "#?value\n").expect("write invalid common lisp fixture"); + fs::write(&unknown_file, "#?value\n").expect("write unknown fixture"); + + let mut cmd = paredit(); + cmd.args(["inspect", "workspace"]) + .arg("--include-unknown") + .arg("--output") + .arg("json") + .arg(&dir) + .assert() + .success() + .stdout(predicate::str::contains("\"file_count\": 4")) + .stdout(predicate::str::contains("\"parsed_count\": 3")) + .stdout(predicate::str::contains("\"parse_error_count\": 1")) + .stdout(predicate::str::contains("\"definition_count\": 2")) + .stdout(predicate::str::contains( + invalid_common_lisp_file.display().to_string(), + )) + .stdout(predicate::str::contains(unknown_file.display().to_string())) + .stdout(predicate::str::contains("\"dialect\": \"unknown\"")); +} + #[test] fn cli_reports_workspace_inventory_for_binary_unknown_files() { let dir = fresh_temp_dir("workspace report-binary");