mirror of
https://github.com/truewhile/MeBox.git
synced 2026-10-06 13:26:38 +08:00
fix: reduce sqlite write pressure during scans
This commit is contained in:
@@ -5,6 +5,7 @@ package database
|
|||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"sync"
|
||||||
|
|
||||||
"github.com/glebarez/sqlite"
|
"github.com/glebarez/sqlite"
|
||||||
"go.uber.org/zap"
|
"go.uber.org/zap"
|
||||||
@@ -38,6 +39,7 @@ func Open(cfg *config.Config, log *zap.Logger) (*gorm.DB, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("gorm open: %w", err)
|
return nil, fmt.Errorf("gorm open: %w", err)
|
||||||
}
|
}
|
||||||
|
installSQLiteWriteGate(db)
|
||||||
sqlDB, err := db.DB()
|
sqlDB, err := db.DB()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("gorm sqldb: %w", err)
|
return nil, fmt.Errorf("gorm sqldb: %w", err)
|
||||||
@@ -51,6 +53,39 @@ func Open(cfg *config.Config, log *zap.Logger) (*gorm.DB, error) {
|
|||||||
return db, nil
|
return db, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func installSQLiteWriteGate(db *gorm.DB) {
|
||||||
|
if db == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
gate := &sqliteWriteGate{}
|
||||||
|
lock := func(tx *gorm.DB) {
|
||||||
|
gate.Lock()
|
||||||
|
}
|
||||||
|
unlock := func(tx *gorm.DB) {
|
||||||
|
gate.Unlock()
|
||||||
|
}
|
||||||
|
_ = db.Callback().Create().Before("gorm:create").Register("mediastation:sqlite_write_lock", lock)
|
||||||
|
_ = db.Callback().Create().After("gorm:create").Register("mediastation:sqlite_write_unlock", unlock)
|
||||||
|
_ = db.Callback().Update().Before("gorm:update").Register("mediastation:sqlite_write_lock", lock)
|
||||||
|
_ = db.Callback().Update().After("gorm:update").Register("mediastation:sqlite_write_unlock", unlock)
|
||||||
|
_ = db.Callback().Delete().Before("gorm:delete").Register("mediastation:sqlite_write_lock", lock)
|
||||||
|
_ = db.Callback().Delete().After("gorm:delete").Register("mediastation:sqlite_write_unlock", unlock)
|
||||||
|
_ = db.Callback().Raw().Before("gorm:raw").Register("mediastation:sqlite_write_lock", lock)
|
||||||
|
_ = db.Callback().Raw().After("gorm:raw").Register("mediastation:sqlite_write_unlock", unlock)
|
||||||
|
}
|
||||||
|
|
||||||
|
type sqliteWriteGate struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
}
|
||||||
|
|
||||||
|
func (g *sqliteWriteGate) Lock() {
|
||||||
|
g.mu.Lock()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (g *sqliteWriteGate) Unlock() {
|
||||||
|
g.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
func buildDSN(cfg *config.Config) string {
|
func buildDSN(cfg *config.Config) string {
|
||||||
dbPath := cfg.Database.DBPath
|
dbPath := cfg.Database.DBPath
|
||||||
if !filepath.IsAbs(dbPath) {
|
if !filepath.IsAbs(dbPath) {
|
||||||
|
|||||||
@@ -309,6 +309,12 @@ func applyMediaQueryFilter(q *gorm.DB, filter MediaQueryFilter) *gorm.DB {
|
|||||||
// 显式写入)。这两个问题都让 EnrichLibrary(WHERE scrape_status='pending')
|
// 显式写入)。这两个问题都让 EnrichLibrary(WHERE scrape_status='pending')
|
||||||
// 永远捞不到数据。
|
// 永远捞不到数据。
|
||||||
func (r *MediaRepository) Upsert(ctx context.Context, m *model.Media) error {
|
func (r *MediaRepository) Upsert(ctx context.Context, m *model.Media) error {
|
||||||
|
return withSQLiteBusyRetry(ctx, func() error {
|
||||||
|
return r.upsert(ctx, m)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *MediaRepository) upsert(ctx context.Context, m *model.Media) error {
|
||||||
var existing model.Media
|
var existing model.Media
|
||||||
err := r.db.WithContext(ctx).Unscoped().Where("path = ?", m.Path).First(&existing).Error
|
err := r.db.WithContext(ctx).Unscoped().Where("path = ?", m.Path).First(&existing).Error
|
||||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
@@ -328,15 +334,16 @@ func (r *MediaRepository) Upsert(ctx context.Context, m *model.Media) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// 已存在:仅刷新文件层面的字段。
|
// 已存在:仅刷新文件层面的字段。
|
||||||
updates := map[string]any{
|
updates := map[string]any{}
|
||||||
"size_bytes": m.SizeBytes,
|
setIfChanged(updates, "size_bytes", existing.SizeBytes, m.SizeBytes)
|
||||||
"duration_sec": m.DurationSec,
|
setIfChanged(updates, "duration_sec", existing.DurationSec, m.DurationSec)
|
||||||
"width": m.Width,
|
setIfChanged(updates, "width", existing.Width, m.Width)
|
||||||
"height": m.Height,
|
setIfChanged(updates, "height", existing.Height, m.Height)
|
||||||
"video_codec": m.VideoCodec,
|
setIfChanged(updates, "video_codec", existing.VideoCodec, m.VideoCodec)
|
||||||
"audio_codec": m.AudioCodec,
|
setIfChanged(updates, "audio_codec", existing.AudioCodec, m.AudioCodec)
|
||||||
"container": m.Container,
|
setIfChanged(updates, "container", existing.Container, m.Container)
|
||||||
"deleted_at": nil,
|
if existing.DeletedAt.Valid {
|
||||||
|
updates["deleted_at"] = nil
|
||||||
}
|
}
|
||||||
// 回填硬链接身份标识,便于后续扫描去重(避免重复识别/多倍占用)。
|
// 回填硬链接身份标识,便于后续扫描去重(避免重复识别/多倍占用)。
|
||||||
if m.FileID != "" && m.FileID != existing.FileID {
|
if m.FileID != "" && m.FileID != existing.FileID {
|
||||||
@@ -347,56 +354,56 @@ func (r *MediaRepository) Upsert(ctx context.Context, m *model.Media) error {
|
|||||||
// 真实剧名。仅在 existing 还停留在 'pending'/'' 时回填扫描标题,
|
// 真实剧名。仅在 existing 还停留在 'pending'/'' 时回填扫描标题,
|
||||||
// 避免覆盖刮削结果。
|
// 避免覆盖刮削结果。
|
||||||
if m.ScrapeStatus == "matched" || existing.ScrapeStatus == "pending" || existing.ScrapeStatus == "" || existing.ScrapeStatus == "no_match" {
|
if m.ScrapeStatus == "matched" || existing.ScrapeStatus == "pending" || existing.ScrapeStatus == "" || existing.ScrapeStatus == "no_match" {
|
||||||
updates["title"] = m.Title
|
setIfChanged(updates, "title", existing.Title, m.Title)
|
||||||
if m.Year > 0 {
|
if m.Year > 0 {
|
||||||
updates["year"] = m.Year
|
setIfChanged(updates, "year", existing.Year, m.Year)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if m.ScrapeStatus == "matched" {
|
if m.ScrapeStatus == "matched" {
|
||||||
updates["scrape_status"] = m.ScrapeStatus
|
setIfChanged(updates, "scrape_status", existing.ScrapeStatus, m.ScrapeStatus)
|
||||||
if m.OriginalName != "" {
|
if m.OriginalName != "" {
|
||||||
updates["original_name"] = m.OriginalName
|
setIfChanged(updates, "original_name", existing.OriginalName, m.OriginalName)
|
||||||
}
|
}
|
||||||
if m.PosterURL != "" {
|
if m.PosterURL != "" {
|
||||||
updates["poster_url"] = m.PosterURL
|
setIfChanged(updates, "poster_url", existing.PosterURL, m.PosterURL)
|
||||||
}
|
}
|
||||||
if m.BackdropURL != "" {
|
if m.BackdropURL != "" {
|
||||||
updates["backdrop_url"] = m.BackdropURL
|
setIfChanged(updates, "backdrop_url", existing.BackdropURL, m.BackdropURL)
|
||||||
}
|
}
|
||||||
if m.Overview != "" {
|
if m.Overview != "" {
|
||||||
updates["overview"] = m.Overview
|
setIfChanged(updates, "overview", existing.Overview, m.Overview)
|
||||||
}
|
}
|
||||||
if m.Rating > 0 {
|
if m.Rating > 0 {
|
||||||
updates["rating"] = m.Rating
|
setIfChanged(updates, "rating", existing.Rating, m.Rating)
|
||||||
}
|
}
|
||||||
if m.Year > 0 {
|
if m.Year > 0 {
|
||||||
updates["year"] = m.Year
|
setIfChanged(updates, "year", existing.Year, m.Year)
|
||||||
}
|
}
|
||||||
if m.TMDbID > 0 {
|
if m.TMDbID > 0 {
|
||||||
updates["tm_db_id"] = m.TMDbID
|
setIfChanged(updates, "tm_db_id", existing.TMDbID, m.TMDbID)
|
||||||
}
|
}
|
||||||
if m.BangumiID > 0 {
|
if m.BangumiID > 0 {
|
||||||
updates["bangumi_id"] = m.BangumiID
|
setIfChanged(updates, "bangumi_id", existing.BangumiID, m.BangumiID)
|
||||||
}
|
}
|
||||||
if m.Languages != "" {
|
if m.Languages != "" {
|
||||||
updates["languages"] = m.Languages
|
setIfChanged(updates, "languages", existing.Languages, m.Languages)
|
||||||
}
|
}
|
||||||
if m.Countries != "" {
|
if m.Countries != "" {
|
||||||
updates["countries"] = m.Countries
|
setIfChanged(updates, "countries", existing.Countries, m.Countries)
|
||||||
}
|
}
|
||||||
if m.Genres != "" {
|
if m.Genres != "" {
|
||||||
updates["genres"] = m.Genres
|
setIfChanged(updates, "genres", existing.Genres, m.Genres)
|
||||||
}
|
}
|
||||||
if m.NSFW {
|
if m.NSFW && !existing.NSFW {
|
||||||
updates["nsfw"] = true
|
updates["nsfw"] = true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if m.PosterURL != "" {
|
if m.PosterURL != "" {
|
||||||
updates["poster_url"] = m.PosterURL
|
setIfChanged(updates, "poster_url", existing.PosterURL, m.PosterURL)
|
||||||
}
|
}
|
||||||
if m.BackdropURL != "" {
|
if m.BackdropURL != "" {
|
||||||
updates["backdrop_url"] = m.BackdropURL
|
setIfChanged(updates, "backdrop_url", existing.BackdropURL, m.BackdropURL)
|
||||||
}
|
}
|
||||||
if lib := m.LibraryID; lib != "" && lib != existing.LibraryID {
|
if lib := m.LibraryID; lib != "" && lib != existing.LibraryID {
|
||||||
updates["library_id"] = m.LibraryID
|
updates["library_id"] = m.LibraryID
|
||||||
@@ -408,9 +415,13 @@ func (r *MediaRepository) Upsert(ctx context.Context, m *model.Media) error {
|
|||||||
updates["episode_num"] = m.EpisodeNum
|
updates["episode_num"] = m.EpisodeNum
|
||||||
}
|
}
|
||||||
if m.STRMURL != "" {
|
if m.STRMURL != "" {
|
||||||
updates["strm_url"] = m.STRMURL
|
setIfChanged(updates, "strm_url", existing.STRMURL, m.STRMURL)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if len(updates) == 0 {
|
||||||
|
*m = existing
|
||||||
|
return nil
|
||||||
|
}
|
||||||
if err := r.db.WithContext(ctx).Unscoped().Model(&model.Media{}).
|
if err := r.db.WithContext(ctx).Unscoped().Model(&model.Media{}).
|
||||||
Where("id = ?", existing.ID).Updates(updates).Error; err != nil {
|
Where("id = ?", existing.ID).Updates(updates).Error; err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -421,6 +432,12 @@ func (r *MediaRepository) Upsert(ctx context.Context, m *model.Media) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func setIfChanged[T comparable](updates map[string]any, key string, current, next T) {
|
||||||
|
if current != next {
|
||||||
|
updates[key] = next
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// FindByID returns the media row or (nil, nil).
|
// FindByID returns the media row or (nil, nil).
|
||||||
func (r *MediaRepository) FindByID(ctx context.Context, id string) (*model.Media, error) {
|
func (r *MediaRepository) FindByID(ctx context.Context, id string) (*model.Media, error) {
|
||||||
var m model.Media
|
var m model.Media
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package repository
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/glebarez/sqlite"
|
"github.com/glebarez/sqlite"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
@@ -10,6 +11,65 @@ import (
|
|||||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
func TestMediaUpsertSkipsUnchangedExistingRow(t *testing.T) {
|
||||||
|
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := database.AutoMigrate(db); err != nil {
|
||||||
|
t.Fatalf("migrate: %v", err)
|
||||||
|
}
|
||||||
|
repos := New(db)
|
||||||
|
lib := model.Library{Name: "电影", Path: "/media/movie", Type: "movie", Enabled: true}
|
||||||
|
if err := repos.Library.Create(t.Context(), &lib); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
media := model.Media{
|
||||||
|
LibraryID: lib.ID,
|
||||||
|
Title: "已有影片",
|
||||||
|
Path: "/media/movie/existing.mkv",
|
||||||
|
SizeBytes: 1024,
|
||||||
|
DurationSec: 60,
|
||||||
|
Width: 1920,
|
||||||
|
Height: 1080,
|
||||||
|
VideoCodec: "h264",
|
||||||
|
AudioCodec: "aac",
|
||||||
|
Container: "matroska,webm",
|
||||||
|
ScrapeStatus: "pending",
|
||||||
|
}
|
||||||
|
if err := repos.Media.Upsert(t.Context(), &media); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
var before model.Media
|
||||||
|
if err := repos.DB.Where("path = ?", media.Path).First(&before).Error; err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
time.Sleep(10 * time.Millisecond)
|
||||||
|
again := model.Media{
|
||||||
|
LibraryID: lib.ID,
|
||||||
|
Title: before.Title,
|
||||||
|
Path: before.Path,
|
||||||
|
SizeBytes: before.SizeBytes,
|
||||||
|
DurationSec: before.DurationSec,
|
||||||
|
Width: before.Width,
|
||||||
|
Height: before.Height,
|
||||||
|
VideoCodec: before.VideoCodec,
|
||||||
|
AudioCodec: before.AudioCodec,
|
||||||
|
Container: before.Container,
|
||||||
|
ScrapeStatus: before.ScrapeStatus,
|
||||||
|
}
|
||||||
|
if err := repos.Media.Upsert(t.Context(), &again); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
var after model.Media
|
||||||
|
if err := repos.DB.Where("path = ?", media.Path).First(&after).Error; err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if !after.UpdatedAt.Equal(before.UpdatedAt) {
|
||||||
|
t.Fatalf("unchanged upsert touched updated_at: before=%s after=%s", before.UpdatedAt, after.UpdatedAt)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestMediaSearchFilteredSupportsChineseFuzzyTerms(t *testing.T) {
|
func TestMediaSearchFilteredSupportsChineseFuzzyTerms(t *testing.T) {
|
||||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -946,10 +946,16 @@ func (d *DownloadService) syncDownloadTaskProgress(ctx context.Context, torrent
|
|||||||
if strings.TrimSpace(status) == "" {
|
if strings.TrimSpace(status) == "" {
|
||||||
status = matched.Status
|
status = matched.Status
|
||||||
}
|
}
|
||||||
updates := map[string]any{"progress": torrent.Progress}
|
updates := map[string]any{}
|
||||||
if status != "" {
|
if math.Abs(float64(matched.Progress-torrent.Progress)) > 0.0001 {
|
||||||
|
updates["progress"] = torrent.Progress
|
||||||
|
}
|
||||||
|
if status != "" && status != matched.Status {
|
||||||
updates["status"] = status
|
updates["status"] = status
|
||||||
}
|
}
|
||||||
|
if len(updates) == 0 {
|
||||||
|
return
|
||||||
|
}
|
||||||
_ = d.repo.DB.WithContext(ctx).Model(&model.DownloadTask{}).Where("id = ?", matched.ID).Updates(updates).Error
|
_ = d.repo.DB.WithContext(ctx).Model(&model.DownloadTask{}).Where("id = ?", matched.ID).Updates(updates).Error
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/glebarez/sqlite"
|
"github.com/glebarez/sqlite"
|
||||||
"go.uber.org/zap"
|
"go.uber.org/zap"
|
||||||
@@ -47,6 +48,46 @@ func TestDownloadViewsDoNotExposePrivateURL(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestSyncDownloadTaskProgressSkipsUnchangedCompletedTask(t *testing.T) {
|
||||||
|
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := db.AutoMigrate(&model.DownloadTask{}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
repos := repository.New(db)
|
||||||
|
task := &model.DownloadTask{
|
||||||
|
Source: "qbittorrent",
|
||||||
|
URL: "magnet:?xt=urn:btih:test",
|
||||||
|
Title: "Already.Done.S01E01",
|
||||||
|
SavePath: "/downloads",
|
||||||
|
Status: "completed",
|
||||||
|
Progress: 1,
|
||||||
|
}
|
||||||
|
if err := repos.Download.Create(t.Context(), task); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
var before model.DownloadTask
|
||||||
|
if err := db.First(&before, "id = ?", task.ID).Error; err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
time.Sleep(10 * time.Millisecond)
|
||||||
|
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
|
||||||
|
svc.syncDownloadTaskProgress(t.Context(), QBitTorrent{
|
||||||
|
Name: task.Title,
|
||||||
|
Progress: 1,
|
||||||
|
State: "completed",
|
||||||
|
}, tasksByIdentity([]model.DownloadTask{before}))
|
||||||
|
var after model.DownloadTask
|
||||||
|
if err := db.First(&after, "id = ?", task.ID).Error; err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if !after.UpdatedAt.Equal(before.UpdatedAt) {
|
||||||
|
t.Fatalf("unchanged completed torrent touched updated_at: before=%s after=%s", before.UpdatedAt, after.UpdatedAt)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestDownloadCompleteAutoOrganizesContentPath(t *testing.T) {
|
func TestDownloadCompleteAutoOrganizesContentPath(t *testing.T) {
|
||||||
root := t.TempDir()
|
root := t.TempDir()
|
||||||
src := filepath.Join(root, "downloads", "国产剧", "狂飙.S01E01.2023.1080p.mkv")
|
src := filepath.Join(root, "downloads", "国产剧", "狂飙.S01E01.2023.1080p.mkv")
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ import (
|
|||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
"encoding/hex"
|
"encoding/hex"
|
||||||
"errors"
|
"errors"
|
||||||
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/golang-jwt/jwt/v5"
|
"github.com/golang-jwt/jwt/v5"
|
||||||
@@ -37,14 +38,16 @@ type Claims struct {
|
|||||||
|
|
||||||
// TokenService 处理双令牌认证(Access Token + Refresh Token)。
|
// TokenService 处理双令牌认证(Access Token + Refresh Token)。
|
||||||
type TokenService struct {
|
type TokenService struct {
|
||||||
cfg *config.Config
|
cfg *config.Config
|
||||||
log *zap.Logger
|
log *zap.Logger
|
||||||
repo *repository.Container
|
repo *repository.Container
|
||||||
|
delayedStoreMu sync.Mutex
|
||||||
|
delayedStores map[string]struct{}
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewTokenService 创建令牌服务实例。
|
// NewTokenService 创建令牌服务实例。
|
||||||
func NewTokenService(cfg *config.Config, log *zap.Logger, repo *repository.Container) *TokenService {
|
func NewTokenService(cfg *config.Config, log *zap.Logger, repo *repository.Container) *TokenService {
|
||||||
return &TokenService{cfg: cfg, log: log, repo: repo}
|
return &TokenService{cfg: cfg, log: log, repo: repo, delayedStores: make(map[string]struct{})}
|
||||||
}
|
}
|
||||||
|
|
||||||
// TokenPair 包含访问令牌和刷新令牌。
|
// TokenPair 包含访问令牌和刷新令牌。
|
||||||
@@ -110,7 +113,9 @@ func (s *TokenService) issuePair(ctx context.Context, userID, role, tier string,
|
|||||||
zap.String("user_id", userID),
|
zap.String("user_id", userID),
|
||||||
zap.Error(err))
|
zap.Error(err))
|
||||||
}
|
}
|
||||||
go s.storeRefreshTokenEventually(userID, tokenHash, rt.ExpiresAt)
|
if s.trackDelayedStore(userID, tokenHash) {
|
||||||
|
go s.storeRefreshTokenEventually(userID, tokenHash, rt.ExpiresAt)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return &TokenPair{
|
return &TokenPair{
|
||||||
@@ -132,9 +137,12 @@ func (s *TokenService) storeRefreshToken(ctx context.Context, rt *model.RefreshT
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (s *TokenService) storeRefreshTokenEventually(userID, tokenHash string, expiresAt time.Time) {
|
func (s *TokenService) storeRefreshTokenEventually(userID, tokenHash string, expiresAt time.Time) {
|
||||||
delay := 500 * time.Millisecond
|
defer s.untrackDelayedStore(userID, tokenHash)
|
||||||
for attempt := 1; attempt <= 30; attempt++ {
|
delay := 5 * time.Second
|
||||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
for attempt := 1; attempt <= 8; attempt++ {
|
||||||
|
timer := time.NewTimer(delay)
|
||||||
|
<-timer.C
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||||
err := s.storeRefreshToken(ctx, &model.RefreshToken{
|
err := s.storeRefreshToken(ctx, &model.RefreshToken{
|
||||||
UserID: userID,
|
UserID: userID,
|
||||||
TokenHash: tokenHash,
|
TokenHash: tokenHash,
|
||||||
@@ -150,15 +158,13 @@ func (s *TokenService) storeRefreshTokenEventually(userID, tokenHash string, exp
|
|||||||
}
|
}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if s.log != nil && (attempt == 1 || attempt%10 == 0) {
|
if s.log != nil && (attempt == 1 || attempt == 4 || attempt == 8) {
|
||||||
s.log.Warn("refresh token delayed store still waiting",
|
s.log.Warn("refresh token delayed store still waiting",
|
||||||
zap.String("user_id", userID),
|
zap.String("user_id", userID),
|
||||||
zap.Int("attempt", attempt),
|
zap.Int("attempt", attempt),
|
||||||
zap.Error(err))
|
zap.Error(err))
|
||||||
}
|
}
|
||||||
timer := time.NewTimer(delay)
|
if delay < 60*time.Second {
|
||||||
<-timer.C
|
|
||||||
if delay < 10*time.Second {
|
|
||||||
delay *= 2
|
delay *= 2
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -167,6 +173,33 @@ func (s *TokenService) storeRefreshTokenEventually(userID, tokenHash string, exp
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *TokenService) trackDelayedStore(userID, tokenHash string) bool {
|
||||||
|
if s == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
key := userID + "\x00" + tokenHash
|
||||||
|
s.delayedStoreMu.Lock()
|
||||||
|
defer s.delayedStoreMu.Unlock()
|
||||||
|
if s.delayedStores == nil {
|
||||||
|
s.delayedStores = make(map[string]struct{})
|
||||||
|
}
|
||||||
|
if _, ok := s.delayedStores[key]; ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
s.delayedStores[key] = struct{}{}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *TokenService) untrackDelayedStore(userID, tokenHash string) {
|
||||||
|
if s == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
key := userID + "\x00" + tokenHash
|
||||||
|
s.delayedStoreMu.Lock()
|
||||||
|
delete(s.delayedStores, key)
|
||||||
|
s.delayedStoreMu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
func (s *TokenService) maxActiveRefreshTokens(ctx context.Context) int {
|
func (s *TokenService) maxActiveRefreshTokens(ctx context.Context) int {
|
||||||
cfg := loadBotConfig(ctx, s.repo)
|
cfg := loadBotConfig(ctx, s.repo)
|
||||||
if cfg.MaxLoggedClients < 1 {
|
if cfg.MaxLoggedClients < 1 {
|
||||||
|
|||||||
Reference in New Issue
Block a user