...

Source file src/github.com/cybertec-postgresql/pgwatch/v6/internal/webserver/jwt_test.go

Documentation: github.com/cybertec-postgresql/pgwatch/v6/internal/webserver

     1  package webserver
     2  
     3  import (
     4  	"bytes"
     5  	"io"
     6  	"net/http"
     7  	"net/http/httptest"
     8  	"testing"
     9  	"time"
    10  
    11  	"github.com/golang-jwt/jwt/v5"
    12  	jsoniter "github.com/json-iterator/go"
    13  	"github.com/stretchr/testify/assert"
    14  )
    15  
    16  var json = jsoniter.ConfigFastest
    17  
    18  func TestIsCorrectPassword(t *testing.T) {
    19  	ts := &WebUIServer{CmdOpts: CmdOpts{WebUser: "user", WebPassword: "pass"}}
    20  	assert.True(t, ts.IsCorrectPassword(loginReq{Username: "user", Password: "pass"}))
    21  	assert.False(t, ts.IsCorrectPassword(loginReq{Username: "user", Password: "wrong"}))
    22  	assert.True(t, (&WebUIServer{}).IsCorrectPassword(loginReq{})) // empty user/pass disables auth
    23  }
    24  
    25  func TestHandleLogin_POST_Success(t *testing.T) {
    26  	ts := &WebUIServer{CmdOpts: CmdOpts{WebUser: "user", WebPassword: "pass"}}
    27  	body, _ := json.Marshal(map[string]string{"user": "user", "password": "pass"})
    28  	r := httptest.NewRequest(http.MethodPost, "/login", bytes.NewReader(body))
    29  	w := httptest.NewRecorder()
    30  	ts.handleLogin(w, r)
    31  	resp := w.Result()
    32  	assert.Equal(t, http.StatusOK, resp.StatusCode)
    33  	token, _ := io.ReadAll(resp.Body)
    34  	assert.NotEmpty(t, string(token))
    35  }
    36  
    37  func TestHandleLogin_POST_Fail(t *testing.T) {
    38  	ts := &WebUIServer{CmdOpts: CmdOpts{WebUser: "user", WebPassword: "pass"}}
    39  	body, _ := json.Marshal(map[string]string{"user": "user", "password": "wrong"})
    40  	r := httptest.NewRequest(http.MethodPost, "/login", bytes.NewReader(body))
    41  	w := httptest.NewRecorder()
    42  	ts.handleLogin(w, r)
    43  	resp := w.Result()
    44  	assert.Equal(t, http.StatusUnauthorized, resp.StatusCode)
    45  }
    46  
    47  func TestHandleLogin_POST_BadJSON(t *testing.T) {
    48  	ts := &WebUIServer{CmdOpts: CmdOpts{WebUser: "user", WebPassword: "pass"}}
    49  	r := httptest.NewRequest(http.MethodPost, "/login", bytes.NewReader([]byte("notjson")))
    50  	w := httptest.NewRecorder()
    51  	ts.handleLogin(w, r)
    52  	resp := w.Result()
    53  	assert.Equal(t, http.StatusInternalServerError, resp.StatusCode)
    54  }
    55  
    56  func TestHandleLogin_GET(t *testing.T) {
    57  	ts := &WebUIServer{}
    58  	r := httptest.NewRequest(http.MethodGet, "/login", nil)
    59  	w := httptest.NewRecorder()
    60  	ts.handleLogin(w, r)
    61  	resp := w.Result()
    62  	assert.Equal(t, http.StatusMethodNotAllowed, resp.StatusCode)
    63  }
    64  
    65  func TestGenerateAndValidateJWT(t *testing.T) {
    66  	token, err := generateJWT("user1")
    67  	assert.NoError(t, err)
    68  	r := httptest.NewRequest(http.MethodGet, "/", nil)
    69  	r.Header.Set("Token", token)
    70  	assert.NoError(t, validateToken(r))
    71  }
    72  
    73  func TestValidateToken_MissingToken(t *testing.T) {
    74  	r := httptest.NewRequest(http.MethodGet, "/", nil)
    75  	err := validateToken(r)
    76  	assert.Error(t, err)
    77  	assert.Contains(t, err.Error(), "can not find token")
    78  }
    79  
    80  func TestValidateToken_InvalidToken(t *testing.T) {
    81  	r := httptest.NewRequest(http.MethodGet, "/", nil)
    82  	r.Header.Set("Token", "invalidtoken")
    83  	err := validateToken(r)
    84  	assert.Error(t, err)
    85  }
    86  
    87  func TestEnsureAuth_ServeHTTP(t *testing.T) {
    88  	called := false
    89  	h := func(w http.ResponseWriter, _ *http.Request) {
    90  		called = true
    91  		w.WriteHeader(http.StatusTeapot)
    92  	}
    93  	token, _ := generateJWT("user1")
    94  	r := httptest.NewRequest(http.MethodGet, "/", nil)
    95  	r.Header.Set("Token", token)
    96  	w := httptest.NewRecorder()
    97  	NewEnsureAuth(h).ServeHTTP(w, r)
    98  	resp := w.Result()
    99  	assert.Equal(t, http.StatusTeapot, resp.StatusCode)
   100  	assert.True(t, called)
   101  }
   102  
   103  func TestEnsureAuth_ServeHTTP_InvalidToken(t *testing.T) {
   104  	h := func(w http.ResponseWriter, _ *http.Request) {
   105  		w.WriteHeader(http.StatusTeapot)
   106  	}
   107  	r := httptest.NewRequest(http.MethodGet, "/", nil)
   108  	r.Header.Set("Token", "invalidtoken")
   109  	w := httptest.NewRecorder()
   110  	NewEnsureAuth(h).ServeHTTP(w, r)
   111  	resp := w.Result()
   112  	assert.Equal(t, http.StatusUnauthorized, resp.StatusCode)
   113  }
   114  
   115  func TestJWT_Expiration(t *testing.T) {
   116  	tok := jwt.New(jwt.SigningMethodHS256)
   117  	claims := tok.Claims.(jwt.MapClaims)
   118  	claims["authorized"] = true
   119  	claims["username"] = "user"
   120  	claims["exp"] = time.Now().Add(-time.Hour).Unix() // expired
   121  	token, _ := tok.SignedString(jwtSecretKey())
   122  	r := httptest.NewRequest(http.MethodGet, "/", nil)
   123  	r.Header.Set("Token", token)
   124  	err := validateToken(r)
   125  	assert.Error(t, err)
   126  	assert.Contains(t, err.Error(), "token is expired")
   127  }
   128  
   129  func TestJWTSecretKey_RandomAndStable(t *testing.T) {
   130  	// It must be at least 32 bytes of entropy and stable within a process.
   131  	assert.GreaterOrEqual(t, len(jwtSecretKey()), 32)
   132  	assert.Equal(t, jwtSecretKey(), jwtSecretKey())
   133  }
   134  
   135  func TestValidateToken_RejectsHardcodedKey(t *testing.T) {
   136  	// A token forged with the old, publicly known key must be rejected.
   137  	tok := jwt.New(jwt.SigningMethodHS256)
   138  	claims := tok.Claims.(jwt.MapClaims)
   139  	claims["authorized"] = false
   140  	claims["username"] = "intruder"
   141  	claims["exp"] = time.Now().Add(time.Hour).Unix()
   142  	forged, err := tok.SignedString([]byte("5m3R7K4754p4m"))
   143  	assert.NoError(t, err)
   144  	r := httptest.NewRequest(http.MethodGet, "/", nil)
   145  	r.Header.Set("Token", forged)
   146  	assert.Error(t, validateToken(r))
   147  }
   148