Skip to content
Merged
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
58 changes: 28 additions & 30 deletions codeflash/cli_cmds/init_java.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,8 @@
from codeflash.code_utils.shell_utils import get_shell_rc_path, is_powershell
from codeflash.telemetry.posthog_cf import ph

_MAVEN_NS = {"m": "http://maven.apache.org/POM/4.0.0"}


class JavaBuildTool(Enum):
"""Java build tools."""
Expand Down Expand Up @@ -75,23 +77,15 @@ def detect_java_build_tool(project_root: Path) -> JavaBuildTool:
def detect_java_source_root(project_root: Path) -> str:
"""Detect the Java source root directory."""
# Standard Maven/Gradle layout
standard_src = project_root / "src" / "main" / "java"
if standard_src.is_dir():
if (project_root / "src" / "main" / "java").is_dir():
return "src/main/java"

# Try to detect from pom.xml
pom_path = project_root / "pom.xml"
if pom_path.exists():
try:
tree = ET.parse(pom_path)
root = tree.getroot()
# Handle Maven namespace
ns = {"m": "http://maven.apache.org/POM/4.0.0"}
source_dir = root.find(".//m:sourceDirectory", ns)
if source_dir is not None and source_dir.text:
return source_dir.text
except ET.ParseError:
pass
root = _get_pom_root_cached(project_root)
if root is not None:
source_dir = root.find(".//m:sourceDirectory", _MAVEN_NS)
if source_dir is not None and source_dir.text:
return source_dir.text

# Fallback to src directory
if (project_root / "src").is_dir():
Expand All @@ -103,22 +97,15 @@ def detect_java_source_root(project_root: Path) -> str:
def detect_java_test_root(project_root: Path) -> str:
"""Detect the Java test root directory."""
# Standard Maven/Gradle layout
standard_test = project_root / "src" / "test" / "java"
if standard_test.is_dir():
if (project_root / "src" / "test" / "java").is_dir():
return "src/test/java"

# Try to detect from pom.xml
pom_path = project_root / "pom.xml"
if pom_path.exists():
try:
tree = ET.parse(pom_path)
root = tree.getroot()
ns = {"m": "http://maven.apache.org/POM/4.0.0"}
test_source_dir = root.find(".//m:testSourceDirectory", ns)
if test_source_dir is not None and test_source_dir.text:
return test_source_dir.text
except ET.ParseError:
pass
root = _get_pom_root_cached(project_root)
if root is not None:
test_source_dir = root.find(".//m:testSourceDirectory", _MAVEN_NS)
if test_source_dir is not None and test_source_dir.text:
return test_source_dir.text

# Fallback patterns
if (project_root / "test").is_dir():
Expand Down Expand Up @@ -461,10 +448,9 @@ def configure_java_project(setup_info: JavaSetupInfo) -> bool:
test_root = setup_info.test_root_override or detect_java_test_root(curdir)

# Only include non-default values
defaults = {"module-root": "src/main/java", "tests-root": "src/test/java"}
if source_root != defaults["module-root"]:
if source_root != "src/main/java":
config["module-root"] = source_root
if test_root != defaults["tests-root"]:
if test_root != "src/test/java":
config["tests-root"] = test_root

if setup_info.formatter_override is not None and setup_info.formatter_override != ["disabled"]:
Expand Down Expand Up @@ -539,6 +525,18 @@ def get_java_test_command(build_tool: JavaBuildTool) -> str:
return "mvn test"


@lru_cache(maxsize=8)
def _get_pom_root_cached(project_root: Path) -> Union[ET.Element, None]:
pom_path = project_root / "pom.xml"
if not pom_path.exists():
return None
try:
tree = ET.parse(pom_path)
return tree.getroot()
except ET.ParseError:
return None


formatter_warning_shown = False

_SPOTLESS_COMMANDS = {
Expand Down
Loading