From 01ecbb6eedeb9e3a14be43b2f8d7cefd02939448 Mon Sep 17 00:00:00 2001 From: Stanislav Chzhen Date: Mon, 19 May 2025 16:48:02 +0300 Subject: [PATCH] Pull request 2408: AGDNS-2743-auth-tests Merge in DNS/adguard-home from AGDNS-2743-auth-tests to master Squashed commit of the following: commit 86fed98eff28622dbb3b321e50c4ad64b06b501a Merge: aa25b027d 5f8f06715 Author: Stanislav Chzhen Date: Fri May 16 21:58:08 2025 +0300 Merge branch 'master' into AGDNS-2743-auth-tests commit aa25b027d090079f3a716b7e53c111102bc75967 Author: Stanislav Chzhen Date: Fri May 16 17:02:54 2025 +0300 home: imp tests commit a16c1fbe76915306f24d81bd16e9e48211d409f5 Author: Stanislav Chzhen Date: Wed May 14 22:32:45 2025 +0300 home: add tests commit 6109e3575f70fa25e2ecdb8a16b01b1c9718fac2 Author: Stanislav Chzhen Date: Tue May 13 22:23:41 2025 +0300 home: add tests commit 88706e9cf24bb399137e79f8478e1d2f761b251a Author: Stanislav Chzhen Date: Wed May 7 15:46:22 2025 +0300 home: auth tests --- internal/home/authhttp_internal_test.go | 310 ++++++++++++++++++++++++ internal/home/clients_internal_test.go | 3 +- internal/home/home_internal_test.go | 3 + internal/home/tls_internal_test.go | 78 +++--- 4 files changed, 357 insertions(+), 37 deletions(-) diff --git a/internal/home/authhttp_internal_test.go b/internal/home/authhttp_internal_test.go index 9819298a..ddd649b0 100644 --- a/internal/home/authhttp_internal_test.go +++ b/internal/home/authhttp_internal_test.go @@ -1,19 +1,329 @@ package home import ( + "bytes" + "encoding/json" "net/http" + "net/http/httptest" "net/netip" "net/textproto" "net/url" + "os" "path/filepath" "testing" + "time" + "github.com/AdguardTeam/AdGuardHome/internal/aghhttp" "github.com/AdguardTeam/golibs/httphdr" "github.com/AdguardTeam/golibs/testutil" + "github.com/josharian/native" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "golang.org/x/crypto/bcrypt" ) +func TestAuth_ServeHTTP_firstRun(t *testing.T) { + storeGlobals(t) + + globalContext.firstRun = true + + mux := http.NewServeMux() + globalContext.mux = mux + + ctx := testutil.ContextWithTimeout(t, testTimeout) + web, err := initWeb(ctx, options{}, nil, nil, testLogger, nil, false) + require.NoError(t, err) + + globalContext.web = web + + testCases := []struct { + name string + path string + method string + wantCode int + }{{ + name: "root", + path: "/", + method: http.MethodGet, + wantCode: http.StatusFound, + }, { + name: "doh_mobileconfig", + path: "/apple/doh.mobileconfig", + method: http.MethodGet, + wantCode: http.StatusFound, + }, { + name: "dot_mobileconfig", + path: "/apple/dot.mobileconfig", + method: http.MethodGet, + wantCode: http.StatusFound, + }, { + name: "change_language", + path: "/control/i18n/change_language", + method: http.MethodGet, + wantCode: http.StatusFound, + }, { + name: "current_language", + path: "/control/i18n/current_language", + method: http.MethodGet, + wantCode: http.StatusFound, + }, { + name: "check_config", + path: "/control/install/check_config", + method: http.MethodPost, + wantCode: http.StatusBadRequest, + }, { + name: "configure", + path: "/control/install/configure", + method: http.MethodPost, + wantCode: http.StatusBadRequest, + }, { + name: "get_addresses", + path: "/control/install/get_addresses", + method: http.MethodGet, + wantCode: http.StatusOK, + }, { + name: "login", + path: "/control/login", + method: http.MethodPost, + wantCode: http.StatusFound, + }, { + name: "logout", + path: "/control/logout", + method: http.MethodGet, + wantCode: http.StatusFound, + }, { + name: "profile", + path: "/control/profile", + method: http.MethodGet, + wantCode: http.StatusFound, + }, { + name: "profile_update", + path: "/control/profile/update", + method: http.MethodGet, + wantCode: http.StatusFound, + }, { + name: "status", + path: "/control/status", + method: http.MethodGet, + wantCode: http.StatusFound, + }, { + name: "update", + path: "/control/update", + method: http.MethodGet, + wantCode: http.StatusFound, + }, { + name: "version", + path: "/control/version.json", + method: http.MethodGet, + wantCode: http.StatusFound, + }} + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + r := httptest.NewRequest(tc.method, tc.path, nil) + + h, pattern := mux.Handler(r) + require.NotEmpty(t, pattern) + + w := httptest.NewRecorder() + h.ServeHTTP(w, r) + + assert.Equal(t, tc.wantCode, w.Code) + }) + } +} + +func TestAuth_ServeHTTP_auth(t *testing.T) { + storeGlobals(t) + + const ( + testTTL = 60 + + glTokenFileSuffix = "test" + + userName = "name" + userPassword = "password" + ) + + 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) + native.Endian.PutUint32(glFileData, uint32(time.Now().Unix()+testTTL)) + + err = os.WriteFile(glTokenFile, glFileData, 0o644) + require.NoError(t, err) + + sessionsDB := filepath.Join(tempDir, "sessions.db") + + users := []webUser{{ + Name: userName, + PasswordHash: string(passwordHash), + }} + auth := InitAuth(sessionsDB, users, testTTL, nil, nil) + globalContext.auth = auth + + mux := http.NewServeMux() + globalContext.mux = mux + + tlsMgr, err := newTLSManager(testutil.ContextWithTimeout(t, testTimeout), &tlsManagerConfig{ + logger: testLogger, + configModified: func() {}, + }) + require.NoError(t, err) + + ctx := testutil.ContextWithTimeout(t, testTimeout) + web, err := initWeb(ctx, options{}, nil, nil, testLogger, tlsMgr, false) + require.NoError(t, err) + + globalContext.web = web + + loginCookie := generateAuthCookie(t, mux, userName, userPassword) + + testCases := []struct { + name string + path string + method string + wantCode int + }{{ + name: "change_language", + path: "/control/i18n/change_language", + method: http.MethodPost, + wantCode: http.StatusInternalServerError, + }, { + name: "current_language", + path: "/control/i18n/current_language", + method: http.MethodGet, + wantCode: http.StatusOK, + }, { + name: "profile", + path: "/control/profile", + method: http.MethodGet, + wantCode: http.StatusOK, + }, { + name: "profile_update", + path: "/control/profile/update", + method: http.MethodPut, + wantCode: http.StatusBadRequest, + }, { + name: "status", + path: "/control/status", + method: http.MethodGet, + wantCode: http.StatusOK, + }, { + name: "version", + path: "/control/version.json", + method: http.MethodGet, + wantCode: http.StatusOK, + }} + + for _, tc := range testCases { + t.Run(tc.path, func(t *testing.T) { + r := httptest.NewRequest(tc.method, tc.path, nil) + assertHandlerStatusCode(t, mux, r, http.StatusForbidden) + + r = httptest.NewRequest(tc.method, tc.path, nil) + r.SetBasicAuth(userName, userPassword) + assertHandlerStatusCode(t, mux, r, tc.wantCode) + + r = httptest.NewRequest(tc.method, tc.path, nil) + r.AddCookie(loginCookie) + assertHandlerStatusCode(t, mux, r, tc.wantCode) + + GLMode = true + t.Cleanup(func() { GLMode = false }) + + r.AddCookie(&http.Cookie{Name: glCookieName, Value: "test"}) + assertHandlerStatusCode(t, mux, r, tc.wantCode) + }) + } +} + +// generateAuthCookie is a helper function that logs in with the provided +// credentials and returns the resulting authentication cookie. +func generateAuthCookie(t *testing.T, mux *http.ServeMux, name, password string) (ac *http.Cookie) { + t.Helper() + + creds, err := json.Marshal(&loginJSON{Name: name, Password: password}) + require.NoError(t, err) + + r := httptest.NewRequest(http.MethodPost, "/control/login", bytes.NewReader(creds)) + r.Header.Set(httphdr.ContentType, aghhttp.HdrValApplicationJSON) + + w := httptest.NewRecorder() + mux.ServeHTTP(w, r) + + for _, c := range w.Result().Cookies() { + if c.Name == sessionCookieName { + return c + } + } + + return nil +} + +// assertHandlerStatusCode is a helper function that asserts the response status +// code of a HTTP handler. +func assertHandlerStatusCode(t *testing.T, h http.Handler, r *http.Request, wantCode int) { + t.Helper() + + w := httptest.NewRecorder() + h.ServeHTTP(w, r) + + assert.Equal(t, wantCode, w.Code) +} + +func TestAuth_ServeHTTP_logout(t *testing.T) { + storeGlobals(t) + + const ( + testTTL = 60 + + userName = "name" + userPassword = "password" + ) + + passwordHash, err := bcrypt.GenerateFromPassword([]byte(userPassword), bcrypt.DefaultCost) + require.NoError(t, err) + + sessionsDB := filepath.Join(t.TempDir(), "sessions.db") + + users := []webUser{{ + Name: userName, + PasswordHash: string(passwordHash), + }} + auth := InitAuth(sessionsDB, users, testTTL, nil, nil) + globalContext.auth = auth + + mux := http.NewServeMux() + globalContext.mux = mux + + ctx := testutil.ContextWithTimeout(t, testTimeout) + web, err := initWeb(ctx, options{}, nil, nil, testLogger, nil, false) + require.NoError(t, err) + + globalContext.web = web + + loginCookie := generateAuthCookie(t, mux, userName, userPassword) + require.NotNil(t, loginCookie) + + r := httptest.NewRequest(http.MethodGet, "/control/profile", nil) + r.AddCookie(loginCookie) + assertHandlerStatusCode(t, mux, r, http.StatusOK) + + r = httptest.NewRequest(http.MethodGet, "/control/logout", nil) + r.AddCookie(loginCookie) + assertHandlerStatusCode(t, mux, r, http.StatusFound) + + r = httptest.NewRequest(http.MethodGet, "/control/profile", nil) + r.AddCookie(loginCookie) + assertHandlerStatusCode(t, mux, r, http.StatusForbidden) +} + // implements http.ResponseWriter type testResponseWriter struct { hdr http.Header diff --git a/internal/home/clients_internal_test.go b/internal/home/clients_internal_test.go index 92d563f6..899ead65 100644 --- a/internal/home/clients_internal_test.go +++ b/internal/home/clients_internal_test.go @@ -5,7 +5,6 @@ import ( "github.com/AdguardTeam/AdGuardHome/internal/client" "github.com/AdguardTeam/AdGuardHome/internal/filtering" - "github.com/AdguardTeam/golibs/logutil/slogutil" "github.com/AdguardTeam/golibs/testutil" "github.com/stretchr/testify/require" ) @@ -22,7 +21,7 @@ func newClientsContainer(t *testing.T) (c *clientsContainer) { ctx := testutil.ContextWithTimeout(t, testTimeout) err := c.Init( ctx, - slogutil.NewDiscardLogger(), + testLogger, nil, client.EmptyDHCP{}, nil, diff --git a/internal/home/home_internal_test.go b/internal/home/home_internal_test.go index c56f3955..0762cb67 100644 --- a/internal/home/home_internal_test.go +++ b/internal/home/home_internal_test.go @@ -3,9 +3,12 @@ package home import ( "testing" + "github.com/AdguardTeam/golibs/logutil/slogutil" "github.com/AdguardTeam/golibs/testutil" ) +var testLogger = slogutil.NewDiscardLogger() + func TestMain(m *testing.M) { initCmdLineOpts() testutil.DiscardLogOutput(m) diff --git a/internal/home/tls_internal_test.go b/internal/home/tls_internal_test.go index 6ad36329..6ce1782f 100644 --- a/internal/home/tls_internal_test.go +++ b/internal/home/tls_internal_test.go @@ -23,7 +23,6 @@ import ( "github.com/AdguardTeam/AdGuardHome/internal/aghalg" "github.com/AdguardTeam/AdGuardHome/internal/client" "github.com/AdguardTeam/AdGuardHome/internal/dnsforward" - "github.com/AdguardTeam/golibs/logutil/slogutil" "github.com/AdguardTeam/golibs/testutil" "github.com/AdguardTeam/golibs/timeutil" "github.com/stretchr/testify/assert" @@ -65,10 +64,9 @@ kXS9jgARhhiWXJrk func TestValidateCertificates(t *testing.T) { ctx := testutil.ContextWithTimeout(t, testTimeout) - logger := slogutil.NewDiscardLogger() m, err := newTLSManager(ctx, &tlsManagerConfig{ - logger: logger, + logger: testLogger, configModified: func() {}, servePlainDNS: false, }) @@ -113,26 +111,41 @@ func TestValidateCertificates(t *testing.T) { // restores them once the test is complete. // // The global variables are: -// - [configuration.dns] -// - [homeContext.clients.storage] -// - [homeContext.dnsServer] -// - [homeContext.mux] +// - [GLMode] +// - [config] +// - [glFilePrefix] +// - [globalContext.auth] +// - [globalContext.clients.storage] +// - [globalContext.dnsServer] +// - [globalContext.firstRun] +// - [globalContext.mux] +// - [globalContext.web] // // TODO(s.chzhen): Remove this once the TLS manager no longer accesses global // variables. Make tests that use this helper concurrent. func storeGlobals(tb testing.TB) { tb.Helper() + prevGLMode := GLMode prevConfig := config + prefGLFilePrefix := glFilePrefix + auth := globalContext.auth storage := globalContext.clients.storage dnsServer := globalContext.dnsServer + firstRun := globalContext.firstRun mux := globalContext.mux + web := globalContext.web tb.Cleanup(func() { + GLMode = prevGLMode config = prevConfig + glFilePrefix = prefGLFilePrefix + globalContext.auth = auth globalContext.clients.storage = storage globalContext.dnsServer = dnsServer + globalContext.firstRun = firstRun globalContext.mux = mux + globalContext.web = web }) } @@ -207,18 +220,17 @@ func TestTLSManager_Reload(t *testing.T) { config.DNS.Port = 0 var ( - logger = slogutil.NewDiscardLogger() - ctx = testutil.ContextWithTimeout(t, testTimeout) - err error + ctx = testutil.ContextWithTimeout(t, testTimeout) + err error ) globalContext.dnsServer, err = dnsforward.NewServer(dnsforward.DNSCreateParams{ - Logger: logger, + Logger: testLogger, }) require.NoError(t, err) globalContext.clients.storage, err = client.NewStorage(ctx, &client.StorageConfig{ - Logger: logger, + Logger: testLogger, Clock: timeutil.SystemClock{}, }) require.NoError(t, err) @@ -238,7 +250,7 @@ func TestTLSManager_Reload(t *testing.T) { writeCertAndKey(t, certDER, certPath, key, keyPath) m, err := newTLSManager(ctx, &tlsManagerConfig{ - logger: logger, + logger: testLogger, configModified: func() {}, tlsSettings: tlsConfigSettings{ Enabled: true, @@ -249,7 +261,7 @@ func TestTLSManager_Reload(t *testing.T) { }) require.NoError(t, err) - web, err := initWeb(ctx, options{}, nil, nil, logger, nil, false) + web, err := initWeb(ctx, options{}, nil, nil, testLogger, nil, false) require.NoError(t, err) m.setWebAPI(web) @@ -272,13 +284,12 @@ func TestTLSManager_Reload(t *testing.T) { func TestTLSManager_HandleTLSStatus(t *testing.T) { var ( - logger = slogutil.NewDiscardLogger() - ctx = testutil.ContextWithTimeout(t, testTimeout) - err error + ctx = testutil.ContextWithTimeout(t, testTimeout) + err error ) m, err := newTLSManager(ctx, &tlsManagerConfig{ - logger: logger, + logger: testLogger, configModified: func() {}, tlsSettings: tlsConfigSettings{ Enabled: true, @@ -309,19 +320,18 @@ func TestValidateTLSSettings(t *testing.T) { globalContext.mux = http.NewServeMux() var ( - logger = slogutil.NewDiscardLogger() - ctx = testutil.ContextWithTimeout(t, testTimeout) - err error + ctx = testutil.ContextWithTimeout(t, testTimeout) + err error ) m, err := newTLSManager(ctx, &tlsManagerConfig{ - logger: logger, + logger: testLogger, configModified: func() {}, servePlainDNS: false, }) require.NoError(t, err) - web, err := initWeb(ctx, options{}, nil, nil, logger, nil, false) + web, err := initWeb(ctx, options{}, nil, nil, testLogger, nil, false) require.NoError(t, err) m.setWebAPI(web) @@ -409,13 +419,12 @@ func TestTLSManager_HandleTLSValidate(t *testing.T) { globalContext.mux = http.NewServeMux() var ( - logger = slogutil.NewDiscardLogger() - ctx = testutil.ContextWithTimeout(t, testTimeout) - err error + ctx = testutil.ContextWithTimeout(t, testTimeout) + err error ) m, err := newTLSManager(ctx, &tlsManagerConfig{ - logger: logger, + logger: testLogger, configModified: func() {}, tlsSettings: tlsConfigSettings{ Enabled: true, @@ -426,7 +435,7 @@ func TestTLSManager_HandleTLSValidate(t *testing.T) { }) require.NoError(t, err) - web, err := initWeb(ctx, options{}, nil, nil, logger, nil, false) + web, err := initWeb(ctx, options{}, nil, nil, testLogger, nil, false) require.NoError(t, err) m.setWebAPI(web) @@ -462,13 +471,12 @@ func TestTLSManager_HandleTLSConfigure(t *testing.T) { storeGlobals(t) var ( - logger = slogutil.NewDiscardLogger() - ctx = testutil.ContextWithTimeout(t, testTimeout) - err error + ctx = testutil.ContextWithTimeout(t, testTimeout) + err error ) globalContext.dnsServer, err = dnsforward.NewServer(dnsforward.DNSCreateParams{ - Logger: logger, + Logger: testLogger, }) require.NoError(t, err) @@ -484,7 +492,7 @@ func TestTLSManager_HandleTLSConfigure(t *testing.T) { require.NoError(t, err) globalContext.clients.storage, err = client.NewStorage(ctx, &client.StorageConfig{ - Logger: logger, + Logger: testLogger, Clock: timeutil.SystemClock{}, }) require.NoError(t, err) @@ -506,7 +514,7 @@ func TestTLSManager_HandleTLSConfigure(t *testing.T) { // Initialize the TLS manager and assert its configuration. m, err := newTLSManager(ctx, &tlsManagerConfig{ - logger: logger, + logger: testLogger, configModified: func() {}, tlsSettings: tlsConfigSettings{ Enabled: true, @@ -517,7 +525,7 @@ func TestTLSManager_HandleTLSConfigure(t *testing.T) { }) require.NoError(t, err) - web, err := initWeb(ctx, options{}, nil, nil, logger, nil, false) + web, err := initWeb(ctx, options{}, nil, nil, testLogger, nil, false) require.NoError(t, err) m.setWebAPI(web)