Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions jobs/ssh_proxy/spec
Original file line number Diff line number Diff line change
Expand Up @@ -87,6 +87,9 @@ properties:
diego.ssh_proxy.idle_connection_timeout_in_seconds:
description: Idle timeout for incoming connections
default: 300
diego.ssh_proxy.max_connection_duration_in_seconds:
description: Maximum lifetime of an SSH connection, in seconds, measured from successful client authentication and including backend connection setup. On expiry, the connection and all its channels are closed abruptly without a reason message to the client. Must be a non-negative integer; 0 means unlimited.
default: 0

diego.ssh_proxy.uaa.url:
description: The domain name of the UAA
Expand Down
10 changes: 9 additions & 1 deletion jobs/ssh_proxy/templates/ssh_proxy.json.erb
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
<% require 'ipaddr' %>
<%=

def parse_ip (ip, var_name)
Expand Down Expand Up @@ -39,6 +40,14 @@
config[:idle_connection_timeout] = "#{value}s"
end

max_connection_duration = p("diego.ssh_proxy.max_connection_duration_in_seconds")
unless max_connection_duration.is_a?(Integer) && max_connection_duration >= 0
raise "diego.ssh_proxy.max_connection_duration_in_seconds must be a non-negative integer; 0 means unlimited"
end
if max_connection_duration > 0
config[:max_connection_duration] = "#{max_connection_duration}s"
end

config[:bbs_address] = "https://" + p("diego.ssh_proxy.bbs.api_location")
config[:bbs_client_cert] = "/var/vcap/jobs/ssh_proxy/config/certs/bbs/client.crt"
config[:bbs_client_key] = "/var/vcap/jobs/ssh_proxy/config/certs/bbs/client.key"
Expand Down Expand Up @@ -122,4 +131,3 @@

config.to_json
%>

75 changes: 75 additions & 0 deletions spec/ssh_proxy_template_spec.rb
Original file line number Diff line number Diff line change
@@ -0,0 +1,75 @@
# frozen_string_literal: true

# rubocop: disable Metrics/BlockLength
require 'rspec'
require 'json'
require 'bosh/template/test'

describe 'ssh_proxy' do
let(:release_path) { File.join(File.dirname(__FILE__), '..') }
let(:release) { Bosh::Template::Test::ReleaseDir.new(release_path) }
let(:job) { release.job('ssh_proxy') }
let(:deployment_manifest_fragment) do
{
'diego' => {
'ssh_proxy' => {
'host_key' => 'HOST KEY',
'bbs' => {
'ca_cert' => 'BBS CA CERT',
'client_cert' => 'BBS CLIENT CERT',
'client_key' => 'BBS CLIENT KEY'
}
}
},
'loggregator' => {
'ca_cert' => 'LOGGREGATOR CA CERT',
'cert' => 'LOGGREGATOR CERT',
'key' => 'LOGGREGATOR KEY'
}
}
end

describe 'ssh_proxy.json.erb' do
let(:template) { job.template('config/ssh_proxy.json') }
let(:rendered_config) { JSON.parse(template.render(deployment_manifest_fragment)) }

context 'when max_connection_duration_in_seconds is not configured' do
it 'omits the connection duration to allow unlimited sessions' do
expect(rendered_config).not_to have_key('max_connection_duration')
end
end

context 'when max_connection_duration_in_seconds is configured' do
before do
deployment_manifest_fragment['diego']['ssh_proxy']['max_connection_duration_in_seconds'] = 86_400
end

it 'renders the duration in seconds' do
expect(rendered_config['max_connection_duration']).to eq('86400s')
end
end

context 'when max_connection_duration_in_seconds is zero' do
before do
deployment_manifest_fragment['diego']['ssh_proxy']['max_connection_duration_in_seconds'] = 0
end

it 'omits the connection duration to allow unlimited sessions' do
expect(rendered_config).not_to have_key('max_connection_duration')
end
end

[-1, 1.5, '', '3600', false].each do |value|
context "when max_connection_duration_in_seconds is #{value.inspect}" do
before do
deployment_manifest_fragment['diego']['ssh_proxy']['max_connection_duration_in_seconds'] = value
end

it 'rejects values that are not non-negative integers' do
expect { rendered_config }.to raise_error(/must be a non-negative integer/)
end
end
end
end
end
# rubocop: enable Metrics/BlockLength
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,7 @@ type SSHProxyConfig struct {
LoggregatorConfig loggingclient.Config `json:"loggregator"`
CommunicationTimeout durationjson.Duration `json:"communication_timeout,omitempty"`
IdleConnectionTimeout durationjson.Duration `json:"idle_connection_timeout,omitempty"`
MaxConnectionDuration durationjson.Duration `json:"max_connection_duration,omitempty"`
ConnectToInstanceAddress bool `json:"connect_to_instance_address"`

BackendsTLSEnabled bool `json:"backends_tls_enabled,omitempty"`
Expand All @@ -68,6 +69,9 @@ func NewSSHProxyConfig(configPath string) (SSHProxyConfig, error) {
if err != nil {
return SSHProxyConfig{}, err
}
if proxyConfig.MaxConnectionDuration < 0 {
return SSHProxyConfig{}, errors.New("max_connection_duration must not be negative")
}
Comment thread
rkoster marked this conversation as resolved.

return proxyConfig, nil
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,7 @@ var _ = Describe("SSHProxyConfig", func() {
"debug_address": "5.5.5.5:9090",
"connect_to_instance_address": true,
"idle_connection_timeout": "5ms",
"max_connection_duration": "1h",

"backends_tls_enabled": true,
"backends_tls_ca_certificates": "./some_filepath/ca.crt",
Expand Down Expand Up @@ -106,6 +107,7 @@ var _ = Describe("SSHProxyConfig", func() {
AllowedHostKeyAlgorithms: "hostkeyalg1,hostkeyalg2,hostkeyalg3",
ConnectToInstanceAddress: true,
IdleConnectionTimeout: durationjson.Duration(5 * time.Millisecond),
MaxConnectionDuration: durationjson.Duration(time.Hour),
LagerConfig: lagerflags.LagerConfig{
LogLevel: lagerflags.DEBUG,
},
Expand All @@ -127,6 +129,53 @@ var _ = Describe("SSHProxyConfig", func() {
})
})

Context("when the max connection duration is negative", func() {
BeforeEach(func() {
configData = `{"max_connection_duration": "-1s"}`
})

It("returns an error", func() {
_, err := config.NewSSHProxyConfig(configFilePath)
Expect(err).To(MatchError("max_connection_duration must not be negative"))
})
})

Context("when the max connection duration is omitted", func() {
BeforeEach(func() {
configData = `{}`
})

It("defaults to zero for unlimited sessions", func() {
proxyConfig, err := config.NewSSHProxyConfig(configFilePath)
Expect(err).NotTo(HaveOccurred())
Expect(proxyConfig.MaxConnectionDuration).To(Equal(durationjson.Duration(0)))
})
})

Context("when the max connection duration is zero", func() {
BeforeEach(func() {
configData = `{"max_connection_duration": "0s"}`
})

It("accepts zero for unlimited sessions", func() {
proxyConfig, err := config.NewSSHProxyConfig(configFilePath)
Expect(err).NotTo(HaveOccurred())
Expect(proxyConfig.MaxConnectionDuration).To(Equal(durationjson.Duration(0)))
})
})

Context("when the max connection duration is greater than one hour", func() {
BeforeEach(func() {
configData = `{"max_connection_duration": "24h"}`
})

It("accepts the configured duration", func() {
proxyConfig, err := config.NewSSHProxyConfig(configFilePath)
Expect(err).NotTo(HaveOccurred())
Expect(proxyConfig.MaxConnectionDuration).To(Equal(durationjson.Duration(24 * time.Hour)))
})
})

Context("when the file does not contain valid json", func() {
BeforeEach(func() {
configData = "{{"
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -62,7 +62,7 @@ func main() {
logger.Error("failed-to-get-tls-config", err)
os.Exit(1)
}
sshProxy := proxy.New(logger, proxySSHServerConfig, metronClient, tlsConfig)
sshProxy := proxy.New(logger, proxySSHServerConfig, metronClient, tlsConfig, time.Duration(sshProxyConfig.MaxConnectionDuration))
server := server.NewServer(logger, sshProxyConfig.Address, sshProxy, time.Duration(sshProxyConfig.IdleConnectionTimeout))

healthCheckHandler := healthcheck.NewHandler(logger)
Expand Down
58 changes: 48 additions & 10 deletions src/code.cloudfoundry.org/diego-ssh/proxy/proxy.go
Original file line number Diff line number Diff line change
@@ -1,13 +1,15 @@
package proxy

import (
"context"
"crypto/tls"
"encoding/json"
"errors"
"fmt"
"net"
"strings"
"sync"
"time"
"unicode/utf8"

loggingclient "code.cloudfoundry.org/diego-logging-client"
Expand Down Expand Up @@ -47,21 +49,24 @@ type Proxy struct {
connections int
metronClient loggingclient.IngressClient

tlsConfig *tls.Config
tlsConfig *tls.Config
maxConnectionDuration time.Duration
}

func New(
logger lager.Logger,
serverConfig *ssh.ServerConfig,
metronClient loggingclient.IngressClient,
tlsConfig *tls.Config,
maxConnectionDuration time.Duration,
) *Proxy {
return &Proxy{
logger: logger,
serverConfig: serverConfig,
connectionLock: &sync.Mutex{},
metronClient: metronClient,
tlsConfig: tlsConfig,
logger: logger,
serverConfig: serverConfig,
connectionLock: &sync.Mutex{},
metronClient: metronClient,
tlsConfig: tlsConfig,
maxConnectionDuration: maxConnectionDuration,
}
}

Expand All @@ -75,10 +80,26 @@ func (p *Proxy) HandleConnection(netConn net.Conn) {
}
defer serverConn.Close()

clientConn, clientChannels, clientRequests, err := NewClientConn(logger, serverConn.Permissions, p.tlsConfig)
ctx := context.Background()
if p.maxConnectionDuration > 0 {
Comment thread
rkoster marked this conversation as resolved.
var cancel context.CancelFunc
ctx, cancel = context.WithTimeout(ctx, p.maxConnectionDuration)
defer cancel()
stop := context.AfterFunc(ctx, func() {
logger.Info("maximum-connection-duration-reached", lager.Data{"duration-in-seconds": p.maxConnectionDuration.Seconds()})
_ = serverConn.Close()
})
defer stop()
}

clientConn, clientChannels, clientRequests, err := NewClientConn(ctx, logger, serverConn.Permissions, p.tlsConfig)
if err != nil {
return
}
if ctx.Done() != nil {
stop := context.AfterFunc(ctx, func() { _ = clientConn.Close() })
defer stop()
}

logMessage := extractLogMessage(logger, serverConn.Permissions)

Expand Down Expand Up @@ -320,7 +341,7 @@ func Wait(logger lager.Logger, waiters ...Waiter) {
wg.Wait()
}

func NewClientConn(logger lager.Logger, permissions *ssh.Permissions, tlsConfig *tls.Config) (ssh.Conn, <-chan ssh.NewChannel, <-chan *ssh.Request, error) {
func NewClientConn(ctx context.Context, logger lager.Logger, permissions *ssh.Permissions, tlsConfig *tls.Config) (ssh.Conn, <-chan ssh.NewChannel, <-chan *ssh.Request, error) {
if permissions == nil || permissions.CriticalOptions == nil {
err := errors.New("Invalid permissions from authentication")
logger.Error("permissions-and-critical-options-required", err)
Expand All @@ -344,7 +365,8 @@ func NewClientConn(logger lager.Logger, permissions *ssh.Permissions, tlsConfig
dialer := func() (net.Conn, error) {
tlsConfig := tlsConfigWithServerName(tlsConfig, targetConfig.ServerCertDomainSAN)
if tlsConfig != nil && targetConfig.TLSAddress != "" {
nConn, err := tls.Dial("tcp", targetConfig.TLSAddress, tlsConfig)
tlsDialer := &tls.Dialer{Config: tlsConfig}
nConn, err := tlsDialer.DialContext(ctx, "tcp", targetConfig.TLSAddress)
if err == nil {
return nConn, nil
}
Expand All @@ -353,9 +375,12 @@ func NewClientConn(logger lager.Logger, permissions *ssh.Permissions, tlsConfig
"tcp_address": targetConfig.TLSAddress,
"server_cert_domain_san": targetConfig.ServerCertDomainSAN,
})
if ctx.Err() != nil {
return nil, ctx.Err()
}
}

nConn, err := net.Dial("tcp", targetConfig.Address)
nConn, err := (&net.Dialer{}).DialContext(ctx, "tcp", targetConfig.Address)
if err != nil {
logger.Error("dial-failed", err, lager.Data{
"address": targetConfig.Address,
Expand All @@ -370,6 +395,18 @@ func NewClientConn(logger lager.Logger, permissions *ssh.Permissions, tlsConfig
if err != nil {
return nil, nil, nil, err
}
// SSH handshakes do not accept a context. Closing the socket interrupts a
// stalled handshake when the same deadline used for dialing expires.
if ctx.Done() != nil {
stop := context.AfterFunc(ctx, func() { _ = nConn.Close() })
defer stop()
}
handshakeComplete := false
defer func() {
if !handshakeComplete {
_ = nConn.Close()
}
}()

logger.Info("connected-to-backend", lager.Data{
"backend-address": nConn.RemoteAddr().String(),
Expand Down Expand Up @@ -427,6 +464,7 @@ func NewClientConn(logger lager.Logger, permissions *ssh.Permissions, tlsConfig
return nil, nil, nil, err
}

handshakeComplete = true
return conn, ch, req, nil
}

Expand Down
Loading
Loading