ContinousAstra_v1 / aegis.py
AGofficial's picture
Upload 21 files
7bb8aac verified
Raw History Blame Contribute Delete
21.5 kB
"""Aegis 0.1: closed AST interpreter. No source reaches Python evaluation.
Only source text and explicit input strings enter this module's interpreter.
Resource bounds are application-level bounds, not an OS security boundary.
"""
from dataclasses import dataclass
import json
import math
import re
import time
@dataclass(frozen=True)
class Limits:
source_bytes: int = 65536
line_length: int = 2048
identifier_length: int = 64
nesting: int = 32
tokens: int = 8192
steps: int = 20000
memory_bytes: int = 2_000_000
string_bytes: int = 16384
input_bytes: int = 16384
output_bytes: int = 65536
integer_bits: int = 4096
exponent: int = 4096
input_requests: int = 20
seconds: float = 1.0
class Fault(Exception):
def __init__(self, kind, token, message):
self.error = dict(kind=kind, line=token.line, column=token.column, message=message)
@dataclass(frozen=True)
class Token:
kind: str
value: object
line: int
column: int
@dataclass(frozen=True)
class Node:
op: str
token: Token
args: tuple
depth: int = 1
ORIGIN = Token("", "", 1, 1)
RESERVED = {"if", "else", "and", "or", "not", "true", "false", "input", "print"}
NUMBER = re.compile(r"(?:[0-9]+\.[0-9]+(?:[eE][+-]?[0-9]+)?|[0-9]+[eE][+-]?[0-9]+|[0-9]+)")
IDENT = re.compile(r"[a-z][a-z0-9_]*")
def envelope(status="ok", output="", error=None, **extra):
return dict(status=status, output=output, error=error, **extra)
def failure(kind, message):
return envelope("error", error=dict(kind=kind, line=1, column=1, message=message))
class Budget:
def __init__(self, limits):
self.limits = limits
self.steps = self.memory = 0
self.remaining = limits.seconds
self.deadline = time.monotonic() + self.remaining
def tick(self, token=ORIGIN, memory=0):
self.steps += 1
self.memory += memory
if (self.steps > self.limits.steps or self.memory > self.limits.memory_bytes
or self.remaining <= 0 or time.monotonic() >= self.deadline):
raise Fault("limit_error", token, "program resource limit exceeded")
def value(self, value, token):
if type(value) is int:
if value.bit_length() > self.limits.integer_bits:
raise Fault("limit_error", token, "integer magnitude limit exceeded")
size = 32 + (value.bit_length() + 7) // 8
elif type(value) is float:
if not math.isfinite(value):
raise Fault("value_error", token, "non-finite numeric result")
size = 32
elif type(value) is str:
try:
size = len(value.encode("utf-8"))
except UnicodeError:
raise Fault("value_error", token, "invalid Unicode text") from None
if size > self.limits.string_bytes:
raise Fault("limit_error", token, "string size limit exceeded")
size += 64
else:
size = 32
self.tick(token, size)
return value
def lex(source, budget):
lim = budget.limits
if isinstance(source, bytes):
if len(source) > lim.source_bytes:
raise Fault("limit_error", ORIGIN, "source size limit exceeded")
try:
source = source.decode("utf-8")
except UnicodeError:
raise Fault("syntax_error", ORIGIN, "source must be UTF-8") from None
if type(source) is not str:
raise Fault("syntax_error", ORIGIN, "source must be text")
if len(source) > lim.source_bytes:
raise Fault("limit_error", ORIGIN, "source size limit exceeded")
try:
size = len(source.encode("utf-8"))
except UnicodeError:
raise Fault("syntax_error", ORIGIN, "source must be UTF-8") from None
if size > lim.source_bytes:
raise Fault("limit_error", ORIGIN, "source size limit exceeded")
budget.tick(memory=size * 4)
tokens, stack, unit = [], [0], None
def add(kind, value, line, column):
t = Token(kind, value, line, column)
budget.tick(t, 256)
if len(tokens) >= lim.tokens:
raise Fault("limit_error", t, "token limit exceeded")
tokens.append(t)
lines = source.replace("\r\n", "\n").split("\n")
for line_no, line in enumerate(lines, 1):
t = Token("", "", line_no, 1)
if len(line) > lim.line_length:
raise Fault("limit_error", t, "line length limit exceeded")
if "\t" in line:
raise Fault("indentation_error", t, "tabs are forbidden")
if not line.strip(" ") or line.lstrip(" ").startswith("#"):
continue
indent = len(line) - len(line.lstrip(" "))
if indent > stack[-1]:
if unit is None:
unit = indent
if unit not in (2, 4) or indent != stack[-1] + unit:
raise Fault("indentation_error", t, "use consistent two or four space indentation")
stack.append(indent)
if len(stack) - 1 > lim.nesting:
raise Fault("limit_error", t, "nesting limit exceeded")
add("INDENT", None, line_no, 1)
else:
while indent < stack[-1]:
stack.pop()
add("DEDENT", None, line_no, 1)
if indent != stack[-1]:
raise Fault("indentation_error", t, "indentation does not match a block")
pos = indent
while pos < len(line):
c = line[pos]
if c == " ":
pos += 1
continue
if c == "#":
break
start = pos
t = Token("", "", line_no, pos + 1)
if c == '"':
pos += 1
escaped = False
while pos < len(line):
if line[pos] == '"' and not escaped:
break
if line[pos] == "\\" and not escaped:
escaped = True
else:
escaped = False
pos += 1
if pos == len(line):
raise Fault("syntax_error", t, "unterminated string")
pos += 1
try:
value = json.loads(line[start:pos])
except ValueError:
raise Fault("syntax_error", t, "invalid string escape or character") from None
add("literal", budget.value(value, t), line_no, start + 1)
elif "0" <= c <= "9":
match = NUMBER.match(line, pos)
raw = match.group()
pos = match.end()
if pos < len(line) and (line[pos].isalnum() or line[pos] in "_."):
raise Fault("syntax_error", t, "malformed number")
if "." in raw or "e" in raw.lower():
value = float(raw)
else:
if len(raw) > min(1234, int(lim.integer_bits * .302) + 2):
raise Fault("limit_error", t, "integer literal limit exceeded")
value = int(raw)
add("literal", budget.value(value, t), line_no, start + 1)
elif "a" <= c <= "z":
match = IDENT.match(line, pos)
word, pos = match.group(), match.end()
if len(word) > lim.identifier_length:
raise Fault("limit_error", t, "identifier length limit exceeded")
add(word if word in RESERVED else "IDENT", word, line_no, start + 1)
elif line[pos:pos+2] in ("==", "!=", "<=", ">="):
add(line[pos:pos+2], None, line_no, pos + 1)
pos += 2
elif c in "=<>+-*/^():":
add(c, None, line_no, pos + 1)
pos += 1
else:
raise Fault("syntax_error", t, "unsupported character")
add("NEWLINE", None, line_no, len(line) + 1)
for _ in stack[1:]:
add("DEDENT", None, len(lines), 1)
add("EOF", None, len(lines), 1)
return tokens
class Parser:
PRECEDENCE = {"or": 1, "and": 2, "==": 3, "!=": 3, "<": 3, "<=": 3,
">": 3, ">=": 3, "+": 4, "-": 4, "*": 5, "/": 5, "^": 6}
def __init__(self, tokens, budget):
self.tokens, self.budget, self.pos = tokens, budget, 0
@property
def current(self):
return self.tokens[self.pos]
def take(self, kind=None):
t = self.current
if kind and t.kind != kind:
category = "indentation_error" if t.kind in ("INDENT", "DEDENT") or kind == "INDENT" else "syntax_error"
raise Fault(category, t, "expected " + kind)
self.pos += 1
return t
def node(self, op, token, *args):
depth = 1 + max((x.depth for x in args if isinstance(x, Node)), default=0)
if depth > self.budget.limits.nesting:
raise Fault("limit_error", token, "expression nesting limit exceeded")
self.budget.tick(token, 256)
return Node(op, token, args, depth)
def expression(self, minimum=1, depth=0):
if depth > self.budget.limits.nesting:
raise Fault("limit_error", self.current, "expression nesting limit exceeded")
t = self.take()
if t.kind in ("not", "-"):
# v0.1 allows one unary operator; parentheses allow another.
left = self.node("unary", t, t.kind, self.primary(depth + 1))
else:
self.pos -= 1
left = self.primary(depth + 1)
compared = False
while self.PRECEDENCE.get(self.current.kind, 0) >= minimum:
op = self.take()
precedence = self.PRECEDENCE[op.kind]
if precedence == 3:
if compared:
raise Fault("syntax_error", op, "comparison chaining is forbidden")
compared = True
right = self.expression(precedence if op.kind == "^" else precedence + 1, depth + 1)
left = self.node("binary", op, op.kind, left, right)
return left
def primary(self, depth):
t = self.take()
if t.kind == "literal":
return self.node("literal", t, t.value)
if t.kind in ("true", "false"):
return self.node("literal", t, t.kind == "true")
if t.kind == "IDENT":
return self.node("name", t, t.value)
if t.kind == "input":
self.take("(")
self.take(")")
return self.node("input", t)
if t.kind == "(":
result = self.expression(depth=depth)
self.take(")")
return result
raise Fault("syntax_error", t, "expected expression")
def block(self, depth=0):
if depth > self.budget.limits.nesting:
raise Fault("limit_error", self.current, "block nesting limit exceeded")
nodes = []
while self.current.kind not in ("EOF", "DEDENT"):
t = self.take()
if t.kind == "IDENT":
self.take("=")
nodes.append(self.node("assign", t, t.value, self.expression()))
self.take("NEWLINE")
elif t.kind == "print":
self.take("(")
nodes.append(self.node("print", t, self.expression()))
self.take(")")
self.take("NEWLINE")
elif t.kind == "if":
condition = self.expression()
yes = self.suite(depth)
no = ()
if self.current.kind == "else":
self.take()
no = self.suite(depth)
nodes.append(self.node("if", t, condition, yes, no))
else:
raise Fault("indentation_error" if t.kind == "INDENT" else "syntax_error", t, "expected statement")
return tuple(nodes)
def suite(self, depth):
self.take(":")
self.take("NEWLINE")
self.take("INDENT")
result = self.block(depth + 1)
if not result:
raise Fault("syntax_error", self.current, "block must contain a statement")
self.take("DEDENT")
return result
NUMERIC = {"int", "float"}
def binary_type(op, a, b, token):
if op in ("==", "!="):
return "bool"
if op in ("and", "or") and a == b == "bool":
return "bool"
if op in ("<", "<=", ">", ">=") and (a == b or {a, b} <= NUMERIC):
return "bool"
if op == "+" and a == b == "str":
return "str"
if op in ("+", "-", "*", "/", "^") and {a, b} <= NUMERIC:
return "float" if op == "/" or "float" in (a, b) else "int"
raise Fault("type_error", token, "incompatible operand types")
def validate_expr(node, env, budget):
budget.tick(node.token)
op, args = node.op, node.args
if op == "literal":
return {type(args[0]).__name__}
if op == "input":
return {"str"}
if op == "name":
if args[0] not in env:
raise Fault("name_error", node.token, "variable is not definitely assigned")
return env[args[0]]
if op == "unary":
types = validate_expr(args[1], env, budget)
if not types <= ({"bool"} if args[0] == "not" else NUMERIC):
raise Fault("type_error", node.token, "invalid unary operand")
return types
left = validate_expr(args[1], env, budget)
right = validate_expr(args[2], env, budget)
types = {binary_type(args[0], a, b, node.token) for a in left for b in right}
# Integer negative powers produce floats; track both possible results.
if args[0] == "^" and "int" in types:
types.add("float")
return types
def validate(nodes, env, budget):
for node in nodes:
budget.tick(node.token)
if node.op == "assign":
env[node.args[0]] = validate_expr(node.args[1], env, budget)
elif node.op == "print":
validate_expr(node.args[0], env, budget)
else:
if validate_expr(node.args[0], env, budget) != {"bool"}:
raise Fault("type_error", node.args[0].token, "if condition must be bool")
yes = validate(node.args[1], dict(env), budget)
no = validate(node.args[2], dict(env), budget)
env = {key: yes[key] | no[key] for key in yes.keys() & no.keys()}
return env
class Session:
"""One fresh run. Suspends at input without replaying prior statements."""
def __init__(self, source, limits=None):
self.budget = Budget(limits or Limits())
self.env, self.output = {}, ""
self.output_size = self.requests = 0
self.waiting = None
self.done = False
self.initial_error = None
try:
parser = Parser(lex(source, self.budget), self.budget)
nodes = parser.block()
parser.take("EOF")
validate(nodes, {}, self.budget)
self.runner = self.block(nodes)
except Fault as exc:
self.initial_error = envelope("error", error=exc.error)
except Exception:
self.initial_error = failure("host_error", "program could not be prepared")
self.budget.remaining = max(0, self.budget.deadline - time.monotonic())
def expression(self, node):
self.budget.tick(node.token)
op, args = node.op, node.args
if op == "literal":
return args[0]
if op == "name":
return self.env[args[0]]
if op == "input":
self.requests += 1
if self.requests > self.budget.limits.input_requests:
raise Fault("limit_error", node.token, "input request limit exceeded")
value = yield node.token
return self.budget.value(value, node.token)
if op == "unary":
value = yield from self.expression(args[1])
return self.budget.value(not value if args[0] == "not" else -value, node.token)
operator = args[0]
left = yield from self.expression(args[1])
if operator == "and" and not left:
return False
if operator == "or" and left:
return True
right = yield from self.expression(args[2])
return self.calculate(operator, left, right, node.token)
def calculate(self, op, a, b, token):
lim = self.budget.limits
if op in ("==", "!="):
equal = type(a) is type(b) and a == b
return equal if op == "==" else not equal
if op in ("and", "or"):
return b
if op in ("<", "<=", ">", ">="):
return {"<": lambda: a < b, "<=": lambda: a <= b,
">": lambda: a > b, ">=": lambda: a >= b}[op]()
if type(a) is str:
if len(a.encode("utf-8")) + len(b.encode("utf-8")) > lim.string_bytes:
raise Fault("limit_error", token, "string size limit exceeded")
else:
if op == "^":
if abs(b) > lim.exponent:
raise Fault("limit_error", token, "exponent limit exceeded")
if type(a) is int and type(b) is int and b >= 0 and abs(a) > 1:
if (abs(a).bit_length() - 1) * b + 1 > lim.integer_bits:
raise Fault("limit_error", token, "integer magnitude limit exceeded")
if a < 0 and type(b) is float and not b.is_integer():
raise Fault("value_error", token, "complex numbers are forbidden")
if op == "*" and type(a) is type(b) is int:
if a and b and a.bit_length() + b.bit_length() - 1 > lim.integer_bits:
raise Fault("limit_error", token, "integer magnitude limit exceeded")
if type(a) is float or type(b) is float:
try:
a, b = float(a), float(b)
except OverflowError:
raise Fault("value_error", token, "numeric promotion overflow") from None
try:
value = {"+": lambda: a + b, "-": lambda: a - b, "*": lambda: a * b,
"/": lambda: a / b, "^": lambda: a ** b}[op]()
except ZeroDivisionError:
raise Fault("value_error", token, "division by zero") from None
except (OverflowError, ValueError):
raise Fault("value_error", token, "numeric result out of range") from None
return self.budget.value(value, token)
def block(self, nodes):
for node in nodes:
self.budget.tick(node.token)
if node.op == "assign":
self.env[node.args[0]] = yield from self.expression(node.args[1])
self.budget.tick(node.token, 128)
elif node.op == "print":
value = yield from self.expression(node.args[0])
text = ("true" if value else "false") if type(value) is bool else str(value)
size = len(text.encode("utf-8")) + 1
if self.output_size + size > self.budget.limits.output_bytes:
raise Fault("limit_error", node.token, "output size limit exceeded")
self.budget.tick(node.token, size * 4)
self.output += text + "\n"
self.output_size += size
else:
condition = yield from self.expression(node.args[0])
yield from self.block(node.args[1] if condition else node.args[2])
def advance(self, value=None):
if self.done:
return failure("value_error", "run has already ended")
if self.initial_error:
self.done = True
return self.initial_error
self.budget.deadline = time.monotonic() + self.budget.remaining
try:
self.budget.tick()
if self.waiting is not None:
if type(value) is not str:
raise Fault("type_error", self.waiting, "input must be a string")
if len(value) > self.budget.limits.input_bytes:
raise Fault("limit_error", self.waiting, "input size limit exceeded")
try:
size = len(value.encode("utf-8"))
except UnicodeError:
raise Fault("value_error", self.waiting, "invalid Unicode input") from None
if size > self.budget.limits.input_bytes:
raise Fault("limit_error", self.waiting, "input size limit exceeded")
self.waiting = self.runner.send(value)
else:
self.waiting = next(self.runner)
return envelope("input_required", self.output, request=dict(index=self.requests,
line=self.waiting.line, column=self.waiting.column))
except StopIteration:
self.done = True
return envelope(output=self.output)
except Fault as exc:
self.done = True
return envelope("error", error=exc.error)
except Exception:
self.done = True
return failure("host_error", "program execution failed")
finally:
self.budget.remaining = max(0, self.budget.deadline - time.monotonic())
if self.done:
self.runner.close()
self.env.clear()
def run(source, inputs=(), limits=None):
session = Session(source, limits)
result = session.advance()
for value in inputs:
if result["status"] != "input_required":
break
result = session.advance(value)
return result