Skip to content

Commit f43d690

Browse files
authored
refactor: rework scheme validation in oauth controller and frontend (#1026)
1 parent 50c25e4 commit f43d690

7 files changed

Lines changed: 107 additions & 158 deletions

File tree

‎frontend/src/lib/hooks/redirect-uri.ts‎

Lines changed: 1 addition & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -75,19 +75,6 @@ export const useRedirectUri = (
7575
};
7676
};
7777

78-
// ported from internal/controller/oauth_controller.go
79-
const getEffectivePort = (url: URL): string => {
80-
if (url.port) {
81-
return url.port;
82-
}
83-
84-
if (url.protocol == "https:") {
85-
return "443";
86-
}
87-
88-
return "80";
89-
};
90-
9178
// https://www.geeksforgeeks.org/javascript/how-to-check-if-a-string-is-a-valid-ip-address-format-in-javascript
9279
const isIP = (str: string): boolean => {
9380
const ipv4 =
@@ -114,7 +101,7 @@ export const isTrustedDomain = (
114101
return false;
115102
}
116103

117-
if (getEffectivePort(url) != getEffectivePort(appUrl)) {
104+
if (url.port != appUrl.port) {
118105
return false;
119106
}
120107

‎internal/controller/oauth_controller.go‎

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -331,8 +331,7 @@ func (controller *OAuthController) isRedirectSafe(redirectURI string) bool {
331331

332332
controller.log.App.Debug().Err(err).Msg("Failed to validate redirect URI")
333333

334-
if errors.Is(err, validators.ErrInvalidURL) ||
335-
errors.Is(err, validators.ErrPortMismatch) {
334+
if !errors.Is(err, validators.ErrHostnameMismatch) {
336335
return false
337336
}
338337

‎internal/controller/oauth_controller_test.go‎

Lines changed: 0 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -81,14 +81,6 @@ func TestOAuthController_isRedirectSafe(t *testing.T) {
8181
redirectURI: "https://sub.example.com",
8282
expected: false,
8383
},
84-
{
85-
description: "Different scheme returns false",
86-
appURL: "https://tinyauth.example.com",
87-
cookieDomain: "example.com",
88-
subdomainsEnabled: true,
89-
redirectURI: "http://tinyauth.example.com",
90-
expected: false,
91-
},
9284
{
9385
description: "Different port returns false",
9486
appURL: "https://tinyauth.example.com",

‎internal/service/access_controls_service.go‎

Lines changed: 1 addition & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -50,9 +50,7 @@ func (service *AccessControlsService) lookupStaticACLs(domain string) *model.App
5050
service.log.App.Debug().Str("name", app).Msg("Found matching container by domain")
5151
return &config
5252
}
53-
if !errors.Is(err, validators.ErrHostnameMismatch) &&
54-
!errors.Is(err, validators.ErrPortMismatch) &&
55-
!errors.Is(err, validators.ErrSchemeMismatch) {
53+
if !errors.Is(err, validators.ErrHostnameMismatch) {
5654
service.log.App.Debug().Str("name", app).Err(err).Msg("Domain validation failed")
5755
}
5856
}

‎pkg/README.md‎

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,10 @@
1+
# Public packages
2+
3+
This directory contains packages that can be used by
4+
other projects.
5+
6+
While we try to maintain a consistent API, no promises
7+
can be made for non-breaking changes throughout updates
8+
as we constantly need to make changes to comply with the
9+
needs of Tinyauth. We advise pinning the version of the
10+
package you wish to use.

‎pkg/validators/domain_validator.go‎

Lines changed: 57 additions & 46 deletions
Original file line numberDiff line numberDiff line change
@@ -10,14 +10,13 @@ import (
1010
"fmt"
1111
"net"
1212
"net/url"
13-
"slices"
1413
"strings"
1514

1615
"golang.org/x/net/idna"
1716
)
1817

18+
// Errors
1919
var (
20-
ErrInvalidURL = fmt.Errorf("invalid url")
2120
ErrSchemeMismatch = fmt.Errorf("scheme mismatch")
2221
ErrPortMismatch = fmt.Errorf("port mismatch")
2322
ErrHostnameMismatch = fmt.Errorf("hostname mismatch")
@@ -29,8 +28,7 @@ type DomainValidatorOptions struct {
2928
WithScheme bool
3029
// Ensure domains have the same port.
3130
WithPort bool
32-
// Specify a list of allowed schemes IF WithScheme is set to true.
33-
// Leave empty to allow any scheme.
31+
// Specify a list of allowed schemes if WithScheme is set to true.
3432
AllowedSchemes []string
3533
}
3634

@@ -48,53 +46,74 @@ func NewDomainValidator(opts DomainValidatorOptions) *DomainValidator {
4846
}
4947
}
5048

51-
func (v *DomainValidator) getURL(i string) (*url.URL, error) {
52-
u, err := url.Parse(i)
53-
54-
if !v.opts.WithScheme && (err != nil || u.Host == "") {
55-
u, err = url.Parse("tinyauth://" + i)
49+
func (v *DomainValidator) checkScheme(rawURL string) error {
50+
if !v.opts.WithScheme {
51+
return nil
5652
}
5753

58-
if err != nil {
59-
return nil, fmt.Errorf("failed to parse input url: %w", err)
54+
if len(v.opts.AllowedSchemes) == 0 {
55+
return fmt.Errorf("allowed schemes must be specified")
6056
}
6157

62-
if u.Host == "" {
63-
return nil, ErrInvalidURL
58+
for _, scheme := range v.opts.AllowedSchemes {
59+
if strings.HasPrefix(strings.ToLower(rawURL), strings.ToLower(scheme)+"://") {
60+
return nil
61+
}
6462
}
6563

66-
if v.opts.WithPort && u.Port() == "" && (u.Scheme != "http" && u.Scheme != "https") {
67-
return nil, fmt.Errorf("port validation is enabled but port is missing in input url and schemes are not enabled")
64+
return fmt.Errorf("invalid scheme")
65+
66+
}
67+
68+
func (v *DomainValidator) getURL(i string) (*url.URL, error) {
69+
if i == "" {
70+
return nil, fmt.Errorf("url cannot be empty")
6871
}
6972

7073
if v.opts.WithScheme {
71-
// Empty scheme means that we parsed the url with the tinyauth:// placeholder
72-
if u.Scheme == "tinyauth" {
73-
return nil, fmt.Errorf("input url is missing scheme")
74+
err := v.checkScheme(i)
75+
76+
if err != nil {
77+
return nil, fmt.Errorf("invalid scheme: %w", err)
7478
}
75-
if len(v.opts.AllowedSchemes) > 0 && !slices.Contains(v.opts.AllowedSchemes, u.Scheme) {
76-
return nil, fmt.Errorf("scheme %s not allowed", u.Scheme)
79+
80+
u, err := url.Parse(i)
81+
82+
if err != nil {
83+
return nil, fmt.Errorf("failed to parse input url: %w", err)
84+
}
85+
86+
if u.Host == "" || u.Scheme == "" {
87+
return nil, fmt.Errorf("missing host or scheme in url: %s", i)
7788
}
89+
90+
return u, nil
7891
}
7992

80-
return u, nil
81-
}
93+
rawURL := i
8294

83-
func (v *DomainValidator) getEffectivePort(u *url.URL) (string, bool) {
84-
if u.Port() != "" {
85-
return u.Port(), true
95+
if !strings.Contains(i, "://") {
96+
// From godoc: [scheme:][//[userinfo@]host][/]path[?query][#fragment]
97+
// So, we can omit the colon and tell the Go URL lib that we want
98+
// to parse the URL without the scheme. If we don't do this,
99+
// the URL lib will parse our entire domain as the path.
100+
rawURL = "//" + i
86101
}
87-
switch u.Scheme {
88-
case "http":
89-
return "80", true
90-
case "https":
91-
return "443", true
92-
default:
93-
return "", false
102+
103+
u, err := url.Parse(rawURL)
104+
105+
if err != nil {
106+
return nil, fmt.Errorf("failed to parse host: %w", err)
94107
}
108+
109+
if u.Host == "" {
110+
return nil, fmt.Errorf("missing host in url: %s", i)
111+
}
112+
113+
return u, nil
95114
}
96115

97-
func (v *DomainValidator) formatHostname(hostname string) (string, error) {
116+
func (v *DomainValidator) getHostname(hostname string) (string, error) {
98117
hostname = strings.ToLower(hostname)
99118
hostname = strings.TrimSuffix(hostname, ".")
100119
if net.ParseIP(hostname) != nil {
@@ -133,26 +152,18 @@ func (v *DomainValidator) Validate(expected, actual string) error {
133152
}
134153

135154
if v.opts.WithPort {
136-
eup, ok := v.getEffectivePort(eu)
137-
if !ok {
138-
return fmt.Errorf("failed to get effective port for url: %s", eu.String())
139-
}
140-
aup, ok := v.getEffectivePort(au)
141-
if !ok {
142-
return fmt.Errorf("failed to get effective port for url: %s", au.String())
143-
}
144-
if eup != aup {
155+
if eu.Port() != au.Port() {
145156
return ErrPortMismatch
146157
}
147158
}
148159

149-
euf, err := v.formatHostname(eu.Hostname())
160+
euf, err := v.getHostname(eu.Hostname())
150161

151162
if err != nil {
152163
return err
153164
}
154165

155-
auf, err := v.formatHostname(au.Hostname())
166+
auf, err := v.getHostname(au.Hostname())
156167

157168
if err != nil {
158169
return err
@@ -165,7 +176,7 @@ func (v *DomainValidator) Validate(expected, actual string) error {
165176
return nil
166177
}
167178

168-
// SafeHostname uses the internal validation for domains that Validator uses
179+
// SafeHostname uses the internal validation for domains that the validator uses
169180
// to parse a hostname. It ensures the input URL is a valid URL, that a host
170181
// is present and that the hostname is lowercased and without a trailing dot.
171182
func (v *DomainValidator) SafeHostname(input string) (string, error) {
@@ -175,5 +186,5 @@ func (v *DomainValidator) SafeHostname(input string) (string, error) {
175186
return "", err
176187
}
177188

178-
return v.formatHostname(u.Hostname())
189+
return v.getHostname(u.Hostname())
179190
}

0 commit comments

Comments
 (0)