Skip to content

Commit 888e165

Browse files
committed
fix: add missing TOTP routes (enable, recovery-codes/regenerate, login verify/recovery)
1 parent 16af57e commit 888e165

1 file changed

Lines changed: 265 additions & 2 deletions

File tree

internal/routes/totp.go

Lines changed: 265 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@ import (
55
"net/http"
66

77
"github.com/go-chi/chi/v5"
8+
"github.com/golang-jwt/jwt/v5"
89
"github.com/pquerna/otp/totp"
910
"github.com/skip2/go-qrcode"
1011
"zerohost/dashboard/internal/database"
@@ -18,9 +19,14 @@ func RegisterTOTPRoutes(r chi.Router) {
1819
h := &TOTPHandler{}
1920

2021
r.Post("/totp/setup", middleware.AuthenticateToken(http.HandlerFunc(h.SetupTOTP)))
21-
r.Post("/totp/verify", middleware.AuthenticateToken(http.HandlerFunc(h.VerifyTOTP)))
22+
r.Post("/totp/enable", middleware.AuthenticateToken(http.HandlerFunc(h.EnableTOTP)))
2223
r.Post("/totp/disable", middleware.AuthenticateToken(http.HandlerFunc(h.DisableTOTP)))
2324
r.Get("/totp/status", middleware.AuthenticateToken(http.HandlerFunc(h.TOTPStatus)))
25+
26+
r.Post("/totp/recovery-codes/regenerate", middleware.AuthenticateToken(http.HandlerFunc(h.RegenerateRecoveryCodes)))
27+
28+
r.Post("/totp/verify", http.HandlerFunc(h.VerifyTOTPLogin))
29+
r.Post("/totp/recovery", http.HandlerFunc(h.RecoveryTOTP))
2430
}
2531

2632
func (h *TOTPHandler) SetupTOTP(w http.ResponseWriter, r *http.Request) {
@@ -91,15 +97,272 @@ func (h *TOTPHandler) VerifyTOTP(w http.ResponseWriter, r *http.Request) {
9197
writeJSON(w, http.StatusOK, map[string]interface{}{"success": true})
9298
}
9399

100+
func (h *TOTPHandler) EnableTOTP(w http.ResponseWriter, r *http.Request) {
101+
user := middleware.GetUser(r)
102+
103+
var body struct {
104+
Code string `json:"code"`
105+
}
106+
json.NewDecoder(r.Body).Decode(&body)
107+
108+
if body.Code == "" {
109+
jsonError(w, "Verification code is required", http.StatusBadRequest)
110+
return
111+
}
112+
113+
var secret string
114+
database.DB.QueryRow("SELECT totp_secret FROM users WHERE id = ?", user.UserID).Scan(&secret)
115+
if secret == "" {
116+
jsonError(w, "TOTP is not set up", http.StatusBadRequest)
117+
return
118+
}
119+
120+
if !totp.Validate(body.Code, secret) {
121+
jsonError(w, "Invalid verification code", http.StatusBadRequest)
122+
return
123+
}
124+
125+
codes := services.GenerateRecoveryCodes(8)
126+
hashedCodes := make([]string, len(codes))
127+
for i, c := range codes {
128+
hashedCodes[i] = services.HashRecoveryCode(c)
129+
}
130+
hashedJSON, _ := json.Marshal(hashedCodes)
131+
132+
database.DB.Exec("UPDATE users SET totp_enabled = TRUE, totp_verified = TRUE, totp_secret = ?, recovery_codes = ? WHERE id = ?",
133+
secret, string(hashedJSON), user.UserID)
134+
services.LogActivity(user.UserID, "totp_enabled", "TOTP two-factor authentication enabled", nil)
135+
136+
writeJSON(w, http.StatusOK, map[string]interface{}{
137+
"recoveryCodes": codes,
138+
})
139+
}
140+
141+
func (h *TOTPHandler) RegenerateRecoveryCodes(w http.ResponseWriter, r *http.Request) {
142+
user := middleware.GetUser(r)
143+
144+
codes := services.GenerateRecoveryCodes(8)
145+
hashedCodes := make([]string, len(codes))
146+
for i, c := range codes {
147+
hashedCodes[i] = services.HashRecoveryCode(c)
148+
}
149+
hashedJSON, _ := json.Marshal(hashedCodes)
150+
151+
database.DB.Exec("UPDATE users SET recovery_codes = ? WHERE id = ?", string(hashedJSON), user.UserID)
152+
services.LogActivity(user.UserID, "recovery_codes_regenerated", "Recovery codes regenerated", nil)
153+
154+
writeJSON(w, http.StatusOK, map[string]interface{}{
155+
"recoveryCodes": codes,
156+
})
157+
}
158+
94159
func (h *TOTPHandler) DisableTOTP(w http.ResponseWriter, r *http.Request) {
95160
user := middleware.GetUser(r)
96161

97-
database.DB.Exec("UPDATE users SET totp_enabled = FALSE, totp_verified = FALSE, totp_secret = NULL WHERE id = ?", user.UserID)
162+
var body struct {
163+
Password string `json:"password"`
164+
}
165+
json.NewDecoder(r.Body).Decode(&body)
166+
167+
var hash string
168+
database.DB.QueryRow("SELECT password_hash FROM users WHERE id = ?", user.UserID).Scan(&hash)
169+
if !verifyPassword(body.Password, hash) {
170+
jsonError(w, "Password is incorrect", http.StatusUnauthorized)
171+
return
172+
}
173+
174+
database.DB.Exec("UPDATE users SET totp_enabled = FALSE, totp_verified = FALSE, totp_secret = NULL, recovery_codes = NULL WHERE id = ?", user.UserID)
98175
services.LogActivity(user.UserID, "totp_disabled", "TOTP two-factor authentication disabled", nil)
99176

100177
writeJSON(w, http.StatusOK, map[string]interface{}{"success": true})
101178
}
102179

180+
func (h *TOTPHandler) VerifyTOTPLogin(w http.ResponseWriter, r *http.Request) {
181+
var body struct {
182+
Code string `json:"code"`
183+
TempToken string `json:"tempToken"`
184+
}
185+
json.NewDecoder(r.Body).Decode(&body)
186+
187+
if body.Code == "" || body.TempToken == "" {
188+
jsonError(w, "Code and temp token are required", http.StatusBadRequest)
189+
return
190+
}
191+
192+
claims := &middleware.UserClaims{}
193+
token, err := jwt.ParseWithClaims(body.TempToken, claims, func(token *jwt.Token) (interface{}, error) {
194+
return []byte(middleware.JWTSecret), nil
195+
})
196+
if err != nil || !token.Valid || !claims.TOTPTemp {
197+
jsonError(w, "Invalid or expired temp token", http.StatusForbidden)
198+
return
199+
}
200+
201+
var secret string
202+
database.DB.QueryRow("SELECT totp_secret FROM users WHERE id = ?", claims.UserID).Scan(&secret)
203+
if secret == "" {
204+
jsonError(w, "TOTP is not configured", http.StatusBadRequest)
205+
return
206+
}
207+
208+
if !totp.Validate(body.Code, secret) {
209+
jsonError(w, "Invalid verification code", http.StatusBadRequest)
210+
return
211+
}
212+
213+
var u struct {
214+
Email string
215+
Username string
216+
PteroUserID *int64
217+
FirstName *string
218+
LastName *string
219+
IsAdmin bool
220+
Restricted bool
221+
EmailVerified bool
222+
TokenVersion int
223+
}
224+
database.DB.QueryRow(
225+
"SELECT email, username, ptero_user_id, first_name, last_name, is_admin, restricted, email_verified, token_version FROM users WHERE id = ?",
226+
claims.UserID,
227+
).Scan(&u.Email, &u.Username, &u.PteroUserID, &u.FirstName, &u.LastName, &u.IsAdmin, &u.Restricted, &u.EmailVerified, &u.TokenVersion)
228+
229+
firstName := ""
230+
lastName := ""
231+
if u.FirstName != nil {
232+
firstName = *u.FirstName
233+
}
234+
if u.LastName != nil {
235+
lastName = *u.LastName
236+
}
237+
238+
userToken, _ := middleware.GenerateToken(middleware.UserClaims{
239+
UserID: claims.UserID,
240+
Email: u.Email,
241+
Username: u.Username,
242+
PteroID: u.PteroUserID,
243+
IsAdmin: u.IsAdmin,
244+
Restricted: u.Restricted,
245+
TokenVersion: u.TokenVersion,
246+
})
247+
248+
writeJSON(w, http.StatusOK, map[string]interface{}{
249+
"token": userToken,
250+
"user": map[string]interface{}{
251+
"id": claims.UserID,
252+
"email": u.Email,
253+
"username": u.Username,
254+
"pteroId": u.PteroUserID,
255+
"firstName": firstName,
256+
"lastName": lastName,
257+
"isAdmin": u.IsAdmin,
258+
"restricted": u.Restricted,
259+
"emailVerified": u.EmailVerified,
260+
},
261+
})
262+
}
263+
264+
func (h *TOTPHandler) RecoveryTOTP(w http.ResponseWriter, r *http.Request) {
265+
var body struct {
266+
Code string `json:"code"`
267+
TempToken string `json:"tempToken"`
268+
}
269+
json.NewDecoder(r.Body).Decode(&body)
270+
271+
if body.Code == "" || body.TempToken == "" {
272+
jsonError(w, "Code and temp token are required", http.StatusBadRequest)
273+
return
274+
}
275+
276+
claims := &middleware.UserClaims{}
277+
token, err := jwt.ParseWithClaims(body.TempToken, claims, func(token *jwt.Token) (interface{}, error) {
278+
return []byte(middleware.JWTSecret), nil
279+
})
280+
if err != nil || !token.Valid || !claims.TOTPTemp {
281+
jsonError(w, "Invalid or expired temp token", http.StatusForbidden)
282+
return
283+
}
284+
285+
var codesJSON *string
286+
database.DB.QueryRow("SELECT recovery_codes FROM users WHERE id = ?", claims.UserID).Scan(&codesJSON)
287+
if codesJSON == nil || *codesJSON == "" {
288+
jsonError(w, "No recovery codes available", http.StatusBadRequest)
289+
return
290+
}
291+
292+
var hashedCodes []string
293+
json.Unmarshal([]byte(*codesJSON), &hashedCodes)
294+
295+
valid := false
296+
var remaining []string
297+
for _, hc := range hashedCodes {
298+
if !valid && services.HashRecoveryCode(body.Code) == hc {
299+
valid = true
300+
continue
301+
}
302+
remaining = append(remaining, hc)
303+
}
304+
305+
if !valid {
306+
jsonError(w, "Invalid recovery code", http.StatusBadRequest)
307+
return
308+
}
309+
310+
remainingJSON, _ := json.Marshal(remaining)
311+
database.DB.Exec("UPDATE users SET recovery_codes = ? WHERE id = ?", string(remainingJSON), claims.UserID)
312+
313+
var u struct {
314+
Email string
315+
Username string
316+
PteroUserID *int64
317+
FirstName *string
318+
LastName *string
319+
IsAdmin bool
320+
Restricted bool
321+
EmailVerified bool
322+
TokenVersion int
323+
}
324+
database.DB.QueryRow(
325+
"SELECT email, username, ptero_user_id, first_name, last_name, is_admin, restricted, email_verified, token_version FROM users WHERE id = ?",
326+
claims.UserID,
327+
).Scan(&u.Email, &u.Username, &u.PteroUserID, &u.FirstName, &u.LastName, &u.IsAdmin, &u.Restricted, &u.EmailVerified, &u.TokenVersion)
328+
329+
firstName := ""
330+
lastName := ""
331+
if u.FirstName != nil {
332+
firstName = *u.FirstName
333+
}
334+
if u.LastName != nil {
335+
lastName = *u.LastName
336+
}
337+
338+
userToken, _ := middleware.GenerateToken(middleware.UserClaims{
339+
UserID: claims.UserID,
340+
Email: u.Email,
341+
Username: u.Username,
342+
PteroID: u.PteroUserID,
343+
IsAdmin: u.IsAdmin,
344+
Restricted: u.Restricted,
345+
TokenVersion: u.TokenVersion,
346+
})
347+
348+
services.LogActivity(claims.UserID, "recovery_code_used", "Recovery code used to sign in", nil)
349+
350+
writeJSON(w, http.StatusOK, map[string]interface{}{
351+
"token": userToken,
352+
"user": map[string]interface{}{
353+
"id": claims.UserID,
354+
"email": u.Email,
355+
"username": u.Username,
356+
"pteroId": u.PteroUserID,
357+
"firstName": firstName,
358+
"lastName": lastName,
359+
"isAdmin": u.IsAdmin,
360+
"restricted": u.Restricted,
361+
"emailVerified": u.EmailVerified,
362+
},
363+
})
364+
}
365+
103366
func (h *TOTPHandler) TOTPStatus(w http.ResponseWriter, r *http.Request) {
104367
user := middleware.GetUser(r)
105368

0 commit comments

Comments
 (0)