mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
fix(tailnet): enforce valid agent and client addresses (#12197)
This adds the ability for `TunnelAuth` to also authorize incoming wireguard node IPs, preventing agents from reporting anything other than their static IP generated from the agent ID.
This commit is contained in:
@@ -30,7 +30,7 @@ type connIO struct {
|
||||
responses chan<- *proto.CoordinateResponse
|
||||
bindings chan<- binding
|
||||
tunnels chan<- tunnel
|
||||
auth agpl.TunnelAuth
|
||||
auth agpl.CoordinateeAuth
|
||||
mu sync.Mutex
|
||||
closed bool
|
||||
disconnected bool
|
||||
@@ -50,7 +50,7 @@ func newConnIO(coordContext context.Context,
|
||||
responses chan<- *proto.CoordinateResponse,
|
||||
id uuid.UUID,
|
||||
name string,
|
||||
auth agpl.TunnelAuth,
|
||||
auth agpl.CoordinateeAuth,
|
||||
) *connIO {
|
||||
peerCtx, cancel := context.WithCancel(peerCtx)
|
||||
now := time.Now().Unix()
|
||||
@@ -126,6 +126,11 @@ var errDisconnect = xerrors.New("graceful disconnect")
|
||||
|
||||
func (c *connIO) handleRequest(req *proto.CoordinateRequest) error {
|
||||
c.logger.Debug(c.peerCtx, "got request")
|
||||
err := c.auth.Authorize(req)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("authorize request: %w", err)
|
||||
}
|
||||
|
||||
if req.UpdateSelf != nil {
|
||||
c.logger.Debug(c.peerCtx, "got node update", slog.F("node", req.UpdateSelf))
|
||||
b := binding{
|
||||
@@ -147,9 +152,6 @@ func (c *connIO) handleRequest(req *proto.CoordinateRequest) error {
|
||||
// doesn't just happily continue thinking everything is fine.
|
||||
return err
|
||||
}
|
||||
if !c.auth.Authorize(dst) {
|
||||
return xerrors.New("unauthorized tunnel")
|
||||
}
|
||||
t := tunnel{
|
||||
tKey: tKey{
|
||||
src: c.UniqueID(),
|
||||
|
||||
@@ -224,7 +224,7 @@ func (c *pgCoord) Close() error {
|
||||
}
|
||||
|
||||
func (c *pgCoord) Coordinate(
|
||||
ctx context.Context, id uuid.UUID, name string, a agpl.TunnelAuth,
|
||||
ctx context.Context, id uuid.UUID, name string, a agpl.CoordinateeAuth,
|
||||
) (
|
||||
chan<- *proto.CoordinateRequest, <-chan *proto.CoordinateResponse,
|
||||
) {
|
||||
|
||||
@@ -5,10 +5,12 @@ import (
|
||||
"database/sql"
|
||||
"io"
|
||||
"net"
|
||||
"net/netip"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
agpltest "github.com/coder/coder/v2/tailnet/test"
|
||||
|
||||
"github.com/google/uuid"
|
||||
@@ -113,6 +115,144 @@ func TestPGCoordinatorSingle_AgentWithoutClients(t *testing.T) {
|
||||
assertEventuallyLost(ctx, t, store, agent.id)
|
||||
}
|
||||
|
||||
func TestPGCoordinatorSingle_AgentInvalidIP(t *testing.T) {
|
||||
t.Parallel()
|
||||
if !dbtestutil.WillUsePostgres() {
|
||||
t.Skip("test only with postgres")
|
||||
}
|
||||
store, ps := dbtestutil.NewDB(t)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitSuperLong)
|
||||
defer cancel()
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
coordinator, err := tailnet.NewPGCoord(ctx, logger, ps, store)
|
||||
require.NoError(t, err)
|
||||
defer coordinator.Close()
|
||||
|
||||
agent := newTestAgent(t, coordinator, "agent")
|
||||
defer agent.close()
|
||||
agent.sendNode(&agpl.Node{
|
||||
Addresses: []netip.Prefix{
|
||||
netip.PrefixFrom(agpl.IP(), 128),
|
||||
},
|
||||
PreferredDERP: 10,
|
||||
})
|
||||
|
||||
// The agent connection should be closed immediately after sending an invalid addr
|
||||
testutil.RequireRecvCtx(ctx, t, agent.closeChan)
|
||||
assertEventuallyLost(ctx, t, store, agent.id)
|
||||
}
|
||||
|
||||
func TestPGCoordinatorSingle_AgentInvalidIPBits(t *testing.T) {
|
||||
t.Parallel()
|
||||
if !dbtestutil.WillUsePostgres() {
|
||||
t.Skip("test only with postgres")
|
||||
}
|
||||
store, ps := dbtestutil.NewDB(t)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitSuperLong)
|
||||
defer cancel()
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
coordinator, err := tailnet.NewPGCoord(ctx, logger, ps, store)
|
||||
require.NoError(t, err)
|
||||
defer coordinator.Close()
|
||||
|
||||
agent := newTestAgent(t, coordinator, "agent")
|
||||
defer agent.close()
|
||||
agent.sendNode(&agpl.Node{
|
||||
Addresses: []netip.Prefix{
|
||||
netip.PrefixFrom(agpl.IPFromUUID(agent.id), 64),
|
||||
},
|
||||
PreferredDERP: 10,
|
||||
})
|
||||
|
||||
// The agent connection should be closed immediately after sending an invalid addr
|
||||
testutil.RequireRecvCtx(ctx, t, agent.closeChan)
|
||||
assertEventuallyLost(ctx, t, store, agent.id)
|
||||
}
|
||||
|
||||
func TestPGCoordinatorSingle_AgentValidIP(t *testing.T) {
|
||||
t.Parallel()
|
||||
if !dbtestutil.WillUsePostgres() {
|
||||
t.Skip("test only with postgres")
|
||||
}
|
||||
store, ps := dbtestutil.NewDB(t)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitSuperLong)
|
||||
defer cancel()
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
coordinator, err := tailnet.NewPGCoord(ctx, logger, ps, store)
|
||||
require.NoError(t, err)
|
||||
defer coordinator.Close()
|
||||
|
||||
agent := newTestAgent(t, coordinator, "agent")
|
||||
defer agent.close()
|
||||
agent.sendNode(&agpl.Node{
|
||||
Addresses: []netip.Prefix{
|
||||
netip.PrefixFrom(agpl.IPFromUUID(agent.id), 128),
|
||||
},
|
||||
PreferredDERP: 10,
|
||||
})
|
||||
require.Eventually(t, func() bool {
|
||||
agents, err := store.GetTailnetPeers(ctx, agent.id)
|
||||
if err != nil && !xerrors.Is(err, sql.ErrNoRows) {
|
||||
t.Fatalf("database error: %v", err)
|
||||
}
|
||||
if len(agents) == 0 {
|
||||
return false
|
||||
}
|
||||
node := new(proto.Node)
|
||||
err = gProto.Unmarshal(agents[0].Node, node)
|
||||
assert.NoError(t, err)
|
||||
assert.EqualValues(t, 10, node.PreferredDerp)
|
||||
return true
|
||||
}, testutil.WaitShort, testutil.IntervalFast)
|
||||
err = agent.close()
|
||||
require.NoError(t, err)
|
||||
<-agent.errChan
|
||||
<-agent.closeChan
|
||||
assertEventuallyLost(ctx, t, store, agent.id)
|
||||
}
|
||||
|
||||
func TestPGCoordinatorSingle_AgentValidIPLegacy(t *testing.T) {
|
||||
t.Parallel()
|
||||
if !dbtestutil.WillUsePostgres() {
|
||||
t.Skip("test only with postgres")
|
||||
}
|
||||
store, ps := dbtestutil.NewDB(t)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitSuperLong)
|
||||
defer cancel()
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
coordinator, err := tailnet.NewPGCoord(ctx, logger, ps, store)
|
||||
require.NoError(t, err)
|
||||
defer coordinator.Close()
|
||||
|
||||
agent := newTestAgent(t, coordinator, "agent")
|
||||
defer agent.close()
|
||||
agent.sendNode(&agpl.Node{
|
||||
Addresses: []netip.Prefix{
|
||||
netip.PrefixFrom(codersdk.WorkspaceAgentIP, 128),
|
||||
},
|
||||
PreferredDERP: 10,
|
||||
})
|
||||
require.Eventually(t, func() bool {
|
||||
agents, err := store.GetTailnetPeers(ctx, agent.id)
|
||||
if err != nil && !xerrors.Is(err, sql.ErrNoRows) {
|
||||
t.Fatalf("database error: %v", err)
|
||||
}
|
||||
if len(agents) == 0 {
|
||||
return false
|
||||
}
|
||||
node := new(proto.Node)
|
||||
err = gProto.Unmarshal(agents[0].Node, node)
|
||||
assert.NoError(t, err)
|
||||
assert.EqualValues(t, 10, node.PreferredDerp)
|
||||
return true
|
||||
}, testutil.WaitShort, testutil.IntervalFast)
|
||||
err = agent.close()
|
||||
require.NoError(t, err)
|
||||
<-agent.errChan
|
||||
<-agent.closeChan
|
||||
assertEventuallyLost(ctx, t, store, agent.id)
|
||||
}
|
||||
|
||||
func TestPGCoordinatorSingle_AgentWithClient(t *testing.T) {
|
||||
t.Parallel()
|
||||
if !dbtestutil.WillUsePostgres() {
|
||||
|
||||
@@ -52,7 +52,7 @@ func (s *ClientService) ServeMultiAgentClient(ctx context.Context, version strin
|
||||
sub := coord.ServeMultiAgent(id)
|
||||
return ServeWorkspaceProxy(ctx, conn, sub)
|
||||
case 2:
|
||||
auth := agpl.SingleTailnetTunnelAuth{}
|
||||
auth := agpl.SingleTailnetCoordinateeAuth{}
|
||||
streamID := agpl.StreamID{
|
||||
Name: id.String(),
|
||||
ID: id,
|
||||
|
||||
@@ -182,7 +182,7 @@ func TestDialCoordinator(t *testing.T) {
|
||||
// avoid blocking
|
||||
reqs := make(chan *proto.CoordinateRequest, 100)
|
||||
resps := make(chan *proto.CoordinateResponse, 100)
|
||||
mCoord.EXPECT().Coordinate(gomock.Any(), proxyID, gomock.Any(), agpl.SingleTailnetTunnelAuth{}).
|
||||
mCoord.EXPECT().Coordinate(gomock.Any(), proxyID, gomock.Any(), agpl.SingleTailnetCoordinateeAuth{}).
|
||||
Times(1).
|
||||
Return(reqs, resps)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user