Files
ai-agent/internal/services/ai_agent_service.go
T
t 18c9354095 refactor: 将客服后端重构为宿主可嵌入模块
- 注入数据库、运行时配置、统一响应、文件存储和平台 AI 能力,补充业务读写工具与客户快捷操作契约。

- 移除模块内重复的组织、客户、工单、标签、技能、旧工作流、MCP 和迁移实现,将身份权限与业务主体交由宿主管理。

- 使用 libSQL 重构向量存储,并完善图片消息、访客身份、排队调度、企业微信和支持聊天页面。

- 统一 HTTP、DTO 与 WebSocket 的 snake_case 协议,补齐模块初始化、业务动作和公共载荷等回归测试。
2026-08-28 22:23:13 +08:00

477 lines
16 KiB
Go

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(),
})
}