feat: finish project basic structure and core functions
This commit is contained in:
+229
@@ -0,0 +1,229 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"nukumizu-backend/config"
|
||||
"nukumizu-backend/postLog"
|
||||
)
|
||||
|
||||
// TokenInfo holds information about an authenticated session token.
|
||||
type TokenInfo struct {
|
||||
UserID int64 `json:"userID"`
|
||||
Level string `json:"level"`
|
||||
UserName string `json:"userName"`
|
||||
CreatedAt time.Time `json:"createdAt"`
|
||||
LastAccess time.Time `json:"lastAccess"`
|
||||
}
|
||||
|
||||
var (
|
||||
tokenStore = make(map[string]*TokenInfo)
|
||||
tokenStoreLock sync.RWMutex
|
||||
)
|
||||
|
||||
// GenerateToken creates a cryptographically random 32-byte hex token.
|
||||
func GenerateToken() (string, error) {
|
||||
bytes := make([]byte, 32)
|
||||
if _, err := rand.Read(bytes); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return hex.EncodeToString(bytes), nil
|
||||
}
|
||||
|
||||
// AddToken adds a token to the in-memory token store.
|
||||
func AddToken(token string, userID int64, level string, userName string) {
|
||||
tokenStoreLock.Lock()
|
||||
defer tokenStoreLock.Unlock()
|
||||
now := time.Now()
|
||||
tokenStore[token] = &TokenInfo{
|
||||
UserID: userID,
|
||||
Level: level,
|
||||
UserName: userName,
|
||||
CreatedAt: now,
|
||||
LastAccess: now,
|
||||
}
|
||||
}
|
||||
|
||||
// GetTokenInfo retrieves token information from the store.
|
||||
func GetTokenInfo(token string) (*TokenInfo, bool) {
|
||||
tokenStoreLock.RLock()
|
||||
defer tokenStoreLock.RUnlock()
|
||||
info, exists := tokenStore[token]
|
||||
return info, exists
|
||||
}
|
||||
|
||||
// RefreshToken updates the LastAccess time for an active token.
|
||||
func RefreshToken(token string) {
|
||||
tokenStoreLock.Lock()
|
||||
defer tokenStoreLock.Unlock()
|
||||
if info, exists := tokenStore[token]; exists {
|
||||
info.LastAccess = time.Now()
|
||||
}
|
||||
}
|
||||
|
||||
// RemoveToken deletes a token from the store.
|
||||
func RemoveToken(token string) {
|
||||
tokenStoreLock.Lock()
|
||||
defer tokenStoreLock.Unlock()
|
||||
delete(tokenStore, token)
|
||||
}
|
||||
|
||||
// GetUserIDFromRequest extracts the user ID from the request's X-Token header.
|
||||
func GetUserIDFromRequest(r *http.Request) int64 {
|
||||
token := r.Header.Get("X-Token")
|
||||
if token == "" {
|
||||
return 0
|
||||
}
|
||||
tokenInfo, exists := GetTokenInfo(token)
|
||||
if !exists {
|
||||
return 0
|
||||
}
|
||||
return tokenInfo.UserID
|
||||
}
|
||||
|
||||
// GetUserLevelFromRequest extracts the user level from the request's X-Token header.
|
||||
func GetUserLevelFromRequest(r *http.Request) string {
|
||||
token := r.Header.Get("X-Token")
|
||||
if token == "" {
|
||||
return ""
|
||||
}
|
||||
tokenInfo, exists := GetTokenInfo(token)
|
||||
if !exists {
|
||||
return ""
|
||||
}
|
||||
return tokenInfo.Level
|
||||
}
|
||||
|
||||
// Auth is the central authentication and authorization function.
|
||||
// It validates the request method, X-Timestamp header (30min tolerance),
|
||||
// X-Token header, and permission level. Returns true if the request is authorized.
|
||||
//
|
||||
// Permission levels: "None" (public, no token required), "bot", "admin".
|
||||
// When level is "bot", both "bot" and "admin" tokens are accepted.
|
||||
// When level is "admin", only "admin" tokens are accepted.
|
||||
func Auth(w http.ResponseWriter, r *http.Request, targetMethod string, targetLevel string) bool {
|
||||
// Validate HTTP method.
|
||||
if r.Method != targetMethod {
|
||||
SendErrorResponse(w, http.StatusMethodNotAllowed, "method not allowed")
|
||||
return false
|
||||
}
|
||||
|
||||
// Validate X-Timestamp.
|
||||
timestamp := r.Header.Get("X-Timestamp")
|
||||
if !config.IsDebugMode() {
|
||||
if timestamp == "" {
|
||||
SendErrorResponse(w, http.StatusUnauthorized, "missing timestamp")
|
||||
return false
|
||||
}
|
||||
|
||||
ts, err := strconv.ParseInt(timestamp, 10, 64)
|
||||
if err != nil {
|
||||
SendErrorResponse(w, http.StatusUnauthorized, "invalid timestamp")
|
||||
return false
|
||||
}
|
||||
|
||||
now := time.Now().Unix()
|
||||
diff := now - ts
|
||||
if diff < 0 {
|
||||
diff = -diff
|
||||
}
|
||||
// 30 minute tolerance per agent.md.
|
||||
if diff > 1800 {
|
||||
SendErrorResponse(w, http.StatusUnauthorized, "request expired")
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// Public endpoints require no token.
|
||||
if targetLevel == "None" {
|
||||
return true
|
||||
}
|
||||
|
||||
// Validate X-Token.
|
||||
token := r.Header.Get("X-Token")
|
||||
if token == "" {
|
||||
SendErrorResponse(w, http.StatusUnauthorized, "missing token")
|
||||
return false
|
||||
}
|
||||
|
||||
tokenInfo, exists := GetTokenInfo(token)
|
||||
if !exists {
|
||||
SendErrorResponse(w, http.StatusUnauthorized, "invalid token")
|
||||
return false
|
||||
}
|
||||
|
||||
// Check permission level.
|
||||
// "bot" level accepts both "bot" and "admin" tokens.
|
||||
// "admin" level accepts only "admin" tokens.
|
||||
switch targetLevel {
|
||||
case "admin":
|
||||
if tokenInfo.Level != "admin" {
|
||||
SendErrorResponse(w, http.StatusForbidden, "permission denied")
|
||||
return false
|
||||
}
|
||||
case "bot":
|
||||
if tokenInfo.Level != "bot" && tokenInfo.Level != "admin" {
|
||||
SendErrorResponse(w, http.StatusForbidden, "permission denied")
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// Refresh token last access time.
|
||||
RefreshToken(token)
|
||||
return true
|
||||
}
|
||||
|
||||
// CleanExpiredTokens removes tokens that have been idle for over 1 hour.
|
||||
func CleanExpiredTokens() {
|
||||
tokenStoreLock.Lock()
|
||||
defer tokenStoreLock.Unlock()
|
||||
now := time.Now()
|
||||
for token, info := range tokenStore {
|
||||
if now.Sub(info.LastAccess) > 1*time.Hour {
|
||||
delete(tokenStore, token)
|
||||
postLog.Debug("Expired token removed for user: " + info.UserName)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// StartTokenCleaner starts a background goroutine that periodically cleans
|
||||
// expired tokens.
|
||||
func StartTokenCleaner() {
|
||||
ticker := time.NewTicker(10 * time.Minute)
|
||||
go func() {
|
||||
for range ticker.C {
|
||||
CleanExpiredTokens()
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
// SendSuccessResponse sends a standardized JSON success response.
|
||||
func SendSuccessResponse(w http.ResponseWriter, message string, data map[string]interface{}) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
resp := map[string]interface{}{
|
||||
"success": true,
|
||||
}
|
||||
if message != "" {
|
||||
resp["message"] = message
|
||||
}
|
||||
for k, v := range data {
|
||||
resp[k] = v
|
||||
}
|
||||
json.NewEncoder(w).Encode(resp)
|
||||
}
|
||||
|
||||
// SendErrorResponse sends a standardized JSON error response.
|
||||
func SendErrorResponse(w http.ResponseWriter, statusCode int, message string) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(statusCode)
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"success": false,
|
||||
"message": message,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,114 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// RateLimiter implements a simple per-IP rate limiter using a token bucket approach.
|
||||
type RateLimiter struct {
|
||||
mu sync.Mutex
|
||||
visitors map[string]*visitor
|
||||
rate int
|
||||
interval time.Duration
|
||||
}
|
||||
|
||||
type visitor struct {
|
||||
count int
|
||||
lastSeen time.Time
|
||||
}
|
||||
|
||||
var globalLimiter *RateLimiter
|
||||
|
||||
// InitRateLimiter initializes the global rate limiter.
|
||||
// rate is the maximum number of requests per interval.
|
||||
func InitRateLimiter(rate int, interval time.Duration) {
|
||||
globalLimiter = &RateLimiter{
|
||||
visitors: make(map[string]*visitor),
|
||||
rate: rate,
|
||||
interval: interval,
|
||||
}
|
||||
// Start a background cleaner.
|
||||
go func() {
|
||||
for {
|
||||
time.Sleep(interval)
|
||||
globalLimiter.cleanup()
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
func (rl *RateLimiter) cleanup() {
|
||||
rl.mu.Lock()
|
||||
defer rl.mu.Unlock()
|
||||
for ip, v := range rl.visitors {
|
||||
if time.Since(v.lastSeen) > rl.interval {
|
||||
delete(rl.visitors, ip)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (rl *RateLimiter) allow(ip string) bool {
|
||||
rl.mu.Lock()
|
||||
defer rl.mu.Unlock()
|
||||
v, exists := rl.visitors[ip]
|
||||
if !exists {
|
||||
rl.visitors[ip] = &visitor{count: 1, lastSeen: time.Now()}
|
||||
return true
|
||||
}
|
||||
if time.Since(v.lastSeen) > rl.interval {
|
||||
v.count = 1
|
||||
v.lastSeen = time.Now()
|
||||
return true
|
||||
}
|
||||
if v.count >= rl.rate {
|
||||
return false
|
||||
}
|
||||
v.count++
|
||||
return true
|
||||
}
|
||||
|
||||
// RateLimitMiddleware limits the number of requests per IP address.
|
||||
func RateLimitMiddleware(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if globalLimiter == nil {
|
||||
next.ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
ip := r.RemoteAddr
|
||||
if !globalLimiter.allow(ip) {
|
||||
SendErrorResponse(w, http.StatusTooManyRequests, "rate limit exceeded")
|
||||
return
|
||||
}
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
// CORSMiddleware sets CORS headers and handles preflight OPTIONS requests.
|
||||
func CORSMiddleware(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Access-Control-Allow-Origin", "*")
|
||||
w.Header().Set("Access-Control-Allow-Methods", "GET, POST, PUT, DELETE, OPTIONS")
|
||||
w.Header().Set("Access-Control-Allow-Headers", "Content-Type, X-Token, X-Timestamp, Authorization")
|
||||
w.Header().Set("Access-Control-Max-Age", "86400")
|
||||
|
||||
if r.Method == http.MethodOptions {
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
return
|
||||
}
|
||||
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
// XSSProtectionMiddleware sets security headers to help prevent XSS attacks.
|
||||
func XSSProtectionMiddleware(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("X-XSS-Protection", "1; mode=block")
|
||||
w.Header().Set("X-Content-Type-Options", "nosniff")
|
||||
w.Header().Set("X-Frame-Options", "DENY")
|
||||
w.Header().Set("Referrer-Policy", "strict-origin-when-cross-origin")
|
||||
w.Header().Set("Content-Security-Policy", "default-src 'self'; script-src 'self'; style-src 'self' 'unsafe-inline'")
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user