fix(grayswan): omit tools from post-call monitor when request context is empty

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-29 21:27:30 +00:00
parent e5fac49e7b
commit 5b74a0375b
3 changed files with 30 additions and 6 deletions

View file

@ -654,6 +654,8 @@ class GraySwanGuardrail(CustomGuardrail):
skip_system=effective_skip_system_message_for_guardrail(self),
skip_tool=effective_skip_tool_message_for_guardrail(self),
)
if not indices:
return (), None
raw_tools: Final = request_data.get("tools")
tools: Final = (
tuple(raw_tools) if not scan_only_tool_results and isinstance(raw_tools, list) and raw_tools else None

View file

@ -119,15 +119,19 @@ def _chat_provider(message: dict[str, JsonValue]):
def _monitor_bodies(vendor: Wire, expected: int = 1) -> tuple[dict[str, JsonValue], ...]:
collected: list[dict[str, JsonValue]] = [] # mutable-ok: accumulator across polling attempts
collected: tuple[dict[str, JsonValue], ...] = ()
def drain_new() -> tuple[dict[str, JsonValue], ...]:
collected.extend(
_JSON_OBJECT.validate_json(request.body)
for request in vendor.drain()
if request.target == "/cygnal/monitor"
nonlocal collected
collected = ( # rebind-ok: eventually polls this closure, so drained bodies must persist across calls
*collected,
*(
_JSON_OBJECT.validate_json(request.body)
for request in vendor.drain()
if request.target == "/cygnal/monitor"
),
)
return tuple(collected)
return collected
return eventually(drain_new, lambda bodies: len(bodies) >= expected, seconds=30)

View file

@ -771,6 +771,24 @@ async def test_post_call_prefers_request_route_over_logging_call_type() -> None:
assert list(payload["tools"]) == _REQUEST_DATA["tools"]
@pytest.mark.asyncio
async def test_post_call_surface_without_messages_sends_response_only() -> None:
guardrail = _post_call_guardrail()
client = _CapturingClient()
guardrail.async_handler = client
await guardrail.apply_guardrail(
inputs={"texts": ["response text"]},
request_data={**_REQUEST_DATA, "litellm_logging_obj": _LoggingObj("aembedding")},
input_type="response",
logging_obj=_LoggingObj("aembedding"),
)
payload = client.calls[0]["json"]
assert list(payload["messages"]) == [{"role": "assistant", "content": "response text"}]
assert "tools" not in payload
@pytest.mark.asyncio
async def test_post_call_unresolvable_call_type_sends_response_only() -> None:
guardrail = _post_call_guardrail()