Skip to content
Merged
Show file tree
Hide file tree
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
1 change: 1 addition & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -377,6 +377,7 @@ If you get stuck, run `:examples` or `:h`.
- In ODE input, prefer explicit multiplication (`20*y` instead of `20y`) for predictable parsing.
- Common LaTeX wrappers and commands are normalized: `$...$`, `\(...\)`, `\sin`, `\cos`, `\ln`, `\sqrt{...}`, `\frac{a}{b}`
- `name = expr` assigns in REPL session (`ans` is always last result)
- Built-in helper names are reserved for evaluation (for example `sin`, `gamma`, `atan2`, `I`) and cannot be reassigned.
- Undefined symbols raise an error

## Safety limits
Expand Down
57 changes: 54 additions & 3 deletions src/calc/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
from difflib import get_close_matches
from math import log

import sympy as _sympy
from sympy import (
Abs,
Add,
Expand Down Expand Up @@ -50,7 +51,7 @@
function_exponentiation,
implicit_multiplication_application,
parse_expr,
rationalize,
rationalize,
)

x, y, z, t = symbols("x y z t")
Expand Down Expand Up @@ -131,6 +132,56 @@ def _den(expr):
"S": Symbol,
}

SYMPY_EXTRA_ALLOWLIST = (
"I",
"acos",
"acosh",
"apart",
"asin",
"asinh",
"atan",
"atan2",
"atanh",
"binomial",
"cancel",
"collect",
"cosh",
"cot",
"csc",
"expand",
"factor",
"gamma",
"limit",
"powsimp",
"product",
"sec",
"series",
"sinh",
"summation",
"tanh",
"together",
"trigsimp",
)


def _expanded_sympy_locals() -> dict[str, object]:
expanded: dict[str, object] = {}
for name in SYMPY_EXTRA_ALLOWLIST:
if name in LOCALS_DICT:
continue
if not name.isidentifier() or name.startswith("_"):
continue
value = getattr(_sympy, name, None)
if value is None:
continue
if not (callable(value) or isinstance(value, _sympy.Basic)):
continue
expanded[name] = value
return expanded


LOCALS_DICT.update(_expanded_sympy_locals())


# parse_expr internally uses eval. Keep globals minimal and disable builtins.
GLOBAL_DICT = {
Expand All @@ -145,14 +196,14 @@ def _den(expr):
"factorial": factorial,
}

TRANSFORMS = (auto_number, factorial_notation, convert_xor, function_exponentiation, rationalize, )
TRANSFORMS = (auto_number, factorial_notation, convert_xor, function_exponentiation, rationalize)
RELAXED_TRANSFORMS = (
auto_number,
factorial_notation,
convert_xor,
function_exponentiation,
implicit_multiplication_application,
rationalize,
rationalize,
)
MAX_EXPRESSION_CHARS = 2000
BLOCKED_PATTERN = re.compile(r"(__|;|\n|\r)")
Expand Down
26 changes: 25 additions & 1 deletion tests/test_core.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import pytest
from sympy import I, Symbol

from calc.core import evaluate, normalize_expression, reserved_name_suggestion
from calc.core import LOCALS_DICT, evaluate, normalize_expression, reserved_name_suggestion


def test_exact_arithmetic():
Expand Down Expand Up @@ -135,6 +136,20 @@ def test_symbol_helpers_for_coefficient_workflows():
assert "C: 421/15" in out


def test_expanded_sympy_allowlist_functions_are_available():
assert str(evaluate("binomial(5, 2)")) == "10"
assert str(evaluate("gamma(6)")) == "120"
assert str(evaluate("limit(sin(x)/x, x, 0)")) == "1"
assert str(evaluate("atan2(1, 1)")) == "pi/4"
assert str(evaluate("I^2")) == "-1"


def test_expanded_sympy_allowlist_does_not_override_curated_names():
assert LOCALS_DICT["S"] is Symbol
assert str(evaluate('S("A")')) == "A"
assert LOCALS_DICT["I"] is I


def test_numeric_eval():
assert str(evaluate("N(pi, 10)")) == "3.141592654"

Expand Down Expand Up @@ -162,6 +177,15 @@ def test_assignment_and_ans_with_session_locals():
def test_assignment_rejects_reserved_name():
with pytest.raises(ValueError, match="reserved name"):
evaluate("sin = 2", session_locals={})
with pytest.raises(ValueError, match="reserved name"):
evaluate("gamma = 2", session_locals={})
with pytest.raises(ValueError, match="reserved name"):
evaluate("I = 2", session_locals={})


def test_expanded_sympy_allowlist_keeps_local_surface_safe():
assert "__builtins__" not in LOCALS_DICT
assert all(not name.startswith("_") for name in LOCALS_DICT)


def test_reserved_name_suggestion_prefers_close_session_name():
Expand Down
Loading