Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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<Pair> 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")
}

}
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
Expand Down
60 changes: 50 additions & 10 deletions java-class-name/src/scala/cli/javaclassname/JavaParser.scala
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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) {
Expand All @@ -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"))
}
}
Loading