Skip to content

Commit e75605b

Browse files
authored
refactor: move domain check into small helper util (#1000)
1 parent 79bcccb commit e75605b

8 files changed

Lines changed: 504 additions & 33 deletions

File tree

‎Dockerfile‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -39,6 +39,7 @@ RUN go mod download
3939

4040
COPY ./cmd ./cmd
4141
COPY ./internal ./internal
42+
COPY ./pkg ./pkg
4243
COPY --from=frontend-builder /frontend/dist ./internal/assets/dist
4344

4445
RUN CGO_ENABLED=0 go build -tags "${BUILD_TAGS}" -ldflags "${LDFLAGS} \

‎Dockerfile.dev‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@ RUN go install github.com/go-delve/delve/cmd/dlv@v1.26.3
1212

1313
COPY ./cmd ./cmd
1414
COPY ./internal ./internal
15+
COPY ./pkg ./pkg
1516
COPY ./air.toml ./
1617

1718
EXPOSE 3000

‎Dockerfile.distroless‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -39,6 +39,7 @@ RUN go mod download
3939

4040
COPY ./cmd ./cmd/
4141
COPY ./internal ./internal
42+
COPY ./pkg ./pkg
4243
COPY --from=frontend-builder /frontend/dist ./internal/assets/dist
4344

4445
RUN CGO_ENABLED=0 go build -tags "${BUILD_TAGS}" -ldflags "${LDFLAGS} \

‎internal/controller/oauth_controller.go‎

Lines changed: 23 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,9 @@
11
package controller
22

33
import (
4+
"errors"
45
"fmt"
56
"net/http"
6-
"net/url"
77
"strings"
88
"time"
99

@@ -12,6 +12,7 @@ import (
1212
"github.com/tinyauthapp/tinyauth/internal/service"
1313
"github.com/tinyauthapp/tinyauth/internal/utils"
1414
"github.com/tinyauthapp/tinyauth/internal/utils/logger"
15+
"github.com/tinyauthapp/tinyauth/pkg/validators"
1516
"go.uber.org/dig"
1617

1718
"github.com/gin-gonic/gin"
@@ -311,54 +312,46 @@ func (controller *OAuthController) getCookieDomain() string {
311312
}
312313

313314
func (controller *OAuthController) isRedirectSafe(redirectURI string) bool {
314-
u, err := url.Parse(redirectURI)
315+
v := validators.NewDomainValidator(validators.DomainValidatorOptions{
316+
WithScheme: true,
317+
WithPort: true,
318+
})
319+
320+
_, err := v.SafeHostname(controller.runtime.AppURL)
315321

316322
if err != nil {
317-
controller.log.App.Error().Err(err).Msg("Failed to parse redirect URI")
323+
controller.log.App.Error().Err(err).Msg("App URL is invalid, cannot validate redirect URI")
318324
return false
319325
}
320326

321-
if u.Scheme == "" || u.Host == "" {
322-
controller.log.App.Warn().Msg("Redirect URI has invalid scheme or host")
323-
return false
327+
err = v.Validate(redirectURI, controller.runtime.AppURL)
328+
329+
if err == nil {
330+
return true
324331
}
325332

326-
au, err := url.Parse(controller.runtime.AppURL)
333+
controller.log.App.Debug().Err(err).Msg("Failed to validate redirect URI")
327334

328-
if err != nil {
329-
controller.log.App.Error().Err(err).Msg("Failed to parse app URL")
335+
if errors.Is(err, validators.ErrInvalidURL) ||
336+
errors.Is(err, validators.ErrSchemeMismatch) ||
337+
errors.Is(err, validators.ErrPortMismatch) {
330338
return false
331339
}
332340

333-
if u.Scheme != au.Scheme {
334-
controller.log.App.Warn().Msg("Redirect URI scheme does not match app URL scheme")
341+
if !controller.config.Auth.SubdomainsEnabled {
335342
return false
336343
}
337344

338-
getEffectivePort := func(u *url.URL) string {
339-
if u.Port() != "" {
340-
return u.Port()
341-
}
342-
if u.Scheme == "https" {
343-
return "443"
344-
}
345-
return "80"
346-
}
347-
348-
if getEffectivePort(u) != getEffectivePort(au) {
349-
controller.log.App.Warn().Msg("Redirect URI port does not match app URL port")
350-
return false
351-
}
345+
v = validators.NewDomainValidator(validators.DomainValidatorOptions{})
352346

353-
if strings.EqualFold(u.Hostname(), au.Hostname()) {
354-
return true
355-
}
347+
hostname, err := v.SafeHostname(redirectURI)
356348

357-
if !controller.config.Auth.SubdomainsEnabled {
349+
if err != nil {
350+
controller.log.App.Error().Err(err).Msg("Failed to get safe hostname from redirect URI")
358351
return false
359352
}
360353

361-
if strings.HasSuffix(strings.ToLower(u.Hostname()), "."+strings.ToLower(controller.runtime.CookieDomain)) {
354+
if strings.HasSuffix(hostname, "."+strings.ToLower(controller.runtime.CookieDomain)) {
362355
return true
363356
}
364357

‎internal/controller/oauth_controller_test.go‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,7 @@ import (
99
"github.com/tinyauthapp/tinyauth/internal/utils/logger"
1010
)
1111

12-
func TestOAuthControllerIsRedirectSafe(t *testing.T) {
12+
func TestOAuthController_isRedirectSafe(t *testing.T) {
1313
log := logger.NewLogger().WithTestConfig()
1414
log.Init()
1515

‎internal/service/access_controls_service.go‎

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,12 @@
11
package service
22

33
import (
4+
"errors"
45
"strings"
56

67
"github.com/tinyauthapp/tinyauth/internal/model"
78
"github.com/tinyauthapp/tinyauth/internal/utils/logger"
9+
"github.com/tinyauthapp/tinyauth/pkg/validators"
810
"go.uber.org/dig"
911
)
1012

@@ -38,13 +40,19 @@ func NewAccessControlsService(i AccessControlServiceInput) *AccessControlsServic
3840
func (service *AccessControlsService) lookupStaticACLs(domain string) *model.App {
3941
var nameMatch *model.App
4042

43+
v := validators.NewDomainValidator(validators.DomainValidatorOptions{})
44+
4145
// First try to find a matching app by domain, then fallback to matching by app name (subdomain)
4246
for app, config := range service.config.Apps {
43-
if config.Config.Domain == domain {
47+
err := v.Validate(config.Config.Domain, domain)
48+
if err == nil {
4449
service.log.App.Debug().Str("name", app).Msg("Found matching container by domain")
4550
return &config
4651
}
47-
if strings.SplitN(domain, ".", 2)[0] == app {
52+
if !errors.Is(err, validators.ErrHostnameMismatch) {
53+
service.log.App.Debug().Str("name", app).Err(err).Msg("Domain validation failed")
54+
}
55+
if strings.HasPrefix(strings.ToLower(domain), strings.ToLower(app+".")) {
4856
service.log.App.Debug().Str("name", app).Msg("Found matching container by app name")
4957
nameMatch = &config
5058
}

‎pkg/validators/domain_validator.go‎

Lines changed: 179 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,179 @@
1+
// Package validators provides validators for various types of data.
2+
//
3+
// Domain validator is a simple utility that ensures two domains are exact
4+
// matches while ensuring that techniques used to bypass such checks do
5+
// not impact the validation.
6+
7+
package validators
8+
9+
import (
10+
"fmt"
11+
"net"
12+
"net/url"
13+
"slices"
14+
"strings"
15+
16+
"golang.org/x/net/idna"
17+
)
18+
19+
var (
20+
ErrInvalidURL = fmt.Errorf("invalid url")
21+
ErrSchemeMismatch = fmt.Errorf("scheme mismatch")
22+
ErrPortMismatch = fmt.Errorf("port mismatch")
23+
ErrHostnameMismatch = fmt.Errorf("hostname mismatch")
24+
)
25+
26+
// DomainValidatorOptions is a set of options for DomainValidator.
27+
type DomainValidatorOptions struct {
28+
// Ensure domains have the same scheme.
29+
WithScheme bool
30+
// Ensure domains have the same port.
31+
WithPort bool
32+
// Specify a list of allowed schemes IF WithScheme is set to true.
33+
// Leave empty to allow any scheme.
34+
AllowedSchemes []string
35+
}
36+
37+
// DomainValidator is a simple utility that ensures two domains are exact
38+
// matches while ensuring that techniques used to bypass such checks do
39+
// not impact the validation.
40+
type DomainValidator struct {
41+
opts DomainValidatorOptions
42+
}
43+
44+
// NewDomainValidator creates a new DomainValidator.
45+
func NewDomainValidator(opts DomainValidatorOptions) *DomainValidator {
46+
return &DomainValidator{
47+
opts: opts,
48+
}
49+
}
50+
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)
56+
}
57+
58+
if err != nil {
59+
return nil, fmt.Errorf("failed to parse input url: %w", err)
60+
}
61+
62+
if u.Host == "" {
63+
return nil, ErrInvalidURL
64+
}
65+
66+
if v.opts.WithPort && !v.opts.WithScheme && u.Port() == "" {
67+
return nil, fmt.Errorf("port validation is enabled but port is missing in input url and schemes are not enabled")
68+
}
69+
70+
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+
}
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)
77+
}
78+
}
79+
80+
return u, nil
81+
}
82+
83+
func (v *DomainValidator) getEffectivePort(u *url.URL) (string, bool) {
84+
if u.Port() != "" {
85+
return u.Port(), true
86+
}
87+
switch u.Scheme {
88+
case "http":
89+
return "80", true
90+
case "https":
91+
return "443", true
92+
default:
93+
return "", false
94+
}
95+
}
96+
97+
func (v *DomainValidator) formatHostname(hostname string) (string, error) {
98+
hostname = strings.ToLower(hostname)
99+
hostname = strings.TrimSuffix(hostname, ".")
100+
if net.ParseIP(hostname) != nil {
101+
return "", fmt.Errorf("ip addresses are not supported")
102+
}
103+
hostname, err := idna.Lookup.ToASCII(hostname)
104+
if err != nil {
105+
return "", fmt.Errorf("failed to convert hostname to ascii: %w", err)
106+
}
107+
return hostname, nil
108+
}
109+
110+
// Validate ensures that two domains are exact matches with the
111+
// options defined in the DomainValidatorOptions. It ensures that the
112+
// inputs are proper URLs and contain a host. It lowercases the hostnames
113+
// and removes the trailing dot. Finally, it checks that the hostnames are
114+
// equal unless WithScheme or WithPort is set to true where it also
115+
// validates the scheme and port respectively.
116+
func (v *DomainValidator) Validate(expected, actual string) error {
117+
eu, err := v.getURL(expected)
118+
119+
if err != nil {
120+
return err
121+
}
122+
123+
au, err := v.getURL(actual)
124+
125+
if err != nil {
126+
return err
127+
}
128+
129+
if v.opts.WithScheme {
130+
if eu.Scheme != au.Scheme {
131+
return ErrSchemeMismatch
132+
}
133+
}
134+
135+
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 {
145+
return ErrPortMismatch
146+
}
147+
}
148+
149+
euf, err := v.formatHostname(eu.Hostname())
150+
151+
if err != nil {
152+
return err
153+
}
154+
155+
auf, err := v.formatHostname(au.Hostname())
156+
157+
if err != nil {
158+
return err
159+
}
160+
161+
if euf != auf {
162+
return ErrHostnameMismatch
163+
}
164+
165+
return nil
166+
}
167+
168+
// SafeHostname uses the internal validation for domains that Validator uses
169+
// to parse a hostname. It ensures the input URL is a valid URL, that a host
170+
// is present and that the hostname is lowercased and without a trailing dot.
171+
func (v *DomainValidator) SafeHostname(input string) (string, error) {
172+
u, err := v.getURL(input)
173+
174+
if err != nil {
175+
return "", err
176+
}
177+
178+
return v.formatHostname(u.Hostname())
179+
}

0 commit comments

Comments
 (0)