fix(backend): harden final MVP invariants
This commit is contained in:
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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))),
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"errors"
|
||||
"log"
|
||||
"net/http"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -14,13 +15,13 @@ 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"`
|
||||
StorageKey string `json:"storageKey"`
|
||||
CreatedAt time.Time `json:"createdAt"`
|
||||
UpdatedAt time.Time `json:"updatedAt"`
|
||||
}
|
||||
@@ -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(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -1,13 +1,13 @@
|
||||
package files
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
"senlinai-agent/backend/internal/models"
|
||||
@@ -17,8 +17,9 @@ type Service struct {
|
||||
root string
|
||||
db *gorm.DB
|
||||
createTemp func(string, string) (stagedFile, error)
|
||||
rename func(string, string) 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
|
||||
}
|
||||
@@ -42,44 +44,72 @@ 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,
|
||||
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)
|
||||
for attempt := 0; attempt < storageKeyAttempts; attempt++ {
|
||||
storageKey, err := s.newStorageKey()
|
||||
if err != nil {
|
||||
s.cleanupOwnedFiles(file.Name())
|
||||
return StoredFile{}, err
|
||||
}
|
||||
return StoredFile{OriginalName: cleanName, RelativePath: relative, AbsolutePath: absolute}, nil
|
||||
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
|
||||
}
|
||||
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)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
_, err := service.Save(12, "broken.bin", strings.NewReader("content"))
|
||||
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
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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"`
|
||||
|
||||
@@ -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 !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, err
|
||||
}
|
||||
tag := &models.SenlinAgentTag{ProjectID: projectID, Name: name}
|
||||
return tag, models.DBService.Create(tag).Error
|
||||
if !retryableTagWrite(lastErr) {
|
||||
return nil, lastErr
|
||||
}
|
||||
time.Sleep(time.Duration(attempt+1) * time.Millisecond)
|
||||
}
|
||||
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
|
||||
}
|
||||
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) {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -236,14 +236,15 @@ func findOrCreateProjectTag(tx *gorm.DB, projectID uint, value string) (*uint, *
|
||||
if name == "" {
|
||||
return nil, nil, "", nil
|
||||
}
|
||||
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 {
|
||||
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
|
||||
}
|
||||
} else if err != nil {
|
||||
var tag models.SenlinAgentTag
|
||||
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
|
||||
|
||||
150
backend/internal/models/migrations.go
Normal file
150
backend/internal/models/migrations.go
Normal file
@@ -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
|
||||
}
|
||||
84
backend/internal/models/migrations_test.go
Normal file
84
backend/internal/models/migrations_test.go
Normal file
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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) {
|
||||
|
||||
230
backend/internal/seed/demo_postgres_test.go
Normal file
230
backend/internal/seed/demo_postgres_test.go
Normal file
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user