Skip to content
Merged
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
2 changes: 2 additions & 0 deletions .env.example
Original file line number Diff line number Diff line change
Expand Up @@ -227,6 +227,8 @@ TINYAUTH_LDAP_GROUPCACHETTL=900

# Enable the OAuth bridge, uses a new way to format OAuth user information.
TINYAUTH_EXPERIMENTAL_OAUTHBRIDGEENABLED=false
# Disable the fallback to forward_auth modules when auth_request or ext_authz fail.
TINYAUTH_EXPERIMENTAL_DISABLEAUTHMODULEFALLBACK=false
Comment thread
coderabbitai[bot] marked this conversation as resolved.

# tailscale config

Expand Down
4 changes: 3 additions & 1 deletion internal/controller/oauth_controller.go
Original file line number Diff line number Diff line change
Expand Up @@ -294,7 +294,9 @@ func (controller *OAuthController) getCookieDomain() string {

func (controller *OAuthController) isRedirectSafe(redirectURI string) bool {
v := validators.NewDomainValidator(validators.DomainValidatorOptions{
WithPort: true,
WithPort: true,
WithScheme: true,
AllowedSchemes: []string{"https", "http"},
})

_, err := v.SafeHostname(controller.runtime.AppURL)
Expand Down
82 changes: 69 additions & 13 deletions internal/controller/proxy_controller.go
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,7 @@ type ProxyContext struct {
type ProxyController struct {
log *logger.Logger
runtime *model.RuntimeConfig
config *model.Config
acls *service.AccessControlsService
auth *service.AuthService
policyEngine *service.PolicyEngine
Expand All @@ -67,6 +68,7 @@ type ProxyControllerInput struct {

Log *logger.Logger
RuntimeConfig *model.RuntimeConfig
Config *model.Config
RouterGroup *gin.RouterGroup `name:"apiRouterGroup"`
ACLsService *service.AccessControlsService
AuthService *service.AuthService
Expand All @@ -77,6 +79,7 @@ func NewProxyController(i ProxyControllerInput) *ProxyController {
controller := &ProxyController{
log: i.Log,
runtime: i.RuntimeConfig,
config: i.Config,
acls: i.ACLsService,
auth: i.AuthService,
policyEngine: i.PolicyEngine,
Expand Down Expand Up @@ -465,6 +468,10 @@ func (controller *ProxyController) getExtAuthzContext(c *gin.Context) (ProxyCont
// We get the path from the query string
path := c.Query("path")

if strings.TrimSpace(path) == "" {
return ProxyContext{}, errors.New("path not found")
}

// For envoy we need to support every method
method := c.Request.Method

Expand All @@ -477,14 +484,22 @@ func (controller *ProxyController) getExtAuthzContext(c *gin.Context) (ProxyCont
}, nil
}

func (controller *ProxyController) determineAuthModules(proxy ProxyType) []AuthModuleType {
func (controller *ProxyController) determineAuthModules(proxy ProxyType, fallbacks bool) []AuthModuleType {
switch proxy {
case Traefik, Caddy:
return []AuthModuleType{ForwardAuth}
case Envoy:
return []AuthModuleType{ExtAuthz, ForwardAuth}
authModules := []AuthModuleType{ExtAuthz}
if fallbacks {
authModules = append(authModules, ForwardAuth)
}
return authModules
case Nginx:
return []AuthModuleType{AuthRequest, ForwardAuth}
authModules := []AuthModuleType{AuthRequest}
if fallbacks {
authModules = append(authModules, ForwardAuth)
}
return authModules
default:
return []AuthModuleType{}
}
Expand Down Expand Up @@ -514,6 +529,39 @@ func (controller *ProxyController) getContextFromAuthModule(c *gin.Context, modu
return ProxyContext{}, fmt.Errorf("unsupported auth module: %v", module)
}

func (controller *ProxyController) authModuleIdentifiersPresent(c *gin.Context, module AuthModuleType) bool {
switch module {
case ForwardAuth:
_, host := controller.getHeader(c, "x-forwarded-host")
_, uri := controller.getHeader(c, "x-forwarded-uri")
return host || uri
Comment thread
coderabbitai[bot] marked this conversation as resolved.
case AuthRequest:
_, ok := controller.getHeader(c, "x-original-url")
return ok
case ExtAuthz:
return strings.TrimSpace(c.Query("path")) != ""
default:
return false
}
}

func (controller *ProxyController) ensureNoMultipleAuthModules(c *gin.Context, authModules []AuthModuleType) error {
present := 0

for _, module := range authModules {
if controller.authModuleIdentifiersPresent(c, module) {
present++
}
}

if present > 1 {
controller.log.App.Warn().Msg("Request carries headers for multiple auth modules, possible spoofing attempt, denying")
return fmt.Errorf("conflicting auth module headers")
}

return nil
}

func (controller *ProxyController) getProxyContext(c *gin.Context) (ProxyContext, error) {
var req Proxy

Expand All @@ -530,26 +578,34 @@ func (controller *ProxyController) getProxyContext(c *gin.Context) (ProxyContext

controller.log.App.Debug().Msgf("Determined proxy type: %v", proxy)

authModules := controller.determineAuthModules(proxy)
authModules := controller.determineAuthModules(proxy, !controller.config.Experimental.DisableAuthModuleFallback)

if len(authModules) == 0 {
return ProxyContext{}, fmt.Errorf("no auth modules supported for proxy: %v", req.Proxy)
}

var ctx ProxyContext
err = controller.ensureNoMultipleAuthModules(c, controller.determineAuthModules(proxy, true))

if err != nil {
return ProxyContext{}, err
}

var ctx *ProxyContext

for _, module := range authModules {
controller.log.App.Debug().Msgf("Trying to get context from auth module %v", module)
ctx, err = controller.getContextFromAuthModule(c, module)
if err == nil {
controller.log.App.Debug().Msgf("Successfully got context from auth module %v", module)
break
authModuleCtx, err := controller.getContextFromAuthModule(c, module)
if err != nil {
controller.log.App.Debug().Msgf("Failed to get context from auth module %v: %v", module, err)
continue
}
controller.log.App.Debug().Msgf("Failed to get context from auth module %v: %v", module, err)
controller.log.App.Debug().Msgf("Successfully got context from auth module %v", module)
ctx = &authModuleCtx
break
}

if err != nil {
return ProxyContext{}, err
if ctx == nil {
return ProxyContext{}, fmt.Errorf("failed to get context from any auth module")
}

// Parse the raw path to populate the cleaned path used for ACLs
Expand Down Expand Up @@ -577,5 +633,5 @@ func (controller *ProxyController) getProxyContext(c *gin.Context) (ProxyContext

ctx.IsBrowser = isBrowser
ctx.ProxyType = proxy
return ctx, nil
return *ctx, nil
}
32 changes: 30 additions & 2 deletions internal/controller/proxy_controller_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -213,7 +213,7 @@ func TestProxyController(t *testing.T) {
description: "Ensure forward auth fallback for envoy",
middlewares: []gin.HandlerFunc{},
run: func(t *testing.T, router *gin.Engine, recorder *httptest.ResponseRecorder) {
req := httptest.NewRequest("HEAD", "/api/auth/envoy?path=/hello", nil)
req := httptest.NewRequest("HEAD", "/api/auth/envoy", nil)
req.Host = ""
req.Header.Set("x-forwarded-host", "test.example.com")
req.Header.Set("x-forwarded-proto", "https")
Expand Down Expand Up @@ -261,7 +261,7 @@ func TestProxyController(t *testing.T) {
description: "Ensure extauthz with envoy non browser returns json",
middlewares: []gin.HandlerFunc{},
run: func(t *testing.T, router *gin.Engine, recorder *httptest.ResponseRecorder) {
req := httptest.NewRequest("HEAD", "/api/auth/envoy?path=/hello", nil)
req := httptest.NewRequest("HEAD", "/api/auth/envoy", nil)
req.Header.Set("x-forwarded-host", "test.example.com")
req.Header.Set("x-forwarded-proto", "https")
req.Header.Set("x-forwarded-uri", "/hello")
Expand Down Expand Up @@ -877,6 +877,32 @@ func TestProxyController(t *testing.T) {
assert.Equal(t, "bar", recorder.Header().Get("x-foo"))
},
},
{
description: "Forward auth and auth request headers should fail for nginx",
run: func(t *testing.T, router *gin.Engine, recorder *httptest.ResponseRecorder) {
req := httptest.NewRequest("GET", "/api/auth/nginx", nil)
req.Header.Set("x-forwarded-host", "foo.example.com")
req.Header.Set("x-forwarded-proto", "https")
req.Header.Set("x-forwarded-uri", "/foo?bar=foo")
req.Header.Set("x-original-url", "https://foo.example.com/foo?bar=foo")
router.ServeHTTP(recorder, req)

assert.Equal(t, http.StatusBadRequest, recorder.Code)
},
},
{
description: "Forward auth and ext authz headers should fail for envoy",
run: func(t *testing.T, router *gin.Engine, recorder *httptest.ResponseRecorder) {
req := httptest.NewRequest("HEAD", "/api/auth/envoy?path=/hello", nil)
req.Host = "foo.example.com"
req.Header.Set("x-forwarded-host", "foo.example.com")
req.Header.Set("x-forwarded-proto", "https")
req.Header.Set("x-forwarded-uri", "/foo?bar=foo")
router.ServeHTTP(recorder, req)

assert.Equal(t, http.StatusBadRequest, recorder.Code)
},
},
}

store := memory.New()
Expand All @@ -892,6 +918,7 @@ func TestProxyController(t *testing.T) {
aclsService := service.NewAccessControlsService(service.AccessControlServiceInput{
Log: log,
Config: &cfg,
Runtime: &runtime,
LabelProvider: nil,
})

Expand Down Expand Up @@ -952,6 +979,7 @@ func TestProxyController(t *testing.T) {
NewProxyController(ProxyControllerInput{
Log: log,
RuntimeConfig: &runtime,
Config: &cfg,
RouterGroup: group,
ACLsService: aclsService,
AuthService: authService,
Expand Down
3 changes: 2 additions & 1 deletion internal/model/config.go
Original file line number Diff line number Diff line change
Expand Up @@ -239,7 +239,8 @@ type LogStreamConfig struct {
}

type ExperimentalConfig struct {
OAuthBridgeEnabled bool `description:"Enable the OAuth bridge, uses a new way to format OAuth user information." yaml:"oauthBridgeEnabled,omitempty"`
OAuthBridgeEnabled bool `description:"Enable the OAuth bridge, uses a new way to format OAuth user information." yaml:"oauthBridgeEnabled,omitempty"`
DisableAuthModuleFallback bool `description:"Disable the fallback to forward_auth modules when auth_request or ext_authz fail." yaml:"disableAuthModuleFallback,omitempty"`
}

type TailscaleConfig struct {
Expand Down
48 changes: 40 additions & 8 deletions internal/service/access_controls_service.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
)

Expand All @@ -17,6 +19,7 @@ type LabelProvider interface {
type AccessControlsService struct {
log *logger.Logger
config *model.Config
runtime *model.RuntimeConfig
labelProvider LabelProvider
}

Expand All @@ -25,6 +28,7 @@ type AccessControlServiceInput struct {

Log *logger.Logger
Config *model.Config
Runtime *model.RuntimeConfig
LabelProvider LabelProvider `optional:"true"`
}

Expand All @@ -33,29 +37,57 @@ 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
var nameMatchedApps []string

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)
Expand All @@ -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")
Expand Down
Loading