diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go index 51bcaaf..10aebd0 100644 --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -1,6 +1,7 @@ package config import ( + "fmt" "os" "path/filepath" "strings" @@ -45,5 +46,52 @@ func LoadFromDir(configDir string) (Config, error) { if cfg.MaxUploadBytes <= 0 { cfg.MaxUploadBytes = 32 << 20 } + if strings.TrimSpace(cfg.StorageDir) == "" { + return Config{}, fmt.Errorf("storage_dir must not be empty") + } + hasAllowedOrigin := false + for _, origin := range cfg.AllowedOrigins { + if strings.TrimSpace(origin) != "" { + hasAllowedOrigin = true + break + } + } + if !hasAllowedOrigin { + return Config{}, fmt.Errorf("allowed_origins must include at least one origin") + } + if err := validateSecret(cfg.Env, "auth_secret", cfg.AuthSecret); err != nil { + return Config{}, err + } + if err := validateSecret(cfg.Env, "ai_key_encryption_secret", cfg.AIKeyEncryptionSecret); err != nil { + return Config{}, err + } return cfg, nil } + +func validateSecret(environment, field, value string) error { + secret := strings.TrimSpace(value) + if secret == "" { + return fmt.Errorf("%s must not be empty", field) + } + if strings.EqualFold(strings.TrimSpace(environment), "production") || strings.EqualFold(strings.TrimSpace(environment), "prod") { + if len(secret) < 32 || isCommonSecret(secret) { + return fmt.Errorf("%s must be at least 32 characters and must not use a development sentinel in production", field) + } + } + return nil +} + +func isCommonSecret(value string) bool { + normalized := strings.ToLower(strings.TrimSpace(value)) + for _, marker := range []string{"change-me", "changeme", "development", "dev-secret", "local-secret", "test-secret", "placeholder"} { + if strings.Contains(normalized, marker) { + return true + } + } + switch normalized { + case "secret", "password", "default", "admin": + return true + default: + return false + } +} diff --git a/backend/internal/config/config_test.go b/backend/internal/config/config_test.go index ad90a9c..e7e8cd0 100644 --- a/backend/internal/config/config_test.go +++ b/backend/internal/config/config_test.go @@ -3,11 +3,38 @@ package config import ( "os" "path/filepath" + "strings" "testing" "github.com/stretchr/testify/require" ) +func TestDevelopmentConfigMatchesLocalWorkspaceContract(t *testing.T) { + t.Setenv("SENLIN_APP_MODE", "dev") + + cfg, err := LoadFromDir(filepath.Join("..", "..", "etc")) + + require.NoError(t, err) + require.Equal(t, "development", cfg.Env) + require.Equal(t, "9150", cfg.Port) + require.Equal(t, "postgres://postgres:postgres@localhost:5432/agent_dev?sslmode=disable", cfg.DSN) + require.NotEmpty(t, strings.TrimSpace(cfg.StorageDir)) + require.Equal(t, int64(32<<20), cfg.MaxUploadBytes) + require.ElementsMatch(t, []string{ + "http://localhost:5173", + "http://127.0.0.1:5173", + "http://localhost:4173", + "http://127.0.0.1:4173", + "http://localhost:4174", + "http://127.0.0.1:4174", + "http://localhost:4175", + "http://127.0.0.1:4175", + "http://localhost:4176", + "http://127.0.0.1:4176", + "http://tauri.localhost", + }, cfg.AllowedOrigins) +} + func TestLoadFromDirDefaultsToDevYAML(t *testing.T) { configDir := t.TempDir() writeConfig(t, configDir, "agent.dev.yaml", "development", "18080", "postgres://dev", "./dev-files", "dev-auth", "dev-system", "dev-ai") @@ -31,7 +58,7 @@ func TestLoadFromDirDefaultsToDevYAML(t *testing.T) { func TestLoadFromDirUsesSENLINAppMode(t *testing.T) { configDir := t.TempDir() writeConfig(t, configDir, "agent.dev.yaml", "development", "18080", "postgres://dev", "./dev-files", "dev-auth", "", "dev-ai") - writeConfig(t, configDir, "agent.prod.yaml", "production", "80", "postgres://prod", "/data/files", "prod-auth", "prod-system", "prod-ai") + writeConfig(t, configDir, "agent.prod.yaml", "production", "80", "postgres://prod", "/data/files", "production-auth-signing-key-2026-safe", "prod-system", "production-ai-encryption-key-2026-safe") t.Setenv("SENLIN_APP_MODE", "prod") cfg, err := LoadFromDir(configDir) @@ -41,12 +68,98 @@ func TestLoadFromDirUsesSENLINAppMode(t *testing.T) { require.Equal(t, "80", cfg.Port) require.Equal(t, "postgres://prod", cfg.DSN) require.Equal(t, "/data/files", cfg.StorageDir) - require.Equal(t, "prod-auth", cfg.AuthSecret) + require.Equal(t, "production-auth-signing-key-2026-safe", cfg.AuthSecret) require.Equal(t, "prod-system", cfg.SystemAIKey) - require.Equal(t, "prod-ai", cfg.AIKeyEncryptionSecret) + require.Equal(t, "production-ai-encryption-key-2026-safe", cfg.AIKeyEncryptionSecret) require.Equal(t, []string{"https://workbench.example.com"}, cfg.AllowedOrigins) } +func TestLoadFromDirRejectsMissingStorageDir(t *testing.T) { + configDir := t.TempDir() + writeConfig(t, configDir, "agent.dev.yaml", "development", "9150", "postgres://agent", "", "dev-auth", "", "dev-ai") + t.Setenv("SENLIN_APP_MODE", "dev") + + _, err := LoadFromDir(configDir) + + require.ErrorContains(t, err, "storage_dir") +} + +func TestLoadFromDirRejectsMissingAllowedOrigins(t *testing.T) { + configDir := t.TempDir() + content := []byte("env: development\n" + + "port: \"9150\"\n" + + "dsn: \"postgres://agent\"\n" + + "storage_dir: \"./data/files\"\n" + + "auth_secret: \"dev-auth\"\n" + + "ai_key_encryption_secret: \"dev-ai\"\n") + require.NoError(t, os.WriteFile(filepath.Join(configDir, "agent.dev.yaml"), content, 0o600)) + t.Setenv("SENLIN_APP_MODE", "dev") + + _, err := LoadFromDir(configDir) + + require.ErrorContains(t, err, "allowed_origins") +} + +func TestLoadFromDirRejectsUnsafeProductionSecrets(t *testing.T) { + tests := []struct { + name string + authSecret string + encryptionSecret string + expectedFieldName string + }{ + {name: "empty auth", authSecret: "", encryptionSecret: "production-ai-encryption-key-2026-safe", expectedFieldName: "auth_secret"}, + {name: "short auth", authSecret: "short", encryptionSecret: "production-ai-encryption-key-2026-safe", expectedFieldName: "auth_secret"}, + {name: "development auth sentinel", authSecret: "development-auth-secret-change-me", encryptionSecret: "production-ai-encryption-key-2026-safe", expectedFieldName: "auth_secret"}, + {name: "dev auth sentinel", authSecret: "dev-secret-dev-secret-dev-secret-000", encryptionSecret: "production-ai-encryption-key-2026-safe", expectedFieldName: "auth_secret"}, + {name: "empty encryption", authSecret: "production-auth-signing-key-2026-safe", encryptionSecret: "", expectedFieldName: "ai_key_encryption_secret"}, + {name: "short encryption", authSecret: "production-auth-signing-key-2026-safe", encryptionSecret: "short", expectedFieldName: "ai_key_encryption_secret"}, + {name: "common encryption sentinel", authSecret: "production-auth-signing-key-2026-safe", encryptionSecret: "change-me-change-me-change-me-change-me", expectedFieldName: "ai_key_encryption_secret"}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + configDir := t.TempDir() + writeConfig(t, configDir, "agent.prod.yaml", "production", "80", "postgres://prod", "/data/files", test.authSecret, "", test.encryptionSecret) + t.Setenv("SENLIN_APP_MODE", "prod") + + _, err := LoadFromDir(configDir) + + require.ErrorContains(t, err, test.expectedFieldName) + }) + } +} + +func TestLoadFromDirRequiresExplicitDevelopmentSecrets(t *testing.T) { + for _, missingField := range []string{"auth", "encryption"} { + t.Run(missingField, func(t *testing.T) { + configDir := t.TempDir() + authSecret, encryptionSecret := "development-auth", "development-encryption" + if missingField == "auth" { + authSecret = "" + } else { + encryptionSecret = "" + } + writeConfig(t, configDir, "agent.dev.yaml", "development", "9150", "postgres://dev", "./files", authSecret, "", encryptionSecret) + t.Setenv("SENLIN_APP_MODE", "dev") + + _, err := LoadFromDir(configDir) + + require.Error(t, err) + }) + } +} + +func TestLoadFromDirAllowsExplicitTestSecrets(t *testing.T) { + configDir := t.TempDir() + writeConfig(t, configDir, "agent.test.yaml", "test", "9150", "postgres://test", "./files", "test-auth", "", "test-ai") + t.Setenv("SENLIN_APP_MODE", "test") + + cfg, err := LoadFromDir(configDir) + + require.NoError(t, err) + require.Equal(t, "test-auth", cfg.AuthSecret) + require.Equal(t, "test-ai", cfg.AIKeyEncryptionSecret) +} + func writeConfig(t *testing.T, dir string, name string, env string, port string, databaseURL string, storageDir string, authSecret string, systemAIKey string, aiKeySecret string) { t.Helper() allowedOrigins := " - http://localhost:5173\n - http://tauri.localhost\n" diff --git a/backend/internal/httpx/router_integration_test.go b/backend/internal/httpx/router_integration_test.go index cf11306..5cb34fd 100644 --- a/backend/internal/httpx/router_integration_test.go +++ b/backend/internal/httpx/router_integration_test.go @@ -6,24 +6,36 @@ import ( "github.com/stretchr/testify/require" "senlinai-agent/backend/internal/config" "senlinai-agent/backend/internal/httpx" + "senlinai-agent/backend/internal/logic/ai" + "senlinai-agent/backend/internal/logic/auth" "senlinai-agent/backend/internal/logic/files" "senlinai-agent/backend/internal/logic/inbox" "senlinai-agent/backend/internal/logic/projects" + "senlinai-agent/backend/internal/logic/search" "senlinai-agent/backend/internal/logic/tasks" ) -func TestBackendFeatureRoutesRegisterTogether(t *testing.T) { +func TestMainRegistrarsRegisterTogether(t *testing.T) { require.NotPanics(t, func() { + cfg := config.Config{ + Env: "test", + AuthSecret: "test-auth-secret", + AIKeyEncryptionSecret: "test-ai-encryption-secret", + } + authService := auth.NewService(cfg.AuthSecret) projectService := projects.NewService() httpx.NewProtectedRouter( - config.Config{Env: "test"}, - func(token string) (uint, error) { return 1, nil }, + cfg, + authService.VerifySession, + auth.NewHandler(authService), projects.NewHandler(projectService), projects.NewTagHandler(projectService), projects.NewCronHandler(projectService), tasks.NewHandler(tasks.NewService()), files.NewHandler(files.NewService(t.TempDir())), inbox.NewHandler(inbox.NewService(inbox.StaticAnalyzer{})), + search.NewHandler(search.NewService(nil)), + ai.NewHandler(ai.NewSessionService(ai.NewGatewayWithSecret(cfg.SystemAIKey, cfg.AIKeyEncryptionSecret))), ) }) } diff --git a/backend/internal/logic/ai/handlers.go b/backend/internal/logic/ai/handlers.go index d31d9aa..d81df19 100644 --- a/backend/internal/logic/ai/handlers.go +++ b/backend/internal/logic/ai/handlers.go @@ -39,8 +39,8 @@ func NewHandler(service *SessionService) *Handler { } func (h *Handler) Register(router gin.IRouter) { - router.GET("/projects/:projectId/ai-sessions", h.list) - router.POST("/projects/:projectId/ai-sessions", h.create) + router.GET("/projects/:id/ai-sessions", h.list) + router.POST("/projects/:id/ai-sessions", h.create) } func (h *Handler) list(c *gin.Context) { @@ -84,7 +84,7 @@ func aiRequestContext(c *gin.Context) (uint, string, bool) { httpx.Error(c, http.StatusUnauthorized, "unauthorized", "未登录或登录已失效") return 0, "", false } - projectIdentity, ok := httpx.IdentityParam(c, "projectId") + projectIdentity, ok := httpx.IdentityParam(c, "id") if !ok { return 0, "", false } diff --git a/backend/internal/logic/ai/rate_limit_postgres_test.go b/backend/internal/logic/ai/rate_limit_postgres_test.go index 9252b77..55e589e 100644 --- a/backend/internal/logic/ai/rate_limit_postgres_test.go +++ b/backend/internal/logic/ai/rate_limit_postgres_test.go @@ -1,3 +1,5 @@ +//go:build integration + package ai import ( @@ -16,10 +18,8 @@ import ( ) func TestPostgresReserveRateLimitIsAtomicAcrossConcurrentConnections(t *testing.T) { - dsn := os.Getenv("DATABASE_URL") - if dsn == "" { - t.Skip("DATABASE_URL is not configured; skipping PostgreSQL AI rate reservation test") - } + dsn := os.Getenv("TEST_DATABASE_URL") + require.NotEmpty(t, dsn, "TEST_DATABASE_URL is required for integration tests and must point to an isolated database") database, err := gorm.Open(postgres.Open(dsn), &gorm.Config{TranslateError: true, Logger: logger.Default.LogMode(logger.Silent)}) require.NoError(t, err) require.NoError(t, models.AutoMigrate(database)) diff --git a/backend/internal/logic/files/handlers.go b/backend/internal/logic/files/handlers.go index 244e659..45cf15b 100644 --- a/backend/internal/logic/files/handlers.go +++ b/backend/internal/logic/files/handlers.go @@ -4,6 +4,7 @@ import ( "errors" "log" "net/http" + "path/filepath" "time" "github.com/gin-gonic/gin" @@ -14,15 +15,15 @@ import ( "senlinai-agent/backend/internal/models" ) -// SourceDTO 是文件资料写接口的稳定响应,仅包含可公开的相对存储路径。 +// SourceDTO 是文件资料写接口的稳定响应,仅返回 opaque storage key,不暴露物理目录布局。 type SourceDTO struct { - ID string `json:"id"` - ProjectID string `json:"projectId"` - Kind string `json:"kind"` - Title string `json:"title"` - FilePath string `json:"filePath"` - CreatedAt time.Time `json:"createdAt"` - UpdatedAt time.Time `json:"updatedAt"` + ID string `json:"id"` + ProjectID string `json:"projectId"` + Kind string `json:"kind"` + Title string `json:"title"` + StorageKey string `json:"storageKey"` + CreatedAt time.Time `json:"createdAt"` + UpdatedAt time.Time `json:"updatedAt"` } // Handler 负责 multipart 边界,文件名清理和本地路径构造始终委托给 Service。 @@ -84,7 +85,7 @@ func (h *Handler) upload(c *gin.Context) { defer file.Close() // 先由文件服务保存内容,再用其返回的相对路径创建资料记录;handler 不拼接任何本地路径。 - stored, err := h.service.Save(project.ID, fileHeader.Filename, file) + stored, err := h.service.Save(project.Identity, fileHeader.Filename, file) if err != nil { httpx.Error(c, http.StatusInternalServerError, "internal_error", "文件保存失败") return @@ -104,7 +105,7 @@ func (h *Handler) upload(c *gin.Context) { func sourceDTO(source models.SenlinAgentSource) SourceDTO { return SourceDTO{ ID: source.Identity, ProjectID: source.ProjectIdentity, Kind: source.Kind, - Title: source.Title, FilePath: source.FilePath, + Title: source.Title, StorageKey: filepath.Base(filepath.FromSlash(source.FilePath)), CreatedAt: source.CreatedAt.UTC(), UpdatedAt: source.UpdatedAt.UTC(), } } diff --git a/backend/internal/logic/files/handlers_test.go b/backend/internal/logic/files/handlers_test.go index d91e61f..10227a5 100644 --- a/backend/internal/logic/files/handlers_test.go +++ b/backend/internal/logic/files/handlers_test.go @@ -11,6 +11,7 @@ import ( "net/http/httptest" "os" "path/filepath" + "strings" "testing" "github.com/gin-gonic/gin" @@ -43,14 +44,17 @@ func TestFileRegistrarSavesBeforeCreatingIdentitySourceDTO(t *testing.T) { var source models.SenlinAgentSource require.NoError(t, database.Where("project_id = ?", project.ID).First(&source).Error) require.FileExists(t, filepath.Join(storageRoot, filepath.FromSlash(source.FilePath))) + require.Contains(t, filepath.ToSlash(source.FilePath), "projects/"+project.Identity+"/") + require.NotContains(t, filepath.ToSlash(source.FilePath), fmt.Sprintf("projects/%d/", project.ID)) var payload map[string]any require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &payload)) require.Equal(t, source.Identity, payload["id"]) require.Equal(t, project.Identity, payload["projectId"]) - require.Equal(t, filepath.ToSlash(source.FilePath), payload["filePath"]) - require.NotContains(t, payload["filePath"], storageRoot) + require.Equal(t, filepath.Base(source.FilePath), payload["storageKey"]) + require.NotContains(t, payload, "filePath") require.NotContains(t, payload, "AbsolutePath") require.NotContains(t, rec.Body.String(), storageRoot) + require.NotContains(t, rec.Body.String(), fmt.Sprintf("projects/%d/", project.ID)) } func TestFileRegistrarChecksOwnershipBeforeWritingFile(t *testing.T) { @@ -139,7 +143,12 @@ func TestFileRegistrarLogsCleanupFailureWithoutLeakingStorageRoot(t *testing.T) tx.AddError(errors.New("forced source create failure")) } })) - service.remove = func(path string) error { return fmt.Errorf("cleanup blocked for %s", path) } + service.remove = func(path string) error { + if strings.HasPrefix(filepath.Base(path), ".upload-") { + return os.Remove(path) + } + return fmt.Errorf("cleanup blocked for %s", path) + } var serverLog bytes.Buffer previousLogOutput := log.Writer() log.SetOutput(&serverLog) diff --git a/backend/internal/logic/files/service.go b/backend/internal/logic/files/service.go index 13727ad..c7d3aba 100644 --- a/backend/internal/logic/files/service.go +++ b/backend/internal/logic/files/service.go @@ -1,24 +1,25 @@ package files import ( + "crypto/rand" + "encoding/hex" "errors" - "fmt" "io" "os" "path/filepath" "strings" - "time" "gorm.io/gorm" "senlinai-agent/backend/internal/models" ) type Service struct { - root string - db *gorm.DB - createTemp func(string, string) (stagedFile, error) - rename func(string, string) error - remove func(string) error + root string + db *gorm.DB + createTemp func(string, string) (stagedFile, error) + publish func(string, string) error + remove func(string) error + newStorageKey func() (string, error) } type stagedFile interface { @@ -29,6 +30,7 @@ type stagedFile interface { type StoredFile struct { OriginalName string + StorageKey string RelativePath string AbsolutePath string } @@ -41,45 +43,73 @@ func NewService(root string, databases ...*gorm.DB) *Service { } return &Service{ root: root, db: database, - createTemp: func(directory, pattern string) (stagedFile, error) { return os.CreateTemp(directory, pattern) }, - rename: os.Rename, - remove: os.Remove, + createTemp: func(directory, pattern string) (stagedFile, error) { return os.CreateTemp(directory, pattern) }, + publish: os.Link, + remove: os.Remove, + newStorageKey: randomStorageKey, } } -// Save 清理客户端文件名,并只返回供持久化的相对路径;绝对路径不得进入 API DTO。 -func (s *Service) Save(projectID uint, originalName string, content io.Reader) (StoredFile, error) { +// Save 使用项目公开 identity 和随机 opaque key 定位文件;硬链接发布提供跨请求的排他创建语义。 +func (s *Service) Save(projectIdentity string, originalName string, content io.Reader) (StoredFile, error) { cleanName := filepath.Base(strings.ReplaceAll(strings.TrimSpace(originalName), "\\", "/")) if cleanName == "." || cleanName == "" { cleanName = "upload.bin" } - relative := filepath.ToSlash(filepath.Join("projects", fmt.Sprint(projectID), fmt.Sprintf("%d-%s", time.Now().UnixNano(), cleanName))) - absolute, err := s.absolutePath(relative) + projectSegment := filepath.Base(strings.ReplaceAll(strings.TrimSpace(projectIdentity), "\\", "/")) + if projectSegment == "." || projectSegment == "" || projectSegment != strings.TrimSpace(projectIdentity) { + return StoredFile{}, ErrSourcePathRequired + } + directoryRelative := filepath.ToSlash(filepath.Join("projects", projectSegment)) + directoryAbsolute, err := s.absolutePath(directoryRelative) if err != nil { return StoredFile{}, err } - if err := os.MkdirAll(filepath.Dir(absolute), 0o755); err != nil { + if err := os.MkdirAll(directoryAbsolute, 0o755); err != nil { return StoredFile{}, err } - // 临时文件与最终文件位于同一目录,Close 成功后再原子替换,避免暴露半写入内容。 - file, err := s.createTemp(filepath.Dir(absolute), ".upload-*") + // 临时文件与最终文件位于同一目录,完整关闭后再排他发布,避免暴露半写入内容。 + file, err := s.createTemp(directoryAbsolute, ".upload-*") if err != nil { return StoredFile{}, err } if _, err := io.Copy(file, content); err != nil { _ = file.Close() - s.cleanupFailedSave(file.Name(), absolute) + s.cleanupOwnedFiles(file.Name()) return StoredFile{}, err } if err := file.Close(); err != nil { - s.cleanupFailedSave(file.Name(), absolute) + s.cleanupOwnedFiles(file.Name()) return StoredFile{}, err } - if err := s.rename(file.Name(), absolute); err != nil { - s.cleanupFailedSave(file.Name(), absolute) - return StoredFile{}, err + for attempt := 0; attempt < storageKeyAttempts; attempt++ { + storageKey, err := s.newStorageKey() + if err != nil { + s.cleanupOwnedFiles(file.Name()) + return StoredFile{}, err + } + relative := filepath.ToSlash(filepath.Join(directoryRelative, storageKey)) + absolute, err := s.absolutePath(relative) + if err != nil { + s.cleanupOwnedFiles(file.Name()) + return StoredFile{}, err + } + if err := s.publish(file.Name(), absolute); err != nil { + if errors.Is(err, os.ErrExist) { + continue + } + s.cleanupOwnedFiles(file.Name()) + return StoredFile{}, err + } + if err := s.remove(file.Name()); err != nil { + // final path 由本次排他发布创建,因此这里只清理本次请求拥有的两个路径。 + s.cleanupOwnedFiles(file.Name(), absolute) + return StoredFile{}, err + } + return StoredFile{OriginalName: cleanName, StorageKey: storageKey, RelativePath: relative, AbsolutePath: absolute}, nil } - return StoredFile{OriginalName: cleanName, RelativePath: relative, AbsolutePath: absolute}, nil + s.cleanupOwnedFiles(file.Name()) + return StoredFile{}, ErrStorageKeyCollision } // Remove 只依据 Service 生成的相对路径定位文件,并尽量移除直至存储根目录的空父目录。 @@ -133,9 +163,10 @@ func (s *Service) absolutePath(relativePath string) (string, error) { return absolute, nil } -func (s *Service) cleanupFailedSave(tempPath, finalPath string) { - _ = s.remove(tempPath) - _ = s.remove(finalPath) +func (s *Service) cleanupOwnedFiles(paths ...string) { + for _, path := range paths { + _ = s.remove(path) + } } func (s *Service) database() *gorm.DB { @@ -148,8 +179,19 @@ func (s *Service) database() *gorm.DB { var ( ErrSourceTitleRequired = errors.New("source title is required") ErrSourcePathRequired = errors.New("source file path is required") + ErrStorageKeyCollision = errors.New("unable to allocate unique storage key") ) +const storageKeyAttempts = 8 + +func randomStorageKey() (string, error) { + buffer := make([]byte, 16) + if _, err := rand.Read(buffer); err != nil { + return "", err + } + return hex.EncodeToString(buffer), nil +} + // CreateSource 只持久化 Save 产生的相对路径,不接受 handler 自行拼接本地路径。 func (s *Service) CreateSource(ownerID uint, project *models.SenlinAgentProject, title string, stored StoredFile) (*models.SenlinAgentSource, error) { title = strings.TrimSpace(title) diff --git a/backend/internal/logic/files/service_test.go b/backend/internal/logic/files/service_test.go index d205d87..ae47a68 100644 --- a/backend/internal/logic/files/service_test.go +++ b/backend/internal/logic/files/service_test.go @@ -6,6 +6,7 @@ import ( "os" "path/filepath" "strings" + "sync" "testing" "github.com/stretchr/testify/require" @@ -32,32 +33,36 @@ func (f *closeFailingFile) Close() error { return errors.New("close failed") } -func TestSaveStoresFileUnderProjectDirectory(t *testing.T) { +func TestSaveStoresFileUnderProjectIdentityWithOpaqueKey(t *testing.T) { service := NewService(t.TempDir()) + projectIdentity := "019b0000-0000-7000-8000-000000000012" - stored, err := service.Save(12, "brief.md", strings.NewReader("hello")) + stored, err := service.Save(projectIdentity, "brief.md", strings.NewReader("hello")) require.NoError(t, err) require.Equal(t, "brief.md", stored.OriginalName) - require.Contains(t, stored.RelativePath, "projects/12/") + require.Contains(t, stored.RelativePath, "projects/"+projectIdentity+"/") + require.NotContains(t, stored.RelativePath, "projects/12/") + require.Regexp(t, `^[a-f0-9]{32}$`, stored.StorageKey) + require.Equal(t, stored.StorageKey, filepath.Base(stored.RelativePath)) require.FileExists(t, stored.AbsolutePath) } func TestSaveNeutralizesPathTraversal(t *testing.T) { service := NewService(t.TempDir()) - stored, err := service.Save(12, "..\\..\\secret.txt", strings.NewReader("hello")) + stored, err := service.Save("019b0000-0000-7000-8000-000000000012", "..\\..\\secret.txt", strings.NewReader("hello")) require.NoError(t, err) require.Equal(t, "secret.txt", stored.OriginalName) - require.Contains(t, stored.RelativePath, "projects/12/") + require.NotContains(t, stored.RelativePath, "secret.txt") } func TestSaveRemovesTemporaryAndPartialFilesWhenCopyFails(t *testing.T) { storageRoot := t.TempDir() service := NewService(storageRoot) - _, err := service.Save(12, "broken.bin", &failingReader{}) + _, err := service.Save("019b0000-0000-7000-8000-000000000012", "broken.bin", &failingReader{}) require.ErrorContains(t, err, "copy failed") requireNoStoredFiles(t, storageRoot) @@ -74,24 +79,63 @@ func TestSaveRemovesTemporaryFileWhenCloseFails(t *testing.T) { return &closeFailingFile{File: file}, nil } - _, err := service.Save(12, "broken.bin", strings.NewReader("content")) + _, err := service.Save("019b0000-0000-7000-8000-000000000012", "broken.bin", strings.NewReader("content")) require.ErrorContains(t, err, "close failed") requireNoStoredFiles(t, storageRoot) } -func TestSaveRemovesTemporaryAndPartialFilesWhenRenameFails(t *testing.T) { +func TestSaveCollisionNeverOverwritesOrDeletesExistingFile(t *testing.T) { storageRoot := t.TempDir() service := NewService(storageRoot) - service.rename = func(oldPath, newPath string) error { - require.NoError(t, os.WriteFile(newPath, []byte("partial final"), 0o600)) - return errors.New("rename failed") + service.newStorageKey = func() (string, error) { return "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", nil } + projectIdentity := "019b0000-0000-7000-8000-000000000012" + first, err := service.Save(projectIdentity, "first.bin", strings.NewReader("first owner")) + require.NoError(t, err) + + _, err = service.Save(projectIdentity, "second.bin", strings.NewReader("second owner")) + + require.ErrorIs(t, err, ErrStorageKeyCollision) + content, readErr := os.ReadFile(first.AbsolutePath) + require.NoError(t, readErr) + require.Equal(t, "first owner", string(content)) + requireOnlyStoredFile(t, storageRoot, first.AbsolutePath) +} + +func TestConcurrentSaveWithSameKeyPublishesExactlyOneOwner(t *testing.T) { + storageRoot := t.TempDir() + service := NewService(storageRoot) + service.newStorageKey = func() (string, error) { return "bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb", nil } + projectIdentity := "019b0000-0000-7000-8000-000000000012" + start := make(chan struct{}) + type result struct { + stored StoredFile + err error } + results := make(chan result, 2) + var ready sync.WaitGroup + ready.Add(2) + for _, body := range []string{"alpha", "beta"} { + go func(content string) { + ready.Done() + <-start + stored, err := service.Save(projectIdentity, content+".bin", strings.NewReader(content)) + results <- result{stored: stored, err: err} + }(body) + } + ready.Wait() + close(start) + firstResult, secondResult := <-results, <-results - _, err := service.Save(12, "broken.bin", strings.NewReader("content")) - - require.ErrorContains(t, err, "rename failed") - requireNoStoredFiles(t, storageRoot) + errors := []error{firstResult.err, secondResult.err} + require.Equal(t, 1, countNilErrors(errors)) + require.Equal(t, 1, countMatchingErrors(errors, ErrStorageKeyCollision)) + winner := firstResult.stored + if firstResult.err != nil { + winner = secondResult.stored + } + require.FileExists(t, winner.AbsolutePath) + requireOnlyStoredFile(t, storageRoot, winner.AbsolutePath) } func requireNoStoredFiles(t *testing.T, storageRoot string) { @@ -106,3 +150,38 @@ func requireNoStoredFiles(t *testing.T, storageRoot string) { return nil })) } + +func requireOnlyStoredFile(t *testing.T, storageRoot, expectedPath string) { + t.Helper() + files := []string{} + require.NoError(t, filepath.WalkDir(storageRoot, func(path string, entry fs.DirEntry, err error) error { + if err != nil { + return err + } + if !entry.IsDir() { + files = append(files, path) + } + return nil + })) + require.Equal(t, []string{expectedPath}, files) +} + +func countNilErrors(values []error) int { + count := 0 + for _, err := range values { + if err == nil { + count++ + } + } + return count +} + +func countMatchingErrors(values []error, target error) int { + count := 0 + for _, err := range values { + if errors.Is(err, target) { + count++ + } + } + return count +} diff --git a/backend/internal/logic/projects/dashboard.go b/backend/internal/logic/projects/dashboard.go deleted file mode 100644 index 141428d..0000000 --- a/backend/internal/logic/projects/dashboard.go +++ /dev/null @@ -1,24 +0,0 @@ -package projects - -import "senlinai-agent/backend/internal/models" - -// Dashboard 保留旧概览查询给尚未迁移的 handler;所有统计仍严格限定项目所有者。 -func (s *Service) Dashboard(ownerID uint, projectID uint) (Dashboard, error) { - if err := ensureProjectOwner(ownerID, projectID); err != nil { - return Dashboard{}, err - } - dashboard := Dashboard{ProjectID: projectID} - if err := models.DBService.Model(&models.SenlinAgentInboxItem{}).Where("project_id = ? AND status = ?", projectID, "open").Count(&dashboard.PendingInboxCount).Error; err != nil { - return dashboard, err - } - if err := models.DBService.Model(&models.SenlinAgentTask{}).Where("project_id = ? AND status <> ?", projectID, "done").Count(&dashboard.OpenTaskCount).Error; err != nil { - return dashboard, err - } - if err := models.DBService.Model(&models.SenlinAgentNote{}).Where("project_id = ?", projectID).Count(&dashboard.RecentNoteCount).Error; err != nil { - return dashboard, err - } - if err := models.DBService.Model(&models.SenlinAgentAISession{}).Where("project_id = ?", projectID).Count(&dashboard.RecentSessionCount).Error; err != nil { - return dashboard, err - } - return dashboard, nil -} diff --git a/backend/internal/logic/projects/dto.go b/backend/internal/logic/projects/dto.go index 261dd2b..793404b 100644 --- a/backend/internal/logic/projects/dto.go +++ b/backend/internal/logic/projects/dto.go @@ -41,14 +41,6 @@ type CreateCronPlanInput struct { NextRunAt *time.Time } -type Dashboard struct { - ProjectID uint `json:"project_id"` - PendingInboxCount int64 `json:"pending_inbox_count"` - OpenTaskCount int64 `json:"open_task_count"` - RecentNoteCount int64 `json:"recent_note_count"` - RecentSessionCount int64 `json:"recent_session_count"` -} - // WorkspaceDTO 是项目工作区首屏聚合响应,所有数据库对象均以公开 identity 关联。 type WorkspaceDTO struct { Project WorkspaceProjectDTO `json:"project"` diff --git a/backend/internal/logic/projects/object_crud.go b/backend/internal/logic/projects/object_crud.go index 397558c..492bef9 100644 --- a/backend/internal/logic/projects/object_crud.go +++ b/backend/internal/logic/projects/object_crud.go @@ -3,8 +3,10 @@ package projects import ( "errors" "strings" + "time" "gorm.io/gorm" + "gorm.io/gorm/clause" "senlinai-agent/backend/internal/models" ) @@ -28,7 +30,7 @@ func (s *Service) CreateCronPlan(ownerID uint, projectID uint, input CreateCronP } plan := &models.SenlinAgentCronPlan{ ProjectID: projectID, CreatedBy: ownerID, Title: title, Schedule: schedule, - Enabled: input.Enabled, NextRunAt: input.NextRunAt, LastResult: "Not run yet", + Enabled: input.Enabled, NextRunAt: input.NextRunAt, } return plan, models.DBService.Create(plan).Error } @@ -45,16 +47,40 @@ func (s *Service) CreateTag(projectID uint, name string) (*models.SenlinAgentTag if name == "" { return nil, ErrTagNameRequired } - var existing models.SenlinAgentTag - err := models.DBService.Where("project_id = ? AND name = ?", projectID, name).First(&existing).Error - if err == nil { - return &existing, nil + var lastErr error + for attempt := 0; attempt < 5; attempt++ { + candidate := models.SenlinAgentTag{ProjectID: projectID, Name: name} + lastErr = models.DBService.Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "project_id"}, {Name: "name"}}, + DoNothing: true, + }).Create(&candidate).Error + if lastErr == nil { + var stored models.SenlinAgentTag + lastErr = models.DBService.Where("project_id = ? AND name = ?", projectID, name).First(&stored).Error + if lastErr == nil { + return &stored, nil + } + } + if !retryableTagWrite(lastErr) { + return nil, lastErr + } + time.Sleep(time.Duration(attempt+1) * time.Millisecond) } - if !errors.Is(err, gorm.ErrRecordNotFound) { - return nil, err + return nil, lastErr +} + +func retryableTagWrite(err error) bool { + if errors.Is(err, gorm.ErrDuplicatedKey) || errors.Is(err, gorm.ErrRecordNotFound) || strings.Contains(strings.ToLower(err.Error()), "database is locked") { + return true } - tag := &models.SenlinAgentTag{ProjectID: projectID, Name: name} - return tag, models.DBService.Create(tag).Error + var sqlState interface{ SQLState() string } + if errors.As(err, &sqlState) { + switch sqlState.SQLState() { + case "40001", "40P01", "55P03": + return true + } + } + return false } func (s *Service) ListProjectTags(ownerID uint, projectID uint) ([]models.SenlinAgentTag, error) { diff --git a/backend/internal/logic/projects/service_test.go b/backend/internal/logic/projects/service_test.go index 02564ef..2e4edc8 100644 --- a/backend/internal/logic/projects/service_test.go +++ b/backend/internal/logic/projects/service_test.go @@ -3,6 +3,9 @@ package projects import ( "errors" "fmt" + "path/filepath" + "sync" + "sync/atomic" "testing" "time" @@ -11,6 +14,7 @@ import ( "github.com/stretchr/testify/require" "gorm.io/driver/postgres" "gorm.io/gorm" + "gorm.io/gorm/logger" "senlinai-agent/backend/internal/models" ) @@ -106,36 +110,6 @@ func TestCreateProjectPersistsWorkspaceMetadata(t *testing.T) { require.Equal(t, "RSS 采集和线索沉淀", workspace.Project.Description) } -func TestDashboardCountsOnlyRequestedProject(t *testing.T) { - database := newTestDB(t) - service := NewService() - first, err := service.CreateProject(1, "Alpha", "") - require.NoError(t, err) - second, err := service.CreateProject(1, "Beta", "") - require.NoError(t, err) - require.NoError(t, database.Create(&models.SenlinAgentInboxItem{ProjectID: first.ID, CreatedBy: 1, SourceType: "text", Status: "open"}).Error) - require.NoError(t, database.Create(&models.SenlinAgentInboxItem{ProjectID: second.ID, CreatedBy: 1, SourceType: "text", Status: "open"}).Error) - require.NoError(t, database.Create(&models.SenlinAgentTask{ProjectID: first.ID, CreatedBy: 1, Title: "A", Status: "open"}).Error) - require.NoError(t, database.Create(&models.SenlinAgentTask{ProjectID: second.ID, CreatedBy: 1, Title: "B", Status: "open"}).Error) - - dashboard, err := service.Dashboard(1, first.ID) - - require.NoError(t, err) - require.Equal(t, int64(1), dashboard.PendingInboxCount) - require.Equal(t, int64(1), dashboard.OpenTaskCount) -} - -func TestDashboardRejectsProjectOwnedByAnotherUser(t *testing.T) { - newTestDB(t) - service := NewService() - project, err := service.CreateProject(2, "Beta", "") - require.NoError(t, err) - - _, err = service.Dashboard(1, project.ID) - - require.Error(t, err) -} - func TestWorkspaceMatchesFrontendContract(t *testing.T) { database := newTestDB(t) service := NewService() @@ -230,6 +204,64 @@ func TestCreateProjectTagRejectsProjectOwnedByAnotherUser(t *testing.T) { require.Error(t, err) } +func TestCreateProjectTagConcurrentlyReturnsOneProjectScopedTag(t *testing.T) { + databasePath := filepath.ToSlash(filepath.Join(t.TempDir(), "tags.db")) + database, err := gorm.Open(sqlite.Open(fmt.Sprintf("file:%s?_pragma=busy_timeout(10000)&_pragma=journal_mode(WAL)", databasePath)), &gorm.Config{TranslateError: true, Logger: logger.Default.LogMode(logger.Silent)}) + require.NoError(t, err) + sqlDB, err := database.DB() + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, sqlDB.Close()) }) + require.NoError(t, models.AutoMigrate(database)) + models.DBService = database + service := NewService() + project, err := service.CreateProjectWithInput(1, CreateProjectRequest{Name: "Concurrent", Identifier: "CONCURRENT"}) + require.NoError(t, err) + + arrived := make(chan struct{}, 2) + release := make(chan struct{}) + var synchronizedQueries atomic.Int32 + require.NoError(t, database.Callback().Create().Before("gorm:create").Register("test:synchronize_tag_inserts", func(tx *gorm.DB) { + if tx.Statement.Schema == nil || tx.Statement.Schema.Table != (models.SenlinAgentTag{}).TableName() { + return + } + if synchronizedQueries.Add(1) <= 2 { + arrived <- struct{}{} + <-release + } + })) + + type result struct { + tag *models.SenlinAgentTag + err error + } + results := make(chan result, 2) + var workers sync.WaitGroup + workers.Add(2) + for range 2 { + go func() { + defer workers.Done() + tag, err := service.CreateProjectTag(1, project.ID, "UI") + results <- result{tag: tag, err: err} + }() + } + <-arrived + <-arrived + close(release) + workers.Wait() + close(results) + + identities := map[string]struct{}{} + for value := range results { + require.NoError(t, value.err) + require.NotNil(t, value.tag) + identities[value.tag.Identity] = struct{}{} + } + require.Len(t, identities, 1) + var count int64 + require.NoError(t, database.Model(&models.SenlinAgentTag{}).Where("project_id = ? AND name = ?", project.ID, "UI").Count(&count).Error) + require.Equal(t, int64(1), count) +} + func TestCreateCronPlanAddsPlanToOwnedProject(t *testing.T) { database := newTestDB(t) service := NewService() @@ -252,6 +284,7 @@ func TestCreateCronPlanAddsPlanToOwnedProject(t *testing.T) { require.Equal(t, "0 9 * * *", plan.Schedule) require.True(t, plan.Enabled) require.NotNil(t, plan.NextRunAt) + require.Empty(t, plan.LastResult, "创建计划只保存元数据,不能伪造执行结果") workspace, err := service.Workspace(1, project.Identity) require.NoError(t, err) require.Len(t, workspace.CronPlans, 1) diff --git a/backend/internal/logic/projects/workspace.go b/backend/internal/logic/projects/workspace.go index e35e915..746b144 100644 --- a/backend/internal/logic/projects/workspace.go +++ b/backend/internal/logic/projects/workspace.go @@ -269,7 +269,7 @@ func (s *Service) workspaceCronPlans(projectID uint, projectIdentity string) ([] items = append(items, WorkspaceCronPlanDTO{ ID: plan.Identity, ProjectID: projectIdentity, Title: plan.Title, Schedule: plan.Schedule, NextRun: utcOptionalTime(plan.NextRunAt), Enabled: plan.Enabled, - LastResult: defaultString(plan.LastResult, "Not run yet"), Owner: owner, + LastResult: strings.TrimSpace(plan.LastResult), Owner: owner, }) } return items, nil diff --git a/backend/internal/logic/projects/workspace_test.go b/backend/internal/logic/projects/workspace_test.go index d305790..6d3bf9e 100644 --- a/backend/internal/logic/projects/workspace_test.go +++ b/backend/internal/logic/projects/workspace_test.go @@ -3,7 +3,6 @@ package projects import ( "encoding/json" "fmt" - "os" "testing" "time" @@ -123,12 +122,6 @@ func TestWorkspaceDoneTaskDoesNotInventCompletedAt(t *testing.T) { require.Contains(t, string(payload), `"completedAt":null`) } -func TestWorkspaceFileDoesNotOwnDashboard(t *testing.T) { - source, err := os.ReadFile("workspace.go") - require.NoError(t, err) - require.NotContains(t, string(source), "func (s *Service) Dashboard") -} - func TestWorkspaceRejectsProjectOwnedByAnotherUserByIdentity(t *testing.T) { newTestDB(t) service := NewService() diff --git a/backend/internal/logic/search/service_postgres_test.go b/backend/internal/logic/search/service_postgres_test.go index 70e4ef2..5360a3c 100644 --- a/backend/internal/logic/search/service_postgres_test.go +++ b/backend/internal/logic/search/service_postgres_test.go @@ -1,3 +1,5 @@ +//go:build integration + package search import ( @@ -109,10 +111,8 @@ func TestPostgresSearchUsesStableFairLimitAndRuneBoundedSnippets(t *testing.T) { func newPostgresSearchTestDB(t *testing.T) *gorm.DB { t.Helper() - dsn := os.Getenv("DATABASE_URL") - if dsn == "" { - t.Skip("DATABASE_URL is not configured; skipping PostgreSQL search integration test") - } + dsn := os.Getenv("TEST_DATABASE_URL") + require.NotEmpty(t, dsn, "TEST_DATABASE_URL is required for integration tests and must point to an isolated database") database, err := gorm.Open(postgres.Open(dsn), &gorm.Config{TranslateError: true}) require.NoError(t, err) transaction := database.Begin() diff --git a/backend/internal/logic/tasks/concurrency_postgres_test.go b/backend/internal/logic/tasks/concurrency_postgres_test.go index 78a8074..ee688a0 100644 --- a/backend/internal/logic/tasks/concurrency_postgres_test.go +++ b/backend/internal/logic/tasks/concurrency_postgres_test.go @@ -1,3 +1,5 @@ +//go:build integration + package tasks import ( @@ -14,10 +16,8 @@ import ( ) func TestPostgresMovePreventsOldProjectShareFromBeingInsertedConcurrently(t *testing.T) { - dsn := os.Getenv("DATABASE_URL") - if dsn == "" { - t.Skip("DATABASE_URL is not configured; skipping PostgreSQL row-lock concurrency test") - } + dsn := os.Getenv("TEST_DATABASE_URL") + require.NotEmpty(t, dsn, "TEST_DATABASE_URL is required for integration tests and must point to an isolated database") database, err := gorm.Open(postgres.Open(dsn), &gorm.Config{TranslateError: true}) require.NoError(t, err) require.NoError(t, models.AutoMigrate(database)) diff --git a/backend/internal/logic/tasks/service.go b/backend/internal/logic/tasks/service.go index 27a6b56..3c78054 100644 --- a/backend/internal/logic/tasks/service.go +++ b/backend/internal/logic/tasks/service.go @@ -236,14 +236,15 @@ func findOrCreateProjectTag(tx *gorm.DB, projectID uint, value string) (*uint, * if name == "" { return nil, nil, "", nil } + candidate := models.SenlinAgentTag{ProjectID: projectID, Name: name} + if err := tx.Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "project_id"}, {Name: "name"}}, + DoNothing: true, + }).Create(&candidate).Error; err != nil { + return nil, nil, "", err + } var tag models.SenlinAgentTag - err := tx.Where("project_id = ? AND name = ?", projectID, name).First(&tag).Error - if errors.Is(err, gorm.ErrRecordNotFound) { - tag = models.SenlinAgentTag{ProjectID: projectID, Name: name} - if err := tx.Create(&tag).Error; err != nil { - return nil, nil, "", err - } - } else if err != nil { + if err := tx.Where("project_id = ? AND name = ?", projectID, name).First(&tag).Error; err != nil { return nil, nil, "", err } return &tag.ID, &tag.Identity, tag.Name, nil diff --git a/backend/internal/models/migrations.go b/backend/internal/models/migrations.go new file mode 100644 index 0000000..6146fac --- /dev/null +++ b/backend/internal/models/migrations.go @@ -0,0 +1,150 @@ +package models + +import ( + "encoding/json" + "fmt" + "strings" + "time" + + "gorm.io/gorm" +) + +// SenlinAgentSchemaMigration 记录已完成的数据修复及审计摘要,确保升级可重入且可追溯。 +type SenlinAgentSchemaMigration struct { + Version int `gorm:"primaryKey;autoIncrement:false"` + Name string `gorm:"not null"` + Details string `gorm:"type:text;not null"` + AppliedAt time.Time `gorm:"not null"` +} + +func (SenlinAgentSchemaMigration) TableName() string { + return "senlin_agent_schema_migrations" +} + +type versionedMigration struct { + version int + name string + run func(*gorm.DB) (map[string]int, error) +} + +func runVersionedMigrations(database *gorm.DB) error { + migrations := []versionedMigration{ + {version: 1, name: "normalize_project_identifiers", run: normalizeLegacyProjectIdentifiers}, + {version: 2, name: "deduplicate_project_tags", run: deduplicateLegacyProjectTags}, + } + for _, migration := range migrations { + var applied int64 + if err := database.Model(&SenlinAgentSchemaMigration{}).Where("version = ?", migration.version).Count(&applied).Error; err != nil { + return err + } + if applied > 0 { + continue + } + if err := database.Transaction(func(tx *gorm.DB) error { + details, err := migration.run(tx) + if err != nil { + return err + } + encoded, err := json.Marshal(details) + if err != nil { + return err + } + return tx.Create(&SenlinAgentSchemaMigration{ + Version: migration.version, Name: migration.name, Details: string(encoded), AppliedAt: time.Now().UTC(), + }).Error + }); err != nil { + return fmt.Errorf("schema migration %d (%s): %w", migration.version, migration.name, err) + } + } + return nil +} + +func normalizeLegacyProjectIdentifiers(tx *gorm.DB) (map[string]int, error) { + var projects []SenlinAgentProject + if err := tx.Select("id", "owner_id", "identifier").Order("owner_id asc, id asc").Find(&projects).Error; err != nil { + return nil, err + } + + keepers := make(map[uint]map[string]uint) + reserved := make(map[uint]map[string]struct{}) + for _, project := range projects { + identifier := strings.TrimSpace(project.Identifier) + if identifier == "" { + continue + } + if keepers[project.OwnerID] == nil { + keepers[project.OwnerID] = make(map[string]uint) + reserved[project.OwnerID] = make(map[string]struct{}) + } + if _, exists := keepers[project.OwnerID][identifier]; !exists { + keepers[project.OwnerID][identifier] = project.ID + } + reserved[project.OwnerID][identifier] = struct{}{} + } + + updated := 0 + for _, project := range projects { + if reserved[project.OwnerID] == nil { + reserved[project.OwnerID] = make(map[string]struct{}) + } + identifier := strings.TrimSpace(project.Identifier) + candidate := identifier + if identifier == "" { + candidate = uniqueIdentifier(fmt.Sprintf("project-%d", project.ID), reserved[project.OwnerID]) + } else if keepers[project.OwnerID][identifier] != project.ID { + candidate = uniqueIdentifier(fmt.Sprintf("%s-%d", identifier, project.ID), reserved[project.OwnerID]) + } + reserved[project.OwnerID][candidate] = struct{}{} + if candidate == project.Identifier { + continue + } + if err := tx.Model(&SenlinAgentProject{}).Where("id = ?", project.ID).Update("identifier", candidate).Error; err != nil { + return nil, err + } + updated++ + } + if err := tx.Exec(`CREATE UNIQUE INDEX IF NOT EXISTS uidx_senlin_agent_projects_owner_identifier ON senlin_agent_projects (owner_id, identifier)`).Error; err != nil { + return nil, err + } + return map[string]int{"audited": len(projects), "updated": updated}, nil +} + +func uniqueIdentifier(base string, reserved map[string]struct{}) string { + candidate := base + for suffix := 2; ; suffix++ { + if _, exists := reserved[candidate]; !exists { + return candidate + } + candidate = fmt.Sprintf("%s-%d", base, suffix) + } +} + +func deduplicateLegacyProjectTags(tx *gorm.DB) (map[string]int, error) { + var tags []SenlinAgentTag + if err := tx.Select("id", "identity", "project_id", "name").Order("project_id asc, name asc, id asc").Find(&tags).Error; err != nil { + return nil, err + } + keepers := make(map[string]SenlinAgentTag) + deduplicated := 0 + for _, tag := range tags { + key := fmt.Sprintf("%d\x00%s", tag.ProjectID, tag.Name) + keeper, exists := keepers[key] + if !exists { + keepers[key] = tag + continue + } + if err := tx.Model(&SenlinAgentTask{}).Where("tag_id = ?", tag.ID).Updates(map[string]any{ + "tag_id": keeper.ID, "tag_identity": keeper.Identity, + }).Error; err != nil { + return nil, err + } + if err := tx.Where("id = ?", tag.ID).Delete(&SenlinAgentTag{}).Error; err != nil { + return nil, err + } + deduplicated++ + } + if err := tx.Exec(`CREATE UNIQUE INDEX IF NOT EXISTS uidx_senlin_agent_tags_project_name ON senlin_agent_tags (project_id, name)`).Error; err != nil { + return nil, err + } + return map[string]int{"audited": len(tags), "deduplicated": deduplicated}, nil +} diff --git a/backend/internal/models/migrations_test.go b/backend/internal/models/migrations_test.go new file mode 100644 index 0000000..3d319e0 --- /dev/null +++ b/backend/internal/models/migrations_test.go @@ -0,0 +1,84 @@ +package models + +import ( + "fmt" + "testing" + + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +func TestAutoMigrateUpgradesLegacyProjectsAndTagsWithoutLosingAssociations(t *testing.T) { + database, err := gorm.Open(sqlite.Open(fmt.Sprintf("file:%s?mode=memory&cache=shared", t.Name())), &gorm.Config{TranslateError: true}) + require.NoError(t, err) + + legacySchema := []string{ + `CREATE TABLE senlin_agent_projects (id integer primary key, owner_id integer not null, name text not null, identifier text)`, + `CREATE TABLE senlin_agent_tags (id integer primary key, identity text, project_id integer not null, name text not null)`, + `CREATE TABLE senlin_agent_tasks (id integer primary key, project_id integer not null, created_by integer not null, tag_id integer, tag_identity text, title text not null, status text not null)`, + `INSERT INTO senlin_agent_projects (id, owner_id, name, identifier) VALUES + (1, 7, 'Blank one', ''), + (2, 7, 'Blank two', NULL), + (3, 7, 'Keeper', 'DUP'), + (4, 7, 'Duplicate', 'DUP'), + (5, 7, 'Reserved suffix', 'DUP-4'), + (6, 8, 'Other owner', 'DUP')`, + `INSERT INTO senlin_agent_tags (id, identity, project_id, name) VALUES + (10, 'tag-keeper', 1, 'UI'), + (11, 'tag-duplicate', 1, 'UI'), + (12, 'tag-other-project', 2, 'UI')`, + `INSERT INTO senlin_agent_tasks (id, project_id, created_by, tag_id, tag_identity, title, status) + VALUES (20, 1, 7, 11, 'tag-duplicate', 'Keep association', 'open')`, + } + for _, statement := range legacySchema { + require.NoError(t, database.Exec(statement).Error) + } + + require.NoError(t, AutoMigrate(database)) + + var projects []SenlinAgentProject + require.NoError(t, database.Order("id asc").Find(&projects).Error) + require.Len(t, projects, 6, "项目迁移不得删除 legacy 行") + require.Equal(t, []string{"project-1", "project-2", "DUP", "DUP-4-2", "DUP-4", "DUP"}, projectIdentifiers(projects)) + + var tags []SenlinAgentTag + require.NoError(t, database.Order("id asc").Find(&tags).Error) + require.Equal(t, []uint{10, 12}, tagIDs(tags), "同项目同名标签应保留最早记录,不影响其他项目") + var task SenlinAgentTask + require.NoError(t, database.First(&task, 20).Error) + require.NotNil(t, task.TagID) + require.Equal(t, uint(10), *task.TagID) + require.NotNil(t, task.TagIdentity) + require.Equal(t, "tag-keeper", *task.TagIdentity) + + var migrationCount int64 + require.NoError(t, database.Table("senlin_agent_schema_migrations").Count(&migrationCount).Error) + require.Equal(t, int64(2), migrationCount) + var auditDetails []string + require.NoError(t, database.Table("senlin_agent_schema_migrations").Order("version asc").Pluck("details", &auditDetails).Error) + require.Contains(t, auditDetails[0], `"updated":3`) + require.Contains(t, auditDetails[1], `"deduplicated":1`) + require.NoError(t, AutoMigrate(database), "versioned migrations must be safe to run again") + require.NoError(t, database.Table("senlin_agent_schema_migrations").Count(&migrationCount).Error) + require.Equal(t, int64(2), migrationCount) + + require.Error(t, database.Exec(`INSERT INTO senlin_agent_projects (owner_id, name, identifier) VALUES (7, 'still duplicate', 'DUP')`).Error) + require.Error(t, database.Exec(`INSERT INTO senlin_agent_tags (project_id, name) VALUES (1, 'UI')`).Error) +} + +func projectIdentifiers(projects []SenlinAgentProject) []string { + result := make([]string, 0, len(projects)) + for _, project := range projects { + result = append(result, project.Identifier) + } + return result +} + +func tagIDs(tags []SenlinAgentTag) []uint { + result := make([]uint, 0, len(tags)) + for _, tag := range tags { + result = append(result, tag.ID) + } + return result +} diff --git a/backend/internal/models/new.go b/backend/internal/models/new.go index f001481..67ddfb8 100644 --- a/backend/internal/models/new.go +++ b/backend/internal/models/new.go @@ -23,7 +23,8 @@ func New(dsn string) error { } func AutoMigrate(database *gorm.DB) error { - return database.AutoMigrate( + if err := database.AutoMigrate( + &SenlinAgentSchemaMigration{}, &SenlinAgentUser{}, &SenlinAgentProject{}, &SenlinAgentInboxItem{}, @@ -40,5 +41,8 @@ func AutoMigrate(database *gorm.DB) error { &SenlinAgentAICallLog{}, &SenlinAgentAIRateBucket{}, &SenlinAgentTaskShare{}, - ) + ); err != nil { + return err + } + return runVersionedMigrations(database) } diff --git a/backend/internal/models/postgres_integration_test.go b/backend/internal/models/postgres_integration_test.go index 7a6a517..81b96be 100644 --- a/backend/internal/models/postgres_integration_test.go +++ b/backend/internal/models/postgres_integration_test.go @@ -7,16 +7,15 @@ import ( "testing" "github.com/stretchr/testify/require" + "gorm.io/driver/postgres" + "gorm.io/gorm" ) func TestPostgresPing(t *testing.T) { - dsn := os.Getenv("DATABASE_DSN") - if dsn == "" { - dsn = os.Getenv("DATABASE_URL") - } - require.NotEmpty(t, dsn, "DATABASE_DSN or DATABASE_URL is required for integration tests") + dsn := os.Getenv("TEST_DATABASE_URL") + require.NotEmpty(t, dsn, "TEST_DATABASE_URL is required for integration tests and must point to an isolated database") - database, err := Open(dsn) + database, err := gorm.Open(postgres.Open(dsn), &gorm.Config{TranslateError: true}) require.NoError(t, err) sqlDB, err := database.DB() require.NoError(t, err) diff --git a/backend/internal/models/project.go b/backend/internal/models/project.go index f76e169..80cffea 100644 --- a/backend/internal/models/project.go +++ b/backend/internal/models/project.go @@ -5,10 +5,10 @@ import "time" type SenlinAgentProject struct { ID uint `gorm:"primaryKey"` Identity string `gorm:"type:char(36);uniqueIndex"` - OwnerID uint `gorm:"index;not null;uniqueIndex:uidx_senlin_agent_projects_owner_identifier,priority:1"` + OwnerID uint `gorm:"index;not null"` OwnerIdentity string `gorm:"type:char(36);index"` Name string `gorm:"not null"` - Identifier string `gorm:"index;uniqueIndex:uidx_senlin_agent_projects_owner_identifier,priority:2"` + Identifier string `gorm:"index"` Icon string Background string Description string diff --git a/backend/internal/seed/demo.go b/backend/internal/seed/demo.go index fecfc2e..b578aa9 100644 --- a/backend/internal/seed/demo.go +++ b/backend/internal/seed/demo.go @@ -190,9 +190,9 @@ func seedProjectA1(tx *gorm.DB, userID uint, projectID uint) error { } } for _, plan := range []models.SenlinAgentCronPlan{ - {ProjectID: projectID, CreatedBy: userID, Title: "每日 Inbox 摘要", Schedule: "0 9 * * *", NextRunAt: &nextRun, Enabled: true, LastResult: "上次执行成功"}, - {ProjectID: projectID, CreatedBy: userID, Title: "每周任务风险扫描", Schedule: "0 10 * * 1", NextRunAt: &dueSoon, Enabled: true, LastResult: "等待下次执行"}, - {ProjectID: projectID, CreatedBy: userID, Title: "月底知识归档提醒", Schedule: "0 18 28 * *", NextRunAt: &dueLater, Enabled: false, LastResult: "已暂停"}, + {ProjectID: projectID, CreatedBy: userID, Title: "每日 Inbox 摘要", Schedule: "0 9 * * *", NextRunAt: &nextRun, Enabled: true}, + {ProjectID: projectID, CreatedBy: userID, Title: "每周任务风险扫描", Schedule: "0 10 * * 1", NextRunAt: &dueSoon, Enabled: true}, + {ProjectID: projectID, CreatedBy: userID, Title: "月底知识归档提醒", Schedule: "0 18 28 * *", NextRunAt: &dueLater, Enabled: false}, } { if err := firstOrCreateCronPlan(tx, plan); err != nil { return err @@ -219,7 +219,14 @@ func seedCompactProject(tx *gorm.DB, userID uint, projectID uint) error { } func firstOrCreateTag(tx *gorm.DB, projectID uint, name string) error { - return firstOrCreate(tx, &models.SenlinAgentTag{}, "project_id = ? AND name = ?", []any{projectID, name}, models.SenlinAgentTag{ProjectID: projectID, Name: name}) + candidate := models.SenlinAgentTag{ProjectID: projectID, Name: name} + if err := tx.Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "project_id"}, {Name: "name"}}, + DoNothing: true, + }).Create(&candidate).Error; err != nil { + return err + } + return tx.Where("project_id = ? AND name = ?", projectID, name).First(&models.SenlinAgentTag{}).Error } func tagID(tx *gorm.DB, projectID uint, name string) (uint, error) { diff --git a/backend/internal/seed/demo_postgres_test.go b/backend/internal/seed/demo_postgres_test.go new file mode 100644 index 0000000..bf9aa7a --- /dev/null +++ b/backend/internal/seed/demo_postgres_test.go @@ -0,0 +1,230 @@ +//go:build integration + +package seed + +import ( + "context" + "fmt" + "os" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/require" + "golang.org/x/crypto/bcrypt" + "gorm.io/driver/postgres" + "gorm.io/gorm" + "gorm.io/gorm/logger" + "senlinai-agent/backend/internal/models" +) + +type demoSeedRunContextKey struct{} + +type demoSeedRunMarker string + +const ( + firstDemoSeedRun demoSeedRunMarker = "first" + secondDemoSeedRun demoSeedRunMarker = "second" +) + +func TestPostgresDemoSeedSerializesConcurrentRunsForSameOwner(t *testing.T) { + dsn := os.Getenv("TEST_DATABASE_URL") + require.NotEmpty(t, dsn, "TEST_DATABASE_URL is required for integration tests and must point to an isolated database") + database, err := gorm.Open(postgres.Open(dsn), &gorm.Config{ + TranslateError: true, + Logger: logger.Default.LogMode(logger.Silent), + }) + require.NoError(t, err) + require.NoError(t, models.AutoMigrate(database)) + + suffix := fmt.Sprint(time.Now().UnixNano()) + email := "postgres-seed-" + suffix + "@example.com" + passwordHash, err := bcrypt.GenerateFromPassword([]byte("original-password"), bcrypt.MinCost) + require.NoError(t, err) + user := models.SenlinAgentUser{Email: email, DisplayName: "保留用户名称", PasswordHash: string(passwordHash), Role: "user"} + require.NoError(t, database.Create(&user).Error) + projects := []models.SenlinAgentProject{ + {OwnerID: user.ID, Name: "项目 A1", Identifier: "A1"}, + {OwnerID: user.ID, Name: "项目 A2", Identifier: "A2"}, + {OwnerID: user.ID, Name: "数据中台", Identifier: "DATA"}, + {OwnerID: user.ID, Name: "运营自动化", Identifier: "AUTO"}, + } + for index := range projects { + require.NoError(t, database.Create(&projects[index]).Error) + } + uiTag := models.SenlinAgentTag{ProjectID: projects[0].ID, Name: "UI"} + require.NoError(t, database.Create(&uiTag).Error) + require.NoError(t, database.Create(&models.SenlinAgentTask{ + ProjectID: projects[0].ID, CreatedBy: user.ID, TagID: &uiTag.ID, + Title: "智能报表导出功能", Status: "open", + }).Error) + t.Cleanup(func() { cleanupPostgresDemoSeed(database, user.ID) }) + + firstAtChildQuery := make(chan struct{}, 1) + secondOwnerQueryAttempted := make(chan struct{}, 1) + secondEnteredChildQuery := make(chan struct{}, 1) + releaseFirst := make(chan struct{}) + var firstChildOnce sync.Once + var releaseOnce sync.Once + releaseFirstRun := func() { releaseOnce.Do(func() { close(releaseFirst) }) } + callbackName := "test:coordinate_concurrent_seed_" + suffix + require.NoError(t, database.Callback().Query().Before("gorm:query").Register(callbackName, func(tx *gorm.DB) { + if tx.Statement.Schema == nil { + return + } + marker, ok := tx.Statement.Context.Value(demoSeedRunContextKey{}).(demoSeedRunMarker) + if !ok { + return + } + if tx.Statement.Schema.Table == (models.SenlinAgentUser{}).TableName() { + if _, locking := tx.Statement.Clauses["FOR"]; locking && marker == secondDemoSeedRun { + signalDemoSeedStage(secondOwnerQueryAttempted) + } + return + } + switch marker { + case firstDemoSeedRun: + firstChildOnce.Do(func() { + signalDemoSeedStage(firstAtChildQuery) + <-releaseFirst + }) + case secondDemoSeedRun: + signalDemoSeedStage(secondEnteredChildQuery) + } + })) + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + var wait sync.WaitGroup + firstResult := make(chan error, 1) + secondResult := make(chan error, 1) + t.Cleanup(func() { + cancel() + releaseFirstRun() + wait.Wait() + database.Callback().Query().Remove(callbackName) + }) + + wait.Add(1) + go func() { + defer wait.Done() + firstContext := context.WithValue(ctx, demoSeedRunContextKey{}, firstDemoSeedRun) + _, err := Demo(database.WithContext(firstContext), DemoOptions{Email: email, DisplayName: "不应覆盖", Password: "new-password"}) + firstResult <- err + }() + waitForDemoSeedStage(t, ctx, firstAtChildQuery, "first Demo to hold the owner lock and reach its first child query") + + wait.Add(1) + go func() { + defer wait.Done() + secondContext := context.WithValue(ctx, demoSeedRunContextKey{}, secondDemoSeedRun) + _, err := Demo(database.WithContext(secondContext), DemoOptions{Email: email, DisplayName: "不应覆盖", Password: "new-password"}) + secondResult <- err + }() + waitForDemoSeedStage(t, ctx, secondOwnerQueryAttempted, "second Demo to attempt the owner FOR UPDATE query") + + select { + case <-secondEnteredChildQuery: + t.Fatal("second Demo entered a child query before the first released the owner lock") + case err := <-secondResult: + t.Fatalf("second Demo returned before the first released the owner lock: %v", err) + case <-time.After(200 * time.Millisecond): + case <-ctx.Done(): + t.Fatal("timed out while proving the second Demo is blocked on the owner lock") + } + releaseFirstRun() + require.NoError(t, waitForDemoSeedResult(t, ctx, firstResult, "first Demo result")) + require.NoError(t, waitForDemoSeedResult(t, ctx, secondResult, "second Demo result")) + + var storedUser models.SenlinAgentUser + require.NoError(t, database.Where("email = ?", email).First(&storedUser).Error) + require.Equal(t, "保留用户名称", storedUser.DisplayName) + require.NoError(t, bcrypt.CompareHashAndPassword([]byte(storedUser.PasswordHash), []byte("original-password"))) + require.Error(t, bcrypt.CompareHashAndPassword([]byte(storedUser.PasswordHash), []byte("new-password"))) + + projectIDs := []uint{projects[0].ID, projects[1].ID, projects[2].ID, projects[3].ID} + assertPostgresDemoCount(t, database, &models.SenlinAgentProject{}, "owner_id = ?", []any{user.ID}, 4) + assertPostgresDemoCount(t, database, &models.SenlinAgentTag{}, "project_id = ?", []any{projects[0].ID}, 4) + assertPostgresDemoCount(t, database, &models.SenlinAgentTask{}, "project_id IN ?", []any{projectIDs}, 6) + assertPostgresDemoCount(t, database, &models.SenlinAgentInboxItem{}, "project_id IN ?", []any{projectIDs}, 7) + assertPostgresDemoCount(t, database, &models.SenlinAgentAISession{}, "project_id = ?", []any{projects[0].ID}, 4) + assertPostgresDemoCount(t, database, &models.SenlinAgentNote{}, "project_id = ?", []any{projects[0].ID}, 2) + assertPostgresDemoCount(t, database, &models.SenlinAgentSource{}, "project_id = ?", []any{projects[0].ID}, 2) + assertPostgresDemoCount(t, database, &models.SenlinAgentCronPlan{}, "project_id = ?", []any{projects[0].ID}, 3) + assertPostgresDemoCount(t, database, &models.SenlinAgentProjectChannel{}, "project_id = ?", []any{projects[0].ID}, 2) +} + +func signalDemoSeedStage(stage chan<- struct{}) { + select { + case stage <- struct{}{}: + default: + } +} + +func waitForDemoSeedStage(t *testing.T, ctx context.Context, stage <-chan struct{}, description string) { + t.Helper() + select { + case <-stage: + case <-ctx.Done(): + t.Fatalf("timed out waiting for %s", description) + } +} + +func waitForDemoSeedResult(t *testing.T, ctx context.Context, result <-chan error, description string) error { + t.Helper() + select { + case err := <-result: + return err + case <-ctx.Done(): + t.Fatalf("timed out waiting for %s", description) + return ctx.Err() + } +} + +func assertPostgresDemoCount(t *testing.T, database *gorm.DB, model any, where string, args []any, expected int64) { + t.Helper() + var count int64 + require.NoError(t, database.Model(model).Where(where, args...).Count(&count).Error) + require.Equal(t, expected, count) +} + +func cleanupPostgresDemoSeed(database *gorm.DB, userID uint) { + var projects []models.SenlinAgentProject + if database.Where("owner_id = ?", userID).Find(&projects).Error != nil { + return + } + projectIDs := make([]uint, 0, len(projects)) + for _, project := range projects { + projectIDs = append(projectIDs, project.ID) + } + if len(projectIDs) != 0 { + var tasks []models.SenlinAgentTask + database.Where("project_id IN ?", projectIDs).Find(&tasks) + taskIDs := make([]uint, 0, len(tasks)) + for _, task := range tasks { + taskIDs = append(taskIDs, task.ID) + } + if len(taskIDs) != 0 { + database.Where("task_id IN ?", taskIDs).Delete(&models.SenlinAgentTaskShare{}) + } + var inboxItems []models.SenlinAgentInboxItem + database.Where("project_id IN ?", projectIDs).Find(&inboxItems) + inboxIDs := make([]uint, 0, len(inboxItems)) + for _, item := range inboxItems { + inboxIDs = append(inboxIDs, item.ID) + } + if len(inboxIDs) != 0 { + database.Where("inbox_item_id IN ?", inboxIDs).Delete(&models.SenlinAgentInboxSuggestion{}) + } + database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentProjectEvent{}) + database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentTask{}) + database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentNote{}) + database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentSource{}) + database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentAISession{}) + database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentTag{}) + database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentProjectChannel{}) + database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentCronPlan{}) + database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentInboxItem{}) + database.Where("id IN ?", projectIDs).Delete(&models.SenlinAgentProject{}) + } + database.Delete(&models.SenlinAgentUser{}, userID) +} diff --git a/backend/internal/seed/demo_test.go b/backend/internal/seed/demo_test.go index ef90984..6d0e10b 100644 --- a/backend/internal/seed/demo_test.go +++ b/backend/internal/seed/demo_test.go @@ -37,6 +37,9 @@ func TestDemoSeedCreatesFrontendWorkspaceData(t *testing.T) { require.GreaterOrEqual(t, len(workspace.AISessions), 4) require.GreaterOrEqual(t, len(workspace.NotesSources), 3) require.GreaterOrEqual(t, len(workspace.CronPlans), 3) + for _, plan := range workspace.CronPlans { + require.Empty(t, plan.LastResult, "演示数据不得声称计划已执行") + } require.Len(t, workspace.Channels, 8) }