Skip to content
22 changes: 18 additions & 4 deletions internal/middleware/context_middleware.go
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Comment thread
steveiliop56 marked this conversation as resolved.
}

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 {
Expand Down Expand Up @@ -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)
Expand Down
65 changes: 53 additions & 12 deletions internal/middleware/context_middleware_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,6 @@ package middleware

import (
"context"
"encoding/base64"
"net/http"
"net/http/httptest"
"testing"
Expand All @@ -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) {
Expand All @@ -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()
Expand All @@ -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)
Expand Down Expand Up @@ -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)
Expand All @@ -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)
Expand All @@ -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)
Expand All @@ -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)
Expand All @@ -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)
Expand All @@ -238,14 +242,51 @@ 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)
assert.Equal(t, "testuser", userCtx.GetUsername())
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()
Expand Down
9 changes: 9 additions & 0 deletions internal/utils/security_utils.go
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ import (
"errors"
"fmt"
"net"
"net/http"
"regexp"
"strings"

Expand Down Expand Up @@ -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()
}
Loading