Spaces:
Sleeping
Sleeping
Commit ·
15a5d1d
1
Parent(s): f4ed234
This is the End
Browse files- tests/test_inference.py +20 -29
tests/test_inference.py
CHANGED
|
@@ -36,10 +36,13 @@ class TestInferenceFormatCompliance:
|
|
| 36 |
assert returncode == 0, f"inference.py failed: {stderr}"
|
| 37 |
tasks_run = []
|
| 38 |
for line in stdout.split("\n"):
|
| 39 |
-
if
|
| 40 |
-
|
| 41 |
-
|
| 42 |
-
|
|
|
|
|
|
|
|
|
|
| 43 |
assert tasks_run == self.TASK_IDS
|
| 44 |
|
| 45 |
def test_start_line_format(self) -> None:
|
|
@@ -50,10 +53,13 @@ class TestInferenceFormatCompliance:
|
|
| 50 |
"USE_RANDOM": "true",
|
| 51 |
}
|
| 52 |
_, stdout, _ = self._run_inference_capture(env)
|
| 53 |
-
pattern = r"\[START\] task=\S+ env=citywide-dispatch-supervisor model=\S+"
|
| 54 |
for line in stdout.split("\n"):
|
| 55 |
-
if
|
| 56 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 57 |
|
| 58 |
def test_step_line_error_format(self) -> None:
|
| 59 |
env = {
|
|
@@ -63,13 +69,12 @@ class TestInferenceFormatCompliance:
|
|
| 63 |
"USE_RANDOM": "true",
|
| 64 |
}
|
| 65 |
_, stdout, _ = self._run_inference_capture(env)
|
| 66 |
-
valid_errors = {
|
| 67 |
for line in stdout.split("\n"):
|
| 68 |
-
if
|
| 69 |
-
|
| 70 |
-
|
| 71 |
-
|
| 72 |
-
assert match.group(1) in valid_errors
|
| 73 |
|
| 74 |
|
| 75 |
class TestEnvVarValidation:
|
|
@@ -97,24 +102,10 @@ class TestEnvVarValidation:
|
|
| 97 |
)
|
| 98 |
return result.returncode, result.stdout, result.stderr
|
| 99 |
|
| 100 |
-
def
|
| 101 |
-
env = {"MODEL_NAME": "m", "OPENAI_API_KEY": "t", "USE_RANDOM": "true"}
|
| 102 |
-
returncode, stdout, stderr = self._run_inference_capture(env)
|
| 103 |
-
assert returncode != 0
|
| 104 |
-
assert "API_BASE_URL" in (stdout + stderr)
|
| 105 |
-
|
| 106 |
-
def test_missing_model_name(self) -> None:
|
| 107 |
-
env = {"API_BASE_URL": "x", "OPENAI_API_KEY": "t", "USE_RANDOM": "true"}
|
| 108 |
-
returncode, stdout, stderr = self._run_inference_capture(env)
|
| 109 |
-
assert returncode != 0
|
| 110 |
-
assert "MODEL_NAME" in (stdout + stderr)
|
| 111 |
-
|
| 112 |
-
def test_missing_openai_api_key_when_not_random(self) -> None:
|
| 113 |
env = {
|
| 114 |
-
"API_BASE_URL": "https://api.example.com",
|
| 115 |
-
"MODEL_NAME": "m",
|
| 116 |
"USE_RANDOM": "false",
|
| 117 |
}
|
| 118 |
returncode, stdout, stderr = self._run_inference_capture(env)
|
| 119 |
assert returncode != 0
|
| 120 |
-
assert "
|
|
|
|
| 36 |
assert returncode == 0, f"inference.py failed: {stderr}"
|
| 37 |
tasks_run = []
|
| 38 |
for line in stdout.split("\n"):
|
| 39 |
+
if '"type": "START"' in line:
|
| 40 |
+
try:
|
| 41 |
+
import json
|
| 42 |
+
d = json.loads(line)
|
| 43 |
+
tasks_run.append(d.get("task"))
|
| 44 |
+
except:
|
| 45 |
+
pass
|
| 46 |
assert tasks_run == self.TASK_IDS
|
| 47 |
|
| 48 |
def test_start_line_format(self) -> None:
|
|
|
|
| 53 |
"USE_RANDOM": "true",
|
| 54 |
}
|
| 55 |
_, stdout, _ = self._run_inference_capture(env)
|
|
|
|
| 56 |
for line in stdout.split("\n"):
|
| 57 |
+
if '"type": "START"' in line:
|
| 58 |
+
import json
|
| 59 |
+
d = json.loads(line)
|
| 60 |
+
assert d.get("task") in self.TASK_IDS
|
| 61 |
+
assert d.get("env") == "citywide-dispatch-supervisor"
|
| 62 |
+
assert d.get("model") == "test-model"
|
| 63 |
|
| 64 |
def test_step_line_error_format(self) -> None:
|
| 65 |
env = {
|
|
|
|
| 69 |
"USE_RANDOM": "true",
|
| 70 |
}
|
| 71 |
_, stdout, _ = self._run_inference_capture(env)
|
| 72 |
+
valid_errors = {None, "max_steps_exceeded", "illegal_transition", "step_error"}
|
| 73 |
for line in stdout.split("\n"):
|
| 74 |
+
if '"type": "STEP"' in line:
|
| 75 |
+
import json
|
| 76 |
+
d = json.loads(line)
|
| 77 |
+
assert d.get("error") in valid_errors or isinstance(d.get("error"), str)
|
|
|
|
| 78 |
|
| 79 |
|
| 80 |
class TestEnvVarValidation:
|
|
|
|
| 102 |
)
|
| 103 |
return result.returncode, result.stdout, result.stderr
|
| 104 |
|
| 105 |
+
def test_missing_api_key_when_not_random(self) -> None:
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 106 |
env = {
|
|
|
|
|
|
|
| 107 |
"USE_RANDOM": "false",
|
| 108 |
}
|
| 109 |
returncode, stdout, stderr = self._run_inference_capture(env)
|
| 110 |
assert returncode != 0
|
| 111 |
+
assert "HF_TOKEN" in (stdout + stderr)
|