diff --git a/cache/afterDelete.go b/cache/afterDelete.go index cd3f3f7..57e8204 100644 --- a/cache/afterDelete.go +++ b/cache/afterDelete.go @@ -38,20 +38,58 @@ func AfterDelete(cache *Gorm2Cache) func(db *gorm.DB) { if err != nil { cache.Logger.CtxError(ctx, "[AfterDelete] invalidating cache for primary keys: %v error: %v", primaryKeys, err) - return + } else { + cache.Logger.CtxInfo(ctx, "[AfterDelete] invalidating cache for primary keys: %v finished.", primaryKeys) } - cache.Logger.CtxInfo(ctx, "[AfterDelete] invalidating cache for primary keys: %v finished.", primaryKeys) } else { cache.Logger.CtxInfo(ctx, "[AfterDelete] now start to invalidate all primary cache for table: %s", tableName) err := cache.InvalidateAllPrimaryCache(ctx, tableName) if err != nil { cache.Logger.CtxError(ctx, "[AfterDelete] invalidating primary cache for table %s error: %v", tableName, err) - return + } else { + cache.Logger.CtxInfo(ctx, "[AfterDelete] invalidating all primary cache for table: %s finished.", tableName) } - cache.Logger.CtxInfo(ctx, "[AfterDelete] invalidating all primary cache for table: %s finished.", tableName) } + // 失效unique键缓存 + // 尝试从WHERE子句中提取unique键 + uniqueKeysMap, _ := getUniqueKeysFromWhereClause(db) + if len(uniqueKeysMap) > 0 { + for indexName, uniqueKeys := range uniqueKeysMap { + if len(uniqueKeys) > 0 { + cache.Logger.CtxInfo(ctx, "[AfterDelete] now start to invalidate unique cache for index %s keys: %+v", indexName, uniqueKeys) + err := cache.BatchInvalidateUniqueCache(ctx, tableName, indexName, uniqueKeys) + if err != nil { + cache.Logger.CtxError(ctx, "[AfterDelete] invalidating unique cache for index %s keys %v error: %v", + indexName, uniqueKeys, err) + } else { + cache.Logger.CtxInfo(ctx, "[AfterDelete] invalidating unique cache for index %s keys: %+v finished.", indexName, uniqueKeys) + } + } + } + } else { + // 如果没有从WHERE子句提取到unique键,失效所有unique键缓存 + s := db.Statement.Schema + if s == nil && db.Statement.Model != nil { + stmt := &gorm.Statement{DB: db} + if err := stmt.Parse(db.Statement.Model); err == nil { + s = stmt.Schema + } + } + if s != nil { + allUniqueIndexes := getAllUniqueIndexes(s) + for indexName := range allUniqueIndexes { + cache.Logger.CtxInfo(ctx, "[AfterDelete] now start to invalidate all unique cache for index %s", indexName) + err := cache.InvalidateAllUniqueCache(ctx, tableName, indexName) + if err != nil { + cache.Logger.CtxError(ctx, "[AfterDelete] invalidating all unique cache for index %s error: %v", indexName, err) + } else { + cache.Logger.CtxInfo(ctx, "[AfterDelete] invalidating all unique cache for index %s finished.", indexName) + } + } + } + } } }() diff --git a/cache/afterUpdate.go b/cache/afterUpdate.go index 0deb1c2..a3085ab 100644 --- a/cache/afterUpdate.go +++ b/cache/afterUpdate.go @@ -40,18 +40,57 @@ func AfterUpdate(cache *Gorm2Cache) func(db *gorm.DB) { if err != nil { cache.Logger.CtxError(ctx, "[AfterUpdate] invalidating primary cache for key %v error: %v", primaryKeys, err) - return + } else { + cache.Logger.CtxInfo(ctx, "[AfterUpdate] invalidating cache for primary keys: %+v finished.", primaryKeys) } - cache.Logger.CtxInfo(ctx, "[AfterUpdate] invalidating cache for primary keys: %+v finished.", primaryKeys) } else { cache.Logger.CtxInfo(ctx, "[AfterUpdate] now start to invalidate all primary cache for table: %s", tableName) err := cache.InvalidateAllPrimaryCache(ctx, tableName) if err != nil { cache.Logger.CtxError(ctx, "[AfterUpdate] invalidating primary cache for table %s error: %v", tableName, err) - return + } else { + cache.Logger.CtxInfo(ctx, "[AfterUpdate] invalidating all primary cache for table: %s finished.", tableName) + } + } + + // 失效unique键缓存 + // 尝试从WHERE子句中提取unique键 + uniqueKeysMap, _ := getUniqueKeysFromWhereClause(db) + if len(uniqueKeysMap) > 0 { + for indexName, uniqueKeys := range uniqueKeysMap { + if len(uniqueKeys) > 0 { + cache.Logger.CtxInfo(ctx, "[AfterUpdate] now start to invalidate unique cache for index %s keys: %+v", indexName, uniqueKeys) + err := cache.BatchInvalidateUniqueCache(ctx, tableName, indexName, uniqueKeys) + if err != nil { + cache.Logger.CtxError(ctx, "[AfterUpdate] invalidating unique cache for index %s keys %v error: %v", + indexName, uniqueKeys, err) + } else { + cache.Logger.CtxInfo(ctx, "[AfterUpdate] invalidating unique cache for index %s keys: %+v finished.", indexName, uniqueKeys) + } + } + } + } else { + // 如果没有从WHERE子句提取到unique键,失效所有unique键缓存 + s := db.Statement.Schema + if s == nil && db.Statement.Model != nil { + stmt := &gorm.Statement{DB: db} + if err := stmt.Parse(db.Statement.Model); err == nil { + s = stmt.Schema + } + } + if s != nil { + allUniqueIndexes := getAllUniqueIndexes(s) + for indexName := range allUniqueIndexes { + cache.Logger.CtxInfo(ctx, "[AfterUpdate] now start to invalidate all unique cache for index %s", indexName) + err := cache.InvalidateAllUniqueCache(ctx, tableName, indexName) + if err != nil { + cache.Logger.CtxError(ctx, "[AfterUpdate] invalidating all unique cache for index %s error: %v", indexName, err) + } else { + cache.Logger.CtxInfo(ctx, "[AfterUpdate] invalidating all unique cache for index %s finished.", indexName) + } + } } - cache.Logger.CtxInfo(ctx, "[AfterUpdate] invalidating all primary cache for table: %s finished.", tableName) } } }() diff --git a/cache/cache.go b/cache/cache.go index 52eb52e..616cd91 100644 --- a/cache/cache.go +++ b/cache/cache.go @@ -117,12 +117,14 @@ func (c *Gorm2Cache) InvalidateSearchCache(ctx context.Context, tableName string } func (c *Gorm2Cache) InvalidatePrimaryCache(ctx context.Context, tableName string, primaryKey string) error { + // primaryKey 已经是最终格式(单个值或已用":"连接的联合主键),直接传入 return c.cache.DeleteKey(ctx, util.GenPrimaryCacheKey(c.InstanceId, tableName, primaryKey)) } func (c *Gorm2Cache) BatchInvalidatePrimaryCache(ctx context.Context, tableName string, primaryKeys []string) error { cacheKeys := make([]string, 0, len(primaryKeys)) for _, primaryKey := range primaryKeys { + // primaryKey 已经是最终格式(单个值或已用":"连接的联合主键),直接传入 cacheKeys = append(cacheKeys, util.GenPrimaryCacheKey(c.InstanceId, tableName, primaryKey)) } return c.cache.BatchDeleteKeys(ctx, cacheKeys) @@ -135,6 +137,7 @@ func (c *Gorm2Cache) InvalidateAllPrimaryCache(ctx context.Context, tableName st func (c *Gorm2Cache) BatchPrimaryKeyExists(ctx context.Context, tableName string, primaryKeys []string) (bool, error) { cacheKeys := make([]string, 0, len(primaryKeys)) for _, primaryKey := range primaryKeys { + // primaryKey 已经是最终格式(单个值或已用":"连接的联合主键),直接传入 cacheKeys = append(cacheKeys, util.GenPrimaryCacheKey(c.InstanceId, tableName, primaryKey)) } return c.cache.BatchKeyExist(ctx, cacheKeys) @@ -146,10 +149,14 @@ func (c *Gorm2Cache) SearchKeyExists(ctx context.Context, tableName string, SQL } func (c *Gorm2Cache) BatchSetPrimaryKeyCache(ctx context.Context, tableName string, kvs []util.Kv) error { - for idx, kv := range kvs { - kvs[idx].Key = util.GenPrimaryCacheKey(c.InstanceId, tableName, kv.Key) + cacheKvs := make([]util.Kv, 0, len(kvs)) + for _, kv := range kvs { + cacheKvs = append(cacheKvs, util.Kv{ + Key: util.GenPrimaryCacheKey(c.InstanceId, tableName, kv.Key), + Value: kv.Value, + }) } - return c.cache.BatchSetKeys(ctx, kvs) + return c.cache.BatchSetKeys(ctx, cacheKvs) } func (c *Gorm2Cache) SetSearchCache(ctx context.Context, cacheValue string, tableName string, @@ -169,7 +176,51 @@ func (c *Gorm2Cache) GetSearchCache(ctx context.Context, tableName string, sql s func (c *Gorm2Cache) BatchGetPrimaryCache(ctx context.Context, tableName string, primaryKeys []string) ([]string, error) { cacheKeys := make([]string, 0, len(primaryKeys)) for _, primaryKey := range primaryKeys { + // primaryKey 已经是最终格式(单个值或已用":"连接的联合主键),直接传入 cacheKeys = append(cacheKeys, util.GenPrimaryCacheKey(c.InstanceId, tableName, primaryKey)) } return c.cache.BatchGetValues(ctx, cacheKeys) } + +// BatchGetUniqueCache 批量获取unique键缓存 +func (c *Gorm2Cache) BatchGetUniqueCache(ctx context.Context, tableName string, uniqueIndexName string, uniqueKeys []string) ([]string, error) { + cacheKeys := make([]string, 0, len(uniqueKeys)) + for _, uniqueKey := range uniqueKeys { + // uniqueKey 已经是最终格式(单个值或已用":"连接的联合unique键),直接传入 + cacheKeys = append(cacheKeys, util.GenUniqueCacheKey(c.InstanceId, tableName, uniqueIndexName, uniqueKey)) + } + return c.cache.BatchGetValues(ctx, cacheKeys) +} + +// BatchSetUniqueCache 批量设置 unique 键缓存。不会修改调用方传入的 kvs。 +func (c *Gorm2Cache) BatchSetUniqueCache(ctx context.Context, tableName string, uniqueIndexName string, kvs []util.Kv) error { + cacheKvs := make([]util.Kv, 0, len(kvs)) + for _, kv := range kvs { + cacheKvs = append(cacheKvs, util.Kv{ + Key: util.GenUniqueCacheKey(c.InstanceId, tableName, uniqueIndexName, kv.Key), + Value: kv.Value, + }) + } + return c.cache.BatchSetKeys(ctx, cacheKvs) +} + +// InvalidateUniqueCache 失效unique键缓存 +func (c *Gorm2Cache) InvalidateUniqueCache(ctx context.Context, tableName string, uniqueIndexName string, uniqueKey string) error { + // uniqueKey 已经是最终格式(单个值或已用":"连接的联合unique键),直接传入 + return c.cache.DeleteKey(ctx, util.GenUniqueCacheKey(c.InstanceId, tableName, uniqueIndexName, uniqueKey)) +} + +// BatchInvalidateUniqueCache 批量失效unique键缓存 +func (c *Gorm2Cache) BatchInvalidateUniqueCache(ctx context.Context, tableName string, uniqueIndexName string, uniqueKeys []string) error { + cacheKeys := make([]string, 0, len(uniqueKeys)) + for _, uniqueKey := range uniqueKeys { + // uniqueKey 已经是最终格式(单个值或已用":"连接的联合unique键),直接传入 + cacheKeys = append(cacheKeys, util.GenUniqueCacheKey(c.InstanceId, tableName, uniqueIndexName, uniqueKey)) + } + return c.cache.BatchDeleteKeys(ctx, cacheKeys) +} + +// InvalidateAllUniqueCache 失效所有unique键缓存 +func (c *Gorm2Cache) InvalidateAllUniqueCache(ctx context.Context, tableName string, uniqueIndexName string) error { + return c.cache.DeleteKeysWithPrefix(ctx, util.GenUniqueCachePrefix(c.InstanceId, tableName, uniqueIndexName)) +} diff --git a/cache/cache_integration_test.go b/cache/cache_integration_test.go index d34559a..2a961ba 100644 --- a/cache/cache_integration_test.go +++ b/cache/cache_integration_test.go @@ -1,8 +1,12 @@ package cache import ( + "context" "os" + "strings" + "sync" "testing" + "time" "github.com/asjdf/gorm-cache/config" "github.com/asjdf/gorm-cache/storage" @@ -703,3 +707,144 @@ func TestQueryHandler_Bind_WithError(t *testing.T) { } } +// recordingStorage 记录 DeleteKeysWithPrefix 的调用,用于断言 schema-less 时是否仍失效 unique 缓存 +type recordingStorage struct { + *storage.Memory + mu sync.Mutex + deletedPrefixes []string +} + +func (r *recordingStorage) DeleteKeysWithPrefix(ctx context.Context, keyPrefix string) error { + r.mu.Lock() + r.deletedPrefixes = append(r.deletedPrefixes, keyPrefix) + r.mu.Unlock() + return r.Memory.DeleteKeysWithPrefix(ctx, keyPrefix) +} + +// testUserWithUnique 带 unique 索引的模型,用于 schema-less fallback 测试 +type testUserWithUnique struct { + ID uint `gorm:"primaryKey"` + Email string `gorm:"uniqueIndex:idx_email"` + Username string `gorm:"uniqueIndex:idx_username"` +} + +func (testUserWithUnique) TableName() string { + return "test_users_unique" +} + +// TestAfterDelete_SchemaNil_StillInvalidatesUniqueCache 复现:当 Schema 为 nil 但 Model 有 unique 索引时, +// fallback 路径应通过解析 Model 得到 schema 并失效所有 unique 缓存,否则会残留过期 unique 缓存。 +func TestAfterDelete_SchemaNil_StillInvalidatesUniqueCache(t *testing.T) { + rec := &recordingStorage{Memory: storage.NewMem(storage.DefaultMemStoreConfig)} + cfg := &config.CacheConfig{ + CacheStorage: rec, + CacheTTL: 1000, + DebugMode: false, + InvalidateWhenUpdate: true, + CacheLevel: config.CacheLevelAll, + Tables: []string{"test_users_unique"}, + } + cache := &Gorm2Cache{Config: cfg, stats: &stats{}} + cache.Init() + + db := setupTestDB(t) + // 模拟 schema-less 场景:Schema 为 nil,但 Model 指向带 unique 索引的模型 + db.Statement.Schema = nil + db.Statement.Model = &testUserWithUnique{} + db.Statement.Table = "test_users_unique" + db.RowsAffected = 1 + db.Error = nil + db.Statement.Context = context.Background() + // 不设置 WHERE,使 getUniqueKeysFromWhereClause 返回空,走「失效所有 unique」的 fallback 路径 + + hook := AfterDelete(cache) + hook(db) + + // 回调内是 goroutine,等待执行完 + time.Sleep(200 * time.Millisecond) + + rec.mu.Lock() + prefixes := append([]string(nil), rec.deletedPrefixes...) + rec.mu.Unlock() + + // 应至少对两个 unique 索引做 DeleteKeysWithPrefix(idx_email, idx_username) + var uniquePrefixCount int + for _, p := range prefixes { + if strings.Contains(p, ":u:") && strings.Contains(p, "test_users_unique") { + uniquePrefixCount++ + } + } + if uniquePrefixCount < 2 { + t.Errorf("schema-less fallback should invalidate all unique caches (expected >= 2 unique prefix deletes), got %d, prefixes: %v", uniquePrefixCount, prefixes) + } + // 同时应包含 idx_email 与 idx_username 的 prefix(util.GenUniqueCachePrefix 格式) + hasEmail := false + hasUsername := false + for _, p := range prefixes { + if strings.Contains(p, "idx_email") { + hasEmail = true + } + if strings.Contains(p, "idx_username") { + hasUsername = true + } + } + if !hasEmail || !hasUsername { + t.Errorf("expected unique prefix deletes for idx_email and idx_username, got prefixes: %v", prefixes) + } +} + +// TestAfterUpdate_SchemaNil_StillInvalidatesUniqueCache 与 AfterDelete 对称:Schema 为 nil 时 update 也应失效 unique 缓存 +func TestAfterUpdate_SchemaNil_StillInvalidatesUniqueCache(t *testing.T) { + rec := &recordingStorage{Memory: storage.NewMem(storage.DefaultMemStoreConfig)} + cfg := &config.CacheConfig{ + CacheStorage: rec, + CacheTTL: 1000, + DebugMode: false, + InvalidateWhenUpdate: true, + CacheLevel: config.CacheLevelAll, + Tables: []string{"test_users_unique"}, + } + cache := &Gorm2Cache{Config: cfg, stats: &stats{}} + cache.Init() + + db := setupTestDB(t) + db.Statement.Schema = nil + db.Statement.Model = &testUserWithUnique{} + db.Statement.Table = "test_users_unique" + db.RowsAffected = 1 + db.Error = nil + db.Statement.Context = context.Background() + + hook := AfterUpdate(cache) + hook(db) + + time.Sleep(200 * time.Millisecond) + + rec.mu.Lock() + prefixes := append([]string(nil), rec.deletedPrefixes...) + rec.mu.Unlock() + + var uniquePrefixCount int + for _, p := range prefixes { + if strings.Contains(p, ":u:") && strings.Contains(p, "test_users_unique") { + uniquePrefixCount++ + } + } + if uniquePrefixCount < 2 { + t.Errorf("schema-less fallback should invalidate all unique caches (expected >= 2 unique prefix deletes), got %d, prefixes: %v", uniquePrefixCount, prefixes) + } + hasEmail := false + hasUsername := false + for _, p := range prefixes { + if strings.Contains(p, "idx_email") { + hasEmail = true + } + if strings.Contains(p, "idx_username") { + hasUsername = true + } + } + if !hasEmail || !hasUsername { + t.Errorf("expected unique prefix deletes for idx_email and idx_username, got prefixes: %v", prefixes) + } +} + diff --git a/cache/dockertest_integration_test.go b/cache/dockertest_integration_test.go new file mode 100644 index 0000000..f47a4b4 --- /dev/null +++ b/cache/dockertest_integration_test.go @@ -0,0 +1,1028 @@ +package cache + +import ( + "errors" + "fmt" + "net" + "os" + "strings" + "sync" + "testing" + "time" + + "github.com/asjdf/gorm-cache/config" + "github.com/asjdf/gorm-cache/storage" + "github.com/ory/dockertest/v3" + "github.com/ory/dockertest/v3/docker" + "gorm.io/driver/mysql" + "gorm.io/driver/postgres" + "gorm.io/gorm" + "gorm.io/gorm/logger" +) + +// 联合主键测试模型 +type UserRole struct { + UserID int64 `gorm:"primaryKey;column:user_id"` + RoleID int64 `gorm:"primaryKey;column:role_id"` + Name string `gorm:"column:name"` +} + +func (UserRole) TableName() string { + return "user_roles" +} + +// Unique键测试模型 +type User struct { + ID uint `gorm:"primaryKey;column:id"` + Email string `gorm:"uniqueIndex:idx_email;column:email;size:255"` + Username string `gorm:"uniqueIndex:idx_username;column:username;size:100"` + Name string `gorm:"column:name;size:255"` +} + +func (User) TableName() string { + return "users" +} + +// 联合Unique键测试模型 +type UserSession struct { + ID uint `gorm:"primaryKey;column:id"` + UserID int64 `gorm:"uniqueIndex:idx_user_token;column:user_id"` + Token string `gorm:"uniqueIndex:idx_user_token;column:token;size:255"` + ExpiresAt time.Time `gorm:"column:expires_at"` +} + +func (UserSession) TableName() string { + return "user_sessions" +} + +var ( + mysqlPool *dockertest.Pool + mysqlResource *dockertest.Resource + mysqlDSN string + setupMySQLOnce sync.Once + cleanupMySQLOnce sync.Once + mysqlSetupErr error + + pgPool *dockertest.Pool + pgResource *dockertest.Resource + pgDSN string + setupPGOnce sync.Once + cleanupPGOnce sync.Once + pgSetupErr error +) + +func setupMySQL(t *testing.T) *gorm.DB { + setupMySQLOnce.Do(func() { + var err error + mysqlPool, err = dockertest.NewPool("") + if err != nil { + mysqlSetupErr = fmt.Errorf("could not connect to docker: %w", err) + return + } + + mysqlResource, err = mysqlPool.RunWithOptions(&dockertest.RunOptions{ + Repository: "mysql", + Tag: "8.0", + Env: []string{ + "MYSQL_ROOT_PASSWORD=testpass", + "MYSQL_DATABASE=testdb", + }, + PortBindings: map[docker.Port][]docker.PortBinding{ + "3306/tcp": {{HostPort: "0"}}, + }, + }, func(config *docker.HostConfig) { + config.AutoRemove = true + config.RestartPolicy = docker.RestartPolicy{Name: "no"} + }) + if err != nil { + mysqlSetupErr = fmt.Errorf("could not start MySQL resource: %w", err) + return + } + + host := mysqlResource.GetHostPort("3306/tcp") + mysqlDSN = fmt.Sprintf("root:testpass@tcp(%s)/testdb?charset=utf8mb4&parseTime=True&loc=Local", host) + + mysqlPool.MaxWait = 120 * time.Second + if err := mysqlPool.Retry(func() error { + conn, openErr := gorm.Open(mysql.Open(mysqlDSN), &gorm.Config{ + Logger: logger.Default.LogMode(logger.Silent), + }) + if openErr != nil { + return openErr + } + sqlDB, openErr := conn.DB() + if openErr != nil { + return openErr + } + err := sqlDB.Ping() + _ = sqlDB.Close() + return err + }); err != nil { + mysqlSetupErr = fmt.Errorf("could not connect to MySQL: %w", err) + return + } + }) + if mysqlSetupErr != nil { + t.Fatalf("MySQL setup failed: %v", mysqlSetupErr) + } + + db, err := gorm.Open(mysql.Open(mysqlDSN), &gorm.Config{ + Logger: logger.Default.LogMode(logger.Silent), + }) + if err != nil { + t.Fatalf("open MySQL DB failed: %v", err) + } + sqlDB, err := db.DB() + if err != nil { + t.Fatalf("get sql.DB from gorm failed: %v", err) + } + sqlDB.SetMaxOpenConns(10) + sqlDB.SetMaxIdleConns(5) + sqlDB.SetConnMaxLifetime(time.Hour) + t.Cleanup(func() { + _ = sqlDB.Close() + }) + + if err := db.Exec("DROP TABLE IF EXISTS user_roles, users, user_sessions").Error; err != nil { + t.Fatalf("failed to drop MySQL tables: %v", err) + } + if err := db.AutoMigrate(&UserRole{}, &User{}, &UserSession{}); err != nil { + t.Fatalf("Auto migrate error: %v", err) + } + return db +} + +func setupPostgreSQL(t *testing.T) *gorm.DB { + setupPGOnce.Do(func() { + var err error + pgPool, err = dockertest.NewPool("") + if err != nil { + pgSetupErr = fmt.Errorf("could not connect to docker: %w", err) + return + } + + pgResource, err = pgPool.RunWithOptions(&dockertest.RunOptions{ + Repository: "postgres", + Tag: "15-alpine", + Env: []string{ + "POSTGRES_PASSWORD=testpass", + "POSTGRES_DB=testdb", + }, + PortBindings: map[docker.Port][]docker.PortBinding{ + "5432/tcp": {{HostPort: "0"}}, + }, + }, func(config *docker.HostConfig) { + config.AutoRemove = true + config.RestartPolicy = docker.RestartPolicy{Name: "no"} + }) + if err != nil { + pgSetupErr = fmt.Errorf("could not start PostgreSQL resource: %w", err) + return + } + + hostPort := pgResource.GetHostPort("5432/tcp") + host, port, err := net.SplitHostPort(hostPort) + if err != nil { + pgSetupErr = fmt.Errorf("could not parse host:port: %w", err) + return + } + if strings.Contains(host, ":") && !strings.HasPrefix(host, "[") { + host = "[" + host + "]" + } + pgDSN = fmt.Sprintf("host=%s port=%s user=postgres password=testpass dbname=testdb sslmode=disable", host, port) + + pgPool.MaxWait = 120 * time.Second + if err := pgPool.Retry(func() error { + conn, openErr := gorm.Open(postgres.Open(pgDSN), &gorm.Config{ + Logger: logger.Default.LogMode(logger.Silent), + }) + if openErr != nil { + return openErr + } + sqlDB, openErr := conn.DB() + if openErr != nil { + return openErr + } + err := sqlDB.Ping() + _ = sqlDB.Close() + return err + }); err != nil { + pgSetupErr = fmt.Errorf("could not connect to PostgreSQL: %w", err) + return + } + }) + if pgSetupErr != nil { + t.Fatalf("PostgreSQL setup failed: %v", pgSetupErr) + } + + db, err := gorm.Open(postgres.Open(pgDSN), &gorm.Config{ + Logger: logger.Default.LogMode(logger.Silent), + }) + if err != nil { + t.Fatalf("open PostgreSQL DB failed: %v", err) + } + sqlDB, err := db.DB() + if err != nil { + t.Fatalf("get sql.DB from gorm failed: %v", err) + } + sqlDB.SetMaxOpenConns(10) + sqlDB.SetMaxIdleConns(5) + sqlDB.SetConnMaxLifetime(time.Hour) + t.Cleanup(func() { + _ = sqlDB.Close() + }) + + if err := db.Exec("DROP TABLE IF EXISTS user_roles, users, user_sessions CASCADE").Error; err != nil { + t.Fatalf("failed to drop PostgreSQL tables: %v", err) + } + if err := db.AutoMigrate(&UserRole{}, &User{}, &UserSession{}); err != nil { + t.Fatalf("Auto migrate error: %v", err) + } + return db +} + +func TestMain(m *testing.M) { + code := m.Run() + + // Cleanup: 每测试独立 *gorm.DB 由 t.Cleanup 关闭,此处仅回收容器 + cleanupMySQLOnce.Do(func() { + if mysqlResource != nil && mysqlPool != nil { + _ = mysqlPool.Purge(mysqlResource) + } + }) + cleanupPGOnce.Do(func() { + if pgResource != nil && pgPool != nil { + _ = pgPool.Purge(pgResource) + } + }) + + os.Exit(code) +} + +// waitForCondition 在 timeout 内按 interval 轮询 predicate,为真则返回 nil,超时返回 error。 +// 用于替代固定 time.Sleep,避免 CI 上因负载导致的偶发失败。 +func waitForCondition(interval, timeout time.Duration, predicate func() bool) error { + deadline := time.Now().Add(timeout) + for time.Now().Before(deadline) { + if predicate() { + return nil + } + time.Sleep(interval) + } + return fmt.Errorf("condition not met within %v", timeout) +} + +// 辅助函数:创建标准缓存配置 +func createTestCache(tables []string) (Cache, error) { + return NewGorm2Cache(&config.CacheConfig{ + CacheLevel: config.CacheLevelAll, + CacheStorage: storage.NewMem(storage.DefaultMemStoreConfig), + InvalidateWhenUpdate: true, + CacheTTL: 5000, + CacheMaxItemCnt: 100, + DebugMode: false, + Tables: tables, + }) +} + +func TestWaitForCondition(t *testing.T) { + t.Run("succeeds when predicate true immediately", func(t *testing.T) { + err := waitForCondition(5*time.Millisecond, 50*time.Millisecond, func() bool { return true }) + if err != nil { + t.Errorf("expected nil, got %v", err) + } + }) + t.Run("succeeds when predicate becomes true", func(t *testing.T) { + n := 0 + err := waitForCondition(5*time.Millisecond, 100*time.Millisecond, func() bool { + n++ + return n >= 3 + }) + if err != nil { + t.Errorf("expected nil after 3 polls, got %v", err) + } + }) + t.Run("returns error when timeout", func(t *testing.T) { + err := waitForCondition(5*time.Millisecond, 20*time.Millisecond, func() bool { return false }) + if err == nil { + t.Error("expected error on timeout") + } + }) +} + +// 测试联合主键 - MySQL +func TestCompositePrimaryKey_MySQL(t *testing.T) { + db := setupMySQL(t) + + cache, err := createTestCache([]string{"user_roles"}) + if err != nil { + t.Fatalf("Failed to create cache: %v", err) + } + if err := db.Use(cache); err != nil { + t.Fatalf("failed to register cache plugin: %v", err) + } + + // 创建测试数据 + userRoles := []UserRole{ + {UserID: 1, RoleID: 1, Name: "Admin"}, + {UserID: 1, RoleID: 2, Name: "User"}, + {UserID: 2, RoleID: 1, Name: "Admin"}, + } + if err := db.Create(&userRoles).Error; err != nil { + t.Fatalf("Failed to create user roles: %v", err) + } + + // 第一次查询 - 应该从数据库读取并缓存 + var result1 []UserRole + if err := db.Where("user_id = ? AND role_id = ?", 1, 1).Find(&result1).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + if len(result1) != 1 || result1[0].Name != "Admin" { + t.Errorf("Expected 1 result with name 'Admin', got %d results", len(result1)) + } + + // 第二次查询 - 应该从缓存读取 + var result2 []UserRole + if err := db.Where("user_id = ? AND role_id = ?", 1, 1).Find(&result2).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + if len(result2) != 1 || result2[0].Name != "Admin" { + t.Errorf("Expected 1 result with name 'Admin', got %d results", len(result2)) + } + + // 测试IN查询 + var result3 []UserRole + if err := db.Where("user_id = ? AND role_id IN (?)", 1, []int64{1, 2}).Find(&result3).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + if len(result3) != 2 { + t.Errorf("Expected 2 results, got %d", len(result3)) + } + + // 更新数据 - 应该失效缓存 + if err := db.Model(&UserRole{}).Where("user_id = ? AND role_id = ?", 1, 1).Update("name", "SuperAdmin").Error; err != nil { + t.Fatalf("Failed to update: %v", err) + } + + // 再次查询 - 应该从数据库读取新数据 + var result4 []UserRole + if err := db.Where("user_id = ? AND role_id = ?", 1, 1).Find(&result4).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + if len(result4) != 1 { + t.Fatalf("Expected 1 result, got %d", len(result4)) + } + if result4[0].Name != "SuperAdmin" { + t.Errorf("Expected name 'SuperAdmin', got %s", result4[0].Name) + } +} + +// 测试联合主键 - PostgreSQL +func TestCompositePrimaryKey_PostgreSQL(t *testing.T) { + db := setupPostgreSQL(t) + + cache, err := createTestCache([]string{"user_roles"}) + if err != nil { + t.Fatalf("Failed to create cache: %v", err) + } + if err := db.Use(cache); err != nil { + t.Fatalf("failed to register cache plugin: %v", err) + } + + // 创建测试数据 + userRoles := []UserRole{ + {UserID: 1, RoleID: 1, Name: "Admin"}, + {UserID: 1, RoleID: 2, Name: "User"}, + {UserID: 2, RoleID: 1, Name: "Admin"}, + } + if err := db.Create(&userRoles).Error; err != nil { + t.Fatalf("Failed to create user roles: %v", err) + } + + // 第一次查询 + var result1 []UserRole + if err := db.Where("user_id = ? AND role_id = ?", 1, 1).Find(&result1).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + if len(result1) != 1 || result1[0].Name != "Admin" { + t.Errorf("Expected 1 result with name 'Admin', got %d results", len(result1)) + } + + // 第二次查询 - 应该从缓存读取 + var result2 []UserRole + if err := db.Where("user_id = ? AND role_id = ?", 1, 1).Find(&result2).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + if len(result2) != 1 || result2[0].Name != "Admin" { + t.Errorf("Expected 1 result with name 'Admin', got %d results", len(result2)) + } +} + +// 测试Unique键 - MySQL +func TestUniqueKey_MySQL(t *testing.T) { + db := setupMySQL(t) + + cache, err := createTestCache([]string{"users"}) + if err != nil { + t.Fatalf("Failed to create cache: %v", err) + } + if err := db.Use(cache); err != nil { + t.Fatalf("failed to register cache plugin: %v", err) + } + + // 创建测试数据 + users := []User{ + {Email: "user1@example.com", Username: "user1", Name: "User 1"}, + {Email: "user2@example.com", Username: "user2", Name: "User 2"}, + } + if err := db.Create(&users).Error; err != nil { + t.Fatalf("Failed to create users: %v", err) + } + + // 第一次查询 - 通过email unique键查询 + var result1 User + if err := db.Where("email = ?", "user1@example.com").First(&result1).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + if result1.Name != "User 1" { + t.Errorf("Expected name 'User 1', got %s", result1.Name) + } + + // 第二次查询 - 应该从缓存读取 + var result2 User + if err := db.Where("email = ?", "user1@example.com").First(&result2).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + if result2.Name != "User 1" { + t.Errorf("Expected name 'User 1', got %s", result2.Name) + } + + // 通过username unique键查询 + var result3 User + if err := db.Where("username = ?", "user2").First(&result3).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + if result3.Name != "User 2" { + t.Errorf("Expected name 'User 2', got %s", result3.Name) + } +} + +// 测试联合Unique键 - MySQL +func TestCompositeUniqueKey_MySQL(t *testing.T) { + db := setupMySQL(t) + + cache, err := NewGorm2Cache(&config.CacheConfig{ + CacheLevel: config.CacheLevelAll, + CacheStorage: storage.NewMem(storage.DefaultMemStoreConfig), + InvalidateWhenUpdate: true, + CacheTTL: 5000, + CacheMaxItemCnt: 100, + DebugMode: false, + Tables: []string{"user_sessions"}, + }) + if err != nil { + t.Fatalf("Failed to create cache: %v", err) + } + if err := db.Use(cache); err != nil { + t.Fatalf("failed to register cache plugin: %v", err) + } + + // 创建测试数据 + sessions := []UserSession{ + {UserID: 1, Token: "token1", ExpiresAt: time.Now().Add(24 * time.Hour)}, + {UserID: 2, Token: "token2", ExpiresAt: time.Now().Add(24 * time.Hour)}, + } + if err := db.Create(&sessions).Error; err != nil { + t.Fatalf("Failed to create sessions: %v", err) + } + + // 第一次查询 - 通过联合unique键查询 + var result1 UserSession + if err := db.Where("user_id = ? AND token = ?", 1, "token1").First(&result1).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + if result1.UserID != 1 { + t.Errorf("Expected UserID 1, got %d", result1.UserID) + } + + // 第二次查询 - 应该从缓存读取 + var result2 UserSession + if err := db.Where("user_id = ? AND token = ?", 1, "token1").First(&result2).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + if result2.UserID != 1 { + t.Errorf("Expected UserID 1, got %d", result2.UserID) + } +} + +// 测试Unique键 - PostgreSQL +func TestUniqueKey_PostgreSQL(t *testing.T) { + db := setupPostgreSQL(t) + + cache, err := NewGorm2Cache(&config.CacheConfig{ + CacheLevel: config.CacheLevelAll, + CacheStorage: storage.NewMem(storage.DefaultMemStoreConfig), + InvalidateWhenUpdate: true, + CacheTTL: 5000, + CacheMaxItemCnt: 100, + DebugMode: false, + Tables: []string{"users"}, + }) + if err != nil { + t.Fatalf("Failed to create cache: %v", err) + } + if err := db.Use(cache); err != nil { + t.Fatalf("failed to register cache plugin: %v", err) + } + + // 创建测试数据 + users := []User{ + {Email: "user1@example.com", Username: "user1", Name: "User 1"}, + {Email: "user2@example.com", Username: "user2", Name: "User 2"}, + } + if err := db.Create(&users).Error; err != nil { + t.Fatalf("Failed to create users: %v", err) + } + + // 第一次查询 + var result1 User + if err := db.Where("email = ?", "user1@example.com").First(&result1).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + if result1.Name != "User 1" { + t.Errorf("Expected name 'User 1', got %s", result1.Name) + } + + // 第二次查询 - 应该从缓存读取 + var result2 User + if err := db.Where("email = ?", "user1@example.com").First(&result2).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + if result2.Name != "User 1" { + t.Errorf("Expected name 'User 1', got %s", result2.Name) + } +} + +// 测试缓存失效 - MySQL +func TestCacheInvalidation_MySQL(t *testing.T) { + db := setupMySQL(t) + + cache, err := createTestCache([]string{"user_roles", "users"}) + if err != nil { + t.Fatalf("Failed to create cache: %v", err) + } + if err := db.Use(cache); err != nil { + t.Fatalf("failed to register cache plugin: %v", err) + } + + // 创建测试数据 + userRole := UserRole{UserID: 1, RoleID: 1, Name: "Admin"} + if err := db.Create(&userRole).Error; err != nil { + t.Fatalf("Failed to create user role: %v", err) + } + + // 第一次查询 - 缓存 + var result1 UserRole + if err := db.Where("user_id = ? AND role_id = ?", 1, 1).First(&result1).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + + // 更新数据 + if err := db.Model(&UserRole{}).Where("user_id = ? AND role_id = ?", 1, 1).Update("name", "SuperAdmin").Error; err != nil { + t.Fatalf("Failed to update: %v", err) + } + + // 再次查询 - 应该获取新数据 + var result2 UserRole + if err := db.Where("user_id = ? AND role_id = ?", 1, 1).First(&result2).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + if result2.Name != "SuperAdmin" { + t.Errorf("Expected name 'SuperAdmin', got %s", result2.Name) + } + + // 删除数据 + if err := db.Where("user_id = ? AND role_id = ?", 1, 1).Delete(&UserRole{}).Error; err != nil { + t.Fatalf("Failed to delete: %v", err) + } + + // 查询已删除的数据 - 应该返回错误 + var result3 UserRole + if err := db.Where("user_id = ? AND role_id = ?", 1, 1).First(&result3).Error; err == nil { + t.Error("Expected error for deleted record, got nil") + } +} + +// 测试缓存统计 - MySQL +func TestCacheStats_MySQL(t *testing.T) { + db := setupMySQL(t) + + cache, err := createTestCache([]string{"users"}) + if err != nil { + t.Fatalf("Failed to create cache: %v", err) + } + if err := db.Use(cache); err != nil { + t.Fatalf("failed to register cache plugin: %v", err) + } + + // 创建测试数据 + user := User{Email: "test@example.com", Username: "test", Name: "Test User"} + if err := db.Create(&user).Error; err != nil { + t.Fatalf("Failed to create user: %v", err) + } + + // 第一次查询 - 应该从数据库读取 + var result1 User + if err := db.Where("id = ?", user.ID).First(&result1).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + + // 第二次查询 - 应该从缓存读取 + var result2 User + if err := db.Where("id = ?", user.ID).First(&result2).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + + // 验证查询结果正确 + if result1.ID != user.ID || result2.ID != user.ID { + t.Errorf("Expected user ID %d, got result1=%d, result2=%d", user.ID, result1.ID, result2.ID) + } + if result1.Name != "Test User" || result2.Name != "Test User" { + t.Errorf("Expected name 'Test User', got result1=%s, result2=%s", result1.Name, result2.Name) + } + + // 检查统计 - 至少应该有查询发生 + lookupCount := cache.LookupCount() + if lookupCount == 0 { + t.Errorf("Expected at least 1 lookup, got 0") + } +} + +// 测试缓存一致性 - 综合测试(创建、更新、删除、联合主键、Unique键) +func TestCacheConsistency_Comprehensive_MySQL(t *testing.T) { + db := setupMySQL(t) + + cache, err := createTestCache([]string{"user_roles", "users"}) + if err != nil { + t.Fatalf("Failed to create cache: %v", err) + } + if err := db.Use(cache); err != nil { + t.Fatalf("failed to register cache plugin: %v", err) + } + + t.Run("CreateAndQuery", func(t *testing.T) { + // 创建数据后立即查询 + userRole := UserRole{UserID: 10, RoleID: 20, Name: "NewRole"} + if err := db.Create(&userRole).Error; err != nil { + t.Fatalf("Failed to create: %v", err) + } + + var result UserRole + if err := db.Where("user_id = ? AND role_id = ?", 10, 20).First(&result).Error; err != nil { + t.Fatalf("Failed to query after create: %v", err) + } + if result.Name != "NewRole" || result.UserID != 10 || result.RoleID != 20 { + t.Errorf("Expected NewRole(10,20), got %s(%d,%d)", result.Name, result.UserID, result.RoleID) + } + }) + + t.Run("UpdateCompositeKey", func(t *testing.T) { + // 联合主键更新 + userRole := UserRole{UserID: 100, RoleID: 200, Name: "Original"} + if err := db.Create(&userRole).Error; err != nil { + t.Fatalf("Failed to create: %v", err) + } + + // 查询缓存 + if err := db.Where("user_id = ? AND role_id = ?", 100, 200).First(&UserRole{}).Error; err != nil { + t.Fatalf("failed to warm cache: %v", err) + } + + // 更新 + if err := db.Model(&UserRole{}).Where("user_id = ? AND role_id = ?", 100, 200).Update("name", "Updated").Error; err != nil { + t.Fatalf("Failed to update: %v", err) + } + + if err := waitForCondition(15*time.Millisecond, 2*time.Second, func() bool { + var r UserRole + if db.Where("user_id = ? AND role_id = ?", 100, 200).First(&r).Error != nil { + return false + } + return r.Name == "Updated" + }); err != nil { + t.Fatalf("cache did not reflect update: %v", err) + } + + // 验证一致性 + var result UserRole + if err := db.Where("user_id = ? AND role_id = ?", 100, 200).First(&result).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + if result.Name != "Updated" { + t.Errorf("Expected 'Updated', got '%s'", result.Name) + } + + // 验证数据库 + sqlDB, _ := db.DB() + var dbName string + row := sqlDB.QueryRow("SELECT name FROM user_roles WHERE user_id = ? AND role_id = ?", 100, 200) + if err := row.Scan(&dbName); err != nil { + t.Fatalf("Failed to query database: %v", err) + } + if result.Name != dbName { + t.Errorf("Cache inconsistency: cache='%s', db='%s'", result.Name, dbName) + } + }) + + t.Run("UpdateUniqueKey", func(t *testing.T) { + // Unique键更新 + user := User{Email: "unique@test.com", Username: "unique", Name: "Original"} + if err := db.Create(&user).Error; err != nil { + t.Fatalf("Failed to create: %v", err) + } + + // 通过unique键查询缓存 + if err := db.Where("email = ?", "unique@test.com").First(&User{}).Error; err != nil { + t.Fatalf("failed to warm cache: %v", err) + } + + // 更新 + if err := db.Model(&User{}).Where("id = ?", user.ID).Update("name", "Updated").Error; err != nil { + t.Fatalf("Failed to update: %v", err) + } + + if err := waitForCondition(15*time.Millisecond, 2*time.Second, func() bool { + var r User + if db.Where("email = ?", "unique@test.com").First(&r).Error != nil { + return false + } + return r.Name == "Updated" + }); err != nil { + t.Fatalf("cache did not reflect update: %v", err) + } + + // 验证 + var result User + if err := db.Where("email = ?", "unique@test.com").First(&result).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + if result.Name != "Updated" { + t.Errorf("Expected 'Updated', got '%s'", result.Name) + } + }) + + t.Run("Delete", func(t *testing.T) { + // 删除测试 + userRole := UserRole{UserID: 200, RoleID: 300, Name: "ToDelete"} + if err := db.Create(&userRole).Error; err != nil { + t.Fatalf("Failed to create: %v", err) + } + + // 查询缓存 + if err := db.Where("user_id = ? AND role_id = ?", 200, 300).First(&UserRole{}).Error; err != nil { + t.Fatalf("failed to warm cache: %v", err) + } + + // 删除 + if err := db.Where("user_id = ? AND role_id = ?", 200, 300).Delete(&UserRole{}).Error; err != nil { + t.Fatalf("Failed to delete: %v", err) + } + + if err := waitForCondition(15*time.Millisecond, 2*time.Second, func() bool { + var r UserRole + err := db.Where("user_id = ? AND role_id = ?", 200, 300).First(&r).Error + return errors.Is(err, gorm.ErrRecordNotFound) + }); err != nil { + t.Fatalf("cache did not reflect delete: %v", err) + } + + // 验证已删除 + var result UserRole + if err := db.Where("user_id = ? AND role_id = ?", 200, 300).First(&result).Error; err == nil { + t.Error("Expected error for deleted record") + } else if !errors.Is(err, gorm.ErrRecordNotFound) { + t.Errorf("Expected ErrRecordNotFound, got %v", err) + } + }) + + t.Run("MultipleUpdates", func(t *testing.T) { + // 多次更新 + userRole := UserRole{UserID: 300, RoleID: 400, Name: "V1"} + if err := db.Create(&userRole).Error; err != nil { + t.Fatalf("Failed to create: %v", err) + } + + updates := []string{"V2", "V3", "V4"} + for _, newName := range updates { + if err := db.Model(&UserRole{}).Where("user_id = ? AND role_id = ?", 300, 400).Update("name", newName).Error; err != nil { + t.Fatalf("Failed to update: %v", err) + } + time.Sleep(50 * time.Millisecond) + + var result UserRole + if err := db.Where("user_id = ? AND role_id = ?", 300, 400).First(&result).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + if result.Name != newName { + t.Errorf("Expected '%s', got '%s'", newName, result.Name) + } + } + }) +} + +// 测试缓存一致性 - 高级场景(批量更新、并发、删除重建、Unique键交叉查询) +func TestCacheConsistency_Advanced_MySQL(t *testing.T) { + db := setupMySQL(t) + + cache, err := createTestCache([]string{"user_roles", "users"}) + if err != nil { + t.Fatalf("Failed to create cache: %v", err) + } + if err := db.Use(cache); err != nil { + t.Fatalf("failed to register cache plugin: %v", err) + } + + t.Run("BatchUpdate", func(t *testing.T) { + // 批量更新 + userRoles := []UserRole{ + {UserID: 700, RoleID: 701, Name: "Role1"}, + {UserID: 700, RoleID: 702, Name: "Role2"}, + } + if err := db.Create(&userRoles).Error; err != nil { + t.Fatalf("Failed to create: %v", err) + } + + db.Where("user_id = ?", 700).Find(&[]UserRole{}) + + if err := db.Model(&UserRole{}).Where("user_id = ?", 700).Update("name", "BatchUpdated").Error; err != nil { + t.Fatalf("Failed to batch update: %v", err) + } + + time.Sleep(100 * time.Millisecond) + + var results []UserRole + if err := db.Where("user_id = ?", 700).Find(&results).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + if len(results) != 2 { + t.Fatalf("Expected 2 results, got %d", len(results)) + } + for _, result := range results { + if result.Name != "BatchUpdated" { + t.Errorf("Expected 'BatchUpdated', got '%s'", result.Name) + } + } + }) + + t.Run("ConcurrentReadWrite", func(t *testing.T) { + // 并发读写 + userRole := UserRole{UserID: 1000, RoleID: 2000, Name: "Initial"} + if err := db.Create(&userRole).Error; err != nil { + t.Fatalf("Failed to create: %v", err) + } + + const numGoroutines = 5 + const numUpdates = 3 + done := make(chan bool, numGoroutines) + errCh := make(chan error, numGoroutines) + + for i := 0; i < numGoroutines; i++ { + go func(id int) { + defer func() { done <- true }() + for j := 0; j < numUpdates; j++ { + newName := fmt.Sprintf("Update-%d-%d", id, j) + if err := db.Model(&UserRole{}).Where("user_id = ? AND role_id = ?", 1000, 2000).Update("name", newName).Error; err != nil { + errCh <- err + return + } + time.Sleep(10 * time.Millisecond) + } + }(i) + } + + for i := 0; i < numGoroutines; i++ { + <-done + } + close(errCh) + for err := range errCh { + t.Fatalf("concurrent update failed: %v", err) + } + + if err := waitForCondition(20*time.Millisecond, 3*time.Second, func() bool { + var r UserRole + if db.Where("user_id = ? AND role_id = ?", 1000, 2000).First(&r).Error != nil { + return false + } + sqlDB, _ := db.DB() + var dbName string + if sqlDB.QueryRow("SELECT name FROM user_roles WHERE user_id = ? AND role_id = ?", 1000, 2000).Scan(&dbName) != nil { + return false + } + return r.Name == dbName + }); err != nil { + t.Fatalf("cache consistency after concurrent updates: %v", err) + } + + var result UserRole + if err := db.Where("user_id = ? AND role_id = ?", 1000, 2000).First(&result).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + + sqlDB, _ := db.DB() + var dbName string + row := sqlDB.QueryRow("SELECT name FROM user_roles WHERE user_id = ? AND role_id = ?", 1000, 2000) + if err := row.Scan(&dbName); err != nil { + t.Fatalf("Failed to query database: %v", err) + } + if result.Name != dbName { + t.Errorf("Cache inconsistency: cache='%s', db='%s'", result.Name, dbName) + } + }) + + t.Run("DeleteAndRecreate", func(t *testing.T) { + // 删除后重新创建 + userRole1 := UserRole{UserID: 2000, RoleID: 3000, Name: "First"} + if err := db.Create(&userRole1).Error; err != nil { + t.Fatalf("Failed to create: %v", err) + } + + if err := db.Where("user_id = ? AND role_id = ?", 2000, 3000).First(&UserRole{}).Error; err != nil { + t.Fatalf("failed to warm cache: %v", err) + } + + if err := db.Where("user_id = ? AND role_id = ?", 2000, 3000).Delete(&UserRole{}).Error; err != nil { + t.Fatalf("Failed to delete: %v", err) + } + + if err := waitForCondition(15*time.Millisecond, 2*time.Second, func() bool { + var r UserRole + return errors.Is(db.Where("user_id = ? AND role_id = ?", 2000, 3000).First(&r).Error, gorm.ErrRecordNotFound) + }); err != nil { + t.Fatalf("cache did not reflect delete: %v", err) + } + + var result2 UserRole + if err := db.Where("user_id = ? AND role_id = ?", 2000, 3000).First(&result2).Error; err == nil { + t.Error("Expected error for deleted record") + } + + userRole2 := UserRole{UserID: 2000, RoleID: 3000, Name: "Second"} + if err := db.Create(&userRole2).Error; err != nil { + t.Fatalf("Failed to recreate: %v", err) + } + + if err := waitForCondition(15*time.Millisecond, 2*time.Second, func() bool { + var r UserRole + if db.Where("user_id = ? AND role_id = ?", 2000, 3000).First(&r).Error != nil { + return false + } + return r.Name == "Second" + }); err != nil { + t.Fatalf("cache did not reflect recreate: %v", err) + } + + var result3 UserRole + if err := db.Where("user_id = ? AND role_id = ?", 2000, 3000).First(&result3).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + if result3.Name != "Second" { + t.Errorf("Expected 'Second', got '%s'", result3.Name) + } + }) + + t.Run("UniqueKeyCrossQuery", func(t *testing.T) { + // Unique键交叉查询 + user := User{Email: "cross@test.com", Username: "crossuser", Name: "Original"} + if err := db.Create(&user).Error; err != nil { + t.Fatalf("Failed to create: %v", err) + } + + if err := db.Where("email = ?", "cross@test.com").First(&User{}).Error; err != nil { + t.Fatalf("failed to warm cache by email: %v", err) + } + if err := db.Where("username = ?", "crossuser").First(&User{}).Error; err != nil { + t.Fatalf("failed to warm cache by username: %v", err) + } + + if err := db.Model(&User{}).Where("id = ?", user.ID).Update("name", "Updated").Error; err != nil { + t.Fatalf("Failed to update: %v", err) + } + + time.Sleep(100 * time.Millisecond) + + var result3, result4 User + if err := db.Where("email = ?", "cross@test.com").First(&result3).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + if err := db.Where("username = ?", "crossuser").First(&result4).Error; err != nil { + t.Fatalf("Failed to query: %v", err) + } + + if result3.Name != "Updated" || result4.Name != "Updated" { + t.Errorf("Expected 'Updated', got email='%s', username='%s'", result3.Name, result4.Name) + } + if result3.Name != result4.Name { + t.Errorf("Inconsistent results: email='%s', username='%s'", result3.Name, result4.Name) + } + }) +} diff --git a/cache/helpers.go b/cache/helpers.go index 6e5e2af..c5a5b41 100644 --- a/cache/helpers.go +++ b/cache/helpers.go @@ -9,13 +9,34 @@ import ( "gorm.io/gorm" "gorm.io/gorm/clause" + "gorm.io/gorm/schema" ) -// getPrimaryKeysFromWhereClause try to find primary keys from Eq and IN exprs in WHERE clause, -// and get objects that are being operated -func getPrimaryKeysFromWhereClause(db *gorm.DB) []string { - primaryKeys := make([]string, 0) +// getPrimaryKeyFields 获取所有主键字段(支持联合主键) +// 返回的字段已按主键顺序排列 +func getPrimaryKeyFields(s *schema.Schema) []*schema.Field { + if s == nil { + return nil + } + // 使用Schema的PrimaryFields,它已经按顺序排列 + if len(s.PrimaryFields) > 0 { + return s.PrimaryFields + } + // 如果没有PrimaryFields,则从Fields中查找(兼容旧版本) + primaryKeyFields := make([]*schema.Field, 0) + for _, field := range s.Fields { + if field.PrimaryKey { + primaryKeyFields = append(primaryKeyFields, field) + } + } + return primaryKeyFields +} +// getPrimaryKeysFromWhereClause 从WHERE子句中提取主键值,支持联合主键 +// 返回格式:对于单个主键,返回["value1", "value2"](多个记录) +// 对于联合主键,返回["field1:value1:field2:value2", "field1:value3:field2:value4"](多个记录) +// 注意:联合主键的key格式为按字段顺序排列的值,用":"分隔 +func getPrimaryKeysFromWhereClause(db *gorm.DB) []string { cla, ok := db.Statement.Clauses["WHERE"] if !ok { return nil @@ -24,48 +45,64 @@ func getPrimaryKeysFromWhereClause(db *gorm.DB) []string { if !ok { return nil } - dbName := "" if db.Statement.Schema == nil { return nil } - for _, field := range db.Statement.Schema.Fields { - if field.PrimaryKey { - dbName = field.DBName - break - } - } - if len(dbName) == 0 { + + primaryKeyFields := getPrimaryKeyFields(db.Statement.Schema) + if len(primaryKeyFields) == 0 { return nil } + + // 收集WHERE子句中的字段值 + fieldValuesMap := make(map[string][]string) // key: fieldName, value: []values for _, expr := range where.Exprs { eqExpr, ok := expr.(clause.Eq) if ok { - if getColNameFromColumn(eqExpr.Column) == dbName { - primaryKeys = append(primaryKeys, fmt.Sprintf("%v", eqExpr.Value)) - } + fieldName := getColNameFromColumn(eqExpr.Column) + fieldValuesMap[fieldName] = append(fieldValuesMap[fieldName], fmt.Sprintf("%v", eqExpr.Value)) continue } inExpr, ok := expr.(clause.IN) if ok { - if getColNameFromColumn(inExpr.Column) == dbName { - for _, val := range inExpr.Values { - primaryKeys = append(primaryKeys, fmt.Sprintf("%v", val)) - } + fieldName := getColNameFromColumn(inExpr.Column) + values := make([]string, 0, len(inExpr.Values)) + for _, val := range inExpr.Values { + values = append(values, fmt.Sprintf("%v", val)) } + fieldValuesMap[fieldName] = append(fieldValuesMap[fieldName], values...) + continue } exprStruct, ok := expr.(clause.Expr) if ok { ttype := getExprType(exprStruct) - //fmt.Printf("expr: %+v, ttype: %s\n", exprStruct, ttype) if ttype == "in" || ttype == "eq" { fieldName := getColNameFromExpr(exprStruct, ttype) - if fieldName == dbName { - pKeys := getPrimaryKeysFromExpr(exprStruct, ttype) - primaryKeys = append(primaryKeys, pKeys...) - } + pKeys := getPrimaryKeysFromExpr(exprStruct, ttype) + fieldValuesMap[fieldName] = append(fieldValuesMap[fieldName], pKeys...) } } } + + // 检查是否所有主键字段都有值 + for _, field := range primaryKeyFields { + if len(fieldValuesMap[field.DBName]) == 0 { + return nil // 缺少某个主键字段的值 + } + } + + // 生成主键key列表 + // 对于单个主键:直接返回所有值 + if len(primaryKeyFields) == 1 { + return uniqueStringSlice(fieldValuesMap[primaryKeyFields[0].DBName]) + } + + // 对于联合主键:生成所有字段值的笛卡尔积(如 user_id=1, role_id IN (1,2,3) -> "1:1","1:2","1:3") + valueSlices := make([][]string, 0, len(primaryKeyFields)) + for _, field := range primaryKeyFields { + valueSlices = append(valueSlices, fieldValuesMap[field.DBName]) + } + primaryKeys := generateCartesianProduct(valueSlices, ":") return uniqueStringSlice(primaryKeys) } @@ -80,32 +117,264 @@ func getColNameFromColumn(col interface{}) string { } } +// hasOtherClauseExceptPrimaryField 检查WHERE子句中是否有除了主键字段之外的其他条件 +// 支持联合主键 func hasOtherClauseExceptPrimaryField(db *gorm.DB) bool { cla, ok := db.Statement.Clauses["WHERE"] if !ok { return false } where, ok := cla.Expression.(clause.Where) - dbName := "" - for _, field := range db.Statement.Schema.Fields { - if field.PrimaryKey { - dbName = field.DBName + if !ok { + return false + } + if db.Statement.Schema == nil { + return true + } + + primaryKeyFields := getPrimaryKeyFields(db.Statement.Schema) + if len(primaryKeyFields) == 0 { + return true // 没有主键,返回true跳过缓存 + } + + // 构建主键字段名集合 + primaryKeyFieldSet := make(map[string]struct{}, len(primaryKeyFields)) + for _, field := range primaryKeyFields { + primaryKeyFieldSet[field.DBName] = struct{}{} + } + + // 检查每个表达式 + for _, expr := range where.Exprs { + eqExpr, ok := expr.(clause.Eq) + if ok { + fieldName := getColNameFromColumn(eqExpr.Column) + if _, isPrimaryKey := primaryKeyFieldSet[fieldName]; !isPrimaryKey { + return true + } + continue + } + inExpr, ok := expr.(clause.IN) + if ok { + fieldName := getColNameFromColumn(inExpr.Column) + if _, isPrimaryKey := primaryKeyFieldSet[fieldName]; !isPrimaryKey { + return true + } + continue + } + exprStruct, ok := expr.(clause.Expr) + if ok { + ttype := getExprType(exprStruct) + if ttype == "in" || ttype == "eq" { + fieldName := getColNameFromExpr(exprStruct, ttype) + if _, isPrimaryKey := primaryKeyFieldSet[fieldName]; !isPrimaryKey { + return true + } + continue + } + return true + } + // 其他类型的表达式,视为有其他条件 + return true + } + return false +} + +// getUniqueIndexFields 获取指定unique索引的字段列表 +func getUniqueIndexFields(s *schema.Schema, indexName string) []*schema.Field { + if s == nil { + return nil + } + allIndexes := s.ParseIndexes() + for _, index := range allIndexes { + if index.Name == indexName && index.Class == "UNIQUE" { + fields := make([]*schema.Field, 0, len(index.Fields)) + for _, fieldOption := range index.Fields { + if fieldOption.Field != nil { + fields = append(fields, fieldOption.Field) + } + } + return fields } } - if len(dbName) == 0 { - return true // return true to skip cache + return nil +} + +// getAllUniqueIndexes 获取所有unique索引 +func getAllUniqueIndexes(s *schema.Schema) map[string]*schema.Index { + if s == nil { + return nil } + allIndexes := s.ParseIndexes() + uniqueIndexes := make(map[string]*schema.Index) + for _, index := range allIndexes { + if index.Class == "UNIQUE" { + uniqueIndexes[index.Name] = index + } + } + if len(uniqueIndexes) == 0 { + return nil + } + return uniqueIndexes +} + +// getUniqueKeysFromWhereClause 从WHERE子句中提取unique键值 +// 返回: keys map[uniqueIndexName][]string,indexes map[uniqueIndexName]*schema.Index(仅 ParseIndexes 一次,供调用方复用) +func getUniqueKeysFromWhereClause(db *gorm.DB) (map[string][]string, map[string]*schema.Index) { + cla, ok := db.Statement.Clauses["WHERE"] + if !ok { + return nil, nil + } + where, ok := cla.Expression.(clause.Where) + if !ok { + return nil, nil + } + if db.Statement.Schema == nil { + return nil, nil + } + + uniqueIndexes := getAllUniqueIndexes(db.Statement.Schema) + if len(uniqueIndexes) == 0 { + return nil, nil + } + + // 收集WHERE子句中的字段值 + fieldValuesMap := make(map[string][]string) // key: fieldName, value: []values for _, expr := range where.Exprs { eqExpr, ok := expr.(clause.Eq) if ok { - if getColNameFromColumn(eqExpr.Column) != dbName { + fieldName := getColNameFromColumn(eqExpr.Column) + fieldValuesMap[fieldName] = append(fieldValuesMap[fieldName], fmt.Sprintf("%v", eqExpr.Value)) + continue + } + inExpr, ok := expr.(clause.IN) + if ok { + fieldName := getColNameFromColumn(inExpr.Column) + values := make([]string, 0, len(inExpr.Values)) + for _, val := range inExpr.Values { + values = append(values, fmt.Sprintf("%v", val)) + } + fieldValuesMap[fieldName] = append(fieldValuesMap[fieldName], values...) + continue + } + exprStruct, ok := expr.(clause.Expr) + if ok { + ttype := getExprType(exprStruct) + if ttype == "in" || ttype == "eq" { + fieldName := getColNameFromExpr(exprStruct, ttype) + pKeys := getPrimaryKeysFromExpr(exprStruct, ttype) + fieldValuesMap[fieldName] = append(fieldValuesMap[fieldName], pKeys...) + } + } + } + + // 检查每个unique索引,看是否所有字段都有值 + result := make(map[string][]string) + for indexName, index := range uniqueIndexes { + // 检查该unique索引的所有字段是否都有值 + allFieldsHaveValues := true + for _, fieldOption := range index.Fields { + if fieldOption.Field == nil { + allFieldsHaveValues = false + break + } + if len(fieldValuesMap[fieldOption.Field.DBName]) == 0 { + allFieldsHaveValues = false + break + } + } + if !allFieldsHaveValues { + continue + } + + // 生成unique键key列表 + if len(index.Fields) == 1 { + // 单个字段的unique键 + if index.Fields[0].Field != nil { + result[indexName] = uniqueStringSlice(fieldValuesMap[index.Fields[0].Field.DBName]) + } + } else { + // 联合unique键:笛卡尔积 + if len(index.Fields) == 0 { + continue + } + valueSlices := make([][]string, 0, len(index.Fields)) + skip := false + for _, fieldOption := range index.Fields { + if fieldOption.Field == nil { + skip = true + break + } + vals := fieldValuesMap[fieldOption.Field.DBName] + if len(vals) == 0 { + skip = true + break + } + valueSlices = append(valueSlices, vals) + } + if !skip && len(valueSlices) == len(index.Fields) { + uniqueKeys := generateCartesianProduct(valueSlices, ":") + if len(uniqueKeys) > 0 { + result[indexName] = uniqueStringSlice(uniqueKeys) + } + } + } + } + + if len(result) == 0 { + return nil, nil + } + return result, uniqueIndexes +} + +// hasOtherClauseExceptUniqueField 检查WHERE子句中是否有除了指定unique索引字段之外的其他条件。 +// 若 uniqueIndex 非 nil 则直接使用其 Fields,避免重复 ParseIndexes。 +func hasOtherClauseExceptUniqueField(db *gorm.DB, uniqueIndexName string, uniqueIndex *schema.Index) bool { + cla, ok := db.Statement.Clauses["WHERE"] + if !ok { + return false + } + where, ok := cla.Expression.(clause.Where) + if !ok { + return false + } + if db.Statement.Schema == nil { + return true + } + + var uniqueFields []*schema.Field + if uniqueIndex != nil { + for _, fo := range uniqueIndex.Fields { + if fo.Field != nil { + uniqueFields = append(uniqueFields, fo.Field) + } + } + } else { + uniqueFields = getUniqueIndexFields(db.Statement.Schema, uniqueIndexName) + } + if len(uniqueFields) == 0 { + return true + } + + // 构建unique字段名集合 + uniqueFieldSet := make(map[string]struct{}, len(uniqueFields)) + for _, field := range uniqueFields { + uniqueFieldSet[field.DBName] = struct{}{} + } + + // 检查每个表达式 + for _, expr := range where.Exprs { + eqExpr, ok := expr.(clause.Eq) + if ok { + fieldName := getColNameFromColumn(eqExpr.Column) + if _, isUniqueField := uniqueFieldSet[fieldName]; !isUniqueField { return true } continue } inExpr, ok := expr.(clause.IN) if ok { - if getColNameFromColumn(inExpr.Column) != dbName { + fieldName := getColNameFromColumn(inExpr.Column) + if _, isUniqueField := uniqueFieldSet[fieldName]; !isUniqueField { return true } continue @@ -115,14 +384,13 @@ func hasOtherClauseExceptPrimaryField(db *gorm.DB) bool { ttype := getExprType(exprStruct) if ttype == "in" || ttype == "eq" { fieldName := getColNameFromExpr(exprStruct, ttype) - if fieldName != dbName { + if _, isUniqueField := uniqueFieldSet[fieldName]; !isUniqueField { return true } continue } return true } - fmt.Printf("expr: %+v\n", expr) return true } return false @@ -214,6 +482,9 @@ func getPrimaryKeysFromExpr(expr clause.Expr, ttype string) []string { return primaryKeys } +// getObjectsAfterLoad 从查询结果中提取主键和对象,支持联合主键 +// 返回的主键格式:对于单个主键,返回["value1", "value2"] +// 对于联合主键,返回["value1:value2", "value3:value4"](按主键字段顺序) func getObjectsAfterLoad(db *gorm.DB) (primaryKeys []string, objects []interface{}) { primaryKeys = make([]string, 0) values := make([]reflect.Value, 0) @@ -229,30 +500,157 @@ func getObjectsAfterLoad(db *gorm.DB) (primaryKeys []string, objects []interface values = append(values, destValue) } - var valueOf func(context.Context, reflect.Value) (value interface{}, zero bool) = nil - if db.Statement.Schema != nil { - for _, field := range db.Statement.Schema.Fields { - if field.PrimaryKey { - valueOf = field.ValueOf - break - } + if db.Statement.Schema == nil { + // 没有schema,无法提取主键,只返回对象 + objects = make([]interface{}, 0, len(values)) + for _, elemValue := range values { + objects = append(objects, elemValue.Interface()) } + return primaryKeys, objects + } + + primaryKeyFields := getPrimaryKeyFields(db.Statement.Schema) + if len(primaryKeyFields) == 0 { + // 没有主键字段,只返回对象 + objects = make([]interface{}, 0, len(values)) + for _, elemValue := range values { + objects = append(objects, elemValue.Interface()) + } + return primaryKeys, objects } objects = make([]interface{}, 0, len(values)) for _, elemValue := range values { - if valueOf != nil { - primaryKey, isZero := valueOf(context.Background(), elemValue) - if isZero { + // 提取主键值 + if len(primaryKeyFields) == 1 { + // 单个主键 + valueOf := primaryKeyFields[0].ValueOf + if valueOf != nil { + primaryKey, isZero := valueOf(context.Background(), elemValue) + if isZero { + continue + } + primaryKeys = append(primaryKeys, fmt.Sprintf("%v", primaryKey)) + } + } else { + // 联合主键:必须所有字段都能取到非零值,且 ValueOf 非 nil + keyParts := make([]string, 0, len(primaryKeyFields)) + valid := true + for _, field := range primaryKeyFields { + valueOf := field.ValueOf + if valueOf == nil { + valid = false + break + } + primaryKey, isZero := valueOf(context.Background(), elemValue) + if isZero { + valid = false + break + } + keyParts = append(keyParts, fmt.Sprintf("%v", primaryKey)) + } + if !valid || len(keyParts) != len(primaryKeyFields) { continue } - primaryKeys = append(primaryKeys, fmt.Sprintf("%v", primaryKey)) + primaryKeys = append(primaryKeys, strings.Join(keyParts, ":")) } objects = append(objects, elemValue.Interface()) } return primaryKeys, objects } +// getUniqueKeysFromObjects 从对象中提取unique键值 +// 返回: map[uniqueIndexName]map[objectIndex]uniqueKey +func getUniqueKeysFromObjects(db *gorm.DB, objects []interface{}) map[string]map[int]string { + if db.Statement.Schema == nil { + return nil + } + + uniqueIndexes := getAllUniqueIndexes(db.Statement.Schema) + if len(uniqueIndexes) == 0 { + return nil + } + + result := make(map[string]map[int]string) + for indexName, index := range uniqueIndexes { + indexKeys := make(map[int]string) + for objIdx, obj := range objects { + objValue := reflect.Indirect(reflect.ValueOf(obj)) + if objValue.Kind() != reflect.Struct { + continue + } + + keyParts := make([]string, 0, len(index.Fields)) + allZero := true + for _, fieldOption := range index.Fields { + if fieldOption.Field == nil { + allZero = true + break + } + valueOf := fieldOption.Field.ValueOf + if valueOf != nil { + fieldValue, isZero := valueOf(context.Background(), objValue) + if isZero { + allZero = true + break + } + allZero = false + keyParts = append(keyParts, fmt.Sprintf("%v", fieldValue)) + } + } + if !allZero && len(keyParts) == len(index.Fields) { + indexKeys[objIdx] = strings.Join(keyParts, ":") + } + } + if len(indexKeys) > 0 { + result[indexName] = indexKeys + } + } + + if len(result) == 0 { + return nil + } + return result +} + +// generateCartesianProduct 对多列值做笛卡尔积,每行用 sep 连接成一条 key。 +// 例如 valueSlices = [["1"], ["1","2","3"]] -> ["1:1","1:2","1:3"]。 +func generateCartesianProduct(valueSlices [][]string, sep string) []string { + if len(valueSlices) == 0 { + return nil + } + n := 1 + for _, s := range valueSlices { + if len(s) == 0 { + return nil + } + n *= len(s) + } + result := make([]string, 0, n) + idx := make([]int, len(valueSlices)) + for { + parts := make([]string, len(valueSlices)) + for i, s := range valueSlices { + parts[i] = s[idx[i]] + } + result = append(result, strings.Join(parts, sep)) + // next combination + j := len(valueSlices) - 1 + for j >= 0 { + idx[j]++ + if idx[j] < len(valueSlices[j]) { + break + } + idx[j] = 0 + j-- + } + if j < 0 { + break + } + } + return result +} + func uniqueStringSlice(slice []string) []string { retSlice := make([]string, 0) mmap := make(map[string]struct{}) diff --git a/cache/helpers_integration_test.go b/cache/helpers_integration_test.go index 32bf4af..76150e3 100644 --- a/cache/helpers_integration_test.go +++ b/cache/helpers_integration_test.go @@ -1,7 +1,11 @@ package cache import ( + "context" + "errors" + "sync" "testing" + "time" "github.com/asjdf/gorm-cache/config" "github.com/asjdf/gorm-cache/storage" @@ -447,6 +451,69 @@ func TestAfterQuery_WithSingleFlight(t *testing.T) { // Singleflight should prevent duplicate queries } +// TestSingleFlight_ContextCancel_DoesNotAffectWaiters 复现并防护「一人取消,全家报错」: +// leader 使用已取消的 context 抢到执行权时,查库/缓存应用独立 background context, +// 否则 leader 的 cancel 会连累所有等待同一 key 的请求拿到 context.Canceled。 +func TestSingleFlight_ContextCancel_DoesNotAffectWaiters(t *testing.T) { + db := setupTestDB(t) + memStorage := storage.NewMem(storage.DefaultMemStoreConfig) + cache := &Gorm2Cache{ + Config: &config.CacheConfig{ + CacheStorage: memStorage, + CacheTTL: 1000, + DebugMode: false, + CacheLevel: config.CacheLevelAll, + Tables: []string{"test_users"}, + }, + stats: &stats{}, + } + cache.Init() + cache.Initialize(db) + + user := TestUser{Name: "User1", Age: 25} + if err := db.Create(&user).Error; err != nil { + t.Fatalf("create user: %v", err) + } + + // 模拟「急性子」:leader 用极短超时 context 先抢到 singleflight,查库过程中其 ctx 已过期 + singleFlightLeaderDelayForTest = 50 * time.Millisecond + defer func() { singleFlightLeaderDelayForTest = 0 }() + + ctxLeader, cancelLeader := context.WithTimeout(context.Background(), 1*time.Millisecond) + defer cancelLeader() + + var resultLeader, resultWaiter []TestUser + var errLeader, errWaiter error + var wg sync.WaitGroup + wg.Add(2) + go func() { + defer wg.Done() + tx := db.WithContext(ctxLeader) + errLeader = tx.Where("name = ?", "User1").Find(&resultLeader).Error + }() + time.Sleep(5 * time.Millisecond) // leader 已抢到锁并在 delay 中,waiter 将等待同一 call + go func() { + defer wg.Done() + tx := db.WithContext(context.Background()) + errWaiter = tx.Where("name = ?", "User1").Find(&resultWaiter).Error + }() + wg.Wait() + + // 正确行为:waiter 的 context 从未被取消/超时,不应拿到 leader 的 context 错误(一人取消/超时,全家报错) + if errWaiter == nil { + if len(resultWaiter) != 1 || resultWaiter[0].Name != "User1" { + t.Errorf("waiter should get one record: got %d, err=%v", len(resultWaiter), errWaiter) + } + } else { + if errors.Is(errWaiter, context.Canceled) || errors.Is(errWaiter, context.DeadlineExceeded) { + t.Errorf("一人取消全家报错:waiter 不应因 leader 的 context 取消/超时而得到相同错误, err=%v", errWaiter) + } + } + // leader 自己拿到 cancel 或超时是可以接受的,不强制断言 + _ = errLeader + _ = resultLeader +} + func TestAfterQuery_WithAsyncWrite(t *testing.T) { db := setupTestDB(t) memStorage := storage.NewMem(storage.DefaultMemStoreConfig) diff --git a/cache/helpers_test.go b/cache/helpers_test.go index 334e83b..27353c1 100644 --- a/cache/helpers_test.go +++ b/cache/helpers_test.go @@ -1,6 +1,7 @@ package cache import ( + "sync" "testing" "gorm.io/gorm" @@ -8,7 +9,6 @@ import ( "gorm.io/gorm/schema" ) - func TestGetColNameFromColumn(t *testing.T) { tests := []struct { name string @@ -237,15 +237,15 @@ func TestHasOtherClauseExceptPrimaryField(t *testing.T) { } s.Fields = []*schema.Field{ { - DBName: "id", + DBName: "id", PrimaryKey: true, }, { - DBName: "name", + DBName: "name", PrimaryKey: false, }, } - + tests := []struct { name string setup func(*gorm.DB) @@ -309,3 +309,249 @@ func TestHasOtherClauseExceptPrimaryField(t *testing.T) { } } +func TestGetPrimaryKeyFields(t *testing.T) { + tests := []struct { + name string + schema *schema.Schema + expectedCount int + expectedFields []string + }{ + { + name: "single primary key", + schema: &schema.Schema{ + Fields: []*schema.Field{ + {DBName: "id", PrimaryKey: true}, + {DBName: "name", PrimaryKey: false}, + }, + }, + expectedCount: 1, + expectedFields: []string{"id"}, + }, + { + name: "composite primary key", + schema: &schema.Schema{ + Fields: []*schema.Field{ + {DBName: "user_id", PrimaryKey: true}, + {DBName: "role_id", PrimaryKey: true}, + {DBName: "name", PrimaryKey: false}, + }, + PrimaryFields: []*schema.Field{ + {DBName: "user_id", PrimaryKey: true}, + {DBName: "role_id", PrimaryKey: true}, + }, + }, + expectedCount: 2, + expectedFields: []string{"user_id", "role_id"}, + }, + { + name: "no primary key", + schema: &schema.Schema{ + Fields: []*schema.Field{ + {DBName: "name", PrimaryKey: false}, + }, + }, + expectedCount: 0, + expectedFields: []string{}, + }, + { + name: "nil schema", + schema: nil, + expectedCount: 0, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := getPrimaryKeyFields(tt.schema) + if len(result) != tt.expectedCount { + t.Errorf("expected %d primary key fields, got %d", tt.expectedCount, len(result)) + return + } + for i, expectedField := range tt.expectedFields { + if i < len(result) && result[i].DBName != expectedField { + t.Errorf("expected field %s at index %d, got %s", expectedField, i, result[i].DBName) + } + } + }) + } +} + +// TestGetPrimaryKeysFromWhereClause_CompositeKeyCartesianProduct ensures +// composite primary key with differing value counts (e.g. user_id=1, role_id IN (1,2,3)) +// yields Cartesian product ["1:1","1:2","1:3"], not just ["1:1"]. +func TestGetPrimaryKeysFromWhereClause_CompositeKeyCartesianProduct(t *testing.T) { + s := &schema.Schema{ + Table: "user_roles", + } + s.Fields = []*schema.Field{ + {DBName: "user_id", PrimaryKey: true}, + {DBName: "role_id", PrimaryKey: true}, + {DBName: "name", PrimaryKey: false}, + } + db := &gorm.DB{ + Statement: &gorm.Statement{ + Schema: s, + }, + } + db.Statement.Clauses = map[string]clause.Clause{ + "WHERE": { + Expression: clause.Where{ + Exprs: []clause.Expression{ + clause.Eq{Column: "user_id", Value: 1}, + clause.IN{Column: "role_id", Values: []interface{}{1, 2, 3}}, + }, + }, + }, + } + got := getPrimaryKeysFromWhereClause(db) + want := []string{"1:1", "1:2", "1:3"} + if len(got) != len(want) { + t.Fatalf("expected %d keys, got %d: %v", len(want), len(got), got) + } + seen := make(map[string]bool) + for _, k := range got { + seen[k] = true + } + for _, w := range want { + if !seen[w] { + t.Errorf("missing expected key %q, got %v", w, got) + } + } +} + +// userWithUniqueIndexes is a model with unique indexes for testing getAllUniqueIndexes. +type userWithUniqueIndexes struct { + ID uint `gorm:"primaryKey"` + Email string `gorm:"uniqueIndex:idx_email"` + Username string `gorm:"uniqueIndex:idx_username"` +} + +func TestGetAllUniqueIndexes(t *testing.T) { + tests := []struct { + name string + schema *schema.Schema + expectedCount int + expectedNames []string + }{ + { + name: "with unique indexes from parsed schema", + schema: func() *schema.Schema { + s, _ := schema.Parse(&userWithUniqueIndexes{}, &sync.Map{}, schema.NamingStrategy{}) + return s + }(), + expectedCount: 2, + expectedNames: []string{"idx_email", "idx_username"}, + }, + { + name: "nil schema", + schema: nil, + expectedCount: 0, + expectedNames: nil, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := getAllUniqueIndexes(tt.schema) + if result == nil && tt.expectedCount > 0 { + t.Errorf("expected %d unique indexes, got nil", tt.expectedCount) + return + } + if result != nil && len(result) != tt.expectedCount { + t.Errorf("expected %d unique indexes, got %d", tt.expectedCount, len(result)) + } + for _, wantName := range tt.expectedNames { + if _, ok := result[wantName]; !ok { + t.Errorf("expected unique index %q in result %v", wantName, result) + } + } + }) + } +} + +func TestGetUniqueIndexFields(t *testing.T) { + tests := []struct { + name string + schema *schema.Schema + indexName string + expectedCount int + }{ + { + name: "nil schema", + schema: nil, + indexName: "idx_email", + expectedCount: 0, + }, + { + name: "non-existent index", + schema: &schema.Schema{ + Table: "users", + }, + indexName: "non_existent", + expectedCount: 0, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := getUniqueIndexFields(tt.schema, tt.indexName) + if len(result) != tt.expectedCount { + t.Errorf("expected %d fields, got %d", tt.expectedCount, len(result)) + } + }) + } +} + +func TestHasOtherClauseExceptUniqueField(t *testing.T) { + // Note: hasOtherClauseExceptUniqueField requires actual unique index information + // from schema.ParseIndexes(), which is complex to set up in unit tests. + // This function is better tested in integration tests. + // Here we just test basic cases. + + tests := []struct { + name string + setup func(*gorm.DB) + indexName string + expected bool + }{ + { + name: "nil schema", + setup: func(db *gorm.DB) { + db.Statement.Schema = nil + db.Statement.Clauses = map[string]clause.Clause{ + "WHERE": { + Expression: clause.Where{ + Exprs: []clause.Expression{ + clause.Eq{Column: "email", Value: "test@example.com"}, + }, + }, + }, + } + }, + indexName: "idx_email", + expected: true, // nil schema returns true + }, + { + name: "no WHERE clause", + setup: func(db *gorm.DB) { + db.Statement.Schema = &schema.Schema{Table: "users"} + db.Statement.Clauses = map[string]clause.Clause{} + }, + indexName: "idx_email", + expected: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + db := &gorm.DB{ + Statement: &gorm.Statement{}, + } + tt.setup(db) + result := hasOtherClauseExceptUniqueField(db, tt.indexName, nil) + if result != tt.expected { + t.Errorf("expected %v, got %v", tt.expected, result) + } + }) + } +} diff --git a/cache/query.go b/cache/query.go index 9a17386..8fb1ae2 100644 --- a/cache/query.go +++ b/cache/query.go @@ -1,12 +1,14 @@ package cache import ( + "context" "errors" "fmt" "reflect" "strconv" "strings" "sync" + "time" "github.com/asjdf/gorm-cache/config" "github.com/asjdf/gorm-cache/storage" @@ -19,6 +21,14 @@ import ( // singleFlight 流程设计 // 根据key lock住,等待结果。query before之前,会先判断是否有key,如果有,就等待结果,如果没有,就执行query before,然后执行query,然后把结果放到key里面,然后unlock,然后返回结果。 // 等待完成后 进行一手返回 然后err设置为err.singleflightHit,afterQuery结束的时候进行一手检查 +// +// 为避免「一人取消,全家报错」:抢到执行权的 leader 必须用独立的 background context 查库/读缓存, +// 否则若 leader 的请求 context 先被 cancel,会连累所有等待同一 key 的请求拿到同一个 cancel 错误。 +const singleflightContextTimeout = 30 * time.Second + +// singleFlightLeaderDelayForTest 仅用于单测:leader 抢到锁后睡眠时长,便于复现「一人取消全家报错」。 +// 由 _test 中赋值为非 0(如 50ms),测试结束后置回 0。 +var singleFlightLeaderDelayForTest time.Duration func newQueryHandler(c *Gorm2Cache) *queryHandler { return &queryHandler{cache: c} @@ -103,6 +113,14 @@ func (h *queryHandler) BeforeQuery() func(db *gorm.DB) { h.singleFlight.m[singleFlightKey] = c h.singleFlight.mu.Unlock() db.InstanceSet("gorm:cache:query:single_flight_call", c) + if singleFlightLeaderDelayForTest > 0 { + time.Sleep(singleFlightLeaderDelayForTest) + } + + // 避免「一人取消,全家报错」:leader 用独立 background context,不随首个请求的 cancel 而中断 + bgCtx, cancel := context.WithTimeout(context.Background(), singleflightContextTimeout) + db.Statement.Context = bgCtx + db.InstanceSet("gorm:cache:query:single_flight_cancel", cancel) tryPrimaryCache := func() (hit bool) { primaryKeys := getPrimaryKeysFromWhereClause(db) @@ -119,10 +137,10 @@ func (h *queryHandler) BeforeQuery() func(db *gorm.DB) { return } - // primary cache hit - cacheValues, err := cache.BatchGetPrimaryCache(ctx, tableName, primaryKeys) + // primary cache hit (use db.Statement.Context: when leader it is bgCtx to avoid cascading cancel) + cacheValues, err := cache.BatchGetPrimaryCache(db.Statement.Context, tableName, primaryKeys) if err != nil { - cache.Logger.CtxError(ctx, "[BeforeQuery] get primary cache value for key %v error: %v", primaryKeys, err) + cache.Logger.CtxError(db.Statement.Context, "[BeforeQuery] get primary cache value for key %v error: %v", primaryKeys, err) db.Error = nil return } @@ -155,12 +173,66 @@ func (h *queryHandler) BeforeQuery() func(db *gorm.DB) { return } + tryUniqueCache := func() (hit bool) { + uniqueKeysMap, uniqueIndexesMap := getUniqueKeysFromWhereClause(db) + if len(uniqueKeysMap) == 0 { + return + } + + // 尝试每个unique索引 + for uniqueIndexName, uniqueKeys := range uniqueKeysMap { + cache.Logger.CtxInfo(ctx, "[BeforeQuery] parse unique keys for index %s = %v", uniqueIndexName, uniqueKeys) + + if len(uniqueKeys) == 0 { + continue + } + + // 检查是否有其他条件(传入 index 避免重复 ParseIndexes) + hasOtherClauseInWhere := hasOtherClauseExceptUniqueField(db, uniqueIndexName, uniqueIndexesMap[uniqueIndexName]) + if hasOtherClauseInWhere { + continue + } + + // unique cache hit (use db.Statement.Context: when leader it is bgCtx to avoid cascading cancel) + cacheValues, err := cache.BatchGetUniqueCache(db.Statement.Context, tableName, uniqueIndexName, uniqueKeys) + if err != nil { + cache.Logger.CtxError(db.Statement.Context, "[BeforeQuery] get unique cache value for index %s key %v error: %v", uniqueIndexName, uniqueKeys, err) + continue + } + if len(cacheValues) != len(uniqueKeys) { + continue + } + finalValue := "" + + destKind := reflect.Indirect(reflect.ValueOf(db.Statement.Dest)).Kind() + if destKind == reflect.Struct && len(cacheValues) == 1 { + finalValue = cacheValues[0] + } else if (destKind == reflect.Array || destKind == reflect.Slice) && len(cacheValues) >= 1 { + finalValue = "[" + strings.Join(cacheValues, ",") + "]" + } + if len(finalValue) == 0 { + cache.Logger.CtxError(ctx, "[BeforeQuery] length of unique cache values and dest not matched") + continue + } + + err = json.Unmarshal([]byte(finalValue), db.Statement.Dest) + if err != nil { + cache.Logger.CtxError(ctx, "[BeforeQuery] unmarshal unique cache final value error: %v", err) + continue + } + db.Error = util.PrimaryCacheHit // 复用PrimaryCacheHit,因为逻辑相同 + hit = true + return + } + return + } + trySearchCache := func() (hit bool) { - // search cache hit - cacheValue, err := cache.GetSearchCache(ctx, tableName, sql, db.Statement.Vars...) + // search cache hit (use db.Statement.Context: when leader it is bgCtx to avoid cascading cancel) + cacheValue, err := cache.GetSearchCache(db.Statement.Context, tableName, sql, db.Statement.Vars...) if err != nil { if !errors.Is(err, storage.ErrCacheNotFound) { - cache.Logger.CtxError(ctx, "[BeforeQuery] get cache value for sql %s error: %v", sql, err) + cache.Logger.CtxError(db.Statement.Context, "[BeforeQuery] get cache value for sql %s error: %v", sql, err) } db.Error = nil return @@ -204,6 +276,11 @@ func (h *queryHandler) BeforeQuery() func(db *gorm.DB) { hit = true return } + // 尝试unique键缓存 + if tryUniqueCache() { + hit = true + return + } } if cache.Config.CacheLevel == config.CacheLevelAll || cache.Config.CacheLevel == config.CacheLevelOnlySearch { if !hit && trySearchCache() { @@ -322,9 +399,49 @@ func (h *queryHandler) AfterQuery() func(db *gorm.DB) { cache.Logger.CtxError(ctx, "[AfterQuery] batch set primary key cache for key %v error: %v", primaryKeys, err) } + + // cache unique cache data + uniqueKeysMap := getUniqueKeysFromObjects(db, objects) + if len(uniqueKeysMap) > 0 { + for indexName, objKeysMap := range uniqueKeysMap { + uniqueKvs := make([]util.Kv, 0, len(objKeysMap)) + for objIdx, uniqueKey := range objKeysMap { + if objIdx < len(objects) { + jsonStr, err := json.Marshal(objects[objIdx]) + if err != nil { + cache.Logger.CtxError(ctx, "[AfterQuery] object %v cannot marshal for unique cache, not cached", objects[objIdx]) + continue + } + uniqueKvs = append(uniqueKvs, util.Kv{ + Key: uniqueKey, + Value: string(jsonStr), + }) + } + } + if len(uniqueKvs) > 0 { + cache.Logger.CtxInfo(ctx, "[AfterQuery] start to set unique cache for index %s count=%d", indexName, len(uniqueKvs)) + err := cache.BatchSetUniqueCache(ctx, tableName, indexName, uniqueKvs) + if err != nil { + cache.Logger.CtxError(ctx, "[AfterQuery] batch set unique cache for index %s error: %v", indexName, err) + } + } + } + } } }() - if !cache.Config.AsyncWrite { + if cache.Config.AsyncWrite { + // 异步写时不在主路径 Wait,由后台 goroutine 在写完后 cancel,避免 fillCallAfterQuery 提前 cancel 导致写缓存被中止 + if cancelObj, hasCancel := db.InstanceGet("gorm:cache:query:single_flight_cancel"); hasCancel { + if cancel, ok := cancelObj.(context.CancelFunc); ok { + go func() { + wg.Wait() + cancel() + }() + // 替换为 no-op,fillCallAfterQuery 统一调用 cancel 时不会提前中止异步写;cache hit 时仍是真实 cancel,会释放 timer + db.InstanceSet("gorm:cache:query:single_flight_cancel", context.CancelFunc(func() {})) + } + } + } else { wg.Wait() } return @@ -385,6 +502,13 @@ func (h *queryHandler) fillCallAfterQuery(db *gorm.DB) { } return } + // 释放 singleflight leader 使用的 background context,避免 timer 泄漏 + // AsyncWrite 且走异步写路径时已把 cancel 替换为 no-op;cache hit 时此处调用真实 cancel,避免 30s timer 泄漏 + if cancelObj, hasCancel := db.InstanceGet("gorm:cache:query:single_flight_cancel"); hasCancel { + if cancel, ok := cancelObj.(context.CancelFunc); ok { + cancel() + } + } c.dest = db.Statement.Dest c.rowsAffected = db.RowsAffected c.err = db.Error diff --git a/cache/singleflight.go b/cache/singleflight.go index cc2cad3..7f2df84 100644 --- a/cache/singleflight.go +++ b/cache/singleflight.go @@ -43,4 +43,3 @@ func (g *Group) Forget(key string) { delete(g.m, key) g.mu.Unlock() } - diff --git a/go.mod b/go.mod index c2c6c9b..170ce78 100644 --- a/go.mod +++ b/go.mod @@ -1,6 +1,6 @@ module github.com/asjdf/gorm-cache -go 1.22 +go 1.24.0 toolchain go1.24.2 @@ -13,11 +13,14 @@ require ( github.com/ory/dockertest/v3 v3.12.0 github.com/redis/go-redis/v9 v9.0.2 github.com/smartystreets/goconvey v1.7.2 - gorm.io/gorm v1.24.5 + gorm.io/driver/mysql v1.6.0 + gorm.io/driver/postgres v1.6.0 + gorm.io/gorm v1.30.0 ) require ( dario.cat/mergo v1.0.0 // indirect + filippo.io/edwards25519 v1.1.0 // indirect github.com/Azure/go-ansiterm v0.0.0-20230124172434-306776ec8161 // indirect github.com/Microsoft/go-winio v0.6.2 // indirect github.com/Nvveen/Gotty v0.0.0-20120604004816-cd527374f1e5 // indirect @@ -31,15 +34,21 @@ require ( github.com/docker/go-units v0.5.0 // indirect github.com/dustin/go-humanize v1.0.1 // indirect github.com/glebarez/go-sqlite v1.20.3 // indirect + github.com/go-sql-driver/mysql v1.8.1 // indirect github.com/go-viper/mapstructure/v2 v2.1.0 // indirect github.com/gogo/protobuf v1.3.2 // indirect github.com/google/shlex v0.0.0-20191202100458-e7afc7fbc510 // indirect github.com/google/uuid v1.3.0 // indirect github.com/gopherjs/gopherjs v1.17.2 // indirect github.com/hashicorp/errwrap v1.0.0 // indirect + github.com/jackc/pgpassfile v1.0.0 // indirect + github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect + github.com/jackc/pgx/v5 v5.6.0 // indirect + github.com/jackc/puddle/v2 v2.2.2 // indirect github.com/jinzhu/inflection v1.0.0 // indirect github.com/jinzhu/now v1.1.5 // indirect github.com/jtolds/gls v4.20.0+incompatible // indirect + github.com/kr/text v0.2.0 // indirect github.com/mattn/go-isatty v0.0.17 // indirect github.com/moby/docker-image-spec v1.3.1 // indirect github.com/moby/sys/user v0.3.0 // indirect @@ -51,12 +60,16 @@ require ( github.com/opencontainers/runc v1.2.3 // indirect github.com/pkg/errors v0.9.1 // indirect github.com/remyoudompheng/bigfft v0.0.0-20230126093431-47fa9a501578 // indirect + github.com/rogpeppe/go-internal v1.14.1 // indirect github.com/sirupsen/logrus v1.9.3 // indirect github.com/smartystreets/assertions v1.13.0 // indirect github.com/xeipuuv/gojsonpointer v0.0.0-20190905194746-02993c407bfb // indirect github.com/xeipuuv/gojsonreference v0.0.0-20180127040603-bd5ef7bd5415 // indirect github.com/xeipuuv/gojsonschema v1.2.0 // indirect - golang.org/x/sys v0.28.0 // indirect + golang.org/x/crypto v0.48.0 // indirect + golang.org/x/sync v0.19.0 // indirect + golang.org/x/sys v0.41.0 // indirect + golang.org/x/text v0.34.0 // indirect gopkg.in/yaml.v2 v2.4.0 // indirect modernc.org/libc v1.22.2 // indirect modernc.org/mathutil v1.5.0 // indirect diff --git a/go.sum b/go.sum index b7e57f9..57137c9 100644 --- a/go.sum +++ b/go.sum @@ -20,6 +20,7 @@ github.com/cespare/xxhash/v2 v2.2.0 h1:DC2CZ1Ep5Y4k3ZQ899DldepgrayRUGE6BBZ/cd9Cj github.com/cespare/xxhash/v2 v2.2.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= github.com/containerd/continuity v0.4.5 h1:ZRoN1sXq9u7V6QoHMcVWGhOwDFqZ4B9i5H6un1Wh0x4= github.com/containerd/continuity v0.4.5/go.mod h1:/lNJvtJKUQStBzpVQ1+rasXO1LAWtUQssk28EZvJ3nE= +github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ33E= github.com/creack/pty v1.1.18 h1:n56/Zwd5o6whRC5PMGretI4IdRLlmBXYNjScPaBgsbY= github.com/creack/pty v1.1.18/go.mod h1:MOBLtS5ELjhRRrroQr9kyvTxUAFNvYEK993ew/Vr4O4= github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= @@ -63,9 +64,16 @@ github.com/hashicorp/errwrap v1.0.0 h1:hLrqtEDnRye3+sgx6z4qVLNuviH3MR5aQ0ykNJa/U github.com/hashicorp/errwrap v1.0.0/go.mod h1:YH+1FKiLXxHSkmPseP+kNlulaMuP3n2brvKWEqk/Jc4= github.com/hashicorp/go-multierror v1.1.1 h1:H5DkEtf6CXdFp0N0Em5UCwQpXMWke8IA0+lD48awMYo= github.com/hashicorp/go-multierror v1.1.1/go.mod h1:iw975J/qwKPdAO1clOe2L8331t/9/fmwbPZ6JB6eMoM= +github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM= +github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM= +github.com/jackc/pgx/v5 v5.6.0 h1:SWJzexBzPL5jb0GEsrPMLIsi/3jOo7RHlzTjcAeDrPY= +github.com/jackc/pgx/v5 v5.6.0/go.mod h1:DNZ/vlrUnhWCoFGxHAG8U2ljioxukquj7utPDgtQdTw= +github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo= +github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4= github.com/jinzhu/inflection v1.0.0 h1:K317FqzuhWc8YvSVlFMCCUb36O/S9MCKRDI7QkRKD/E= github.com/jinzhu/inflection v1.0.0/go.mod h1:h+uFLlag+Qp1Va5pdKtLDYj+kHp5pxUVkryuEj+Srlc= -github.com/jinzhu/now v1.1.4/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8= github.com/jinzhu/now v1.1.5 h1:/o9tlHleP7gOFmsnYNz3RGnqzefHA47wQpKrrdTIwXQ= github.com/jinzhu/now v1.1.5/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8= github.com/json-iterator/go v1.1.12 h1:PV8peI4a0ysnczrg+LtxykD8LfKY9ML6u2jnxaEnrnM= @@ -76,6 +84,10 @@ github.com/karlseguin/ccache/v3 v3.0.3 h1:cz+3tSdTrovp00xHPP3Y6ca/YuSl5kchhYG83w github.com/karlseguin/ccache/v3 v3.0.3/go.mod h1:qxC372+Qn+IBj8Pe3KvGjHPj0sWwEF7AeZVhsNPZ6uY= github.com/kisielk/errcheck v1.5.0/go.mod h1:pFxgyoBC7bSaBwPgfKdkLd5X25qrDl4LWUI2bnpBCr8= github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+oQHNcck= +github.com/kr/pretty v0.3.0 h1:WgNl7dwNpEZ6jJ9k1snq4pZsg7DOEN8hP9Xw0Tsjwk0= +github.com/kr/pretty v0.3.0/go.mod h1:640gp4NfQd8pI5XOwp5fnNeVWj67G7CFk/SaSQn7NBk= +github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= +github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= github.com/lib/pq v1.10.9 h1:YXG7RB+JIjhP29X+OtkiDnYaXQwpS4JEWq7dtCCRUEw= github.com/lib/pq v1.10.9/go.mod h1:AlVN5x4E4T544tWzH6hKfbfQvm3HdbOxrmggDNAPY9o= github.com/mattn/go-isatty v0.0.17 h1:BTarxUcIeDqL27Mc+vyvdWYSL28zpIhv3RoTdsLMPng= @@ -107,6 +119,8 @@ github.com/redis/go-redis/v9 v9.0.2/go.mod h1:/xDTe9EF1LM61hek62Poq2nzQSGj0xSrEt github.com/remyoudompheng/bigfft v0.0.0-20200410134404-eec4a21b6bb0/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= github.com/remyoudompheng/bigfft v0.0.0-20230126093431-47fa9a501578 h1:VstopitMQi3hZP0fzvnsLmzXZdQGc4bEcgu24cp+d4M= github.com/remyoudompheng/bigfft v0.0.0-20230126093431-47fa9a501578/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= +github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ= +github.com/rogpeppe/go-internal v1.14.1/go.mod h1:MaRKkUm5W0goXpeCfT7UZI6fk/L7L7so1lCWt35ZSgc= github.com/sirupsen/logrus v1.9.3 h1:dueUQJ1C2q9oE3F7wvmSGAaVtTmUizReu6fjN8uqzbQ= github.com/sirupsen/logrus v1.9.3/go.mod h1:naHLuLoDiP4jHNo9R0sCBMtWGeIprob74mVsIT4qYEQ= github.com/smartystreets/assertions v1.2.0/go.mod h1:tcbTF8ujkAEcZ8TElKY+i30BzYlVhC/LOxJk7iOWnoo= @@ -131,6 +145,8 @@ github.com/yuin/goldmark v1.2.1/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9dec golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI= golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= +golang.org/x/crypto v0.48.0 h1:/VRzVqiRSggnhY7gNRxPauEQ5Drw9haKdM0jqfcCFts= +golang.org/x/crypto v0.48.0/go.mod h1:r0kV5h3qnFPlQnBSrULhlsRfryS2pmewsg+XfMgkVos= golang.org/x/mod v0.2.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA= golang.org/x/mod v0.3.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA= golang.org/x/net v0.0.0-20190311183353-d8887717615a/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= @@ -141,16 +157,20 @@ golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwY golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20190911185100-cd5d95a43a6e/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= +golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4= +golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI= golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20200930185726-fdedc70b468f/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20210616094352-59db8d763f22/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220715151400-c0bba94af5f8/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220811171246-fbc7d0a398ab/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/sys v0.28.0 h1:Fksou7UEQUWlKvIdsqzJmUmCX3cZuD2+P3XyyzwMhlA= -golang.org/x/sys v0.28.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= +golang.org/x/sys v0.41.0 h1:Ivj+2Cp/ylzLiEU89QhWblYnOE9zerudt9Ftecq2C6k= +golang.org/x/sys v0.41.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks= golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= +golang.org/x/text v0.34.0 h1:oL/Qq0Kdaqxa1KbNeMKwQq0reLCCaFtqu2eNuSeNHbk= +golang.org/x/text v0.34.0/go.mod h1:homfLqTYRFyVYemLBFl5GgL/DWEiH5wcsQ5gSh1yziA= golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= golang.org/x/tools v0.0.0-20190328211700-ab21143f2384/go.mod h1:LCzVGOaR6xXOjkQ3onu1FJEFr0SW1gC7cKk1uF8kGRs= golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo= @@ -160,15 +180,20 @@ golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8T golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= golang.org/x/xerrors v0.0.0-20200804184101-5ec99f83aff1/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= -gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q= gopkg.in/yaml.v2 v2.4.0 h1:D8xgwECY7CYvx+Y2n4sBz93Jn9JRvxdiyyo8CTfuKaY= gopkg.in/yaml.v2 v2.4.0/go.mod h1:RDklbk79AGWmwhnvt/jBztapEOGDOx6ZbXqjP6csGnQ= gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= -gorm.io/gorm v1.24.5 h1:g6OPREKqqlWq4kh/3MCQbZKImeB9e6Xgc4zD+JgNZGE= -gorm.io/gorm v1.24.5/go.mod h1:DVrVomtaYTbqs7gB/x2uVvqnXzv0nqjB396B8cG4dBA= +gorm.io/driver/mysql v1.6.0 h1:eNbLmNTpPpTOVZi8MMxCi2aaIm0ZpInbORNXDwyLGvg= +gorm.io/driver/mysql v1.6.0/go.mod h1:D/oCC2GWK3M/dqoLxnOlaNKmXz8WNTfcS9y5ovaSqKo= +gorm.io/driver/postgres v1.6.0 h1:2dxzU8xJ+ivvqTRph34QX+WrRaJlmfyPqXmoGVjMBa4= +gorm.io/driver/postgres v1.6.0/go.mod h1:vUw0mrGgrTK+uPHEhAdV4sfFELrByKVGnaVRkXDhtWo= +gorm.io/gorm v1.30.0 h1:qbT5aPv1UH8gI99OsRlvDToLxW5zR7FzS9acZDOZcgs= +gorm.io/gorm v1.30.0/go.mod h1:8Z33v652h4//uMA76KjeDH8mJXPm1QNCYrMeatR0DOE= gotest.tools/v3 v3.5.1 h1:EENdUnS3pdur5nybKYIh2Vfgc8IUNBjxDPSjtiJcOzU= gotest.tools/v3 v3.5.1/go.mod h1:isy3WKz7GK6uNw/sbHzfKBLvlvXwUyV06n6brMxxopU= modernc.org/libc v1.22.2 h1:4U7v51GyhlWqQmwCHj28Rdq2Yzwk55ovjFrdPjs8Hb0= diff --git a/util/key.go b/util/key.go index a75a36e..0b3be39 100644 --- a/util/key.go +++ b/util/key.go @@ -1,9 +1,11 @@ package util import ( + "encoding/base64" "fmt" "math/rand" "reflect" + "sort" "strings" ) @@ -18,14 +20,76 @@ func GenInstanceId() string { return string(str) } -func GenPrimaryCacheKey(instanceId string, tableName string, primaryKey string) string { - return fmt.Sprintf("%s:%s:p:%s:%s", GormCachePrefix, instanceId, tableName, primaryKey) +// joinKeyParts 对多段做 base64 编码后用 sep 连接,避免 "a:b","c" 与 "a","b:c" 碰撞 +func joinKeyParts(parts []string, sep string) string { + if len(parts) == 0 { + return "" + } + if len(parts) == 1 { + return parts[0] + } + encoded := make([]string, len(parts)) + for i, p := range parts { + encoded[i] = base64.RawURLEncoding.EncodeToString([]byte(p)) + } + return strings.Join(encoded, sep) +} + +// GenPrimaryCacheKey 生成主键缓存key,支持单个主键和联合主键 +// 联合主键各段会做 base64 编码,避免含 ":" 的值产生 key 碰撞 +func GenPrimaryCacheKey(instanceId string, tableName string, primaryKeyValues ...string) string { + key := joinKeyParts(primaryKeyValues, ":") + return fmt.Sprintf("%s:%s:p:%s:%s", GormCachePrefix, instanceId, tableName, key) +} + +// GenPrimaryCacheKeyFromMap 从map生成主键缓存key,支持联合主键 +// primaryKeyMap: key为字段名,value为字段值,会自动按字段名排序以保证一致性 +func GenPrimaryCacheKeyFromMap(instanceId string, tableName string, primaryKeyMap map[string]string) string { + keys := make([]string, 0, len(primaryKeyMap)) + for k := range primaryKeyMap { + keys = append(keys, k) + } + // 排序以保证key的一致性 + sort.Strings(keys) + values := make([]string, 0, len(keys)) + for _, k := range keys { + values = append(values, primaryKeyMap[k]) + } + return GenPrimaryCacheKey(instanceId, tableName, values...) } func GenPrimaryCachePrefix(instanceId string, tableName string) string { return GormCachePrefix + ":" + instanceId + ":p:" + tableName } +// GenUniqueCacheKey 生成unique键缓存key,支持单个unique键和联合unique键 +// 联合unique键各段会做 base64 编码,避免含 ":" 的值产生 key 碰撞 +func GenUniqueCacheKey(instanceId string, tableName string, uniqueIndexName string, uniqueKeyValues ...string) string { + key := joinKeyParts(uniqueKeyValues, ":") + return fmt.Sprintf("%s:%s:u:%s:%s:%s", GormCachePrefix, instanceId, tableName, uniqueIndexName, key) +} + +// GenUniqueCacheKeyFromMap 从map生成unique键缓存key,支持联合unique键 +// uniqueKeyMap: key为字段名,value为字段值,会自动按字段名排序以保证一致性 +func GenUniqueCacheKeyFromMap(instanceId string, tableName string, uniqueIndexName string, uniqueKeyMap map[string]string) string { + keys := make([]string, 0, len(uniqueKeyMap)) + for k := range uniqueKeyMap { + keys = append(keys, k) + } + // 排序以保证key的一致性 + sort.Strings(keys) + values := make([]string, 0, len(keys)) + for _, k := range keys { + values = append(values, uniqueKeyMap[k]) + } + return GenUniqueCacheKey(instanceId, tableName, uniqueIndexName, values...) +} + +// GenUniqueCachePrefix 生成unique键缓存前缀 +func GenUniqueCachePrefix(instanceId string, tableName string, uniqueIndexName string) string { + return fmt.Sprintf("%s:%s:u:%s:%s", GormCachePrefix, instanceId, tableName, uniqueIndexName) +} + func GenSearchCacheKey(instanceId string, tableName string, sql string, vars ...interface{}) string { buf := strings.Builder{} buf.WriteString(sql) diff --git a/util/key_test.go b/util/key_test.go index 7e02e61..09d5bbf 100644 --- a/util/key_test.go +++ b/util/key_test.go @@ -2,6 +2,7 @@ package util import ( "reflect" + "strings" "testing" ) @@ -24,13 +25,103 @@ func TestGenInstanceId(t *testing.T) { func TestGenPrimaryCacheKey(t *testing.T) { instanceId := "test123" tableName := "users" - primaryKey := "1" - key := GenPrimaryCacheKey(instanceId, tableName, primaryKey) - expected := "gormcache:test123:p:users:1" + tests := []struct { + name string + primaryKeyVals []string + expected string + }{ + { + name: "single primary key", + primaryKeyVals: []string{"1"}, + expected: "gormcache:test123:p:users:1", + }, + { + name: "composite primary key", + primaryKeyVals: []string{"1", "2"}, + expected: "gormcache:test123:p:users:MQ:Mg", // base64-encoded parts to avoid ":" collision + }, + { + name: "composite primary key with three fields", + primaryKeyVals: []string{"1", "2", "3"}, + expected: "gormcache:test123:p:users:MQ:Mg:Mw", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + key := GenPrimaryCacheKey(instanceId, tableName, tt.primaryKeyVals...) + if key != tt.expected { + t.Errorf("expected %s, got %s", tt.expected, key) + } + }) + } +} + +func TestGenPrimaryCacheKeyFromMap(t *testing.T) { + instanceId := "test123" + tableName := "users" + + tests := []struct { + name string + primaryKeyMap map[string]string + expected string + }{ + { + name: "single primary key", + primaryKeyMap: map[string]string{ + "id": "1", + }, + expected: "gormcache:test123:p:users:1", + }, + { + name: "composite primary key", + primaryKeyMap: map[string]string{ + "user_id": "1", + "role_id": "2", + }, + expected: "gormcache:test123:p:users:Mg:MQ", // sorted: role_id, user_id -> values 2,1 -> base64 + }, + { + name: "composite primary key with three fields", + primaryKeyMap: map[string]string{ + "a": "1", + "b": "2", + "c": "3", + }, + expected: "gormcache:test123:p:users:MQ:Mg:Mw", // sorted: a, b, c + }, + } - if key != expected { - t.Errorf("expected %s, got %s", expected, key) + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + key := GenPrimaryCacheKeyFromMap(instanceId, tableName, tt.primaryKeyMap) + // 验证key格式正确,包含所有值(顺序可能因map遍历顺序而不同,但排序后应该一致) + if len(key) == 0 { + t.Error("expected non-empty key") + } + // 验证前缀正确 + expectedPrefix := "gormcache:test123:p:users:" + if !strings.Contains(key, expectedPrefix) { + t.Errorf("key should contain prefix %s, got %s", expectedPrefix, key) + } + if key != tt.expected { + t.Errorf("expected %s, got %s", tt.expected, key) + } + }) + } +} + + +// TestGenPrimaryCacheKey_NoCollision ensures composite key parts are encoded so that +// ["a:b","c"] and ["a","b:c"] do not produce the same cache key. +func TestGenPrimaryCacheKey_NoCollision(t *testing.T) { + instanceId := "test123" + tableName := "users" + k1 := GenPrimaryCacheKey(instanceId, tableName, "a:b", "c") + k2 := GenPrimaryCacheKey(instanceId, tableName, "a", "b:c") + if k1 == k2 { + t.Errorf("composite keys must not collide: both produced %q", k1) } } @@ -237,3 +328,108 @@ func TestGenSearchCacheKey_ReflectValue(t *testing.T) { } } +func TestGenUniqueCacheKey(t *testing.T) { + instanceId := "test123" + tableName := "users" + uniqueIndexName := "idx_email" + + tests := []struct { + name string + uniqueKeyVals []string + expected string + }{ + { + name: "single unique key", + uniqueKeyVals: []string{"user@example.com"}, + expected: "gormcache:test123:u:users:idx_email:user@example.com", + }, + { + name: "composite unique key", + uniqueKeyVals: []string{"user@example.com", "123"}, + expected: "gormcache:test123:u:users:idx_email:dXNlckBleGFtcGxlLmNvbQ:MTIz", // base64-encoded + }, + { + name: "composite unique key with three fields", + uniqueKeyVals: []string{"1", "2", "3"}, + expected: "gormcache:test123:u:users:idx_email:MQ:Mg:Mw", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + key := GenUniqueCacheKey(instanceId, tableName, uniqueIndexName, tt.uniqueKeyVals...) + if key != tt.expected { + t.Errorf("expected %s, got %s", tt.expected, key) + } + }) + } +} + +func TestGenUniqueCacheKeyFromMap(t *testing.T) { + instanceId := "test123" + tableName := "users" + uniqueIndexName := "idx_email" + + tests := []struct { + name string + uniqueKeyMap map[string]string + expected string + }{ + { + name: "single unique key", + uniqueKeyMap: map[string]string{ + "email": "user@example.com", + }, + expected: "gormcache:test123:u:users:idx_email:user@example.com", + }, + { + name: "composite unique key", + uniqueKeyMap: map[string]string{ + "email": "user@example.com", + "code": "123", + }, + expected: "gormcache:test123:u:users:idx_email:MTIz:dXNlckBleGFtcGxlLmNvbQ", // sorted: code, email -> base64 + }, + { + name: "composite unique key with three fields", + uniqueKeyMap: map[string]string{ + "a": "1", + "b": "2", + "c": "3", + }, + expected: "gormcache:test123:u:users:idx_email:MQ:Mg:Mw", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + key := GenUniqueCacheKeyFromMap(instanceId, tableName, uniqueIndexName, tt.uniqueKeyMap) + // 验证key格式正确 + if len(key) == 0 { + t.Error("expected non-empty key") + } + // 验证前缀正确 + expectedPrefix := "gormcache:test123:u:users:idx_email:" + if !strings.Contains(key, expectedPrefix) { + t.Errorf("key should contain prefix %s, got %s", expectedPrefix, key) + } + if key != tt.expected { + t.Errorf("expected %s, got %s", tt.expected, key) + } + }) + } +} + +func TestGenUniqueCachePrefix(t *testing.T) { + instanceId := "test123" + tableName := "users" + uniqueIndexName := "idx_email" + + prefix := GenUniqueCachePrefix(instanceId, tableName, uniqueIndexName) + expected := "gormcache:test123:u:users:idx_email" + + if prefix != expected { + t.Errorf("expected %s, got %s", expected, prefix) + } +} +