diff --git a/internal/data/config_store.go b/internal/data/config_store.go index 037ed57..8e5ee4d 100644 --- a/internal/data/config_store.go +++ b/internal/data/config_store.go @@ -382,6 +382,12 @@ func (d *Data) reloadConfig(ctx context.Context) error { } useRedis := next.Admin.System != nil && next.Admin.System.UseRedis candidateRedis := openRedis(next.Data.Redis, useRedis, d.logger()) + candidateRedisAccepted := false + defer func() { + if !candidateRedisAccepted && candidateRedis != nil { + _ = candidateRedis.Close() + } + }() useMongo := next.Admin.System != nil && next.Admin.System.UseMongo candidateMongo, mongoErr := openMongo(next.Data.Mongo, useMongo) if mongoErr != nil { @@ -397,6 +403,16 @@ func (d *Data) reloadConfig(ctx context.Context) error { if err != nil { return err } + candidateDBListAccepted := false + defer func() { + if !candidateDBListAccepted { + closeDatabaseList(candidateDBList) + } + }() + integrationConfigs, err := readIntegrationRuntime(candidateDB) + if err != nil { + return fmt.Errorf("reload integration runtime: %w", err) + } d.gormDB.replace(candidateDB, d.enqueueDataScopeAudit) d.databaseReady.Store(databaseReady) @@ -409,14 +425,16 @@ func (d *Data) reloadConfig(ctx context.Context) error { d.mongo.replace(candidateMongo) mongoAccepted = true } + closeCandidate = false + candidateDBListAccepted = true + candidateRedisAccepted = true d.runtime.Replace(next.Data, next.Admin) - if err = d.loadIntegrationRuntime(candidateDB); err != nil { - return fmt.Errorf("reload integration runtime: %w", err) + if d.integrations != nil { + d.integrations.Replace(integrationConfigs) } if d.storage != nil { d.storage.Replace(candidateStorage) } - closeCandidate = false return nil } diff --git a/internal/data/data.go b/internal/data/data.go index 00568e2..4ac5bd3 100644 --- a/internal/data/data.go +++ b/internal/data/data.go @@ -176,6 +176,34 @@ func NewData(runtime *conf.Runtime, appLogger *slog.Logger, storageManager *stor c.Database = &conf.Data_Database{} } d := &Data{runtime: runtime, integrations: runtimeconfig.NewStore(), appLogger: appLogger, storage: storageManager, catalog: catalog} + var stopConfigWatcher func() + var cleanupOnce sync.Once + cleanup := func() { + cleanupOnce.Do(func() { + if stopConfigWatcher != nil { + stopConfigWatcher() + } + if d.auditLog != nil { + d.auditLog.Close() + } + if d.gormDB != nil { + d.gormDB.close() + } + closeDatabaseList(d.dbList) + if d.redis != nil { + d.redis.close() + } + if d.mongo != nil { + d.mongo.close() + } + }) + } + initialized := false + defer func() { + if !initialized { + cleanup() + } + }() usingFallback := !databaseConnectionConfigured(c.Database) var db *gorm.DB var err error @@ -199,8 +227,6 @@ func NewData(runtime *conf.Runtime, appLogger *slog.Logger, storageManager *stor d.auditLog = newDataScopeAuditWriter(d, appLogger) d.dbList, err = openDatabaseList(c.DatabaseList, appLogger) if err != nil { - d.auditLog.Close() - d.gormDB.close() return nil, nil, err } for _, item := range d.dbList { @@ -256,15 +282,8 @@ func NewData(runtime *conf.Runtime, appLogger *slog.Logger, storageManager *stor mongoClient = nil } d.mongo = newReloadableMongo(mongoClient) - stopConfigWatcher := d.watchConfig() - cleanup := func() { - stopConfigWatcher() - d.auditLog.Close() - d.gormDB.close() - closeDatabaseList(d.dbList) - d.redis.close() - d.mongo.close() - } + stopConfigWatcher = d.watchConfig() + initialized = true return d, cleanup, nil } diff --git a/internal/data/initialization_backend.go b/internal/data/initialization_backend.go index de1f7fd..e0419af 100644 --- a/internal/data/initialization_backend.go +++ b/internal/data/initialization_backend.go @@ -209,7 +209,12 @@ func (d *Data) InitializeDatabase(ctx context.Context, input *system.DatabaseCon if err := d.persistDatabaseConfig(config, signingKey); err != nil { return fmt.Errorf("persist database configuration: %w", err) } + integrationConfigs, err := readIntegrationRuntime(candidate) + if err != nil { + return fmt.Errorf("initialize integration runtime: %w", err) + } d.activateDatabase(candidate, config) + activated = true currentData, currentAdmin := d.runtime.Values() if currentAdmin == nil { currentAdmin = &conf.AdminBackend{} @@ -221,9 +226,8 @@ func (d *Data) InitializeDatabase(ctx context.Context, input *system.DatabaseCon currentAdmin.Storage = storageConfig currentAdmin.Email = emailConfig d.runtime.Replace(currentData, currentAdmin) - if err = d.loadIntegrationRuntime(candidate); err != nil { - return fmt.Errorf("initialize integration runtime: %w", err) + if d.integrations != nil { + d.integrations.Replace(integrationConfigs) } - activated = true return nil } diff --git a/internal/data/integration/integration_config.go b/internal/data/integration/integration_config.go index 9f0dbc5..46589ee 100644 --- a/internal/data/integration/integration_config.go +++ b/internal/data/integration/integration_config.go @@ -27,10 +27,6 @@ func (ConfigPO) TableName() string { return "sys_integration_configs" } type integrationConfigRepo struct{ data Provider } -type integrationRuntimeProvider interface { - IntegrationRuntime() *runtimeconfig.Store -} - func NewIntegrationConfigRepo(data Provider) integrationbiz.IntegrationConfigRepo { return &integrationConfigRepo{data: data} } @@ -110,8 +106,8 @@ func (r *integrationConfigRepo) publish(kind, provider string, enabled bool, val } func integrationRuntime(provider Provider) *runtimeconfig.Store { - if value, ok := provider.(integrationRuntimeProvider); ok { - return value.IntegrationRuntime() + if provider != nil { + return provider.IntegrationRuntime() } return nil } diff --git a/internal/data/task/task.go b/internal/data/task/task.go index c462b30..5e112f5 100644 --- a/internal/data/task/task.go +++ b/internal/data/task/task.go @@ -127,6 +127,9 @@ func taskLogFromPO(v taskLogPO) *taskbiz.TimedTaskLog { return &taskbiz.TimedTaskLog{ID: v.ID, CreatedAt: v.CreatedAt, UpdatedAt: v.UpdatedAt, TaskID: v.TaskID, TaskName: v.TaskName, TriggerType: v.TriggerType, StartedAt: v.StartedAt, FinishedAt: v.FinishedAt, DurationMS: v.DurationMS, Status: v.Status, ErrorMsg: v.ErrorMsg, Output: v.Output} } func (r *taskRepo) ListTaskLogs(ctx context.Context, page, size int, taskID uint, status string) ([]*taskbiz.TimedTaskLog, int64, error) { + if !r.data.DatabaseReady() { + return []*taskbiz.TimedTaskLog{}, 0, nil + } db := r.data.DB().WithContext(ctx).Model(&taskLogPO{}) if taskID != 0 { db = db.Where("task_id = ?", taskID) diff --git a/internal/data/task/task_test.go b/internal/data/task/task_test.go index 3f456e6..6b951fe 100644 --- a/internal/data/task/task_test.go +++ b/internal/data/task/task_test.go @@ -15,3 +15,14 @@ func TestTaskRepoListIsEmptyBeforeDatabaseInitialization(t *testing.T) { t.Fatalf("bootstrap tasks = (%d, %d), want empty", total, len(items)) } } + +func TestTaskRepoLogsAreEmptyBeforeDatabaseInitialization(t *testing.T) { + repo := NewTaskRepo(&Data{}) + items, total, err := repo.ListTaskLogs(context.Background(), 0, 0, 0, "") + if err != nil { + t.Fatal(err) + } + if total != 0 || len(items) != 0 { + t.Fatalf("bootstrap task logs = (%d, %d), want empty", total, len(items)) + } +} diff --git a/internal/integration/cache/cache.go b/internal/integration/cache/cache.go index 67fc5fa..ef0116b 100644 --- a/internal/integration/cache/cache.go +++ b/internal/integration/cache/cache.go @@ -31,7 +31,12 @@ func New(provider RedisProvider) system.Cache { return &Store{provider: provider, memory: make(map[string]memoryEntry)} } -func (s *Store) client() redis.UniversalClient { return s.provider.RedisClient() } +func (s *Store) client() redis.UniversalClient { + if s == nil || s.provider == nil { + return nil + } + return s.provider.RedisClient() +} func (s *Store) Get(ctx context.Context, key string) (string, bool, error) { if client := s.client(); client != nil { @@ -101,6 +106,7 @@ func (s *Store) Increment(ctx context.Context, key string, expiration time.Durat entry, ok := s.memory[key] if ok && !entry.expiresAt.IsZero() && time.Now().After(entry.expiresAt) { ok = false + entry = memoryEntry{} } value := int64(0) if ok { @@ -108,8 +114,12 @@ func (s *Store) Increment(ctx context.Context, key string, expiration time.Durat } value++ entry.value = strconv.FormatInt(value, 10) - if !ok && expiration > 0 { - entry.expiresAt = time.Now().Add(expiration) + if !ok { + if expiration > 0 { + entry.expiresAt = time.Now().Add(expiration) + } else { + entry.expiresAt = time.Time{} + } } s.memory[key] = entry return value, nil diff --git a/internal/integration/cache/cache_test.go b/internal/integration/cache/cache_test.go new file mode 100644 index 0000000..5e56537 --- /dev/null +++ b/internal/integration/cache/cache_test.go @@ -0,0 +1,42 @@ +package cache + +import ( + "context" + "testing" + "time" +) + +func TestStoreWithoutRedisProviderUsesMemoryFallback(t *testing.T) { + store := New(nil) + ctx := context.Background() + + if err := store.Set(ctx, "key", "value", time.Minute); err != nil { + t.Fatalf("Set() error = %v", err) + } + if value, ok, err := store.Get(ctx, "key"); err != nil || !ok || value != "value" { + t.Fatalf("Get() = %q, %v, %v", value, ok, err) + } + if value, err := store.Increment(ctx, "counter", time.Minute); err != nil || value != 1 { + t.Fatalf("Increment() = %d, %v", value, err) + } + if err := store.Delete(ctx, "key"); err != nil { + t.Fatalf("Delete() error = %v", err) + } +} + +func TestStoreIncrementRecreatesExpiredKeyWithoutStaleExpiry(t *testing.T) { + store := New(nil) + ctx := context.Background() + + if _, err := store.Increment(ctx, "counter", time.Millisecond); err != nil { + t.Fatalf("initial Increment() error = %v", err) + } + time.Sleep(5 * time.Millisecond) + if value, err := store.Increment(ctx, "counter", 0); err != nil || value != 1 { + t.Fatalf("expired Increment() = %d, %v", value, err) + } + time.Sleep(5 * time.Millisecond) + if value, ok, err := store.Get(ctx, "counter"); err != nil || !ok || value != "1" { + t.Fatalf("Get() after recreation = %q, %v, %v", value, ok, err) + } +} diff --git a/internal/integration/email/email.go b/internal/integration/email/email.go index 41bb386..65e2f58 100644 --- a/internal/integration/email/email.go +++ b/internal/integration/email/email.go @@ -36,16 +36,30 @@ func (r *emailRepo) email() *conf.AdminBackend_Email { } func (r *emailRepo) Enabled() bool { - config := r.email() - return config != nil && config.Host != "" && config.From != "" && config.Secret != "" && config.Port > 0 + return emailEnabled(r.email()) } func (r *emailRepo) DefaultRecipients() []string { config := r.email() - if config == nil { + if config == nil || strings.TrimSpace(config.To) == "" { return nil } - return []string{config.To} + return []string{strings.TrimSpace(config.To)} +} + +func emailEnabled(config *conf.AdminBackend_Email) bool { + return config != nil && strings.TrimSpace(config.Host) != "" && + strings.TrimSpace(config.From) != "" && strings.TrimSpace(config.Secret) != "" && config.Port > 0 +} + +func normalizeRecipients(values []string) []string { + result := make([]string, 0, len(values)) + for _, value := range values { + if value = strings.TrimSpace(value); value != "" { + result = append(result, value) + } + } + return result } func cleanHeader(value string) string { @@ -53,19 +67,23 @@ func cleanHeader(value string) string { } func (r *emailRepo) Send(ctx context.Context, to []string, subject, body string) error { - if !r.Enabled() { + config := r.email() + if !emailEnabled(config) { return errors.New("邮件服务未配置") } + to = normalizeRecipients(to) if len(to) == 0 { return errors.New("收件人不能为空") } - config := r.email() + if ctx == nil { + ctx = context.Background() + } address := net.JoinHostPort(config.Host, fmt.Sprint(config.Port)) dialer := &net.Dialer{Timeout: 10 * time.Second} var conn net.Conn var err error if config.IsSsl { - conn, err = tls.DialWithDialer(dialer, "tcp", address, &tls.Config{ServerName: config.Host, MinVersion: tls.VersionTLS12}) + conn, err = (&tls.Dialer{NetDialer: dialer, Config: &tls.Config{ServerName: config.Host, MinVersion: tls.VersionTLS12}}).DialContext(ctx, "tcp", address) } else { conn, err = dialer.DialContext(ctx, "tcp", address) } diff --git a/internal/integration/email/email_test.go b/internal/integration/email/email_test.go new file mode 100644 index 0000000..9459b35 --- /dev/null +++ b/internal/integration/email/email_test.go @@ -0,0 +1,24 @@ +package email + +import ( + "context" + "testing" + + "kra/internal/conf" +) + +func TestDefaultRecipientsIgnoreEmptyConfiguration(t *testing.T) { + runtime := conf.NewRuntime(&conf.Data{}, &conf.AdminBackend{Email: &conf.AdminBackend_Email{To: " "}}) + repo := &emailRepo{runtime: runtime} + if got := repo.DefaultRecipients(); len(got) != 0 { + t.Fatalf("DefaultRecipients() = %#v, want empty", got) + } +} + +func TestSendRejectsEmptyRecipientsBeforeNetworkDial(t *testing.T) { + runtime := conf.NewRuntime(&conf.Data{}, &conf.AdminBackend{Email: &conf.AdminBackend_Email{Host: "smtp.example.com", From: "from@example.com", Secret: "secret", Port: 25}}) + repo := &emailRepo{runtime: runtime} + if err := repo.Send(context.Background(), []string{" ", ""}, "subject", "body"); err == nil { + t.Fatal("Send() accepted empty recipients") + } +} diff --git a/internal/integration/payment/adapter.go b/internal/integration/payment/adapter.go index 0047775..ab21059 100644 --- a/internal/integration/payment/adapter.go +++ b/internal/integration/payment/adapter.go @@ -4,6 +4,7 @@ import ( "context" "fmt" bizpayment "kra/internal/biz/payment" + "strings" ) // Adapter is the provider boundary used by the payment repository. Provider @@ -15,42 +16,38 @@ type Adapter interface { Callback(context.Context, *bizpayment.PaymentCallback, map[string]any) (*bizpayment.PaymentResult, error) } -// New constructs the SDK-backed adapter for a configured provider. +type adapterFactory func() Adapter + +// adapterFactories is the single payment integration registration point. +// Keep constructors zero-state: provider configuration belongs to each call, +// so an adapter can never accidentally retain secrets or order data. +var adapterFactories = map[string]adapterFactory{ + bizpayment.PaymentAlipay: func() Adapter { return &alipayAdapter{} }, + bizpayment.PaymentAlipayV3: func() Adapter { return &alipayV3Adapter{} }, + bizpayment.PaymentWechatV2: func() Adapter { return &wechatV2Adapter{} }, + bizpayment.PaymentWechatV3: func() Adapter { return &wechatV3Adapter{} }, + bizpayment.PaymentApple: func() Adapter { return &appleAdapter{} }, + bizpayment.PaymentDouyin: func() Adapter { return &douyinAdapter{} }, + bizpayment.PaymentQQ: func() Adapter { return &qqAdapter{} }, + bizpayment.PaymentAllinPay: func() Adapter { return &allinpayAdapter{} }, + bizpayment.PaymentLakala: func() Adapter { return &lakalaAdapter{} }, + bizpayment.PaymentPayPal: func() Adapter { return &paypalAdapter{} }, + bizpayment.PaymentSaobei: func() Adapter { return &saobeiAdapter{} }, + bizpayment.PaymentChinaums: func() Adapter { return newVendorAdapter(bizpayment.PaymentChinaums, vendorChinaums) }, + bizpayment.PaymentSFT: func() Adapter { return newVendorAdapter(bizpayment.PaymentSFT, vendorSFT) }, + bizpayment.PaymentSuperPay: func() Adapter { return newVendorAdapter(bizpayment.PaymentSuperPay, vendorSupperPay) }, + bizpayment.PaymentWechatGame: func() Adapter { return newVendorAdapter(bizpayment.PaymentWechatGame, vendorWechatGame) }, + bizpayment.PaymentDouyinGame: func() Adapter { return newVendorAdapter(bizpayment.PaymentDouyinGame, vendorDouyinGame) }, +} + +// New constructs the SDK-backed adapter for a configured provider. Provider +// identifiers are normalized at this I/O boundary so direct callers get the +// same behavior as the configuration and business layers. func New(provider string) (Adapter, error) { - switch provider { - case bizpayment.PaymentAlipay: - return &alipayAdapter{}, nil - case bizpayment.PaymentAlipayV3: - return &alipayV3Adapter{}, nil - case bizpayment.PaymentWechatV2: - return &wechatV2Adapter{}, nil - case bizpayment.PaymentWechatV3: - return &wechatV3Adapter{}, nil - case bizpayment.PaymentApple: - return &appleAdapter{}, nil - case bizpayment.PaymentDouyin: - return &douyinAdapter{}, nil - case bizpayment.PaymentQQ: - return &qqAdapter{}, nil - case bizpayment.PaymentAllinPay: - return &allinpayAdapter{}, nil - case bizpayment.PaymentLakala: - return &lakalaAdapter{}, nil - case bizpayment.PaymentPayPal: - return &paypalAdapter{}, nil - case bizpayment.PaymentSaobei: - return &saobeiAdapter{}, nil - case bizpayment.PaymentChinaums: - return newVendorAdapter(provider, vendorChinaums), nil - case bizpayment.PaymentSFT: - return newVendorAdapter(provider, vendorSFT), nil - case bizpayment.PaymentSuperPay: - return newVendorAdapter(provider, vendorSupperPay), nil - case bizpayment.PaymentWechatGame: - return newVendorAdapter(provider, vendorWechatGame), nil - case bizpayment.PaymentDouyinGame: - return newVendorAdapter(provider, vendorDouyinGame), nil - default: + provider = strings.ToLower(strings.TrimSpace(provider)) + factory, ok := adapterFactories[provider] + if !ok { return nil, fmt.Errorf("支付渠道 %s 没有适配器", provider) } + return factory(), nil } diff --git a/internal/integration/payment/adapter_test.go b/internal/integration/payment/adapter_test.go index 52cdcb3..a537f6b 100644 --- a/internal/integration/payment/adapter_test.go +++ b/internal/integration/payment/adapter_test.go @@ -18,3 +18,40 @@ func TestEverySupportedProviderHasAdapter(t *testing.T) { }) } } + +func TestNewNormalizesProviderIdentifier(t *testing.T) { + for _, provider := range bizpayment.SupportedPaymentProviders { + adapter, err := New(" " + provider + " ") + if err != nil { + t.Fatalf("New(%q): %v", provider, err) + } + if adapter == nil { + t.Fatalf("New(%q) returned nil", provider) + } + } +} + +func TestNewRejectsUnknownProvider(t *testing.T) { + if _, err := New("unknown-provider"); err == nil { + t.Fatal("unknown provider unexpectedly accepted") + } +} + +func TestNonceIsUniqueForConcurrentRequests(t *testing.T) { + const count = 1000 + values := make(chan string, count) + for i := 0; i < count; i++ { + go func() { values <- nonce() }() + } + seen := make(map[string]struct{}, count) + for i := 0; i < count; i++ { + value := <-values + if value == "" { + t.Fatal("nonce returned an empty value") + } + if _, exists := seen[value]; exists { + t.Fatalf("nonce collision for %q", value) + } + seen[value] = struct{}{} + } +} diff --git a/internal/integration/payment/wechat_v2.go b/internal/integration/payment/wechat_v2.go index 89a8ca0..3d72762 100644 --- a/internal/integration/payment/wechat_v2.go +++ b/internal/integration/payment/wechat_v2.go @@ -3,8 +3,10 @@ package payment import ( "bytes" "context" + "crypto/rand" "crypto/tls" "crypto/x509" + "encoding/hex" "encoding/json" "errors" "fmt" @@ -570,7 +572,15 @@ func paymentHTTPClient(c map[string]any) (*http.Client, error) { return &http.Client{Timeout: 20 * time.Second, Transport: &http.Transport{TLSClientConfig: &tls.Config{Certificates: []tls.Certificate{cert}, RootCAs: pool, MinVersion: tls.VersionTLS12}}}, nil } -func nonce() string { return fmt.Sprintf("%d", time.Now().UnixNano()) } +// nonce returns a compact request token. Clock-only values can collide when +// several payment requests are prepared in the same scheduler tick. +func nonce() string { + var raw [8]byte + if _, err := rand.Read(raw[:]); err == nil { + return hex.EncodeToString(raw[:]) + } + return fmt.Sprintf("%d", time.Now().UnixNano()) +} func stringOr(extra map[string]any, key, fallback string) string { if value, ok := extra[key].(string); ok { return value diff --git a/internal/integration/runtimeconfig/store.go b/internal/integration/runtimeconfig/store.go index 4d399ac..efa4553 100644 --- a/internal/integration/runtimeconfig/store.go +++ b/internal/integration/runtimeconfig/store.go @@ -3,6 +3,7 @@ package runtimeconfig import ( + "bytes" "encoding/json" "strings" "sync" @@ -41,6 +42,13 @@ func cloneConfig(config Config) Config { return config } +func sameConfig(left, right Config) bool { + return left.Kind == right.Kind && + left.Provider == right.Provider && + left.Enabled == right.Enabled && + bytes.Equal(left.Values, right.Values) +} + func (s *Store) Get(kind, provider string) (Config, bool) { if s == nil { return Config{}, false @@ -105,11 +113,20 @@ func (s *Store) Replace(configs []Config) { s.mu.Unlock() changed := make(map[string]Config, len(previous)+len(next)) - for key, config := range previous { - changed[key] = Config{Kind: config.Kind, Provider: config.Provider} + for key, previousConfig := range previous { + nextConfig, exists := next[key] + if !exists { + changed[key] = Config{Kind: previousConfig.Kind, Provider: previousConfig.Provider} + continue + } + if !sameConfig(previousConfig, nextConfig) { + changed[key] = nextConfig + } } - for key, config := range next { - changed[key] = config + for key, nextConfig := range next { + if _, exists := previous[key]; !exists { + changed[key] = nextConfig + } } for _, item := range listeners { if config, ok := changed[configKey(item.kind, item.provider)]; ok { diff --git a/internal/integration/runtimeconfig/store_test.go b/internal/integration/runtimeconfig/store_test.go index 964da35..8327ab1 100644 --- a/internal/integration/runtimeconfig/store_test.go +++ b/internal/integration/runtimeconfig/store_test.go @@ -3,6 +3,7 @@ package runtimeconfig import ( "encoding/json" "testing" + "time" ) func TestStoreSetDeleteAndSubscribe(t *testing.T) { @@ -28,3 +29,37 @@ func TestStoreSetDeleteAndSubscribe(t *testing.T) { t.Fatalf("delete update = %#v", update) } } + +func TestStoreReplaceSkipsUnchangedValues(t *testing.T) { + store := NewStore() + updates := make(chan Config, 2) + stop := store.Subscribe("mq", "rabbitmq", func(config Config) { updates <- config }) + defer stop() + + config := Config{Kind: "mq", Provider: "rabbitmq", Enabled: true, Values: json.RawMessage(`{"host":"localhost"}`)} + store.Set(config) + select { + case <-updates: + case <-time.After(time.Second): + t.Fatal("initial set notification was not delivered") + } + + store.Replace([]Config{config}) + select { + case update := <-updates: + t.Fatalf("unchanged replace emitted notification: %#v", update) + case <-time.After(20 * time.Millisecond): + } + + changed := config + changed.Values = json.RawMessage(`{"host":"other"}`) + store.Replace([]Config{changed}) + select { + case update := <-updates: + if string(update.Values) != string(changed.Values) { + t.Fatalf("changed replace = %#v", update) + } + case <-time.After(time.Second): + t.Fatal("changed replace notification was not delivered") + } +} diff --git a/internal/integration/storage/aliyun_storage.go b/internal/integration/storage/aliyun_storage.go index 8b7dc4b..aedee5c 100644 --- a/internal/integration/storage/aliyun_storage.go +++ b/internal/integration/storage/aliyun_storage.go @@ -69,6 +69,10 @@ func (s *aliyunStorage) Compose(ctx context.Context, names []string, destination return composeFiles(ctx, s, names, destination) } func (s *aliyunStorage) DeletePrefix(ctx context.Context, prefix string) error { + prefix, err := normalizeDeletePrefix(prefix) + if err != nil { + return err + } cursor := "" for { items, next, more, err := s.List(ctx, prefix, cursor, 1000) @@ -83,7 +87,10 @@ func (s *aliyunStorage) DeletePrefix(ctx context.Context, prefix string) error { if !more { return nil } - cursor = next + cursor, err = advanceDeletePrefixCursor(cursor, next, more) + if err != nil { + return err + } } } func (s *aliyunStorage) List(_ context.Context, prefix, cursor string, limit int) ([]*system.StoredFile, string, bool, error) { diff --git a/internal/integration/storage/aws_storage.go b/internal/integration/storage/aws_storage.go index 68d3a67..881066d 100644 --- a/internal/integration/storage/aws_storage.go +++ b/internal/integration/storage/aws_storage.go @@ -106,8 +106,13 @@ func (s *awsStorage) Compose(ctx context.Context, names []string, destination st return composeFiles(ctx, s, names, destination) } func (s *awsStorage) DeletePrefix(ctx context.Context, prefix string) error { + prefix, err := normalizeDeletePrefix(prefix) + if err != nil { + return err + } + cursor := "" for { - items, _, more, err := s.List(ctx, prefix, "", 1000) + items, next, more, err := s.List(ctx, prefix, cursor, 1000) if err != nil { return err } @@ -116,9 +121,13 @@ func (s *awsStorage) DeletePrefix(ctx context.Context, prefix string) error { return err } } - if !more || len(items) == 0 { + if !more { return nil } + cursor, err = advanceDeletePrefixCursor(cursor, next, more) + if err != nil { + return err + } } } func (s *awsStorage) List(ctx context.Context, prefix, cursor string, limit int) ([]*system.StoredFile, string, bool, error) { diff --git a/internal/integration/storage/compose.go b/internal/integration/storage/compose.go index ab63aca..a79710a 100644 --- a/internal/integration/storage/compose.go +++ b/internal/integration/storage/compose.go @@ -43,6 +43,14 @@ func composeStreams( errCh <- nil }() putErr := put(ctx, destination, reader) + // A failed destination may stop reading before the producer reaches EOF. + // Close the pipe reader in that case so the producer's next write returns + // instead of leaving the goroutine blocked forever. + if putErr != nil { + _ = reader.CloseWithError(putErr) + } else { + _ = reader.Close() + } composeErr := <-errCh if putErr != nil { _ = remove(ctx, destination) diff --git a/internal/integration/storage/compose_test.go b/internal/integration/storage/compose_test.go index 8659073..75db552 100644 --- a/internal/integration/storage/compose_test.go +++ b/internal/integration/storage/compose_test.go @@ -8,6 +8,7 @@ import ( "io" "strings" "testing" + "time" ) func TestComposeStreams(t *testing.T) { @@ -48,3 +49,62 @@ func TestComposeStreamsRemovesPartialDestination(t *testing.T) { t.Fatalf("compose error = %v, removed = %v", err, removed) } } + +func TestComposeStreamsUnblocksProducerWhenDestinationFails(t *testing.T) { + putErr := errors.New("destination failed") + done := make(chan error, 1) + go func() { + _, err := composeStreams(context.Background(), []string{"large"}, func(context.Context, string) (io.ReadCloser, error) { + return io.NopCloser(strings.NewReader(strings.Repeat("x", 1<<20))), nil + }, func(context.Context, string, io.Reader) error { + return putErr + }, func(context.Context, string) error { return nil }, "out") + done <- err + }() + + select { + case err := <-done: + if !errors.Is(err, putErr) { + t.Fatalf("compose error = %v, want %v", err, putErr) + } + case <-time.After(time.Second): + t.Fatal("compose remained blocked after destination failure") + } +} + +func TestNormalizeDeletePrefix(t *testing.T) { + tests := []struct { + name string + value string + valid string + wantErr bool + }{ + {name: "canonicalizes", value: "/uploads/chunks/", valid: "uploads/chunks"}, + {name: "cleans duplicate separators", value: "uploads//chunks", valid: "uploads/chunks"}, + {name: "rejects empty", value: " / ", wantErr: true}, + {name: "rejects traversal", value: "uploads/../", wantErr: true}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := normalizeDeletePrefix(tt.value) + if tt.wantErr { + if err == nil { + t.Fatalf("normalizeDeletePrefix(%q) succeeded with %q", tt.value, got) + } + return + } + if err != nil || got != tt.valid { + t.Fatalf("normalizeDeletePrefix(%q) = %q, %v; want %q", tt.value, got, err, tt.valid) + } + }) + } +} + +func TestAdvanceDeletePrefixCursor(t *testing.T) { + if _, err := advanceDeletePrefixCursor("cursor", "cursor", true); err == nil { + t.Fatal("same cursor should fail") + } + if next, err := advanceDeletePrefixCursor("", "next", true); err != nil || next != "next" { + t.Fatalf("advance cursor = %q, %v", next, err) + } +} diff --git a/internal/integration/storage/huawei_storage.go b/internal/integration/storage/huawei_storage.go index 7c6d871..26df589 100644 --- a/internal/integration/storage/huawei_storage.go +++ b/internal/integration/storage/huawei_storage.go @@ -68,6 +68,10 @@ func (s *huaweiStorage) Compose(ctx context.Context, names []string, destination return composeFiles(ctx, s, names, destination) } func (s *huaweiStorage) DeletePrefix(ctx context.Context, prefix string) error { + prefix, err := normalizeDeletePrefix(prefix) + if err != nil { + return err + } cursor := "" for { items, next, more, err := s.List(ctx, prefix, cursor, 1000) @@ -82,7 +86,10 @@ func (s *huaweiStorage) DeletePrefix(ctx context.Context, prefix string) error { if !more { return nil } - cursor = next + cursor, err = advanceDeletePrefixCursor(cursor, next, more) + if err != nil { + return err + } } } func (s *huaweiStorage) List(_ context.Context, prefix, cursor string, limit int) ([]*system.StoredFile, string, bool, error) { diff --git a/internal/integration/storage/local.go b/internal/integration/storage/local.go index a4e1721..fa49bc2 100644 --- a/internal/integration/storage/local.go +++ b/internal/integration/storage/local.go @@ -138,7 +138,11 @@ func (s *fileStorage) Compose(ctx context.Context, names []string, destination s return composeFiles(ctx, s, names, destination) } func (s *fileStorage) DeletePrefix(ctx context.Context, prefix string) error { - path, err := s.resolve(strings.TrimSuffix(prefix, "/") + "/placeholder") + prefix, err := normalizeDeletePrefix(prefix) + if err != nil { + return err + } + path, err := s.resolve(prefix + "/placeholder") if err != nil { return err } diff --git a/internal/integration/storage/local_test.go b/internal/integration/storage/local_test.go new file mode 100644 index 0000000..01ba80e --- /dev/null +++ b/internal/integration/storage/local_test.go @@ -0,0 +1,24 @@ +package storage + +import ( + "context" + "os" + "path/filepath" + "testing" +) + +func TestLocalDeletePrefixRejectsEmptyPrefix(t *testing.T) { + root := t.TempDir() + sentinel := filepath.Join(root, "keep.txt") + if err := os.WriteFile(sentinel, []byte("keep"), 0o600); err != nil { + t.Fatalf("write sentinel: %v", err) + } + storage := &fileStorage{root: root, urlPrefix: "/files"} + + if err := storage.DeletePrefix(context.Background(), " /"); err == nil { + t.Fatal("DeletePrefix() accepted an empty prefix") + } + if _, err := os.Stat(sentinel); err != nil { + t.Fatalf("sentinel was removed: %v", err) + } +} diff --git a/internal/integration/storage/prefix.go b/internal/integration/storage/prefix.go new file mode 100644 index 0000000..19b7628 --- /dev/null +++ b/internal/integration/storage/prefix.go @@ -0,0 +1,41 @@ +package storage + +import ( + "errors" + "path" + "strings" +) + +// normalizeDeletePrefix rejects requests that could accidentally target the +// storage root and returns the canonical key form used by all backends. +func normalizeDeletePrefix(prefix string) (string, error) { + prefix = strings.TrimSpace(strings.ReplaceAll(prefix, "\\", "/")) + prefix = strings.Trim(prefix, "/") + if prefix == "" { + return "", errors.New("storage delete prefix is required") + } + for _, part := range strings.Split(prefix, "/") { + if part == ".." { + return "", errors.New("invalid storage delete prefix") + } + } + clean := path.Clean(prefix) + if clean == "." || clean == ".." || strings.HasPrefix(clean, "../") { + return "", errors.New("invalid storage delete prefix") + } + return clean, nil +} + +// advanceDeletePrefixCursor prevents a backend that reports a truncated page +// without a new cursor from making DeletePrefix loop forever. +func advanceDeletePrefixCursor(current, next string, more bool) (string, error) { + if !more { + return "", nil + } + current = strings.TrimSpace(current) + next = strings.TrimSpace(next) + if next == "" || next == current { + return "", errors.New("storage delete prefix pagination made no progress") + } + return next, nil +} diff --git a/internal/integration/storage/qiniu_storage.go b/internal/integration/storage/qiniu_storage.go index 0b9ca73..4c6f693 100644 --- a/internal/integration/storage/qiniu_storage.go +++ b/internal/integration/storage/qiniu_storage.go @@ -96,8 +96,13 @@ func (s *qiniuStorage) Compose(ctx context.Context, names []string, destination return composeFiles(ctx, s, names, destination) } func (s *qiniuStorage) DeletePrefix(ctx context.Context, prefix string) error { + prefix, err := normalizeDeletePrefix(prefix) + if err != nil { + return err + } + cursor := "" for { - items, _, more, err := s.List(ctx, prefix, "", 1000) + items, next, more, err := s.List(ctx, prefix, cursor, 1000) if err != nil { return err } @@ -109,6 +114,10 @@ func (s *qiniuStorage) DeletePrefix(ctx context.Context, prefix string) error { if !more { return nil } + cursor, err = advanceDeletePrefixCursor(cursor, next, more) + if err != nil { + return err + } } } func (s *qiniuStorage) List(ctx context.Context, prefix, cursor string, limit int) ([]*system.StoredFile, string, bool, error) { diff --git a/internal/integration/storage/s3_storage.go b/internal/integration/storage/s3_storage.go index 53ad0c0..344c4c9 100644 --- a/internal/integration/storage/s3_storage.go +++ b/internal/integration/storage/s3_storage.go @@ -98,6 +98,10 @@ func (s *s3Storage) Compose(ctx context.Context, names []string, destination str return composeFiles(ctx, s, names, destination) } func (s *s3Storage) DeletePrefix(ctx context.Context, prefix string) error { + prefix, err := normalizeDeletePrefix(prefix) + if err != nil { + return err + } items := s.client.ListObjects(ctx, s.bucket, minio.ListObjectsOptions{Prefix: s.key(prefix), Recursive: true}) for item := range items { if item.Err != nil { diff --git a/internal/integration/storage/tencent_storage.go b/internal/integration/storage/tencent_storage.go index 7cdb32a..a71e85b 100644 --- a/internal/integration/storage/tencent_storage.go +++ b/internal/integration/storage/tencent_storage.go @@ -82,6 +82,10 @@ func (s *tencentStorage) Compose(ctx context.Context, names []string, destinatio return composeFiles(ctx, s, names, destination) } func (s *tencentStorage) DeletePrefix(ctx context.Context, prefix string) error { + prefix, err := normalizeDeletePrefix(prefix) + if err != nil { + return err + } cursor := "" for { items, next, more, err := s.List(ctx, prefix, cursor, 1000) @@ -96,7 +100,10 @@ func (s *tencentStorage) DeletePrefix(ctx context.Context, prefix string) error if !more { return nil } - cursor = next + cursor, err = advanceDeletePrefixCursor(cursor, next, more) + if err != nil { + return err + } } } func (s *tencentStorage) List(ctx context.Context, prefix, cursor string, limit int) ([]*system.StoredFile, string, bool, error) {