Module refinery.lib.scripts.vba.deobfuscation.simplify

VBA expression simplification and constant folding transforms.

Expand source code Browse git
"""
VBA expression simplification and constant folding transforms.
"""
from __future__ import annotations

import operator

from typing import Callable

from refinery.lib.scripts import Transformer, set_child
from refinery.lib.scripts.vba.deobfuscation.builtins import VBA_BUILTIN_CONSTANTS
from refinery.lib.scripts.vba.deobfuscation.helpers import (
    apply_removals,
    body_lists,
    constant_args,
    is_literal,
    is_nan_or_inf,
    make_integer_literal,
    make_numeric_literal,
    make_string_literal,
    module_compare_mode,
    numeric_value,
    string_value,
    value_to_node,
    vba_int_div,
    vba_mod,
)
from refinery.lib.scripts.vba.deobfuscation.names import (
    CHR_NAMES,
    CompareMode,
    Value,
    dispatch_builtin,
)
from refinery.lib.scripts.vba.model import (
    VbaBinaryExpression,
    VbaBooleanLiteral,
    VbaCallExpression,
    VbaConstDeclaration,
    VbaEmptyLiteral,
    VbaForEachStatement,
    VbaForStatement,
    VbaIdentifier,
    VbaLetStatement,
    VbaModule,
    VbaOnErrorAction,
    VbaOnErrorStatement,
    VbaParenExpression,
    VbaProcedureDeclaration,
    VbaStringLiteral,
    VbaUnaryExpression,
)

_BINARY_OPS: dict[str, Callable] = {
    '+'  : operator.add,
    '-'  : operator.sub,
    '*'  : operator.mul,
    '/'  : operator.truediv,
}

_INTEGER_OPS: dict[str, Callable] = {
    '\\' : vba_int_div,
    'Mod': vba_mod,
}


class _EvaluationFailed(Exception):
    pass


def _try_evaluate_call(node: VbaCallExpression, compare_mode: CompareMode = CompareMode.BINARY) -> Value:
    """
    Try to statically evaluate a VBA builtin call with constant arguments. Raises
    `_EvaluationFailed` if the call cannot be evaluated; otherwise returns the evaluated Python
    value (may be `None` for VBA Empty). The `compare_mode` flag carries the module `Option
    Compare` mode.
    """
    if not isinstance(node.callee, VbaIdentifier):
        raise _EvaluationFailed
    values = constant_args(node.arguments)
    if values is None:
        raise _EvaluationFailed
    name = node.callee.name.lower()
    try:
        matched, result = dispatch_builtin(name, values, compare_mode)
    except (ValueError, OverflowError, TypeError, IndexError):
        raise _EvaluationFailed
    if not matched:
        raise _EvaluationFailed
    return result


def _has_oern(body: list) -> bool:
    return any(
        isinstance(s, VbaOnErrorStatement) and s.action is VbaOnErrorAction.RESUME_NEXT
        for s in body
    )


class VbaSimplifications(Transformer):

    def __init__(self):
        super().__init__()
        self._assigned_names: set[str] = set()
        self._oern_bodies: set[int] = set()
        self._compare_mode = CompareMode.BINARY

    def visit(self, node):
        if isinstance(node, VbaModule):
            self._collect_context(node)
            super().visit(node)
            if self._remove_self_assignments(node):
                self.mark_changed()
            return None
        return super().visit(node)

    def _collect_context(self, module: VbaModule):
        self._assigned_names = set(VBA_BUILTIN_CONSTANTS)
        self._oern_bodies = set()
        self._compare_mode = module_compare_mode(module)
        for n in module.walk():
            if isinstance(n, VbaLetStatement) and isinstance(n.target, VbaIdentifier):
                self._assigned_names.add(n.target.name.lower())
            elif isinstance(n, VbaConstDeclaration):
                for d in n.declarators:
                    self._assigned_names.add(d.name.lower())
            elif isinstance(n, (VbaForStatement, VbaForEachStatement)):
                if isinstance(n.variable, VbaIdentifier):
                    self._assigned_names.add(n.variable.name.lower())
            if isinstance(n, VbaProcedureDeclaration):
                if n.params:
                    for p in n.params:
                        self._assigned_names.add(p.name.lower())
                if n.name:
                    self._assigned_names.add(n.name.lower())
                if n.body and _has_oern(n.body):
                    self._oern_bodies.add(id(n.body))
        if module.body and _has_oern(module.body):
            self._oern_bodies.add(id(module.body))

    @staticmethod
    def _remove_self_assignments(module: VbaModule) -> bool:
        removals: list[tuple[int, list]] = []
        for body in body_lists(module):
            for idx, stmt in enumerate(body):
                if (
                    isinstance(stmt, VbaLetStatement)
                    and isinstance(stmt.target, VbaIdentifier)
                    and isinstance(stmt.value, VbaIdentifier)
                    and stmt.target.name.lower() == stmt.value.name.lower()
                ):
                    removals.append((idx, body))
        return apply_removals(removals)

    def _is_oern_undefined(self, node) -> bool:
        if not isinstance(node, VbaIdentifier):
            return False
        if node.name.lower() in self._assigned_names:
            return False
        parent = node.parent
        while parent is not None:
            if isinstance(parent, VbaProcedureDeclaration):
                return id(parent.body) in self._oern_bodies
            if isinstance(parent, VbaModule):
                return id(parent.body) in self._oern_bodies
            parent = parent.parent
        return False

    def visit_VbaBinaryExpression(self, node: VbaBinaryExpression):
        self.generic_visit(node)
        if node.left is None or node.right is None:
            return None
        if node.operator in ('&', '+'):
            result = self._fold_string_concat(node)
            if result is not None:
                return result
        return self._fold_numeric_binary(node)

    def _fold_string_concat(self, node: VbaBinaryExpression):
        if node.operator == '&':
            if isinstance(node.right, (VbaEmptyLiteral, VbaStringLiteral)) and not node.right.value:
                return node.left
            if isinstance(node.left, (VbaEmptyLiteral, VbaStringLiteral)) and not node.left.value:
                return node.right
        if self._is_oern_undefined(node.left) and string_value(node.right) is not None:
            return node.right
        if self._is_oern_undefined(node.right) and string_value(node.left) is not None:
            return node.left
        lhs = string_value(node.left)
        rhs = string_value(node.right)
        if lhs is not None and rhs is not None:
            return make_string_literal(lhs + rhs)
        if rhs is not None:
            if (
                isinstance(node.left, VbaBinaryExpression)
                and node.left.operator in ('&', '+')
            ):
                inner_right_str = string_value(node.left.right)
                if inner_right_str is not None:
                    set_child(node.left, 'right', make_string_literal(inner_right_str + rhs))
                    return node.left
        if lhs is not None:
            inner = node.right
            while (
                isinstance(inner, VbaBinaryExpression)
                and inner.operator in ('&', '+')
                and isinstance(inner.left, VbaBinaryExpression)
                and inner.left.operator in ('&', '+')
            ):
                inner = inner.left
            if (
                isinstance(inner, VbaBinaryExpression)
                and inner.operator in ('&', '+')
            ):
                inner_left_str = string_value(inner.left)
                if inner_left_str is not None:
                    set_child(inner, 'left', make_string_literal(lhs + inner_left_str))
                    return node.right
        return None

    @staticmethod
    def _fold_numeric_binary(node: VbaBinaryExpression):
        lhs = numeric_value(node.left)
        rhs = numeric_value(node.right)
        if lhs is not None and rhs is not None:
            fn = _BINARY_OPS.get(node.operator)
            if fn is not None:
                try:
                    result = fn(lhs, rhs)
                except (ZeroDivisionError, ValueError, OverflowError):
                    return None
                if is_nan_or_inf(result):
                    return None
                return make_numeric_literal(result)
            fn = _INTEGER_OPS.get(node.operator)
            if fn is not None:
                try:
                    result = fn(lhs, rhs)
                except (ZeroDivisionError, ValueError, OverflowError):
                    return None
                return make_integer_literal(int(result))
            if node.operator == '^':
                try:
                    result = lhs ** rhs
                except (ZeroDivisionError, ValueError, OverflowError):
                    return None
                return make_numeric_literal(result)
        return None

    def visit_VbaCallExpression(self, node: VbaCallExpression):
        self.generic_visit(node)
        try:
            result = _try_evaluate_call(node, self._compare_mode)
        except _EvaluationFailed:
            return None
        if isinstance(result, str) and len(result) == 1 and not result.isprintable():
            if (
                isinstance(node.callee, VbaIdentifier)
                and node.callee.name.lower() in CHR_NAMES
            ):
                return None
        return value_to_node(result)

    def visit_VbaIdentifier(self, node: VbaIdentifier):
        value = VBA_BUILTIN_CONSTANTS.get(node.name.lower())
        if value is None:
            return None
        return make_integer_literal(value)

    def visit_VbaParenExpression(self, node: VbaParenExpression):
        self.generic_visit(node)
        inner = node.expression
        if inner is None:
            return None
        if isinstance(node.parent, VbaCallExpression) and node in node.parent.arguments:
            return None
        if isinstance(inner, (VbaIdentifier, VbaParenExpression)) or is_literal(inner):
            return inner
        return None

    def visit_VbaUnaryExpression(self, node: VbaUnaryExpression):
        self.generic_visit(node)
        if node.operand is None:
            return None
        op = node.operator
        if op == '-':
            val = numeric_value(node.operand)
            if val is not None:
                return make_numeric_literal(-val)
        if op == 'Not':
            if isinstance(node.operand, VbaBooleanLiteral):
                return VbaBooleanLiteral(value=not node.operand.value)
            val = numeric_value(node.operand)
            if isinstance(val, int):
                return make_integer_literal(~val)
        return None

Classes

class VbaSimplifications

In-place tree rewriter. Each visit method may return a replacement node or None to keep the original. Tracks whether any transformation was applied via the changed flag.

When a models cache is attached by the pipeline, setting changed truthy invalidates it, so a transform that mutates the tree never leaves a stale model behind for the next consumer. The same set advances the global mutation epoch, which protects the Node.children memo exactly as far as the model caches.

Expand source code Browse git
class VbaSimplifications(Transformer):

    def __init__(self):
        super().__init__()
        self._assigned_names: set[str] = set()
        self._oern_bodies: set[int] = set()
        self._compare_mode = CompareMode.BINARY

    def visit(self, node):
        if isinstance(node, VbaModule):
            self._collect_context(node)
            super().visit(node)
            if self._remove_self_assignments(node):
                self.mark_changed()
            return None
        return super().visit(node)

    def _collect_context(self, module: VbaModule):
        self._assigned_names = set(VBA_BUILTIN_CONSTANTS)
        self._oern_bodies = set()
        self._compare_mode = module_compare_mode(module)
        for n in module.walk():
            if isinstance(n, VbaLetStatement) and isinstance(n.target, VbaIdentifier):
                self._assigned_names.add(n.target.name.lower())
            elif isinstance(n, VbaConstDeclaration):
                for d in n.declarators:
                    self._assigned_names.add(d.name.lower())
            elif isinstance(n, (VbaForStatement, VbaForEachStatement)):
                if isinstance(n.variable, VbaIdentifier):
                    self._assigned_names.add(n.variable.name.lower())
            if isinstance(n, VbaProcedureDeclaration):
                if n.params:
                    for p in n.params:
                        self._assigned_names.add(p.name.lower())
                if n.name:
                    self._assigned_names.add(n.name.lower())
                if n.body and _has_oern(n.body):
                    self._oern_bodies.add(id(n.body))
        if module.body and _has_oern(module.body):
            self._oern_bodies.add(id(module.body))

    @staticmethod
    def _remove_self_assignments(module: VbaModule) -> bool:
        removals: list[tuple[int, list]] = []
        for body in body_lists(module):
            for idx, stmt in enumerate(body):
                if (
                    isinstance(stmt, VbaLetStatement)
                    and isinstance(stmt.target, VbaIdentifier)
                    and isinstance(stmt.value, VbaIdentifier)
                    and stmt.target.name.lower() == stmt.value.name.lower()
                ):
                    removals.append((idx, body))
        return apply_removals(removals)

    def _is_oern_undefined(self, node) -> bool:
        if not isinstance(node, VbaIdentifier):
            return False
        if node.name.lower() in self._assigned_names:
            return False
        parent = node.parent
        while parent is not None:
            if isinstance(parent, VbaProcedureDeclaration):
                return id(parent.body) in self._oern_bodies
            if isinstance(parent, VbaModule):
                return id(parent.body) in self._oern_bodies
            parent = parent.parent
        return False

    def visit_VbaBinaryExpression(self, node: VbaBinaryExpression):
        self.generic_visit(node)
        if node.left is None or node.right is None:
            return None
        if node.operator in ('&', '+'):
            result = self._fold_string_concat(node)
            if result is not None:
                return result
        return self._fold_numeric_binary(node)

    def _fold_string_concat(self, node: VbaBinaryExpression):
        if node.operator == '&':
            if isinstance(node.right, (VbaEmptyLiteral, VbaStringLiteral)) and not node.right.value:
                return node.left
            if isinstance(node.left, (VbaEmptyLiteral, VbaStringLiteral)) and not node.left.value:
                return node.right
        if self._is_oern_undefined(node.left) and string_value(node.right) is not None:
            return node.right
        if self._is_oern_undefined(node.right) and string_value(node.left) is not None:
            return node.left
        lhs = string_value(node.left)
        rhs = string_value(node.right)
        if lhs is not None and rhs is not None:
            return make_string_literal(lhs + rhs)
        if rhs is not None:
            if (
                isinstance(node.left, VbaBinaryExpression)
                and node.left.operator in ('&', '+')
            ):
                inner_right_str = string_value(node.left.right)
                if inner_right_str is not None:
                    set_child(node.left, 'right', make_string_literal(inner_right_str + rhs))
                    return node.left
        if lhs is not None:
            inner = node.right
            while (
                isinstance(inner, VbaBinaryExpression)
                and inner.operator in ('&', '+')
                and isinstance(inner.left, VbaBinaryExpression)
                and inner.left.operator in ('&', '+')
            ):
                inner = inner.left
            if (
                isinstance(inner, VbaBinaryExpression)
                and inner.operator in ('&', '+')
            ):
                inner_left_str = string_value(inner.left)
                if inner_left_str is not None:
                    set_child(inner, 'left', make_string_literal(lhs + inner_left_str))
                    return node.right
        return None

    @staticmethod
    def _fold_numeric_binary(node: VbaBinaryExpression):
        lhs = numeric_value(node.left)
        rhs = numeric_value(node.right)
        if lhs is not None and rhs is not None:
            fn = _BINARY_OPS.get(node.operator)
            if fn is not None:
                try:
                    result = fn(lhs, rhs)
                except (ZeroDivisionError, ValueError, OverflowError):
                    return None
                if is_nan_or_inf(result):
                    return None
                return make_numeric_literal(result)
            fn = _INTEGER_OPS.get(node.operator)
            if fn is not None:
                try:
                    result = fn(lhs, rhs)
                except (ZeroDivisionError, ValueError, OverflowError):
                    return None
                return make_integer_literal(int(result))
            if node.operator == '^':
                try:
                    result = lhs ** rhs
                except (ZeroDivisionError, ValueError, OverflowError):
                    return None
                return make_numeric_literal(result)
        return None

    def visit_VbaCallExpression(self, node: VbaCallExpression):
        self.generic_visit(node)
        try:
            result = _try_evaluate_call(node, self._compare_mode)
        except _EvaluationFailed:
            return None
        if isinstance(result, str) and len(result) == 1 and not result.isprintable():
            if (
                isinstance(node.callee, VbaIdentifier)
                and node.callee.name.lower() in CHR_NAMES
            ):
                return None
        return value_to_node(result)

    def visit_VbaIdentifier(self, node: VbaIdentifier):
        value = VBA_BUILTIN_CONSTANTS.get(node.name.lower())
        if value is None:
            return None
        return make_integer_literal(value)

    def visit_VbaParenExpression(self, node: VbaParenExpression):
        self.generic_visit(node)
        inner = node.expression
        if inner is None:
            return None
        if isinstance(node.parent, VbaCallExpression) and node in node.parent.arguments:
            return None
        if isinstance(inner, (VbaIdentifier, VbaParenExpression)) or is_literal(inner):
            return inner
        return None

    def visit_VbaUnaryExpression(self, node: VbaUnaryExpression):
        self.generic_visit(node)
        if node.operand is None:
            return None
        op = node.operator
        if op == '-':
            val = numeric_value(node.operand)
            if val is not None:
                return make_numeric_literal(-val)
        if op == 'Not':
            if isinstance(node.operand, VbaBooleanLiteral):
                return VbaBooleanLiteral(value=not node.operand.value)
            val = numeric_value(node.operand)
            if isinstance(val, int):
                return make_integer_literal(~val)
        return None

Ancestors

Methods

def visit(self, node)
Expand source code Browse git
def visit(self, node):
    if isinstance(node, VbaModule):
        self._collect_context(node)
        super().visit(node)
        if self._remove_self_assignments(node):
            self.mark_changed()
        return None
    return super().visit(node)
def visit_VbaBinaryExpression(self, node)
Expand source code Browse git
def visit_VbaBinaryExpression(self, node: VbaBinaryExpression):
    self.generic_visit(node)
    if node.left is None or node.right is None:
        return None
    if node.operator in ('&', '+'):
        result = self._fold_string_concat(node)
        if result is not None:
            return result
    return self._fold_numeric_binary(node)
def visit_VbaCallExpression(self, node)
Expand source code Browse git
def visit_VbaCallExpression(self, node: VbaCallExpression):
    self.generic_visit(node)
    try:
        result = _try_evaluate_call(node, self._compare_mode)
    except _EvaluationFailed:
        return None
    if isinstance(result, str) and len(result) == 1 and not result.isprintable():
        if (
            isinstance(node.callee, VbaIdentifier)
            and node.callee.name.lower() in CHR_NAMES
        ):
            return None
    return value_to_node(result)
def visit_VbaIdentifier(self, node)
Expand source code Browse git
def visit_VbaIdentifier(self, node: VbaIdentifier):
    value = VBA_BUILTIN_CONSTANTS.get(node.name.lower())
    if value is None:
        return None
    return make_integer_literal(value)
def visit_VbaParenExpression(self, node)
Expand source code Browse git
def visit_VbaParenExpression(self, node: VbaParenExpression):
    self.generic_visit(node)
    inner = node.expression
    if inner is None:
        return None
    if isinstance(node.parent, VbaCallExpression) and node in node.parent.arguments:
        return None
    if isinstance(inner, (VbaIdentifier, VbaParenExpression)) or is_literal(inner):
        return inner
    return None
def visit_VbaUnaryExpression(self, node)
Expand source code Browse git
def visit_VbaUnaryExpression(self, node: VbaUnaryExpression):
    self.generic_visit(node)
    if node.operand is None:
        return None
    op = node.operator
    if op == '-':
        val = numeric_value(node.operand)
        if val is not None:
            return make_numeric_literal(-val)
    if op == 'Not':
        if isinstance(node.operand, VbaBooleanLiteral):
            return VbaBooleanLiteral(value=not node.operand.value)
        val = numeric_value(node.operand)
        if isinstance(val, int):
            return make_integer_literal(~val)
    return None

Inherited members