Files
backend/handlers.go
T
NanamiAdmin d49090a2ab refactor(auth): extract timestamp validation to separate function
Move timestamp validation logic from handlers to ValidateTimeStamp function in auth.go to avoid code duplication and improve maintainability
2026-02-27 21:21:18 +08:00

294 lines
7.0 KiB
Go

package main
import (
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"strconv"
"super-frpc/postLog"
"time"
)
type RegisterRequest struct {
Username string `json:"username"`
Passwd string `json:"passwd"`
TimeStamp int64 `json:"timeStamp"`
Type string `json:"type"`
}
type LoginRequest struct {
Username string `json:"username"`
Passwd string `json:"passwd"`
TimeStamp int64 `json:"timeStamp"`
}
type Response struct {
Success bool `json:"success"`
Message string `json:"message,omitempty"`
Data interface{} `json:"data,omitempty"`
}
func RegisterHandler(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
SendErrorResponse(w, http.StatusMethodNotAllowed, "invalid request method")
return
}
body, err := io.ReadAll(r.Body)
if err != nil {
SendErrorResponse(w, http.StatusBadRequest, "failed to read request body")
return
}
defer r.Body.Close()
var req RegisterRequest
if err := json.Unmarshal(body, &req); err != nil {
SendErrorResponse(w, http.StatusBadRequest, "invalid request format")
return
}
if req.Username == "" || req.Passwd == "" {
SendErrorResponse(w, http.StatusBadRequest, "username and password are required")
return
}
if err := ValidateTimeStamp(req.TimeStamp); err != nil {
SendErrorResponse(w, http.StatusBadRequest, err.Error())
return
}
if !isValidInput(req.Username) || !isValidInput(req.Passwd) {
SendErrorResponse(w, http.StatusBadRequest, "invalid input: contains illegal characters")
return
}
if !isValidPassword(req.Passwd) {
SendErrorResponse(w, http.StatusBadRequest, "password does not meet complexity requirements (must contain uppercase, lowercase, digit, and special character)")
return
}
userType := req.Type
if userType == "" {
userType = "visitor"
}
validTypes := map[string]bool{
"superuser": true,
"admin": true,
"visitor": true,
}
if !validTypes[userType] {
SendErrorResponse(w, http.StatusBadRequest, "invalid user type")
return
}
userID, err := AddUser(req.Username, req.Passwd, userType)
if err != nil {
SendErrorResponse(w, http.StatusInternalServerError, err.Error())
return
}
user, err := GetUserByID(userID)
if err != nil {
SendErrorResponse(w, http.StatusInternalServerError, "failed to retrieve user after registration")
return
}
SendSuccessResponse(w, "user registered successfully", map[string]interface{}{
"userID": user.UserID,
"username": user.Username,
"type": user.Type,
})
}
func LoginHandler(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
SendErrorResponse(w, http.StatusMethodNotAllowed, "invalid request method")
return
}
body, err := io.ReadAll(r.Body)
if err != nil {
SendErrorResponse(w, http.StatusBadRequest, "failed to read request body")
return
}
defer r.Body.Close()
var req LoginRequest
if err := json.Unmarshal(body, &req); err != nil {
SendErrorResponse(w, http.StatusBadRequest, "invalid request format")
return
}
if req.Username == "" || req.Passwd == "" {
SendErrorResponse(w, http.StatusBadRequest, "username and password are required")
return
}
if err := ValidateTimeStamp(req.TimeStamp); err != nil {
SendErrorResponse(w, http.StatusBadRequest, err.Error())
return
}
if !isValidInput(req.Username) || !isValidInput(req.Passwd) {
SendErrorResponse(w, http.StatusBadRequest, "invalid input: contains illegal characters")
return
}
user, err := GetUserByUsername(req.Username)
if err != nil {
SendErrorResponse(w, http.StatusUnauthorized, "invalid username or password")
return
}
if !verifyPassword(req.Passwd, user.Passwd) {
SendErrorResponse(w, http.StatusUnauthorized, "invalid username or password")
return
}
existingTokenInfo, err := GetTokenInfo(user.UserID)
if err == nil && existingTokenInfo != nil {
SendErrorResponse(w, http.StatusConflict, "user is already logged in")
return
}
token, err := GenerateToken(user.UserID)
if err != nil {
SendErrorResponse(w, http.StatusInternalServerError, "failed to generate token")
return
}
SendSuccessResponse(w, "login successful", map[string]interface{}{
"token": token,
"userID": user.UserID,
"username": user.Username,
"type": user.Type,
})
}
func SendErrorResponse(w http.ResponseWriter, statusCode int, message string) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(statusCode)
resp := Response{
Success: false,
Message: message,
}
jsonResp, err := json.Marshal(resp)
if err != nil {
postLog.Error(fmt.Sprintf("failed to marshal error response: %v", err))
return
}
w.Write(jsonResp)
}
func SendSuccessResponse(w http.ResponseWriter, message string, data interface{}) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusOK)
resp := Response{
Success: true,
Message: message,
Data: data,
}
jsonResp, err := json.Marshal(resp)
if err != nil {
postLog.Error(fmt.Sprintf("failed to marshal success response: %v", err))
return
}
w.Write(jsonResp)
}
func ValidateRequest(w http.ResponseWriter, r *http.Request, requiredFields ...string) (int, string, error) {
body, err := io.ReadAll(r.Body)
if err != nil {
return 0, "", errors.New("failed to read request body")
}
defer r.Body.Close()
return ValidateRequestWithBody(w, r, body, requiredFields...)
}
func ValidateRequestWithBody(w http.ResponseWriter, r *http.Request, body []byte, requiredFields ...string) (int, string, error) {
var reqMap map[string]interface{}
if err := json.Unmarshal(body, &reqMap); err != nil {
return 0, "", errors.New("invalid request format")
}
token, ok := reqMap["token"].(string)
if !ok || token == "" {
return 0, "", errors.New("token is required")
}
timeStamp, ok := reqMap["timeStamp"].(float64)
if !ok {
return 0, "", errors.New("timeStamp is required")
}
if err := ValidateTimeStamp(int64(timeStamp)); err != nil {
return 0, "", err
}
userID, err := extractUserIDFromToken(token)
if err != nil {
return 0, "", err
}
if err := ValidateToken(userID, token); err != nil {
return 0, "", err
}
for _, field := range requiredFields {
if _, ok := reqMap[field]; !ok {
return 0, "", fmt.Errorf("required field %s is missing", field)
}
}
return userID, token, nil
}
func GetUserType(userID int) (string, error) {
user, err := GetUserByID(userID)
if err != nil {
return "", err
}
return user.Type, nil
}
func CheckPermission(userID int, requiredTypes ...string) error {
userType, err := GetUserType(userID)
if err != nil {
return err
}
for _, t := range requiredTypes {
if userType == t {
return nil
}
}
return errors.New("permission denied")
}
func GetClientIP(r *http.Request) string {
forwarded := r.Header.Get("X-Forwarded-For")
if forwarded != "" {
return forwarded
}
return r.RemoteAddr
}
func LogRequest(r *http.Request, userID int) {
postLog.Info(fmt.Sprintf("[%s] %s %s - UserID: %d - IP: %s",
time.Now().Format("2006-01-02 15:04:05"),
r.Method,
r.URL.Path,
userID,
GetClientIP(r),
))
}
func IntToString(i int) string {
return strconv.Itoa(i)
}