diff --git a/java-class-name-tests/src/test/scala/cli/javaclassname/JavaClassNameTests.scala b/java-class-name-tests/src/test/scala/cli/javaclassname/JavaClassNameTests.scala index 5843f21..fdb67d8 100644 --- a/java-class-name-tests/src/test/scala/cli/javaclassname/JavaClassNameTests.scala +++ b/java-class-name-tests/src/test/scala/cli/javaclassname/JavaClassNameTests.scala @@ -216,4 +216,43 @@ class JavaClassNameTests extends munit.FunSuite { assertEquals(extractClassName("Main.java", content), "Main") } + // Java 21+ compact source files (JEP 445/463/512) have no + // top-level type, so `javac` names the implicit class after the source file. + test("unnamed class with top-level main") { + val content = + """String greeting = "hi"; + | + |void main() { + | System.out.println(greeting); + |} + |""".stripMargin + assertEquals(extractClassName("Unnamed.java", content), "Unnamed") + } + + test("unnamed class with a helper class and a record after main") { + val content = + """import java.util.List; + | + |void main() { + | System.out.println(new Helper().greet(List.of(new Pair(1, "a")))); + |} + | + |class Helper { + | String greet(List ps) { return "hi " + ps; } + |} + | + |record Pair(int n, String s) {} + |""".stripMargin + assertEquals(extractClassName("WithHelpers.java", content), "WithHelpers") + } + + test("unnamed class with a field before main") { + val content = + """static final int N = 1; + |String greeting = "hi"; + |void main() { System.out.println(greeting + N); } + |""".stripMargin + assertEquals(extractClassName("FieldFirst.java", content), "FieldFirst") + } + } diff --git a/java-class-name/src/scala/cli/javaclassname/JavaClassName.scala b/java-class-name/src/scala/cli/javaclassname/JavaClassName.scala index 943701d..aa2d0c9 100644 --- a/java-class-name/src/scala/cli/javaclassname/JavaClassName.scala +++ b/java-class-name/src/scala/cli/javaclassname/JavaClassName.scala @@ -25,7 +25,7 @@ object JavaClassName { sys.exit(1) } val content = Files.readAllBytes(p) - val classNameOpt = JavaParser.parseRootPublicClassName(content) + val classNameOpt = JavaParser.rootClassName(content, p.getFileName.toString) for (className <- classNameOpt) println(className) } diff --git a/java-class-name/src/scala/cli/javaclassname/JavaParser.scala b/java-class-name/src/scala/cli/javaclassname/JavaParser.scala index 1efe927..78dd4e0 100644 --- a/java-class-name/src/scala/cli/javaclassname/JavaParser.scala +++ b/java-class-name/src/scala/cli/javaclassname/JavaParser.scala @@ -23,6 +23,24 @@ object JavaParser { extends OutlineJavaParser(source) { override def ObjectTpt(): untpd.Tree = javaLangDot(tpnme.Object) + /** Set when the source is a Java 21+ compact source file (JEP 512, formerly "unnamed classes"): + * top-level fields / methods with no enclosing type, which `javac` wraps in an implicit class + * named after the source file. + */ + var isCompactUnit: Boolean = false + + /** Type bodies are skipped by `typeBody` (and our `enumDecl`), so the stock parser only calls + * `termDecl` for top-level members of compact source files. The stock implementation types + * `void` via `defn.UnitType` and crashes with an NPE, and `compilationUnit` returns + * `EmptyTree` for compact units anyway. We don't need the members, so we flag the unit as + * compact and skip to the end of the source. + */ + override def termDecl(start: Int, mods: untpd.Modifiers, parentToken: Int): List[untpd.Tree] = { + isCompactUnit = true + while in.token != JavaTokens.EOF do in.nextToken() + List(untpd.EmptyTree) // non-empty, so `compilationUnit` treats the unit as compact + } + /** Primitive types show up in record headers (and method signatures), e.g. `record R(int a)`. * The stock implementation resolves them via `defn.IntType` & co, which crashes with an NPE * without initialized definitions. We only need class names, so any untyped placeholder tree @@ -63,12 +81,23 @@ object JavaParser { ): untpd.Template = super.makeTemplate(parents, stats, tparams, needsDummyConstr = false) } - private def parseOutline(byteContent: Array[Byte]): untpd.Tree = { + private enum Outline { + case Types(stats: List[untpd.Tree]) + case Compact + } + + private def parseOutline(byteContent: Array[Byte]): Outline = { given Context = ContextBase().initialCtx.fresh val virtualFile = VirtualFile("placeholder.java", byteContent) val sourceFile = SourceFile(virtualFile, Codec.UTF8) val outlineParser = UntypedOutlineJavaParser(sourceFile) - outlineParser.parse() + val tree = outlineParser.parse() + if outlineParser.isCompactUnit then Outline.Compact + else + Outline.Types(tree match { + case pd: Trees.PackageDef[_] => pd.stats + case _ => Nil + }) } extension (mdef: untpd.DefTree) { @@ -82,13 +111,24 @@ object JavaParser { mdef.mods.privateWithin.isEmpty && !mdef.mods.isOneOf(Flags.Private | Flags.Protected) } + private def publicRootTypeName(stats: List[untpd.Tree]): Option[String] = + stats.collectFirst { + case mdef: ModuleDef if mdef.isPublic => mdef.name.toString + } + + /** The name of the first public top-level type declared in the source, if any. */ def parseRootPublicClassName(byteContent: Array[Byte]): Option[String] = - Option(parseOutline(byteContent)) - .flatMap { - case pd: Trees.PackageDef[_] => Some(pd.stats) - case _ => None - } - .flatMap(_.collectFirst { - case mdef: ModuleDef if mdef.isPublic => mdef.name.toString - }) + parseOutline(byteContent) match { + case Outline.Types(stats) => publicRootTypeName(stats) + case Outline.Compact => None + } + + /** The class name `javac` would produce a class file for: the first public top-level type, or for + * a Java 21+ compact source file (top-level `main` & co, JEP 512) the source file name stem. + */ + def rootClassName(byteContent: Array[Byte], sourceFileName: String): Option[String] = + parseOutline(byteContent) match { + case Outline.Types(stats) => publicRootTypeName(stats) + case Outline.Compact => Some(sourceFileName.stripSuffix(".java")) + } }