Skip to content

Commit 4b63b11

Browse files
fix(adkrest): limit /run_live WebSocket message size (#1664)
* fix(adkrest): limit /run_live WebSocket message size * Address review: 16 MiB limit, unexported constant, mechanism test Matches adk-python's uvicorn ws_max_size default (16 MiB). Unexports the constant since nothing currently overrides it. Adds a mechanism-level regression test for the SetReadLimit call: this checkout's Agent interface has no way, via its public API, to build a fake agent that supports live sessions (agent.New()'s returned type does not implement the liveAgent interface RunLive() requires), so a full RunLiveHandler integration test isn't achievable here without separate, larger live-agent test scaffolding. This test instead exercises the exact websocket.Upgrader + SetReadLimit configuration the handler uses, in isolation. * Address review: replace mechanism test with handler-level test Removes runtime_live_limit_test.go (needed only because this branch lacked live-agent test scaffolding at the time it was written) and adds TestRunLiveHandlerEnforcesMessageSizeLimit to runtime_live_test.go instead, using mockLiveAgent/dialRunLiveHandler to exercise RunLiveHandler end to end, per review. golangci-lint's two reported issues are not addressed by this commit -- see PR discussion. --------- Co-authored-by: prasanna8585 <prasanna8585@users.noreply.github.com> Co-authored-by: wolo <wolo@google.com>
1 parent 65eff6f commit 4b63b11

2 files changed

Lines changed: 52 additions & 0 deletions

File tree

‎server/adkrest/controllers/runtime.go‎

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -100,6 +100,16 @@ type RuntimeAPIControllerConfig struct {
100100
CheckOrigin func(*http.Request) bool
101101
}
102102

103+
// maxLiveMessageBytes is the read limit RunLiveHandler applies to a single
104+
// client-sent WebSocket message. Matches uvicorn's ws_max_size default,
105+
// which adk-python's dev servers (adk web, adk api_server) leave unset, so
106+
// a message either server accepts, the other does too.
107+
//
108+
// Unexported: nothing currently overrides it. If that's needed later, add
109+
// a field to RuntimeAPIControllerConfig where zero means this default, the
110+
// same way ServerConfig.MaxPayloadSize works.
111+
const maxLiveMessageBytes = 16 << 20 // 16 MiB
112+
103113
// NewRuntimeAPIController creates the controller for the Runtime API.
104114
//
105115
// Deprecated: use [NewRuntimeAPIControllerWithConfig], which does not have to
@@ -423,6 +433,9 @@ func (c *RuntimeAPIController) RunLiveHandler(rw http.ResponseWriter, req *http.
423433
_ = ws.Close()
424434
}()
425435

436+
// The upgrade bypasses MaxBytesMiddleware, and gorilla/websocket has no default limit.
437+
ws.SetReadLimit(maxLiveMessageBytes)
438+
426439
sendClose := func(code int, reason string) {
427440
_ = ws.WriteMessage(websocket.CloseMessage, websocket.FormatCloseMessage(code, truncateCloseReason(reason)))
428441
_ = ws.SetReadDeadline(time.Now().Add(time.Second))

‎server/adkrest/controllers/runtime_live_test.go‎

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -605,3 +605,42 @@ func TestTruncateCloseReason(t *testing.T) {
605605
t.Errorf("close frame payload = %d bytes, want at most 125", got)
606606
}
607607
}
608+
609+
func TestRunLiveHandlerEnforcesMessageSizeLimit(t *testing.T) {
610+
for _, tc := range []struct {
611+
name string
612+
size int
613+
forward bool
614+
}{
615+
{name: "at limit is forwarded", size: maxLiveMessageBytes, forward: true},
616+
{name: "one byte over is rejected", size: maxLiveMessageBytes + 1},
617+
} {
618+
t.Run(tc.name, func(t *testing.T) {
619+
liveSession := newRecordingLiveSession()
620+
conn, handlerDone := dialRunLiveHandler(t, func(agent.InvocationContext) (agent.LiveSession, iter.Seq2[*session.Event, error], error) {
621+
return liveSession, func(yield func(*session.Event, error) bool) { <-liveSession.closed }, nil
622+
})
623+
if err := conn.WriteMessage(websocket.BinaryMessage, make([]byte, tc.size)); err != nil {
624+
t.Fatalf("WriteMessage() failed: %v", err)
625+
}
626+
627+
if tc.forward {
628+
got := waitForLiveRequest(t, liveSession)
629+
if blob, ok := got.RealtimeInput.(*genai.Blob); !ok || len(blob.Data) != tc.size {
630+
t.Fatalf("RealtimeInput = %T, want a *genai.Blob of %d bytes", got.RealtimeInput, tc.size)
631+
}
632+
return
633+
}
634+
635+
if closeErr := readCloseError(t, conn); closeErr.Code != websocket.CloseMessageTooBig {
636+
t.Fatalf("close code = %d, want %d", closeErr.Code, websocket.CloseMessageTooBig)
637+
}
638+
select {
639+
case req := <-liveSession.requests:
640+
t.Fatalf("oversized message reached the live session: %v", req)
641+
default:
642+
}
643+
waitForHandlerExit(t, handlerDone)
644+
})
645+
}
646+
}

0 commit comments

Comments
 (0)