Files
SkillCompiler/data/skills-bench/tasks/python-scala-translation/environment/translate.py
T
2026-09-04 14:58:42 +08:00

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")