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
12 changes: 11 additions & 1 deletion src/cmcp_gateway/mcp/proxy.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,8 @@

from __future__ import annotations

import hashlib
import json
import logging
from dataclasses import dataclass
from datetime import UTC, datetime
Expand Down Expand Up @@ -94,6 +96,7 @@ def _build_cedar_context(
entry = self._catalog.lookup(tool_name)
return {
"tool_name": tool_name,
"arguments": arguments,
"server_identity": entry.server.url if entry else "",
"compliance_domain": entry.compliance_domain if entry else "external",
"baa_covered": (not entry.requires_baa) if entry else False,
Expand Down Expand Up @@ -166,6 +169,9 @@ async def call_tool(
sensitivity_before = self._session.max_sensitivity
would_have_denied = False

_payload_bytes = json.dumps(arguments, sort_keys=True, separators=(",", ":")).encode()
request_payload_hash = f"sha256:{hashlib.sha256(_payload_bytes).hexdigest()}"

# Step 1: catalog lookup
entry = self._catalog.lookup(tool_name)
if entry is None:
Expand All @@ -177,7 +183,7 @@ async def call_tool(
server_identity=None,
policy_decision="deny",
policy_rule_matched="catalog_miss",
request_payload_hash=None,
request_payload_hash=request_payload_hash,
session_sensitivity_before=sensitivity_before,
session_sensitivity_after=self._session.max_sensitivity,
)
Expand Down Expand Up @@ -216,6 +222,7 @@ async def call_tool(
server_identity=entry.server.url,
policy_decision="deny",
policy_rule_matched=str(exc),
request_payload_hash=request_payload_hash,
session_sensitivity_before=sensitivity_before,
session_sensitivity_after=self._session.max_sensitivity,
)
Expand Down Expand Up @@ -258,6 +265,7 @@ async def call_tool(
server_identity=entry.server.url,
policy_decision="deny",
policy_rule_matched=f"agt_gateway:{type(exc).__name__}",
request_payload_hash=request_payload_hash,
session_sensitivity_before=sensitivity_before,
session_sensitivity_after=self._session.max_sensitivity,
)
Expand Down Expand Up @@ -315,6 +323,7 @@ async def call_tool(
server_identity=entry.server.url,
policy_decision="deny",
policy_rule_matched=egress_deny_reason,
request_payload_hash=request_payload_hash,
session_sensitivity_before=sensitivity_before,
session_sensitivity_after=self._session.max_sensitivity,
)
Expand Down Expand Up @@ -343,6 +352,7 @@ async def call_tool(
policy_decision=policy_decision,
policy_rule_matched=policy_rule,
latency_us=latency_us,
request_payload_hash=request_payload_hash,
session_sensitivity_before=sensitivity_before,
session_sensitivity_after=self._session.max_sensitivity,
)
Expand Down
61 changes: 61 additions & 0 deletions tests/unit/test_mcp_proxy.py
Original file line number Diff line number Diff line change
Expand Up @@ -192,3 +192,64 @@ async def test_proxy_result_contains_audit_entry_hash():
proxy, _, chain = _make_proxy()
result = await proxy.call_tool("c1", "test.tool", {})
assert result.audit_entry_hash == chain.chain_tip


# ── POLICY-004: Cedar context includes arguments ──────────────────────────────

@pytest.mark.asyncio
async def test_cedar_context_includes_arguments():
"""POLICY-004 — arguments must appear in the Cedar context so policies can inspect them."""
evaluator = _make_evaluator()
proxy, _, _ = _make_proxy(evaluator=evaluator)
args = {"patient_id": "p-123", "action": "read"}
await proxy.call_tool("c1", "test.tool", args)
ctx = evaluator.evaluate.call_args[0][0]
assert ctx["arguments"] == args


# ── POLICY-005: request_payload_hash in all audit entries ────────────────────

@pytest.mark.asyncio
async def test_audit_payload_hash_present_on_allow():
"""POLICY-005 — successful calls must record request_payload_hash in the audit entry."""
proxy, _, chain = _make_proxy()
await proxy.call_tool("c1", "test.tool", {"k": "v"})
entry = next(e for e in reversed(chain.entries) if e.entry_type == "tool_call")
assert entry.request_payload_hash is not None
assert entry.request_payload_hash.startswith("sha256:")


@pytest.mark.asyncio
async def test_audit_payload_hash_present_on_catalog_deny():
"""POLICY-005 — catalog-miss denials must record request_payload_hash."""
proxy, _, chain = _make_proxy()
await proxy.call_tool("c1", "ghost.tool", {"x": 1})
entry = next(e for e in reversed(chain.entries) if e.entry_type == "tool_call")
assert entry.request_payload_hash is not None
assert entry.request_payload_hash.startswith("sha256:")


@pytest.mark.asyncio
async def test_audit_payload_hash_present_on_cedar_deny():
"""POLICY-005 — Cedar policy denials must record request_payload_hash."""
evaluator = _make_evaluator(allow=False)
proxy, _, chain = _make_proxy(evaluator=evaluator)
await proxy.call_tool("c1", "test.tool", {"secret": "leak"})
entry = next(e for e in reversed(chain.entries) if e.entry_type == "tool_call")
assert entry.request_payload_hash is not None
assert entry.request_payload_hash.startswith("sha256:")


@pytest.mark.asyncio
async def test_audit_payload_hash_is_canonical_sha256():
"""POLICY-005 — payload hash must be sha256 of canonical JSON (sort_keys, no spaces)."""
import hashlib
import json

proxy, _, chain = _make_proxy()
args = {"b": 2, "a": 1}
await proxy.call_tool("c1", "test.tool", args)
entry = next(e for e in reversed(chain.entries) if e.entry_type == "tool_call")
expected_bytes = json.dumps(args, sort_keys=True, separators=(",", ":")).encode()
expected = f"sha256:{hashlib.sha256(expected_bytes).hexdigest()}"
assert entry.request_payload_hash == expected
Loading