mirror of
https://git.vectorsigma.ru/public/AdGuardHome.git
synced 2026-08-03 16:59:30 +00:00
home: add tests
This commit is contained in:
@@ -54,7 +54,7 @@ type authConfig struct {
|
|||||||
|
|
||||||
// rateLimiter manages the rate limiting for login attempts. It must not be
|
// rateLimiter manages the rate limiting for login attempts. It must not be
|
||||||
// nil.
|
// nil.
|
||||||
rateLimiter loginRaateLimiter
|
rateLimiter loginRateLimiter
|
||||||
|
|
||||||
// trustedProxies is a set of subnets considered as trusted.
|
// trustedProxies is a set of subnets considered as trusted.
|
||||||
trustedProxies netutil.SubnetSet
|
trustedProxies netutil.SubnetSet
|
||||||
@@ -75,12 +75,27 @@ type authConfig struct {
|
|||||||
|
|
||||||
// auth stores web user information and handles authentication.
|
// auth stores web user information and handles authentication.
|
||||||
type auth struct {
|
type auth struct {
|
||||||
logger *slog.Logger
|
// logger is used to log the operation of the auth module.
|
||||||
rateLimiter loginRaateLimiter
|
logger *slog.Logger
|
||||||
|
|
||||||
|
// rateLimiter manages rate limiting for login attempts.
|
||||||
|
rateLimiter loginRateLimiter
|
||||||
|
|
||||||
|
// trustedProxies is a set of subnets considered trusted.
|
||||||
trustedProxies netutil.SubnetSet
|
trustedProxies netutil.SubnetSet
|
||||||
sessions aghuser.SessionStorage
|
|
||||||
users aghuser.DB
|
// sessions stores web users' sessions.
|
||||||
isGLiNet bool
|
sessions aghuser.SessionStorage
|
||||||
|
|
||||||
|
// users stores user credentials.
|
||||||
|
users aghuser.DB
|
||||||
|
|
||||||
|
// isGLiNet indicates whether GLiNet mode is enabled.
|
||||||
|
isGLiNet bool
|
||||||
|
|
||||||
|
// isUserless indicates that there are no users defined in the configuration
|
||||||
|
// file.
|
||||||
|
isUserless bool
|
||||||
}
|
}
|
||||||
|
|
||||||
// newAuth returns the new properly initialized *auth.
|
// newAuth returns the new properly initialized *auth.
|
||||||
@@ -111,6 +126,7 @@ func newAuth(ctx context.Context, conf *authConfig) (a *auth, err error) {
|
|||||||
sessions: s,
|
sessions: s,
|
||||||
users: userDB,
|
users: userDB,
|
||||||
isGLiNet: conf.isGLiNet,
|
isGLiNet: conf.isGLiNet,
|
||||||
|
isUserless: len(conf.users) == 0,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -174,6 +190,8 @@ func (a *auth) addUser(ctx context.Context, u *webUser, password string) (err er
|
|||||||
panic(err)
|
panic(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
a.isUserless = false
|
||||||
|
|
||||||
a.logger.DebugContext(ctx, "added user", "login", u.Name)
|
a.logger.DebugContext(ctx, "added user", "login", u.Name)
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
@@ -320,7 +320,7 @@ type authMiddlewareDefaultConfig struct {
|
|||||||
logger *slog.Logger
|
logger *slog.Logger
|
||||||
|
|
||||||
// rateLimiter manages the rate limiting for login attempts.
|
// rateLimiter manages the rate limiting for login attempts.
|
||||||
rateLimiter loginRaateLimiter
|
rateLimiter loginRateLimiter
|
||||||
|
|
||||||
// trustedProxies is a set of subnets considered as trusted.
|
// trustedProxies is a set of subnets considered as trusted.
|
||||||
//
|
//
|
||||||
@@ -340,7 +340,7 @@ type authMiddlewareDefaultConfig struct {
|
|||||||
// passes it with the context.
|
// passes it with the context.
|
||||||
type authMiddlewareDefault struct {
|
type authMiddlewareDefault struct {
|
||||||
logger *slog.Logger
|
logger *slog.Logger
|
||||||
rateLimiter loginRaateLimiter
|
rateLimiter loginRateLimiter
|
||||||
trustedProxies netutil.SubnetSet
|
trustedProxies netutil.SubnetSet
|
||||||
sessions aghuser.SessionStorage
|
sessions aghuser.SessionStorage
|
||||||
users aghuser.DB
|
users aghuser.DB
|
||||||
|
|||||||
@@ -9,8 +9,8 @@ import (
|
|||||||
// cache.
|
// cache.
|
||||||
const failedAuthTTL = 1 * time.Minute
|
const failedAuthTTL = 1 * time.Minute
|
||||||
|
|
||||||
// loginRaateLimiter is an interface for rate limiting login attempts.
|
// loginRateLimiter is an interface for rate limiting login attempts.
|
||||||
type loginRaateLimiter interface {
|
type loginRateLimiter interface {
|
||||||
// check returns the duration of time left until a user is unblocked.
|
// check returns the duration of time left until a user is unblocked.
|
||||||
// A non-positive result indicates that the user is not blocked.
|
// A non-positive result indicates that the user is not blocked.
|
||||||
check(usrID string) (left time.Duration)
|
check(usrID string) (left time.Duration)
|
||||||
@@ -66,7 +66,7 @@ func newAuthRateLimiter(blockDur time.Duration, maxAttempts uint) (ab *authRateL
|
|||||||
}
|
}
|
||||||
|
|
||||||
// type check
|
// type check
|
||||||
var _ loginRaateLimiter = (*authRateLimiter)(nil)
|
var _ loginRateLimiter = (*authRateLimiter)(nil)
|
||||||
|
|
||||||
// cleanupLocked checks each blocked users removing ones with expired TTL. For
|
// cleanupLocked checks each blocked users removing ones with expired TTL. For
|
||||||
// internal use only.
|
// internal use only.
|
||||||
|
|||||||
@@ -863,7 +863,7 @@ func initUsers(
|
|||||||
baseLogger *slog.Logger,
|
baseLogger *slog.Logger,
|
||||||
isGLiNet bool,
|
isGLiNet bool,
|
||||||
) (auth *auth, err error) {
|
) (auth *auth, err error) {
|
||||||
var rateLimiter loginRaateLimiter
|
var rateLimiter loginRateLimiter
|
||||||
if config.AuthAttempts > 0 && config.AuthBlockMin > 0 {
|
if config.AuthAttempts > 0 && config.AuthBlockMin > 0 {
|
||||||
blockDur := time.Duration(config.AuthBlockMin) * time.Minute
|
blockDur := time.Duration(config.AuthBlockMin) * time.Minute
|
||||||
rateLimiter = newAuthRateLimiter(blockDur, config.AuthAttempts)
|
rateLimiter = newAuthRateLimiter(blockDur, config.AuthAttempts)
|
||||||
|
|||||||
@@ -47,10 +47,15 @@ type profileJSON struct {
|
|||||||
// handleGetProfile is the handler for GET /control/profile endpoint.
|
// handleGetProfile is the handler for GET /control/profile endpoint.
|
||||||
func (web *webAPI) handleGetProfile(w http.ResponseWriter, r *http.Request) {
|
func (web *webAPI) handleGetProfile(w http.ResponseWriter, r *http.Request) {
|
||||||
var name string
|
var name string
|
||||||
u, ok := webUserFromContext(r.Context())
|
|
||||||
// There may be no user in the context if the configuration file defines no
|
if !(web.auth.isUserless || web.auth.isGLiNet) {
|
||||||
// users.
|
u, ok := webUserFromContext(r.Context())
|
||||||
if ok {
|
if !ok {
|
||||||
|
w.WriteHeader(http.StatusUnauthorized)
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
name = string(u.Login)
|
name = string(u.Login)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
119
internal/home/profilehttp_internal_test.go
Normal file
119
internal/home/profilehttp_internal_test.go
Normal file
@@ -0,0 +1,119 @@
|
|||||||
|
package home
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/AdguardTeam/AdGuardHome/internal/agh"
|
||||||
|
"github.com/AdguardTeam/golibs/testutil"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"golang.org/x/crypto/bcrypt"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestWeb_HandleGetProfile(t *testing.T) {
|
||||||
|
storeGlobals(t)
|
||||||
|
|
||||||
|
const (
|
||||||
|
testTTL = 60
|
||||||
|
|
||||||
|
glTokenFileSuffix = "test"
|
||||||
|
|
||||||
|
userName = "name"
|
||||||
|
userPassword = "password"
|
||||||
|
|
||||||
|
path = "/control/profile"
|
||||||
|
)
|
||||||
|
|
||||||
|
passwordHash, err := bcrypt.GenerateFromPassword([]byte(userPassword), bcrypt.DefaultCost)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
tempDir := t.TempDir()
|
||||||
|
glFilePrefix = tempDir + "/gl_token_"
|
||||||
|
glTokenFile := glFilePrefix + glTokenFileSuffix
|
||||||
|
|
||||||
|
glFileData := make([]byte, 4)
|
||||||
|
binary.NativeEndian.PutUint32(glFileData, uint32(time.Now().Unix()+testTTL))
|
||||||
|
|
||||||
|
err = os.WriteFile(glTokenFile, glFileData, 0o644)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
sessionsDB := filepath.Join(tempDir, "sessions.db")
|
||||||
|
|
||||||
|
user := &webUser{
|
||||||
|
Name: userName,
|
||||||
|
PasswordHash: string(passwordHash),
|
||||||
|
}
|
||||||
|
|
||||||
|
auth, err := newAuth(testutil.ContextWithTimeout(t, testTimeout), &authConfig{
|
||||||
|
baseLogger: testLogger,
|
||||||
|
rateLimiter: emptyRateLimiter{},
|
||||||
|
trustedProxies: nil,
|
||||||
|
dbFilename: sessionsDB,
|
||||||
|
users: nil,
|
||||||
|
sessionTTL: testTTL * time.Second,
|
||||||
|
isGLiNet: false,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
t.Cleanup(func() { auth.close(testutil.ContextWithTimeout(t, testTimeout)) })
|
||||||
|
|
||||||
|
globalContext.mux = http.NewServeMux()
|
||||||
|
|
||||||
|
tlsMgr, err := newTLSManager(testutil.ContextWithTimeout(t, testTimeout), &tlsManagerConfig{
|
||||||
|
logger: testLogger,
|
||||||
|
confModifier: agh.EmptyConfigModifier{},
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
web, err := initWeb(
|
||||||
|
testutil.ContextWithTimeout(t, testTimeout),
|
||||||
|
options{},
|
||||||
|
nil,
|
||||||
|
nil,
|
||||||
|
testLogger,
|
||||||
|
tlsMgr,
|
||||||
|
auth,
|
||||||
|
agh.EmptyConfigModifier{},
|
||||||
|
false,
|
||||||
|
)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
globalContext.web = web
|
||||||
|
|
||||||
|
mux := auth.middleware().Wrap(globalContext.mux)
|
||||||
|
|
||||||
|
require.True(t, t.Run("userless", func(t *testing.T) {
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
r := httptest.NewRequest(http.MethodGet, path, nil)
|
||||||
|
|
||||||
|
web.handleGetProfile(w, r)
|
||||||
|
assert.Equal(t, http.StatusOK, w.Code)
|
||||||
|
}))
|
||||||
|
|
||||||
|
require.True(t, t.Run("add_user", func(t *testing.T) {
|
||||||
|
ctx := testutil.ContextWithTimeout(t, testTimeout)
|
||||||
|
err = auth.addUser(ctx, user, userPassword)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
r := httptest.NewRequest(http.MethodGet, path, nil)
|
||||||
|
|
||||||
|
web.handleGetProfile(w, r)
|
||||||
|
assert.Equal(t, http.StatusUnauthorized, w.Code)
|
||||||
|
|
||||||
|
w = httptest.NewRecorder()
|
||||||
|
r = httptest.NewRequest(http.MethodGet, path, nil)
|
||||||
|
|
||||||
|
loginCookie := generateAuthCookie(t, mux, userName, userPassword)
|
||||||
|
r.AddCookie(loginCookie)
|
||||||
|
|
||||||
|
web.handleGetProfile(w, r)
|
||||||
|
assert.Equal(t, http.StatusUnauthorized, w.Code)
|
||||||
|
}))
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user