package services import ( "context" "encoding/json" "slices" "strings" "time" "code.tczkiot.com/wlw/ai-agent/internal/models" "code.tczkiot.com/wlw/ai-agent/internal/pkg/dto" "code.tczkiot.com/wlw/ai-agent/internal/pkg/dto/request" "code.tczkiot.com/wlw/ai-agent/internal/pkg/enums" "code.tczkiot.com/wlw/ai-agent/internal/pkg/errorsx" "code.tczkiot.com/wlw/ai-agent/internal/pkg/utils" "code.tczkiot.com/wlw/ai-agent/internal/repositories" "code.tczkiot.com/wlw/ai-agent/internal/pkg/httpx/params" "github.com/mlogclub/simple/sqls" "gorm.io/gorm" ) var AIAgentService = newAIAgentService() const defaultNewAgentRolloutPercent = 5 func newAIAgentService() *aIAgentService { return &aIAgentService{} } type aIAgentService struct { } func (s *aIAgentService) Get(id int64) *models.AIAgent { if id <= 0 { return nil } return repositories.AIAgentRepository.Get(sqls.DB(), id) } func (s *aIAgentService) Take(where ...interface{}) *models.AIAgent { return repositories.AIAgentRepository.Take(sqls.DB(), where...) } func (s *aIAgentService) Find(cnd *sqls.Cnd) []models.AIAgent { return repositories.AIAgentRepository.Find(sqls.DB(), cnd) } func (s *aIAgentService) FindOne(cnd *sqls.Cnd) *models.AIAgent { return repositories.AIAgentRepository.FindOne(sqls.DB(), cnd) } func (s *aIAgentService) FindPageByParams(params *params.QueryParams) (list []models.AIAgent, paging *sqls.Paging) { return repositories.AIAgentRepository.FindPageByParams(sqls.DB(), params) } func (s *aIAgentService) FindPageByCnd(cnd *sqls.Cnd) (list []models.AIAgent, paging *sqls.Paging) { return repositories.AIAgentRepository.FindPageByCnd(sqls.DB(), cnd) } func (s *aIAgentService) Count(cnd *sqls.Cnd) int64 { return repositories.AIAgentRepository.Count(sqls.DB(), cnd) } func (s *aIAgentService) FindByIds(ids []int64) []models.AIAgent { return repositories.AIAgentRepository.FindByIds(sqls.DB(), ids) } func (s *aIAgentService) CreateAIAgent(req request.CreateAIAgentRequest, operator *dto.AuthPrincipal) (*models.AIAgent, error) { if operator == nil { return nil, errorsx.UnauthorizedI18n("error.auth.expired") } item, err := s.buildAIAgentModel(0, req) if err != nil { return nil, err } item.Status = enums.StatusOk item.SortNo = 0 item.AuditFields = utils.BuildAuditFields(operator) if err := repositories.AIAgentRepository.Create(sqls.DB(), item); err != nil { return nil, err } return item, nil } func (s *aIAgentService) UpdateAIAgent(req request.UpdateAIAgentRequest, operator *dto.AuthPrincipal) error { if operator == nil { return errorsx.UnauthorizedI18n("error.auth.expired") } current := s.Get(req.ID) if current == nil { return errorsx.InvalidParamI18n("error.e0002") } item, err := s.buildAIAgentModel(req.ID, req.CreateAIAgentRequest) if err != nil { return err } columns := map[string]any{ "name": item.Name, "avatar": item.Avatar, "description": item.Description, "ai_config_id": item.AIConfigID, "max_steps": item.MaxSteps, "context_window": item.ContextWindow, "tool_policy": item.ToolPolicy, "knowledge_policy": item.KnowledgePolicy, "service_mode": item.ServiceMode, "system_prompt": item.SystemPrompt, "welcome_message": item.WelcomeMessage, "reply_timeout_seconds": item.ReplyTimeoutSeconds, "rollout_percent": item.RolloutPercent, "team_ids": item.TeamIDs, "handoff_mode": item.HandoffMode, "fallback_mode": item.FallbackMode, "fallback_message": item.FallbackMessage, "knowledge_ids": item.KnowledgeIDs, "update_user_id": operator.UserID, "update_user_name": operator.Username, "updated_at": time.Now(), } if item.RolloutPercent != current.RolloutPercent { columns["previous_rollout_percent"] = current.RolloutPercent } return repositories.AIAgentRepository.Updates(sqls.DB(), req.ID, columns) } func (s *aIAgentService) DeleteAIAgent(id int64, operator *dto.AuthPrincipal) error { current := s.Get(id) if current == nil { return errorsx.InvalidParamI18n("error.e0002") } if ChannelService.Take("ai_agent_id = ?", id) != nil { return errorsx.ForbiddenI18n("error.e0185") } return repositories.AIAgentRepository.Updates(sqls.DB(), id, map[string]any{ "status": enums.StatusDeleted, "update_user_id": operator.UserID, "update_user_name": operator.Username, "updated_at": time.Now(), }) } // PublishAIAgent snapshots the complete Agent capability set before it can // receive traffic. func (s *aIAgentService) PublishAIAgent(id int64, operator *dto.AuthPrincipal) (*models.AgentRevision, error) { if operator == nil { return nil, errorsx.UnauthorizedI18n("error.auth.expired") } var revision *models.AgentRevision err := sqls.WithTransaction(func(ctx *sqls.TxContext) error { agent := repositories.AIAgentRepository.Get(ctx.Tx, id) if agent == nil || agent.Status != enums.StatusOk { return errorsx.InvalidParamI18n("error.e0002") } if err := s.validatePublishableAgent(ctx.Tx, agent); err != nil { return err } var err error revision, err = AgentRevisionService.PublishSnapshot(ctx.Tx, agent, operator) if err != nil { return err } return repositories.AIAgentRepository.Updates(ctx.Tx, agent.ID, map[string]any{ "published_revision_id": revision.ID, "update_user_id": operator.UserID, "update_user_name": operator.Username, "updated_at": time.Now(), }) }) if err != nil { return nil, err } return revision, nil } func (s *aIAgentService) validatePublishableAgent(db *gorm.DB, agent *models.AIAgent) error { if agent == nil { return errorsx.InvalidParam("ai agent is required before publishing") } platform, err := PlatformAIService.IsPlatform(context.Background()) if err != nil { return errorsx.InvalidParam("failed to resolve AI model source") } if !platform && agent.AIConfigID <= 0 { return errorsx.InvalidParam("ai agent model configuration is required before publishing") } if !platform { config := repositories.AIConfigRepository.Get(db, agent.AIConfigID) if config == nil || config.Status != enums.StatusOk { return errorsx.InvalidParam("ai agent model configuration is unavailable") } } if _, err := s.normalizeToolPolicy(agent.ToolPolicy); err != nil { return err } return nil } // RollbackAIAgent switches an Agent back to a previously published immutable // revision. It never rewrites the historical snapshot itself. func (s *aIAgentService) RollbackAIAgent(id, revisionID int64, operator *dto.AuthPrincipal) error { if operator == nil { return errorsx.UnauthorizedI18n("error.auth.expired") } if id <= 0 || revisionID <= 0 { return errorsx.InvalidParam("agent id and revision id are required") } return sqls.WithTransaction(func(ctx *sqls.TxContext) error { agent := repositories.AIAgentRepository.Get(ctx.Tx, id) if agent == nil || agent.Status != enums.StatusOk { return errorsx.InvalidParamI18n("error.e0002") } revision := repositories.AgentRevisionRepository.Get(ctx.Tx, revisionID) if revision == nil || revision.AgentID != agent.ID || revision.Status != enums.StatusOk { return errorsx.InvalidParam("agent revision does not exist") } updates := map[string]any{ "published_revision_id": revision.ID, "update_user_id": operator.UserID, "update_user_name": operator.Username, "updated_at": time.Now(), } return repositories.AIAgentRepository.Updates(ctx.Tx, agent.ID, updates) }) } // RollbackAIAgentRollout restores the prior Agent rollout percentage and // swaps it into history, allowing operators to undo and redo one rollout // change without rewriting an immutable AgentRevision. func (s *aIAgentService) RollbackAIAgentRollout(id int64, operator *dto.AuthPrincipal) error { if operator == nil { return errorsx.UnauthorizedI18n("error.auth.expired") } if id <= 0 { return errorsx.InvalidParam("agent id is required") } return sqls.WithTransaction(func(ctx *sqls.TxContext) error { agent := repositories.AIAgentRepository.Get(ctx.Tx, id) if agent == nil || agent.Status != enums.StatusOk { return errorsx.InvalidParamI18n("error.e0002") } if agent.PreviousRolloutPercent < 1 || agent.PreviousRolloutPercent > 100 { return errorsx.InvalidParam("agent rollout has no previous value to restore") } return repositories.AIAgentRepository.Updates(ctx.Tx, agent.ID, map[string]any{ "rollout_percent": agent.PreviousRolloutPercent, "previous_rollout_percent": agent.RolloutPercent, "update_user_id": operator.UserID, "update_user_name": operator.Username, "updated_at": time.Now(), }) }) } func (s *aIAgentService) buildAIAgentModel(id int64, req request.CreateAIAgentRequest) (*models.AIAgent, error) { name := strings.TrimSpace(req.Name) if name == "" { return nil, errorsx.InvalidParamI18n("error.e0005") } if exists := s.Take("name = ? AND id <> ?", name, id); exists != nil { return nil, errorsx.InvalidParamI18n("error.e0006") } avatar := strings.TrimSpace(req.Avatar) if len(avatar) > 1024 { return nil, errorsx.InvalidParam("ai agent avatar URL must not exceed 1024 characters") } platform, err := PlatformAIService.IsPlatform(context.Background()) if err != nil { return nil, errorsx.InvalidParam("failed to resolve AI model source") } if platform && req.AIConfigID <= 0 && id > 0 { if current := s.Get(id); current != nil { req.AIConfigID = current.AIConfigID } } if !platform { if req.AIConfigID <= 0 { return nil, errorsx.InvalidParamI18n("error.e0010") } aiConfig := AIConfigService.Get(req.AIConfigID) if aiConfig == nil { return nil, errorsx.InvalidParamI18n("error.e0009") } if aiConfig.Status != enums.StatusOk { return nil, errorsx.InvalidParamI18n("error.e0011") } } if req.MaxSteps == 0 { req.MaxSteps = 6 } if req.MaxSteps < 1 || req.MaxSteps > 8 { return nil, errorsx.InvalidParam("ai agent max steps must be between 1 and 8") } if req.ContextWindow < 0 { return nil, errorsx.InvalidParam("ai agent context window must not be negative") } toolPolicy, err := s.normalizeToolPolicy(req.ToolPolicy) if err != nil { return nil, err } // AI Agent 只描述 AI 的能力与转人工策略;渠道决定使用人工还是 AI 接待。 req.ServiceMode = enums.IMConversationServiceModeAIFirst teamIDs, err := s.normalizeTeamIDs(req.TeamIDs) if err != nil { return nil, err } if !slices.Contains(enums.AIAgentHandoffModeValues, enums.AIAgentHandoffMode(req.HandoffMode)) { return nil, errorsx.InvalidParamI18n("error.e0336") } if req.FallbackMode == 0 { req.FallbackMode = enums.AIAgentFallbackModeNoAnswer } if !slices.Contains(enums.AIAgentFallbackModeValues, enums.AIAgentFallbackMode(req.FallbackMode)) { return nil, errorsx.InvalidParamI18n("error.e0123") } if enums.AIAgentHandoffMode(req.HandoffMode) == enums.AIAgentHandoffModeDefaultTeamPool && len(teamIDs) == 0 { return nil, errorsx.InvalidParamI18n("error.e0347") } if req.ReplyTimeoutSeconds < 0 { return nil, errorsx.InvalidParamI18n("error.e0144") } if req.RolloutPercent == 0 { req.RolloutPercent = defaultNewAgentRolloutPercent } if req.RolloutPercent < 1 || req.RolloutPercent > 100 { return nil, errorsx.InvalidParam("ai agent rollout percent must be between 1 and 100") } knowledgeBaseIDs, err := s.normalizeKnowledgeBaseIDs(req.KnowledgeBaseIDs) if err != nil { return nil, err } return &models.AIAgent{ Name: name, Avatar: avatar, Description: strings.TrimSpace(req.Description), AIConfigID: req.AIConfigID, MaxSteps: req.MaxSteps, ContextWindow: req.ContextWindow, ToolPolicy: toolPolicy, KnowledgePolicy: strings.TrimSpace(req.KnowledgePolicy), ServiceMode: req.ServiceMode, SystemPrompt: strings.TrimSpace(req.SystemPrompt), WelcomeMessage: strings.TrimSpace(req.WelcomeMessage), ReplyTimeoutSeconds: req.ReplyTimeoutSeconds, RolloutPercent: req.RolloutPercent, TeamIDs: utils.JoinInt64s(teamIDs), HandoffMode: req.HandoffMode, FallbackMode: req.FallbackMode, FallbackMessage: strings.TrimSpace(req.FallbackMessage), KnowledgeIDs: utils.JoinInt64s(knowledgeBaseIDs), }, nil } type normalizedAIAgentToolPolicy struct { MaxTotalCalls int `json:"max_total_calls,omitempty"` MaxArgumentBytes int `json:"max_argument_bytes,omitempty"` AllowedRiskLevels []string `json:"allowed_risk_levels,omitempty"` } func (s *aIAgentService) normalizeToolPolicy(raw string) (string, error) { raw = strings.TrimSpace(raw) if raw == "" { return "", nil } policy := normalizedAIAgentToolPolicy{} if err := json.Unmarshal([]byte(raw), &policy); err != nil { return "", errorsx.InvalidParam("ai agent tool policy must be valid JSON") } if policy.MaxTotalCalls < 0 || policy.MaxTotalCalls > 8 { return "", errorsx.InvalidParam("ai agent tool policy max_total_calls must be between 1 and 8") } if policy.MaxArgumentBytes < 0 || policy.MaxArgumentBytes > 64*1024 { return "", errorsx.InvalidParam("ai agent tool policy max_argument_bytes must be between 1 and 65536") } seen := make(map[string]struct{}, len(policy.AllowedRiskLevels)) riskLevels := make([]string, 0, len(policy.AllowedRiskLevels)) for _, level := range policy.AllowedRiskLevels { level = strings.ToLower(strings.TrimSpace(level)) if level == "" { continue } if level != "read" && level != "write" { return "", errorsx.InvalidParam("ai agent tool policy contains an invalid risk level") } if _, exists := seen[level]; exists { continue } seen[level] = struct{}{} riskLevels = append(riskLevels, level) } policy.AllowedRiskLevels = riskLevels data, err := json.Marshal(policy) if err != nil { return "", errorsx.InvalidParam("ai agent tool policy is invalid") } return string(data), nil } func (s *aIAgentService) normalizeKnowledgeBaseIDs(input []int64) ([]int64, error) { ret := make([]int64, 0, len(input)) seen := make(map[int64]struct{}) for _, id := range input { if id <= 0 { continue } if _, exists := seen[id]; exists { continue } knowledgeBase := KnowledgeBaseService.Get(id) if knowledgeBase == nil || knowledgeBase.Status != enums.StatusOk { return nil, errorsx.InvalidParam("knowledge base is not available") } seen[id] = struct{}{} ret = append(ret, id) } slices.Sort(ret) return ret, nil } func (s *aIAgentService) normalizeTeamIDs(input []int64) ([]int64, error) { ret := make([]int64, 0, len(input)) seen := make(map[int64]struct{}) for _, id := range input { if id <= 0 { continue } if _, exists := seen[id]; exists { continue } team := AgentTeamService.Get(id) if team == nil || team.Status == enums.StatusDeleted { continue } // if team.Status != enums.StatusOk { // return nil, errorsx.InvalidParamI18n("error.e0173") // } seen[id] = struct{}{} ret = append(ret, id) } slices.Sort(ret) return ret, nil } func (s *aIAgentService) UpdateSort(ids []int64) error { return sqls.WithTransaction(func(ctx *sqls.TxContext) error { for i, id := range ids { if err := repositories.AIAgentRepository.UpdateColumn(ctx.Tx, id, "sort_no", i+1); err != nil { return err } } return nil }) } func (s *aIAgentService) UpdateStatus(id int64, status int, operator *dto.AuthPrincipal) error { if operator == nil { return errorsx.UnauthorizedI18n("error.auth.expired") } current := s.Get(id) if current == nil { return errorsx.InvalidParamI18n("error.e0002") } if status != int(enums.StatusOk) && status != int(enums.StatusDisabled) { return errorsx.InvalidParamI18n("error.e0254") } return repositories.AIAgentRepository.Updates(sqls.DB(), id, map[string]any{ "status": status, "update_user_id": operator.UserID, "update_user_name": operator.Username, "updated_at": time.Now(), }) }