forked from UKGovernmentBEIS/inspect_ai
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtest_extensions.py
More file actions
135 lines (107 loc) · 4.51 KB
/
Copy pathtest_extensions.py
File metadata and controls
135 lines (107 loc) · 4.51 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
import importlib
import pytest
from pydantic_core import to_jsonable_python
from test_helpers.tools import list_files
from test_helpers.utils import ensure_test_package_installed, skip_if_trio
from inspect_ai import Task, eval_async
from inspect_ai._util.registry import _registry
from inspect_ai.dataset import Sample
from inspect_ai.hooks import Hooks
from inspect_ai.hooks import hooks as register_hooks
from inspect_ai.hooks._hooks import emit_run_start
from inspect_ai.model import ChatMessageUser, GenerateConfig, get_model
from inspect_ai.scorer import includes
from inspect_ai.solver import generate, use_tools
from inspect_ai.util import SandboxEnvironmentSpec
@skip_if_trio
async def test_extension_model():
# ensure the package is installed
ensure_test_package_installed()
# call the model
mdl = get_model("custom/gpt7")
result = await mdl.generate(
[ChatMessageUser(content="hello")], [], "none", GenerateConfig()
)
assert result.completion == "Hello from gpt7"
@pytest.mark.anyio
async def test_extension_sandboxenv():
# ensure the package is installed
ensure_test_package_installed()
# run a task using the sandboxenv
task = Task(
dataset=[
Sample(
input="Please use the list_files tool to list the files in the current directory"
)
],
solver=[use_tools(list_files()), generate()],
scorer=includes(),
sandbox="podman",
)
await eval_async(task, model="mockllm/model")
@pytest.mark.slow
@pytest.mark.anyio
async def test_extension_sandboxenv_with_specialised_config():
# ensure the package is installed
ensure_test_package_installed()
module = importlib.import_module("inspect_package.sandboxenv.podman")
PodmanSandboxEnvironmentConfig = module.PodmanSandboxEnvironmentConfig
# run a task using the sandboxenv
task = Task(
dataset=[
Sample(
input="Please use the list_files tool to list the files in the current directory"
)
],
solver=[use_tools(list_files()), generate()],
scorer=includes(),
sandbox=SandboxEnvironmentSpec(
"podman", PodmanSandboxEnvironmentConfig(socket_path="/path/to/socket")
),
)
logs = await eval_async(task, model="mockllm/model")
# Ensure that the PodmanSandboxEnvironmentConfig object is serializable.
to_jsonable_python(logs[0].eval, exclude_none=True, fallback=lambda _x: None)
def test_can_roundtrip_specialised_config():
ensure_test_package_installed()
module = importlib.import_module("inspect_package.sandboxenv.podman")
PodmanSandboxEnvironmentConfig = module.PodmanSandboxEnvironmentConfig
# Historical issue: the SandboxEnvironmentSpec type was unable to determine which
# sandbox-specific config Pydantic model to instantiate when deserializing from
# JSON.
spec = SandboxEnvironmentSpec(
type="podman",
config=PodmanSandboxEnvironmentConfig(socket_path="/path/to/socket"),
)
json_str = spec.model_dump_json()
recreated = SandboxEnvironmentSpec.model_validate_json(json_str)
assert recreated == spec
assert recreated.config == spec.config
assert isinstance(recreated.config, PodmanSandboxEnvironmentConfig)
def test_can_load_log_file_for_unavailable_sandbox_environment():
json_str = """{"type":"unavailable","config":{"key":"value"}}"""
recreated = SandboxEnvironmentSpec.model_validate_json(json_str)
assert isinstance(recreated.config, dict)
def test_supports_str_config():
spec = SandboxEnvironmentSpec(type="podman", config="/path/to/socket")
json_str = spec.model_dump_json()
recreated = SandboxEnvironmentSpec.model_validate_json(json_str)
assert recreated == spec
assert recreated.config == spec.config
assert isinstance(recreated.config, str)
async def test_hooks():
# Pre-register a hook so the test exercises the case where
# `emit_run_start` runs while the registry already contains entries
# from the current process. Order-of-iteration regressions in the hook
# registry would surface here.
@register_hooks(name="preexisting_test_hook", description="test")
class PreexistingHook(Hooks):
pass
ensure_test_package_installed()
module = importlib.import_module("inspect_package.hooks.custom")
module.run_ids = []
try:
await emit_run_start(eval_set_id=None, run_id="42", tasks=[])
finally:
_registry.pop("hooks:preexisting_test_hook", None)
assert module.run_ids == ["42"]