package service import ( "context" "errors" "fmt" "sync" "time" "gen2d/internal/model" "gorm.io/gorm" ) var ( projectService *ProjectService projectServiceOnce sync.Once ) // ProjectService 工程管理服务。 type ProjectService struct { db *gorm.DB } // InitProjectService 初始化工程管理服务。 func InitProjectService(db *gorm.DB) { projectServiceOnce.Do(func() { projectService = &ProjectService{db: db} }) } // GetProjectService 获取工程管理服务实例。 func GetProjectService() *ProjectService { return projectService } // CreateProject 创建工程。 func (s *ProjectService) CreateProject(ctx context.Context, userID uint, name string, style map[string]string) (*model.ProjectResponse, error) { project := &model.Project{ UserID: userID, Name: name, CreatedAt: time.Now(), } if err := s.db.WithContext(ctx).Create(project).Error; err != nil { return nil, fmt.Errorf("create project failed: %w", err) } // 保存风格键值对 if err := s.saveStyle(ctx, project.ID, style); err != nil { // 回滚工程创建 s.db.WithContext(ctx).Delete(project) return nil, fmt.Errorf("save style failed: %w", err) } return s.toProjectResponse(project, style), nil } // ListProjects 获取用户工程列表。 func (s *ProjectService) ListProjects(ctx context.Context, userID uint, page, pageSize int) (*model.ProjectsListResponse, error) { var projects []model.Project var total int64 offset := (page - 1) * pageSize if err := s.db.WithContext(ctx).Model(&model.Project{}).Where("user_id = ?", userID).Count(&total).Error; err != nil { return nil, fmt.Errorf("count projects failed: %w", err) } if err := s.db.WithContext(ctx). Where("user_id = ?", userID). Order("created_at DESC"). Limit(pageSize). Offset(offset). Find(&projects).Error; err != nil { return nil, fmt.Errorf("list projects failed: %w", err) } response := &model.ProjectsListResponse{ Total: int(total), Projects: make([]model.ProjectResponse, len(projects)), } for i, p := range projects { style, _ := s.getStyle(ctx, p.ID) response.Projects[i] = *s.toProjectResponse(&p, style) } return response, nil } // GetProject 获取工程详情。 func (s *ProjectService) GetProject(ctx context.Context, userID uint, projectID uint) (*model.ProjectResponse, error) { var project model.Project if err := s.db.WithContext(ctx).Where("id = ? AND user_id = ?", projectID, userID).First(&project).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return nil, errors.New("project not found") } return nil, fmt.Errorf("get project failed: %w", err) } style, _ := s.getStyle(ctx, project.ID) return s.toProjectResponse(&project, style), nil } // UpdateProject 更新工程信息。 func (s *ProjectService) UpdateProject(ctx context.Context, userID uint, projectID uint, name string) (*model.ProjectResponse, error) { var project model.Project if err := s.db.WithContext(ctx).Where("id = ? AND user_id = ?", projectID, userID).First(&project).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return nil, errors.New("project not found") } return nil, fmt.Errorf("get project failed: %w", err) } project.Name = name if err := s.db.WithContext(ctx).Save(&project).Error; err != nil { return nil, fmt.Errorf("update project failed: %w", err) } style, _ := s.getStyle(ctx, project.ID) return s.toProjectResponse(&project, style), nil } // DeleteProject 删除工程。 func (s *ProjectService) DeleteProject(ctx context.Context, userID uint, projectID uint) error { var project model.Project if err := s.db.WithContext(ctx).Where("id = ? AND user_id = ?", projectID, userID).First(&project).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return errors.New("project not found") } return fmt.Errorf("get project failed: %w", err) } return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { // 删除风格记录 if err := tx.Where("project_id = ?", projectID).Delete(&model.ProjectStyleRecord{}).Error; err != nil { return fmt.Errorf("delete style records failed: %w", err) } // TODO: 删除任务和素材记录 // TODO: 删除七牛云对象 // 删除工程 if err := tx.Delete(&project).Error; err != nil { return fmt.Errorf("delete project failed: %w", err) } return nil }) } // GetStyle 获取工程风格。 func (s *ProjectService) GetStyle(ctx context.Context, userID uint, projectID uint) (map[string]string, error) { var project model.Project if err := s.db.WithContext(ctx).Where("id = ? AND user_id = ?", projectID, userID).First(&project).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return nil, errors.New("project not found") } return nil, fmt.Errorf("get project failed: %w", err) } style, _ := s.getStyle(ctx, project.ID) return style, nil } // UpdateStyle 更新工程风格。 func (s *ProjectService) UpdateStyle(ctx context.Context, userID uint, projectID uint, kvPairs map[string]string) error { var project model.Project if err := s.db.WithContext(ctx).Where("id = ? AND user_id = ?", projectID, userID).First(&project).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return errors.New("project not found") } return fmt.Errorf("get project failed: %w", err) } return s.saveStyle(ctx, projectID, kvPairs) } // GetTasks 获取工程下的任务列表。 func (s *ProjectService) GetTasks(ctx context.Context, userID uint, projectID uint) ([]model.TaskResponse, error) { var project model.Project if err := s.db.WithContext(ctx).Where("id = ? AND user_id = ?", projectID, userID).First(&project).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return nil, errors.New("project not found") } return nil, fmt.Errorf("get project failed: %w", err) } var tasks []model.Task if err := s.db.WithContext(ctx). Where("project_id = ?", projectID). Order("created_at DESC"). Find(&tasks).Error; err != nil { return nil, fmt.Errorf("get tasks failed: %w", err) } response := make([]model.TaskResponse, len(tasks)) for i, t := range tasks { response[i] = model.TaskResponse{ ID: fmt.Sprintf("%d", t.ID), ProjectID: fmt.Sprintf("%d", t.ProjectID), Prompt: t.Prompt, AssetType: t.AssetType, Status: t.Status, Stage: t.Stage, Progress: t.Progress, RetryCount: t.RetryCount, Error: t.Error, CreatedAt: t.CreatedAt.Format("2006-01-02T15:04:05Z07:00"), UpdatedAt: t.UpdatedAt.Format("2006-01-02T15:04:05Z07:00"), } } return response, nil } func (s *ProjectService) saveStyle(ctx context.Context, projectID uint, style map[string]string) error { return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { // 删除现有风格记录 if err := tx.Where("project_id = ?", projectID).Delete(&model.ProjectStyleRecord{}).Error; err != nil { return err } // 插入新记录 for key, value := range style { record := &model.ProjectStyleRecord{ ProjectID: projectID, Key: key, Value: value, } if err := tx.Create(record).Error; err != nil { return err } } return nil }) } func (s *ProjectService) getStyle(ctx context.Context, projectID uint) (map[string]string, error) { var records []model.ProjectStyleRecord if err := s.db.WithContext(ctx).Where("project_id = ?", projectID).Find(&records).Error; err != nil { return nil, err } style := make(map[string]string) for _, r := range records { style[r.Key] = r.Value } return style, nil } func (s *ProjectService) toProjectResponse(project *model.Project, style map[string]string) *model.ProjectResponse { // 获取任务数量 var taskCount int64 s.db.Model(&model.Task{}).Where("project_id = ?", project.ID).Count(&taskCount) styleResp := &model.ProjectStyleResponse{ KvPairs: style, } return &model.ProjectResponse{ ID: fmt.Sprintf("%d", project.ID), Name: project.Name, Style: styleResp, CreatedAt: project.CreatedAt.Format("2006-01-02T15:04:05Z07:00"), UpdatedAt: project.UpdatedAt.Format("2006-01-02T15:04:05Z07:00"), TaskCount: int(taskCount), } }