Thanks to visit codestin.com
Credit goes to github.com

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
9 changes: 1 addition & 8 deletions coderd/aitasks.go
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,6 @@ import (
"encoding/json"
"errors"
"fmt"
"net"
"net/http"
"net/url"
"slices"
Expand Down Expand Up @@ -1086,13 +1085,7 @@ func (api *API) authAndDoWithTaskAppClient(
}
defer release()

client := &http.Client{
Transport: &http.Transport{
DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
return agentConn.DialContext(ctx, network, addr)
},
},
}
client := agentConn.AppHTTPClient()
return do(ctx, client, parsedURL)
}

Expand Down
1 change: 1 addition & 0 deletions coderd/tailnet.go
Original file line number Diff line number Diff line change
Expand Up @@ -298,6 +298,7 @@ func (s *ServerTailnet) AgentConn(ctx context.Context, agentID uuid.UUID) (works
conn = workspacesdk.NewAgentConn(s.conn, workspacesdk.AgentConnOptions{
AgentID: agentID,
CloseFunc: func() error { return workspacesdk.ErrSkipClose },
Logger: s.logger,
})

// Since we now have an open conn, be careful to close it if we error
Expand Down
43 changes: 35 additions & 8 deletions codersdk/workspacesdk/agentconn.go
Original file line number Diff line number Diff line change
Expand Up @@ -102,6 +102,7 @@ type AgentConn interface {
DebugMagicsock(ctx context.Context) ([]byte, error)
DebugManifest(ctx context.Context) ([]byte, error)
DialContext(ctx context.Context, network string, addr string) (net.Conn, error)
AppHTTPClient() *http.Client
GetPeerDiagnostics() tailnet.PeerDiagnostics
ListContainers(ctx context.Context) (codersdk.WorkspaceAgentListContainersResponse, error)
ListMCPTools(ctx context.Context) (ListMCPToolsResponse, error)
Expand Down Expand Up @@ -158,6 +159,7 @@ func (c *agentConn) SetExtraHeaders(h http.Header) {
type AgentConnOptions struct {
AgentID uuid.UUID
CloseFunc func() error
Logger slog.Logger
}

func (c *agentConn) agentAddress() netip.Addr {
Expand Down Expand Up @@ -370,6 +372,24 @@ func (c *agentConn) DialContext(ctx context.Context, network string, addr string
}
}

// AppHTTPClient returns an HTTP client for reaching HTTP apps served by this
// workspace agent. Redirects are blocked to prevent misuse.
func (c *agentConn) AppHTTPClient() *http.Client {
return &http.Client{
CheckRedirect: func(*http.Request, []*http.Request) error {
return http.ErrUseLastResponse
},
Transport: &http.Transport{
// Disable keep-alives so these short-lived clients don't leave
// idle connections (and their goroutines) lingering after they're
// discarded.
DisableKeepAlives: true,
// Host locked to agent, port from URL.
DialContext: c.DialContext,
},
}
}

// ListeningPorts lists the ports that are currently in use by the workspace.
func (c *agentConn) ListeningPorts(ctx context.Context) (codersdk.WorkspaceAgentListeningPortsResponse, error) {
ctx, span := tracing.StartSpan(ctx)
Expand Down Expand Up @@ -1386,7 +1406,12 @@ func (c *agentConn) apiRequest(ctx context.Context, method, path string, body in
// scoped to a single request: its transport cancels in-flight dials
// once reqCtx ends.
func (c *agentConn) apiClient(reqCtx context.Context) *http.Client {
agentAddr := netip.AddrPortFrom(c.agentAddress(), AgentHTTPAPIServerPort)
return &http.Client{
// Redirects are blocked to prevent misuse.
CheckRedirect: func(*http.Request, []*http.Request) error {
return http.ErrUseLastResponse
},
Transport: &http.Transport{
// Disable keep alives as we're usually only making a single
// request, and this triggers goleak in tests
Expand All @@ -1400,11 +1425,17 @@ func (c *agentConn) apiClient(reqCtx context.Context) *http.Client {
if err != nil {
return nil, xerrors.Errorf("split host port %q: %w", addr, err)
}

// Verify that the port is TailnetStatisticsPort.
if port != strconv.Itoa(AgentHTTPAPIServerPort) {
return nil, xerrors.Errorf("request %q does not appear to be for http api", addr)
}
if reqAddr, err := netip.ParseAddr(host); err != nil || reqAddr != agentAddr.Addr() {
c.opts.Logger.Warn(ctx, "blocked workspace agent API request to unintended host",
slog.F("agent_id", c.opts.AgentID),
slog.F("request_host", host),
slog.F("intended_agent_addr", agentAddr.Addr()),
)
return nil, xerrors.Errorf("request host %q does not match intended agent %q", host, agentAddr.Addr())
}
Comment thread
ethanndickson marked this conversation as resolved.

// http.Transport detaches ctx from the request context so
// a pending dial can outlive its request and serve future
Expand All @@ -1422,12 +1453,8 @@ func (c *agentConn) apiClient(reqCtx context.Context) *http.Client {
return nil, xerrors.Errorf("workspace agent not reachable in time: %v", ctx.Err())
}

ipAddr, err := netip.ParseAddr(host)
if err != nil {
return nil, xerrors.Errorf("parse host addr: %w", err)
}

conn, err := c.Conn.DialContextTCP(ctx, netip.AddrPortFrom(ipAddr, AgentHTTPAPIServerPort))
// Always dial the pinned agent address, never the request host.
conn, err := c.Conn.DialContextTCP(ctx, agentAddr)
if err != nil {
return nil, xerrors.Errorf("dial http api: %w", err)
}
Expand Down
198 changes: 198 additions & 0 deletions codersdk/workspacesdk/agentconn_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,15 +2,25 @@ package workspacesdk_test

import (
"context"
"errors"
"fmt"
"net"
"net/http"
"net/netip"
"strings"
"sync/atomic"
"testing"

"github.com/google/uuid"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.uber.org/goleak"
"tailscale.com/tailcfg"

"github.com/coder/coder/v2/codersdk/workspacesdk"
"github.com/coder/coder/v2/tailnet"
"github.com/coder/coder/v2/tailnet/proto"
"github.com/coder/coder/v2/tailnet/tailnettest"
"github.com/coder/coder/v2/testutil"
)

Expand Down Expand Up @@ -64,3 +74,191 @@ func TestAgentConn_DialBoundedByRequestContext(t *testing.T) {

goleak.VerifyNone(t, ignoreCurrent)
}

func TestAgentConnRejectsCrossAgentRedirects(t *testing.T) {
t.Parallel()

derpMap, _ := tailnettest.RunDERPAndSTUN(t)
cases := []struct {
name string
status int
invoke func(context.Context, workspacesdk.AgentConn) error
}{
{
name: "get 302",
status: http.StatusFound,
invoke: func(ctx context.Context, conn workspacesdk.AgentConn) error {
_, err := conn.ListeningPorts(ctx)
return err
},
},
{
name: "post 307",
status: http.StatusTemporaryRedirect,
invoke: func(ctx context.Context, conn workspacesdk.AgentConn) error {
return conn.WriteFile(ctx, "/tmp/attacker", strings.NewReader("redirect-body"))
},
},
{
name: "post 308",
status: http.StatusPermanentRedirect,
invoke: func(ctx context.Context, conn workspacesdk.AgentConn) error {
return conn.WriteFile(ctx, "/tmp/attacker", strings.NewReader("redirect-body"))
},
},
}

for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()

ctx := testutil.Context(t, testutil.WaitMedium)

clientID := uuid.New()
attackerID := uuid.New()
victimID := uuid.New()
clientConn, _ := newTailnetConn(t, derpMap, clientID, "client")
attackerConn, attackerIP := newTailnetConn(t, derpMap, attackerID, "attacker")
victimConn, victimIP := newTailnetConn(t, derpMap, victimID, "victim")
stitchTailnet(t, map[uuid.UUID]*tailnet.Conn{
clientID: clientConn,
attackerID: attackerConn,
victimID: victimConn,
})

var victimHit atomic.Bool
victimRouter := http.NewServeMux()
victimRouter.HandleFunc("/api/v0/listening-ports", func(rw http.ResponseWriter, _ *http.Request) {
victimHit.Store(true)
rw.Header().Set("Content-Type", "application/json")
_, _ = rw.Write([]byte(`{"ports":[]}`))
})
victimRouter.HandleFunc("/api/v0/write-file", func(rw http.ResponseWriter, _ *http.Request) {
victimHit.Store(true)
rw.WriteHeader(http.StatusOK)
})
serveTailnetHTTP(t, victimConn, victimRouter)

victimBaseURL := fmt.Sprintf("http://[%s]:%d", victimIP, workspacesdk.AgentHTTPAPIServerPort)
attackerRouter := http.NewServeMux()
attackerRouter.HandleFunc("/", func(rw http.ResponseWriter, r *http.Request) {
http.Redirect(rw, r, victimBaseURL+r.URL.RequestURI(), tc.status)
})
serveTailnetHTTP(t, attackerConn, attackerRouter)

require.True(t, clientConn.AwaitReachable(ctx, attackerIP))
require.True(t, clientConn.AwaitReachable(ctx, victimIP))

conn := workspacesdk.NewAgentConn(clientConn, workspacesdk.AgentConnOptions{
AgentID: attackerID,
})

err := tc.invoke(ctx, conn)
require.Error(t, err)
require.False(t, victimHit.Load())
})
}
}

// TestAgentConnAppHTTPClientRefusesRedirects verifies the app HTTP client does
// not follow redirects.
func TestAgentConnAppHTTPClientRefusesRedirects(t *testing.T) {
t.Parallel()

tailnetConn, err := tailnet.NewConn(&tailnet.Options{
Addresses: []netip.Prefix{tailnet.TailscaleServicePrefix.RandomPrefix()},
Logger: testutil.Logger(t),
})
require.NoError(t, err)
t.Cleanup(func() {
_ = tailnetConn.Close()
})

conn := workspacesdk.NewAgentConn(tailnetConn, workspacesdk.AgentConnOptions{
AgentID: uuid.New(),
})

client := conn.AppHTTPClient()
require.NotNil(t, client.CheckRedirect)
req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, "http://example.invalid/", nil)
require.NoError(t, err)
require.ErrorIs(t, client.CheckRedirect(req, nil), http.ErrUseLastResponse)
}

func newTailnetConn(t *testing.T, derpMap *tailcfg.DERPMap, id uuid.UUID, name string) (*tailnet.Conn, netip.Addr) {
t.Helper()

addr := tailnet.TailscaleServicePrefix.AddrFromUUID(id)
conn, err := tailnet.NewConn(&tailnet.Options{
ID: id,
Addresses: []netip.Prefix{netip.PrefixFrom(addr, 128)},
Logger: testutil.Logger(t).Named(name),
DERPMap: derpMap,
})
require.NoError(t, err)
t.Cleanup(func() {
assert.NoError(t, conn.Close())
})

return conn, addr
}

func serveTailnetHTTP(t *testing.T, conn *tailnet.Conn, handler http.Handler) {
t.Helper()

ln, err := conn.Listen("tcp", fmt.Sprintf(":%d", workspacesdk.AgentHTTPAPIServerPort))
require.NoError(t, err)

server := &http.Server{Handler: handler, ReadHeaderTimeout: testutil.WaitShort}
t.Cleanup(func() {
assert.NoError(t, server.Close())
assert.NoError(t, ln.Close())
})

go func() {
err := server.Serve(ln)
if err != nil && !errors.Is(err, net.ErrClosed) && !errors.Is(err, http.ErrServerClosed) {
assert.NoError(t, err)
}
}()
}

// stitchTailnet cross-programs every conn's node into every other conn, the
// N-peer analog of tailnet's stitch test helper, so the peers can reach each
// other without a coordinator.
func stitchTailnet(t *testing.T, conns map[uuid.UUID]*tailnet.Conn) {
t.Helper()

sendNode := func(srcID uuid.UUID, node *tailnet.Node) {
protoNode, err := tailnet.NodeToProto(node)
if !assert.NoError(t, err) {
return
}
for dstID, dst := range conns {
if dstID == srcID {
continue
}
err = dst.UpdatePeers([]*proto.CoordinateResponse_PeerUpdate{{
Id: srcID[:],
Node: protoNode,
Kind: proto.CoordinateResponse_PeerUpdate_NODE,
}})
assert.NoError(t, err)
}
}

for srcID, src := range conns {
src.SetNodeCallback(func(node *tailnet.Node) {
sendNode(srcID, node)
})
if node := src.Node(); node != nil {
sendNode(srcID, node)
}
}

t.Cleanup(func() {
for _, conn := range conns {
conn.SetNodeCallback(nil)
}
})
}
14 changes: 14 additions & 0 deletions codersdk/workspacesdk/agentconnmock/agentconnmock.go

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

1 change: 1 addition & 0 deletions codersdk/workspacesdk/workspacesdk.go
Original file line number Diff line number Diff line change
Expand Up @@ -300,6 +300,7 @@ func (c *Client) DialAgent(dialCtx context.Context, agentID uuid.UUID, options *
<-controller.Closed()
return conn.Close()
},
Logger: options.Logger,
})

if !agentConn.AwaitReachable(dialCtx) {
Expand Down
Loading
Loading