TDPB Desktop Agent (#63154)

* Implement support for TDPB in Desktop Agent.

* Support backwards compatibility with TDP clients. This is unfortunately needed because Teleport Connect establishes a TLS tunnel directly from the client to the agent, which prevents us from intercepting and translating messages on the proxy.

* Streamline protocol detection and decoder selection a bit.

* update no-op client and fix lint warnings

* Fix intermittent connection errors caused by improper handling of withheld messages during MFA. Withheld messages should (in some cases) be translated before forwarding.

* Improve some comments and log messages. Fix unnecessary indention.
This commit is contained in:
rhammonds-teleport
2026-02-17 18:54:39 +00:00
committed by GitHub
parent f04c97125d
commit 883c482139
19 changed files with 2212 additions and 2054 deletions
@@ -553,6 +553,13 @@ message DesktopRecording {
// DelayMilliseconds is the delay in milliseconds from the start of the session
int64 DelayMilliseconds = 3 [(gogoproto.jsontag) = "ms"]; // JSON tag intentionally matches SessionPrintEvent
// TDPBMessage is the encoded TDPB message.
// Only one of [Message, TDPBMessage] will be set.
bytes TDPBMessage = 4 [
(gogoproto.nullable) = true,
(gogoproto.jsontag) = "tdpb_message"
];
}
// DesktopClipboardReceive is emitted when Teleport receives
+1449 -1400
View File
File diff suppressed because it is too large Load Diff
+2
View File
@@ -41,6 +41,7 @@ import (
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/services"
"github.com/gravitational/teleport/lib/srv/desktop"
"github.com/gravitational/teleport/lib/srv/desktop/tdp/protocol/tdpb"
"github.com/gravitational/teleport/lib/utils"
)
@@ -177,6 +178,7 @@ func (process *TeleportProcess) initWindowsDesktopServiceRegistered(logger *slog
return trace.Wrap(err)
}
tlsConfig.ClientAuth = tls.RequireAndVerifyClientCert
tlsConfig.NextProtos = []string{tdpb.ProtocolName}
// Populate the correct CAs for the incoming client connection.
tlsConfig.GetConfigForClient = func(info *tls.ClientHelloInfo) (*tls.Config, error) {
var clusterName string
+26 -24
View File
@@ -25,11 +25,13 @@ import (
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
tdpbv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/desktop/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/types/events"
libevents "github.com/gravitational/teleport/lib/events"
"github.com/gravitational/teleport/lib/services"
"github.com/gravitational/teleport/lib/srv/desktop/tdp/protocol/legacy"
"github.com/gravitational/teleport/lib/srv/desktop/tdp/protocol/tdpb"
"github.com/gravitational/teleport/lib/tlsca"
)
@@ -201,8 +203,8 @@ func (d *desktopSessionAuditor) makeClipboardReceive(length int32) *events.Deskt
// In the happy path, no event is emitted here, but details from the announcement
// are cached for future audit events. An event is returned only if there was
// an error.
func (d *desktopSessionAuditor) onSharedDirectoryAnnounce(m legacy.SharedDirectoryAnnounce) *events.DesktopSharedDirectoryStart {
err := d.auditCache.SetName(directoryID(m.DirectoryID), directoryName(m.Name))
func (d *desktopSessionAuditor) onSharedDirectoryAnnounce(m *tdpb.SharedDirectoryAnnounce) *events.DesktopSharedDirectoryStart {
err := d.auditCache.SetName(directoryID(m.DirectoryId), directoryName(m.Name))
if err == nil {
// no work to do yet, but data is cached for future events
return nil
@@ -228,21 +230,21 @@ func (d *desktopSessionAuditor) onSharedDirectoryAnnounce(m legacy.SharedDirecto
},
DesktopAddr: d.desktop.GetAddr(),
DirectoryName: m.Name,
DirectoryID: m.DirectoryID,
DirectoryID: m.DirectoryId,
DesktopName: d.desktop.GetName(),
}
}
// makeSharedDirectoryStart creates a DesktopSharedDirectoryStart event.
func (d *desktopSessionAuditor) makeSharedDirectoryStart(m legacy.SharedDirectoryAcknowledge) *events.DesktopSharedDirectoryStart {
func (d *desktopSessionAuditor) makeSharedDirectoryStart(m *tdpb.SharedDirectoryAcknowledge) *events.DesktopSharedDirectoryStart {
code := libevents.DesktopSharedDirectoryStartCode
name, ok := d.auditCache.GetName(directoryID(m.DirectoryID))
name, ok := d.auditCache.GetName(directoryID(m.DirectoryId))
if !ok {
code = libevents.DesktopSharedDirectoryStartFailureCode
name = "unknown"
}
if m.ErrCode != legacy.ErrCodeNil {
if m.ErrorCode != legacy.ErrCodeNil {
code = libevents.DesktopSharedDirectoryStartFailureCode
}
@@ -256,10 +258,10 @@ func (d *desktopSessionAuditor) makeSharedDirectoryStart(m legacy.SharedDirector
UserMetadata: d.identity.GetUserMetadata(),
SessionMetadata: d.getSessionMetadata(),
ConnectionMetadata: d.getConnectionMetadata(),
Status: statusFromErrCode(m.ErrCode),
Status: statusFromErrCode(m.ErrorCode),
DesktopAddr: d.desktop.GetAddr(),
DirectoryName: string(name),
DirectoryID: m.DirectoryID,
DirectoryID: m.DirectoryId,
DesktopName: d.desktop.GetName(),
}
}
@@ -268,12 +270,12 @@ func (d *desktopSessionAuditor) makeSharedDirectoryStart(m legacy.SharedDirector
// In the happy path, no event is emitted here, but details from the operation
// are cached for future audit events. An event is returned only if there was
// an error.
func (d *desktopSessionAuditor) onSharedDirectoryReadRequest(m legacy.SharedDirectoryReadRequest) *events.DesktopSharedDirectoryRead {
did := directoryID(m.DirectoryID)
func (d *desktopSessionAuditor) onSharedDirectoryReadRequest(completion completionID, directory directoryID, m *tdpbv1.SharedDirectoryRequest_Read) *events.DesktopSharedDirectoryRead {
did := directory
path := m.Path
offset := m.Offset
err := d.auditCache.SetReadRequestInfo(completionID(m.CompletionID), readRequestInfo{
err := d.auditCache.SetReadRequestInfo(completion, readRequestInfo{
directoryID: did,
path: path,
offset: offset,
@@ -314,7 +316,7 @@ func (d *desktopSessionAuditor) onSharedDirectoryReadRequest(m legacy.SharedDire
}
// makeSharedDirectoryReadResponse creates a DesktopSharedDirectoryRead audit event.
func (d *desktopSessionAuditor) makeSharedDirectoryReadResponse(m legacy.SharedDirectoryReadResponse) *events.DesktopSharedDirectoryRead {
func (d *desktopSessionAuditor) makeSharedDirectoryReadResponse(completion completionID, errorCode uint32, m *tdpbv1.SharedDirectoryResponse_Read) *events.DesktopSharedDirectoryRead {
var did directoryID
var name directoryName
@@ -324,7 +326,7 @@ func (d *desktopSessionAuditor) makeSharedDirectoryReadResponse(m legacy.SharedD
code := libevents.DesktopSharedDirectoryReadCode
// Gather info from the audit cache
info, ok := d.auditCache.TakeReadRequestInfo(completionID(m.CompletionID))
info, ok := d.auditCache.TakeReadRequestInfo(completion)
if ok {
did = info.directoryID
// Only search for the directory name if we retrieved the directory ID from the audit cache.
@@ -341,7 +343,7 @@ func (d *desktopSessionAuditor) makeSharedDirectoryReadResponse(m legacy.SharedD
name = "unknown"
}
if m.ErrCode != legacy.ErrCodeNil {
if errorCode != legacy.ErrCodeNil {
code = libevents.DesktopSharedDirectoryWriteFailureCode
}
@@ -355,12 +357,12 @@ func (d *desktopSessionAuditor) makeSharedDirectoryReadResponse(m legacy.SharedD
UserMetadata: d.identity.GetUserMetadata(),
SessionMetadata: d.getSessionMetadata(),
ConnectionMetadata: d.getConnectionMetadata(),
Status: statusFromErrCode(m.ErrCode),
Status: statusFromErrCode(errorCode),
DesktopAddr: d.desktop.GetAddr(),
DirectoryName: string(name),
DirectoryID: uint32(did),
Path: path,
Length: m.ReadDataLength,
Length: uint32(len(m.Data)),
Offset: offset,
DesktopName: d.desktop.GetName(),
}
@@ -370,13 +372,13 @@ func (d *desktopSessionAuditor) makeSharedDirectoryReadResponse(m legacy.SharedD
// In the happy path, no event is emitted here, but details from the operation
// are cached for future audit events. An event is returned only if there was
// an error.
func (d *desktopSessionAuditor) onSharedDirectoryWriteRequest(m legacy.SharedDirectoryWriteRequest) *events.DesktopSharedDirectoryWrite {
did := directoryID(m.DirectoryID)
func (d *desktopSessionAuditor) onSharedDirectoryWriteRequest(completion completionID, directory directoryID, m *tdpbv1.SharedDirectoryRequest_Write) *events.DesktopSharedDirectoryWrite {
did := directory
path := m.Path
offset := m.Offset
err := d.auditCache.SetWriteRequestInfo(
completionID(m.CompletionID),
completion,
writeRequestInfo{
directoryID: did,
path: path,
@@ -412,13 +414,13 @@ func (d *desktopSessionAuditor) onSharedDirectoryWriteRequest(m legacy.SharedDir
DirectoryName: string(name),
DirectoryID: uint32(did),
Path: path,
Length: m.WriteDataLength,
Length: uint32(len(m.Data)),
Offset: offset,
}
}
// makeSharedDirectoryWriteResponse creates a DesktopSharedDirectoryWrite audit event.
func (d *desktopSessionAuditor) makeSharedDirectoryWriteResponse(m legacy.SharedDirectoryWriteResponse) *events.DesktopSharedDirectoryWrite {
func (d *desktopSessionAuditor) makeSharedDirectoryWriteResponse(completion completionID, errorCode uint32, m *tdpbv1.SharedDirectoryResponse_Write) *events.DesktopSharedDirectoryWrite {
var did directoryID
var name directoryName
@@ -427,7 +429,7 @@ func (d *desktopSessionAuditor) makeSharedDirectoryWriteResponse(m legacy.Shared
code := libevents.DesktopSharedDirectoryWriteCode
// Gather info from the audit cache
info, ok := d.auditCache.TakeWriteRequestInfo(completionID(m.CompletionID))
info, ok := d.auditCache.TakeWriteRequestInfo(completion)
if ok {
did = info.directoryID
// Only search for the directory name if we retrieved the directoryID from the audit cache.
@@ -444,7 +446,7 @@ func (d *desktopSessionAuditor) makeSharedDirectoryWriteResponse(m legacy.Shared
name = "unknown"
}
if m.ErrCode != legacy.ErrCodeNil {
if errorCode != legacy.ErrCodeNil {
code = libevents.DesktopSharedDirectoryWriteFailureCode
}
@@ -458,7 +460,7 @@ func (d *desktopSessionAuditor) makeSharedDirectoryWriteResponse(m legacy.Shared
UserMetadata: d.identity.GetUserMetadata(),
SessionMetadata: d.getSessionMetadata(),
ConnectionMetadata: d.getConnectionMetadata(),
Status: statusFromErrCode(m.ErrCode),
Status: statusFromErrCode(errorCode),
DesktopAddr: d.desktop.GetAddr(),
DirectoryName: string(name),
DirectoryID: uint32(did),
+49 -69
View File
@@ -29,11 +29,13 @@ import (
"github.com/jonboulle/clockwork"
"github.com/stretchr/testify/require"
tdpbv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/desktop/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/types/events"
libevents "github.com/gravitational/teleport/lib/events"
"github.com/gravitational/teleport/lib/reversetunnel"
"github.com/gravitational/teleport/lib/srv/desktop/tdp/protocol/legacy"
"github.com/gravitational/teleport/lib/srv/desktop/tdp/protocol/tdpb"
"github.com/gravitational/teleport/lib/tlsca"
logutils "github.com/gravitational/teleport/lib/utils/log"
)
@@ -257,16 +259,16 @@ func TestDesktopSharedDirectoryStartEvent(t *testing.T) {
if test.sendsAnnounce {
// SharedDirectoryAnnounce initializes the nameCache.
audit.onSharedDirectoryAnnounce(legacy.SharedDirectoryAnnounce{
DirectoryID: uint32(testDirectoryID),
audit.onSharedDirectoryAnnounce(&tdpb.SharedDirectoryAnnounce{
DirectoryId: uint32(testDirectoryID),
Name: testDirName,
})
}
// SharedDirectoryAcknowledge causes the event to be emitted
startEvent := audit.makeSharedDirectoryStart(legacy.SharedDirectoryAcknowledge{
DirectoryID: uint32(testDirectoryID),
ErrCode: test.errCode,
startEvent := audit.makeSharedDirectoryStart(&tdpb.SharedDirectoryAcknowledge{
DirectoryId: uint32(testDirectoryID),
ErrorCode: test.errCode,
})
baseEvent := &events.DesktopSharedDirectoryStart{
@@ -392,29 +394,24 @@ func TestDesktopSharedDirectoryReadEvent(t *testing.T) {
if test.sendsAnnounce {
// SharedDirectoryAnnounce initializes the name cache
audit.onSharedDirectoryAnnounce(legacy.SharedDirectoryAnnounce{
DirectoryID: uint32(testDirectoryID),
audit.onSharedDirectoryAnnounce(&tdpb.SharedDirectoryAnnounce{
DirectoryId: uint32(testDirectoryID),
Name: testDirName,
})
}
if test.sendsReq {
// SharedDirectoryReadRequest initializes the readRequestCache.
audit.onSharedDirectoryReadRequest(legacy.SharedDirectoryReadRequest{
CompletionID: uint32(testCompletionID),
DirectoryID: uint32(testDirectoryID),
Path: testFilePath,
Offset: testOffset,
Length: testLength,
audit.onSharedDirectoryReadRequest(testCompletionID, testDirectoryID, &tdpbv1.SharedDirectoryRequest_Read{
Path: testFilePath,
Offset: testOffset,
Length: testLength,
})
}
// SharedDirectoryReadResponse causes the event to be emitted.
readEvent := audit.makeSharedDirectoryReadResponse(legacy.SharedDirectoryReadResponse{
CompletionID: uint32(testCompletionID),
ErrCode: test.errCode,
ReadDataLength: testLength,
ReadData: []byte{}, // irrelevant in this context
readEvent := audit.makeSharedDirectoryReadResponse(testCompletionID, test.errCode, &tdpbv1.SharedDirectoryResponse_Read{
Data: make([]byte, testLength), // slice contents are irrelevant in this context
})
baseEvent := &events.DesktopSharedDirectoryRead{
@@ -542,27 +539,23 @@ func TestDesktopSharedDirectoryWriteEvent(t *testing.T) {
if test.sendsAnnounce {
// SharedDirectoryAnnounce initializes the nameCache.
audit.onSharedDirectoryAnnounce(legacy.SharedDirectoryAnnounce{
DirectoryID: uint32(testDirectoryID),
audit.onSharedDirectoryAnnounce(&tdpb.SharedDirectoryAnnounce{
DirectoryId: uint32(testDirectoryID),
Name: testDirName,
})
}
if test.sendsReq {
// SharedDirectoryWriteRequest initializes the writeRequestCache.
audit.onSharedDirectoryWriteRequest(legacy.SharedDirectoryWriteRequest{
CompletionID: uint32(testCompletionID),
DirectoryID: uint32(testDirectoryID),
Path: testFilePath,
Offset: testOffset,
WriteDataLength: testLength,
audit.onSharedDirectoryWriteRequest(testCompletionID, testDirectoryID, &tdpbv1.SharedDirectoryRequest_Write{
Path: testFilePath,
Offset: testOffset,
Data: make([]byte, testLength),
})
}
// SharedDirectoryWriteResponse causes the event to be emitted.
writeEvent := audit.makeSharedDirectoryWriteResponse(legacy.SharedDirectoryWriteResponse{
CompletionID: uint32(testCompletionID),
ErrCode: test.errCode,
writeEvent := audit.makeSharedDirectoryWriteResponse(testCompletionID, test.errCode, &tdpbv1.SharedDirectoryResponse_Write{
BytesWritten: testLength,
})
@@ -620,8 +613,8 @@ func TestDesktopSharedDirectoryStartEventAuditCacheMax(t *testing.T) {
fillReadRequestCache(&audit.auditCache, testDirectoryID)
// Send a SharedDirectoryAnnounce
startEvent := audit.onSharedDirectoryAnnounce(legacy.SharedDirectoryAnnounce{
DirectoryID: uint32(testDirectoryID),
startEvent := audit.onSharedDirectoryAnnounce(&tdpb.SharedDirectoryAnnounce{
DirectoryId: uint32(testDirectoryID),
Name: testDirName,
})
require.NotNil(t, startEvent)
@@ -666,8 +659,8 @@ func TestDesktopSharedDirectoryReadEventAuditCacheMax(t *testing.T) {
id, audit := setup(testDesktop)
// Send a SharedDirectoryAnnounce
audit.onSharedDirectoryAnnounce(legacy.SharedDirectoryAnnounce{
DirectoryID: uint32(testDirectoryID),
audit.onSharedDirectoryAnnounce(&tdpb.SharedDirectoryAnnounce{
DirectoryId: uint32(testDirectoryID),
Name: testDirName,
})
@@ -675,12 +668,10 @@ func TestDesktopSharedDirectoryReadEventAuditCacheMax(t *testing.T) {
fillReadRequestCache(&audit.auditCache, testDirectoryID)
// SharedDirectoryReadRequest should cause a failed audit event.
readEvent := audit.onSharedDirectoryReadRequest(legacy.SharedDirectoryReadRequest{
CompletionID: uint32(testCompletionID),
DirectoryID: uint32(testDirectoryID),
Path: testFilePath,
Offset: testOffset,
Length: testLength,
readEvent := audit.onSharedDirectoryReadRequest(testCompletionID, testDirectoryID, &tdpbv1.SharedDirectoryRequest_Read{
Path: testFilePath,
Offset: testOffset,
Length: testLength,
})
require.NotNil(t, readEvent)
@@ -724,19 +715,17 @@ func TestDesktopSharedDirectoryReadEventAuditCacheMax(t *testing.T) {
func TestDesktopSharedDirectoryWriteEventAuditCacheMax(t *testing.T) {
id, audit := setup(testDesktop)
audit.onSharedDirectoryAnnounce(legacy.SharedDirectoryAnnounce{
DirectoryID: uint32(testDirectoryID),
audit.onSharedDirectoryAnnounce(&tdpb.SharedDirectoryAnnounce{
DirectoryId: uint32(testDirectoryID),
Name: testDirName,
})
fillReadRequestCache(&audit.auditCache, testDirectoryID)
writeEvent := audit.onSharedDirectoryWriteRequest(legacy.SharedDirectoryWriteRequest{
CompletionID: uint32(testCompletionID),
DirectoryID: uint32(testDirectoryID),
Path: testFilePath,
Offset: testOffset,
WriteDataLength: testLength,
writeEvent := audit.onSharedDirectoryWriteRequest(testCompletionID, testDirectoryID, &tdpbv1.SharedDirectoryRequest_Write{
Path: testFilePath,
Offset: testOffset,
Data: make([]byte, testLength),
})
require.NotNil(t, writeEvent, "audit event should have been generated")
@@ -780,8 +769,8 @@ func TestAuditCacheLifecycle(t *testing.T) {
_, audit := setup(testDesktop)
// SharedDirectoryAnnounce initializes the nameCache.
audit.onSharedDirectoryAnnounce(legacy.SharedDirectoryAnnounce{
DirectoryID: uint32(testDirectoryID),
audit.onSharedDirectoryAnnounce(&tdpb.SharedDirectoryAnnounce{
DirectoryId: uint32(testDirectoryID),
Name: testDirName,
})
@@ -796,22 +785,18 @@ func TestAuditCacheLifecycle(t *testing.T) {
require.False(t, ok)
// A SharedDirectoryReadRequest should add a corresponding entry in the readRequestCache.
audit.onSharedDirectoryReadRequest(legacy.SharedDirectoryReadRequest{
CompletionID: uint32(testCompletionID),
DirectoryID: uint32(testDirectoryID),
Path: testFilePath,
Offset: testOffset,
Length: testLength,
audit.onSharedDirectoryReadRequest(testCompletionID, testDirectoryID, &tdpbv1.SharedDirectoryRequest_Read{
Path: testFilePath,
Offset: testOffset,
Length: testLength,
})
require.Equal(t, 2, audit.auditCache.totalItems())
// A SharedDirectoryWriteRequest should add a corresponding entry in the writeRequestCache.
audit.onSharedDirectoryWriteRequest(legacy.SharedDirectoryWriteRequest{
CompletionID: uint32(testCompletionID),
DirectoryID: uint32(testDirectoryID),
Path: testFilePath,
Offset: testOffset,
WriteDataLength: testLength,
audit.onSharedDirectoryWriteRequest(testCompletionID, testDirectoryID, &tdpbv1.SharedDirectoryRequest_Write{
Path: testFilePath,
Offset: testOffset,
Data: make([]byte, testLength),
})
require.Equal(t, 3, audit.auditCache.totalItems())
@@ -822,18 +807,13 @@ func TestAuditCacheLifecycle(t *testing.T) {
require.Contains(t, audit.auditCache.writeRequestCache, testCompletionID)
// SharedDirectoryReadResponse should cause the entry in the readRequestCache to be cleaned up.
audit.makeSharedDirectoryReadResponse(legacy.SharedDirectoryReadResponse{
CompletionID: uint32(testCompletionID),
ErrCode: legacy.ErrCodeNil,
ReadDataLength: testLength,
ReadData: []byte{}, // irrelevant in this context
audit.makeSharedDirectoryReadResponse(testCompletionID, legacy.ErrCodeNil, &tdpbv1.SharedDirectoryResponse_Read{
Data: make([]byte, testLength), // slice contents are irrelevant in this context
})
require.Equal(t, 2, audit.auditCache.totalItems())
// SharedDirectoryWriteResponse should cause the entry in the writeRequestCache to be cleaned up.
audit.makeSharedDirectoryWriteResponse(legacy.SharedDirectoryWriteResponse{
CompletionID: uint32(testCompletionID),
ErrCode: legacy.ErrCodeNil,
audit.makeSharedDirectoryWriteResponse(testCompletionID, legacy.ErrCodeNil, &tdpbv1.SharedDirectoryResponse_Write{
BytesWritten: testLength,
})
require.Equal(t, 1, audit.auditCache.totalItems())
File diff suppressed because it is too large Load Diff
@@ -23,11 +23,14 @@ import (
"context"
"image/png"
"log/slog"
"slices"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/srv/desktop/tdp"
"github.com/gravitational/teleport/lib/srv/desktop/tdp/protocol/legacy"
"github.com/gravitational/teleport/lib/srv/desktop/tdp/protocol/tdpb"
)
// LicenseStore implements client-side license storage for Microsoft
@@ -48,10 +51,6 @@ type Config struct {
// AuthorizeFn is called to authorize a user connecting to a Windows desktop.
AuthorizeFn func(login string) error
// Conn handles TDP messages between Windows Desktop Service
// and a Teleport Proxy.
Conn *tdp.Conn
// Encoder is an optional override for PNG encoding.
Encoder *png.Encoder
@@ -90,6 +89,9 @@ type Config struct {
// AD indicates whether the desktop is part of an Active Directory domain.
AD bool
// The desktop protocol version used by the client (TDP or TDPB).
ClientProtocol string
}
//nolint:unused // used in client.go that is behind desktop_access_rdp build flag
@@ -97,9 +99,6 @@ func (c *Config) checkAndSetDefaults() error {
if c.Addr == "" {
return trace.BadParameter("missing Addr in rdpclient.Config")
}
if c.Conn == nil {
return trace.BadParameter("missing Conn in rdpclient.Config")
}
if c.AuthorizeFn == nil {
return trace.BadParameter("missing AuthorizeFn in rdpclient.Config")
}
@@ -109,6 +108,9 @@ func (c *Config) checkAndSetDefaults() error {
if c.Encoder == nil {
c.Encoder = tdp.PNGEncoder()
}
if !slices.Contains([]string{tdpb.ProtocolName, legacy.ProtocolName}, c.ClientProtocol) {
return trace.BadParameter("missing ClientProtocol in rdpclient.Config")
}
c.Logger = c.Logger.With("rdp_addr", c.Addr)
return nil
}
+3 -1
View File
@@ -28,6 +28,8 @@ import (
"context"
"errors"
"time"
"github.com/gravitational/teleport/lib/srv/desktop/tdp"
)
// Client is the dummy RDP client.
@@ -37,7 +39,7 @@ type Client struct {
// New creates and connects a new Client based on opts.
//
//nolint:staticcheck // SA4023. False positive, depends on build tags.
func New(cfg Config) (*Client, error) {
func New(_ *tdp.Conn, cfg Config) (*Client, error) {
return nil, errors.New("the real rdpclient.Client implementation was not included in this build")
}
+19 -39
View File
@@ -21,13 +21,14 @@ package rdpclient
import (
"bytes"
"io"
"log/slog"
"testing"
"github.com/stretchr/testify/require"
"github.com/gravitational/teleport/lib/srv/desktop/tdp"
"github.com/gravitational/teleport/lib/srv/desktop/tdp/protocol/legacy"
"github.com/gravitational/teleport/lib/srv/desktop/tdp/protocol/tdpb"
)
type fakeConn struct {
@@ -54,63 +55,42 @@ func (f *fakeConn) AddMessage(message tdp.Message) error {
func TestClientNew_EOF(t *testing.T) {
f := fakeConn{}
err := f.AddMessage(legacy.ClientUsername{Username: "user"})
require.NoError(t, err)
conn := tdp.NewConn(&f, legacy.Decode)
conn := tdp.NewConn(&f, tdp.DecoderAdapter(tdpb.DecodePermissive))
_, err = New(createConfig(conn))
require.EqualError(t, err, "EOF")
_, err := New(conn, createConfig())
require.ErrorIs(t, err, io.EOF)
}
func TestClientNew_NoKeyboardLayout(t *testing.T) {
f := fakeConn{}
err := f.AddMessage(legacy.ClientUsername{Username: "user"})
require.NoError(t, err)
err = f.AddMessage(legacy.ClientScreenSpec{
Width: 100,
Height: 100,
})
require.NoError(t, err)
err = f.AddMessage(legacy.ClientScreenSpec{
Width: 100,
Height: 100,
})
err := f.AddMessage(&tdpb.ClientHello{Username: "user"})
require.NoError(t, err)
conn := tdp.NewConn(&f, legacy.Decode)
conn := tdp.NewConn(&f, tdp.DecoderAdapter(tdpb.DecodePermissive))
_, err = New(createConfig(conn))
_, err = New(conn, createConfig())
require.NoError(t, err)
}
func TestClientNew_KeyboardLayout(t *testing.T) {
f := fakeConn{}
err := f.AddMessage(legacy.ClientUsername{Username: "user"})
require.NoError(t, err)
err = f.AddMessage(legacy.ClientScreenSpec{
Width: 100,
Height: 100,
})
require.NoError(t, err)
err = f.AddMessage(legacy.ClientKeyboardLayout{})
require.NoError(t, err)
err = f.AddMessage(legacy.ClientScreenSpec{
Width: 100,
Height: 100,
})
err := f.AddMessage(&tdpb.ClientHello{Username: "user", KeyboardLayout: 1})
require.NoError(t, err)
conn := tdp.NewConn(&f, legacy.Decode)
conn := tdp.NewConn(&f, tdp.DecoderAdapter(tdpb.DecodePermissive))
_, err = New(createConfig(conn))
_, err = New(conn, createConfig())
require.NoError(t, err)
}
func createConfig(conn *tdp.Conn) Config {
func createConfig() Config {
return Config{
Addr: "example.com",
AuthorizeFn: func(login string) error { return nil },
Conn: conn,
Logger: slog.Default(),
Addr: "example.com",
AuthorizeFn: func(login string) error { return nil },
Logger: slog.Default(),
Width: 1,
Height: 1,
ClientProtocol: tdpb.ProtocolName,
}
}
@@ -42,6 +42,9 @@ import (
"github.com/gravitational/teleport/lib/web/mfajson"
)
// ProtocolName is the identifier for the TDP protocol.
const ProtocolName = "teleport-tdp"
// MessageType identifies the type of the message.
type MessageType byte
@@ -33,6 +33,9 @@ import (
"github.com/gravitational/teleport/lib/srv/desktop/tdp"
)
// ProtocolName is the identifier for the TDPB protocol.
const ProtocolName = "teleport-tdpb-1.0"
// ErrUnknownMessage is returned when an unknown message is decoded.
var ErrUnknownMessage = errors.New("decoded unknown TDPB message")
@@ -271,6 +271,13 @@ func TranslateToLegacy(msg tdp.Message) ([]tdp.Message, error) {
return nil, trace.WrapWithMessage(err, "Cannot parse uuid bytes from ping")
}
return []tdp.Message{legacy.Ping{UUID: id}}, nil
case *ServerHello:
return []tdp.Message{legacy.ConnectionActivated{
IOChannelID: uint16(m.ActivationSpec.IoChannelId),
UserChannelID: uint16(m.ActivationSpec.UserChannelId),
ScreenWidth: uint16(m.ActivationSpec.ScreenWidth),
ScreenHeight: uint16(m.ActivationSpec.ScreenHeight),
}}, nil
default:
return nil, trace.Errorf("Could not translate to TDP. Encountered unexpected message type %T", m)
}
+113 -91
View File
@@ -41,6 +41,7 @@ import (
"github.com/gravitational/teleport"
apidefaults "github.com/gravitational/teleport/api/defaults"
tdpbv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/desktop/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/types/events"
"github.com/gravitational/teleport/api/utils/clientutils"
@@ -60,6 +61,7 @@ import (
"github.com/gravitational/teleport/lib/srv/desktop/rdp/rdpclient"
"github.com/gravitational/teleport/lib/srv/desktop/tdp"
"github.com/gravitational/teleport/lib/srv/desktop/tdp/protocol/legacy"
"github.com/gravitational/teleport/lib/srv/desktop/tdp/protocol/tdpb"
"github.com/gravitational/teleport/lib/tlsca"
"github.com/gravitational/teleport/lib/utils"
"github.com/gravitational/teleport/lib/utils/dns"
@@ -624,33 +626,64 @@ func (s *WindowsService) Serve(plainLis net.Listener) error {
}
}
func newErrorSender(protocol string, conn *tdp.Conn, logger *slog.Logger) func(string) {
if protocol == tdpb.ProtocolName {
return func(message string) {
if err := conn.WriteMessage(&tdpb.Alert{Message: message, Severity: tdpbv1.AlertSeverity_ALERT_SEVERITY_ERROR}); err != nil {
logger.ErrorContext(context.Background(), "Failed to send TDPB error message", "error", err, "message", message)
}
}
}
return func(message string) {
if err := conn.WriteMessage(&legacy.Alert{Message: message, Severity: legacy.SeverityError}); err != nil {
logger.ErrorContext(context.Background(), "Failed to send TDP error message", "error", err, "message", message)
}
}
}
// handleConnection handles TLS connections from a Teleport proxy.
// It authenticates and authorizes the connection, and then begins
// translating the TDP messages from the proxy into native RDP.
func (s *WindowsService) handleConnection(proxyConn *tls.Conn) {
log := s.cfg.Logger
tdpConn := tdp.NewConn(proxyConn, legacy.Decode)
// Ensure TLS handshake is complete so that we can the read ALPN result.
if err := proxyConn.Handshake(); err != nil {
log.ErrorContext(context.Background(), "Failed to complete TLS handshake")
return
}
// Figure out which protocol the client is using
clientProtocol := proxyConn.ConnectionState().NegotiatedProtocol
var decoder tdp.Decoder
switch clientProtocol {
case tdpb.ProtocolName:
decoder = tdp.DecoderAdapter(tdpb.DecodePermissive)
case "":
clientProtocol = legacy.ProtocolName
decoder = legacy.Decode
default:
log.ErrorContext(context.Background(), "Unknown client protocol selection", "protocol", clientProtocol)
return
}
tdpConn := tdp.NewConn(proxyConn, decoder)
defer tdpConn.Close()
// Inline function to enforce that we are centralizing TDP Error sending in this function.
sendTDPError := func(message string) {
if err := tdpConn.WriteMessage(legacy.Alert{Message: message, Severity: legacy.SeverityError}); err != nil {
log.ErrorContext(context.Background(), "Failed to send TDP error message", "error", err)
}
}
// Inline function to enforce that we are centralizing TDP/TDPB Error sending in this function.
sendError := newErrorSender(clientProtocol, tdpConn, log)
// Check connection limits.
remoteAddr, _, err := net.SplitHostPort(proxyConn.RemoteAddr().String())
if err != nil {
log.ErrorContext(context.Background(), "Could not parse client IP", "addr", proxyConn.RemoteAddr().String(), "error", err)
sendTDPError("Internal error.")
sendError("Internal error.")
return
}
log = log.With("client_ip", remoteAddr)
if err := s.cfg.ConnLimiter.AcquireConnection(remoteAddr); err != nil {
log.WarnContext(context.Background(), "Connection limit exceeded, rejecting connection")
sendTDPError("Connection limit exceeded.")
sendError("Connection limit exceeded.")
return
}
defer s.cfg.ConnLimiter.ReleaseConnection(remoteAddr)
@@ -659,7 +692,7 @@ func (s *WindowsService) handleConnection(proxyConn *tls.Conn) {
ctx, err := s.middleware.WrapContextWithUser(s.closeCtx, proxyConn)
if err != nil {
log.WarnContext(ctx, "mTLS authentication failed for incoming connection", "error", err)
sendTDPError("Connection authentication failed.")
sendError("Connection authentication failed.")
return
}
log.DebugContext(ctx, "Authenticated Windows desktop connection")
@@ -667,7 +700,7 @@ func (s *WindowsService) handleConnection(proxyConn *tls.Conn) {
authContext, err := s.cfg.Authorizer.Authorize(ctx)
if err != nil {
log.WarnContext(ctx, "authorization failed for Windows desktop connection", "error", err)
sendTDPError("Connection authorization failed.")
sendError("Connection authorization failed.")
return
}
@@ -690,12 +723,12 @@ func (s *WindowsService) handleConnection(proxyConn *tls.Conn) {
}))
if err != nil {
log.WarnContext(ctx, "Failed to fetch desktop by name", "error", err)
sendTDPError("Teleport failed to find the requested desktop in its database.")
sendError("Teleport failed to find the requested desktop in its database.")
return
}
if len(desktops) == 0 {
log.ErrorContext(ctx, "desktop not found", "host_uuid", s.cfg.Heartbeat.HostUUID, "name", desktopName)
sendTDPError(fmt.Sprintf("Could not find desktop %v.", desktopName))
sendError(fmt.Sprintf("Could not find desktop %v.", desktopName))
return
}
desktop := desktops[0]
@@ -704,19 +737,19 @@ func (s *WindowsService) handleConnection(proxyConn *tls.Conn) {
log.DebugContext(ctx, "Connecting to Windows desktop")
defer log.DebugContext(ctx, "Windows desktop disconnected")
if err := s.connectRDP(ctx, log, tdpConn, desktop, authContext); err != nil {
if err := s.connectRDP(ctx, log, tdpConn, desktop, authContext, clientProtocol); err != nil {
log.ErrorContext(context.Background(), "RDP connection failed", "error", err)
msg := "RDP connection failed."
var um trace.UserMessager
if errors.As(err, &um) {
msg = um.UserMessage()
}
sendTDPError(msg)
sendError(msg)
return
}
}
func (s *WindowsService) connectRDP(ctx context.Context, log *slog.Logger, tdpConn *tdp.Conn, desktop types.WindowsDesktop, authCtx *authz.Context) error {
func (s *WindowsService) connectRDP(ctx context.Context, log *slog.Logger, tdpConn *tdp.Conn, desktop types.WindowsDesktop, authCtx *authz.Context, clientProtocol string) error {
identity := authCtx.Identity.GetIdentity()
log = log.With("teleport_user", identity.Username, "desktop_addr", desktop.GetAddr(), "ad", !desktop.NonAD())
@@ -838,17 +871,16 @@ func (s *WindowsService) connectRDP(ctx context.Context, log *slog.Logger, tdpCo
}
log = log.With("kdc_addr", kdcAddr, "nla", nla)
log.InfoContext(context.Background(), "initiating RDP client")
log.InfoContext(context.Background(), "initiating RDP client", "client_protocol", clientProtocol)
//nolint:staticcheck // SA4023. False positive, depends on build tags.
rdpc, err := rdpclient.New(rdpclient.Config{
rdpc, err := rdpclient.New(tdpConn, rdpclient.Config{
LicenseStore: s.cfg.LicenseStore,
HostID: s.cfg.Heartbeat.HostUUID,
Logger: log,
Addr: addr.String(),
ComputerName: computerName,
KDCAddr: kdcAddr,
Conn: tdpConn,
AuthorizeFn: authorize,
AllowClipboard: authCtx.Checker.DesktopClipboard(),
AllowDirectorySharing: authCtx.Checker.DesktopDirectorySharing(),
@@ -857,6 +889,7 @@ func (s *WindowsService) connectRDP(ctx context.Context, log *slog.Logger, tdpCo
Height: height,
AD: !desktop.NonAD(),
NLA: nla,
ClientProtocol: clientProtocol,
})
// before we check the error above, we grab the Windows user so that
// future audit events include the proper username
@@ -974,6 +1007,30 @@ func populateCertMetadata(metadata *events.WindowsCertificateMetadata, cert *x50
metadata.EnhancedKeyUsage = enhancedKeyUsages
}
func (s *WindowsService) recordEvent(ctx context.Context, t time.Time, delay int64, m tdp.Message, data []byte, recorder libevents.SessionPreparerRecorder) {
e := &events.DesktopRecording{
Metadata: events.Metadata{
Type: libevents.DesktopRecordingEvent,
Time: t,
},
TDPBMessage: data,
DelayMilliseconds: delay,
}
if len(data) > libevents.MaxProtoMessageSizeBytes {
// Technically a PNG frame is unbounded and could be too big for a single protobuf.
// In practice though, Windows limits RDP bitmaps to 64x64 pixels, and we compress
// the PNGs before they get here, so most PNG frames are under 500 bytes. The largest
// ones are around 2000 bytes. Anything approaching the limit of a single protobuf
// is likely some sort of DoS attempt and not legitimate RDP traffic, so we don't log it.
s.cfg.Logger.WarnContext(ctx, "refusing to record message", "len", len(data), "type", logutils.TypeAttr(m))
} else {
if err := libevents.SetupAndRecordEvent(ctx, recorder, e); err != nil {
s.cfg.Logger.WarnContext(ctx, "could not record desktop recording event", "error", err)
}
}
}
func (s *WindowsService) makeTDPSendHandler(
ctx context.Context,
recorder libevents.SessionPreparerRecorder,
@@ -981,45 +1038,22 @@ func (s *WindowsService) makeTDPSendHandler(
tdpConn *tdp.Conn,
audit *desktopSessionAuditor,
) func(m tdp.Message, b []byte) {
return func(m tdp.Message, b []byte) {
switch b[0] {
case byte(legacy.TypeRDPConnectionInitialized), byte(legacy.TypeRDPFastPathPDU), byte(legacy.TypePNG2Frame),
byte(legacy.TypePNGFrame), byte(legacy.TypeError), byte(legacy.TypeAlert):
e := &events.DesktopRecording{
Metadata: events.Metadata{
Type: libevents.DesktopRecordingEvent,
Time: s.cfg.Clock.Now().UTC().Round(time.Millisecond),
},
Message: b,
DelayMilliseconds: delay(),
}
if e.Size() > libevents.MaxProtoMessageSizeBytes {
// Technically a PNG frame is unbounded and could be too big for a single protobuf.
// In practice though, Windows limits RDP bitmaps to 64x64 pixels, and we compress
// the PNGs before they get here, so most PNG frames are under 500 bytes. The largest
// ones are around 2000 bytes. Anything approaching the limit of a single protobuf
// is likely some sort of DoS attempt and not legitimate RDP traffic, so we don't log it.
s.cfg.Logger.WarnContext(ctx, "refusing to record PNG frame, image too large", "len", len(b))
} else {
if err := libevents.SetupAndRecordEvent(ctx, recorder, e); err != nil {
s.cfg.Logger.WarnContext(ctx, "could not record desktop recording event", "error", err)
}
}
case byte(legacy.TypeClipboardData):
if clip, ok := m.(legacy.ClipboardData); ok {
// the TDP send handler emits a clipboard receive event, because we
// received clipboard data from the remote desktop and are sending
// it on the TDP connection
rxEvent := audit.makeClipboardReceive(int32(len(clip)))
s.emit(ctx, rxEvent)
}
case byte(legacy.TypeSharedDirectoryAcknowledge):
if message, ok := m.(legacy.SharedDirectoryAcknowledge); ok {
s.emit(ctx, audit.makeSharedDirectoryStart(message))
}
case byte(legacy.TypeSharedDirectoryReadRequest):
if message, ok := m.(legacy.SharedDirectoryReadRequest); ok {
errorEvent := audit.onSharedDirectoryReadRequest(message)
return func(msg tdp.Message, data []byte) {
switch m := msg.(type) {
case *tdpb.ServerHello, *tdpb.FastPathPDU, *tdpb.PNGFrame, *tdpb.Alert:
s.recordEvent(ctx, s.cfg.Clock.Now().UTC().Round(time.Millisecond), delay(), m, data, recorder)
case *tdpb.ClipboardData:
// the TDP send handler emits a clipboard receive event, because we
// received clipboard data from the remote desktop and are sending
// it on the TDP connection
rxEvent := audit.makeClipboardReceive(int32(len(m.Data)))
s.emit(ctx, rxEvent)
case *tdpb.SharedDirectoryAcknowledge:
s.emit(ctx, audit.makeSharedDirectoryStart(m))
case *tdpb.SharedDirectoryRequest:
switch req := m.Operation.(type) {
case *tdpbv1.SharedDirectoryRequest_Write_:
errorEvent := audit.onSharedDirectoryWriteRequest(completionID(m.CompletionId), directoryID(m.DirectoryId), req.Write)
if errorEvent != nil {
// if we can't audit due to a full cache, abort the connection
// as a security measure
@@ -1028,10 +1062,8 @@ func (s *WindowsService) makeTDPSendHandler(
}
s.emit(ctx, errorEvent)
}
}
case byte(legacy.TypeSharedDirectoryWriteRequest):
if message, ok := m.(legacy.SharedDirectoryWriteRequest); ok {
errorEvent := audit.onSharedDirectoryWriteRequest(message)
case *tdpbv1.SharedDirectoryRequest_Read_:
errorEvent := audit.onSharedDirectoryReadRequest(completionID(m.CompletionId), directoryID(m.DirectoryId), req.Read)
if errorEvent != nil {
// if we can't audit due to a full cache, abort the connection
// as a security measure
@@ -1054,36 +1086,21 @@ func (s *WindowsService) makeTDPReceiveHandler(
) func(m tdp.Message) {
return func(m tdp.Message) {
switch msg := m.(type) {
case legacy.ClientScreenSpec, legacy.MouseButton, legacy.MouseMove:
case *tdpb.ClientScreenSpec, *tdpb.MouseButton, *tdpb.MouseMove:
b, err := m.Encode()
if err != nil {
s.cfg.Logger.WarnContext(ctx, "could not emit desktop recording event", "error", err)
}
e := &events.DesktopRecording{
Metadata: events.Metadata{
Type: libevents.DesktopRecordingEvent,
Time: s.cfg.Clock.Now().UTC().Round(time.Millisecond),
},
Message: b,
DelayMilliseconds: delay(),
}
if e.Size() > libevents.MaxProtoMessageSizeBytes {
// screen spec, mouse button, and mouse move are fixed size messages,
// so they cannot exceed the maximum size
s.cfg.Logger.WarnContext(ctx, "refusing to record message", "len", len(b), "type", logutils.TypeAttr(m))
} else {
if err := libevents.SetupAndRecordEvent(ctx, recorder, e); err != nil {
s.cfg.Logger.WarnContext(ctx, "could not record desktop recording event", "error", err)
}
}
case legacy.ClipboardData:
s.recordEvent(ctx, s.cfg.Clock.Now().UTC().Round(time.Millisecond), delay(), m, b, recorder)
case *tdpb.ClipboardData:
// the TDP receive handler emits a clipboard send event, because we
// received clipboard data from the user (over TDP) and are sending
// it to the remote desktop
sendEvent := audit.makeClipboardSend(int32(len(msg)))
sendEvent := audit.makeClipboardSend(int32(len(msg.Data)))
s.emit(ctx, sendEvent)
case legacy.SharedDirectoryAnnounce:
errorEvent := audit.onSharedDirectoryAnnounce(m.(legacy.SharedDirectoryAnnounce))
case *tdpb.SharedDirectoryAnnounce:
errorEvent := audit.onSharedDirectoryAnnounce(m.(*tdpb.SharedDirectoryAnnounce))
if errorEvent != nil {
// if we can't audit due to a full cache, abort the connection
// as a security measure
@@ -1093,12 +1110,15 @@ func (s *WindowsService) makeTDPReceiveHandler(
}
s.emit(ctx, errorEvent)
}
case legacy.SharedDirectoryReadResponse:
case *tdpb.SharedDirectoryResponse:
// shared directory audit events can be noisy, so we use a compactor
// to retain and delay them in an attempt to coalesce contiguous events
audit.compactor.handleRead(ctx, audit.makeSharedDirectoryReadResponse(msg))
case legacy.SharedDirectoryWriteResponse:
audit.compactor.handleWrite(ctx, audit.makeSharedDirectoryWriteResponse(msg))
switch op := msg.Operation.(type) {
case *tdpbv1.SharedDirectoryResponse_Read_:
audit.compactor.handleRead(ctx, audit.makeSharedDirectoryReadResponse(completionID(msg.CompletionId), msg.ErrorCode, op.Read))
case *tdpbv1.SharedDirectoryResponse_Write_:
audit.compactor.handleWrite(ctx, audit.makeSharedDirectoryWriteResponse(completionID(msg.CompletionId), msg.ErrorCode, op.Write))
}
}
}
}
@@ -1361,10 +1381,12 @@ type monitorErrorSender struct {
}
func (m *monitorErrorSender) WriteString(s string) (n int, err error) {
if err := m.tdpConn.WriteMessage(legacy.Alert{Message: s, Severity: legacy.SeverityError}); err != nil {
return 0, trace.Wrap(err, "sending TDP error message")
if err := m.tdpConn.WriteMessage(&tdpb.Alert{
Severity: tdpbv1.AlertSeverity_ALERT_SEVERITY_ERROR,
Message: s,
}); err != nil {
return 0, trace.Wrap(err, "sending TDPB error message")
}
return len(s), nil
}
+22 -20
View File
@@ -34,10 +34,13 @@ import (
"testing"
"time"
"github.com/google/go-cmp/cmp"
"github.com/jonboulle/clockwork"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"google.golang.org/protobuf/testing/protocmp"
tdpbv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/desktop/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/types/events"
"github.com/gravitational/teleport/lib/auth/authclient"
@@ -47,7 +50,7 @@ import (
"github.com/gravitational/teleport/lib/modules"
"github.com/gravitational/teleport/lib/service/servicecfg"
"github.com/gravitational/teleport/lib/srv/desktop/tdp"
"github.com/gravitational/teleport/lib/srv/desktop/tdp/protocol/legacy"
"github.com/gravitational/teleport/lib/srv/desktop/tdp/protocol/tdpb"
"github.com/gravitational/teleport/lib/tlsca"
logutils "github.com/gravitational/teleport/lib/utils/log"
"github.com/gravitational/teleport/lib/utils/log/logtest"
@@ -310,21 +313,19 @@ func TestEmitsRecordingEventsOnSend(t *testing.T) {
emitter := &eventstest.MockRecorderEmitter{}
emitterPreparer := libevents.WithNoOpPreparer(emitter)
// a fake PNG Frame message
encoded := []byte{byte(legacy.TypePNGFrame), 0x01, 0x02}
delay := func() int64 { return 0 }
handler := s.makeTDPSendHandler(context.Background(), emitterPreparer, delay, nil /* conn */, nil /* auditor */)
// the handler accepts both the message structure and its encoded form,
// but our logic only depends on the encoded form, so pass a nil message
handler(nil /* message */, encoded)
msg := &tdpb.PNGFrame{Data: []byte{0x01, 0x02}}
encoded, err := msg.Encode()
require.NoError(t, err)
handler(msg, encoded)
e := emitter.LastEvent()
require.NotNil(t, e)
dr, ok := e.(*events.DesktopRecording)
require.True(t, ok)
require.Equal(t, encoded, dr.Message)
require.Equal(t, encoded, dr.TDPBMessage)
}
func TestSkipsExtremelyLargePNGs(t *testing.T) {
@@ -341,15 +342,14 @@ func TestSkipsExtremelyLargePNGs(t *testing.T) {
// a fake PNG Frame message, which is way too big to be legitimate
maliciousPNG := make([]byte, libevents.MaxProtoMessageSizeBytes+1)
rand.Read(maliciousPNG)
maliciousPNG[0] = byte(legacy.TypePNGFrame)
png := &tdpb.PNGFrame{Data: maliciousPNG}
encoded, err := png.Encode()
require.NoError(t, err)
delay := func() int64 { return 0 }
handler := s.makeTDPSendHandler(context.Background(), emitterPreparer, delay, nil /* conn */, nil /* auditor */)
// the handler accepts both the message structure and its encoded form,
// but our logic only depends on the encoded form, so pass a nil message
var msg tdp.Message
handler(msg, maliciousPNG)
handler(png, encoded)
require.Nil(t, emitter.LastEvent())
}
@@ -367,9 +367,9 @@ func TestEmitsRecordingEventsOnReceive(t *testing.T) {
delay := func() int64 { return 0 }
handler := s.makeTDPReceiveHandler(context.Background(), emitterPreparer, delay, nil /* conn */, nil /* auditor */)
msg := legacy.MouseButton{
Button: legacy.LeftMouseButton,
State: legacy.ButtonPressed,
msg := &tdpb.MouseButton{
Button: tdpbv1.MouseButtonType_MOUSE_BUTTON_TYPE_LEFT,
Pressed: true,
}
handler(msg)
@@ -377,9 +377,9 @@ func TestEmitsRecordingEventsOnReceive(t *testing.T) {
require.NotNil(t, e)
dr, ok := e.(*events.DesktopRecording)
require.True(t, ok)
decoded, err := legacy.Decode(bytes.NewBuffer(dr.Message))
decoded, err := tdpb.DecodePermissive(bytes.NewBuffer(dr.TDPBMessage))
require.NoError(t, err)
require.Equal(t, msg, decoded)
require.Empty(t, cmp.Diff((*tdpbv1.MouseButton)(msg), (*tdpbv1.MouseButton)(decoded.(*tdpb.MouseButton)), protocmp.Transform()))
}
func TestEmitsClipboardSendEvents(t *testing.T) {
@@ -404,7 +404,9 @@ func TestEmitsClipboardSendEvents(t *testing.T) {
rand.Read(fakeClipboardData)
start := s.cfg.Clock.Now().UTC()
msg := legacy.ClipboardData(fakeClipboardData)
msg := &tdpb.ClipboardData{
Data: fakeClipboardData,
}
handler(msg)
e := emitter.LastEvent()
@@ -440,7 +442,7 @@ func TestEmitsClipboardReceiveEvents(t *testing.T) {
rand.Read(fakeClipboardData)
start := s.cfg.Clock.Now().UTC()
msg := legacy.ClipboardData(fakeClipboardData)
msg := &tdpb.ClipboardData{Data: fakeClipboardData}
encoded, err := msg.Encode()
require.NoError(t, err)
handler(msg, encoded)
+2 -1
View File
@@ -106,6 +106,7 @@ import (
"github.com/gravitational/teleport/lib/services"
"github.com/gravitational/teleport/lib/services/readonly"
"github.com/gravitational/teleport/lib/session"
"github.com/gravitational/teleport/lib/srv/desktop/tdp/protocol/tdpb"
"github.com/gravitational/teleport/lib/utils"
logutils "github.com/gravitational/teleport/lib/utils/log"
"github.com/gravitational/teleport/lib/web/app"
@@ -1054,7 +1055,7 @@ func (h *Handler) bindDefaultEndpoints() {
h.GET("/webapi/sites/:site/desktopservices", h.WithClusterAuth(h.clusterDesktopServicesGet))
h.GET("/webapi/sites/:site/desktops/:desktopName", h.WithClusterAuth(h.getDesktopHandle))
// GET /webapi/sites/:site/desktops/:desktopName/connect?username=<username>&width=<width>&height=<height>
h.GET("/webapi/sites/:site/desktops/:desktopName/connect/ws", h.WithClusterAuthWebSocket(h.desktopConnectHandle, WithSubprotocols(protocolTDPB)))
h.GET("/webapi/sites/:site/desktops/:desktopName/connect/ws", h.WithClusterAuthWebSocket(h.desktopConnectHandle, WithSubprotocols(tdpb.ProtocolName)))
// GET /webapi/sites/:site/desktopplayback/:sid/ws
h.GET("/webapi/sites/:site/desktopplayback/:sid/ws", h.WithClusterAuthWebSocket(h.desktopPlaybackHandle))
h.GET("/webapi/sites/:site/desktops/:desktopName/active", h.WithClusterAuth(h.desktopIsActive))
+35 -13
View File
@@ -57,7 +57,6 @@ import (
const (
tdpbQueryParameter = "tdpb"
protocolTDPB = "teleport-tdpb-1.0"
protocolTDP = "teleport-tdp"
)
@@ -226,7 +225,26 @@ func (t *tdpHandshaker) forwardTDPB(w io.Writer, username string, _ bool) error
hello.KeyboardLayout = t.keyboardLayout.KeyboardLayout
}
return trace.Wrap(sendAll(w, append([]tdp.Message{hello}, t.withheld...)))
withheld, err := translateAll(t.withheld, tdpb.TranslateToModern)
if err != nil {
return trace.Wrap(err)
}
return trace.Wrap(sendAll(w, append([]tdp.Message{hello}, withheld...)))
}
func translateAll(messages []tdp.Message, translate func(tdp.Message) ([]tdp.Message, error)) ([]tdp.Message, error) {
translated := make([]tdp.Message, 0, len(messages))
for _, msg := range messages {
out, err := translate(msg)
if err != nil {
return nil, trace.Wrap(err)
}
if len(out) > 0 {
translated = append(translated, out...)
}
}
return translated, nil
}
// implements handshaker for TDPB clients
@@ -292,7 +310,11 @@ func (t *tdpbHandshaker) forwardTDP(w io.Writer, username string, forwardKeyboar
messages = append(messages, legacy.ClientKeyboardLayout{KeyboardLayout: t.hello.KeyboardLayout})
}
return sendAll(w, append(messages, t.withheld...))
withheld, err := translateAll(t.withheld, tdpb.TranslateToLegacy)
if err != nil {
return trace.Wrap(err)
}
return sendAll(w, append(messages, withheld...))
}
func (t *tdpbHandshaker) forwardTDPB(w io.Writer, username string, _ bool) error {
@@ -319,7 +341,7 @@ type handshaker interface {
// creates a handshaker instance that interops with either TDP or TDPB clients
func newHandshaker(protocol string, ws *websocket.Conn) handshaker {
if protocol == protocolTDPB {
if protocol == tdpb.ProtocolName {
return &tdpbHandshaker{
connection: &desktopWebsocketAdapter{conn: ws},
}
@@ -435,7 +457,7 @@ func (h *Handler) createDesktopConnection(
serverProtocol = protocolTDP
sendKeyboardLayout, _ := utils.MinVerWithoutPreRelease(version, "18.0.0")
err = handshaker.forwardTDP(serviceConnTLS, username, sendKeyboardLayout)
case protocolTDPB:
case tdpb.ProtocolName:
err = handshaker.forwardTDPB(serviceConnTLS, username, true /* unused */)
default:
err = trace.BadParameter("Unknown desktop agent protocol %v", serverProtocol)
@@ -576,7 +598,7 @@ func (h *Handler) createDesktopTLSConfig(
}
tlsConfig.Certificates = []tls.Certificate{certConf}
tlsConfig.NextProtos = []string{protocolTDPB}
tlsConfig.NextProtos = []string{tdpb.ProtocolName}
// Pass target desktop name via SNI.
tlsConfig.ServerName = desktopName + SNISuffix
return tlsConfig, nil
@@ -650,8 +672,8 @@ func readClientProtocol(r *http.Request) (string, error) {
switch tdpbVersion {
case "":
return protocolTDP, nil
case protocolTDPB:
return protocolTDPB, nil
case tdpb.ProtocolName:
return tdpb.ProtocolName, nil
default:
return "", trace.BadParameter("unknown TDPB version %q", tdpbVersion)
}
@@ -733,7 +755,7 @@ func (d desktopPinger) pingTDPB(ctx context.Context) error {
}
func newConn(rwc io.ReadWriteCloser, protocol string) *tdp.Conn {
if protocol == protocolTDPB {
if protocol == tdpb.ProtocolName {
return tdp.NewConn(rwc, tdp.DecoderAdapter(tdpb.DecodePermissive))
}
return tdp.NewConn(rwc, legacy.Decode)
@@ -790,16 +812,16 @@ func (p desktopWebsocketProxy) run(ctx context.Context) error {
needTranslation := p.clientProtocol != p.serverProtocol
if needTranslation {
// Translation is needed
if p.serverProtocol == protocolTDPB {
p.log.InfoContext(ctx, "Proxying desktop connection with translation", "server_dialect", protocolTDPB, "client_dialect", protocolTDP)
// Server speaks TDPB
if p.serverProtocol == tdpb.ProtocolName {
p.log.InfoContext(ctx, "Proxying desktop connection with translation", "server_dialect", tdpb.ProtocolName, "client_dialect", protocolTDP)
// Agent speaks TDPB
// Translate to TDPB when writing to the server. Intercept pings when reading from the server.
serverConn = tdp.NewReadWriteInterceptor(serverConn, nil, tdpb.TranslateToModern)
// Client speaks TDP
// Translate to TDP (legacy) when writing to this connection
clientConn = tdp.NewReadWriteInterceptor(clientConn, nil, tdpb.TranslateToLegacy)
} else {
p.log.InfoContext(ctx, "Proxying desktop connection with translation", "server_dialect", protocolTDP, "client_dialect", protocolTDPB)
p.log.InfoContext(ctx, "Proxying desktop connection with translation", "server_dialect", protocolTDP, "client_dialect", tdpb.ProtocolName)
// Agent speaks TDP
// Translate to TDPB when reading from this connection.
serverConn = tdp.NewReadWriteInterceptor(serverConn, nil, tdpb.TranslateToLegacy)
+49 -7
View File
@@ -202,7 +202,7 @@ func TestProxyConnection(t *testing.T) {
{
name: "tdp-tdpb",
clientProtocol: protocolTDP,
serverProtocol: protocolTDPB,
serverProtocol: tdpb.ProtocolName,
version: "17.5.0",
clientFn: tdpClient,
echoFn: tdpbEchoServer,
@@ -217,7 +217,7 @@ func TestProxyConnection(t *testing.T) {
},
{
name: "tdpb-tdp",
clientProtocol: protocolTDPB,
clientProtocol: tdpb.ProtocolName,
serverProtocol: protocolTDP,
version: "17.5.0",
clientFn: tdpbClient,
@@ -225,8 +225,8 @@ func TestProxyConnection(t *testing.T) {
},
{
name: "tdpb-tdpb",
clientProtocol: protocolTDPB,
serverProtocol: protocolTDPB,
clientProtocol: tdpb.ProtocolName,
serverProtocol: tdpb.ProtocolName,
version: "17.5.0",
clientFn: tdpbClient,
echoFn: tdpbEchoServer,
@@ -234,7 +234,7 @@ func TestProxyConnection(t *testing.T) {
{
name: "tdp-tdpb-no-latency-monitor",
clientProtocol: protocolTDP,
serverProtocol: protocolTDPB,
serverProtocol: tdpb.ProtocolName,
/* server version does not support latency monitoring */
version: "17.0.0",
clientFn: tdpClient,
@@ -339,7 +339,7 @@ func TestHandshaker(t *testing.T) {
defer handshaker.Close()
defer client.Close()
shaker := newHandshaker(protocolTDPB, handshaker)
shaker := newHandshaker(tdpb.ProtocolName, handshaker)
done := make(chan error)
go func() {
@@ -355,7 +355,6 @@ func TestHandshaker(t *testing.T) {
KeyboardLayout: 12,
}
require.NoError(t, tdp.EncodeTo(&clientConn, hello))
require.NoError(t, tdp.EncodeTo(&clientConn, &tdpb.PNGFrame{Data: []byte("somedata")}))
// Should succeed
require.NoError(t, <-done)
@@ -399,6 +398,49 @@ func TestHandshaker(t *testing.T) {
require.ErrorIs(t, err, io.EOF)
})
t.Run("withheld-tdpb-messages-are-translated", func(t *testing.T) {
tdpbHandshaker := tdpbHandshaker{
hello: &tdpb.ClientHello{
ScreenSpec: &tdpbv1.ClientScreenSpec{
Width: 10,
Height: 10,
},
KeyboardLayout: 1,
},
withheld: []tdp.Message{&tdpb.MouseMove{X: 1, Y: 2}},
}
buf := bytes.NewBuffer(nil)
require.NoError(t, tdpbHandshaker.forwardTDP(buf, "someuser", false))
username, err := legacy.Decode(buf)
require.IsType(t, legacy.ClientUsername{}, username)
require.NoError(t, err)
screenSpec, err := legacy.Decode(buf)
require.NoError(t, err)
require.IsType(t, legacy.ClientScreenSpec{}, screenSpec)
pngFrame, err := legacy.Decode(buf)
require.NoError(t, err)
require.IsType(t, legacy.MouseMove{}, pngFrame)
})
t.Run("withheld-tdp-messages-are-translated", func(t *testing.T) {
tdbHandshaker := tdpHandshaker{
screenSpec: legacy.ClientScreenSpec{
Width: 10,
Height: 10,
},
withheld: []tdp.Message{legacy.MouseMove{X: 1, Y: 2}},
}
buf := bytes.NewBuffer(nil)
require.NoError(t, tdbHandshaker.forwardTDPB(buf, "someuser", false))
hello, err := tdpb.DecodePermissive(buf)
require.IsType(t, &tdpb.ClientHello{}, hello)
require.NoError(t, err)
pngFrame, err := tdpb.DecodePermissive(buf)
require.NoError(t, err)
require.IsType(t, &tdpb.MouseMove{}, pngFrame)
})
}
func TestDesktopWebsocketAdapter(t *testing.T) {
+11 -2
View File
@@ -368,8 +368,17 @@ export class TdpClient extends EventEmitter<EventMap> {
// processMessage should be await-ed when called,
// so that its internal await-or-not logic is obeyed.
async processMessage(buffer: ArrayBufferLike): Promise<void> {
const result = this.codec.decodeMessage(buffer);
async processMessage(
buffer: ArrayBufferLike,
codecOverride?: Codec
): Promise<void> {
let codec = this.codec;
if (codecOverride) {
// Allow the caller to override the codec.
codec = codecOverride;
}
const result = codec.decodeMessage(buffer);
if (!result) {
// Codec implementations *should* return an 'unknown' result kind
// instead of undefined, but double check anyway for safety.
@@ -19,6 +19,7 @@
import {
ClientScreenSpec,
selectDirectoryInBrowser,
TdpbCodec,
TdpClient,
TdpClientEvent,
} from 'shared/libs/tdp';
@@ -57,6 +58,7 @@ export class PlayerClient extends TdpClient {
private sendTimeUpdates = true;
private lastUpdateTime = 0;
private timeout = null;
private tdpbCodec = new TdpbCodec();
constructor({ url, setTime, setPlayerStatus, setStatusText }) {
super(
@@ -174,7 +176,16 @@ export class PlayerClient extends TdpClient {
this.scheduleNextUpdate(json.ms);
}
await super.processMessage(base64ToArrayBuffer(json.message));
// Handle TDPB recordings by switching to the TDPB codec
if (json.tdpb_message !== undefined) {
await super.processMessage(
base64ToArrayBuffer(json.tdpb_message),
this.tdpbCodec
);
} else {
// Handle TDP recordings
await super.processMessage(base64ToArrayBuffer(json.message));
}
}
}