fix(security): scope MCP read tools to the requesting user's filter (GHSA-p66q-2gfp-8v55) (#4864)

Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
Amir Raminfar
2026-07-24 16:43:21 -07:00
committed by GitHub
parent 9c665d42a9
commit df5a60479c
3 changed files with 60 additions and 4 deletions
+7 -1
View File
@@ -15,6 +15,12 @@ type contextKey string
const remoteUser contextKey = "remoteUser" const remoteUser contextKey = "remoteUser"
// WithUser returns a copy of ctx carrying user so downstream handlers can
// resolve it via UserFromContext.
func WithUser(ctx context.Context, user User) context.Context {
return context.WithValue(ctx, remoteUser, user)
}
type proxyAuthContext struct { type proxyAuthContext struct {
headerUser string headerUser string
headerEmail string headerEmail string
@@ -55,7 +61,7 @@ func (p *proxyAuthContext) AuthMiddleware(next http.Handler) http.Handler {
userRoles = ParseRole(r.Header.Get(p.headerRoles)) userRoles = ParseRole(r.Header.Get(p.headerRoles))
} }
user := newUser(r.Header.Get(p.headerUser), r.Header.Get(p.headerEmail), r.Header.Get(p.headerName), containerFilter, userRoles) user := newUser(r.Header.Get(p.headerUser), r.Header.Get(p.headerEmail), r.Header.Get(p.headerName), containerFilter, userRoles)
ctx := context.WithValue(r.Context(), remoteUser, user) ctx := WithUser(r.Context(), user)
next.ServeHTTP(w, r.WithContext(ctx)) next.ServeHTTP(w, r.WithContext(ctx))
} else { } else {
next.ServeHTTP(w, r) next.ServeHTTP(w, r)
+17 -3
View File
@@ -8,6 +8,7 @@ import (
"strings" "strings"
"time" "time"
"github.com/amir20/dozzle/internal/auth"
"github.com/amir20/dozzle/internal/container" "github.com/amir20/dozzle/internal/container"
container_support "github.com/amir20/dozzle/internal/support/container" container_support "github.com/amir20/dozzle/internal/support/container"
"github.com/modelcontextprotocol/go-sdk/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
@@ -48,6 +49,19 @@ func NewServer(hostService HostService, labels container.ContainerLabels, versio
return s return s
} }
// resolveLabels returns the container label filter to apply for the requesting
// user. When authentication is enabled the user's own filter is carried on the
// request context; using it (instead of the server-global labels) keeps MCP
// reads scoped to the same containers the user can see through the normal log
// and stats paths. Falls back to the server-global labels when no per-user
// filter is present (e.g. auth disabled).
func (s *Server) resolveLabels(ctx context.Context) container.ContainerLabels {
if user := auth.UserFromContext(ctx); user != nil && user.ContainerLabels.Exists() {
return user.ContainerLabels
}
return s.labels
}
// Handler returns an http.Handler for the MCP streamable HTTP transport. // Handler returns an http.Handler for the MCP streamable HTTP transport.
func (s *Server) Handler() http.Handler { func (s *Server) Handler() http.Handler {
return mcp.NewStreamableHTTPHandler(func(r *http.Request) *mcp.Server { return mcp.NewStreamableHTTPHandler(func(r *http.Request) *mcp.Server {
@@ -117,7 +131,7 @@ func (s *Server) registerTools() {
// --- Tool Handlers --- // --- Tool Handlers ---
func (s *Server) handleListContainers(ctx context.Context, _ *mcp.CallToolRequest, params *listContainersParams) (*mcp.CallToolResult, any, error) { func (s *Server) handleListContainers(ctx context.Context, _ *mcp.CallToolRequest, params *listContainersParams) (*mcp.CallToolResult, any, error) {
containers, errs := s.hostService.ListAllContainers(s.labels) containers, errs := s.hostService.ListAllContainers(s.resolveLabels(ctx))
for _, err := range errs { for _, err := range errs {
if err != nil { if err != nil {
log.Warn().Err(err).Msg("partial failure listing containers from a host") log.Warn().Err(err).Msg("partial failure listing containers from a host")
@@ -216,7 +230,7 @@ func parseStream(stream *string) (container.StdType, *mcp.CallToolResult) {
// sinceMinutes (defaulting to 5). On a user-facing failure it returns a non-nil // sinceMinutes (defaulting to 5). On a user-facing failure it returns a non-nil
// error result; otherwise the caller must call cancel when done. // error result; otherwise the caller must call cancel when done.
func (s *Server) fetchLogs(ctx context.Context, host, containerID string, stream *string, sinceMinutes *int) (<-chan *container.LogEvent, context.CancelFunc, *mcp.CallToolResult) { func (s *Server) fetchLogs(ctx context.Context, host, containerID string, stream *string, sinceMinutes *int) (<-chan *container.LogEvent, context.CancelFunc, *mcp.CallToolResult) {
containerSvc, err := s.hostService.FindContainer(host, containerID, s.labels) containerSvc, err := s.hostService.FindContainer(host, containerID, s.resolveLabels(ctx))
if err != nil { if err != nil {
return nil, nil, errorResult(fmt.Sprintf("container not found: %v", err)) return nil, nil, errorResult(fmt.Sprintf("container not found: %v", err))
} }
@@ -435,7 +449,7 @@ func (s *Server) handleGetContainerStats(ctx context.Context, _ *mcp.CallToolReq
}, nil, nil }, nil, nil
} }
containerSvc, err := s.hostService.FindContainer(params.Host, params.ContainerID, s.labels) containerSvc, err := s.hostService.FindContainer(params.Host, params.ContainerID, s.resolveLabels(ctx))
if err != nil { if err != nil {
return &mcp.CallToolResult{ return &mcp.CallToolResult{
Content: []mcp.Content{&mcp.TextContent{Text: fmt.Sprintf("container not found: %v", err)}}, Content: []mcp.Content{&mcp.TextContent{Text: fmt.Sprintf("container not found: %v", err)}},
+36
View File
@@ -8,6 +8,7 @@ import (
"testing" "testing"
"time" "time"
"github.com/amir20/dozzle/internal/auth"
"github.com/amir20/dozzle/internal/container" "github.com/amir20/dozzle/internal/container"
container_support "github.com/amir20/dozzle/internal/support/container" container_support "github.com/amir20/dozzle/internal/support/container"
"github.com/amir20/dozzle/internal/utils" "github.com/amir20/dozzle/internal/utils"
@@ -23,6 +24,10 @@ type mockHostService struct {
listErrs []error listErrs []error
logEvents []*container.LogEvent logEvents []*container.LogEvent
logErr error logErr error
// gotLabels records the last label filter passed to a lookup so tests can
// assert the requesting user's filter is applied instead of the global one.
gotLabels container.ContainerLabels
} }
type stubClientService struct { type stubClientService struct {
@@ -54,6 +59,7 @@ func (s *stubClientService) RawLogs(context.Context, container.Container, time.T
} }
func (m *mockHostService) FindContainer(host string, id string, labels container.ContainerLabels) (*container_support.ContainerService, error) { func (m *mockHostService) FindContainer(host string, id string, labels container.ContainerLabels) (*container_support.ContainerService, error) {
m.gotLabels = labels
if m.findErr != nil { if m.findErr != nil {
return nil, m.findErr return nil, m.findErr
} }
@@ -67,6 +73,7 @@ func (m *mockHostService) FindContainer(host string, id string, labels container
} }
func (m *mockHostService) ListAllContainers(labels container.ContainerLabels) ([]container.Container, []error) { func (m *mockHostService) ListAllContainers(labels container.ContainerLabels) ([]container.Container, []error) {
m.gotLabels = labels
return m.containers, m.listErrs return m.containers, m.listErrs
} }
@@ -74,6 +81,35 @@ func (m *mockHostService) Hosts() []container.Host {
return m.hosts return m.hosts
} }
func TestReadToolsUseRequestingUsersFilter(t *testing.T) {
svc := &mockHostService{
containers: []container.Container{{ID: "abc123", Name: "web", Host: "local"}},
}
// Server-global filter. A restricted user must NOT read through this.
global := container.ContainerLabels{"com.example.scope": {"admin"}}
s := NewServer(svc, global, "test")
userLabels := container.ContainerLabels{"com.example.scope": {"tenant-a"}}
ctx := auth.WithUser(context.Background(), auth.User{ContainerLabels: userLabels})
// list_containers scopes to the user's filter.
_, _, err := s.handleListContainers(ctx, nil, &listContainersParams{})
require.NoError(t, err)
assert.Equal(t, userLabels, svc.gotLabels, "list_containers should use the requesting user's filter")
// get_container_stats (FindContainer path) scopes to the user's filter.
svc.gotLabels = nil
_, _, err = s.handleGetContainerStats(ctx, nil, &getContainerStatsParams{Host: "local", ContainerID: "abc123"})
require.NoError(t, err)
assert.Equal(t, userLabels, svc.gotLabels, "get_container_stats should use the requesting user's filter")
// With no user on the context (e.g. auth disabled) it falls back to the global filter.
svc.gotLabels = nil
_, _, err = s.handleListContainers(context.Background(), nil, &listContainersParams{})
require.NoError(t, err)
assert.Equal(t, global, svc.gotLabels, "falls back to server-global filter when unauthenticated")
}
func TestListContainers(t *testing.T) { func TestListContainers(t *testing.T) {
svc := &mockHostService{ svc := &mockHostService{
containers: []container.Container{ containers: []container.Container{