diff --git a/cmd/server/main.go b/cmd/server/main.go index dd484c1..e3c9402 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -5,9 +5,11 @@ import ( "net/http" "os" "shagram/internal/api" + "shagram/internal/auth" "shagram/internal/db" "shagram/internal/models" "shagram/internal/websocket" + "time" "github.com/gin-gonic/gin" ) @@ -29,6 +31,27 @@ func main() { router := gin.Default() + router.POST("/api/auth/login", func(c *gin.Context) { + var req struct { + Username string `json:"username"` + } + if err := c.ShouldBindJSON(&req); err != nil || req.Username == "" { + c.JSON(http.StatusBadRequest, gin.H{"error": "username required"}) + return + } + + token, err := auth.NewAccessToken(req.Username, time.Hour) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + + c.JSON(200, gin.H{"access_token": token}) + }) + router.GET("/api/me", auth.Middleware(), func(c *gin.Context) { + username := c.GetString(auth.CtxUsernameKey) + c.JSON(200, gin.H{"username": username}) + }) router.GET("/ws/:room", func(c *gin.Context) { api.WebSocketHandler(hub, database)(c) }) diff --git a/deploy/shagram/compose.yaml b/deploy/shagram/compose.yaml index 2644abd..03731f5 100644 --- a/deploy/shagram/compose.yaml +++ b/deploy/shagram/compose.yaml @@ -6,6 +6,7 @@ services: container_name: shagram-app environment: - DATABASE_PATH=/app/data/shagram.db + - JWT_SECRET=${JWT_SECRET} volumes: - shagram_data:/app/data expose: diff --git a/go.mod b/go.mod index 797ebf1..a639755 100644 --- a/go.mod +++ b/go.mod @@ -14,6 +14,7 @@ require ( github.com/go-playground/validator/v10 v10.27.0 // indirect github.com/goccy/go-json v0.10.2 // indirect github.com/goccy/go-yaml v1.18.0 // indirect + github.com/golang-jwt/jwt/v5 v5.3.0 // indirect github.com/gorilla/websocket v1.5.3 // indirect github.com/json-iterator/go v1.1.12 // indirect github.com/klauspost/cpuid/v2 v2.3.0 // indirect diff --git a/go.sum b/go.sum index 73752d5..34c72ca 100644 --- a/go.sum +++ b/go.sum @@ -22,6 +22,8 @@ github.com/goccy/go-json v0.10.2 h1:CrxCmQqYDkv1z7lO7Wbh2HN93uovUHgrECaO5ZrCXAU= github.com/goccy/go-json v0.10.2/go.mod h1:6MelG93GURQebXPDq3khkgXZkazVtN9CRI+MGFi0w8I= github.com/goccy/go-yaml v1.18.0 h1:8W7wMFS12Pcas7KU+VVkaiCng+kG8QiFeFwzFb+rwuw= github.com/goccy/go-yaml v1.18.0/go.mod h1:XBurs7gK8ATbW4ZPGKgcbrY1Br56PdM69F7LkFRi1kA= +github.com/golang-jwt/jwt/v5 v5.3.0 h1:pv4AsKCKKZuqlgs5sUmn4x8UlGa0kEVt/puTpKx9vvo= +github.com/golang-jwt/jwt/v5 v5.3.0/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE= github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg= github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg= github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE= diff --git a/internal/auth/jwt.go b/internal/auth/jwt.go new file mode 100644 index 0000000..65575e8 --- /dev/null +++ b/internal/auth/jwt.go @@ -0,0 +1,64 @@ +package auth + +import ( + "errors" + "os" + "time" + + "github.com/golang-jwt/jwt/v5" +) + +type Claims struct { + Username string `json:"username"` + jwt.RegisteredClaims +} + +func secret() ([]byte, error) { + s := os.Getenv("JWT_SECRET") + if s == "" { + return nil, errors.New("JWT_SECRET is not set") + } + return []byte(s), nil +} + +func NewAccessToken(username string, ttl time.Duration) (string, error) { + key, err := secret() + if err != nil { + return "", err + } + now := time.Now() + claims := Claims{ + Username: username, + RegisteredClaims: jwt.RegisteredClaims{ + Subject: username, + IssuedAt: jwt.NewNumericDate(now), + ExpiresAt: jwt.NewNumericDate(now.Add(ttl)), + }, + } + + t := jwt.NewWithClaims(jwt.SigningMethodHS256, claims) + return t.SignedString(key) +} + +func ParseAccessToken(tokenString string) (*Claims, error) { + key, err := secret() + if err != nil { + return nil, err + } + + token, err := jwt.ParseWithClaims(tokenString, &Claims{}, func(token *jwt.Token) (any, error) { + return key, nil + }, + jwt.WithValidMethods([]string{jwt.SigningMethodHS256.Alg()}), + ) + if err != nil { + return nil, err + } + + claims, ok := token.Claims.(*Claims) + if !ok || !token.Valid { + return nil, errors.New("invalid token") + } + + return claims, nil +} diff --git a/internal/auth/middleware.go b/internal/auth/middleware.go new file mode 100644 index 0000000..d6bd55a --- /dev/null +++ b/internal/auth/middleware.go @@ -0,0 +1,41 @@ +package auth + +import ( + "net/http" + "strings" + + "github.com/gin-gonic/gin" +) + +const CtxUsernameKey = "username" + +func Middleware() gin.HandlerFunc { + return func(c *gin.Context) { + h := c.GetHeader("Authorization") + if h == "" { + c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "missing Authorization header"}) + return + } + + const prefix = "Bearer " + if !strings.HasPrefix(h, prefix) { + c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "expected Bearer token"}) + return + } + + raw := strings.TrimSpace(strings.TrimPrefix(h, prefix)) + if raw == "" { + c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "empty token"}) + return + } + + claims, err := ParseAccessToken(raw) + if err != nil { + c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "invalid token"}) + return + } + + c.Set(CtxUsernameKey, claims.Username) + c.Next() + } +}