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{}))
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()
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
131 assert.GreaterOrEqual(t, len(jwtSecretKey()), 32)
132 assert.Equal(t, jwtSecretKey(), jwtSecretKey())
133 }
134
135 func TestValidateToken_RejectsHardcodedKey(t *testing.T) {
136
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