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 }