mirror of
https://github.com/gravitational/teleport.git
synced 2026-09-24 16:17:11 +08:00
Refactor hostname resolution for SSH connections via the WebUI (#35773)
* Refactor hostname resolution for SSH connections via the WebUI
Previously all hostname resolution was performed by sending either
the hostname or server uuid to Auth in a GetSSHTargets request.
The returned information was then sent to the UI so that the console
tab can be updated to include `user@host`. While this worked, it
required a round trip to Auth that had to iterate through the entire
set of nodes to do the resolution. During a concurrent session test
that tried to spawn several thousand web sessions it was discovered
that this cause a massive spike in CPU on Auth.
There are two reasons that Auth was used to resolve the hostname.
1) The proxy doesn't perform RBAC on behalf of the user. Instead
the proxy has an Auth client with the users identity.
2) The UI used to provide an input box that users could manually
enter the target host.
Since we no longer allow open dialing and the input box to enter
a connection string manually has been removed we no longer need
to worry about the second case. To avoid the round trip to auth
we can use the local Proxy cache to look up the hostname if we
wait until AFTER the ssh connection to the target is established.
At this point we can be sure that the user does have access to
the target since the node allowed the connection.
* Fix terminal tests
The testing utilities used to mimic the web ui half of the websocket
made assumptions that the session metadata would be the first message
provided. However, now that hostname resolution occurs after the
initial SSH connection has been established it is possible for other
messages to be sent first. This caused some tests to fail and others
to hang forever.
The duplicated terminal logic from (WebSuite) makeTerminal and
(testProxy makeTerminal has now consolidated into a single test
terminal utility. The new `terminal` now wraps a TerminalStream
which allows test code to make use of the message processing logic
already in place there instead of having to write custom logic
per test. All a test needs to do is provide handlers for any messages
that it desires to introspect.
This commit is contained in:
@@ -36,7 +36,6 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gogo/protobuf/proto"
|
||||
"github.com/gorilla/websocket"
|
||||
"github.com/gravitational/roundtrip"
|
||||
"github.com/gravitational/trace"
|
||||
@@ -64,7 +63,6 @@ import (
|
||||
"github.com/gravitational/teleport/lib/service"
|
||||
"github.com/gravitational/teleport/lib/service/servicecfg"
|
||||
"github.com/gravitational/teleport/lib/services"
|
||||
"github.com/gravitational/teleport/lib/session"
|
||||
"github.com/gravitational/teleport/lib/sshutils"
|
||||
"github.com/gravitational/teleport/lib/tlsca"
|
||||
"github.com/gravitational/teleport/lib/utils"
|
||||
@@ -1552,33 +1550,7 @@ func (w *WebClient) SSH(termReq web.TerminalRequest) (*web.TerminalStream, error
|
||||
}
|
||||
|
||||
defer resp.Body.Close()
|
||||
ty, raw, err := ws.ReadMessage()
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
if ty != websocket.BinaryMessage {
|
||||
return nil, trace.BadParameter("unexpected websocket message; got %d want %d", ty, websocket.BinaryMessage)
|
||||
}
|
||||
|
||||
var env web.Envelope
|
||||
err = proto.Unmarshal(raw, &env)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
type siteSessionGenerateResponse struct {
|
||||
Session session.Session `json:"session"`
|
||||
}
|
||||
|
||||
var sessResp siteSessionGenerateResponse
|
||||
err = json.Unmarshal([]byte(env.Payload), &sessResp)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
stream := web.NewTerminalStream(context.Background(), ws, utils.NewLoggerForTests())
|
||||
return stream, nil
|
||||
return web.NewTerminalStream(context.Background(), web.TerminalStreamConfig{WS: ws}), nil
|
||||
}
|
||||
|
||||
// AddClientCredentials adds authenticated credentials to a client.
|
||||
|
||||
@@ -39,7 +39,6 @@ import (
|
||||
"github.com/gravitational/teleport/api/types"
|
||||
"github.com/gravitational/teleport/lib/client"
|
||||
"github.com/gravitational/teleport/lib/session"
|
||||
"github.com/gravitational/teleport/lib/utils"
|
||||
"github.com/gravitational/teleport/lib/web"
|
||||
)
|
||||
|
||||
@@ -193,7 +192,7 @@ func connectToHost(ctx context.Context, tc *client.TeleportClient, webSession *w
|
||||
return nil, trace.BadParameter("unexpected websocket message received %d", ty)
|
||||
}
|
||||
|
||||
stream := web.NewTerminalStream(ctx, ws, utils.NewLogger())
|
||||
stream := web.NewTerminalStream(ctx, web.TerminalStreamConfig{WS: ws})
|
||||
return stream, trace.Wrap(err)
|
||||
}
|
||||
|
||||
|
||||
+8
-37
@@ -2932,7 +2932,7 @@ func (h *Handler) siteNodeConnect(
|
||||
clusterName := site.GetName()
|
||||
if req.SessionID.IsZero() {
|
||||
// An existing session ID was not provided so we need to create a new one.
|
||||
sessionData, err = h.generateSession(ctx, clt, &req, clusterName, sessionCtx)
|
||||
sessionData, err = h.generateSession(&req, clusterName, sessionCtx)
|
||||
if err != nil {
|
||||
h.log.WithError(err).Debug("Unable to generate new ssh session.")
|
||||
return nil, trace.Wrap(err)
|
||||
@@ -2980,8 +2980,8 @@ func (h *Handler) siteNodeConnect(
|
||||
terminalConfig := TerminalHandlerConfig{
|
||||
Term: req.Term,
|
||||
SessionCtx: sessionCtx,
|
||||
AuthProvider: clt,
|
||||
LocalAuthProvider: h.auth.accessPoint,
|
||||
UserAuthClient: clt,
|
||||
LocalAccessPoint: h.auth.accessPoint,
|
||||
DisplayLogin: displayLogin,
|
||||
SessionData: sessionData,
|
||||
KeepAliveInterval: keepAliveInterval,
|
||||
@@ -3012,11 +3012,11 @@ func (h *Handler) siteNodeConnect(
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (h *Handler) generateSession(ctx context.Context, clt auth.ClientI, req *TerminalRequest, clusterName string, scx *SessionContext) (session.Session, error) {
|
||||
func (h *Handler) generateSession(req *TerminalRequest, clusterName string, scx *SessionContext) (session.Session, error) {
|
||||
owner := scx.cfg.User
|
||||
h.log.Infof("Generating new session for %s\n", clusterName)
|
||||
|
||||
host, err := findByHost(ctx, clt, req.Server)
|
||||
host, port, err := serverHostPort(req.Server)
|
||||
if err != nil {
|
||||
return session.Session{}, trace.Wrap(err)
|
||||
}
|
||||
@@ -3031,10 +3031,10 @@ func (h *Handler) generateSession(ctx context.Context, clt auth.ClientI, req *Te
|
||||
return session.Session{
|
||||
Kind: types.SSHSessionKind,
|
||||
Login: req.Login,
|
||||
ServerID: host.id,
|
||||
ServerID: host,
|
||||
ClusterName: clusterName,
|
||||
ServerHostname: host.hostName,
|
||||
ServerHostPort: host.port,
|
||||
ServerHostname: host,
|
||||
ServerHostPort: port,
|
||||
Moderated: accessEvaluator.IsModerated(),
|
||||
ID: session.NewID(),
|
||||
Created: time.Now().UTC(),
|
||||
@@ -3093,35 +3093,6 @@ func findByQuery(ctx context.Context, clt auth.ClientI, query string) ([]hostInf
|
||||
return hosts, nil
|
||||
}
|
||||
|
||||
// findByHost return a host matching by the host name.
|
||||
func findByHost(ctx context.Context, clt auth.ClientI, serverName string) (*hostInfo, error) {
|
||||
initialHost, initialPort, err := serverHostPort(serverName)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
rsp, err := clt.GetSSHTargets(ctx, &proto.GetSSHTargetsRequest{
|
||||
Host: initialHost,
|
||||
Port: strconv.Itoa(initialPort),
|
||||
})
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
var host hostInfo
|
||||
if len(rsp.Servers) == 1 {
|
||||
host.hostName = rsp.Servers[0].GetHostname()
|
||||
host.id = rsp.Servers[0].GetName()
|
||||
host.port = 0
|
||||
} else {
|
||||
host.hostName = initialHost
|
||||
host.id = initialHost
|
||||
host.port = initialPort
|
||||
}
|
||||
|
||||
return &host, nil
|
||||
}
|
||||
|
||||
// fetchExistingSession fetches an active or pending SSH session by the SessionID passed in the TerminalRequest.
|
||||
func (h *Handler) fetchExistingSession(ctx context.Context, clt auth.ClientI, req *TerminalRequest, siteName string) (session.Session, types.SessionTracker, error) {
|
||||
sessionID, err := session.ParseID(req.SessionID.String())
|
||||
|
||||
+277
-512
File diff suppressed because it is too large
Load Diff
+16
-13
@@ -239,14 +239,14 @@ func (h *Handler) executeCommand(
|
||||
|
||||
commandHandlerConfig := CommandHandlerConfig{
|
||||
SessionCtx: sessionCtx,
|
||||
AuthProvider: clt,
|
||||
UserAuthClient: clt,
|
||||
SessionData: sessionData,
|
||||
KeepAliveInterval: keepAliveInterval,
|
||||
ProxyHostPort: h.ProxyHostPort(),
|
||||
InteractiveCommand: interactiveCommand,
|
||||
Router: h.cfg.Router,
|
||||
TracerProvider: h.cfg.TracerProvider,
|
||||
LocalAuthProvider: h.auth.accessPoint,
|
||||
LocalAccessPoint: h.auth.accessPoint,
|
||||
mfaFuncCache: mfaCacheFn,
|
||||
buffer: buffer,
|
||||
}
|
||||
@@ -472,13 +472,13 @@ func newCommandHandler(ctx context.Context, cfg CommandHandlerConfig) (*commandH
|
||||
"session_id": cfg.SessionData.ID.String(),
|
||||
}),
|
||||
ctx: cfg.SessionCtx,
|
||||
authProvider: cfg.AuthProvider,
|
||||
userAuthClient: cfg.UserAuthClient,
|
||||
sessionData: cfg.SessionData,
|
||||
keepAliveInterval: cfg.KeepAliveInterval,
|
||||
proxyHostPort: cfg.ProxyHostPort,
|
||||
interactiveCommand: cfg.InteractiveCommand,
|
||||
router: cfg.Router,
|
||||
localAuthProvider: cfg.LocalAuthProvider,
|
||||
localAccessPoint: cfg.LocalAccessPoint,
|
||||
tracer: cfg.tracer,
|
||||
},
|
||||
mfaAuthCache: cfg.mfaFuncCache,
|
||||
@@ -490,8 +490,8 @@ func newCommandHandler(ctx context.Context, cfg CommandHandlerConfig) (*commandH
|
||||
type CommandHandlerConfig struct {
|
||||
// SessionCtx is the context for the user's web session.
|
||||
SessionCtx *SessionContext
|
||||
// AuthProvider is used to fetch nodes and sessions from the backend.
|
||||
AuthProvider AuthProvider
|
||||
// UserAuthClient is used to fetch nodes and sessions from the backend via the users' identity.
|
||||
UserAuthClient UserAuthClient
|
||||
// SessionData is the data to send to the client on the initial session creation.
|
||||
SessionData session.Session
|
||||
// KeepAliveInterval is the interval for sending ping frames to a web client.
|
||||
@@ -506,9 +506,12 @@ type CommandHandlerConfig struct {
|
||||
Router *proxy.Router
|
||||
// TracerProvider is used to create the tracer
|
||||
TracerProvider oteltrace.TracerProvider
|
||||
// LocalAuthProvider is used to fetch user information from the
|
||||
// local cluster when connecting to agentless nodes.
|
||||
LocalAuthProvider agentless.AuthProvider
|
||||
// LocalAccessPoint is the subset of the Proxy cache required to
|
||||
// look up information from the local cluster. This should not
|
||||
// be used for anything that requires RBAC on behalf of the user.
|
||||
// Anything requests that should be made on behalf of the user should
|
||||
// use [UserAuthClient].
|
||||
LocalAccessPoint localAccessPoint
|
||||
// tracer is used to create spans
|
||||
tracer oteltrace.Tracer
|
||||
// mfaFuncCache is used to cache the MFA auth method
|
||||
@@ -534,8 +537,8 @@ func (t *CommandHandlerConfig) CheckAndSetDefaults() error {
|
||||
return trace.BadParameter("server: missing server")
|
||||
}
|
||||
|
||||
if t.AuthProvider == nil {
|
||||
return trace.BadParameter("AuthProvider must be provided")
|
||||
if t.UserAuthClient == nil {
|
||||
return trace.BadParameter("UserAuthClient must be provided")
|
||||
}
|
||||
|
||||
if t.SessionCtx == nil {
|
||||
@@ -550,8 +553,8 @@ func (t *CommandHandlerConfig) CheckAndSetDefaults() error {
|
||||
t.TracerProvider = tracing.DefaultProvider()
|
||||
}
|
||||
|
||||
if t.LocalAuthProvider == nil {
|
||||
return trace.BadParameter("LocalAuthProvider must be provided")
|
||||
if t.LocalAccessPoint == nil {
|
||||
return trace.BadParameter("localAccessPoint must be provided")
|
||||
}
|
||||
|
||||
if t.mfaFuncCache == nil {
|
||||
|
||||
+111
-51
@@ -96,9 +96,9 @@ type TerminalRequest struct {
|
||||
ParticipantMode types.SessionParticipantMode `json:"mode"`
|
||||
}
|
||||
|
||||
// AuthProvider is a subset of the full Auth API.
|
||||
type AuthProvider interface {
|
||||
GetNodes(ctx context.Context, namespace string) ([]types.Server, error)
|
||||
// UserAuthClient is a subset of the Auth API that performs
|
||||
// operations on behalf of the user so that the correct RBAC is applied.
|
||||
type UserAuthClient interface {
|
||||
GetSessionEvents(namespace string, sid session.ID, after int) ([]events.EventFields, error)
|
||||
GetSessionTracker(ctx context.Context, sessionID string) (types.SessionTracker, error)
|
||||
IsMFARequired(ctx context.Context, req *authproto.IsMFARequiredRequest) (*authproto.IsMFARequiredResponse, error)
|
||||
@@ -125,8 +125,8 @@ func NewTerminal(ctx context.Context, cfg TerminalHandlerConfig) (*TerminalHandl
|
||||
"session_id": cfg.SessionData.ID.String(),
|
||||
}),
|
||||
ctx: cfg.SessionCtx,
|
||||
authProvider: cfg.AuthProvider,
|
||||
localAuthProvider: cfg.LocalAuthProvider,
|
||||
userAuthClient: cfg.UserAuthClient,
|
||||
localAccessPoint: cfg.LocalAccessPoint,
|
||||
sessionData: cfg.SessionData,
|
||||
keepAliveInterval: cfg.KeepAliveInterval,
|
||||
proxyHostPort: cfg.ProxyHostPort,
|
||||
@@ -151,11 +151,14 @@ type TerminalHandlerConfig struct {
|
||||
Term session.TerminalParams
|
||||
// SessionCtx is the context for the users web session.
|
||||
SessionCtx *SessionContext
|
||||
// AuthProvider is used to fetch nodes and sessions from the backend.
|
||||
AuthProvider AuthProvider
|
||||
// LocalAuthProvider is used to fetch user information from the
|
||||
// local cluster when connecting to agentless nodes.
|
||||
LocalAuthProvider agentless.AuthProvider
|
||||
// UserAuthClient is used to fetch nodes and sessions from the backend.
|
||||
UserAuthClient UserAuthClient
|
||||
// LocalAccessPoint is the subset of the Proxy cache required to
|
||||
// look up information from the local cluster. This should not
|
||||
// be used for anything that requires RBAC on behalf of the user.
|
||||
// Requests that should be made on behalf of the user should
|
||||
// use [UserAuthClient].
|
||||
LocalAccessPoint localAccessPoint
|
||||
// DisplayLogin is the login name to display in the UI.
|
||||
DisplayLogin string
|
||||
// SessionData is the data to send to the client on the initial session creation.
|
||||
@@ -209,12 +212,12 @@ func (t *TerminalHandlerConfig) CheckAndSetDefaults() error {
|
||||
return trace.BadParameter("term: bad dimensions(%dx%d)", t.Term.W, t.Term.H)
|
||||
}
|
||||
|
||||
if t.AuthProvider == nil {
|
||||
return trace.BadParameter("AuthProvider must be provided")
|
||||
if t.UserAuthClient == nil {
|
||||
return trace.BadParameter("UserAuthClient must be provided")
|
||||
}
|
||||
|
||||
if t.LocalAuthProvider == nil {
|
||||
return trace.BadParameter("LocalAuthProvider must be provided")
|
||||
if t.LocalAccessPoint == nil {
|
||||
return trace.BadParameter("localAccessPoint must be provided")
|
||||
}
|
||||
|
||||
if t.SessionCtx == nil {
|
||||
@@ -244,8 +247,8 @@ type sshBaseHandler struct {
|
||||
log *logrus.Entry
|
||||
// ctx is a web session context for the currently logged-in user.
|
||||
ctx *SessionContext
|
||||
// authProvider is used to fetch nodes and sessions from the backend.
|
||||
authProvider AuthProvider
|
||||
// userAuthClient is used to fetch nodes and sessions from the backend via the users' identity.
|
||||
userAuthClient UserAuthClient
|
||||
// proxyHostPort is the address of the server to connect to.
|
||||
proxyHostPort string
|
||||
// proxyPublicAddr is the public web proxy address.
|
||||
@@ -260,13 +263,24 @@ type sshBaseHandler struct {
|
||||
router *proxy.Router
|
||||
// tracer creates spans
|
||||
tracer oteltrace.Tracer
|
||||
// localAuthProvider is used to fetch user information from the
|
||||
// local cluster when connecting to agentless nodes.
|
||||
localAuthProvider agentless.AuthProvider
|
||||
// localAccessPoint is the subset of the Proxy cache required to
|
||||
// look up information from the local cluster. This should not
|
||||
// be used for anything that requires RBAC on behalf of the user.
|
||||
// Requests that should be made on behalf of the user should
|
||||
// use [UserAuthClient].
|
||||
localAccessPoint localAccessPoint
|
||||
// interactiveCommand is a command to execute.
|
||||
interactiveCommand []string
|
||||
}
|
||||
|
||||
// localAccessPoint is a subset of the cache used to look up
|
||||
// various cluster details.
|
||||
type localAccessPoint interface {
|
||||
GetUser(ctx context.Context, username string, withSecrets bool) (types.User, error)
|
||||
GetRole(ctx context.Context, name string) (types.Role, error)
|
||||
GetNode(ctx context.Context, namespace, name string) (types.Server, error)
|
||||
}
|
||||
|
||||
// TerminalHandler connects together an SSH session with a web-based
|
||||
// terminal via a web socket.
|
||||
type TerminalHandler struct {
|
||||
@@ -333,43 +347,60 @@ func (t *TerminalHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
var sessionMetadataResponse []byte
|
||||
t.handler(ws, r)
|
||||
}
|
||||
|
||||
func (t *TerminalHandler) writeSessionData(ctx context.Context) error {
|
||||
envelope := &Envelope{
|
||||
Version: defaults.WebsocketVersion,
|
||||
Type: defaults.WebsocketSessionMetadata,
|
||||
}
|
||||
|
||||
sessionDataTemp := t.sessionData
|
||||
|
||||
// If the displayLogin is set then use it in the session metadata instead of the
|
||||
// login name used in the SSH connection. This is specifically for the use case
|
||||
// when joining a session to avoid displaying "-teleport-internal-join" as the username.
|
||||
if t.displayLogin != "" {
|
||||
sessionDataTemp := t.sessionData
|
||||
sessionDataTemp.Login = t.displayLogin
|
||||
sessionMetadataResponse, err = json.Marshal(siteSessionGenerateResponse{Session: sessionDataTemp})
|
||||
sessionMetadataResponse, err := json.Marshal(siteSessionGenerateResponse{Session: sessionDataTemp})
|
||||
if err != nil {
|
||||
t.sendError("unable to marshal session response", err, t.stream.ws)
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
envelope.Payload = string(sessionMetadataResponse)
|
||||
} else {
|
||||
sessionMetadataResponse, err = json.Marshal(siteSessionGenerateResponse{Session: t.sessionData})
|
||||
}
|
||||
// The Proxy cache is used to retrieve the server and resolve the hostname here instead
|
||||
// of the user auth client to avoid a round trip to the Auth server. This would normally
|
||||
// not be ok since this bypasses user RBAC, however, since at this point we have already
|
||||
// established a connection to the target host via the user identity, the user MUST have
|
||||
// access to the target host.
|
||||
server, err := t.localAccessPoint.GetNode(ctx, apidefaults.Namespace, sessionDataTemp.ServerID)
|
||||
if err != nil {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
sessionDataTemp.ServerHostname = server.GetHostname()
|
||||
|
||||
if err != nil {
|
||||
t.sendError("unable to marshal session response", err, ws)
|
||||
return
|
||||
}
|
||||
|
||||
envelope := &Envelope{
|
||||
Version: defaults.WebsocketVersion,
|
||||
Type: defaults.WebsocketSessionMetadata,
|
||||
Payload: string(sessionMetadataResponse),
|
||||
sessionMetadataResponse, err := json.Marshal(siteSessionGenerateResponse{Session: sessionDataTemp})
|
||||
if err != nil {
|
||||
t.sendError("unable to marshal session response", err, t.stream.ws)
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
envelope.Payload = string(sessionMetadataResponse)
|
||||
}
|
||||
|
||||
envelopeBytes, err := proto.Marshal(envelope)
|
||||
if err != nil {
|
||||
t.sendError("unable to marshal session data event for web client", err, ws)
|
||||
return
|
||||
t.sendError("unable to marshal session data event for web client", err, t.stream.ws)
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
|
||||
err = ws.WriteMessage(websocket.BinaryMessage, envelopeBytes)
|
||||
if err != nil {
|
||||
t.sendError("unable to write message to socket", err, ws)
|
||||
return
|
||||
if err := t.stream.ws.WriteMessage(websocket.BinaryMessage, envelopeBytes); err != nil {
|
||||
t.sendError("unable to write message to socket", err, t.stream.ws)
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
|
||||
t.handler(ws, r)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Close the websocket stream.
|
||||
@@ -405,7 +436,7 @@ func (t *TerminalHandler) handler(ws *websocket.Conn, r *http.Request) {
|
||||
tctx := oteltrace.ContextWithRemoteSpanContext(context.Background(), oteltrace.SpanContextFromContext(r.Context()))
|
||||
ctx, cancel := context.WithCancel(tctx)
|
||||
defer cancel()
|
||||
t.stream = NewTerminalStream(ctx, ws, t.log)
|
||||
t.stream = NewTerminalStream(ctx, TerminalStreamConfig{WS: ws, Logger: t.log})
|
||||
|
||||
// Create a Teleport client, if not able to, show the reason to the user in
|
||||
// the terminal.
|
||||
@@ -515,7 +546,7 @@ func (t *TerminalHandler) makeClient(ctx context.Context, stream *TerminalStream
|
||||
|
||||
// issueSessionMFACerts performs the mfa ceremony to retrieve new certs that can be
|
||||
// used to access nodes which require per-session mfa. The ceremony is performed directly
|
||||
// to make use of the authProvider already established for the session instead of leveraging
|
||||
// to make use of the userAuthClient already established for the session instead of leveraging
|
||||
// the TeleportClient which would require dialing the auth server a second time.
|
||||
func (t *sshBaseHandler) issueSessionMFACerts(ctx context.Context, tc *client.TeleportClient, wsStream *WSStream) ([]ssh.AuthMethod, error) {
|
||||
ctx, span := t.tracer.Start(ctx, "terminal/issueSessionMFACerts")
|
||||
@@ -559,7 +590,7 @@ func (t *sshBaseHandler) issueSessionMFACerts(ctx context.Context, tc *client.Te
|
||||
}
|
||||
|
||||
key, _, err = client.PerformMFACeremony(ctx, client.PerformMFACeremonyParams{
|
||||
CurrentAuthClient: t.authProvider,
|
||||
CurrentAuthClient: t.userAuthClient,
|
||||
RootAuthClient: t.ctx.cfg.RootClient,
|
||||
MFAPrompt: mfa.PromptFunc(func(ctx context.Context, chal *authproto.MFAAuthenticateChallenge) (*authproto.MFAAuthenticateResponse, error) {
|
||||
span.AddEvent("prompting user with mfa challenge")
|
||||
@@ -631,7 +662,7 @@ func (t *sshBaseHandler) connectToHost(ctx context.Context, ws WSConn, tc *clien
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
signer := agentless.SignerFromSSHCertificate(cert, t.localAuthProvider, tc.SiteName, tc.Username)
|
||||
signer := agentless.SignerFromSSHCertificate(cert, t.localAccessPoint, tc.SiteName, tc.Username)
|
||||
|
||||
type clientRes struct {
|
||||
clt *client.NodeClient
|
||||
@@ -751,11 +782,15 @@ func (t *TerminalHandler) streamTerminal(ctx context.Context, tc *client.Telepor
|
||||
|
||||
defer nc.Close()
|
||||
|
||||
if err := t.writeSessionData(ctx); err != nil {
|
||||
t.log.WithError(err).Warn("Unable to stream terminal - failure sending session data")
|
||||
}
|
||||
|
||||
var beforeStart func(io.Writer)
|
||||
if t.participantMode == types.SessionModeratorMode {
|
||||
beforeStart = func(out io.Writer) {
|
||||
nc.OnMFA = func() {
|
||||
if err := t.presenceChecker(ctx, out, t.authProvider, t.sessionData.ID.String(), promptMFAChallenge(t.stream.WSStream, protobufMFACodec{})); err != nil {
|
||||
if err := t.presenceChecker(ctx, out, t.userAuthClient, t.sessionData.ID.String(), promptMFAChallenge(t.stream.WSStream, protobufMFACodec{})); err != nil {
|
||||
t.log.WithError(err).Warn("Unable to stream terminal - failure performing presence checks")
|
||||
return
|
||||
}
|
||||
@@ -968,20 +1003,45 @@ func NewWStream(ctx context.Context, ws WSConn, log logrus.FieldLogger, handlers
|
||||
return w
|
||||
}
|
||||
|
||||
// TerminalStreamConfig contains dependencies of a TerminalStream.
|
||||
type TerminalStreamConfig struct {
|
||||
// The websocket to operate over. Required.
|
||||
WS WSConn
|
||||
// A logger to emit log messages. Optional.
|
||||
Logger logrus.FieldLogger
|
||||
// A custom set of handlers to process messages received
|
||||
// over the websocket. Optional.
|
||||
Handlers map[string]WSHandlerFunc
|
||||
}
|
||||
|
||||
// NewTerminalStream creates a stream that manages reading and writing
|
||||
// data over the provided [websocket.Conn]
|
||||
func NewTerminalStream(ctx context.Context, ws WSConn, log logrus.FieldLogger) *TerminalStream {
|
||||
func NewTerminalStream(ctx context.Context, cfg TerminalStreamConfig) *TerminalStream {
|
||||
t := &TerminalStream{
|
||||
sessionReadyC: make(chan struct{}),
|
||||
}
|
||||
|
||||
handlers := map[string]WSHandlerFunc{
|
||||
defaults.WebsocketResize: t.handleWindowResize,
|
||||
defaults.WebsocketFileTransferRequest: t.handleFileTransferRequest,
|
||||
defaults.WebsocketFileTransferDecision: t.handleFileTransferDecision,
|
||||
if cfg.Handlers == nil {
|
||||
cfg.Handlers = map[string]WSHandlerFunc{}
|
||||
}
|
||||
|
||||
t.WSStream = NewWStream(ctx, ws, log, handlers)
|
||||
if _, ok := cfg.Handlers[defaults.WebsocketResize]; !ok {
|
||||
cfg.Handlers[defaults.WebsocketResize] = t.handleWindowResize
|
||||
}
|
||||
|
||||
if _, ok := cfg.Handlers[defaults.WebsocketFileTransferRequest]; !ok {
|
||||
cfg.Handlers[defaults.WebsocketFileTransferRequest] = t.handleFileTransferRequest
|
||||
}
|
||||
|
||||
if _, ok := cfg.Handlers[defaults.WebsocketFileTransferDecision]; !ok {
|
||||
cfg.Handlers[defaults.WebsocketFileTransferDecision] = t.handleFileTransferDecision
|
||||
}
|
||||
|
||||
if cfg.Logger == nil {
|
||||
cfg.Logger = utils.NewLogger()
|
||||
}
|
||||
|
||||
t.WSStream = NewWStream(ctx, cfg.WS, cfg.Logger, cfg.Handlers)
|
||||
|
||||
return t
|
||||
}
|
||||
|
||||
+176
-4
@@ -16,23 +16,32 @@
|
||||
* along with this program. If not, see <http://www.gnu.org/licenses/>.
|
||||
*/
|
||||
|
||||
package web_test
|
||||
package web
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gogo/protobuf/proto"
|
||||
"github.com/gorilla/websocket"
|
||||
"github.com/gravitational/roundtrip"
|
||||
"github.com/gravitational/trace"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/gravitational/teleport/api/types"
|
||||
"github.com/gravitational/teleport/lib/client"
|
||||
"github.com/gravitational/teleport/lib/defaults"
|
||||
"github.com/gravitational/teleport/lib/session"
|
||||
"github.com/gravitational/teleport/lib/utils"
|
||||
"github.com/gravitational/teleport/lib/web"
|
||||
)
|
||||
|
||||
// TestTerminalReadFromClosedConn verifies that Teleport recovers
|
||||
@@ -49,7 +58,7 @@ func TestTerminalReadFromClosedConn(t *testing.T) {
|
||||
t.Errorf("couldn't upgrade websocket connection: %v", err)
|
||||
}
|
||||
|
||||
envelope := web.Envelope{
|
||||
envelope := Envelope{
|
||||
Type: defaults.WebsocketRaw,
|
||||
Payload: "hello",
|
||||
}
|
||||
@@ -66,7 +75,9 @@ func TestTerminalReadFromClosedConn(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
defer resp.Body.Close()
|
||||
|
||||
stream := web.NewTerminalStream(context.Background(), conn, utils.NewLoggerForTests())
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
t.Cleanup(cancel)
|
||||
stream := NewTerminalStream(ctx, TerminalStreamConfig{WS: conn, Logger: utils.NewLoggerForTests()})
|
||||
|
||||
// close the stream before we attempt to read from it,
|
||||
// this will produce a net.ErrClosed error on the read
|
||||
@@ -75,3 +86,164 @@ func TestTerminalReadFromClosedConn(t *testing.T) {
|
||||
_, err = io.Copy(io.Discard, stream)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
type terminal struct {
|
||||
ws *websocket.Conn
|
||||
stream *TerminalStream
|
||||
|
||||
sessionC chan session.Session
|
||||
}
|
||||
|
||||
type connectConfig struct {
|
||||
pack *authPack
|
||||
host string
|
||||
proxy string
|
||||
sessionID session.ID
|
||||
participantMode types.SessionParticipantMode
|
||||
keepAliveInterval time.Duration
|
||||
mfaCeremony func(challenge client.MFAAuthenticateChallenge) []byte
|
||||
handlers map[string]WSHandlerFunc
|
||||
}
|
||||
|
||||
func connectToHost(ctx context.Context, cfg connectConfig) (*terminal, error) {
|
||||
req := TerminalRequest{
|
||||
Server: cfg.host,
|
||||
Login: cfg.pack.login,
|
||||
Term: session.TerminalParams{
|
||||
W: 100,
|
||||
H: 100,
|
||||
},
|
||||
SessionID: cfg.sessionID,
|
||||
ParticipantMode: cfg.participantMode,
|
||||
KeepAliveInterval: cfg.keepAliveInterval,
|
||||
}
|
||||
|
||||
data, err := json.Marshal(req)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
u := url.URL{
|
||||
Host: cfg.proxy,
|
||||
Scheme: client.WSS,
|
||||
Path: "/v1/webapi/sites/-current-/connect",
|
||||
}
|
||||
|
||||
q := u.Query()
|
||||
q.Set("params", string(data))
|
||||
q.Set(roundtrip.AccessTokenQueryParam, cfg.pack.session.Token)
|
||||
u.RawQuery = q.Encode()
|
||||
|
||||
header := http.Header{}
|
||||
header.Add("Origin", "http://localhost")
|
||||
for _, cookie := range cfg.pack.cookies {
|
||||
header.Add("Cookie", cookie.String())
|
||||
}
|
||||
|
||||
dialer := websocket.Dialer{
|
||||
TLSClientConfig: &tls.Config{
|
||||
InsecureSkipVerify: true,
|
||||
},
|
||||
}
|
||||
|
||||
ws, resp, err := dialer.Dial(u.String(), header)
|
||||
if err != nil {
|
||||
var sb strings.Builder
|
||||
sb.WriteString("websocket dial")
|
||||
if resp != nil {
|
||||
fmt.Fprintf(&sb, "; status code %v;", resp.StatusCode)
|
||||
fmt.Fprintf(&sb, "headers: %v; body: ", resp.Header)
|
||||
io.Copy(&sb, resp.Body)
|
||||
}
|
||||
return nil, trace.Wrap(err, sb.String())
|
||||
}
|
||||
if err := resp.Body.Close(); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
t := &terminal{ws: ws, sessionC: make(chan session.Session, 1)}
|
||||
|
||||
// If MFA is expected, it should be performed prior to creating
|
||||
// the TerminalStream to avoid messages being handled by multiple
|
||||
// readers.
|
||||
if cfg.mfaCeremony != nil {
|
||||
if err := t.performMFACeremony(cfg.mfaCeremony); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
}
|
||||
|
||||
if cfg.handlers == nil {
|
||||
cfg.handlers = map[string]WSHandlerFunc{}
|
||||
}
|
||||
|
||||
if _, ok := cfg.handlers[defaults.WebsocketSessionMetadata]; !ok {
|
||||
cfg.handlers[defaults.WebsocketSessionMetadata] = func(ctx context.Context, envelope Envelope) {
|
||||
if envelope.Type != defaults.WebsocketSessionMetadata {
|
||||
return
|
||||
}
|
||||
|
||||
var sessResp siteSessionGenerateResponse
|
||||
if err := json.Unmarshal([]byte(envelope.Payload), &sessResp); err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
t.sessionC <- sessResp.Session
|
||||
}
|
||||
}
|
||||
|
||||
t.stream = NewTerminalStream(ctx, TerminalStreamConfig{
|
||||
WS: ws,
|
||||
Logger: utils.NewLogger(),
|
||||
Handlers: cfg.handlers,
|
||||
})
|
||||
|
||||
return t, nil
|
||||
}
|
||||
|
||||
func (t *terminal) GetSession() session.Session {
|
||||
sess := <-t.sessionC
|
||||
t.sessionC <- sess
|
||||
|
||||
return sess
|
||||
}
|
||||
|
||||
func (t *terminal) Close() error {
|
||||
return t.stream.Close()
|
||||
}
|
||||
|
||||
func (t *terminal) Write(p []byte) (int, error) {
|
||||
return t.stream.Write(p)
|
||||
}
|
||||
|
||||
func (t *terminal) Read(p []byte) (int, error) {
|
||||
return t.stream.Read(p)
|
||||
}
|
||||
|
||||
func (t *terminal) performMFACeremony(ceremonyFn func(challenge client.MFAAuthenticateChallenge) []byte) error {
|
||||
// Wait for websocket authn challenge event.
|
||||
ty, raw, err := t.ws.ReadMessage()
|
||||
if err != nil {
|
||||
return trace.Wrap(err, "reading ws message")
|
||||
}
|
||||
|
||||
if ty != websocket.BinaryMessage {
|
||||
return trace.BadParameter("got unexpected websocket message type %d", ty)
|
||||
}
|
||||
|
||||
var env Envelope
|
||||
if err := proto.Unmarshal(raw, &env); err != nil {
|
||||
return trace.Wrap(err, "unmarshalling envelope")
|
||||
}
|
||||
|
||||
var challenge client.MFAAuthenticateChallenge
|
||||
if err := json.Unmarshal([]byte(env.Payload), &challenge); err != nil {
|
||||
return trace.Wrap(err, "unmarshalling challenge")
|
||||
}
|
||||
|
||||
// Send response over ws.
|
||||
if err := t.ws.WriteMessage(websocket.BinaryMessage, ceremonyFn(challenge)); err != nil {
|
||||
return trace.Wrap(err, "sending challenge response")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user