From 9d494429c12ba3a299d53dde9147e0d94b88f5da Mon Sep 17 00:00:00 2001 From: soundgoof Date: Fri, 7 Aug 2026 09:03:12 +0200 Subject: [PATCH] Add configured RDP destination mappings --- README.md | 7 ++ cmd/rdpgw/config/configuration.go | 116 ++++++++++++++++++++----- cmd/rdpgw/config/configuration_test.go | 72 ++++++++++++++- cmd/rdpgw/main.go | 23 +++-- cmd/rdpgw/protocol/gateway.go | 36 ++++++++ cmd/rdpgw/protocol/gateway_test.go | 38 ++++++++ cmd/rdpgw/protocol/process.go | 8 +- 7 files changed, 267 insertions(+), 33 deletions(-) diff --git a/README.md b/README.md index 94d334d..2da9c48 100644 --- a/README.md +++ b/README.md @@ -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) diff --git a/cmd/rdpgw/config/configuration.go b/cmd/rdpgw/config/configuration.go index 493ec2d..d226b8d 100644 --- a/cmd/rdpgw/config/configuration.go +++ b/cmd/rdpgw/config/configuration.go @@ -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" @@ -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). @@ -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"` @@ -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 { @@ -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) diff --git a/cmd/rdpgw/config/configuration_test.go b/cmd/rdpgw/config/configuration_test.go index 05d93b0..3311f03 100644 --- a/cmd/rdpgw/config/configuration_test.go +++ b/cmd/rdpgw/config/configuration_test.go @@ -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 @@ -165,4 +235,4 @@ func TestHeaderConfigValidation(t *testing.T) { } }) } -} \ No newline at end of file +} diff --git a/cmd/rdpgw/main.go b/cmd/rdpgw/main.go index 8729c89..3191e0c 100644 --- a/cmd/rdpgw/main.go +++ b/cmd/rdpgw/main.go @@ -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) @@ -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), @@ -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, @@ -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 { diff --git a/cmd/rdpgw/protocol/gateway.go b/cmd/rdpgw/protocol/gateway.go index e38b18e..6fcff7a 100644 --- a/cmd/rdpgw/protocol/gateway.go +++ b/cmd/rdpgw/protocol/gateway.go @@ -3,6 +3,7 @@ package protocol import ( "context" "errors" + "fmt" "log" "net" "net/http" @@ -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 @@ -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 @@ -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{} diff --git a/cmd/rdpgw/protocol/gateway_test.go b/cmd/rdpgw/protocol/gateway_test.go index 779966f..ffcf4a8 100644 --- a/cmd/rdpgw/protocol/gateway_test.go +++ b/cmd/rdpgw/protocol/gateway_test.go @@ -2,6 +2,7 @@ package protocol import ( "bufio" + "errors" "net" "net/http" "net/http/httptest" @@ -13,6 +14,43 @@ import ( "github.com/patrickmn/go-cache" ) +func TestOpenHostUsesDestinationMapping(t *testing.T) { + var dialed string + gw := Gateway{ + DestinationMappings: map[string]string{ + "rdp.example.com:3389": "192.0.2.42:3389", + }, + dialTimeout: func(network, address string, timeout time.Duration) (net.Conn, error) { + dialed = address + return nil, errors.New("stop after recording destination") + }, + } + + _, destination, err := gw.openHost("RDP.EXAMPLE.COM.:3389") + if err == nil { + t.Fatal("openHost unexpectedly succeeded") + } + if destination != "192.0.2.42:3389" || dialed != destination { + t.Fatalf("openHost dialed %q and returned %q", dialed, destination) + } +} + +func TestOpenHostLeavesUnmappedDestinationUnchanged(t *testing.T) { + var dialed string + gw := Gateway{ + dialTimeout: func(network, address string, timeout time.Duration) (net.Conn, error) { + dialed = address + return nil, errors.New("stop after recording destination") + }, + } + + const host = "legacy.example.com:3389" + _, destination, _ := gw.openHost(host) + if destination != host || dialed != host { + t.Fatalf("openHost dialed %q and returned %q, want %q", dialed, destination, host) + } +} + func TestHeaderHasToken(t *testing.T) { cases := []struct { name string diff --git a/cmd/rdpgw/protocol/process.go b/cmd/rdpgw/protocol/process.go index ceaf16e..f6b3121 100644 --- a/cmd/rdpgw/protocol/process.go +++ b/cmd/rdpgw/protocol/process.go @@ -10,7 +10,6 @@ import ( "log" "net" "strconv" - "time" "github.com/bolkedebruin/rdpgw/cmd/rdpgw/identity" ) @@ -138,14 +137,15 @@ func (p *Processor) Process(ctx context.Context) error { return fmt.Errorf("%x: denied by security policy", E_PROXY_RAP_ACCESSDENIED) } } - log.Printf("Establishing connection to RDP server: %s", host) - p.tunnel.rwc, err = net.DialTimeout("tcp", host, time.Second*15) + rwc, destination, err := p.gw.openHost(host) + log.Printf("Establishing connection to RDP server: %s via %s", host, destination) if err != nil { - log.Printf("Error connecting to %s, %s", host, err) + log.Printf("Error connecting to %s via %s, %s", host, destination, err) msg := p.channelResponse(E_PROXY_INTERNALERROR) p.tunnel.Write(msg) return err } + p.tunnel.rwc = rwc p.tunnel.TargetServer = host log.Printf("Connection established") msg := p.channelResponse(ERROR_SUCCESS)