diff --git a/tailcat_files.go b/tailcat_files.go index c294b78a4..9882d8c5a 100644 --- a/tailcat_files.go +++ b/tailcat_files.go @@ -45,7 +45,8 @@ type FileService struct { // SSHOptions configures the SSH server returned by // [Server.SSHConnHandler]. type SSHOptions struct { - // Shell enables shell and exec sessions. + // Shell enables shell and exec sessions, plus local TCP forwarding to + // loopback IP addresses on the server. Shell bool // Exec, if non-empty, is a command (program and arguments) that diff --git a/tailcat_ssh.go b/tailcat_ssh.go index 4dcec1185..b33bb15bb 100644 --- a/tailcat_ssh.go +++ b/tailcat_ssh.go @@ -51,6 +51,7 @@ func (s *Server) HandleTailscaleSSHConn(c net.Conn) { // sends a command, it is run by the user's shell (PowerShell on // Windows); otherwise an interactive login shell is started with a // PTY. The SFTP subsystem is served per opts.Files; see [SSHOptions]. +// Local TCP forwarding is also allowed to the server's loopback IP addresses. // With opts.Exec, every session instead runs that one command, and // nothing else is offered. func (s *Server) SSHConnHandler(opts SSHOptions) func(net.Conn) { @@ -86,11 +87,21 @@ func (s *Server) SSHConnHandler(opts SSHOptions) func(net.Conn) { subsystems["sftp"] = h } srv := &ssh.Server{ - Handler: handler, - PublicKeyHandler: publicKeyHandler, - ChannelHandlers: map[string]ssh.ChannelHandler{"session": ssh.DefaultSessionHandler}, + Handler: handler, + PublicKeyHandler: publicKeyHandler, + ChannelHandlers: map[string]ssh.ChannelHandler{ + "session": ssh.DefaultSessionHandler, + "direct-tcpip": ssh.DirectTCPIPHandler, + }, RequestHandlers: map[string]ssh.RequestHandler{}, SubsystemHandlers: subsystems, + LocalPortForwardingCallback: func(ctx ssh.Context, host string, port uint32) bool { + // File-only and forced-command services must not grant access + // to other local services. Shell users already have that access. + // Require IP literals so forwarding never depends on DNS resolution. + return opts.Shell && port > 0 && port <= 65535 && + net.ParseIP(host).IsLoopback() + }, } if publicKeyHandler == nil { srv.NoClientAuthHandler = func(ctx ssh.Context) error { return nil } diff --git a/tailcat_ssh_test.go b/tailcat_ssh_test.go index 2d74c0822..fd6366d0c 100644 --- a/tailcat_ssh_test.go +++ b/tailcat_ssh_test.go @@ -10,8 +10,11 @@ import ( "context" "crypto/ed25519" "crypto/rand" + "errors" "io" "net" + "net/http" + "net/http/httptest" "runtime" "strings" "testing" @@ -99,6 +102,56 @@ func (e *testSSHEnv) dialSSHClient(t *testing.T, config *gossh.ClientConfig) (*g return gossh.NewClient(sshConn, chans, reqs), nil } +func TestSSHLocalTCPForwarding(t *testing.T) { + t.Parallel() + const payload = "forwarded over SSH" + target := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + io.WriteString(w, payload) + })) + defer target.Close() + _, port, err := net.SplitHostPort(strings.TrimPrefix(target.URL, "http://")) + if err != nil { + t.Fatal(err) + } + for _, tt := range []struct { + name, host string + opts tailcat.SSHOptions + allow bool + }{ + {"Loopback", "127.0.0.1", tailcat.SSHOptions{Shell: true}, true}, + {"Localhost", "localhost", tailcat.SSHOptions{Shell: true}, false}, + {"RemoteIP", "192.0.2.1", tailcat.SSHOptions{Shell: true}, false}, + {"Hostname", "example.invalid", tailcat.SSHOptions{Shell: true}, false}, + {"FilesOnly", "127.0.0.1", tailcat.SSHOptions{Files: &tailcat.FileService{Dir: t.TempDir(), Mode: tailcat.FileServeRO}}, false}, + {"ForcedCommand", "127.0.0.1", tailcat.SSHOptions{Shell: true, Exec: []string{"unused-command"}}, false}, + } { + t.Run(tt.name, func(t *testing.T) { + c := setupSSHEnv(t, tt.opts).sshClient(t) + transport := &http.Transport{DialContext: c.DialContext} + defer transport.CloseIdleConnections() + client := &http.Client{Transport: transport, Timeout: 10 * time.Second} + res, err := client.Get("http://" + net.JoinHostPort(tt.host, port)) + if err == nil { + defer res.Body.Close() + } + if !tt.allow { + var openErr *gossh.OpenChannelError + if !errors.As(err, &openErr) || openErr.Reason != gossh.Prohibited { + t.Fatalf("forwarding error = %v; want administratively prohibited", err) + } + return + } + if err != nil { + t.Fatal(err) + } + got, err := io.ReadAll(res.Body) + if err != nil || string(got) != payload { + t.Fatalf("forwarded response = %q, %v; want %q", got, err, payload) + } + }) + } +} + func TestSSHPublicKeyAuthentication(t *testing.T) { t.Parallel() newSigner := func() gossh.Signer {