neuro-symbolic-coder-13M / math_encoder.py
DexopT's picture
Upload 13 files
5337457 verified
Raw
History Blame Contribute Delete
25.6 kB
# file: math_encoder.py
"""Comprehensive Python-to-mathematical-specification converter.
Converts any Python source code to a structured mathematical representation
(Dict / JSON) via recursive AST traversal. Handles all major Python AST node
types: control flow, data structures, assignments, expressions, classes,
imports, and more.
Typical usage::
converter = CodeToMathConverter()
math_spec = converter.convert("def f(x): return x + 1")
# math_spec is a nested dict
"""
import ast
import json
from typing import Any, Dict, List, Optional, Union
class CodeToMathConverter:
"""Converts Python source code to a mathematical/logical specification.
Traverses the Python AST and produces a structured dictionary
representation that captures the semantic structure of the code
in a language-neutral mathematical form suitable for downstream
processing (tokenisation, embedding, training pipelines, etc.).
Attributes:
ops_map: Mapping from AST binary/augmented-assignment operators to
normalised string names.
comp_ops_map: Mapping from AST comparison operators to normalised
string names.
bool_ops_map: Mapping from boolean operator types to 'AND' / 'OR'.
unary_ops_map: Mapping from unary operator types to 'NOT' / 'NEG' /
'POS' / 'INVERT'.
"""
# ------------------------------------------------------------------
# Initialisation
# ------------------------------------------------------------------
def __init__(self) -> None:
# Binary and augmented-assignment operators
self.ops_map: Dict[type, str] = {
ast.Add: "ADD",
ast.Sub: "SUB",
ast.Mult: "MUL",
ast.Div: "DIV",
ast.Mod: "MOD",
ast.Pow: "POW",
ast.FloorDiv: "FLOORDIV",
ast.LShift: "LSHIFT",
ast.RShift: "RSHIFT",
ast.BitOr: "BITOR",
ast.BitAnd: "BITAND",
ast.BitXor: "BITXOR",
ast.MatMult: "MATMUL",
}
# Comparison operators (original name preserved for compat)
self.comp_ops_map: Dict[type, str] = {
ast.Eq: "EQ",
ast.NotEq: "NEQ",
ast.Lt: "LT",
ast.LtE: "LTE",
ast.Gt: "GT",
ast.GtE: "GTE",
ast.Is: "IS",
ast.IsNot: "ISNOT",
ast.In: "IN",
ast.NotIn: "NOTIN",
}
# Boolean operators (and / or)
self.bool_ops_map: Dict[type, str] = {
ast.And: "AND",
ast.Or: "OR",
}
# Unary operators
self.unary_ops_map: Dict[type, str] = {
ast.Not: "NOT",
ast.USub: "NEG",
ast.UAdd: "POS",
ast.Invert: "INVERT",
}
# ------------------------------------------------------------------
# Public API
# ------------------------------------------------------------------
def convert(self, code: str) -> Dict[str, Any]:
"""Convert Python source code to a mathematical specification.
Args:
code: A string containing valid Python source code.
Returns:
A nested dictionary representing the mathematical structure
of the code. On parse failure returns
``{"error": "invalid_syntax"}``.
"""
try:
tree = ast.parse(code)
return self._traverse(tree)
except SyntaxError:
return {"error": "invalid_syntax"}
# ------------------------------------------------------------------
# Recursive traversal
# ------------------------------------------------------------------
def _traverse(self, node: Any) -> Any:
"""Recursively traverse an AST node and convert to math representation.
Handles individual AST nodes, lists of AST nodes, and ``None``
values. Unknown node types are recorded with their Python class
name and source line number so callers can diagnose coverage gaps.
Args:
node: An AST node, a ``list`` of AST nodes, or ``None``.
Returns:
A ``dict``, a ``list`` of dicts, or ``None`` for ``None`` input.
"""
# --- Sentinel / collection values ---
if node is None:
return None
if isinstance(node, list):
return [self._traverse(item) for item in node]
# ================================================================
# Module
# ================================================================
if isinstance(node, ast.Module):
return {
"type": "module",
"body": self._traverse(node.body),
}
# ================================================================
# Control Flow
# ================================================================
elif isinstance(node, ast.If):
return {
"type": "if",
"test": self._traverse(node.test),
"body": self._traverse(node.body),
"orelse": self._traverse(node.orelse),
}
elif isinstance(node, ast.For):
return {
"type": "for",
"target": self._traverse(node.target),
"iter": self._traverse(node.iter),
"body": self._traverse(node.body),
"orelse": self._traverse(node.orelse) if node.orelse else [],
}
elif isinstance(node, ast.While):
return {
"type": "while",
"test": self._traverse(node.test),
"body": self._traverse(node.body),
"orelse": self._traverse(node.orelse) if node.orelse else [],
}
elif isinstance(node, ast.Break):
return {"type": "break"}
elif isinstance(node, ast.Continue):
return {"type": "continue"}
elif isinstance(node, ast.Try):
return {
"type": "try",
"body": self._traverse(node.body),
"handlers": self._traverse(node.handlers),
"orelse": self._traverse(node.orelse) if node.orelse else [],
"finalbody": self._traverse(node.finalbody),
}
elif isinstance(node, ast.ExceptHandler):
return {
"type": "except",
"exc_type": self._traverse(node.type) if node.type else None,
"name": node.name,
"body": self._traverse(node.body),
}
elif isinstance(node, ast.Raise):
return {
"type": "raise",
"exc": self._traverse(node.exc) if node.exc else None,
"cause": self._traverse(node.cause) if node.cause else None,
}
elif isinstance(node, ast.With):
return {
"type": "with",
"items": self._traverse(node.items),
"body": self._traverse(node.body),
}
elif isinstance(node, ast.withitem):
return {
"type": "withitem",
"context_expr": self._traverse(node.context_expr),
"optional_vars": self._traverse(node.optional_vars),
}
# ================================================================
# Data Structures
# ================================================================
elif isinstance(node, ast.List):
return {
"type": "list",
"elements": self._traverse(node.elts),
}
elif isinstance(node, ast.Dict):
pairs: List[Dict[str, Any]] = []
for k, v in zip(node.keys, node.values):
pairs.append({
"key": self._traverse(k),
"value": self._traverse(v),
})
return {"type": "dict", "pairs": pairs}
elif isinstance(node, ast.Tuple):
return {
"type": "tuple",
"elements": self._traverse(node.elts),
}
elif isinstance(node, ast.Set):
return {
"type": "set",
"elements": self._traverse(node.elts),
}
elif isinstance(node, ast.ListComp):
return {
"type": "list_comprehension",
"elt": self._traverse(node.elt),
"generators": self._traverse(node.generators),
}
elif isinstance(node, ast.SetComp):
return {
"type": "set_comprehension",
"elt": self._traverse(node.elt),
"generators": self._traverse(node.generators),
}
elif isinstance(node, ast.DictComp):
return {
"type": "dict_comprehension",
"key": self._traverse(node.key),
"value": self._traverse(node.value),
"generators": self._traverse(node.generators),
}
elif isinstance(node, ast.GeneratorExp):
return {
"type": "generator_expression",
"elt": self._traverse(node.elt),
"generators": self._traverse(node.generators),
}
elif isinstance(node, ast.comprehension):
return {
"type": "comp",
"target": self._traverse(node.target),
"iter": self._traverse(node.iter),
"ifs": self._traverse(node.ifs),
}
# ================================================================
# Assignments
# ================================================================
elif isinstance(node, ast.Assign):
return {
"type": "assign",
"targets": self._traverse(node.targets),
"value": self._traverse(node.value),
}
elif isinstance(node, ast.AugAssign):
return {
"type": "augassign",
"target": self._traverse(node.target),
"op": self.ops_map.get(type(node.op), "UNKNOWN_OP"),
"value": self._traverse(node.value),
}
elif isinstance(node, ast.AnnAssign):
return {
"type": "annassign",
"target": self._traverse(node.target),
"annotation": self._traverse(node.annotation),
"value": self._traverse(node.value) if node.value else None,
}
# ================================================================
# Expressions
# ================================================================
elif isinstance(node, ast.Call):
keywords: List[Dict[str, Any]] = []
for kw in node.keywords:
keywords.append({
"arg": kw.arg,
"value": self._traverse(kw.value),
})
return {
"type": "call",
"func": self._traverse(node.func),
"args": self._traverse(node.args),
"keywords": keywords,
}
elif isinstance(node, ast.Attribute):
return {
"type": "attribute",
"value": self._traverse(node.value),
"attr": node.attr,
}
elif isinstance(node, ast.Subscript):
return {
"type": "subscript",
"value": self._traverse(node.value),
"slice": self._traverse(node.slice),
}
elif isinstance(node, ast.Slice):
return {
"type": "slice",
"lower": self._traverse(node.lower) if node.lower else None,
"upper": self._traverse(node.upper) if node.upper else None,
"step": self._traverse(node.step) if node.step else None,
}
elif isinstance(node, ast.BoolOp):
return {
"type": "bool_op",
"op": self.bool_ops_map.get(type(node.op), "UNKNOWN_BOOLOP"),
"values": self._traverse(node.values),
}
elif isinstance(node, ast.UnaryOp):
return {
"type": "unary_op",
"op": self.unary_ops_map.get(type(node.op), "UNKNOWN_UNARYOP"),
"operand": self._traverse(node.operand),
}
elif isinstance(node, ast.IfExp):
return {
"type": "if_exp",
"test": self._traverse(node.test),
"body": self._traverse(node.body),
"orelse": self._traverse(node.orelse),
}
elif isinstance(node, ast.Lambda):
return {
"type": "lambda",
"args": self._traverse(node.args),
"body": self._traverse(node.body),
}
elif isinstance(node, ast.arguments):
# Shared by FunctionDef and Lambda
kw_defaults: List[Any] = []
for d in node.kw_defaults:
kw_defaults.append(self._traverse(d))
defaults: List[Any] = []
for d in node.defaults:
defaults.append(self._traverse(d))
return {
"type": "arguments",
"args": [arg.arg for arg in node.args],
"defaults": defaults,
"vararg": node.vararg.arg if node.vararg else None,
"kwarg": node.kwarg.arg if node.kwarg else None,
"kwonlyargs": [a.arg for a in node.kwonlyargs],
"kw_defaults": kw_defaults,
}
elif isinstance(node, ast.Starred):
return {
"type": "starred",
"value": self._traverse(node.value),
}
elif isinstance(node, ast.FormattedValue):
return {
"type": "formatted_value",
"value": self._traverse(node.value),
"conversion": node.conversion,
"format_spec": (
self._traverse(node.format_spec)
if node.format_spec
else None
),
}
elif isinstance(node, ast.JoinedStr):
return {
"type": "fstring",
"values": self._traverse(node.values),
}
# ================================================================
# Other
# ================================================================
elif isinstance(node, ast.ClassDef):
return {
"type": "class",
"name": node.name,
"bases": self._traverse(node.bases),
"keywords": self._traverse(node.keywords),
"body": self._traverse(node.body),
"decorator_list": self._traverse(node.decorator_list),
}
elif isinstance(node, ast.FunctionDef):
return {
"type": "function",
"name": node.name,
"args": self._traverse(node.args),
"returns": (
self._traverse(node.returns)
if node.returns
else self._infer_return_type(node)
),
"body": self._traverse(node.body),
"decorator_list": self._traverse(node.decorator_list),
}
elif isinstance(node, ast.AsyncFunctionDef):
return {
"type": "async_function",
"name": node.name,
"args": self._traverse(node.args),
"returns": (
self._traverse(node.returns)
if node.returns
else self._infer_return_type(node)
),
"body": self._traverse(node.body),
"decorator_list": self._traverse(node.decorator_list),
}
elif isinstance(node, ast.Return):
return {
"type": "return",
"value": self._traverse(node.value),
}
elif isinstance(node, ast.Yield):
return {
"type": "yield",
"value": self._traverse(node.value) if node.value else None,
}
elif isinstance(node, ast.YieldFrom):
return {
"type": "yield_from",
"value": self._traverse(node.value),
}
elif isinstance(node, ast.Import):
names: List[Dict[str, Optional[str]]] = []
for alias in node.names:
names.append({"name": alias.name, "asname": alias.asname})
return {"type": "import", "names": names}
elif isinstance(node, ast.ImportFrom):
names: List[Dict[str, Optional[str]]] = []
for alias in node.names:
names.append({"name": alias.name, "asname": alias.asname})
return {
"type": "import_from",
"module": node.module,
"names": names,
"level": node.level,
}
elif isinstance(node, ast.Pass):
return {"type": "pass"}
elif isinstance(node, ast.Global):
return {"type": "global", "names": node.names}
elif isinstance(node, ast.Nonlocal):
return {"type": "nonlocal", "names": node.names}
elif isinstance(node, ast.Assert):
return {
"type": "assert",
"test": self._traverse(node.test),
"msg": self._traverse(node.msg) if node.msg else None,
}
elif isinstance(node, ast.Delete):
return {
"type": "delete",
"targets": self._traverse(node.targets),
}
elif isinstance(node, ast.Expr):
# Standalone expression in statement position,
# e.g. ``print("hello")`` — unwrap the inner value.
return {
"type": "expr",
"value": self._traverse(node.value),
}
# ================================================================
# Leaf / expression nodes
# ================================================================
elif isinstance(node, ast.Name):
return {"type": "variable", "name": node.id}
elif isinstance(node, ast.Constant):
return {"type": "constant", "value": node.value}
elif isinstance(node, ast.BinOp):
return {
"type": "binary_op",
"op": self.ops_map.get(type(node.op), "UNKNOWN_OP"),
"left": self._traverse(node.left),
"right": self._traverse(node.right),
}
elif isinstance(node, ast.Compare):
return {
"type": "comparison",
"op": self.comp_ops_map.get(
type(node.ops[0]), "UNKNOWN_CMPOP"
),
"left": self._traverse(node.left),
"comparator": self._traverse(node.comparators[0]),
}
# ================================================================
# Fallback – unknown node type
# ================================================================
return {
"type": "unknown_node",
"node_type": type(node).__name__,
"lineno": getattr(node, "lineno", -1),
}
# ------------------------------------------------------------------
# Type inference helper
# ------------------------------------------------------------------
def _infer_return_type(self, func_node: ast.FunctionDef) -> str:
"""Infer the return type of a function from its body.
Uses simple heuristics based on the first ``return`` statement's
value type. Can be enhanced with type annotations or LLM-based
inference in the future.
Args:
func_node: An ``ast.FunctionDef`` node.
Returns:
A string describing the inferred return type (e.g. ``"number"``,
``"list"``, ``"str"``, …).
"""
# Explicit return annotation – defer to _traverse
if func_node.returns:
traversed = self._traverse(func_node.returns)
if isinstance(traversed, dict):
return traversed.get("name", str(traversed))
return str(traversed)
for stmt in func_node.body:
if isinstance(stmt, ast.Return):
if stmt.value is None:
return "none"
if isinstance(stmt.value, ast.BinOp):
return "number"
elif isinstance(stmt.value, ast.List):
return "list"
elif isinstance(stmt.value, ast.Dict):
return "dict"
elif isinstance(stmt.value, ast.Set):
return "set"
elif isinstance(stmt.value, ast.Tuple):
return "tuple"
elif isinstance(stmt.value, ast.Constant):
if isinstance(stmt.value.value, bool):
return "bool"
elif isinstance(stmt.value.value, (int, float)):
return "number"
elif isinstance(stmt.value.value, str):
return "str"
elif stmt.value.value is None:
return "none"
return type(stmt.value.value).__name__
elif isinstance(stmt.value, ast.Call):
return "call"
elif isinstance(stmt.value, ast.Name):
return "variable"
return "unknown"
# ======================================================================
# Test cases
# ======================================================================
def test_math_encoder() -> None:
"""Run all test cases for CodeToMathConverter."""
converter = CodeToMathConverter()
def _pprint(obj: Any) -> str:
return json.dumps(obj, indent=2, default=str)
def _title(text: str) -> None:
print(f"\n{'=' * 60}")
print(f" {text}")
print(f"{'=' * 60}")
def _test(label: str, code: str) -> None:
_title(label)
print(f"Input code:\n{code.strip()}")
result = converter.convert(code)
print(f"\nMath spec:\n{_pprint(result)}")
# ----------------------------------------------------------------
# 1. Simple function with if/else
# ----------------------------------------------------------------
_test("1. Simple function with if/else", """
def classify(x):
if x > 0:
return "positive"
elif x == 0:
return "zero"
else:
return "negative"
""")
# ----------------------------------------------------------------
# 2. For loop over list
# ----------------------------------------------------------------
_test("2. For loop over list", """
def sum_list(items):
total = 0
for item in items:
total += item
return total
""")
# ----------------------------------------------------------------
# 3. List comprehension
# ----------------------------------------------------------------
_test("3. List comprehension", """
squares = [n * n for n in range(10) if n % 2 == 0]
""")
# ----------------------------------------------------------------
# 4. Function with dict return and function calls
# ----------------------------------------------------------------
_test("4. Function with dict return and function calls", """
def build_person(name, age):
return {"name": name, "age": age, "greeting": hello(name)}
def hello(n):
return f"Hello, {n}!"
""")
# ----------------------------------------------------------------
# 5. Class definition with method
# ----------------------------------------------------------------
_test("5. Class definition with method", """
class Counter:
def __init__(self, start=0):
self.value = start
def increment(self):
self.value += 1
return self.value
""")
# ----------------------------------------------------------------
# 6. F-string
# ----------------------------------------------------------------
_test("6. F-string", """
def greet(name, age):
return f"Hello {name}, you are {age} years old"
""")
# ----------------------------------------------------------------
# 7. While loop with break
# ----------------------------------------------------------------
_test("7. While loop with break", """
def find_first(items, target):
i = 0
while i < len(items):
if items[i] == target:
break
i += 1
return i
""")
# ----------------------------------------------------------------
# 8. Try/except
# ----------------------------------------------------------------
_test("8. Try/except", """
def safe_divide(a, b):
try:
return a / b
except ZeroDivisionError:
return None
finally:
print("done")
""")
print(f"\n{'=' * 60}")
print(" All 8 tests completed.")
print(f"{'=' * 60}")
if __name__ == "__main__":
test_math_encoder()