Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
86 changes: 85 additions & 1 deletion projects/egress-gate/tests/gates/test_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,8 +16,16 @@
GateResources,
RegexConfig,
RegexGate,
Utf8BodyGate,
)
from egress_gate.request import (
ExistingHeaderAction,
HttpRequest,
HttpTarget,
RequestContext,
RequestMutations,
WriteHeaderMutation,
)
from egress_gate.request import HttpRequest, HttpTarget, RequestContext
from egress_gate.result import Finding, GateControl, GateEvaluation
from egress_gate.timeout import Timeout

Expand Down Expand Up @@ -107,6 +115,63 @@ def _evaluate(
return GateEvaluation.proceed().model_copy(update={"control": GateControl.DENY})


class _UndeclaredBodyMutationGate(Gate[_RequestConfig, None]):
capabilities = frozenset()
finding_types = ()

def _evaluate(
self,
request: HttpRequest,
*,
timeout: Timeout,
) -> GateEvaluation:
del request, timeout
return GateEvaluation.proceed(
request_mutations=RequestMutations(replacement_body=b"changed")
)


class _UndeclaredHeaderMutationGate(Gate[_RequestConfig, None]):
capabilities = frozenset()
finding_types = ()

def _evaluate(
self,
request: HttpRequest,
*,
timeout: Timeout,
) -> GateEvaluation:
del request, timeout
return GateEvaluation.proceed(
request_mutations=RequestMutations(
header_mutations=(
WriteHeaderMutation(
kind="write",
name="x-openshell-middleware-test",
value="changed",
on_existing=ExistingHeaderAction.OVERWRITE,
),
)
)
)


class _InvalidUtf8ReplacementGate(Utf8BodyGate[_RequestConfig, None]):
capabilities = frozenset({GateCapability.READ_BODY, GateCapability.REPLACE_BODY})
finding_types = ()

def _evaluate_text(
self,
text: str,
*,
timeout: Timeout,
) -> GateEvaluation:
del text, timeout
return GateEvaluation.proceed(
request_mutations=RequestMutations(replacement_body=b"\xff")
)


def _request(*, body: bytes = b"payload", host: str = "example.com") -> HttpRequest:
return HttpRequest(
context=RequestContext(request_id="request-1", sandbox_id="sandbox-1"),
Expand Down Expand Up @@ -145,13 +210,32 @@ def test_gate_public_wrapper_enforces_declared_output_capabilities() -> None:
_RequestConfig(name="test", kind="test-request"), None
).evaluate(_request(), timeout=Timeout.from_seconds(1))

with pytest.raises(GateContractError, match="undeclared body replacement"):
_UndeclaredBodyMutationGate(
_RequestConfig(name="test", kind="test-request"), None
).evaluate(_request(), timeout=Timeout.from_seconds(1))

with pytest.raises(GateContractError, match="undeclared header mutations"):
_UndeclaredHeaderMutationGate(
_RequestConfig(name="test", kind="test-request"), None
).evaluate(_request(), timeout=Timeout.from_seconds(1))

with pytest.raises(GateContractError, match="undeclared finding"):
_CapabilityBypassGate(
_RequestConfig(name="test", kind="test-request"),
None,
).evaluate(_request(), timeout=Timeout.from_seconds(1))


def test_utf8_body_gate_rejects_a_non_utf8_replacement() -> None:
gate = _InvalidUtf8ReplacementGate(
_RequestConfig(name="test", kind="test-request"), None
)

with pytest.raises(GateContractError, match="non-UTF-8 replacement"):
gate.evaluate(_request(), timeout=Timeout.from_seconds(1))


def test_gate_public_wrapper_classifies_invalid_models_as_contract_errors() -> None:
with pytest.raises(GateContractError, match="gate output is invalid"):
_InvalidEvaluationGate(
Expand Down
65 changes: 65 additions & 0 deletions projects/egress-gate/tests/gates/test_regex.py
Original file line number Diff line number Diff line change
Expand Up @@ -417,6 +417,71 @@ def test_replacement_selects_ranked_non_overlapping_winners() -> None:
assert len(evaluation.findings) == 2


@pytest.mark.parametrize(
("rules", "text", "expected"),
[
(
[
{"name": "short", "pattern": "ab", "confidence": "high"},
{"name": "long", "pattern": "abc", "confidence": "high"},
],
"abc",
b"<token>",
),
(
[
{"name": "earlier", "pattern": "ab", "confidence": "high"},
{"name": "later", "pattern": "ba", "confidence": "high"},
],
"aba",
b"<token>a",
),
],
)
def test_equal_confidence_overlap_prefers_length_then_start(
rules: list[dict[str, object]],
text: str,
expected: bytes,
) -> None:
evaluation = _run(
_config(rules, action_kind="replace", template="<{entity}>"),
text,
)

assert evaluation.request_mutations.replacement_body == expected
assert evaluation.findings[0].count == 2


def test_identical_overlap_uses_entity_name_as_a_stable_tie_breaker() -> None:
config = RegexConfig.model_validate(
{
"name": "regex",
"kind": "regex",
"scan": {
"kind": "body",
"action": {"kind": "replace", "template": "<{entity}>"},
},
"pattern_catalog": {
"entities": [
{
"name": "zeta",
"rules": [{"pattern": "abc", "confidence": "high"}],
},
{
"name": "alpha",
"rules": [{"pattern": "abc", "confidence": "high"}],
},
]
},
}
)

evaluation = _run(config, "abc")

assert evaluation.request_mutations.replacement_body == b"<alpha>"
assert [finding.label for finding in evaluation.findings] == ["alpha", "zeta"]


@pytest.mark.parametrize(
"template",
[
Expand Down
51 changes: 51 additions & 0 deletions projects/egress-gate/tests/service/test_grpc_integration.py
Original file line number Diff line number Diff line change
Expand Up @@ -71,6 +71,38 @@ def _evaluation(
)


def _progressive_redaction_config() -> Message:
values = {
"gates": [
{
"name": name,
"kind": "regex",
"scan": {
"kind": "body",
"action": {"kind": "replace", "template": "[{entity}]"},
},
"pattern_catalog": {
"entities": [
{
"name": entity,
"rules": [{"pattern": pattern, "confidence": "high"}],
}
]
},
}
for name, entity, pattern in (
("redact-email", "email", r"alice@example\.com"),
("redact-api-key", "api_key", r"sk-[0-9]+"),
("redact-phone", "phone", r"555-[0-9]{4}"),
)
],
"default_decision": "allow",
}
request = pb2.ValidateConfigRequest()
json_format.ParseDict(values, request.config)
return request.config


@asynccontextmanager
async def _running_stub(
middleware: EgressGateMiddleware,
Expand Down Expand Up @@ -118,6 +150,25 @@ async def test_generated_stub_round_trip_covers_manifest_and_gate_actions() -> N
assert denied.reason_code == "egress_gate_regex_denied"


@pytest.mark.asyncio
async def test_generated_stub_returns_three_gate_progressive_redaction() -> None:
middleware = EgressGateMiddleware(create_builtin_registry())
request = _evaluation(b"email=alice@example.com api_key=sk-123456 phone=555-0100")
request.config.CopyFrom(_progressive_redaction_config())

async with _running_stub(middleware) as (stub, _):
response = await stub.EvaluateHttpRequest(request)

assert response.decision == pb2.DECISION_ALLOW
assert response.has_body is True
assert response.body == b"email=[email] api_key=[api_key] phone=[phone]"
assert [finding.label for finding in response.findings] == [
"email",
"api_key",
"phone",
]


@pytest.mark.asyncio
async def test_generated_stub_maps_invalid_phase_to_invalid_argument() -> None:
middleware = EgressGateMiddleware(create_builtin_registry())
Expand Down
29 changes: 29 additions & 0 deletions projects/egress-gate/tests/test_request.py
Original file line number Diff line number Diff line change
Expand Up @@ -134,11 +134,30 @@ def test_request_mutations_preserve_ordered_discriminated_header_mutations() ->

def test_request_mutations_reject_invalid_bounds() -> None:
mutation = RemoveHeaderMutation(kind="remove", name="x-test")
assert (
len(
RequestMutations(
header_mutations=tuple(mutation for _ in range(MAX_HEADER_MUTATIONS))
).header_mutations
)
== MAX_HEADER_MUTATIONS
)
with pytest.raises(ValidationError):
RequestMutations(
header_mutations=tuple(mutation for _ in range(MAX_HEADER_MUTATIONS + 1))
)

exact_data = RequestMutations(
header_mutations=(
WriteHeaderMutation(
kind="write",
name="x",
value="x" * (MAX_HEADER_MUTATION_DATA_BYTES - 1),
on_existing=ExistingHeaderAction.OVERWRITE,
),
)
)
assert exact_data.header_mutations
with pytest.raises(ValidationError):
RequestMutations(
header_mutations=(
Expand All @@ -151,6 +170,16 @@ def test_request_mutations_reject_invalid_bounds() -> None:
)
)

assert (
len(
RequestMutations(replacement_body=b"x" * MAX_BODY_BYTES).replacement_body
or b""
)
== MAX_BODY_BYTES
)
with pytest.raises(ValidationError):
RequestMutations(replacement_body=b"x" * (MAX_BODY_BYTES + 1))


def test_request_models_reject_non_tuple_sequences_and_extra_fields() -> None:
values: dict[str, object] = {
Expand Down
Loading