mirror of
https://github.com/gravitational/teleport.git
synced 2026-09-24 16:17:11 +08:00
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:
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -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
@@ -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),
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
@@ -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) {
|
||||
|
||||
@@ -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));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user