96 lines
2.8 KiB
Go
96 lines
2.8 KiB
Go
package area
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"os"
|
|
"strings"
|
|
"sync"
|
|
"testing"
|
|
|
|
"gorm.io/driver/postgres"
|
|
"gorm.io/gorm"
|
|
)
|
|
|
|
func TestConcurrentUpdateOnPostgresReturnsConflict(t *testing.T) {
|
|
baseDSN := os.Getenv("SENSE_AREA_TEST_DATABASE_URL")
|
|
if baseDSN == "" {
|
|
t.Skip("set SENSE_AREA_TEST_DATABASE_URL to run the PostgreSQL concurrency test")
|
|
}
|
|
admin, err := gorm.Open(postgres.Open(baseDSN), &gorm.Config{})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
const schema = "sense_area_69_concurrency"
|
|
if err = admin.Exec("DROP SCHEMA IF EXISTS " + schema + " CASCADE").Error; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err = admin.Exec("CREATE SCHEMA " + schema).Error; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
t.Cleanup(func() { admin.Exec("DROP SCHEMA IF EXISTS " + schema + " CASCADE") })
|
|
separator := "?"
|
|
if strings.Contains(baseDSN, "?") {
|
|
separator = "&"
|
|
}
|
|
db, err := gorm.Open(postgres.Open(baseDSN+separator+"search_path="+schema), &gorm.Config{})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err = db.AutoMigrate(&Definition{}, &Version{}); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
for _, statement := range []string{
|
|
`CREATE TABLE sense_devices (id text primary key, name text, location text, status text)`,
|
|
`CREATE TABLE sense_admission_profiles (device_id text, token text, name text, width integer, height integer, encoding text, verification_status text)`,
|
|
`CREATE TABLE sense_media_routes (id text primary key, device_id text, profile_token text)`,
|
|
`INSERT INTO sense_devices VALUES ('device-1','东门摄像机','教学楼东门','active')`,
|
|
`INSERT INTO sense_admission_profiles VALUES ('device-1','main','主码流',1920,1080,'H264','ready')`,
|
|
`INSERT INTO sense_media_routes VALUES ('device-1:main','device-1','main')`,
|
|
} {
|
|
if err = db.Exec(statement).Error; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|
|
service := NewService(db)
|
|
created, err := service.Create(context.Background(), triangleRequest())
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
start := make(chan struct{})
|
|
errorsChannel := make(chan error, 2)
|
|
var wait sync.WaitGroup
|
|
for index := 0; index < 2; index++ {
|
|
wait.Add(1)
|
|
go func() {
|
|
defer wait.Done()
|
|
<-start
|
|
request := triangleRequest()
|
|
request.ExpectedVersion = created.Version
|
|
_, updateErr := service.Update(context.Background(), created.ID, request)
|
|
errorsChannel <- updateErr
|
|
}()
|
|
}
|
|
close(start)
|
|
wait.Wait()
|
|
close(errorsChannel)
|
|
successes, conflicts := 0, 0
|
|
for updateErr := range errorsChannel {
|
|
switch {
|
|
case updateErr == nil:
|
|
successes++
|
|
case errors.Is(updateErr, ErrConflict):
|
|
conflicts++
|
|
default:
|
|
t.Fatalf("unexpected concurrent update error: %v", updateErr)
|
|
}
|
|
}
|
|
if successes != 1 || conflicts != 1 {
|
|
t.Fatalf("successes=%d conflicts=%d", successes, conflicts)
|
|
}
|
|
versions, err := service.Versions(context.Background(), created.ID)
|
|
if err != nil || len(versions) != 2 {
|
|
t.Fatalf("versions=%d err=%v", len(versions), err)
|
|
}
|
|
}
|