feat: add CoderVPN protocol definition & implementation (#14855)

closes #14731

Defines and implements the CoderVPN control protocol, which will be used to communicate with desktop client applications.
This commit is contained in:
Spike Curtis
2024-10-01 19:40:42 +04:00
committed by GitHub
parent 38d8e3ad6a
commit f7ddbb744f
7 changed files with 3554 additions and 0 deletions
+8
View File
@@ -488,6 +488,7 @@ gen: \
agent/proto/agent.pb.go \
provisionersdk/proto/provisioner.pb.go \
provisionerd/proto/provisionerd.pb.go \
vpn/vpn.proto \
coderd/database/dump.sql \
$(DB_GEN_FILES) \
site/src/api/typesGenerated.ts \
@@ -517,6 +518,7 @@ gen/mark-fresh:
agent/proto/agent.pb.go \
provisionersdk/proto/provisioner.pb.go \
provisionerd/proto/provisionerd.pb.go \
vpn/vpn.proto \
coderd/database/dump.sql \
$(DB_GEN_FILES) \
site/src/api/typesGenerated.ts \
@@ -600,6 +602,12 @@ provisionerd/proto/provisionerd.pb.go: provisionerd/proto/provisionerd.proto
--go-drpc_opt=paths=source_relative \
./provisionerd/proto/provisionerd.proto
vpn/vpn.pb.go: vpn/vpn.proto
protoc \
--go_out=. \
--go_opt=paths=source_relative \
./vpn/vpn.proto
site/src/api/typesGenerated.ts: $(wildcard scripts/apitypings/*) $(shell find ./codersdk $(FIND_EXCLUSIONS) -type f -name '*.go')
go run ./scripts/apitypings/ > $@
./scripts/pnpm_install.sh
+135
View File
@@ -0,0 +1,135 @@
package vpn
import (
"context"
"encoding/binary"
"io"
"sync"
"google.golang.org/protobuf/proto"
"cdr.dev/slog"
)
// MaxLength is the largest possible CoderVPN Protocol message size. This is set
// so that a misbehaving peer can't cause us to allocate a huge amount of memory.
const MaxLength = 0x1000000 // 16MiB
// serdes SERializes and DESerializes protobuf messages to and from the conn.
type serdes[S rpcMessage, R receivableRPCMessage[RR], RR any] struct {
ctx context.Context
logger slog.Logger
conn io.ReadWriteCloser
sendCh <-chan S
recvCh chan<- R
closeOnce sync.Once
wg sync.WaitGroup
}
func (s *serdes[_, R, RR]) recvLoop() {
s.logger.Debug(s.ctx, "starting recvLoop")
defer s.closeIdempotent()
defer close(s.recvCh)
for {
var length uint32
if err := binary.Read(s.conn, binary.BigEndian, &length); err != nil {
s.logger.Debug(s.ctx, "failed to read length", slog.Error(err))
return
}
if length > MaxLength {
s.logger.Critical(s.ctx, "message length exceeds max",
slog.F("length", length))
return
}
s.logger.Debug(s.ctx, "about to read message", slog.F("length", length))
mb := make([]byte, length)
if n, err := io.ReadFull(s.conn, mb); err != nil {
s.logger.Debug(s.ctx, "failed to read message",
slog.Error(err),
slog.F("num_bytes_read", n))
return
}
msg := R(new(RR))
if err := proto.Unmarshal(mb, msg); err != nil {
s.logger.Critical(s.ctx, "failed to unmarshal message", slog.Error(err))
return
}
select {
case s.recvCh <- msg:
s.logger.Debug(s.ctx, "passed received message to speaker")
case <-s.ctx.Done():
s.logger.Debug(s.ctx, "recvLoop canceled", slog.Error(s.ctx.Err()))
}
}
}
func (s *serdes[S, _, _]) sendLoop() {
s.logger.Debug(s.ctx, "starting sendLoop")
defer s.closeIdempotent()
for {
select {
case <-s.ctx.Done():
s.logger.Debug(s.ctx, "sendLoop canceled", slog.Error(s.ctx.Err()))
return
case msg, ok := <-s.sendCh:
if !ok {
s.logger.Debug(s.ctx, "sendCh closed")
return
}
mb, err := proto.Marshal(msg)
if err != nil {
s.logger.Critical(s.ctx, "failed to marshal message", slog.Error(err))
return
}
if err := binary.Write(s.conn, binary.BigEndian, uint32(len(mb))); err != nil {
s.logger.Debug(s.ctx, "failed to write length", slog.Error(err))
return
}
if _, err := s.conn.Write(mb); err != nil {
s.logger.Debug(s.ctx, "failed to write message", slog.Error(err))
return
}
}
}
}
func (s *serdes[_, _, _]) closeIdempotent() {
s.closeOnce.Do(func() {
if err := s.conn.Close(); err != nil {
s.logger.Error(s.ctx, "failed to close connection", slog.Error(err))
} else {
s.logger.Info(s.ctx, "closed connection")
}
})
}
func (s *serdes[_, _, _]) Close() error {
s.closeIdempotent()
s.wg.Wait()
return nil
}
func (s *serdes[_, _, _]) start() {
s.wg.Add(2)
go func() {
defer s.wg.Done()
s.recvLoop()
}()
go func() {
defer s.wg.Done()
s.sendLoop()
}()
}
func newSerdes[S rpcMessage, R receivableRPCMessage[RR], RR any](
ctx context.Context, logger slog.Logger, conn io.ReadWriteCloser,
sendCh <-chan S, recvCh chan<- R,
) *serdes[S, R, RR] {
return &serdes[S, R, RR]{
ctx: ctx,
logger: logger.Named("serdes"),
conn: conn,
sendCh: sendCh,
recvCh: recvCh,
}
}
+374
View File
@@ -0,0 +1,374 @@
package vpn
import (
"context"
"fmt"
"io"
"strings"
"sync"
"golang.org/x/xerrors"
"google.golang.org/protobuf/proto"
"cdr.dev/slog"
"github.com/coder/coder/v2/apiversion"
)
type SpeakerRole string
type rpcMessage interface {
proto.Message
GetRpc() *RPC
// EnsureRPC isn't autogenerated, but we'll manually add it for RPC types so that the speaker
// can allocate the RPC.
EnsureRPC() *RPC
}
func (t *TunnelMessage) EnsureRPC() *RPC {
if t.Rpc == nil {
t.Rpc = &RPC{}
}
return t.Rpc
}
func (m *ManagerMessage) EnsureRPC() *RPC {
if m.Rpc == nil {
m.Rpc = &RPC{}
}
return m.Rpc
}
// receivableRPCMessage is an rpcMessage that we can receive, and unmarshal, using generics, from a
// byte stream. proto.Unmarshal requires us to have already allocated the memory for the message
// type we are unmarshalling. All our message types are pointers like *TunnelMessage, so to
// allocate, the compiler needs to know:
//
// a) that the type is a pointer type
// b) what type it is pointing to
//
// So, this generic interface requires that the message is a pointer to the type RR. Then, we pass
// both the receivableRPCMessage and RR as type constraints, so that we'll have access to the
// underlying type when it comes time to allocate it. It's a bit messy, but the alternative is
// reflection, which has its own challenges in understandability.
type receivableRPCMessage[RR any] interface {
rpcMessage
*RR
}
const (
SpeakerRoleManager SpeakerRole = "manager"
SpeakerRoleTunnel SpeakerRole = "tunnel"
)
const (
CurrentMajor = 1
CurrentMinor = 0
)
var CurrentVersion = apiversion.New(CurrentMajor, CurrentMinor)
// speaker is an implementation of the CoderVPN protocol. It handles unary RPCs and their responses,
// as well as the low-level serialization & deserialization to the ReadWriteCloser (rwc).
//
// ┌────────┐ sendCh
// ◄─────│ ◄────────────────────────────────────────────────────────────────── ◄┐
// │ │ ▲ rpc requests
// rwc │ serdes │ │ │ sendReply()
// │ │ ┌───────────────────┐ ┌──────┼──────┐
// ──────► ┼────────► recvFromSerdes() │ rpc │rpc handling │ │
// └────────┘ recvCh │ ┼────────────► ◄──── unaryRPC()
// │ │ responses │ │ │
// │ │ │ │
// │ │ └─────────────┘ ┌ ─ ─│─ ─ ─ ─ ─ ─ ─ ┐
// │ ┼──────────────────────────────────────────► request handling
// └───────────────────┘ requests (outside speaker)
// └ ─ ─ ─ ─ ─ ─ ─ ─ ─ ┘
//
// speaker is implemented as a generic type that accepts the type of message we send (S), the type we receive (R), and
// the underlying type that R points to (RR). The speaker is intended to be wrapped by another, non-generic type for
// the role (manager or tunnel). E.g. Tunnel from this package.
//
// The serdes handles SERialiazation and DESerialization of the low level message types. The wrapping type may send
// non-RPC messages (that is messages that don't expect an explicit reply) by sending on the sendCh.
//
// Unary RPCs are handled by the unaryRPC() function, which handles sending the message and waiting for the response.
//
// recvFromSerdes() reads all incoming messages from the serdes. If they are RPC responses, it dispatches them to the
// waiting unaryRPC() function call, if any. If they are RPC requests or non-RPC messages, it wraps them in a request
// struct and sends them over the requests chan. The manager/tunnel role type must read from this chan and handle
// the requests. If they are RPC types, it should call sendReply() on the request with the reply message.
type speaker[S rpcMessage, R receivableRPCMessage[RR], RR any] struct {
serdes *serdes[S, R, RR]
requests chan *request[S, R]
logger slog.Logger
nextMsgID uint64
ctx context.Context
cancel context.CancelFunc
sendCh chan<- S
recvCh <-chan R
recvLoopDone chan struct{}
mu sync.Mutex
responseChans map[uint64]chan R
}
// newSpeaker creates a new protocol speaker.
func newSpeaker[S rpcMessage, R receivableRPCMessage[RR], RR any](
ctx context.Context, logger slog.Logger, conn io.ReadWriteCloser,
me, them SpeakerRole,
) (
*speaker[S, R, RR], error,
) {
ctx, cancel := context.WithCancel(ctx)
if err := handshake(ctx, conn, logger, me, them); err != nil {
cancel()
return nil, xerrors.Errorf("handshake failed: %w", err)
}
sendCh := make(chan S)
recvCh := make(chan R)
s := &speaker[S, R, RR]{
serdes: newSerdes(ctx, logger, conn, sendCh, recvCh),
logger: logger,
requests: make(chan *request[S, R]),
responseChans: make(map[uint64]chan R),
nextMsgID: 1,
ctx: ctx,
cancel: cancel,
sendCh: sendCh,
recvCh: recvCh,
recvLoopDone: make(chan struct{}),
}
return s, nil
}
// start starts the serialzation/deserialization. It's important this happens
// after any assignments of the speaker to its owning Tunnel or Manager, since
// the mutex is copied and that is not threadsafe.
// nolint: revive
func (s *speaker[_, _, _]) start() {
s.serdes.start()
go s.recvFromSerdes()
}
func (s *speaker[S, R, _]) recvFromSerdes() {
defer close(s.recvLoopDone)
defer close(s.requests)
for {
select {
case <-s.ctx.Done():
s.logger.Debug(s.ctx, "recvFromSerdes context done while waiting for proto", slog.Error(s.ctx.Err()))
return
case msg, ok := <-s.recvCh:
if !ok {
s.logger.Debug(s.ctx, "recvCh is closed")
return
}
rpc := msg.GetRpc()
if rpc != nil && rpc.ResponseTo != 0 {
// this is a unary response
s.tryToDeliverResponse(msg)
continue
}
req := &request[S, R]{
ctx: s.ctx,
msg: msg,
replyCh: s.sendCh,
}
select {
case <-s.ctx.Done():
s.logger.Debug(s.ctx, "recvFromSerdes context done while waiting for request handler", slog.Error(s.ctx.Err()))
return
case s.requests <- req:
}
}
}
}
// nolint: revive
func (s *speaker[_, _, _]) Close() error {
s.cancel()
err := s.serdes.Close()
return err
}
// unaryRPC sends a request/response style RPC over the protocol, waits for the response, then
// returns the response
func (s *speaker[S, R, _]) unaryRPC(ctx context.Context, req S) (resp R, err error) {
rpc := req.EnsureRPC()
msgID, respCh := s.newRPC()
rpc.MsgId = msgID
logger := s.logger.With(slog.F("msg_id", msgID))
select {
case <-ctx.Done():
return resp, ctx.Err()
case <-s.ctx.Done():
return resp, xerrors.Errorf("vpn protocol closed: %w", s.ctx.Err())
case <-s.recvLoopDone:
logger.Debug(s.ctx, "recvLoopDone while sending request")
return resp, io.ErrUnexpectedEOF
case s.sendCh <- req:
logger.Debug(s.ctx, "sent rpc request", slog.F("req", req))
}
select {
case <-ctx.Done():
s.rmResponseChan(msgID)
return resp, ctx.Err()
case <-s.ctx.Done():
s.rmResponseChan(msgID)
return resp, xerrors.Errorf("vpn protocol closed: %w", s.ctx.Err())
case <-s.recvLoopDone:
logger.Debug(s.ctx, "recvLoopDone while waiting for response")
return resp, io.ErrUnexpectedEOF
case resp = <-respCh:
logger.Debug(s.ctx, "got response", slog.F("resp", resp))
return resp, nil
}
}
func (s *speaker[_, R, _]) newRPC() (uint64, chan R) {
s.mu.Lock()
defer s.mu.Unlock()
msgID := s.nextMsgID
s.nextMsgID++
c := make(chan R, 1)
s.responseChans[msgID] = c
return msgID, c
}
func (s *speaker[_, _, _]) rmResponseChan(msgID uint64) {
s.mu.Lock()
defer s.mu.Unlock()
delete(s.responseChans, msgID)
}
func (s *speaker[_, R, _]) tryToDeliverResponse(resp R) {
msgID := resp.GetRpc().GetResponseTo()
s.mu.Lock()
defer s.mu.Unlock()
c, ok := s.responseChans[msgID]
if ok {
c <- resp
// Remove the channel since we delivered a response. This ensures that each response channel
// gets _at most_ one response. Since the channels are buffered with size 1, send will
// never block.
delete(s.responseChans, msgID)
}
}
// handshake performs the initial CoderVPN protocol handshake over the given conn
func handshake(
ctx context.Context, conn io.ReadWriteCloser, logger slog.Logger, me, them SpeakerRole,
) error {
// read and write simultaneously to avoid deadlocking if the conn is not buffered
errCh := make(chan error, 2)
go func() {
ours := headerString(CurrentVersion, me)
_, err := conn.Write([]byte(ours))
logger.Debug(ctx, "wrote out header")
if err != nil {
err = xerrors.Errorf("write header: %w", err)
}
errCh <- err
}()
headerCh := make(chan string, 1)
go func() {
// we can't use bufio.Scanner here because we need to ensure we don't read beyond the
// first newline. So, we'll read one byte at a time. It's inefficient, but the initial
// header is only a few characters, so we'll keep this code simple.
buf := make([]byte, 256)
have := 0
for {
_, err := conn.Read(buf[have : have+1])
if err != nil {
errCh <- xerrors.Errorf("read header: %w", err)
return
}
if buf[have] == '\n' {
logger.Debug(ctx, "got newline header delimiter")
// use have (not have+1) since we don't want the delimiter for verification.
headerCh <- string(buf[:have])
return
}
have++
if have >= len(buf) {
errCh <- xerrors.Errorf("header malformed or too large: %s", string(buf))
return
}
}
}()
writeOK := false
theirHeader := ""
readOK := false
for !(readOK && writeOK) {
select {
case <-ctx.Done():
_ = conn.Close() // ensure our read/write goroutines get a chance to clean up
return ctx.Err()
case err := <-errCh:
if err == nil {
// write goroutine sends nil when completing successfully.
logger.Debug(ctx, "write ok")
writeOK = true
continue
}
_ = conn.Close()
return err
case theirHeader = <-headerCh:
logger.Debug(ctx, "read ok")
readOK = true
}
}
logger.Debug(ctx, "handshake read/write complete", slog.F("their_header", theirHeader))
err := validateHeader(theirHeader, them)
if err != nil {
return xerrors.Errorf("validate header (%s): %w", theirHeader, err)
}
return nil
}
const headerPreamble = "codervpn"
func headerString(version *apiversion.APIVersion, role SpeakerRole) string {
return fmt.Sprintf("%s %s %s\n", headerPreamble, version.String(), role)
}
func validateHeader(header string, expectedRole SpeakerRole) error {
parts := strings.Split(header, " ")
if len(parts) != 3 {
return xerrors.New("wrong number of parts")
}
if parts[0] != headerPreamble {
return xerrors.New("invalid preamble")
}
if err := CurrentVersion.Validate(parts[1]); err != nil {
return xerrors.Errorf("version: %w", err)
}
if parts[2] != string(expectedRole) {
return xerrors.New("unexpected role")
}
return nil
}
type request[S rpcMessage, R rpcMessage] struct {
ctx context.Context
msg R
replyCh chan<- S
}
func (r *request[S, _]) sendReply(reply S) error {
rrpc := reply.EnsureRPC()
mrpc := r.msg.GetRpc()
if mrpc == nil {
return xerrors.Errorf("message didn't want a reply")
}
rrpc.ResponseTo = mrpc.MsgId
select {
case <-r.ctx.Done():
return r.ctx.Err()
case r.replyCh <- reply:
}
return nil
}
+456
View File
@@ -0,0 +1,456 @@
package vpn
import (
"context"
"encoding/binary"
"io"
"net"
"strings"
"testing"
"time"
"github.com/stretchr/testify/require"
"go.uber.org/goleak"
"google.golang.org/protobuf/proto"
"cdr.dev/slog"
"cdr.dev/slog/sloggers/slogtest"
"github.com/coder/coder/v2/testutil"
)
func TestMain(m *testing.M) {
goleak.VerifyTestMain(m)
}
// TestSpeaker_RawPeer tests the speaker with a peer that we simulate by directly making reads and
// writes to the other end of the pipe. There should be at least one test that does this, rather
// than use 2 speakers so that we don't have a bug where we don't adhere to the stated protocol, but
// both sides have the bug and can still communicate.
func TestSpeaker_RawPeer(t *testing.T) {
t.Parallel()
mp, tp := net.Pipe()
defer mp.Close()
defer tp.Close()
ctx := testutil.Context(t, testutil.WaitShort)
// We're going to use deadlines for this test so that we don't hang the main test thread if
// the speaker misbehaves.
err := mp.SetReadDeadline(time.Now().Add(testutil.WaitShort))
require.NoError(t, err)
err = mp.SetWriteDeadline(time.Now().Add(testutil.WaitShort))
require.NoError(t, err)
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
var tun *speaker[*TunnelMessage, *ManagerMessage, ManagerMessage]
errCh := make(chan error, 1)
go func() {
s, err := newSpeaker[*TunnelMessage, *ManagerMessage](ctx, logger, tp, SpeakerRoleTunnel, SpeakerRoleManager)
tun = s
errCh <- err
}()
expectedHandshake := "codervpn 1.0 tunnel\n"
b := make([]byte, 256)
n, err := mp.Read(b)
require.NoError(t, err)
require.Equal(t, expectedHandshake, string(b[:n]))
_, err = mp.Write([]byte("codervpn 1.0 manager\n"))
require.NoError(t, err)
err = testutil.RequireRecvCtx(ctx, t, errCh)
require.NoError(t, err)
tun.start()
// send a message and verify it follows protocol for encoding
testutil.RequireSendCtx(ctx, t, tun.sendCh, &TunnelMessage{
Msg: &TunnelMessage_Start{
Start: &StartResponse{},
},
})
var msgLen uint32
err = binary.Read(mp, binary.BigEndian, &msgLen)
require.NoError(t, err)
msgBuf := make([]byte, msgLen)
n, err = mp.Read(msgBuf)
require.NoError(t, err)
require.Equal(t, msgLen, uint32(n))
msg := new(TunnelMessage)
err = proto.Unmarshal(msgBuf, msg)
require.NoError(t, err)
_, ok := msg.Msg.(*TunnelMessage_Start)
require.True(t, ok)
// Should close the pipe on close of the speaker.
err = tun.Close()
require.NoError(t, err)
_, err = mp.Read(b)
require.ErrorIs(t, err, io.EOF)
}
func TestSpeaker_HandshakeRWFailure(t *testing.T) {
t.Parallel()
mp, tp := net.Pipe()
// immediately close the pipe, so we'll get read & write failures on handshake
_ = mp.Close()
ctx := testutil.Context(t, testutil.WaitShort)
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}).Leveled(slog.LevelDebug)
var tun *speaker[*TunnelMessage, *ManagerMessage, ManagerMessage]
errCh := make(chan error, 1)
go func() {
s, err := newSpeaker[*TunnelMessage, *ManagerMessage](
ctx, logger.Named("tun"), tp, SpeakerRoleTunnel, SpeakerRoleManager,
)
tun = s
errCh <- err
}()
err := testutil.RequireRecvCtx(ctx, t, errCh)
require.ErrorContains(t, err, "handshake failed")
require.Nil(t, tun)
}
func TestSpeaker_HandshakeCtxDone(t *testing.T) {
t.Parallel()
mp, tp := net.Pipe()
defer mp.Close()
defer tp.Close()
testCtx := testutil.Context(t, testutil.WaitShort)
ctx, cancel := context.WithCancel(testCtx)
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}).Leveled(slog.LevelDebug)
var tun *speaker[*TunnelMessage, *ManagerMessage, ManagerMessage]
errCh := make(chan error, 1)
go func() {
s, err := newSpeaker[*TunnelMessage, *ManagerMessage](
ctx, logger.Named("tun"), tp, SpeakerRoleTunnel, SpeakerRoleManager,
)
tun = s
errCh <- err
}()
cancel()
err := testutil.RequireRecvCtx(testCtx, t, errCh)
require.ErrorContains(t, err, "handshake failed")
require.Nil(t, tun)
}
func TestSpeaker_OversizeHandshake(t *testing.T) {
t.Parallel()
mp, tp := net.Pipe()
defer mp.Close()
defer tp.Close()
ctx := testutil.Context(t, testutil.WaitShort)
// We're going to use deadlines for this test so that we don't hang the main test thread if
// the speaker misbehaves.
err := mp.SetReadDeadline(time.Now().Add(testutil.WaitShort))
require.NoError(t, err)
err = mp.SetWriteDeadline(time.Now().Add(testutil.WaitShort))
require.NoError(t, err)
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}).Leveled(slog.LevelDebug)
var tun *speaker[*TunnelMessage, *ManagerMessage, ManagerMessage]
errCh := make(chan error, 1)
go func() {
s, err := newSpeaker[*TunnelMessage, *ManagerMessage](ctx, logger, tp, SpeakerRoleTunnel, SpeakerRoleManager)
tun = s
errCh <- err
}()
expectedHandshake := "codervpn 1.0 tunnel\n"
b := make([]byte, 256)
n, err := mp.Read(b)
require.NoError(t, err)
require.Equal(t, expectedHandshake, string(b[:n]))
badHandshake := strings.Repeat("bad", 256)
_, err = mp.Write([]byte(badHandshake))
require.Error(t, err) // other side closes when we write too much
err = testutil.RequireRecvCtx(ctx, t, errCh)
require.ErrorContains(t, err, "handshake failed")
require.Nil(t, tun)
}
func TestSpeaker_HandshakeInvalid(t *testing.T) {
t.Parallel()
// nolint: paralleltest // no longer need to reinitialize loop vars in go 1.22
for _, tc := range []struct {
name, handshake string
}{
{name: "preamble", handshake: "ssh 1.0 manager\n"},
{name: "2components", handshake: "ssh manager\n"},
{name: "newversion", handshake: "codervpn 1.1 manager\n"},
{name: "oldversion", handshake: "codervpn 0.1 manager\n"},
{name: "unknown_role", handshake: "codervpn 1.0 supervisor\n"},
{name: "unexpected_role", handshake: "codervpn 1.0 tunnel\n"},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
mp, tp := net.Pipe()
defer mp.Close()
defer tp.Close()
ctx := testutil.Context(t, testutil.WaitShort)
// We're going to use deadlines for this test so that we don't hang the main test thread if
// the speaker misbehaves.
err := mp.SetReadDeadline(time.Now().Add(testutil.WaitShort))
require.NoError(t, err)
err = mp.SetWriteDeadline(time.Now().Add(testutil.WaitShort))
require.NoError(t, err)
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}).Leveled(slog.LevelDebug)
var tun *speaker[*TunnelMessage, *ManagerMessage, ManagerMessage]
errCh := make(chan error, 1)
go func() {
s, err := newSpeaker[*TunnelMessage, *ManagerMessage](ctx, logger, tp, SpeakerRoleTunnel, SpeakerRoleManager)
tun = s
errCh <- err
}()
_, err = mp.Write([]byte(tc.handshake))
require.NoError(t, err)
expectedHandshake := "codervpn 1.0 tunnel\n"
b := make([]byte, 256)
n, err := mp.Read(b)
require.NoError(t, err)
require.Equal(t, expectedHandshake, string(b[:n]))
err = testutil.RequireRecvCtx(ctx, t, errCh)
require.ErrorContains(t, err, "validate header")
require.Nil(t, tun)
})
}
}
// TestSpeaker_RawPeer tests the speaker with a peer that we simulate by directly making reads and
// writes to the other end of the pipe. There should be at least one test that does this, rather
// than use 2 speakers so that we don't have a bug where we don't adhere to the stated protocol, but
// both sides have the bug and can still communicate.
func TestSpeaker_CorruptMessage(t *testing.T) {
t.Parallel()
mp, tp := net.Pipe()
defer mp.Close()
defer tp.Close()
ctx := testutil.Context(t, testutil.WaitShort)
// We're going to use deadlines for this test so that we don't hang the main test thread if
// the speaker misbehaves.
err := mp.SetReadDeadline(time.Now().Add(testutil.WaitShort))
require.NoError(t, err)
err = mp.SetWriteDeadline(time.Now().Add(testutil.WaitShort))
require.NoError(t, err)
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}).Leveled(slog.LevelDebug)
var tun *speaker[*TunnelMessage, *ManagerMessage, ManagerMessage]
errCh := make(chan error, 1)
go func() {
s, err := newSpeaker[*TunnelMessage, *ManagerMessage](ctx, logger, tp, SpeakerRoleTunnel, SpeakerRoleManager)
tun = s
errCh <- err
}()
expectedHandshake := "codervpn 1.0 tunnel\n"
b := make([]byte, 256)
n, err := mp.Read(b)
require.NoError(t, err)
require.Equal(t, expectedHandshake, string(b[:n]))
_, err = mp.Write([]byte("codervpn 1.0 manager\n"))
require.NoError(t, err)
err = testutil.RequireRecvCtx(ctx, t, errCh)
require.NoError(t, err)
tun.start()
err = binary.Write(mp, binary.BigEndian, uint32(10))
require.NoError(t, err)
n, err = mp.Write([]byte{0, 0, 0, 0, 0, 0, 0, 0, 0, 0})
require.NoError(t, err)
require.EqualValues(t, 10, n)
// it should hang up on us if we write nonsense
_, err = mp.Read(b)
require.ErrorIs(t, err, io.EOF)
}
func TestSpeaker_unaryRPC_mainline(t *testing.T) {
t.Parallel()
ctx, tun, mgr := setupSpeakers(t)
errCh := make(chan error, 1)
var resp *TunnelMessage
go func() {
r, err := mgr.unaryRPC(ctx, &ManagerMessage{
Msg: &ManagerMessage_Start{
Start: &StartRequest{
CoderUrl: "https://coder.example.com",
},
},
})
resp = r
errCh <- err
}()
req := testutil.RequireRecvCtx(ctx, t, tun.requests)
require.NotEqualValues(t, 0, req.msg.GetRpc().GetMsgId())
require.Equal(t, "https://coder.example.com", req.msg.GetStart().GetCoderUrl())
err := req.sendReply(&TunnelMessage{
Msg: &TunnelMessage_Start{
Start: &StartResponse{},
},
})
require.NoError(t, err)
err = testutil.RequireRecvCtx(ctx, t, errCh)
require.NoError(t, err)
_, ok := resp.Msg.(*TunnelMessage_Start)
require.True(t, ok)
// closing the manager should close the tun.requests channel
err = mgr.Close()
require.NoError(t, err)
select {
case _, ok := <-tun.requests:
require.False(t, ok)
case <-ctx.Done():
t.Fatal("timed out waiting for requests to close")
}
}
func TestSpeaker_unaryRPC_canceled(t *testing.T) {
t.Parallel()
testCtx, tun, mgr := setupSpeakers(t)
ctx, cancel := context.WithCancel(testCtx)
defer cancel()
errCh := make(chan error, 1)
var resp *TunnelMessage
go func() {
r, err := mgr.unaryRPC(ctx, &ManagerMessage{
Msg: &ManagerMessage_Start{
Start: &StartRequest{
CoderUrl: "https://coder.example.com",
},
},
})
resp = r
errCh <- err
}()
req := testutil.RequireRecvCtx(testCtx, t, tun.requests)
require.NotEqualValues(t, 0, req.msg.GetRpc().GetMsgId())
require.Equal(t, "https://coder.example.com", req.msg.GetStart().GetCoderUrl())
cancel()
err := testutil.RequireRecvCtx(testCtx, t, errCh)
require.ErrorIs(t, err, context.Canceled)
require.Nil(t, resp)
err = req.sendReply(&TunnelMessage{
Msg: &TunnelMessage_Start{
Start: &StartResponse{},
},
})
require.NoError(t, err)
}
func TestSpeaker_unaryRPC_hung_up(t *testing.T) {
t.Parallel()
testCtx, tun, mgr := setupSpeakers(t)
ctx, cancel := context.WithCancel(testCtx)
defer cancel()
errCh := make(chan error, 1)
var resp *TunnelMessage
go func() {
r, err := mgr.unaryRPC(ctx, &ManagerMessage{
Msg: &ManagerMessage_Start{
Start: &StartRequest{
CoderUrl: "https://coder.example.com",
},
},
})
resp = r
errCh <- err
}()
req := testutil.RequireRecvCtx(testCtx, t, tun.requests)
require.NotEqualValues(t, 0, req.msg.GetRpc().GetMsgId())
require.Equal(t, "https://coder.example.com", req.msg.GetStart().GetCoderUrl())
// When: Tunnel closes instead of replying.
err := tun.Close()
require.NoError(t, err)
// Then: we should get an error on the RPC.
err = testutil.RequireRecvCtx(testCtx, t, errCh)
require.ErrorIs(t, err, io.ErrUnexpectedEOF)
require.Nil(t, resp)
}
func TestSpeaker_unaryRPC_sendLoop(t *testing.T) {
t.Parallel()
testCtx, tun, mgr := setupSpeakers(t)
ctx, cancel := context.WithCancel(testCtx)
defer cancel()
// When: Tunnel closes before we send the RPC
err := tun.Close()
require.NoError(t, err)
// When: serdes sendloop is closed
// Send a message from the manager. This closes the manager serdes sendloop, since it will error
// when writing the message to the (closed) pipe.
testutil.RequireSendCtx(ctx, t, mgr.sendCh, &ManagerMessage{
Msg: &ManagerMessage_GetPeerUpdate{},
})
// When: we send an RPC
errCh := make(chan error, 1)
var resp *TunnelMessage
go func() {
r, err := mgr.unaryRPC(ctx, &ManagerMessage{
Msg: &ManagerMessage_Start{
Start: &StartRequest{
CoderUrl: "https://coder.example.com",
},
},
})
resp = r
errCh <- err
}()
// Then: we should get an error on the RPC.
err = testutil.RequireRecvCtx(testCtx, t, errCh)
require.ErrorIs(t, err, io.ErrUnexpectedEOF)
require.Nil(t, resp)
}
func setupSpeakers(t *testing.T) (
context.Context, *speaker[*TunnelMessage, *ManagerMessage, ManagerMessage], *speaker[*ManagerMessage, *TunnelMessage, TunnelMessage],
) {
mp, tp := net.Pipe()
t.Cleanup(func() { _ = mp.Close() })
t.Cleanup(func() { _ = tp.Close() })
ctx := testutil.Context(t, testutil.WaitShort)
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
var tun *speaker[*TunnelMessage, *ManagerMessage, ManagerMessage]
var mgr *speaker[*ManagerMessage, *TunnelMessage, TunnelMessage]
errCh := make(chan error, 2)
go func() {
s, err := newSpeaker[*TunnelMessage, *ManagerMessage](
ctx, logger.Named("tun"), tp, SpeakerRoleTunnel, SpeakerRoleManager,
)
tun = s
errCh <- err
}()
go func() {
s, err := newSpeaker[*ManagerMessage, *TunnelMessage](
ctx, logger.Named("mgr"), mp, SpeakerRoleManager, SpeakerRoleTunnel,
)
mgr = s
errCh <- err
}()
err := testutil.RequireRecvCtx(ctx, t, errCh)
require.NoError(t, err)
err = testutil.RequireRecvCtx(ctx, t, errCh)
require.NoError(t, err)
tun.start()
mgr.start()
return ctx, tun, mgr
}
+229
View File
@@ -0,0 +1,229 @@
package vpn
import (
"context"
"database/sql/driver"
"encoding/json"
"fmt"
"io"
"reflect"
"strconv"
"sync"
"unicode"
"golang.org/x/xerrors"
"cdr.dev/slog"
)
type Tunnel struct {
speaker[*TunnelMessage, *ManagerMessage, ManagerMessage]
ctx context.Context
logger slog.Logger
requestLoopDone chan struct{}
logMu sync.Mutex
logs []*TunnelMessage
}
func NewTunnel(
ctx context.Context, logger slog.Logger, conn io.ReadWriteCloser,
) (*Tunnel, error) {
logger = logger.Named("vpn")
s, err := newSpeaker[*TunnelMessage, *ManagerMessage](
ctx, logger, conn, SpeakerRoleTunnel, SpeakerRoleManager)
if err != nil {
return nil, err
}
t := &Tunnel{
// nolint: govet // safe to copy the locks here because we haven't started the speaker
speaker: *(s),
ctx: ctx,
logger: logger,
requestLoopDone: make(chan struct{}),
}
t.speaker.start()
go t.requestLoop()
return t, nil
}
func (t *Tunnel) requestLoop() {
defer close(t.requestLoopDone)
for req := range t.speaker.requests {
if req.msg.Rpc != nil && req.msg.Rpc.MsgId != 0 {
resp := t.handleRPC(req.msg, req.msg.Rpc.MsgId)
if err := req.sendReply(resp); err != nil {
t.logger.Debug(t.ctx, "failed to send RPC reply", slog.Error(err))
}
continue
}
// Not a unary RPC. We don't know of any message types that are neither a response nor a
// unary RPC from the Manager. This shouldn't ever happen because we checked the protocol
// version during the handshake.
t.logger.Critical(t.ctx, "unknown request", slog.F("msg", req.msg))
}
}
// handleRPC handles unary RPCs from the manager.
func (t *Tunnel) handleRPC(req *ManagerMessage, msgID uint64) *TunnelMessage {
resp := &TunnelMessage{}
resp.Rpc = &RPC{ResponseTo: msgID}
switch msg := req.GetMsg().(type) {
case *ManagerMessage_GetPeerUpdate:
// TODO: actually get the peer updates
resp.Msg = &TunnelMessage_PeerUpdate{
PeerUpdate: &PeerUpdate{
UpsertedWorkspaces: nil,
UpsertedAgents: nil,
},
}
return resp
case *ManagerMessage_Start:
startReq := msg.Start
t.logger.Info(t.ctx, "starting CoderVPN tunnel",
slog.F("url", startReq.CoderUrl),
slog.F("tunnel_fd", startReq.TunnelFileDescriptor),
)
// TODO: actually start the tunnel
resp.Msg = &TunnelMessage_Start{
Start: &StartResponse{
Success: true,
},
}
return resp
case *ManagerMessage_Stop:
t.logger.Info(t.ctx, "stopping CoderVPN tunnel")
// TODO: actually stop the tunnel
resp.Msg = &TunnelMessage_Stop{
Stop: &StopResponse{
Success: true,
},
}
err := t.speaker.Close()
if err != nil {
t.logger.Error(t.ctx, "failed to close speaker", slog.Error(err))
} else {
t.logger.Info(t.ctx, "coderVPN tunnel stopped")
}
return resp
default:
t.logger.Warn(t.ctx, "unhandled manager request", slog.F("request", msg))
return resp
}
}
// ApplyNetworkSettings sends a request to the manager to apply the given network settings
func (t *Tunnel) ApplyNetworkSettings(ctx context.Context, ns *NetworkSettingsRequest) error {
msg, err := t.speaker.unaryRPC(ctx, &TunnelMessage{
Msg: &TunnelMessage_NetworkSettings{
NetworkSettings: ns,
},
})
if err != nil {
return xerrors.Errorf("rpc failure: %w", err)
}
resp := msg.GetNetworkSettings()
if !resp.Success {
return xerrors.Errorf("network settings failed: %s", resp.ErrorMessage)
}
return nil
}
var _ slog.Sink = &Tunnel{}
func (t *Tunnel) LogEntry(_ context.Context, e slog.SinkEntry) {
t.logMu.Lock()
defer t.logMu.Unlock()
t.logs = append(t.logs, &TunnelMessage{
Msg: &TunnelMessage_Log{
Log: sinkEntryToPb(e),
},
})
}
func (t *Tunnel) Sync() {
t.logMu.Lock()
logs := t.logs
t.logs = nil
t.logMu.Unlock()
for _, msg := range logs {
select {
case <-t.ctx.Done():
return
case t.sendCh <- msg:
}
}
}
func sinkEntryToPb(e slog.SinkEntry) *Log {
l := &Log{
Level: Log_Level(e.Level),
Message: e.Message,
LoggerNames: e.LoggerNames,
}
for _, field := range e.Fields {
l.Fields = append(l.Fields, &Log_Field{
Name: field.Name,
Value: formatValue(field.Value),
})
}
return l
}
// the following are taken from sloghuman:
func formatValue(v interface{}) string {
if vr, ok := v.(driver.Valuer); ok {
var err error
v, err = vr.Value()
if err != nil {
return fmt.Sprintf("error calling Value: %v", err)
}
}
if v == nil {
return "<nil>"
}
typ := reflect.TypeOf(v)
switch typ.Kind() {
case reflect.Struct, reflect.Map:
byt, err := json.Marshal(v)
if err != nil {
panic(err)
}
return string(byt)
case reflect.Slice:
// Byte slices are optimistically readable.
if typ.Elem().Kind() == reflect.Uint8 {
return fmt.Sprintf("%q", v)
}
fallthrough
default:
return quote(fmt.Sprintf("%+v", v))
}
}
// quotes quotes a string so that it is suitable
// as a key for a map or in general some output that
// cannot span multiple lines or have weird characters.
func quote(key string) string {
// strconv.Quote does not quote an empty string so we need this.
if key == "" {
return `""`
}
var hasSpace bool
for _, r := range key {
if unicode.IsSpace(r) {
hasSpace = true
break
}
}
quoted := strconv.Quote(key)
// If the key doesn't need to be quoted, don't quote it.
// We do not use strconv.CanBackquote because it doesn't
// account tabs.
if !hasSpace && quoted[1:len(quoted)-1] == key {
return key
}
return quoted
}
+2155
View File
File diff suppressed because it is too large Load Diff
+197
View File
@@ -0,0 +1,197 @@
syntax = "proto3";
option go_package = "github.com/coder/coder/v2/vpn";
import "google/protobuf/timestamp.proto";
package vpn;
// The CoderVPN protocol operates over a bidirectional stream between a "manager" and a "tunnel."
// The manager is part of the Coder Desktop application and written in OS native code. It handles
// configuring the VPN and displaying status to the end user. The tunnel is written in Go and
// handles operating the actual tunnel, including reading and writing packets, & communicating with
// the Coder server control plane.
// RPC allows a very simple unary request/response RPC mechanism. The requester generates a unique
// msg_id which it sets on the request, the responder sets response_to that msg_id on the response
// message
message RPC {
uint64 msg_id = 1;
uint64 response_to = 2;
}
// ManagerMessage is a message from the manager (to the tunnel).
message ManagerMessage {
RPC rpc = 1;
oneof msg {
GetPeerUpdate get_peer_update = 2;
NetworkSettingsResponse network_settings = 3;
StartRequest start = 4;
StopRequest stop = 5;
}
}
// TunnelMessage is a message from the tunnel (to the manager).
message TunnelMessage {
RPC rpc = 1;
oneof msg {
Log log = 2;
PeerUpdate peer_update = 3;
NetworkSettingsRequest network_settings = 4;
StartResponse start = 5;
StopResponse stop = 6;
}
}
// Log is a log message generated by the tunnel. The manager should log it to the system log. It is
// one-way tunnel -> manager with no response.
message Log {
enum Level {
// these are designed to match slog levels
DEBUG = 0;
INFO = 1;
WARN = 2;
ERROR = 3;
CRITICAL = 4;
FATAL = 5;
}
Level level = 1;
string message = 2;
repeated string logger_names = 3;
message Field {
string name = 1;
string value = 2;
}
repeated Field fields = 4;
}
// GetPeerUpdate asks for a PeerUpdate with a full set of data.
message GetPeerUpdate {}
// PeerUpdate is an update about workspaces and agents connected via the tunnel. It is generated in
// response to GetPeerUpdate (which dumps the full set). It is also generated on any changes (not in
// response to any request).
message PeerUpdate {
repeated Workspace upserted_workspaces = 1;
repeated Agent upserted_agents = 2;
repeated Workspace deleted_workspaces = 3;
repeated Agent deleted_agents = 4;
}
message Workspace {
bytes id = 1; // UUID
string name = 2;
enum Status {
UNKNOWN = 0;
PENDING = 1;
STARTING = 2;
RUNNING = 3;
STOPPING = 4;
STOPPED = 5;
FAILED = 6;
CANCELING = 7;
CANCELED = 8;
DELETING = 9;
DELETED = 10;
}
Status status = 3;
}
message Agent {
bytes id = 1; // UUID
string name = 2;
bytes workspace_id = 3; // UUID
string fqdn = 4;
repeated string ip_addrs = 5;
// last_handshake is the primary indicator of whether we are connected to a peer. Zero value or
// anything longer than 5 minutes ago means there is a problem.
google.protobuf.Timestamp last_handshake = 6;
}
// NetworkSettingsRequest is based on
// https://developer.apple.com/documentation/networkextension/nepackettunnelnetworksettings for
// macOS. It is a request/response message with response NetworkSettingsResponse
message NetworkSettingsRequest {
uint32 tunnel_overhead_bytes = 1;
uint32 mtu = 2;
message DNSSettings {
repeated string servers = 1;
repeated string search_domains = 2;
// domain_name is the primary domain name of the tunnel
string domain_name = 3;
repeated string match_domains = 4;
// match_domains_no_search specifies if the domains in the matchDomains list should not be
// appended to the resolver’s list of search domains.
bool match_domains_no_search = 5;
}
DNSSettings dns_settings = 3;
string tunnel_remote_address = 4;
message IPv4Settings {
repeated string addrs = 1;
repeated string subnet_masks = 2;
// router is the next-hop router in dotted-decimal format
string router = 3;
message IPv4Route {
string destination = 1;
string mask = 2;
// router is the next-hop router in dotted-decimal format
string router = 3;
}
repeated IPv4Route included_routes = 4;
repeated IPv4Route excluded_routes = 5;
}
IPv4Settings ipv4_settings = 5;
message IPv6Settings {
repeated string addrs = 1;
repeated uint32 prefix_lengths = 2;
message IPv6Route {
string destination = 1;
uint32 prefix_length = 2;
// router is the address of the next-hop
string router = 3;
}
repeated IPv6Route included_routes = 3;
repeated IPv6Route excluded_routes = 4;
}
IPv6Settings ipv6_settings = 6;
}
// NetworkSettingsResponse is the response from the manager to the tunnel for a
// NetworkSettingsRequest
message NetworkSettingsResponse {
bool success = 1;
string error_message = 2;
}
// StartRequest is a request from the manager to start the tunnel. The tunnel replies with a
// StartResponse.
message StartRequest {
int32 tunnel_file_descriptor = 1;
string coder_url = 2;
string api_token = 3;
}
message StartResponse {
bool success = 1;
string error_message = 2;
}
// StopRequest is a request from the manager to stop the tunnel. The tunnel replies with a
// StopResponse.
message StopRequest {}
// StopResponse is a response to stopping the tunnel. After sending this response, the tunnel closes
// its side of the bidirectional stream for writing.
message StopResponse {
bool success = 1;
string error_message = 2;
}