Skip to content
Open
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
7 changes: 7 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -164,6 +164,13 @@ Server:
Hosts:
- localhost:3389
- my-{{ preferred_username }}-host:3389
# Optional hostname-to-IP mappings. The hostname is authorized and remains
# visible to the RDP client, while rdpgw connects to Address. Port defaults
# to 3389. Destinations are automatically included in the host list.
Destinations:
- Hostname: rdp.example.com
Address: 192.0.2.42
Port: 3389
# if true the server randomly selects a host to connect to
# valid options are:
# - roundrobin, which selects a random host from the list (default)
Expand Down
116 changes: 95 additions & 21 deletions cmd/rdpgw/config/configuration.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,9 +3,12 @@ package config
import (
"fmt"
"log"
"net"
"os"
"strconv"
"strings"

"github.com/bolkedebruin/rdpgw/cmd/rdpgw/protocol"
"github.com/bolkedebruin/rdpgw/cmd/rdpgw/security"
"github.com/knadh/koanf/parsers/yaml"
"github.com/knadh/koanf/providers/confmap"
Expand Down Expand Up @@ -80,23 +83,24 @@ type Configuration struct {
}

type ServerConfig struct {
GatewayAddress string `koanf:"gatewayaddress"`
Port int `koanf:"port"`
BindAddress string `koanf:"bindaddress"`
CertFile string `koanf:"certfile"`
KeyFile string `koanf:"keyfile"`
Hosts []string `koanf:"hosts"`
HostSelection string `koanf:"hostselection"`
SessionKey string `koanf:"sessionkey"`
SessionEncryptionKey string `koanf:"sessionencryptionkey"`
SessionStore string `koanf:"sessionstore"`
MaxSessionLength int `koanf:"maxsessionlength"`
SendBuf int `koanf:"sendbuf"`
ReceiveBuf int `koanf:"receivebuf"`
Tls string `koanf:"tls"`
Authentication []string `koanf:"authentication"`
AuthSocket string `koanf:"authsocket"`
BasicAuthTimeout int `koanf:"basicauthtimeout"`
GatewayAddress string `koanf:"gatewayaddress"`
Port int `koanf:"port"`
BindAddress string `koanf:"bindaddress"`
CertFile string `koanf:"certfile"`
KeyFile string `koanf:"keyfile"`
Hosts []string `koanf:"hosts"`
Destinations []DestinationConfig `koanf:"destinations"`
HostSelection string `koanf:"hostselection"`
SessionKey string `koanf:"sessionkey"`
SessionEncryptionKey string `koanf:"sessionencryptionkey"`
SessionStore string `koanf:"sessionstore"`
MaxSessionLength int `koanf:"maxsessionlength"`
SendBuf int `koanf:"sendbuf"`
ReceiveBuf int `koanf:"receivebuf"`
Tls string `koanf:"tls"`
Authentication []string `koanf:"authentication"`
AuthSocket string `koanf:"authsocket"`
BasicAuthTimeout int `koanf:"basicauthtimeout"`
// AllowedDestinationPorts gates the TCP ports `hostselection: any` may
// forward to. Empty defaults to {3389}. Ignored for the curated host
// modes (roundrobin, signed, unsigned).
Expand All @@ -112,6 +116,73 @@ type ServerConfig struct {
TrustedProxies []string `koanf:"trustedproxies"`
}

type DestinationConfig struct {
Hostname string `koanf:"hostname"`
Address string `koanf:"address"`
Port int `koanf:"port"`
}

func (d DestinationConfig) Endpoints() (requested string, dial string, err error) {
hostname := strings.ToLower(strings.TrimSuffix(strings.TrimSpace(d.Hostname), "."))
if hostname == "" || strings.ContainsAny(hostname, " :/[]") {
return "", "", fmt.Errorf("invalid destination hostname %q", d.Hostname)
}
ip := net.ParseIP(strings.TrimSpace(d.Address))
if ip == nil {
return "", "", fmt.Errorf("destination %q has invalid IP address %q", hostname, d.Address)
}
port := d.Port
if port == 0 {
port = 3389
}
if port < 1 || port > 65535 {
return "", "", fmt.Errorf("destination %q has invalid port %d", hostname, d.Port)
}
requested, err = protocol.CanonicalHostPort(net.JoinHostPort(hostname, strconv.Itoa(port)))
if err != nil {
return "", "", err
}
return requested, net.JoinHostPort(ip.String(), strconv.Itoa(port)), nil
}

func (s *ServerConfig) DestinationMappings() (map[string]string, error) {
mappings := make(map[string]string, len(s.Destinations))
for _, destination := range s.Destinations {
requested, dial, err := destination.Endpoints()
if err != nil {
return nil, err
}
if _, exists := mappings[requested]; exists {
return nil, fmt.Errorf("duplicate destination %q", requested)
}
mappings[requested] = dial
}
return mappings, nil
}

func (s *ServerConfig) EffectiveHosts() ([]string, error) {
if _, err := s.DestinationMappings(); err != nil {
return nil, err
}
hosts := append([]string(nil), s.Hosts...)
seen := make(map[string]struct{}, len(hosts)+len(s.Destinations))
for _, host := range hosts {
key := strings.ToLower(host)
if canonical, err := protocol.CanonicalHostPort(host); err == nil {
key = canonical
}
seen[key] = struct{}{}
}
for _, destination := range s.Destinations {
requested, _, _ := destination.Endpoints()
if _, exists := seen[requested]; !exists {
hosts = append(hosts, requested)
seen[requested] = struct{}{}
}
}
return hosts, nil
}

type KerberosConfig struct {
Keytab string `koanf:"keytab"`
Krb5Conf string `koanf:"krb5conf"`
Expand All @@ -124,13 +195,13 @@ type OpenIDConfig struct {
}

type HeaderConfig struct {
UserHeader string `koanf:"userheader"`
UserIdHeader string `koanf:"useridheader"`
EmailHeader string `koanf:"emailheader"`
UserHeader string `koanf:"userheader"`
UserIdHeader string `koanf:"useridheader"`
EmailHeader string `koanf:"emailheader"`
DisplayNameHeader string `koanf:"displaynameheader"`
// TrustedProxies is the CIDR allow-list of upstream proxies allowed to
// stamp UserHeader (and friends). Empty disables header auth at runtime.
TrustedProxies []string `koanf:"trustedproxies"`
TrustedProxies []string `koanf:"trustedproxies"`
}

type RDGCapsConfig struct {
Expand Down Expand Up @@ -260,6 +331,9 @@ func Load(configFile string) Configuration {
k.UnmarshalWithConf("Security", &Conf.Security, koanfTag)
k.UnmarshalWithConf("Client", &Conf.Client, koanfTag)
k.UnmarshalWithConf("Kerberos", &Conf.Kerberos, koanfTag)
if _, err := Conf.Server.DestinationMappings(); err != nil {
log.Fatalf("invalid server destinations: %s", err)
}

if err := checkDefaultSecrets(&Conf); err != nil {
log.Fatalf("refusing to start: %s", err)
Expand Down
72 changes: 71 additions & 1 deletion cmd/rdpgw/config/configuration_test.go
Original file line number Diff line number Diff line change
@@ -1,9 +1,79 @@
package config

import (
"os"
"path/filepath"
"reflect"
"testing"
)

func TestLoadDestinationsFromYAML(t *testing.T) {
originalConf := Conf
Conf = Configuration{}
t.Cleanup(func() { Conf = originalConf })

configFile := filepath.Join(t.TempDir(), "rdpgw.yaml")
data := []byte(`Server:
Destinations:
- Hostname: rdp.example.com
Address: 192.0.2.42
- Hostname: alternate.example.com
Address: 2001:db8::42
Port: 3390
`)
if err := os.WriteFile(configFile, data, 0600); err != nil {
t.Fatalf("writing config file: %v", err)
}

configuration := Load(configFile)
mappings, err := configuration.Server.DestinationMappings()
if err != nil {
t.Fatalf("DestinationMappings returned an error: %v", err)
}
want := map[string]string{
"rdp.example.com:3389": "192.0.2.42:3389",
"alternate.example.com:3390": "[2001:db8::42]:3390",
}
if !reflect.DeepEqual(mappings, want) {
t.Fatalf("DestinationMappings() = %v, want %v", mappings, want)
}
}

func TestDestinationMappingsRejectInvalidConfiguration(t *testing.T) {
tests := []ServerConfig{
{Destinations: []DestinationConfig{{Address: "192.0.2.42"}}},
{Destinations: []DestinationConfig{{Hostname: "rdp.example.com", Address: "not-an-ip"}}},
{Destinations: []DestinationConfig{{Hostname: "rdp.example.com", Address: "192.0.2.42", Port: 70000}}},
{Destinations: []DestinationConfig{
{Hostname: "rdp.example.com", Address: "192.0.2.42"},
{Hostname: "RDP.EXAMPLE.COM.", Address: "192.0.2.43", Port: 3389},
}},
}
for _, server := range tests {
if _, err := server.DestinationMappings(); err == nil {
t.Fatalf("DestinationMappings accepted invalid configuration: %+v", server.Destinations)
}
}
}

func TestEffectiveHostsIncludesCanonicalDestinations(t *testing.T) {
server := ServerConfig{
Hosts: []string{"legacy.example.com:3389", "RDP.EXAMPLE.COM.:3389"},
Destinations: []DestinationConfig{
{Hostname: "rdp.example.com", Address: "192.0.2.42"},
{Hostname: "second.example.com", Address: "192.0.2.43"},
},
}
hosts, err := server.EffectiveHosts()
if err != nil {
t.Fatalf("EffectiveHosts returned an error: %v", err)
}
want := []string{"legacy.example.com:3389", "RDP.EXAMPLE.COM.:3389", "second.example.com:3389"}
if !reflect.DeepEqual(hosts, want) {
t.Fatalf("EffectiveHosts() = %v, want %v", hosts, want)
}
}

func TestHeaderEnabled(t *testing.T) {
cases := []struct {
name string
Expand Down Expand Up @@ -165,4 +235,4 @@ func TestHeaderConfigValidation(t *testing.T) {
}
})
}
}
}
23 changes: 16 additions & 7 deletions cmd/rdpgw/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -78,6 +78,14 @@ func main() {
panic(err)
}
conf = config.Load(opts.ConfigFile)
hosts, err := conf.Server.EffectiveHosts()
if err != nil {
log.Fatalf("Invalid server destinations: %s", err)
}
destinationMappings, err := conf.Server.DestinationMappings()
if err != nil {
log.Fatalf("Invalid server destinations: %s", err)
}

// set callback url and external advertised gateway address
url, err := url.Parse(conf.Server.GatewayAddress)
Expand All @@ -97,7 +105,7 @@ func main() {
security.UserSigningKey = []byte(conf.Security.UserTokenSigningKey)
security.QuerySigningKey = []byte(conf.Security.QueryTokenSigningKey)
security.HostSelection = conf.Server.HostSelection
security.Hosts = conf.Server.Hosts
security.Hosts = hosts

// init session store
web.InitStore([]byte(conf.Server.SessionKey),
Expand All @@ -113,7 +121,7 @@ func main() {
QueryInfo: security.QueryInfo,
QueryTokenIssuer: conf.Security.QueryTokenIssuer,
EnableUserToken: conf.Security.EnableUserToken,
Hosts: conf.Server.Hosts,
Hosts: hosts,
HostSelection: conf.Server.HostSelection,
RdpOpts: web.RdpOpts{
UsernameTemplate: conf.Client.UsernameTemplate,
Expand Down Expand Up @@ -198,11 +206,12 @@ func main() {
DisableAll: conf.Caps.DisableRedirect,
EnableAll: conf.Caps.RedirectAll,
},
IdleTimeout: conf.Caps.IdleTimeout,
SmartCardAuth: conf.Caps.SmartCardAuth,
TokenAuth: conf.Caps.TokenAuth,
ReceiveBuf: conf.Server.ReceiveBuf,
SendBuf: conf.Server.SendBuf,
IdleTimeout: conf.Caps.IdleTimeout,
SmartCardAuth: conf.Caps.SmartCardAuth,
TokenAuth: conf.Caps.TokenAuth,
DestinationMappings: destinationMappings,
ReceiveBuf: conf.Server.ReceiveBuf,
SendBuf: conf.Server.SendBuf,
}

if conf.Caps.TokenAuth {
Expand Down
36 changes: 36 additions & 0 deletions cmd/rdpgw/protocol/gateway.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package protocol
import (
"context"
"errors"
"fmt"
"log"
"net"
"net/http"
Expand All @@ -28,6 +29,18 @@ type CheckPAACookieFunc func(context.Context, string) (bool, error)
type CheckClientNameFunc func(context.Context, string) (bool, error)
type CheckHostFunc func(context.Context, string) (bool, error)

func CanonicalHostPort(hostport string) (string, error) {
host, port, err := net.SplitHostPort(hostport)
if err != nil {
return "", fmt.Errorf("invalid destination %q: %w", hostport, err)
}
host = strings.ToLower(strings.TrimSuffix(host, "."))
if host == "" {
return "", fmt.Errorf("invalid destination %q: hostname is empty", hostport)
}
return net.JoinHostPort(host, port), nil
}

type Gateway struct {
// CheckPAACookie verifies if the PAA cookie sent by the client is valid
CheckPAACookie CheckPAACookieFunc
Expand All @@ -38,6 +51,10 @@ type Gateway struct {
// CheckHost verifies if the client is allowed to connect to the remote host
CheckHost CheckHostFunc

// DestinationMappings maps an authorized hostname and port to the IP address
// and port used for the upstream TCP connection.
DestinationMappings map[string]string

// RedirectFlags sets what devices the client is allowed to redirect to the remote host
RedirectFlags RedirectFlags

Expand All @@ -52,6 +69,25 @@ type Gateway struct {

ReceiveBuf int
SendBuf int

dialTimeout func(network, address string, timeout time.Duration) (net.Conn, error)
}

func (g *Gateway) openHost(host string) (net.Conn, string, error) {
canonical, err := CanonicalHostPort(host)
if err != nil {
return nil, host, err
}
destination := host
if mapped, ok := g.DestinationMappings[canonical]; ok {
destination = mapped
}
dial := g.dialTimeout
if dial == nil {
dial = net.DialTimeout
}
conn, err := dial("tcp", destination, 15*time.Second)
return conn, destination, err
}

var upgrader = websocket.Upgrader{}
Expand Down
Loading
Loading