Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 30d74cbf2b |
@@ -277,6 +277,7 @@ class _SuperviseRpcFake:
|
||||
|
||||
supervise_status: "str | None" = None
|
||||
propose_error: bool = False
|
||||
poll_error: bool = False
|
||||
|
||||
@property
|
||||
def propose_calls(self) -> list[dict[str, str]]:
|
||||
@@ -300,6 +301,8 @@ class _SuperviseRpcFake:
|
||||
self, source_ip: str, identity_token: str, proposal_id: str,
|
||||
) -> dict[str, object]:
|
||||
del source_ip, identity_token, proposal_id
|
||||
if self.poll_error:
|
||||
raise PolicyResolveError("orchestrator down")
|
||||
if self.supervise_status is None:
|
||||
return {"status": "pending"}
|
||||
return {"status": self.supervise_status, "notes": "", "final_file": None}
|
||||
@@ -586,6 +589,16 @@ class TestSuperviseBranch(unittest.TestCase):
|
||||
self.assertEqual(403, flow.response.status_code)
|
||||
self.assertIn("timed out", flow.response.get_text())
|
||||
|
||||
def test_poll_error_during_wait_times_out_and_blocks(self) -> None:
|
||||
# A transient orchestrator error on each poll is retried until the
|
||||
# deadline, then fails closed (blocked) — never forwarded unsupervised.
|
||||
addon = self._supervised_addon("approved")
|
||||
cast(Any, addon._resolver).poll_error = True
|
||||
flow = _Flow(_Request(host="api.example.com", method="POST", body=f"k={_OPENAI_KEY}"))
|
||||
_run_request(addon, flow)
|
||||
assert flow.response is not None
|
||||
self.assertEqual(403, flow.response.status_code)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Inbound DLP on responses
|
||||
|
||||
@@ -533,6 +533,41 @@ class TestDispatchSuperviseAgentRpc(unittest.TestCase):
|
||||
self.assertEqual(200, status)
|
||||
self.assertEqual("unknown", poll["status"])
|
||||
|
||||
# --- request validation (400) + fail-closed (403) ----------------------
|
||||
|
||||
def test_propose_invalid_json_is_400(self) -> None:
|
||||
status, _ = dispatch(self.orch, "POST", "/supervise/propose", b"{not json")
|
||||
self.assertEqual(400, status)
|
||||
|
||||
def test_propose_missing_fields_are_400(self) -> None:
|
||||
rec = self._register()
|
||||
base = {
|
||||
"source_ip": rec.source_ip, "identity_token": rec.identity_token,
|
||||
"tool": TOOL_EGRESS_ALLOW, "proposed_file": "routes:\n", "justification": "j",
|
||||
}
|
||||
for drop in ("source_ip", "proposed_file", "justification"):
|
||||
body = {k: v for k, v in base.items() if k != drop}
|
||||
status, _ = dispatch(self.orch, "POST", "/supervise/propose", _body(body))
|
||||
self.assertEqual(400, status, drop)
|
||||
|
||||
def test_poll_invalid_json_is_400(self) -> None:
|
||||
status, _ = dispatch(self.orch, "POST", "/supervise/poll", b"{not json")
|
||||
self.assertEqual(400, status)
|
||||
|
||||
def test_poll_missing_fields_are_400(self) -> None:
|
||||
rec = self._register()
|
||||
for body in ({"identity_token": rec.identity_token, "proposal_id": "p"},
|
||||
{"source_ip": rec.source_ip, "identity_token": rec.identity_token}):
|
||||
status, _ = dispatch(self.orch, "POST", "/supervise/poll", _body(body))
|
||||
self.assertEqual(400, status)
|
||||
|
||||
def test_poll_unattributed_is_403(self) -> None:
|
||||
status, payload = dispatch(self.orch, "POST", "/supervise/poll", _body({
|
||||
"source_ip": "10.9.9.9", "identity_token": "wrong", "proposal_id": "p",
|
||||
}))
|
||||
self.assertEqual(403, status)
|
||||
self.assertIn("unattributed", str(payload["error"]))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -65,10 +65,16 @@ class _FakeSuperviseResolver:
|
||||
|
||||
def __init__(
|
||||
self, bottle_id: str | None = "dev", raises: bool = False, policy: str = "",
|
||||
poll_raises: bool = False, poll_none: bool = False,
|
||||
) -> None:
|
||||
self.bottle_id = bottle_id
|
||||
self.raises = raises
|
||||
self._policy = policy
|
||||
# propose succeeds, but the later poll fails — models an orchestrator
|
||||
# that becomes unreachable (poll_raises) or unattributes mid-flight
|
||||
# (poll_none) after the proposal was queued.
|
||||
self.poll_raises = poll_raises
|
||||
self.poll_none = poll_none
|
||||
|
||||
def propose_supervise(
|
||||
self, source_ip: str, identity_token: str, *,
|
||||
@@ -90,9 +96,9 @@ class _FakeSuperviseResolver:
|
||||
self, source_ip: str, identity_token: str, proposal_id: str,
|
||||
) -> dict[str, object] | None:
|
||||
del source_ip, identity_token
|
||||
if self.raises:
|
||||
if self.raises or self.poll_raises:
|
||||
raise supervise_server.PolicyResolveError("orchestrator down")
|
||||
if self.bottle_id is None:
|
||||
if self.bottle_id is None or self.poll_none:
|
||||
return None
|
||||
try:
|
||||
response = _sv.read_response(self.bottle_id, proposal_id)
|
||||
@@ -500,6 +506,25 @@ class TestHandleToolsCall(unittest.TestCase):
|
||||
self.assertIn("proposal remains queued", text)
|
||||
self.assertEqual(1, len(_sv.list_pending_proposals("dev")))
|
||||
|
||||
_ALLOW: dict[str, object] = {
|
||||
"name": _sv.TOOL_EGRESS_ALLOW,
|
||||
"arguments": {"routes_yaml": "routes:\n - host: example.com\n", "justification": "x"},
|
||||
}
|
||||
|
||||
def test_poll_unreachable_after_queue_raises_internal(self):
|
||||
# Proposal queues, then the orchestrator becomes unreachable on poll.
|
||||
self.resolver.poll_raises = True
|
||||
with self.assertRaises(_RpcInternalError) as cm:
|
||||
_tools_call(self.resolver, self._ALLOW, ServerConfig(response_timeout_seconds=5))
|
||||
self.assertEqual(ERR_INTERNAL, cm.exception.code)
|
||||
|
||||
def test_poll_unattributed_after_queue_raises_internal(self):
|
||||
# Proposal queues, then poll comes back unattributed (fail-closed).
|
||||
self.resolver.poll_none = True
|
||||
with self.assertRaises(_RpcInternalError) as cm:
|
||||
_tools_call(self.resolver, self._ALLOW, ServerConfig(response_timeout_seconds=5))
|
||||
self.assertEqual(ERR_INTERNAL, cm.exception.code)
|
||||
|
||||
|
||||
class TestResponseTimeoutEnv(unittest.TestCase):
|
||||
def test_unset_uses_default(self):
|
||||
@@ -783,6 +808,18 @@ class TestNonBlockingSupervise(unittest.TestCase):
|
||||
self.assertTrue(result["isError"])
|
||||
self.assertIn("status: unknown", result["content"][0]["text"]) # type: ignore[index]
|
||||
|
||||
def test_check_poll_unreachable_raises_internal(self):
|
||||
self.resolver.poll_raises = True
|
||||
with self.assertRaises(_RpcInternalError) as cm:
|
||||
_check(self.resolver, {"arguments": {"proposal_id": "p"}})
|
||||
self.assertEqual(ERR_INTERNAL, cm.exception.code)
|
||||
|
||||
def test_check_poll_unattributed_raises_internal(self):
|
||||
self.resolver.poll_none = True
|
||||
with self.assertRaises(_RpcInternalError) as cm:
|
||||
_check(self.resolver, {"arguments": {"proposal_id": "p"}})
|
||||
self.assertEqual(ERR_INTERNAL, cm.exception.code)
|
||||
|
||||
def test_check_missing_id_raises(self):
|
||||
with self.assertRaises(_RpcClientError) as cm:
|
||||
_check(self.resolver, {"arguments": {}})
|
||||
|
||||
Reference in New Issue
Block a user