haiku.rag/tests/agents/rlm/test_sandbox.py
2026-02-06 12:07:02 +01:00

490 lines
18 KiB
Python

import pytest
class TestSafeBuiltins:
"""Test that safe builtins are available."""
@pytest.mark.asyncio
async def test_print_available(self, repl_env_empty):
result = await repl_env_empty.execute_async("print('hello')")
assert result.success
assert "hello" in result.stdout
@pytest.mark.asyncio
async def test_len_available(self, repl_env_empty):
result = await repl_env_empty.execute_async("print(len([1, 2, 3]))")
assert result.success
assert "3" in result.stdout
@pytest.mark.asyncio
async def test_range_available(self, repl_env_empty):
result = await repl_env_empty.execute_async("print(list(range(3)))")
assert result.success
assert "[0, 1, 2]" in result.stdout
@pytest.mark.asyncio
async def test_enumerate_available(self, repl_env_empty):
result = await repl_env_empty.execute_async(
"print(list(enumerate(['a', 'b'])))"
)
assert result.success
assert "[(0, 'a'), (1, 'b')]" in result.stdout
@pytest.mark.asyncio
async def test_sorted_available(self, repl_env_empty):
result = await repl_env_empty.execute_async("print(sorted([3, 1, 2]))")
assert result.success
assert "[1, 2, 3]" in result.stdout
@pytest.mark.asyncio
async def test_sum_available(self, repl_env_empty):
result = await repl_env_empty.execute_async("print(sum([1, 2, 3]))")
assert result.success
assert "6" in result.stdout
@pytest.mark.asyncio
async def test_min_max_available(self, repl_env_empty):
result = await repl_env_empty.execute_async(
"print(min([3, 1, 2]), max([3, 1, 2]))"
)
assert result.success
assert "1 3" in result.stdout
@pytest.mark.asyncio
async def test_all_any_available(self, repl_env_empty):
result = await repl_env_empty.execute_async(
"print(all([True, True]), any([False, True]))"
)
assert result.success
assert "True True" in result.stdout
@pytest.mark.asyncio
async def test_dict_list_set_tuple_available(self, repl_env_empty):
result = await repl_env_empty.execute_async(
"print(dict(a=1), list((1,2)), set([1,2,1]), tuple([1,2]))"
)
assert result.success
assert "{'a': 1}" in result.stdout
@pytest.mark.asyncio
async def test_str_int_float_bool_available(self, repl_env_empty):
result = await repl_env_empty.execute_async(
"print(str(1), int('2'), float('3.0'), bool(1))"
)
assert result.success
assert "1 2 3.0 True" in result.stdout
@pytest.mark.asyncio
async def test_zip_map_filter_available(self, repl_env_empty):
result = await repl_env_empty.execute_async(
"print(list(zip([1,2], ['a','b'])), "
"list(map(str, [1,2])), "
"list(filter(lambda x: x > 1, [1,2,3])))"
)
assert result.success
assert "[(1, 'a'), (2, 'b')]" in result.stdout
@pytest.mark.asyncio
async def test_isinstance_type_available(self, repl_env_empty):
result = await repl_env_empty.execute_async(
"print(isinstance(1, int), type([]))"
)
assert result.success
assert "True" in result.stdout
class TestDangerousBuiltinsBlocked:
"""Test that dangerous builtins are blocked."""
@pytest.mark.asyncio
async def test_eval_blocked(self, repl_env_empty):
result = await repl_env_empty.execute_async("eval('1+1')")
assert not result.success
assert "eval" in result.stderr.lower() or "not defined" in result.stderr.lower()
@pytest.mark.asyncio
async def test_exec_blocked(self, repl_env_empty):
result = await repl_env_empty.execute_async("exec('x = 1')")
assert not result.success
assert "exec" in result.stderr.lower() or "not defined" in result.stderr.lower()
@pytest.mark.asyncio
async def test_compile_blocked(self, repl_env_empty):
result = await repl_env_empty.execute_async(
"compile('1+1', '<string>', 'eval')"
)
assert not result.success
assert (
"compile" in result.stderr.lower() or "not defined" in result.stderr.lower()
)
@pytest.mark.asyncio
async def test_open_blocked(self, repl_env_empty):
result = await repl_env_empty.execute_async("open('/etc/passwd')")
assert not result.success
assert "open" in result.stderr.lower() or "not defined" in result.stderr.lower()
@pytest.mark.asyncio
async def test_input_blocked(self, repl_env_empty):
result = await repl_env_empty.execute_async("input('Enter: ')")
assert not result.success
assert (
"input" in result.stderr.lower() or "not defined" in result.stderr.lower()
)
@pytest.mark.asyncio
async def test___import___blocked(self, repl_env_empty):
result = await repl_env_empty.execute_async("__import__('os')")
assert not result.success
@pytest.mark.asyncio
async def test_globals_blocked(self, repl_env_empty):
result = await repl_env_empty.execute_async("globals()")
assert not result.success
@pytest.mark.asyncio
async def test_locals_blocked(self, repl_env_empty):
result = await repl_env_empty.execute_async("locals()")
assert not result.success
@pytest.mark.asyncio
async def test_breakpoint_blocked(self, repl_env_empty):
result = await repl_env_empty.execute_async("breakpoint()")
assert not result.success
@pytest.mark.asyncio
async def test_getattr_setattr_blocked(self, repl_env_empty):
result = await repl_env_empty.execute_async("getattr(object, '__class__')")
assert not result.success
@pytest.mark.asyncio
async def test_delattr_blocked(self, repl_env_empty):
result = await repl_env_empty.execute_async("delattr(object, 'x')")
assert not result.success
class TestAllowedImports:
"""Test that allowed imports work."""
@pytest.mark.asyncio
async def test_json_import(self, repl_env_empty):
result = await repl_env_empty.execute_async(
"import json\nprint(json.dumps({'a': 1}))"
)
assert result.success
assert '{"a": 1}' in result.stdout
@pytest.mark.asyncio
async def test_re_import(self, repl_env_empty):
result = await repl_env_empty.execute_async(
"import re\nprint(re.match(r'\\d+', '123').group())"
)
assert result.success
assert "123" in result.stdout
@pytest.mark.asyncio
async def test_math_import(self, repl_env_empty):
result = await repl_env_empty.execute_async("import math\nprint(math.sqrt(4))")
assert result.success
assert "2.0" in result.stdout
@pytest.mark.asyncio
async def test_statistics_import(self, repl_env_empty):
result = await repl_env_empty.execute_async(
"import statistics\nprint(statistics.mean([1, 2, 3]))"
)
assert result.success
assert "2" in result.stdout
@pytest.mark.asyncio
async def test_collections_import(self, repl_env_empty):
result = await repl_env_empty.execute_async(
"from collections import Counter\nprint(Counter(['a', 'b', 'a']))"
)
assert result.success
assert "'a': 2" in result.stdout
@pytest.mark.asyncio
async def test_itertools_import(self, repl_env_empty):
result = await repl_env_empty.execute_async(
"from itertools import chain\nprint(list(chain([1], [2])))"
)
assert result.success
assert "[1, 2]" in result.stdout
@pytest.mark.asyncio
async def test_functools_import(self, repl_env_empty):
result = await repl_env_empty.execute_async(
"from functools import reduce\nprint(reduce(lambda a, b: a+b, [1,2,3]))"
)
assert result.success
assert "6" in result.stdout
@pytest.mark.asyncio
async def test_datetime_import(self, repl_env_empty):
result = await repl_env_empty.execute_async(
"from datetime import date\nprint(date(2025, 1, 1))"
)
assert result.success
assert "2025-01-01" in result.stdout
class TestDangerousImportsBlocked:
"""Test that dangerous imports are blocked."""
@pytest.mark.asyncio
async def test_os_blocked(self, repl_env_empty):
result = await repl_env_empty.execute_async("import os")
assert not result.success
assert (
"not allowed" in result.stderr.lower() or "error" in result.stderr.lower()
)
@pytest.mark.asyncio
async def test_sys_blocked(self, repl_env_empty):
result = await repl_env_empty.execute_async("import sys")
assert not result.success
@pytest.mark.asyncio
async def test_subprocess_blocked(self, repl_env_empty):
result = await repl_env_empty.execute_async("import subprocess")
assert not result.success
@pytest.mark.asyncio
async def test_shutil_blocked(self, repl_env_empty):
result = await repl_env_empty.execute_async("import shutil")
assert not result.success
@pytest.mark.asyncio
async def test_socket_blocked(self, repl_env_empty):
result = await repl_env_empty.execute_async("import socket")
assert not result.success
@pytest.mark.asyncio
async def test_requests_blocked(self, repl_env_empty):
result = await repl_env_empty.execute_async("import requests")
assert not result.success
@pytest.mark.asyncio
async def test_builtins_module_blocked(self, repl_env_empty):
result = await repl_env_empty.execute_async("import builtins")
assert not result.success
class TestHaikuRAGBridgeFunctions:
"""Test haiku.rag bridge functions in sandbox."""
@pytest.mark.asyncio
async def test_search(self, repl_env_empty):
"""Test search function calls client with correct args."""
from unittest.mock import AsyncMock
from haiku.rag.store.models import SearchResult
mock_results = [
SearchResult(
chunk_id="chunk-1",
document_id="doc-1",
document_title="Test Doc",
document_uri="test://doc",
content="Test content about foxes",
score=0.9,
page_numbers=[1],
headings=["Heading"],
)
]
repl_env_empty.client.search = AsyncMock(return_value=mock_results)
result = await repl_env_empty.execute_async(
"results = search('fox', limit=5)\n"
"print(len(results), results[0]['chunk_id'], 'fox' in results[0]['content'].lower())"
)
assert result.success
assert "1 chunk-1 True" in result.stdout
repl_env_empty.client.search.assert_called_once_with(
"fox", limit=5, filter=None
)
@pytest.mark.asyncio
async def test_list_documents(self, repl_env_empty):
"""Test list_documents returns list structure."""
result = await repl_env_empty.execute_async(
"docs = list_documents()\nprint(type(docs).__name__, len(docs))"
)
assert result.success
assert "list 0" in result.stdout
@pytest.mark.asyncio
async def test_get_document(self, repl_env_empty):
"""Test get_document calls client correctly."""
from unittest.mock import AsyncMock
from haiku.rag.store.models import Document
mock_doc = Document(
id="doc-1",
uri="test://doc",
title="Test Doc",
content="The quick brown fox",
)
repl_env_empty.client.get_document_by_id = AsyncMock(return_value=mock_doc)
result = await repl_env_empty.execute_async(
"doc = get_document('doc-1')\nprint('fox' in doc.lower())"
)
assert result.success
assert "True" in result.stdout
@pytest.mark.asyncio
async def test_get_document_missing(self, repl_env_empty):
"""Test get_document returns None for missing document."""
result = await repl_env_empty.execute_async(
"doc = get_document('Nonexistent')\nprint(doc is None)"
)
assert result.success
assert "True" in result.stdout
@pytest.mark.asyncio
async def test_ask(self, repl_env_empty):
"""Test ask function calls client with correct args."""
from unittest.mock import AsyncMock
repl_env_empty.client.ask = AsyncMock(return_value=("The fox is brown.", []))
result = await repl_env_empty.execute_async(
"answer = ask('What color is the fox?')\nprint('fox' in answer.lower())"
)
assert result.success
assert "True" in result.stdout
repl_env_empty.client.ask.assert_called_once_with(
"What color is the fox?", filter=None
)
class TestSandboxExecution:
"""Test general sandbox execution behavior."""
@pytest.mark.asyncio
async def test_variable_persistence(self, repl_env_empty):
"""Variables persist across executions."""
await repl_env_empty.execute_async("x = 42")
result = await repl_env_empty.execute_async("print(x)")
assert result.success
assert "42" in result.stdout
@pytest.mark.asyncio
async def test_function_definition(self, repl_env_empty):
"""Can define and call functions."""
result = await repl_env_empty.execute_async(
"def add(a, b):\n return a + b\nprint(add(1, 2))"
)
assert result.success
assert "3" in result.stdout
@pytest.mark.asyncio
async def test_class_definition(self, repl_env_empty):
"""Can define and use classes."""
result = await repl_env_empty.execute_async(
"class Point:\n"
" def __init__(self, x, y):\n"
" self.x = x\n"
" self.y = y\n"
"p = Point(1, 2)\n"
"print(p.x, p.y)"
)
assert result.success
assert "1 2" in result.stdout
@pytest.mark.asyncio
async def test_list_comprehension(self, repl_env_empty):
"""List comprehensions work."""
result = await repl_env_empty.execute_async("print([x**2 for x in range(5)])")
assert result.success
assert "[0, 1, 4, 9, 16]" in result.stdout
@pytest.mark.asyncio
async def test_dict_comprehension(self, repl_env_empty):
"""Dict comprehensions work."""
result = await repl_env_empty.execute_async(
"print({x: x**2 for x in range(3)})"
)
assert result.success
assert "{0: 0, 1: 1, 2: 4}" in result.stdout
@pytest.mark.asyncio
async def test_exception_handling(self, repl_env_empty):
"""Can catch and handle exceptions."""
result = await repl_env_empty.execute_async(
"try:\n x = 1/0\nexcept ZeroDivisionError:\n print('caught')"
)
assert result.success
assert "caught" in result.stdout
@pytest.mark.asyncio
async def test_uncaught_exception_reports_error(self, repl_env_empty):
"""Uncaught exceptions are reported."""
result = await repl_env_empty.execute_async("x = 1/0")
assert not result.success
assert "ZeroDivisionError" in result.stderr
@pytest.mark.asyncio
async def test_syntax_error_reports_error(self, repl_env_empty):
"""Syntax errors are reported."""
result = await repl_env_empty.execute_async("def foo(")
assert not result.success
assert "SyntaxError" in result.stderr
@pytest.mark.asyncio
async def test_output_truncation(self, repl_env_empty):
"""Output is truncated if too long."""
repl_env_empty.config.max_output_chars = 100
result = await repl_env_empty.execute_async("print('x' * 1000)")
assert result.success
assert (
len(result.stdout) <= 100 + 50
) # Allow some margin for truncation message
class TestSecurityEscapes:
"""Test that common security escape attempts are blocked."""
@pytest.mark.asyncio
async def test_eval_via_builtins_dict(self, repl_env_empty):
"""Cannot access eval through __builtins__."""
result = await repl_env_empty.execute_async("__builtins__['eval']('1+1')")
assert not result.success
@pytest.mark.asyncio
async def test_import_via_builtins(self, repl_env_empty):
"""Cannot import os through builtins trickery."""
result = await repl_env_empty.execute_async("__builtins__.__import__('os')")
assert not result.success
@pytest.mark.asyncio
async def test_class_bases_escape(self, repl_env_empty):
"""Cannot escape through __class__.__bases__."""
result = await repl_env_empty.execute_async(
"().__class__.__bases__[0].__subclasses__()"
)
assert not result.success
@pytest.mark.asyncio
async def test_code_object_escape(self, repl_env_empty):
"""Cannot create code objects."""
result = await repl_env_empty.execute_async(
"def f(): pass\n"
"type(f.__code__)(0, 0, 0, 0, 0, 0, b'', (), (), (), '', '', 0, b'')"
)
assert not result.success
@pytest.mark.asyncio
async def test_import_system_escape(self, repl_env_empty):
"""Cannot escape through importlib."""
result = await repl_env_empty.execute_async("import importlib")
assert not result.success
@pytest.mark.asyncio
async def test_pickle_escape(self, repl_env_empty):
"""Cannot use pickle for code execution."""
result = await repl_env_empty.execute_async("import pickle")
assert not result.success