Files
ai-agent/internal/services/ai_workflow_service.go
T
2026-06-22 00:11:26 +08:00

220 lines
6.6 KiB
Go

package services
import (
"crypto/sha256"
"encoding/hex"
"encoding/json"
"strings"
"time"
"agent-desk/internal/ai/workflow/dsl"
workflowregistry "agent-desk/internal/ai/workflow/registry"
workflowvalidator "agent-desk/internal/ai/workflow/validator"
"agent-desk/internal/models"
"agent-desk/internal/pkg/dto"
"agent-desk/internal/pkg/dto/request"
"agent-desk/internal/pkg/enums"
"agent-desk/internal/pkg/errorsx"
"agent-desk/internal/pkg/httpx/params"
"agent-desk/internal/pkg/utils"
"agent-desk/internal/repositories"
"github.com/mlogclub/simple/sqls"
)
var AIWorkflowService = newAIWorkflowService()
func newAIWorkflowService() *aiWorkflowService {
return &aiWorkflowService{
registry: workflowregistry.DefaultRegistry(),
}
}
type aiWorkflowService struct {
registry *workflowregistry.Registry
}
func (s *aiWorkflowService) Get(id int64) *models.AIWorkflow {
if id <= 0 {
return nil
}
return repositories.AIWorkflowRepository.Get(sqls.DB(), id)
}
func (s *aiWorkflowService) GetVersion(id int64) *models.AIWorkflowVersion {
if id <= 0 {
return nil
}
return repositories.AIWorkflowVersionRepository.Get(sqls.DB(), id)
}
func (s *aiWorkflowService) FindPageByCnd(cnd *sqls.Cnd) (list []models.AIWorkflow, paging *sqls.Paging) {
return repositories.AIWorkflowRepository.FindPageByCnd(sqls.DB(), cnd)
}
func (s *aiWorkflowService) FindVersionPageByParams(params *params.QueryParams) (list []models.AIWorkflowVersion, paging *sqls.Paging) {
return repositories.AIWorkflowVersionRepository.FindPageByParams(sqls.DB(), params)
}
func (s *aiWorkflowService) ListNodeSpecs() []workflowregistry.NodeSpec {
return s.registry.List()
}
func (s *aiWorkflowService) ValidateDefinition(def dsl.Definition) workflowvalidator.Result {
return workflowvalidator.ValidateDefinition(def, s.registry)
}
func (s *aiWorkflowService) CreateWorkflow(req request.CreateAIWorkflowRequest, operator *dto.AuthPrincipal) (*models.AIWorkflow, error) {
if operator == nil {
return nil, errorsx.UnauthorizedI18n("error.auth.expired")
}
name := strings.TrimSpace(req.Name)
if name == "" {
return nil, errorsx.InvalidParam("workflow name is required")
}
ownerType := normalizeWorkflowOwnerType(req.OwnerType)
if ownerType == "" {
return nil, errorsx.InvalidParam("workflow owner type is required")
}
if req.OwnerID <= 0 {
return nil, errorsx.InvalidParam("workflow owner id is required")
}
definition, err := marshalDefinition(req.Definition)
if err != nil {
return nil, err
}
item := &models.AIWorkflow{
Name: name,
Description: strings.TrimSpace(req.Description),
OwnerType: ownerType,
OwnerID: req.OwnerID,
Status: enums.StatusOk,
DraftDefinition: definition,
AuditFields: utils.BuildAuditFields(operator),
}
if err := repositories.AIWorkflowRepository.Create(sqls.DB(), item); err != nil {
return nil, err
}
return item, nil
}
func (s *aiWorkflowService) UpdateWorkflow(req request.UpdateAIWorkflowRequest, operator *dto.AuthPrincipal) error {
if operator == nil {
return errorsx.UnauthorizedI18n("error.auth.expired")
}
if s.Get(req.ID) == nil {
return errorsx.InvalidParamI18n("error.e0002")
}
name := strings.TrimSpace(req.Name)
if name == "" {
return errorsx.InvalidParam("workflow name is required")
}
ownerType := normalizeWorkflowOwnerType(req.OwnerType)
if ownerType == "" {
return errorsx.InvalidParam("workflow owner type is required")
}
if req.OwnerID <= 0 {
return errorsx.InvalidParam("workflow owner id is required")
}
definition, err := marshalDefinition(req.Definition)
if err != nil {
return err
}
return repositories.AIWorkflowRepository.Updates(sqls.DB(), req.ID, map[string]interface{}{
"name": name,
"description": strings.TrimSpace(req.Description),
"owner_type": ownerType,
"owner_id": req.OwnerID,
"draft_definition": definition,
"update_user_id": operator.UserID,
"update_user_name": operator.Username,
"updated_at": time.Now(),
})
}
func (s *aiWorkflowService) DeleteWorkflow(id int64, operator *dto.AuthPrincipal) error {
if operator == nil {
return errorsx.UnauthorizedI18n("error.auth.expired")
}
if s.Get(id) == nil {
return errorsx.InvalidParamI18n("error.e0002")
}
return repositories.AIWorkflowRepository.Updates(sqls.DB(), id, map[string]interface{}{
"status": enums.StatusDeleted,
"update_user_id": operator.UserID,
"update_user_name": operator.Username,
"updated_at": time.Now(),
})
}
func (s *aiWorkflowService) PublishWorkflow(req request.PublishAIWorkflowRequest, operator *dto.AuthPrincipal) (*models.AIWorkflowVersion, error) {
if operator == nil {
return nil, errorsx.UnauthorizedI18n("error.auth.expired")
}
workflow := s.Get(req.WorkflowID)
if workflow == nil || workflow.Status == enums.StatusDeleted {
return nil, errorsx.InvalidParamI18n("error.e0002")
}
result := s.ValidateDefinition(req.Definition)
if !result.Valid {
return nil, errorsx.InvalidParam("workflow definition is invalid")
}
definition, err := marshalDefinition(req.Definition)
if err != nil {
return nil, err
}
now := time.Now()
var version *models.AIWorkflowVersion
err = sqls.WithTransaction(func(ctx *sqls.TxContext) error {
nextVersion := repositories.AIWorkflowVersionRepository.MaxVersionByWorkflowID(ctx.Tx, req.WorkflowID) + 1
version = &models.AIWorkflowVersion{
WorkflowID: req.WorkflowID,
Version: nextVersion,
Status: enums.StatusOk,
Definition: definition,
DefinitionHash: hashDefinition(definition),
PublishedAt: &now,
PublishedByID: operator.UserID,
PublishedByName: operator.Username,
AuditFields: utils.BuildAuditFields(operator),
}
if err := repositories.AIWorkflowVersionRepository.Create(ctx.Tx, version); err != nil {
return err
}
return repositories.AIWorkflowRepository.Updates(ctx.Tx, req.WorkflowID, map[string]interface{}{
"draft_definition": definition,
"published_version_id": version.ID,
"update_user_id": operator.UserID,
"update_user_name": operator.Username,
"updated_at": now,
})
})
if err != nil {
return nil, err
}
return version, nil
}
func marshalDefinition(def dsl.Definition) (string, error) {
buf, err := json.Marshal(def)
if err != nil {
return "", errorsx.InvalidParam("invalid workflow definition")
}
return string(buf), nil
}
func hashDefinition(definition string) string {
sum := sha256.Sum256([]byte(definition))
return hex.EncodeToString(sum[:])
}
func normalizeWorkflowOwnerType(ownerType string) string {
ownerType = strings.TrimSpace(ownerType)
switch ownerType {
case "ai_agent", "workspace":
return ownerType
default:
return ""
}
}