chore: add Auth function to handle all permission verification; remove all old auth codes
This commit is contained in:
+28
-82
@@ -3,7 +3,6 @@ package main
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"super-frpc/postLog"
|
||||
@@ -50,80 +49,42 @@ func SendSuccessResponse(w http.ResponseWriter, message string, data interface{}
|
||||
w.Write(jsonResp)
|
||||
}
|
||||
|
||||
func ValidateRequest(w http.ResponseWriter, r *http.Request, requiredFields ...string) (int, string, error) { // ValidateRequest validates the request body and header
|
||||
body, err := io.ReadAll(r.Body)
|
||||
func Auth(w http.ResponseWriter, r *http.Request, targetMethod string, allowedUserLevels ...string) (int, error) {
|
||||
if r.Method != targetMethod {
|
||||
return 0, fmt.Errorf("Method not allowed: %s", targetMethod)
|
||||
}
|
||||
|
||||
if !isDebug && !ValidateTimeStamp(r.Header) {
|
||||
return 0, fmt.Errorf("Invalid or missing X-Timestamp in header")
|
||||
}
|
||||
|
||||
userID, err := extractUserIDFromToken(r.Header.Get("X-Token"))
|
||||
if err != nil {
|
||||
return 0, "", fmt.Errorf("Failed to read request body: %w", err)
|
||||
}
|
||||
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, "", fmt.Errorf("Invalid request format: %w", err)
|
||||
return 0, fmt.Errorf("Invalid token format: %w", err)
|
||||
}
|
||||
|
||||
token := r.Header.Get("X-Token")
|
||||
if token == "" {
|
||||
return 0, "", fmt.Errorf("Token is required in header: %s", token)
|
||||
if err := ValidateToken(userID, r.Header.Get("X-Token")); err != nil {
|
||||
return 0, fmt.Errorf("Token validation failed: %w", err)
|
||||
}
|
||||
|
||||
if !isDebug {
|
||||
if !ValidateTimeStamp(r.Header) {
|
||||
return 0, "", fmt.Errorf("Invalid or missing X-Timestamp in header")
|
||||
if len(allowedUserLevels) > 0 {
|
||||
currentUser, err := DBQuerySpecificUser(userID)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("Failed to query user: %w", err)
|
||||
}
|
||||
allowed := false
|
||||
for _, level := range allowedUserLevels {
|
||||
if currentUser.Type == level {
|
||||
allowed = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !allowed {
|
||||
return 0, fmt.Errorf("User level not allowed: required one of %v, got %s", allowedUserLevels, currentUser.Type)
|
||||
}
|
||||
}
|
||||
|
||||
userID, err := extractUserIDFromToken(token)
|
||||
if err != nil {
|
||||
return 0, "", fmt.Errorf("Invalid token format: %w", err)
|
||||
}
|
||||
|
||||
if err := ValidateToken(userID, token); err != nil {
|
||||
return 0, "", fmt.Errorf("Token validation failed: %w", err)
|
||||
}
|
||||
|
||||
for _, field := range requiredFields {
|
||||
if _, ok := reqMap[field]; !ok {
|
||||
return 0, "", fmt.Errorf("required field %s is missing: %s", field, reqMap[field])
|
||||
}
|
||||
}
|
||||
|
||||
return userID, token, nil
|
||||
}
|
||||
|
||||
func ValidateRequestWithHeader(w http.ResponseWriter, r *http.Request, requiredFields ...string) (int, string, error) {
|
||||
token := r.Header.Get("X-Token")
|
||||
if token == "" {
|
||||
return 0, "", fmt.Errorf("Token is required in header: %s", token)
|
||||
}
|
||||
|
||||
if !isDebug {
|
||||
if !ValidateTimeStamp(r.Header) {
|
||||
return 0, "", fmt.Errorf("Invalid or missing X-Timestamp in header")
|
||||
}
|
||||
}
|
||||
|
||||
userID, err := extractUserIDFromToken(token)
|
||||
if err != nil {
|
||||
return 0, "", fmt.Errorf("Invalid token format in header: %w", err)
|
||||
}
|
||||
|
||||
if err := ValidateToken(userID, token); err != nil {
|
||||
return 0, "", fmt.Errorf("Token validation failed in header: %w", err)
|
||||
}
|
||||
|
||||
for _, field := range requiredFields {
|
||||
headerValue := r.Header.Get(fmt.Sprintf("X-%s", field))
|
||||
if headerValue == "" {
|
||||
return 0, "", fmt.Errorf("required field %s is missing in header: %s", field, headerValue)
|
||||
}
|
||||
}
|
||||
|
||||
return userID, token, nil
|
||||
return userID, nil
|
||||
}
|
||||
|
||||
func GetUserType(userID int) (string, error) {
|
||||
@@ -134,21 +95,6 @@ func GetUserType(userID int) (string, error) {
|
||||
return user.Type, nil
|
||||
}
|
||||
|
||||
func CheckPermission(userID int, requiredTypes ...string) error {
|
||||
userType, err := GetUserType(userID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("Failed to check permission: %w", err)
|
||||
}
|
||||
|
||||
for _, t := range requiredTypes {
|
||||
if userType == t {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
return fmt.Errorf("Permission denied for user type %s", userType)
|
||||
}
|
||||
|
||||
func GetClientIP(r *http.Request) string {
|
||||
forwarded := r.Header.Get("X-Forwarded-For")
|
||||
if forwarded != "" {
|
||||
|
||||
Reference in New Issue
Block a user