diff --git a/src/main/java/org/openrewrite/java/testing/assertj/CollapseConsecutiveAssertThatStatements.java b/src/main/java/org/openrewrite/java/testing/assertj/CollapseConsecutiveAssertThatStatements.java index c969ad597..64089d13d 100644 --- a/src/main/java/org/openrewrite/java/testing/assertj/CollapseConsecutiveAssertThatStatements.java +++ b/src/main/java/org/openrewrite/java/testing/assertj/CollapseConsecutiveAssertThatStatements.java @@ -161,12 +161,13 @@ private J.MethodInvocation getCollapsedAssertThat(List consecutiveAss J.MethodInvocation assertThat = (J.MethodInvocation) assertion.getSelect(); assert assertThat != null; J.MethodInvocation newSelect = collapsed == null ? assertThat : collapsed; + List chainedComments = collapsed == null ? emptyList() : st.getPrefix().getComments(); collapsed = assertion.getPadding().withSelect(JRightPadded .build((Expression) newSelect.withPrefix(Space.EMPTY)) - .withAfter(Space.build(chainedIndent, ListUtils.map(st.getPrefix().getComments(), c -> c.withSuffix(chainedIndent))))); + .withAfter(Space.build(chainedIndent, ListUtils.map(chainedComments, c -> c.withSuffix(chainedIndent))))); } - return requireNonNull(collapsed).withPrefix(originalPrefix.withComments(emptyList())); + return requireNonNull(collapsed).withPrefix(originalPrefix); } }); } diff --git a/src/test/java/org/openrewrite/java/testing/assertj/CollapseConsecutiveAssertThatStatementsTest.java b/src/test/java/org/openrewrite/java/testing/assertj/CollapseConsecutiveAssertThatStatementsTest.java index 8a9af0ff9..2b822ba6e 100644 --- a/src/test/java/org/openrewrite/java/testing/assertj/CollapseConsecutiveAssertThatStatementsTest.java +++ b/src/test/java/org/openrewrite/java/testing/assertj/CollapseConsecutiveAssertThatStatementsTest.java @@ -108,8 +108,8 @@ void test() { class MyTest { void test() { List listA = Arrays.asList("a", "b", "c"); + // Comment nor whitespace below duplicated assertThat(listA) - // Comment nor whitespace below duplicated .isNotNull() .hasSize(3) .containsExactly("a", "b", "c"); @@ -444,8 +444,8 @@ void test() { class MyTest { void test() { List listA = Arrays.asList("a", "b", "c"); + // Check not null assertThat(listA) - // Check not null .isNotNull() // Check size is 3 .hasSize(3) @@ -457,4 +457,35 @@ void test() { ) ); } + + @Test + void preservesSectionCommentBeforeCollapsedChain() { + rewriteRun( + java( + """ + import static org.assertj.core.api.Assertions.assertThat; + + class Test { + void verify(String value) { + // then + assertThat(value).isNotNull(); + assertThat(value).isEqualTo("expected"); + } + } + """, + """ + import static org.assertj.core.api.Assertions.assertThat; + + class Test { + void verify(String value) { + // then + assertThat(value) + .isNotNull() + .isEqualTo("expected"); + } + } + """ + ) + ); + } }