refactor!: move go-chi/session into Gitea (#39504)

The `gitea.com/go-chi/session` package only exists for Gitea, so it
moves into `modules/session` to fix its bugs directly. Fixes the flake
in
https://github.com/go-gitea/gitea/actions/runs/36726154500/job/109923538400.

- Sessions are only written back when changed, so a read-only request
can't revert a concurrent change or restore a logged-out session, like
https://github.com/go-macaron/session/commit/ae808a4a4660c802965c834299ab08f167effd12
- The session cookie is only set once a session holds data
- Every backend refreshes the expiry on load and file sessions are
written atomically
- Also fix  https://github.com/go-gitea/gitea/issues/36176

## ⚠️ BREAKING ⚠️

* the `mysql`, `postgres`, `couchbase` and `memcache` session providers
are removed, use `file`, `db` or `redis` instead
* login-related cookies are renamed to `gitea_session` and
`gitea_remember`, if you'd like to use the old names, set `COOKIE_NAME`
and `COOKIE_REMEMBER_NAME` in app.ini

---------

Co-authored-by: wxiaoguang <wxiaoguang@gmail.com>
This commit is contained in:
silverwind
2026-10-02 21:08:14 +02:00
committed by GitHub
parent bea6fcaa84
commit cf89ecd887
34 changed files with 913 additions and 981 deletions
+13 -147
View File
@@ -5,172 +5,38 @@ package session
import (
"context"
"fmt"
"log"
"sync"
"gitea.dev/models/auth"
"gitea.dev/modules/log"
"gitea.dev/modules/timeutil"
"gitea.com/go-chi/session"
)
// DBStore represents a session store implementation based on the DB.
type DBStore struct {
sid string
lock sync.RWMutex
data map[any]any
type dbBackend struct {
maxLifetime int64
}
func dbContext() context.Context {
return context.Background()
}
// NewDBStore creates and returns a DB session store.
func NewDBStore(sid string, kv map[any]any) *DBStore {
return &DBStore{
sid: sid,
data: kv,
}
}
// Set sets value to given key in session.
func (s *DBStore) Set(key, val any) error {
s.lock.Lock()
defer s.lock.Unlock()
s.data[key] = val
return nil
}
// Get gets value by given key in session.
func (s *DBStore) Get(key any) any {
s.lock.RLock()
defer s.lock.RUnlock()
return s.data[key]
}
// Delete delete a key from session.
func (s *DBStore) Delete(key any) error {
s.lock.Lock()
defer s.lock.Unlock()
delete(s.data, key)
return nil
}
// ID returns current session ID.
func (s *DBStore) ID() string {
return s.sid
}
// Release releases resource and save data to provider.
func (s *DBStore) Release() error {
// Skip encoding if the data is empty
if len(s.data) == 0 {
return nil
}
data, err := session.EncodeGob(s.data)
if err != nil {
return err
}
return auth.UpdateSession(dbContext(), s.sid, data)
}
// Flush deletes all session data.
func (s *DBStore) Flush() error {
s.lock.Lock()
defer s.lock.Unlock()
s.data = make(map[any]any)
return nil
}
// DBProvider represents a DB session provider implementation.
type DBProvider struct {
maxLifetime int64
}
// Init initializes DB session provider.
// connStr: username:password@protocol(address)/dbname?param=value
func (p *DBProvider) Init(maxLifetime int64, connStr string) error {
p.maxLifetime = maxLifetime
return nil
}
// Read returns raw session store by session ID.
func (p *DBProvider) Read(sid string) (session.RawStore, error) {
s, err := auth.ReadSession(dbContext(), sid)
if err != nil {
func (b *dbBackend) load(sid string) ([]byte, error) {
sess, exist, err := auth.GetSession(dbContext(), sid)
if err != nil || !exist || sess.LastAccessTime.Add(b.maxLifetime) <= timeutil.TimeStampNow() {
return nil, err
}
var kv map[any]any
if len(s.Data) == 0 || s.Expiry.Add(p.maxLifetime) <= timeutil.TimeStampNow() {
kv = make(map[any]any)
} else {
kv, err = session.DecodeGob(s.Data)
if err != nil {
return nil, err
}
}
return NewDBStore(sid, kv), nil
return sess.Data, auth.UpdateSessionLastAccessTime(dbContext(), sid)
}
// Exist returns true if session with given ID exists.
func (p *DBProvider) Exist(sid string) (bool, error) {
has, err := auth.ExistSession(dbContext(), sid)
if err != nil {
return false, fmt.Errorf("session/DB: error checking existence: %w", err)
}
return has, nil
func (b *dbBackend) save(sid string, data []byte, create bool) error {
return auth.UpdateSession(dbContext(), sid, data, create)
}
// Destroy deletes a session by session ID.
func (p *DBProvider) Destroy(sid string) error {
func (b *dbBackend) destroy(sid string) error {
return auth.DestroySession(dbContext(), sid)
}
// Regenerate regenerates a session store from old session ID to new one.
func (p *DBProvider) Regenerate(oldsid, sid string) (_ session.RawStore, err error) {
s, err := auth.RegenerateSession(dbContext(), oldsid, sid)
if err != nil {
return nil, err
}
var kv map[any]any
if len(s.Data) == 0 || s.Expiry.Add(p.maxLifetime) <= timeutil.TimeStampNow() {
kv = make(map[any]any)
} else {
kv, err = session.DecodeGob(s.Data)
if err != nil {
return nil, err
}
}
return NewDBStore(sid, kv), nil
}
// Count counts and returns number of sessions.
func (p *DBProvider) Count() (int, error) {
total, err := auth.CountSessions(dbContext())
if err != nil {
return 0, fmt.Errorf("session/DB: error counting records: %w", err)
}
return int(total), nil
}
// GC calls GC to clean expired sessions.
func (p *DBProvider) GC() {
if err := auth.CleanupSessions(dbContext(), p.maxLifetime); err != nil {
log.Printf("session/DB: error garbage collecting: %v", err)
func (b *dbBackend) gc() {
if err := auth.CleanupSessions(dbContext(), b.maxLifetime); err != nil {
log.Error("Unable to garbage collect sessions: %v", err)
}
}
func init() {
session.Register("db", &DBProvider{})
}
+115
View File
@@ -0,0 +1,115 @@
// Copyright 2013 Beego Authors
// Copyright 2014 The Macaron Authors
// Copyright 2026 The Gitea Authors. All rights reserved.
// SPDX-License-Identifier: Apache-2.0
package session
import (
"errors"
"fmt"
"io/fs"
"os"
"path/filepath"
"sync"
"time"
"gitea.dev/modules/log"
"gitea.dev/modules/util"
)
type fileBackend struct {
lock sync.RWMutex // exclusive for removals, so they never interleave with a load or save
rootPath string
maxLifetime time.Duration
}
func newFileBackend(rootPath string, maxLifetime int64) *fileBackend {
return &fileBackend{rootPath: filepath.Clean(rootPath), maxLifetime: time.Duration(maxLifetime) * time.Second}
}
func (b *fileBackend) filepath(sid string) string {
return filepath.Join(b.rootPath, sid[0:1], sid[1:2], sid)
}
func ignoreNotExist(err error) error {
if errors.Is(err, fs.ErrNotExist) {
return nil
}
return err
}
func (b *fileBackend) load(sid string) ([]byte, error) {
b.lock.RLock()
defer b.lock.RUnlock()
filename := b.filepath(sid)
stat, err := os.Lstat(filename)
if err != nil {
return nil, ignoreNotExist(err)
}
if !stat.Mode().IsRegular() {
return nil, fmt.Errorf("session file %s is not a regular file", filename)
}
if time.Since(stat.ModTime()) > b.maxLifetime {
return nil, nil
}
data, err := os.ReadFile(filename)
if err == nil {
now := time.Now()
err = os.Chtimes(filename, now, now)
}
return data, ignoreNotExist(err)
}
func (b *fileBackend) save(sid string, data []byte, create bool) error {
b.lock.RLock()
defer b.lock.RUnlock()
filename := b.filepath(sid)
if create {
if err := os.MkdirAll(filepath.Dir(filename), 0o700); err != nil {
return err
}
} else if _, err := os.Lstat(filename); err != nil {
return ignoreNotExist(err)
}
tmpFile, err := os.CreateTemp(filepath.Dir(filename), sid+".*.tmp")
if err != nil {
return err
}
_, err = tmpFile.Write(data)
if err = errors.Join(err, tmpFile.Close()); err == nil {
err = util.RenameWithRetry(tmpFile.Name(), filename)
}
if err != nil {
_ = os.Remove(tmpFile.Name())
}
return err
}
func (b *fileBackend) destroy(sid string) error {
b.lock.Lock()
defer b.lock.Unlock()
return ignoreNotExist(os.Remove(b.filepath(sid)))
}
func (b *fileBackend) expired(path string) bool {
info, err := os.Lstat(path)
return err == nil && time.Since(info.ModTime()) > b.maxLifetime
}
func (b *fileBackend) gc() {
err := filepath.WalkDir(b.rootPath, func(path string, entry fs.DirEntry, err error) error {
if err != nil || entry.IsDir() || !b.expired(path) {
return ignoreNotExist(err)
}
b.lock.Lock()
defer b.lock.Unlock()
if b.expired(path) { // a concurrent load may have refreshed it
err = os.Remove(path)
}
return ignoreNotExist(err)
})
if err != nil {
log.Error("Unable to garbage collect session files: %v", err)
}
}
+16
View File
@@ -0,0 +1,16 @@
// Copyright 2026 The Gitea Authors. All rights reserved.
// SPDX-License-Identifier: MIT
package session
import (
"testing"
"gitea.dev/models/unittest"
_ "gitea.dev/models"
)
func TestMain(m *testing.M) {
unittest.MainTest(m, &unittest.TestOptions{FixtureFiles: []string{}})
}
+84 -36
View File
@@ -4,65 +4,113 @@
package session
import (
"bytes"
"encoding/gob"
"maps"
"net/http"
"sync"
"time"
"gitea.com/go-chi/session"
"gitea.dev/modules/util"
)
type mockMemRawStore struct {
s *session.MemStore
type memoryBackend struct {
lock sync.Mutex
maxLifetime time.Duration
sessions map[string]memorySession
}
var _ session.RawStore = (*mockMemRawStore)(nil)
type memorySession struct {
data []byte
accessed time.Time
}
func (m *mockMemRawStore) Set(k, v any) error {
// We need to use gob to encode the value, to make it have the same behavior as other stores and catch abuses.
// Because gob needs to "Register" the type before it can encode it, and it's unable to decode a struct to "any" so use a map to help to decode the value.
var buf bytes.Buffer
if err := gob.NewEncoder(&buf).Encode(map[string]any{"v": v}); err != nil {
return err
func newMemoryBackend(maxLifetime int64) *memoryBackend {
return &memoryBackend{maxLifetime: time.Duration(maxLifetime) * time.Second, sessions: map[string]memorySession{}}
}
func (b *memoryBackend) expired(sess memorySession) bool {
return time.Since(sess.accessed) > b.maxLifetime
}
func (b *memoryBackend) load(sid string) ([]byte, error) {
b.lock.Lock()
defer b.lock.Unlock()
sess, ok := b.sessions[sid]
if !ok || b.expired(sess) {
return nil, nil
}
return m.s.Set(k, buf.Bytes())
sess.accessed = time.Now()
b.sessions[sid] = sess
return sess.data, nil
}
func (m *mockMemRawStore) Get(k any) (ret any) {
v, ok := m.s.Get(k).([]byte)
if !ok {
return nil
func (b *memoryBackend) save(sid string, data []byte, create bool) error {
b.lock.Lock()
defer b.lock.Unlock()
if _, exists := b.sessions[sid]; exists || create {
b.sessions[sid] = memorySession{data: data, accessed: time.Now()}
}
var w map[string]any
_ = gob.NewDecoder(bytes.NewBuffer(v)).Decode(&w)
return w["v"]
return nil
}
func (m *mockMemRawStore) Delete(k any) error {
return m.s.Delete(k)
func (b *memoryBackend) destroy(sid string) error {
b.lock.Lock()
defer b.lock.Unlock()
delete(b.sessions, sid)
return nil
}
func (m *mockMemRawStore) ID() string {
return m.s.ID()
}
func (m *mockMemRawStore) Release() error {
return m.s.Release()
}
func (m *mockMemRawStore) Flush() error {
return m.s.Flush()
func (b *memoryBackend) gc() {
b.lock.Lock()
defer b.lock.Unlock()
maps.DeleteFunc(b.sessions, func(_ string, sess memorySession) bool { return b.expired(sess) })
}
type mockMemStore struct {
*mockMemRawStore
sid string
data map[any][]byte
}
var _ Store = (*mockMemStore)(nil)
func (m mockMemStore) Destroy(writer http.ResponseWriter, request *http.Request) error {
// NewMockMemStore returns a store encoding each value like the real backends do, to catch values that can't be stored
func NewMockMemStore(sid string) Store {
return &mockMemStore{sid: sid, data: map[any][]byte{}}
}
func (m *mockMemStore) Set(key, value any) error {
encoded, err := util.PackData(map[any]any{key: value})
if err == nil {
m.data[key] = encoded
}
return err
}
func (m *mockMemStore) Get(key any) any {
var decoded map[any]any
_ = util.UnpackData(m.data[key], &decoded)
return decoded[key]
}
func (m *mockMemStore) Delete(key any) error {
delete(m.data, key)
return nil
}
func NewMockMemStore(sid string) Store {
return &mockMemStore{&mockMemRawStore{session.NewMemStore(sid)}}
func (m *mockMemStore) ID() string {
return m.sid
}
func (m *mockMemStore) Release() error {
return nil
}
func (m *mockMemStore) Flush() error {
clear(m.data)
return nil
}
func (m *mockMemStore) Destroy(http.ResponseWriter, *http.Request) error {
return nil
}
func (m *mockMemStore) Regenerate(http.ResponseWriter, *http.Request) {}
+31 -187
View File
@@ -6,214 +6,58 @@
package session
import (
"fmt"
"sync"
"errors"
"time"
"gitea.dev/modules/graceful"
"gitea.dev/modules/nosql"
"gitea.com/go-chi/session"
"github.com/redis/go-redis/v9"
)
// RedisStore represents a redis session store implementation.
type RedisStore struct {
c redis.UniversalClient
prefix, sid string
duration time.Duration
lock sync.RWMutex
data map[any]any
type redisBackend struct {
client redis.UniversalClient
prefix string
maxLifetime time.Duration
}
// NewRedisStore creates and returns a redis session store.
func NewRedisStore(c redis.UniversalClient, prefix, sid string, dur time.Duration, kv map[any]any) *RedisStore {
return &RedisStore{
c: c,
prefix: prefix,
sid: sid,
duration: dur,
data: kv,
func newRedisBackend(config string, maxLifetime int64) (*redisBackend, error) {
uri := nosql.ToRedisURI(config)
b := &redisBackend{
client: nosql.GetManager().GetRedisClient(uri.String()),
prefix: uri.Query().Get("prefix"),
maxLifetime: time.Duration(maxLifetime) * time.Second,
}
return b, b.client.Ping(graceful.GetManager().ShutdownContext()).Err()
}
// Set sets value to given key in session.
func (s *RedisStore) Set(key, val any) error {
s.lock.Lock()
defer s.lock.Unlock()
s.data[key] = val
return nil
}
// Get gets value by given key in session.
func (s *RedisStore) Get(key any) any {
s.lock.RLock()
defer s.lock.RUnlock()
return s.data[key]
}
// Delete delete a key from session.
func (s *RedisStore) Delete(key any) error {
s.lock.Lock()
defer s.lock.Unlock()
delete(s.data, key)
return nil
}
// ID returns current session ID.
func (s *RedisStore) ID() string {
return s.sid
}
// Release releases resource and save data to provider.
func (s *RedisStore) Release() error {
// Skip encoding if the data is empty
if len(s.data) == 0 {
func (b *redisBackend) load(sid string) ([]byte, error) {
ctx := graceful.GetManager().HammerContext()
var get *redis.StringCmd
_, err := b.client.Pipelined(ctx, func(pipe redis.Pipeliner) error {
get = pipe.Get(ctx, b.prefix+sid)
pipe.Expire(ctx, b.prefix+sid, b.maxLifetime)
return nil
})
if errors.Is(err, redis.Nil) {
return nil, nil
}
data, err := session.EncodeGob(s.data)
if err != nil {
return err
}
return s.c.Set(graceful.GetManager().HammerContext(), s.prefix+s.sid, string(data), s.duration).Err()
}
// Flush deletes all session data.
func (s *RedisStore) Flush() error {
s.lock.Lock()
defer s.lock.Unlock()
s.data = make(map[any]any)
return nil
}
// RedisProvider represents a redis session provider implementation.
type RedisProvider struct {
c redis.UniversalClient
duration time.Duration
prefix string
}
// Init initializes redis session provider.
// configs: network=tcp,addr=:6379,password=macaron,db=0,pool_size=100,idle_timeout=180,prefix=session;
func (p *RedisProvider) Init(maxlifetime int64, configs string) (err error) {
p.duration, err = time.ParseDuration(fmt.Sprintf("%ds", maxlifetime))
if err != nil {
return err
}
uri := nosql.ToRedisURI(configs)
for k, v := range uri.Query() {
switch k {
case "prefix":
p.prefix = v[0]
}
}
p.c = nosql.GetManager().GetRedisClient(uri.String())
return p.c.Ping(graceful.GetManager().ShutdownContext()).Err()
}
// Read returns raw session store by session ID.
func (p *RedisProvider) Read(sid string) (session.RawStore, error) {
psid := p.prefix + sid
if exist, err := p.Exist(sid); err == nil && !exist {
if err := p.c.Set(graceful.GetManager().HammerContext(), psid, "", p.duration).Err(); err != nil {
return nil, err
}
} else if err != nil {
return nil, err
}
var kv map[any]any
kvs, err := p.c.Get(graceful.GetManager().HammerContext(), psid).Result()
if err != nil {
return nil, err
}
if len(kvs) == 0 {
kv = make(map[any]any)
} else {
kv, err = session.DecodeGob([]byte(kvs))
if err != nil {
return nil, err
}
}
return NewRedisStore(p.c, p.prefix, sid, p.duration, kv), nil
return get.Bytes()
}
// Exist returns true if session with given ID exists.
func (p *RedisProvider) Exist(sid string) (bool, error) {
v, err := p.c.Exists(graceful.GetManager().HammerContext(), p.prefix+sid).Result()
return err == nil && v == 1, err
func (b *redisBackend) save(sid string, data []byte, create bool) error {
ctx := graceful.GetManager().HammerContext()
if create {
return b.client.Set(ctx, b.prefix+sid, data, b.maxLifetime).Err()
}
return b.client.SetXX(ctx, b.prefix+sid, data, b.maxLifetime).Err()
}
// Destroy deletes a session by session ID.
func (p *RedisProvider) Destroy(sid string) error {
return p.c.Del(graceful.GetManager().HammerContext(), p.prefix+sid).Err()
func (b *redisBackend) destroy(sid string) error {
return b.client.Del(graceful.GetManager().HammerContext(), b.prefix+sid).Err()
}
// Regenerate regenerates a session store from old session ID to new one.
func (p *RedisProvider) Regenerate(oldsid, sid string) (_ session.RawStore, err error) {
poldsid := p.prefix + oldsid
psid := p.prefix + sid
if exist, err := p.Exist(sid); err != nil {
return nil, err
} else if exist {
return nil, fmt.Errorf("new sid '%s' already exists", sid)
}
if exist, err := p.Exist(oldsid); err == nil && !exist {
// Make a fake old session.
if err := p.c.Set(graceful.GetManager().HammerContext(), poldsid, "", p.duration).Err(); err != nil {
return nil, err
}
} else if err != nil {
return nil, err
}
// do not use Rename here, because the old sid and new sid may be in different redis cluster slot.
kvs, err := p.c.Get(graceful.GetManager().HammerContext(), poldsid).Result()
if err != nil {
return nil, err
}
if err = p.c.Del(graceful.GetManager().HammerContext(), poldsid).Err(); err != nil {
return nil, err
}
if err = p.c.Set(graceful.GetManager().HammerContext(), psid, kvs, p.duration).Err(); err != nil {
return nil, err
}
var kv map[any]any
if len(kvs) == 0 {
kv = make(map[any]any)
} else {
kv, err = session.DecodeGob([]byte(kvs))
if err != nil {
return nil, err
}
}
return NewRedisStore(p.c, p.prefix, sid, p.duration, kv), nil
}
// Count counts and returns number of sessions.
func (p *RedisProvider) Count() (int, error) {
size, err := p.c.DBSize(graceful.GetManager().HammerContext()).Result()
return int(size), err
}
// GC calls GC to clean expired sessions.
func (*RedisProvider) GC() {}
func init() {
session.Register("redis", &RedisProvider{})
}
func (*redisBackend) gc() {}
+146
View File
@@ -0,0 +1,146 @@
// Copyright 2013 Beego Authors
// Copyright 2014 The Macaron Authors
// Copyright 2026 The Gitea Authors. All rights reserved.
// SPDX-License-Identifier: Apache-2.0
package session
import (
"context"
"encoding/gob"
"fmt"
"net/http"
"time"
"gitea.dev/modules/graceful"
"gitea.dev/modules/log"
"gitea.dev/modules/setting"
"gitea.dev/modules/util"
)
type backend interface {
load(sid string) ([]byte, error) // returns nil for missing or expired sessions, refreshes the expiry of others
save(sid string, data []byte, create bool) error // without create, only an existing session is updated
destroy(sid string) error
gc()
}
// CHI-SESSION-GOB-REGISTER: packages must gob.Register the types they store at startup, so data stored before a restart still decodes
func init() {
gob.Register([]any{})
gob.Register(map[int]any{})
gob.Register(map[string]any{})
gob.Register(map[any]any{})
gob.Register(map[string]string{})
gob.Register(map[int]string{})
gob.Register(map[int]int{})
gob.Register(map[int]int64{})
}
func newBackend(provider, config string, maxLifetime int64) (backend, error) {
switch provider {
case "memory":
return newMemoryBackend(maxLifetime), nil
case "file":
return newFileBackend(config, maxLifetime), nil
case "redis":
return newRedisBackend(config, maxLifetime)
case "db":
return &dbBackend{maxLifetime: maxLifetime}, nil
}
return nil, fmt.Errorf(`unsupported [session] PROVIDER %q, supported are "memory", "file", "redis" and "db", use "db" or "redis" to replace the removed "mysql", "postgres", "couchbase" and "memcache" providers`, provider)
}
func Sessioner() (func(next http.Handler) http.Handler, error) {
backend, err := newBackend(setting.SessionConfig.Provider, setting.SessionConfig.ProviderConfig, setting.SessionConfig.Maxlifetime)
if err != nil {
return nil, err
}
go runGC(graceful.GetManager().ShutdownContext(), backend, time.Duration(setting.SessionConfig.Gclifetime)*time.Second)
return func(next http.Handler) http.Handler {
return sessionHandler(backend, next)
}, nil
}
func sessionHandler(backend backend, next http.Handler) http.Handler {
return http.HandlerFunc(func(resp http.ResponseWriter, req *http.Request) {
sess, err := startSession(backend, resp, req)
if err != nil {
log.Error("Unable to start session: %v", err)
resp.WriteHeader(http.StatusInternalServerError)
return
}
next.ServeHTTP(resp, req.WithContext(context.WithValue(req.Context(), ContextKey, sess)))
if err := sess.Release(); err != nil {
log.Error("Unable to release session: %v", err)
}
})
}
func startSession(backend backend, resp http.ResponseWriter, req *http.Request) (*store, error) {
sess := &store{backend: backend, resp: resp, data: map[any]any{}}
cookie, err := req.Cookie(setting.SessionConfig.CookieName)
if err != nil || !isValidSessionID(cookie.Value) {
sess.sid = newSessionID()
return sess, nil
}
sess.sid, sess.cookieSID = cookie.Value, cookie.Value
encoded, err := backend.load(sess.sid)
if err != nil {
return nil, err
}
if len(encoded) == 0 {
return sess, nil
}
var data map[any]any
if err := util.UnpackData(encoded, &data); err != nil {
log.Error("Unable to decode session data, starting with an empty session: %v", err)
sess.stored, sess.changed = true, true
} else if len(data) > 0 {
sess.data, sess.stored = data, true
}
return sess, nil
}
func newSessionID() string {
// lower case (in case the file system is case-insensitive) and length=16 (db session's primary key is fixed size 16)
// the entropy is about 36^16 > 80 bits
return util.FastCryptoRandomString(16, "abcdefghijklmnopqrstuvwxyz0123456789")
}
func isValidSessionID(sid string) bool {
if len(sid) != 16 { // db session has a primary key with fixed size 16
return false
}
for i := range len(sid) {
c := sid[i]
valid := (c >= '0' && c <= '9') || (c >= 'a' && c <= 'z')
if !valid {
return false
}
}
return true
}
func newCookie(value string) *http.Cookie {
return &http.Cookie{
Name: setting.SessionConfig.CookieName,
Value: value,
Path: util.IfZero(setting.SessionConfig.CookiePath, "/"),
Domain: setting.SessionConfig.Domain,
Secure: setting.SessionConfig.Secure,
HttpOnly: true,
SameSite: setting.SessionConfig.SameSite,
}
}
func runGC(ctx context.Context, backend backend, interval time.Duration) {
for {
backend.gc()
select {
case <-ctx.Done():
return
case <-time.After(interval):
}
}
}
+264
View File
@@ -0,0 +1,264 @@
// Copyright 2026 The Gitea Authors. All rights reserved.
// SPDX-License-Identifier: MIT
package session
import (
"context"
"errors"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"testing"
"time"
auth_model "gitea.dev/models/auth"
"gitea.dev/models/db"
"gitea.dev/modules/setting"
"gitea.dev/modules/test"
"gitea.dev/modules/timeutil"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
type handleFunc func(resp http.ResponseWriter, req *http.Request, sess Store)
type failingBackend struct {
backend
failDestroy bool
}
func (b *failingBackend) destroy(sid string) error {
if b.failDestroy {
return errors.New("destroy failed")
}
return b.backend.destroy(sid)
}
func TestSession(t *testing.T) {
defer test.MockVariableValue(&setting.SessionConfig.CookiePath, "/sub")()
defer test.MockVariableValue(&setting.SessionConfig.Secure, true)()
_, err := newBackend("mysql", "", 3600)
assert.ErrorContains(t, err, `use "db" or "redis"`)
t.Run("GCStopsOnShutdown", func(t *testing.T) {
ctx, cancel := context.WithCancel(t.Context())
cancel()
runGC(ctx, newMemoryBackend(3600), time.Hour)
})
cases := []struct {
name string
newBackend func(t *testing.T) backend
}{
{
name: "memory",
newBackend: func(*testing.T) backend { return newMemoryBackend(3600) },
},
{
name: "file",
newBackend: func(t *testing.T) backend { return newFileBackend(t.TempDir(), 3600) },
},
{
name: "db",
newBackend: func(*testing.T) backend { return &dbBackend{maxLifetime: 3600} },
},
{
name: "redis",
newBackend: func(t *testing.T) backend {
backend, err := newRedisBackend(test.PrepareTestRedis(t)+"?prefix=gitea-test-session-", 3600)
require.NoError(t, err)
return backend
},
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
backend := &failingBackend{backend: tc.newBackend(t)}
serve := func(sid string, handle handleFunc) *httptest.ResponseRecorder {
req := httptest.NewRequest(http.MethodGet, "/", nil)
if sid != "" {
req.AddCookie(&http.Cookie{Name: setting.SessionConfig.CookieName, Value: sid})
}
resp := httptest.NewRecorder()
sessionHandler(backend, http.HandlerFunc(func(resp http.ResponseWriter, req *http.Request) {
handle(resp, req, GetContextSession(req))
})).ServeHTTP(resp, req)
return resp
}
create := func(t *testing.T) string {
cookies := serve("", func(_ http.ResponseWriter, _ *http.Request, sess Store) {
require.NoError(t, sess.Set("key", "value"))
}).Result().Cookies()
require.Len(t, cookies, 1)
return cookies[0].Value
}
get := func(sid string, key any) (value any) {
serve(sid, func(_ http.ResponseWriter, _ *http.Request, sess Store) { value = sess.Get(key) })
return value
}
t.Run("CookieOnlyOnceSessionHoldsData", func(t *testing.T) {
assert.Empty(t, serve("", func(http.ResponseWriter, *http.Request, Store) {}).Result().Cookies())
resp := serve("../../etc/passwd", func(resp http.ResponseWriter, req *http.Request, sess Store) {
require.NoError(t, sess.Set("key", "value"))
require.NoError(t, sess.Set("other", 1))
http.Redirect(resp, req, "https://example.com/", http.StatusSeeOther)
})
cookies := resp.Result().Cookies()
require.Len(t, cookies, 1)
sid := cookies[0].Value
assert.True(t, isValidSessionID(sid))
assert.Equal(t, "/sub", cookies[0].Path)
assert.True(t, cookies[0].HttpOnly)
assert.True(t, cookies[0].Secure)
assert.Equal(t, http.SameSiteLaxMode, cookies[0].SameSite)
switch sessionBackend := backend.backend.(type) {
case *dbBackend:
now := timeutil.TimeStampNow()
_, err := db.GetEngine(t.Context()).ID(sid).Cols(auth_model.DbSessionLastAccessTime).Update(&auth_model.Session{LastAccessTime: now - 60})
require.NoError(t, err)
_, err = sessionBackend.load(sid)
require.NoError(t, err)
sess, _, err := auth_model.GetSession(t.Context(), sid)
require.NoError(t, err)
assert.GreaterOrEqual(t, sess.LastAccessTime, now)
case *redisBackend:
require.NoError(t, sessionBackend.client.Expire(t.Context(), "gitea-test-session-"+sid, time.Minute).Err())
_, err := sessionBackend.load(sid)
require.NoError(t, err)
ttl, err := sessionBackend.client.TTL(t.Context(), "gitea-test-session-"+sid).Result()
require.NoError(t, err)
assert.Greater(t, ttl, time.Minute)
}
resp = serve(sid, func(_ http.ResponseWriter, _ *http.Request, sess Store) {
assert.Equal(t, sid, sess.ID())
assert.Equal(t, "value", sess.Get("key"))
require.NoError(t, sess.Set("key", "changed"))
})
assert.Empty(t, resp.Result().Cookies())
})
t.Run("RegenerateMovesDataToNewID", func(t *testing.T) {
oldSID := create(t)
var newSID string
resp := serve(oldSID, func(resp http.ResponseWriter, req *http.Request, sess Store) {
sess.Regenerate(resp, req)
newSID = sess.ID()
})
cookies := resp.Result().Cookies()
require.Len(t, cookies, 1)
assert.Equal(t, newSID, cookies[0].Value)
assert.NotEqual(t, oldSID, newSID)
assert.Nil(t, get(oldSID, "key"))
assert.Equal(t, "value", get(newSID, "key"))
resp = serve("malformed", func(resp http.ResponseWriter, req *http.Request, sess Store) { sess.Regenerate(resp, req) })
assert.Empty(t, resp.Result().Cookies())
})
t.Run("DestroyIsNotUndoneByRelease", func(t *testing.T) {
sid := create(t)
var resp *httptest.ResponseRecorder
serve(sid, func(_ http.ResponseWriter, _ *http.Request, concurrent Store) {
resp = serve(sid, func(resp http.ResponseWriter, req *http.Request, sess Store) {
require.NoError(t, sess.Flush())
require.NoError(t, sess.Destroy(resp, req))
})
require.NoError(t, concurrent.Set("key", "changed"))
})
cookies := resp.Result().Cookies()
require.Len(t, cookies, 1)
assert.Equal(t, -1, cookies[0].MaxAge)
assert.Equal(t, "/sub", cookies[0].Path)
assert.True(t, cookies[0].Secure)
encoded, err := backend.load(sid)
require.NoError(t, err)
assert.Nil(t, encoded)
})
t.Run("FailedDestroyIsRetriedByRelease", func(t *testing.T) {
sid := create(t)
serve(sid, func(resp http.ResponseWriter, req *http.Request, sess Store) {
backend.failDestroy = true
require.Error(t, sess.Destroy(resp, req))
backend.failDestroy = false
assert.Nil(t, sess.Get("key"))
})
encoded, err := backend.load(sid)
require.NoError(t, err)
assert.Nil(t, encoded)
})
t.Run("EmptiedSessionIsPersisted", func(t *testing.T) {
for _, empty := range []func(Store) error{Store.Flush, func(sess Store) error { return sess.Delete("key") }} {
sid := create(t)
serve(sid, func(_ http.ResponseWriter, _ *http.Request, sess Store) { require.NoError(t, empty(sess)) })
assert.Nil(t, get(sid, "key"))
}
})
t.Run("UnchangedReleaseKeepsConcurrentChanges", func(t *testing.T) {
sid := create(t)
serve(sid, func(_ http.ResponseWriter, _ *http.Request, reader Store) {
serve(sid, func(_ http.ResponseWriter, _ *http.Request, writer Store) {
require.NoError(t, writer.Set("key", "changed"))
})
require.NoError(t, reader.Delete("missing"))
})
assert.Equal(t, "changed", get(sid, "key"))
})
t.Run("ConcurrentNewSessionWritesThrough", func(t *testing.T) {
sid := newSessionID()
serve(sid, func(_ http.ResponseWriter, _ *http.Request, first Store) {
serve(sid, func(_ http.ResponseWriter, _ *http.Request, second Store) {
require.NoError(t, second.Set("second", 2))
})
require.NoError(t, first.Set("first", 1))
})
assert.Equal(t, 1, get(sid, "first"))
})
t.Run("UndecodableDataReadsAsEmptySession", func(t *testing.T) {
sid := newSessionID()
require.NoError(t, backend.save(sid, []byte("undecodable"), true))
resp := serve(sid, func(_ http.ResponseWriter, _ *http.Request, sess Store) { assert.Nil(t, sess.Get("key")) })
assert.Equal(t, http.StatusOK, resp.Code)
encoded, err := backend.load(sid)
require.NoError(t, err)
assert.Nil(t, encoded)
})
})
}
}
func TestFileBackendWritesAtomicallyAndExpiresByModTime(t *testing.T) {
root := t.TempDir()
backend := newFileBackend(root, 3600)
sid := newSessionID()
filename := filepath.Join(root, sid[0:1], sid[1:2], sid)
require.NoError(t, backend.save(sid, []byte("first"), true))
require.NoError(t, backend.save(sid, []byte("second"), false))
entries, err := os.ReadDir(filepath.Dir(filename))
require.NoError(t, err)
require.Len(t, entries, 1)
assert.Equal(t, sid, entries[0].Name())
encoded, err := backend.load(sid)
require.NoError(t, err)
assert.Equal(t, "second", string(encoded))
expired := time.Now().Add(-2 * time.Hour)
require.NoError(t, os.Chtimes(filename, expired, expired))
encoded, err = backend.load(sid)
require.NoError(t, err)
assert.Nil(t, encoded)
backend.gc()
assert.NoFileExists(t, filename)
}
+116 -20
View File
@@ -5,43 +5,139 @@ package session
import (
"net/http"
"sync"
"gitea.dev/modules/setting"
"gitea.com/go-chi/session"
"gitea.dev/modules/log"
"gitea.dev/modules/util"
)
type RawStore = session.RawStore
type Store interface {
RawStore
Set(key, value any) error
Get(key any) any
Delete(key any) error
ID() string
Release() error
Flush() error
Destroy(http.ResponseWriter, *http.Request) error
Regenerate(http.ResponseWriter, *http.Request)
}
type mockStoreContextKeyStruct struct{}
type store struct {
backend backend
resp http.ResponseWriter
lock sync.RWMutex
sid string
cookieSID string // the session ID the client holds
data map[any]any
stored bool // the backend holds data for sid
changed bool
}
var MockStoreContextKey = mockStoreContextKeyStruct{}
type contextKeyStruct struct{}
// RegenerateSession regenerates the underlying session and returns the new store
func RegenerateSession(resp http.ResponseWriter, req *http.Request) (Store, error) {
var ContextKey = contextKeyStruct{}
func (s *store) Set(key, value any) error {
s.lock.Lock()
defer s.lock.Unlock()
s.data[key] = value
s.changed = true
s.sendCookie(s.resp)
return nil
}
func (s *store) Get(key any) any {
s.lock.RLock()
defer s.lock.RUnlock()
return s.data[key]
}
func (s *store) Delete(key any) error {
s.lock.Lock()
defer s.lock.Unlock()
if _, ok := s.data[key]; ok {
delete(s.data, key)
s.changed = true
}
return nil
}
func (s *store) ID() string {
s.lock.RLock()
defer s.lock.RUnlock()
return s.sid
}
func (s *store) Flush() error {
s.lock.Lock()
defer s.lock.Unlock()
clear(s.data)
s.changed = true
return nil
}
func (s *store) Release() error {
s.lock.Lock()
defer s.lock.Unlock()
var err error
switch {
case !s.changed:
return nil
case len(s.data) > 0:
var data []byte
if data, err = util.PackData(s.data); err == nil {
err = s.backend.save(s.sid, data, !s.stored)
}
case s.stored:
err = s.backend.destroy(s.sid)
}
if err == nil {
s.changed, s.stored = false, len(s.data) > 0
}
return err
}
func (s *store) Destroy(resp http.ResponseWriter, _ *http.Request) error {
s.lock.Lock()
defer s.lock.Unlock()
err := s.backend.destroy(s.sid)
cookie := newCookie("")
cookie.MaxAge = -1
http.SetCookie(resp, cookie)
if err == nil {
s.sid, s.stored = newSessionID(), false
}
s.cookieSID, s.changed = "", err != nil // Release retries a failed destroy
clear(s.data)
return err
}
// Regenerate moves the session data to a new session ID, so an ID known before sign-in is never authenticated
func (s *store) Regenerate(resp http.ResponseWriter, req *http.Request) {
for _, f := range BeforeRegenerateSession {
f(resp, req)
}
if setting.IsInTesting {
if store, ok := req.Context().Value(MockStoreContextKey).(Store); ok {
return store, nil
s.lock.Lock()
defer s.lock.Unlock()
if s.stored {
if err := s.backend.destroy(s.sid); err != nil {
log.Error("Unable to destroy the regenerated session: %v", err)
}
}
return session.RegenerateSession(resp, req)
s.sid, s.stored, s.changed = newSessionID(), false, true
s.sendCookie(resp)
}
func (s *store) sendCookie(resp http.ResponseWriter) {
if s.cookieSID != s.sid && len(s.data) > 0 {
http.SetCookie(resp, newCookie(s.sid))
s.cookieSID = s.sid
}
}
func GetContextSession(req *http.Request) Store {
if setting.IsInTesting {
if store, ok := req.Context().Value(MockStoreContextKey).(Store); ok {
return store
}
}
return session.GetSession(req)
sess, _ := req.Context().Value(ContextKey).(Store)
return sess
}
// BeforeRegenerateSession is a list of functions that are called before a session is regenerated.
-202
View File
@@ -1,202 +0,0 @@
// Copyright 2019 The Gitea Authors. All rights reserved.
// SPDX-License-Identifier: MIT
package session
import (
"fmt"
"sync"
"gitea.dev/modules/json"
"gitea.com/go-chi/session"
couchbase "gitea.com/go-chi/session/couchbase"
memcache "gitea.com/go-chi/session/memcache"
mysql "gitea.com/go-chi/session/mysql"
postgres "gitea.com/go-chi/session/postgres"
)
// VirtualSessionProvider represents a shadowed session provider implementation.
type VirtualSessionProvider struct {
lock sync.RWMutex
provider session.Provider
}
// Init initializes the cookie session provider with the given config.
func (o *VirtualSessionProvider) Init(gcLifetime int64, config string) error {
var opts session.Options
if err := json.Unmarshal([]byte(config), &opts); err != nil {
return err
}
// Note that these options are unprepared so we can't just use NewManager here.
// Nor can we access the provider map in session.
// So we will just have to do this by hand.
// This is only slightly more wrong than modules/setting/session.go:23
switch opts.Provider {
case "memory":
o.provider = &session.MemProvider{}
case "file":
o.provider = &session.FileProvider{}
case "redis":
o.provider = &RedisProvider{}
case "db":
o.provider = &DBProvider{}
case "mysql":
o.provider = &mysql.MysqlProvider{}
case "postgres":
o.provider = &postgres.PostgresProvider{}
case "couchbase":
o.provider = &couchbase.CouchbaseProvider{}
case "memcache":
o.provider = &memcache.MemcacheProvider{}
default:
return fmt.Errorf("VirtualSessionProvider: Unknown Provider: %s", opts.Provider)
}
return o.provider.Init(gcLifetime, opts.ProviderConfig)
}
// Read returns raw session store by session ID.
func (o *VirtualSessionProvider) Read(sid string) (session.RawStore, error) {
o.lock.RLock()
defer o.lock.RUnlock()
if exist, err := o.provider.Exist(sid); err == nil && exist {
return o.provider.Read(sid)
} else if err != nil {
return nil, fmt.Errorf("check if '%s' exist failed: %w", sid, err)
}
kv := make(map[any]any)
return NewVirtualStore(o, sid, kv), nil
}
// Exist returns true if session with given ID exists.
func (o *VirtualSessionProvider) Exist(sid string) (bool, error) {
return true, nil
}
// Destroy deletes a session by session ID.
func (o *VirtualSessionProvider) Destroy(sid string) error {
o.lock.Lock()
defer o.lock.Unlock()
return o.provider.Destroy(sid)
}
// Regenerate regenerates a session store from old session ID to new one.
func (o *VirtualSessionProvider) Regenerate(oldsid, sid string) (session.RawStore, error) {
o.lock.Lock()
defer o.lock.Unlock()
return o.provider.Regenerate(oldsid, sid)
}
// Count counts and returns number of sessions.
func (o *VirtualSessionProvider) Count() (int, error) {
o.lock.RLock()
defer o.lock.RUnlock()
return o.provider.Count()
}
// GC calls GC to clean expired sessions.
func (o *VirtualSessionProvider) GC() {
o.provider.GC()
}
func init() {
session.Register("VirtualSession", &VirtualSessionProvider{})
}
// VirtualStore represents a virtual session store implementation.
type VirtualStore struct {
p *VirtualSessionProvider
sid string
lock sync.RWMutex
data map[any]any
released bool
}
// NewVirtualStore creates and returns a virtual session store.
func NewVirtualStore(p *VirtualSessionProvider, sid string, kv map[any]any) *VirtualStore {
return &VirtualStore{
p: p,
sid: sid,
data: kv,
}
}
// Set sets value to given key in session.
func (s *VirtualStore) Set(key, val any) error {
s.lock.Lock()
defer s.lock.Unlock()
s.data[key] = val
return nil
}
// Get gets value by given key in session.
func (s *VirtualStore) Get(key any) any {
s.lock.RLock()
defer s.lock.RUnlock()
return s.data[key]
}
// Delete delete a key from session.
func (s *VirtualStore) Delete(key any) error {
s.lock.Lock()
defer s.lock.Unlock()
delete(s.data, key)
return nil
}
// ID returns current session ID.
func (s *VirtualStore) ID() string {
return s.sid
}
// Release releases resource and save data to provider.
func (s *VirtualStore) Release() error {
s.lock.Lock()
defer s.lock.Unlock()
// Now need to lock the provider
s.p.lock.Lock()
defer s.p.lock.Unlock()
if len(s.data) > 0 {
// Now ensure that we don't exist!
realProvider := s.p.provider
if !s.released {
if exist, err := realProvider.Exist(s.sid); err == nil && exist {
// This is an error!
return fmt.Errorf("new sid '%s' already exists", s.sid)
} else if err != nil {
return fmt.Errorf("check if '%s' exist failed: %w", s.sid, err)
}
}
realStore, err := realProvider.Read(s.sid)
if err != nil {
return err
}
if err := realStore.Flush(); err != nil {
return err
}
for key, value := range s.data {
if err := realStore.Set(key, value); err != nil {
return err
}
}
err = realStore.Release()
if err == nil {
s.released = true
}
return err
}
return nil
}
// Flush deletes all session data.
func (s *VirtualStore) Flush() error {
s.lock.Lock()
defer s.lock.Unlock()
s.data = make(map[any]any)
return nil
}