176 lines
3.6 KiB
Go
176 lines
3.6 KiB
Go
package session
|
|
|
|
import (
|
|
"context"
|
|
"crypto/rand"
|
|
"encoding/hex"
|
|
"errors"
|
|
"fmt"
|
|
"log"
|
|
"net/http"
|
|
|
|
"omctf.ru/block-game-backend/codegen/ent"
|
|
"omctf.ru/block-game-backend/codegen/ent/setting"
|
|
"omctf.ru/block-game-backend/db"
|
|
|
|
"github.com/gorilla/sessions"
|
|
)
|
|
|
|
func generateRandomKey() (string, error) {
|
|
bytes := make([]byte, 32)
|
|
_, err := rand.Read(bytes)
|
|
if err != nil {
|
|
return "", fmt.Errorf("rand failed: %w", err)
|
|
}
|
|
return hex.EncodeToString(bytes), nil
|
|
}
|
|
|
|
func getOrCreateSessionSecret(ctx context.Context) (string, error) {
|
|
secret, err := db.Client.Setting.Query().
|
|
Where(setting.Key("session_secret")).
|
|
Only(ctx)
|
|
|
|
if err == nil {
|
|
return secret.Value, nil
|
|
}
|
|
|
|
if !ent.IsNotFound(err) {
|
|
return "", fmt.Errorf("unexpected db error: %w", err)
|
|
}
|
|
|
|
log.Println("session secret not found, generating a new one")
|
|
|
|
newSecret, err := generateRandomKey()
|
|
if err != nil {
|
|
return "", fmt.Errorf("generating session secret failed: %w", err)
|
|
}
|
|
|
|
_, err = db.Client.Setting.Create().
|
|
SetKey("session_secret").
|
|
SetValue(newSecret).
|
|
Save(ctx)
|
|
if err != nil {
|
|
return "", fmt.Errorf("saving session secret failed: %w", err)
|
|
}
|
|
|
|
return newSecret, nil
|
|
}
|
|
|
|
var Store *sessions.CookieStore
|
|
|
|
func Initialize() error {
|
|
ctx := context.Background()
|
|
|
|
sessionSecret, err := getOrCreateSessionSecret(ctx)
|
|
if err != nil {
|
|
return fmt.Errorf("getting session secret failed: %w", err)
|
|
}
|
|
|
|
Store = sessions.NewCookieStore([]byte(sessionSecret))
|
|
Store.Options = &sessions.Options{
|
|
Path: "/",
|
|
MaxAge: 86400 * 7, // 7 days
|
|
HttpOnly: true,
|
|
Secure: false, // we don't use https
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
type NotLoggedInError struct{}
|
|
|
|
func (e *NotLoggedInError) Error() string {
|
|
return "Not logged in"
|
|
}
|
|
|
|
func IsNotLoggedIn(err error) bool {
|
|
if err == nil {
|
|
return false
|
|
}
|
|
var e *NotLoggedInError
|
|
return errors.As(err, &e)
|
|
}
|
|
|
|
type InvalidSessionError struct {
|
|
Reason error
|
|
}
|
|
|
|
func (err *InvalidSessionError) Error() string {
|
|
return fmt.Sprintf("invalid session: %s", err.Reason)
|
|
}
|
|
|
|
func (err *InvalidSessionError) Unwrap() error {
|
|
return err.Reason
|
|
}
|
|
|
|
func IsInvalidSession(err error) bool {
|
|
if err == nil {
|
|
return false
|
|
}
|
|
var e *InvalidSessionError
|
|
return errors.As(err, &e)
|
|
}
|
|
|
|
func GetUserId(r *http.Request) (int, error) {
|
|
session, err := Store.Get(r, "auth")
|
|
if err != nil {
|
|
return -1, &InvalidSessionError{Reason: err}
|
|
}
|
|
|
|
id := session.Values["user_id"]
|
|
if id == nil {
|
|
return -1, &NotLoggedInError{}
|
|
}
|
|
|
|
idValue, ok := id.(int)
|
|
if !ok {
|
|
return -1, &InvalidSessionError{Reason: fmt.Errorf("id is not an int")}
|
|
}
|
|
|
|
return idValue, nil
|
|
}
|
|
|
|
func SetUserId(w http.ResponseWriter, r *http.Request, userId int) error {
|
|
session, err := Store.Get(r, "auth")
|
|
if err != nil {
|
|
return &InvalidSessionError{Reason: err}
|
|
}
|
|
|
|
session.Values["user_id"] = userId
|
|
if err := session.Save(r, w); err != nil {
|
|
return fmt.Errorf("saving session failed: %w", err)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
func GetSessionUser(r *http.Request) (*ent.User, error) {
|
|
ctx := r.Context()
|
|
|
|
userId, err := GetUserId(r)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("fetching the user id failed: %w", err)
|
|
}
|
|
|
|
user, err := db.Client.User.Get(ctx, userId)
|
|
if err != nil {
|
|
return nil, &InvalidSessionError{Reason: fmt.Errorf("fetching the user from the database failed: %w", err)}
|
|
}
|
|
|
|
return user, nil
|
|
}
|
|
|
|
func ClearSession(w http.ResponseWriter, r *http.Request) error {
|
|
session, err := Store.Get(r, "auth")
|
|
if err != nil {
|
|
return fmt.Errorf("fetching session failed: %w", err)
|
|
}
|
|
|
|
session.Options.MaxAge = -1
|
|
if err := session.Save(r, w); err != nil {
|
|
return fmt.Errorf("clearing session failed: %w", err)
|
|
}
|
|
|
|
return nil
|
|
}
|