413 lines
16 KiBLFS
Python
413 lines
16 KiBLFS
Python
import re
|
|
|
|
# Input Python code to translate
|
|
INPUT_PYTHON = open("/root/Tokenizer.py").readlines()
|
|
|
|
|
|
class PythonToScalaTranslator:
|
|
"""Translates Python code to idiomatic Scala."""
|
|
|
|
# Type mappings from Python to Scala
|
|
TYPE_MAP = { # noqa: RUF012
|
|
"int": "Int",
|
|
"str": "String",
|
|
"bool": "Boolean",
|
|
"float": "Double",
|
|
"None": "Unit",
|
|
"list": "List",
|
|
"dict": "Map",
|
|
"List": "List",
|
|
"Optional": "Option",
|
|
}
|
|
|
|
def __init__(self, python_code: str):
|
|
self.python_code = python_code
|
|
self.lines = python_code
|
|
|
|
def snake_to_camel(self, name: str) -> str:
|
|
"""Convert snake_case to camelCase."""
|
|
components = name.split("_")
|
|
return components[0] + "".join(x.title() for x in components[1:])
|
|
|
|
def translate_type(self, py_type: str) -> str:
|
|
"""Translate Python type annotation to Scala type."""
|
|
if not py_type:
|
|
return ""
|
|
py_type = py_type.strip()
|
|
|
|
# Handle None union types (Task | None -> Option[Task])
|
|
if "| None" in py_type:
|
|
inner = py_type.replace("| None", "").strip()
|
|
return f"Option[{self.translate_type(inner)}]"
|
|
if "None |" in py_type:
|
|
inner = py_type.replace("None |", "").strip()
|
|
return f"Option[{self.translate_type(inner)}]"
|
|
|
|
# Handle Optional[T]
|
|
if py_type.startswith("Optional["):
|
|
inner = py_type[9:-1]
|
|
return f"Option[{self.translate_type(inner)}]"
|
|
|
|
# Handle List[T]
|
|
if py_type.startswith("List["):
|
|
inner = py_type[5:-1]
|
|
return f"List[{self.translate_type(inner)}]"
|
|
|
|
# Handle dict[K, V]
|
|
if py_type.startswith("dict["):
|
|
inner = py_type[5:-1]
|
|
parts = inner.split(",")
|
|
if len(parts) == 2:
|
|
k, v = parts[0].strip(), parts[1].strip()
|
|
return f"Map[{self.translate_type(k)}, {self.translate_type(v)}]"
|
|
|
|
# Simple type mapping
|
|
return self.TYPE_MAP.get(py_type, py_type)
|
|
|
|
def get_class_body_range(self, start_idx: int) -> int:
|
|
"""Find where a class body ends (next non-indented class/decorator or EOF)."""
|
|
idx = start_idx + 1
|
|
while idx < len(self.lines):
|
|
line = self.lines[idx]
|
|
stripped = line.strip()
|
|
# If non-empty line at column 0 that starts a new class or decorator
|
|
if stripped and not line[0].isspace():
|
|
if stripped.startswith("class ") or stripped.startswith("@"):
|
|
return idx
|
|
idx += 1
|
|
return idx
|
|
|
|
def translate_enum(self, start_idx: int) -> tuple[list[str], int]:
|
|
"""Translate Python Enum to Scala sealed trait."""
|
|
line = self.lines[start_idx]
|
|
match = re.match(r"class\s+(\w+)\s*\(\s*Enum\s*\)", line)
|
|
enum_name = match.group(1)
|
|
|
|
end_idx = self.get_class_body_range(start_idx)
|
|
values = []
|
|
|
|
for idx in range(start_idx + 1, end_idx):
|
|
stripped = self.lines[idx].strip()
|
|
value_match = re.match(r"(\w+)\s*=\s*(\d+)", stripped)
|
|
if value_match:
|
|
name, val = value_match.groups()
|
|
values.append((name.title(), val))
|
|
|
|
output = [
|
|
f"sealed trait {enum_name} extends Ordered[{enum_name}] {{",
|
|
" def value: Int",
|
|
f" def compare(that: {enum_name}): Int = this.value - that.value",
|
|
"}",
|
|
"",
|
|
f"object {enum_name} {{",
|
|
]
|
|
for name, val in values:
|
|
output.append(f" case object {name} extends {enum_name} {{ val value = {val} }}")
|
|
output.append("}")
|
|
|
|
return output, end_idx
|
|
|
|
def translate_dataclass(self, start_idx: int) -> tuple[list[str], int]:
|
|
"""Translate Python @dataclass to Scala case class."""
|
|
line = self.lines[start_idx]
|
|
match = re.match(r"class\s+(\w+)\s*:", line)
|
|
class_name = match.group(1)
|
|
|
|
end_idx = self.get_class_body_range(start_idx)
|
|
fields = []
|
|
|
|
for idx in range(start_idx + 1, end_idx):
|
|
stripped = self.lines[idx].strip()
|
|
field_match = re.match(r"(\w+)\s*:\s*(\w+)(?:\s*=\s*(.+))?", stripped)
|
|
if field_match:
|
|
name, py_type, default = field_match.groups()
|
|
scala_type = self.translate_type(py_type)
|
|
scala_default = None
|
|
if default:
|
|
scala_default = "false" if default == "False" else ("true" if default == "True" else default)
|
|
fields.append((name, scala_type, scala_default))
|
|
|
|
output = [f"case class {class_name}("]
|
|
for i, (name, typ, default) in enumerate(fields):
|
|
comma = "," if i < len(fields) - 1 else ""
|
|
if default:
|
|
output.append(f" {name}: {typ} = {default}{comma}")
|
|
else:
|
|
output.append(f" {name}: {typ}{comma}")
|
|
output.append(")")
|
|
|
|
return output, end_idx
|
|
|
|
def translate_abc(self, start_idx: int) -> tuple[list[str], int]:
|
|
"""Translate Python ABC to Scala trait."""
|
|
line = self.lines[start_idx]
|
|
match = re.match(r"class\s+(\w+)\s*\(\s*ABC\s*\)", line)
|
|
class_name = match.group(1)
|
|
|
|
end_idx = self.get_class_body_range(start_idx)
|
|
methods = []
|
|
|
|
idx = start_idx + 1
|
|
while idx < end_idx:
|
|
stripped = self.lines[idx].strip()
|
|
if stripped == "@abstractmethod":
|
|
idx += 1
|
|
if idx < end_idx:
|
|
method_line = self.lines[idx].strip()
|
|
method_match = re.match(r"def\s+(\w+)\s*\(\s*self\s*(?:,\s*(.+))?\)\s*(?:->\s*(.+))?:", method_line)
|
|
if method_match:
|
|
name = method_match.group(1)
|
|
params_str = method_match.group(2) or ""
|
|
return_type = method_match.group(3) or "None"
|
|
|
|
params = []
|
|
if params_str:
|
|
for p in params_str.split(","):
|
|
p = p.strip()
|
|
if ":" in p:
|
|
pname, ptype = p.split(":", 1)
|
|
params.append((pname.strip(), ptype.strip()))
|
|
|
|
methods.append((name, params, return_type.strip()))
|
|
idx += 1
|
|
|
|
output = [f"trait {class_name} {{"]
|
|
for name, params, return_type in methods:
|
|
scala_name = self.snake_to_camel(name)
|
|
scala_return = self.translate_type(return_type)
|
|
param_strs = [f"{self.snake_to_camel(pn)}: {self.translate_type(pt)}" for pn, pt in params]
|
|
output.append(f' def {scala_name}({", ".join(param_strs)}): {scala_return}')
|
|
output.append("}")
|
|
|
|
return output, end_idx
|
|
|
|
def translate_regular_class(self, start_idx: int) -> tuple[list[str], int]:
|
|
"""Translate regular Python class to Scala class."""
|
|
line = self.lines[start_idx]
|
|
match = re.match(r"class\s+(\w+)\s*(?:\(\s*(\w+)\s*\))?:", line)
|
|
class_name = match.group(1)
|
|
parent_class = match.group(2)
|
|
|
|
end_idx = self.get_class_body_range(start_idx)
|
|
|
|
# Parse all methods in the class
|
|
init_params = []
|
|
instance_vars = []
|
|
methods = []
|
|
|
|
idx = start_idx + 1
|
|
while idx < end_idx:
|
|
stripped = self.lines[idx].strip()
|
|
|
|
# Parse __init__
|
|
if stripped.startswith("def __init__"):
|
|
init_match = re.match(r"def\s+__init__\s*\(\s*self\s*(?:,\s*(.+))?\)", stripped)
|
|
if init_match:
|
|
params_str = init_match.group(1) or ""
|
|
if params_str:
|
|
for p in params_str.split(","):
|
|
p = p.strip()
|
|
if ":" in p:
|
|
pname, ptype = p.split(":", 1)
|
|
init_params.append((pname.strip(), ptype.strip()))
|
|
|
|
# Parse instance variables in __init__ body
|
|
idx += 1
|
|
while idx < end_idx:
|
|
body_line = self.lines[idx].strip()
|
|
if body_line.startswith("def "):
|
|
break
|
|
var_match = re.match(r"self\.(_?\w+)\s*(?::\s*([^=]+))?\s*=\s*(.+)", body_line)
|
|
if var_match:
|
|
vname, vtype, vinit = var_match.groups()
|
|
instance_vars.append((vname, vtype.strip() if vtype else None, vinit.strip()))
|
|
idx += 1
|
|
continue
|
|
|
|
# Parse regular methods
|
|
if stripped.startswith("def ") and not stripped.startswith("def __"):
|
|
method_match = re.match(r"def\s+(\w+)\s*\(\s*self\s*(?:,\s*(.+))?\)\s*(?:->\s*(.+))?:", stripped)
|
|
if method_match:
|
|
name = method_match.group(1)
|
|
params_str = method_match.group(2) or ""
|
|
return_type = method_match.group(3) or "None"
|
|
|
|
params = []
|
|
if params_str:
|
|
for p in params_str.split(","):
|
|
p = p.strip()
|
|
if ":" in p:
|
|
pname, ptype = p.split(":", 1)
|
|
params.append((pname.strip(), ptype.strip()))
|
|
|
|
# Collect method body
|
|
body_lines = []
|
|
idx += 1
|
|
while idx < end_idx:
|
|
body_line = self.lines[idx]
|
|
body_stripped = body_line.strip()
|
|
if body_stripped.startswith("def "):
|
|
break
|
|
if body_stripped:
|
|
body_lines.append(body_stripped)
|
|
idx += 1
|
|
|
|
methods.append({"name": name, "params": params, "return_type": return_type.strip(), "body": body_lines})
|
|
continue
|
|
|
|
idx += 1
|
|
|
|
# Generate Scala class
|
|
extends_str = f" extends {parent_class}" if parent_class else ""
|
|
|
|
# Constructor params
|
|
output = []
|
|
if init_params:
|
|
param_strs = [f"{self.snake_to_camel(pn)}: {self.translate_type(pt)}" for pn, pt in init_params]
|
|
output.append(f'class {class_name}({", ".join(param_strs)}){extends_str} {{')
|
|
else:
|
|
output.append(f"class {class_name}{extends_str} {{")
|
|
|
|
# Instance variables (excluding those that just store constructor params)
|
|
init_param_names = {p[0].lstrip("_") for p in init_params}
|
|
for vname, vtype, vinit in instance_vars:
|
|
clean_name = vname.lstrip("_")
|
|
if clean_name in init_param_names:
|
|
continue
|
|
scala_name = self.snake_to_camel(clean_name)
|
|
scala_type = self.translate_type(vtype) if vtype else "Map[Int, Task]"
|
|
scala_init = "Map.empty" if vinit == "{}" else vinit
|
|
output.append(f" private var {scala_name}: {scala_type} = {scala_init}")
|
|
|
|
if instance_vars and any(vname.lstrip("_") not in init_param_names for vname, _, _ in instance_vars):
|
|
output.append("")
|
|
|
|
# Methods
|
|
for method in methods:
|
|
method_output = self.translate_method(method, class_name)
|
|
output.extend(method_output)
|
|
|
|
output.append("}")
|
|
|
|
return output, end_idx
|
|
|
|
def translate_method(self, method: dict, class_name: str) -> list[str]:
|
|
"""Translate a method to Scala."""
|
|
name = method["name"]
|
|
params = method["params"]
|
|
return_type = method["return_type"]
|
|
|
|
scala_name = self.snake_to_camel(name)
|
|
scala_return = self.translate_type(return_type)
|
|
param_strs = [f"{self.snake_to_camel(pn)}: {self.translate_type(pt)}" for pn, pt in params]
|
|
params_str = ", ".join(param_strs)
|
|
|
|
# Generate idiomatic Scala body based on method semantics
|
|
body = self.translate_method_body(name, method["body"])
|
|
|
|
output = []
|
|
if len(body) == 1:
|
|
output.append(f" def {scala_name}({params_str}): {scala_return} = {body[0]}")
|
|
else:
|
|
output.append(f" def {scala_name}({params_str}): {scala_return} = {{")
|
|
for line in body:
|
|
output.append(f" {line}")
|
|
output.append(" }")
|
|
output.append("")
|
|
|
|
return output
|
|
|
|
def translate_method_body(self, method_name: str, body_lines: list[str]) -> list[str]:
|
|
"""Translate method body to idiomatic Scala."""
|
|
|
|
# Analyze body to determine translation strategy
|
|
body_text = "\n".join(body_lines)
|
|
|
|
# save: self._tasks[task.id] = task -> Map update
|
|
if method_name == "save" and "[task.id]" in body_text:
|
|
return ["tasks = tasks + (task.id -> task)"]
|
|
|
|
# find_by_id: return self._tasks.get(task_id) -> Map.get returns Option
|
|
if method_name == "find_by_id" and ".get(" in body_text:
|
|
return ["tasks.get(taskId)"]
|
|
|
|
# find_all: return list(self._tasks.values()) -> .values.toList
|
|
if method_name == "find_all" and "values()" in body_text:
|
|
return ["tasks.values.toList"]
|
|
|
|
# complete_task: null check pattern -> Option pattern matching
|
|
if method_name == "complete_task" and "is None" in body_text:
|
|
return [
|
|
"repository.findById(taskId) match {",
|
|
" case Some(task) =>",
|
|
" repository.save(task.copy(completed = true))",
|
|
" true",
|
|
" case None => false",
|
|
"}",
|
|
]
|
|
|
|
# get_pending_by_priority: list comprehension + sorted -> filter + sortBy
|
|
if method_name == "get_pending_by_priority" and "sorted(" in body_text:
|
|
return ["repository.findAll", " .filterNot(_.completed)", " .sortBy(_.priority)(Ordering[Priority].reverse)"]
|
|
|
|
# Default: basic translation
|
|
return ["???"]
|
|
|
|
def translate(self) -> str:
|
|
"""Main translation entry point."""
|
|
output = []
|
|
idx = 0
|
|
|
|
while idx < len(self.lines):
|
|
line = self.lines[idx].strip()
|
|
|
|
# Skip imports and empty lines
|
|
if not line or line.startswith("from ") or line.startswith("import "):
|
|
idx += 1
|
|
continue
|
|
|
|
# @dataclass -> case class
|
|
if line == "@dataclass":
|
|
idx += 1
|
|
if idx < len(self.lines):
|
|
result, idx = self.translate_dataclass(idx)
|
|
output.extend(result)
|
|
output.append("")
|
|
continue
|
|
|
|
# Enum -> sealed trait
|
|
if re.match(r"class\s+\w+\s*\(\s*Enum\s*\)", line):
|
|
result, idx = self.translate_enum(idx)
|
|
output.extend(result)
|
|
output.append("")
|
|
continue
|
|
|
|
# ABC -> trait
|
|
if re.match(r"class\s+\w+\s*\(\s*ABC\s*\)", line):
|
|
result, idx = self.translate_abc(idx)
|
|
output.extend(result)
|
|
output.append("")
|
|
continue
|
|
|
|
# Regular class
|
|
if re.match(r"class\s+\w+", line):
|
|
result, idx = self.translate_regular_class(idx)
|
|
output.extend(result)
|
|
output.append("")
|
|
continue
|
|
|
|
idx += 1
|
|
|
|
return "\n".join(output).strip()
|
|
|
|
|
|
# Run translation
|
|
translator = PythonToScalaTranslator(INPUT_PYTHON)
|
|
scala_code = translator.translate()
|
|
|
|
# Write output
|
|
with open("/root/Tokenizer.scala", "w") as f:
|
|
f.write(scala_code)
|
|
|
|
print("Generated testcase.scala")
|