diff --git a/internal/middleware/context_middleware.go b/internal/middleware/context_middleware.go index 3884013d..2cbb046e 100644 --- a/internal/middleware/context_middleware.go +++ b/internal/middleware/context_middleware.go @@ -89,14 +89,26 @@ func (m *ContextMiddleware) Middleware() gin.HandlerFunc { c.Set("context", userContext) c.Next() return - } else { - m.log.App.Debug().Msgf("Error authenticating session cookie: %v", err) } + + m.log.App.Debug().Msgf("Error authenticating session cookie: %v", err) + } + + authHeader := c.GetHeader("x-tinyauth-authorization") + + if authHeader == "" { + authHeader = c.GetHeader("Authorization") } - username, password, ok := c.Request.BasicAuth() + if authHeader != "" { + username, password, ok := utils.ParseBasicAuth(authHeader) + + if !ok { + m.log.App.Debug().Msg("Error authenticating with basic auth") + c.Next() + return + } - if ok { userContext, headers, err := m.basicAuth(username, password) if err != nil { @@ -237,6 +249,8 @@ func (m *ContextMiddleware) cookieAuth(ctx context.Context, uuid string, ip stri return userContext, cookie, nil } +// basicAuth authenticates a local user and returns the user context with +// any response headers to set. func (m *ContextMiddleware) basicAuth(username string, password string) (*model.UserContext, map[string]string, error) { headers := make(map[string]string) userContext := new(model.UserContext) diff --git a/internal/middleware/context_middleware_test.go b/internal/middleware/context_middleware_test.go index 9a2df892..f2193cc9 100644 --- a/internal/middleware/context_middleware_test.go +++ b/internal/middleware/context_middleware_test.go @@ -2,7 +2,6 @@ package middleware import ( "context" - "encoding/base64" "net/http" "net/http/httptest" "testing" @@ -17,7 +16,9 @@ import ( "github.com/tinyauthapp/tinyauth/internal/repository/memory" "github.com/tinyauthapp/tinyauth/internal/service" "github.com/tinyauthapp/tinyauth/internal/test" + "github.com/tinyauthapp/tinyauth/internal/utils" "github.com/tinyauthapp/tinyauth/internal/utils/logger" + "golang.org/x/crypto/bcrypt" ) func TestContextMiddleware(t *testing.T) { @@ -26,9 +27,12 @@ func TestContextMiddleware(t *testing.T) { cfg, runtime := test.CreateTestConfigs(t) - basicAuthHeader := func(username, password string) string { - return "Basic " + base64.StdEncoding.EncodeToString([]byte(username+":"+password)) - } + colonPasswd, err := bcrypt.GenerateFromPassword([]byte("pa:ss"), bcrypt.DefaultCost) + require.NoError(t, err) + runtime.LocalUsers = append(runtime.LocalUsers, model.LocalUser{ + Username: "colonuser", + Password: string(colonPasswd), + }) seedSession := func(t *testing.T, queries repository.Store, params repository.CreateSessionParams) { t.Helper() @@ -51,7 +55,7 @@ func TestContextMiddleware(t *testing.T) { description: "Skip path bypasses auth processing", run: func(t *testing.T, args runArgs) { req := httptest.NewRequest("GET", "/api/healthz", nil) - req.Header.Set("Authorization", basicAuthHeader("testuser", "password")) + req.Header.Set("Authorization", utils.EncodeBasicAuth("testuser", "password")) userCtx, _ := args.do(req) assert.Nil(t, userCtx) @@ -165,7 +169,7 @@ func TestContextMiddleware(t *testing.T) { description: "Valid basic auth sets authenticated local context", run: func(t *testing.T, args runArgs) { req := httptest.NewRequest("GET", "/api/test", nil) - req.Header.Set("Authorization", basicAuthHeader("testuser", "password")) + req.Header.Set("Authorization", "Basic "+utils.EncodeBasicAuth("testuser", "password")) userCtx, _ := args.do(req) require.NotNil(t, userCtx) @@ -178,7 +182,7 @@ func TestContextMiddleware(t *testing.T) { description: "Invalid basic auth password yields no context", run: func(t *testing.T, args runArgs) { req := httptest.NewRequest("GET", "/api/test", nil) - req.Header.Set("Authorization", basicAuthHeader("testuser", "wrongpassword")) + req.Header.Set("Authorization", "Basic "+utils.EncodeBasicAuth("testuser", "wrongpassword")) userCtx, _ := args.do(req) assert.Nil(t, userCtx) @@ -188,7 +192,7 @@ func TestContextMiddleware(t *testing.T) { description: "Basic auth is rejected for users with totp", run: func(t *testing.T, args runArgs) { req := httptest.NewRequest("GET", "/api/test", nil) - req.Header.Set("Authorization", basicAuthHeader("totpuser", "password")) + req.Header.Set("Authorization", "Basic "+utils.EncodeBasicAuth("totpuser", "password")) userCtx, _ := args.do(req) assert.Nil(t, userCtx) @@ -199,12 +203,12 @@ func TestContextMiddleware(t *testing.T) { run: func(t *testing.T, args runArgs) { for range 3 { req := httptest.NewRequest("GET", "/api/test", nil) - req.Header.Set("Authorization", basicAuthHeader("testuser", "wrongpassword")) + req.Header.Set("Authorization", "Basic "+utils.EncodeBasicAuth("testuser", "wrongpassword")) args.do(req) } req := httptest.NewRequest("GET", "/api/test", nil) - req.Header.Set("Authorization", basicAuthHeader("testuser", "password")) + req.Header.Set("Authorization", "Basic "+utils.EncodeBasicAuth("testuser", "password")) userCtx, recorder := args.do(req) assert.Nil(t, userCtx) @@ -226,7 +230,7 @@ func TestContextMiddleware(t *testing.T) { req := httptest.NewRequest("GET", "/api/test", nil) req.AddCookie(&http.Cookie{Name: "tinyauth-session", Value: uuid}) - req.Header.Set("Authorization", basicAuthHeader("totpuser", "password")) + req.Header.Set("Authorization", "Basic "+utils.EncodeBasicAuth("totpuser", "password")) userCtx, _ := args.do(req) require.NotNil(t, userCtx) @@ -238,7 +242,34 @@ func TestContextMiddleware(t *testing.T) { description: "Ensure fallback to basic auth when cookie is missing", run: func(t *testing.T, args runArgs) { req := httptest.NewRequest("GET", "/api/test", nil) - req.Header.Set("Authorization", basicAuthHeader("testuser", "password")) + req.Header.Set("Authorization", "Basic "+utils.EncodeBasicAuth("testuser", "password")) + userCtx, _ := args.do(req) + + require.NotNil(t, userCtx) + assert.Equal(t, "testuser", userCtx.GetUsername()) + assert.True(t, userCtx.Authenticated) + }, + }, + { + description: "Valid x-tinyauth-Authorization sets authenticated local context", + run: func(t *testing.T, args runArgs) { + req := httptest.NewRequest("GET", "/api/test", nil) + req.Header.Set("x-tinyauth-authorization", "Basic "+utils.EncodeBasicAuth("testuser", "password")) + req.SetBasicAuth("testuser", "password") + userCtx, _ := args.do(req) + + require.NotNil(t, userCtx) + assert.Equal(t, model.ProviderLocal, userCtx.Provider) + assert.Equal(t, "testuser", userCtx.GetUsername()) + assert.True(t, userCtx.Authenticated) + }, + }, + { + description: "x-tinyauth-authorization takes priority over authorization", + run: func(t *testing.T, args runArgs) { + req := httptest.NewRequest("GET", "/api/test", nil) + req.Header.Set("x-tinyauth-authorization", "Basic "+utils.EncodeBasicAuth("testuser", "password")) + req.Header.Set("authorization", "Basic "+utils.EncodeBasicAuth("testuser", "wrongpassword")) userCtx, _ := args.do(req) require.NotNil(t, userCtx) @@ -246,6 +277,16 @@ func TestContextMiddleware(t *testing.T) { assert.True(t, userCtx.Authenticated) }, }, + { + description: "x-tinyauth-authorization header being invalid doesn't fail the request", + run: func(t *testing.T, args runArgs) { + req := httptest.NewRequest("GET", "/api/test", nil) + req.Header.Set("x-tinyauth-authorization", "Basic "+utils.EncodeBasicAuth("testuser", "wrongpassword")) + userCtx, _ := args.do(req) + + assert.Nil(t, userCtx) + }, + }, } ctx := context.TODO() diff --git a/internal/utils/security_utils.go b/internal/utils/security_utils.go index 71b59d41..c6bef108 100644 --- a/internal/utils/security_utils.go +++ b/internal/utils/security_utils.go @@ -6,6 +6,7 @@ import ( "errors" "fmt" "net" + "net/http" "regexp" "strings" @@ -116,3 +117,11 @@ func GenerateString(length int) string { rand.Read(src) return base64.RawURLEncoding.EncodeToString(src)[:length] } + +func ParseBasicAuth(auth string) (username, password string, ok bool) { + req := &http.Request{ + Header: make(http.Header), + } + req.Header.Set("Authorization", auth) + return req.BasicAuth() +}