acef46063c
GetTasks 返回的 TaskResponse.ID 使用了数据库自增 ID (t.ID) 而非 ExternalID, 导致前端用错误 ID 调用 GetTask 接口时 external_id 查询不到记录。 改为返回 t.ExternalID 与实际查询字段一致。
273 lines
8.0 KiB
Go
273 lines
8.0 KiB
Go
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: t.ExternalID,
|
|
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),
|
|
}
|
|
}
|