fix: auth module selection (#1089)

This commit is contained in:
Stavros
2026-08-25 17:07:56 +03:00
committed by GitHub
parent be48d712ee
commit 847d8325c7
11 changed files with 250 additions and 99 deletions
+40 -8
View File
@@ -2,11 +2,13 @@ package service
import (
"errors"
"fmt"
"net"
"strings"
"unicode"
"github.com/tinyauthapp/tinyauth/internal/model"
"github.com/tinyauthapp/tinyauth/internal/utils/logger"
"github.com/tinyauthapp/tinyauth/pkg/validators"
"go.uber.org/dig"
)
@@ -17,6 +19,7 @@ type LabelProvider interface {
type AccessControlsService struct {
log *logger.Logger
config *model.Config
runtime *model.RuntimeConfig
labelProvider LabelProvider
}
@@ -25,6 +28,7 @@ type AccessControlServiceInput struct {
Log *logger.Logger
Config *model.Config
Runtime *model.RuntimeConfig
LabelProvider LabelProvider `optional:"true"`
}
@@ -33,12 +37,38 @@ func NewAccessControlsService(i AccessControlServiceInput) *AccessControlsServic
return &AccessControlsService{
log: i.Log,
config: i.Config,
runtime: i.Runtime,
labelProvider: i.LabelProvider,
}
}
func (service *AccessControlsService) ensureAscii(str string) bool {
for i := 0; i < len(str); i++ {
if str[i] > unicode.MaxASCII {
return false
}
}
return true
}
func (service *AccessControlsService) normalizeDomain(domain string) string {
if host, _, err := net.SplitHostPort(domain); err == nil {
domain = host
}
domain = strings.TrimRight(domain, ".")
return strings.ToLower(domain)
}
func (service *AccessControlsService) getACLs(domain string, lookup func(locator func(name string, app *model.App) bool) error) (*model.App, error) {
v := validators.NewDomainValidator(validators.DomainValidatorOptions{})
if !service.ensureAscii(domain) {
return nil, errors.New("domain contains non-ascii characters")
}
normalizedDomain := service.normalizeDomain(domain)
if !strings.HasSuffix(normalizedDomain, "."+service.runtime.CookieDomain) && normalizedDomain != service.runtime.CookieDomain {
return nil, fmt.Errorf("domain does not match cookie domain, expected %s (or a subdomain), got %s", service.runtime.CookieDomain, domain)
}
var domainMatch *model.App
var nameMatch *model.App
@@ -46,16 +76,18 @@ func (service *AccessControlsService) getACLs(domain string, lookup func(locator
locatorFunc := func(name string, app *model.App) bool {
if app.Config.Domain != "" {
err := v.Validate(app.Config.Domain, domain)
if err == nil {
if !service.ensureAscii(app.Config.Domain) {
service.log.App.Warn().Str("name", name).Str("domain", app.Config.Domain).Msg("Domain contains non-ascii characters, skipping")
return false
}
if normalizedDomain == service.normalizeDomain(app.Config.Domain) {
service.log.App.Debug().Str("name", name).Msg("Found matching container by domain")
domainMatch = app
return true
} else if !errors.Is(err, validators.ErrHostnameMismatch) {
service.log.App.Debug().Str("name", name).Err(err).Msg("Domain validation failed")
}
return false
}
if strings.HasPrefix(strings.ToLower(domain), strings.ToLower(name+".")) {
if strings.HasPrefix(normalizedDomain, strings.ToLower(name+".")) {
service.log.App.Debug().Str("name", name).Msg("Found matching container by app name")
nameMatch = app
nameMatchedApps = append(nameMatchedApps, name)
@@ -79,7 +111,7 @@ func (service *AccessControlsService) getACLs(domain string, lookup func(locator
}
if len(nameMatchedApps) > 1 {
service.log.App.Warn().Str("domain", domain).Strs("apps", nameMatchedApps).Msg("Multiple apps matched domain by name, app names must be unique, using last match")
return nil, fmt.Errorf("domain matched multiple apps by name prefix, use explicit domain config")
}
service.log.App.Debug().Str("domain", domain).Msg("Found matching app by app name")
@@ -4,8 +4,10 @@ import (
"errors"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/tinyauthapp/tinyauth/internal/model"
"github.com/tinyauthapp/tinyauth/internal/test"
"github.com/tinyauthapp/tinyauth/internal/utils/logger"
)
@@ -34,14 +36,25 @@ func TestAccessControlsService(t *testing.T) {
log := logger.NewLogger().WithTestConfig()
log.Init()
_, runtime := test.CreateTestConfigs(t)
tests := []struct {
name string
domain string
acls map[string]model.App
want *model.App
name string
domain string
acls map[string]model.App
want *model.App
errorFunc func(t *testing.T, e error)
}{
{
name: "returns ACLs for domain",
domain: "app.example.com",
acls: map[string]model.App{
"foo": {Config: model.AppConfig{Domain: "app.example.com"}},
},
want: &model.App{Config: model.AppConfig{Domain: "app.example.com"}},
},
{
name: "returns ACLs for root domain",
domain: "example.com",
acls: map[string]model.App{
"foo": {Config: model.AppConfig{Domain: "example.com"}},
@@ -65,20 +78,11 @@ func TestAccessControlsService(t *testing.T) {
want: &model.App{Config: model.AppConfig{Domain: "example.com"}},
},
{
name: "returns ACLs for non-ascii domain",
name: "returns error for non-ascii domain",
domain: "bücher.example.com",
acls: map[string]model.App{
"foo": {Config: model.AppConfig{Domain: "bücher.example.com"}},
errorFunc: func(t *testing.T, e error) {
assert.ErrorContains(t, e, "domain contains non-ascii characters")
},
want: &model.App{Config: model.AppConfig{Domain: "bücher.example.com"}},
},
{
name: "returns ACLs for punycode domain and non-ascii config",
domain: "bücher.example.com",
acls: map[string]model.App{
"foo": {Config: model.AppConfig{Domain: "xn--bcher-kva.example.com"}},
},
want: &model.App{Config: model.AppConfig{Domain: "xn--bcher-kva.example.com"}},
},
{
name: "returns ACLs with case-insensitive matching",
@@ -110,6 +114,33 @@ func TestAccessControlsService(t *testing.T) {
acls: map[string]model.App{},
want: nil,
},
{
name: "App in domain not matching with the cookie domain should return nothing with name matching",
domain: "foo.bad_example.com",
acls: map[string]model.App{
"foo": {
Path: model.AppPath{Allow: "/foo"},
},
},
want: nil,
errorFunc: func(t *testing.T, e error) {
assert.ErrorContains(t, e, "domain does not match cookie domain")
},
},
{
name: "App in domain not matching with the cookie domain should return nothing with domain matching",
domain: "foo.bad_example.com",
acls: map[string]model.App{
"foo": {
Path: model.AppPath{Allow: "/foo"},
Config: model.AppConfig{Domain: "foo.bad_example.com"},
},
},
want: nil,
errorFunc: func(t *testing.T, e error) {
assert.ErrorContains(t, e, "domain does not match cookie domain")
},
},
}
// run once for a mock provider
@@ -118,10 +149,15 @@ func TestAccessControlsService(t *testing.T) {
mock := newMockProvider(test.acls, false)
acls := NewAccessControlsService(AccessControlServiceInput{
Log: log,
Runtime: &runtime,
Config: &model.Config{},
LabelProvider: mock,
})
app, err := acls.getACLs(test.domain, mock.Lookup)
if test.errorFunc != nil {
test.errorFunc(t, err)
return
}
require.NoError(t, err)
require.Equal(t, test.want, app)
})
@@ -131,12 +167,17 @@ func TestAccessControlsService(t *testing.T) {
for _, test := range tests {
t.Run(test.name+"(staticACLs)", func(t *testing.T) {
acls := NewAccessControlsService(AccessControlServiceInput{
Log: log,
Log: log,
Runtime: &runtime,
Config: &model.Config{
Apps: test.acls,
},
})
app, err := acls.lookupStaticACLs(test.domain)
if test.errorFunc != nil {
test.errorFunc(t, err)
return
}
require.NoError(t, err)
require.Equal(t, test.want, app)
})
@@ -145,16 +186,32 @@ func TestAccessControlsService(t *testing.T) {
// get acls should return an error when the provider fails
mock := newMockProvider(map[string]model.App{}, true)
acls := NewAccessControlsService(AccessControlServiceInput{
Log: log,
Config: &model.Config{},
Log: log,
Runtime: &runtime,
Config: &model.Config{},
})
_, err := acls.getACLs("example.com", mock.Lookup)
require.Error(t, err)
assert.Error(t, err)
// get acls should return an error when multiple apps with the same domain exist
acls = NewAccessControlsService(AccessControlServiceInput{
Log: log,
Runtime: &runtime,
Config: &model.Config{
Apps: map[string]model.App{
"foo": {Path: model.AppPath{Allow: "/foo"}},
"foo.bar": {Path: model.AppPath{Allow: "/bar"}},
},
},
})
_, err = acls.GetAccessControls("foo.bar.example.com")
assert.ErrorContains(t, err, "domain matched multiple apps by name prefix, use explicit domain config")
// get access controls should get acls from
// static when static acls are configured
acls = NewAccessControlsService(AccessControlServiceInput{
Log: log,
Log: log,
Runtime: &runtime,
Config: &model.Config{
Apps: map[string]model.App{
"foo": {Config: model.AppConfig{Domain: "foo.example.com"}},
@@ -163,12 +220,12 @@ func TestAccessControlsService(t *testing.T) {
})
app, err := acls.GetAccessControls("foo.example.com")
require.NoError(t, err)
require.Equal(t, &model.App{Config: model.AppConfig{Domain: "foo.example.com"}}, app)
assert.Equal(t, &model.App{Config: model.AppConfig{Domain: "foo.example.com"}}, app)
// should return nil for no apps
app, err = acls.GetAccessControls("bar.example.com")
require.NoError(t, err)
require.Nil(t, app)
assert.Nil(t, app)
// Should use label provider if available
mock = newMockProvider(map[string]model.App{
@@ -178,10 +235,11 @@ func TestAccessControlsService(t *testing.T) {
}, false)
acls = NewAccessControlsService(AccessControlServiceInput{
Log: log,
Runtime: &runtime,
Config: &model.Config{},
LabelProvider: mock,
})
app, err = acls.GetAccessControls("bar.example.com")
require.NoError(t, err)
require.Equal(t, &model.App{Config: model.AppConfig{Domain: "bar.example.com"}}, app)
assert.Equal(t, &model.App{Config: model.AppConfig{Domain: "bar.example.com"}}, app)
}