This commit is contained in:
mlogclub
2026-04-09 10:01:23 +08:00
commit efe801b8bf
707 changed files with 110595 additions and 0 deletions
@@ -0,0 +1,451 @@
package engine
import (
"context"
"encoding/json"
"fmt"
"strings"
"cs-agent/internal/ai/rag"
"cs-agent/internal/ai/runtime/internal/impl/adapter"
"cs-agent/internal/ai/runtime/internal/impl/callbacks"
"cs-agent/internal/ai/runtime/internal/impl/factory"
"cs-agent/internal/ai/runtime/internal/impl/retrievers"
"cs-agent/internal/pkg/utils"
"github.com/cloudwego/eino/adk"
"github.com/cloudwego/eino/schema"
"github.com/google/uuid"
)
type Service struct {
agentFactory *factory.AgentFactory
runnerFactory *factory.RunnerFactory
}
func NewService() *Service {
return &Service{
agentFactory: factory.NewAgentFactory(),
runnerFactory: factory.NewRunnerFactory(),
}
}
func (s *Service) Run(ctx context.Context, req Request) (*Summary, error) {
summary := &Summary{
RunID: uuid.NewString(),
Status: "started",
ToolCodes: make([]string, 0),
InvokedToolCodes: make([]string, 0),
}
collector := callbacks.NewRuntimeTraceCollector()
collector.Data.RunID = summary.RunID
if req.AIAgent == nil || req.Conversation == nil || req.UserMessage == nil {
summary.Status = "error"
summary.ErrorMessage = "invalid runtime request"
collector.Data.Status = summary.Status
collector.Data.Error.Message = summary.ErrorMessage
collector.Data.Error.Stage = "prepare"
summary.TraceData = collector.Marshal()
return summary, fmt.Errorf("%s", summary.ErrorMessage)
}
if req.AIConfig == nil {
summary.Status = "error"
summary.ErrorMessage = "ai config is nil"
collector.Data.Status = summary.Status
collector.Data.Error.Message = summary.ErrorMessage
collector.Data.Error.Stage = "prepare"
summary.TraceData = collector.Marshal()
return summary, fmt.Errorf("%s", summary.ErrorMessage)
}
history := adapter.BuildHistoryMessages(req.Conversation.ID, req.UserMessage.ID, 12)
summary.HistoryMessageCount = len(history.Messages)
collector.Data.Input.HistoryMessageCount = len(history.Messages)
collector.Data.Input.KnowledgeBaseIDs = utils.SplitInt64s(req.AIAgent.KnowledgeIDs)
collector.Data.Input.CurrentUserMessagePreview = preview(req.UserMessage.Content, 120)
toolDefs, err := factory.NewToolFactory().BuildMCPTools(req.AIAgent)
if err != nil {
summary.Status = "error"
summary.ErrorMessage = err.Error()
collector.Data.Status = summary.Status
collector.Data.Error.Message = err.Error()
collector.Data.Error.Stage = "prepare"
summary.TraceData = collector.Marshal()
return summary, err
}
toolDefsByModelName := make(map[string]string, len(toolDefs))
for _, item := range toolDefs {
summary.ToolCodes = append(summary.ToolCodes, item.ToolCode)
toolDefsByModelName[item.ModelName] = item.ToolCode
}
for modelName, toolCode := range req.ExtraToolCodes {
toolCode = strings.TrimSpace(toolCode)
modelName = strings.TrimSpace(modelName)
if toolCode == "" || modelName == "" {
continue
}
summary.ToolCodes = appendIfMissing(summary.ToolCodes, toolCode)
toolDefsByModelName[modelName] = toolCode
}
collector.Data.Input.ToolCodes = append(collector.Data.Input.ToolCodes, summary.ToolCodes...)
collector.Data.Model.Provider = string(req.AIConfig.Provider)
collector.Data.Model.Name = req.AIConfig.ModelName
agent, err := s.agentFactory.BuildCustomerServiceAgent(ctx, req.AIAgent, req.AIConfig, toolDefs, req.ExtraTools, req.ExtraToolCodes, collector)
if err != nil {
summary.Status = "error"
summary.ErrorMessage = err.Error()
collector.Data.Status = summary.Status
collector.Data.Error.Message = err.Error()
collector.Data.Error.Stage = "prepare"
summary.TraceData = collector.Marshal()
return summary, err
}
checkPointID := strings.TrimSpace(req.CheckPointID)
if checkPointID == "" {
checkPointID = "eino_cp_" + summary.RunID
}
summary.CheckPointID = checkPointID
runner := s.runnerFactory.Build(ctx, agent, false, true)
if runner == nil {
summary.Status = "error"
summary.ErrorMessage = "failed to build runner"
collector.Data.Status = summary.Status
collector.Data.Error.Message = summary.ErrorMessage
collector.Data.Error.Stage = "prepare"
summary.TraceData = collector.Marshal()
return summary, fmt.Errorf("%s", summary.ErrorMessage)
}
messages := make([]*schema.Message, 0, len(history.Messages)+3)
messages = append(messages, history.Messages...)
retriever := retrievers.NewKnowledgeRetriever(req.AIAgent)
if results, _, retrieveErr := retriever.Retrieve(ctx, strings.TrimSpace(req.UserMessage.Content)); retrieveErr == nil {
summary.RetrieverCount = len(results)
collector.Data.Retriever.Count = len(results)
for _, item := range results {
collector.Data.Retriever.Items = append(collector.Data.Retriever.Items, callbacks.RetrieverTraceItem{
Query: preview(req.UserMessage.Content, 120),
KnowledgeBaseID: item.KnowledgeBaseID,
DocumentID: item.DocumentID,
DocumentTitle: item.DocumentTitle,
Score: float64(item.Score),
})
}
if knowledgeContext := buildKnowledgeContext(results); knowledgeContext != "" {
messages = append(messages, schema.SystemMessage(knowledgeContext))
}
}
messages = append(messages, schema.UserMessage(strings.TrimSpace(req.UserMessage.Content)))
collector.Data.Interrupt.CheckPointID = checkPointID
consumeAgentEvents(runner.Run(ctx, messages, buildRunOptions(checkPointID)...), summary, collector, toolDefsByModelName)
summary.ModelName = req.AIConfig.ModelName
collector.Data.Status = summary.Status
collector.Data.Output.ReplyText = summary.ReplyText
collector.Data.Output.FinishReason = summary.Status
summary.TraceData = collector.Marshal()
return summary, nil
}
func (s *Service) Resume(ctx context.Context, req ResumeRequest) (*Summary, error) {
summary := &Summary{
RunID: uuid.NewString(),
Status: "started",
CheckPointID: strings.TrimSpace(req.CheckPointID),
ToolCodes: make([]string, 0),
InvokedToolCodes: make([]string, 0),
Interrupts: make([]InterruptContextSummary, 0),
}
collector := callbacks.NewRuntimeTraceCollector()
collector.Data.RunID = summary.RunID
collector.Data.Interrupt.CheckPointID = summary.CheckPointID
if req.AIAgent == nil {
summary.Status = "error"
summary.ErrorMessage = "ai agent is nil"
collector.Data.Status = summary.Status
collector.Data.Error.Message = summary.ErrorMessage
collector.Data.Error.Stage = "resume_prepare"
summary.TraceData = collector.Marshal()
return summary, fmt.Errorf("%s", summary.ErrorMessage)
}
if req.AIConfig == nil {
summary.Status = "error"
summary.ErrorMessage = "ai config is nil"
collector.Data.Status = summary.Status
collector.Data.Error.Message = summary.ErrorMessage
collector.Data.Error.Stage = "resume_prepare"
summary.TraceData = collector.Marshal()
return summary, fmt.Errorf("%s", summary.ErrorMessage)
}
if summary.CheckPointID == "" {
summary.Status = "error"
summary.ErrorMessage = "checkpoint id is required"
collector.Data.Status = summary.Status
collector.Data.Error.Message = summary.ErrorMessage
collector.Data.Error.Stage = "resume_prepare"
summary.TraceData = collector.Marshal()
return summary, fmt.Errorf("%s", summary.ErrorMessage)
}
toolDefs, err := factory.NewToolFactory().BuildMCPTools(req.AIAgent)
if err != nil {
summary.Status = "error"
summary.ErrorMessage = err.Error()
collector.Data.Status = summary.Status
collector.Data.Error.Message = err.Error()
collector.Data.Error.Stage = "resume_prepare"
summary.TraceData = collector.Marshal()
return summary, err
}
toolDefsByModelName := make(map[string]string, len(toolDefs))
for _, item := range toolDefs {
summary.ToolCodes = append(summary.ToolCodes, item.ToolCode)
toolDefsByModelName[item.ModelName] = item.ToolCode
}
for modelName, toolCode := range req.ExtraToolCodes {
toolCode = strings.TrimSpace(toolCode)
modelName = strings.TrimSpace(modelName)
if toolCode == "" || modelName == "" {
continue
}
summary.ToolCodes = appendIfMissing(summary.ToolCodes, toolCode)
toolDefsByModelName[modelName] = toolCode
}
collector.Data.Input.ToolCodes = append(collector.Data.Input.ToolCodes, summary.ToolCodes...)
collector.Data.Model.Provider = string(req.AIConfig.Provider)
collector.Data.Model.Name = req.AIConfig.ModelName
agent, err := s.agentFactory.BuildCustomerServiceAgent(ctx, req.AIAgent, req.AIConfig, toolDefs, req.ExtraTools, req.ExtraToolCodes, collector)
if err != nil {
summary.Status = "error"
summary.ErrorMessage = err.Error()
collector.Data.Status = summary.Status
collector.Data.Error.Message = err.Error()
collector.Data.Error.Stage = "resume_prepare"
summary.TraceData = collector.Marshal()
return summary, err
}
runner := s.runnerFactory.Build(ctx, agent, false, true)
if runner == nil {
summary.Status = "error"
summary.ErrorMessage = "failed to build runner"
collector.Data.Status = summary.Status
collector.Data.Error.Message = summary.ErrorMessage
collector.Data.Error.Stage = "resume_prepare"
summary.TraceData = collector.Marshal()
return summary, fmt.Errorf("%s", summary.ErrorMessage)
}
var iter *adk.AsyncIterator[*adk.AgentEvent]
if len(req.ResumeData) > 0 {
iter, err = runner.ResumeWithParams(ctx, summary.CheckPointID, &adk.ResumeParams{Targets: req.ResumeData})
} else {
iter, err = runner.Resume(ctx, summary.CheckPointID)
}
if err != nil {
summary.Status = "error"
summary.ErrorMessage = err.Error()
collector.Data.Status = summary.Status
collector.Data.Error.Message = err.Error()
collector.Data.Error.Stage = "resume_prepare"
summary.TraceData = collector.Marshal()
return summary, err
}
consumeAgentEvents(iter, summary, collector, toolDefsByModelName)
summary.ModelName = req.AIConfig.ModelName
collector.Data.Status = summary.Status
collector.Data.Output.ReplyText = summary.ReplyText
collector.Data.Output.FinishReason = summary.Status
summary.TraceData = collector.Marshal()
return summary, nil
}
func buildRunOptions(checkPointID string) []adk.AgentRunOption {
if strings.TrimSpace(checkPointID) == "" {
return nil
}
return []adk.AgentRunOption{adk.WithCheckPointID(checkPointID)}
}
func consumeAgentEvents(iter *adk.AsyncIterator[*adk.AgentEvent], summary *Summary, collector *callbacks.RuntimeTraceCollector, toolDefsByModelName map[string]string) {
if iter == nil || summary == nil || collector == nil {
return
}
for {
event, ok := iter.Next()
if !ok {
break
}
if event == nil {
continue
}
if event.Err != nil {
summary.Status = "error"
summary.ErrorMessage = event.Err.Error()
collector.Data.Error.Message = event.Err.Error()
collector.Data.Error.Stage = "model"
continue
}
if event.Action != nil && event.Action.Interrupted != nil {
summary.Status = "interrupted"
summary.Interrupted = true
summary.Interrupts = summarizeInterrupts(event.Action.Interrupted.InterruptContexts)
collector.Data.Interrupt.Items = convertInterruptTraceItems(summary.Interrupts)
continue
}
if event.Output == nil || event.Output.MessageOutput == nil {
continue
}
message, getErr := event.Output.MessageOutput.GetMessage()
if getErr != nil || message == nil {
continue
}
switch event.Output.MessageOutput.Role {
case schema.Assistant:
summary.ReplyText = strings.TrimSpace(message.Content)
case schema.Tool:
summary.ToolCallCount++
if toolDefsByModelName != nil {
toolCode := strings.TrimSpace(toolDefsByModelName[message.ToolName])
if toolCode != "" {
summary.InvokedToolCodes = appendIfMissing(summary.InvokedToolCodes, toolCode)
}
}
}
}
if summary.Status == "started" {
if strings.TrimSpace(summary.ReplyText) == "" {
summary.Status = "fallback"
} else {
summary.Status = "completed"
}
}
}
func convertInterruptTraceItems(items []InterruptContextSummary) []callbacks.InterruptTraceContext {
if len(items) == 0 {
return nil
}
ret := make([]callbacks.InterruptTraceContext, 0, len(items))
for _, item := range items {
ret = append(ret, callbacks.InterruptTraceContext{
Type: item.Type,
ID: item.ID,
InfoPreview: item.InfoPreview,
})
}
return ret
}
func previewInterruptInfo(info any) string {
if info == nil {
return ""
}
switch v := info.(type) {
case string:
return preview(v, 200)
default:
data, err := json.Marshal(v)
if err != nil {
return ""
}
return preview(string(data), 200)
}
}
func summarizeInterrupts(items []*adk.InterruptCtx) []InterruptContextSummary {
if len(items) == 0 {
return nil
}
ret := make([]InterruptContextSummary, 0, len(items))
for _, item := range items {
if item == nil {
continue
}
ret = append(ret, InterruptContextSummary{
Type: extractInterruptType(item.Info),
ID: strings.TrimSpace(item.ID),
InfoPreview: previewInterruptInfo(item.Info),
})
}
return ret
}
func extractInterruptType(info any) string {
if info == nil {
return ""
}
switch v := info.(type) {
case map[string]any:
return strings.TrimSpace(getStringFromAnyMap(v, "type"))
default:
return ""
}
}
func getStringFromAnyMap(data map[string]any, key string) string {
value, ok := data[key]
if !ok || value == nil {
return ""
}
switch v := value.(type) {
case string:
return v
default:
return fmt.Sprintf("%v", v)
}
}
func appendIfMissing(items []string, value string) []string {
value = strings.TrimSpace(value)
if value == "" {
return items
}
for _, item := range items {
if strings.TrimSpace(item) == value {
return items
}
}
return append(items, value)
}
func preview(value string, limit int) string {
if limit <= 0 {
return ""
}
value = strings.TrimSpace(value)
runes := []rune(value)
if len(runes) <= limit {
return value
}
return string(runes[:limit]) + "..."
}
func buildKnowledgeContext(items []rag.RetrieveResult) string {
if len(items) == 0 {
return ""
}
var builder strings.Builder
builder.WriteString("以下是可供参考的知识库内容,请优先基于这些内容回答;如果仍不确定,请明确说明并向用户澄清。\n\n")
for i, item := range items {
if i >= 5 {
break
}
builder.WriteString("[知识片段")
builder.WriteString(fmt.Sprintf("%d", i+1))
builder.WriteString("]\n")
if strings.TrimSpace(item.DocumentTitle) != "" {
builder.WriteString("标题: ")
builder.WriteString(strings.TrimSpace(item.DocumentTitle))
builder.WriteString("\n")
}
if strings.TrimSpace(item.Content) != "" {
builder.WriteString("内容: ")
builder.WriteString(strings.TrimSpace(item.Content))
builder.WriteString("\n")
}
builder.WriteString("\n")
}
return strings.TrimSpace(builder.String())
}
@@ -0,0 +1,52 @@
package engine
import (
"cs-agent/internal/models"
einotool "github.com/cloudwego/eino/components/tool"
)
type Request struct {
Conversation *models.Conversation
UserMessage *models.Message
AIAgent *models.AIAgent
AIConfig *models.AIConfig
CheckPointID string
ExtraTools []einotool.BaseTool
ExtraToolCodes map[string]string
}
type ResumeRequest struct {
Conversation *models.Conversation
AIAgent *models.AIAgent
AIConfig *models.AIConfig
CheckPointID string
ResumeData map[string]any
ExtraTools []einotool.BaseTool
ExtraToolCodes map[string]string
}
type InterruptContextSummary struct {
Type string `json:"type,omitempty"`
ID string `json:"id"`
InfoPreview string `json:"infoPreview,omitempty"`
}
type Summary struct {
RunID string
Status string
ReplyText string
ModelName string
PromptTokens int
CompletionTokens int
HistoryMessageCount int
RetrieverCount int
ToolCallCount int
ToolCodes []string
InvokedToolCodes []string
CheckPointID string
Interrupted bool
Interrupts []InterruptContextSummary
TraceData string
ErrorMessage string
}
@@ -0,0 +1,26 @@
package adapter
import "cs-agent/internal/models"
type AIConfigSnapshot struct {
ID int64
Provider string
ModelName string
BaseURL string
MaxOutputTokens int
TimeoutMS int
}
func BuildAIConfigSnapshot(item *models.AIConfig) *AIConfigSnapshot {
if item == nil {
return nil
}
return &AIConfigSnapshot{
ID: item.ID,
Provider: string(item.Provider),
ModelName: item.ModelName,
BaseURL: item.BaseURL,
MaxOutputTokens: item.MaxOutputTokens,
TimeoutMS: item.TimeoutMS,
}
}
@@ -0,0 +1,22 @@
package adapter
import "cs-agent/internal/models"
type ConversationSnapshot struct {
ID int64
AIAgentID int64
LastMessageID int64
CurrentAssigneeID int64
}
func BuildConversationSnapshot(item *models.Conversation) *ConversationSnapshot {
if item == nil {
return nil
}
return &ConversationSnapshot{
ID: item.ID,
AIAgentID: item.AIAgentID,
LastMessageID: item.LastMessageID,
CurrentAssigneeID: item.CurrentAssigneeID,
}
}
@@ -0,0 +1,192 @@
package adapter
import (
"context"
"encoding/json"
"fmt"
"hash/crc32"
"regexp"
"strings"
"cs-agent/internal/ai/mcps"
einojsonschema "github.com/eino-contrib/jsonschema"
einotool "github.com/cloudwego/eino/components/tool"
"github.com/cloudwego/eino/schema"
)
var toolNameSanitizer = regexp.MustCompile(`[^a-zA-Z0-9_]`)
type MCPToolDefinition struct {
ToolCode string
ServerCode string
ToolName string
ModelName string
Title string
Description string
FixedArgs map[string]string
}
type MCPTool struct {
definition MCPToolDefinition
info *schema.ToolInfo
}
func NewMCPTool(definition MCPToolDefinition, metadata *mcps.ToolInfo) *MCPTool {
return &MCPTool{
definition: definition,
info: buildToolInfo(definition, metadata),
}
}
var _ einotool.InvokableTool = (*MCPTool)(nil)
func (t *MCPTool) Info(ctx context.Context) (*schema.ToolInfo, error) {
if t == nil || t.info == nil {
return nil, nil
}
return t.info, nil
}
func (t *MCPTool) InvokableRun(ctx context.Context, argumentsInJSON string, opts ...einotool.Option) (string, error) {
if t == nil {
return "", fmt.Errorf("mcp tool is nil")
}
arguments, err := parseArguments(argumentsInJSON)
if err != nil {
return "", err
}
arguments = mergeFixedArguments(arguments, t.definition.FixedArgs)
result, err := mcps.Runtime.CallTool(ctx, t.definition.ServerCode, t.definition.ToolName, arguments)
if err != nil {
return "", err
}
return buildToolResultSummary(result), nil
}
func buildToolInfo(definition MCPToolDefinition, metadata *mcps.ToolInfo) *schema.ToolInfo {
desc := strings.TrimSpace(definition.Description)
if desc == "" && metadata != nil {
desc = strings.TrimSpace(metadata.Description)
}
title := strings.TrimSpace(definition.Title)
if title == "" && metadata != nil {
title = strings.TrimSpace(metadata.Title)
}
if title != "" && desc != "" {
desc = title + "\n\n" + desc
} else if title != "" {
desc = title
}
if desc == "" {
desc = "Call MCP tool " + strings.TrimSpace(definition.ToolCode)
}
info := &schema.ToolInfo{
Name: BuildModelToolName(definition),
Desc: desc,
Extra: map[string]any{
"toolCode": definition.ToolCode,
"serverCode": definition.ServerCode,
"toolName": definition.ToolName,
},
}
if js := buildParamsSchema(metadata); js != nil {
info.ParamsOneOf = schema.NewParamsOneOfByJSONSchema(js)
}
return info
}
func buildParamsSchema(metadata *mcps.ToolInfo) *einojsonschema.Schema {
if metadata == nil || metadata.InputSchema == nil {
return genericObjectSchema()
}
raw, err := json.Marshal(metadata.InputSchema)
if err != nil || len(raw) == 0 {
return genericObjectSchema()
}
js := &einojsonschema.Schema{}
if err := json.Unmarshal(raw, js); err != nil {
return genericObjectSchema()
}
return js
}
func genericObjectSchema() *einojsonschema.Schema {
return &einojsonschema.Schema{
Version: einojsonschema.Version,
Type: "object",
AdditionalProperties: &einojsonschema.Schema{},
}
}
func parseArguments(argumentsInJSON string) (map[string]any, error) {
argumentsInJSON = strings.TrimSpace(argumentsInJSON)
if argumentsInJSON == "" {
return map[string]any{}, nil
}
args := make(map[string]any)
if err := json.Unmarshal([]byte(argumentsInJSON), &args); err != nil {
return nil, fmt.Errorf("invalid tool arguments: %w", err)
}
return args, nil
}
func mergeFixedArguments(arguments map[string]any, fixedArgs map[string]string) map[string]any {
if len(arguments) == 0 && len(fixedArgs) == 0 {
return map[string]any{}
}
ret := make(map[string]any, len(arguments)+len(fixedArgs))
for key, value := range arguments {
ret[key] = value
}
for key, value := range fixedArgs {
ret[key] = strings.TrimSpace(value)
}
return ret
}
func buildToolResultSummary(result *mcps.ToolCallResult) string {
if result == nil {
return ""
}
lines := make([]string, 0, len(result.Content)+2)
if result.IsError {
lines = append(lines, "tool returned an error")
}
if result.StructuredContent != nil {
if data, err := json.Marshal(result.StructuredContent); err == nil {
lines = append(lines, string(data))
}
}
for _, item := range result.Content {
switch item.Type {
case "text":
if text := strings.TrimSpace(item.Text); text != "" {
lines = append(lines, text)
}
default:
if item.Data == nil {
continue
}
if data, err := json.Marshal(item.Data); err == nil {
lines = append(lines, string(data))
}
}
}
return strings.TrimSpace(strings.Join(lines, "\n"))
}
func BuildModelToolName(definition MCPToolDefinition) string {
if strings.TrimSpace(definition.ModelName) != "" {
return strings.TrimSpace(definition.ModelName)
}
base := "mcp_" + strings.TrimSpace(definition.ServerCode) + "_" + strings.TrimSpace(definition.ToolName)
base = toolNameSanitizer.ReplaceAllString(base, "_")
base = strings.Trim(base, "_")
if base == "" {
base = "mcp_tool"
}
checksum := crc32.ChecksumIEEE([]byte(definition.ToolCode))
return fmt.Sprintf("%s_%08x", base, checksum)
}
@@ -0,0 +1,69 @@
package adapter
import (
"strings"
"cs-agent/internal/models"
"cs-agent/internal/pkg/enums"
"cs-agent/internal/repositories"
"github.com/cloudwego/eino/schema"
"github.com/mlogclub/simple/sqls"
)
const defaultHistoryLimit = 12
type HistoryBuildResult struct {
Messages []*schema.Message
RawItems []models.Message
}
func BuildHistoryMessages(conversationID int64, currentMessageID int64, limit int) HistoryBuildResult {
if conversationID <= 0 {
return HistoryBuildResult{}
}
if limit <= 0 {
limit = defaultHistoryLimit
}
items := repositories.MessageRepository.Find(sqls.DB(), sqls.NewCnd().
Eq("conversation_id", conversationID).
Desc("id").
Limit(limit+1))
for i, j := 0, len(items)-1; i < j; i, j = i+1, j-1 {
items[i], items[j] = items[j], items[i]
}
ret := HistoryBuildResult{
Messages: make([]*schema.Message, 0, len(items)),
RawItems: make([]models.Message, 0, len(items)),
}
for _, item := range items {
if item.ID == currentMessageID {
continue
}
msg := BuildSchemaMessage(&item)
if msg == nil {
continue
}
ret.RawItems = append(ret.RawItems, item)
ret.Messages = append(ret.Messages, msg)
}
return ret
}
func BuildSchemaMessage(item *models.Message) *schema.Message {
if item == nil {
return nil
}
content := strings.TrimSpace(item.Content)
if content == "" {
return nil
}
switch item.SenderType {
case enums.IMSenderTypeCustomer:
return schema.UserMessage(content)
case enums.IMSenderTypeAI, enums.IMSenderTypeAgent:
return schema.AssistantMessage(content, nil)
default:
return nil
}
}
@@ -0,0 +1,55 @@
package agents
import (
"context"
"fmt"
"github.com/cloudwego/eino/adk"
)
type CustomerServiceAgent struct {
Inner adk.Agent
}
var _ adk.ResumableAgent = (*CustomerServiceAgent)(nil)
func (a *CustomerServiceAgent) Name(ctx context.Context) string {
if a == nil || a.Inner == nil {
return "customer_service_agent"
}
return a.Inner.Name(ctx)
}
func (a *CustomerServiceAgent) Description(ctx context.Context) string {
if a == nil || a.Inner == nil {
return "customer service chat agent"
}
return a.Inner.Description(ctx)
}
func (a *CustomerServiceAgent) Run(ctx context.Context, input *adk.AgentInput, options ...adk.AgentRunOption) *adk.AsyncIterator[*adk.AgentEvent] {
if a == nil || a.Inner == nil {
iter, gen := adk.NewAsyncIteratorPair[*adk.AgentEvent]()
gen.Send(&adk.AgentEvent{Err: context.Canceled})
gen.Close()
return iter
}
return a.Inner.Run(ctx, input, options...)
}
func (a *CustomerServiceAgent) Resume(ctx context.Context, info *adk.ResumeInfo, options ...adk.AgentRunOption) *adk.AsyncIterator[*adk.AgentEvent] {
if a == nil || a.Inner == nil {
iter, gen := adk.NewAsyncIteratorPair[*adk.AgentEvent]()
gen.Send(&adk.AgentEvent{Err: fmt.Errorf("customer service agent is not initialized")})
gen.Close()
return iter
}
ra, ok := a.Inner.(adk.ResumableAgent)
if !ok {
iter, gen := adk.NewAsyncIteratorPair[*adk.AgentEvent]()
gen.Send(&adk.AgentEvent{Err: fmt.Errorf("inner agent %q does not implement resumable agent", a.Inner.Name(ctx))})
gen.Close()
return iter
}
return ra.Resume(ctx, info, options...)
}
@@ -0,0 +1,85 @@
package callbacks
import (
"context"
"encoding/json"
"strings"
"time"
einotool "github.com/cloudwego/eino/components/tool"
"github.com/cloudwego/eino/adk"
)
type ToolMetadata struct {
ToolCode string
ServerCode string
ToolName string
}
type RuntimeTraceHandler struct {
*adk.BaseChatModelAgentMiddleware
collector *RuntimeTraceCollector
toolMetadataBy map[string]ToolMetadata
}
func NewRuntimeTraceHandler(collector *RuntimeTraceCollector, toolMetadataBy map[string]ToolMetadata) *RuntimeTraceHandler {
return &RuntimeTraceHandler{
BaseChatModelAgentMiddleware: &adk.BaseChatModelAgentMiddleware{},
collector: collector,
toolMetadataBy: toolMetadataBy,
}
}
func (h *RuntimeTraceHandler) WrapInvokableToolCall(_ context.Context, endpoint adk.InvokableToolCallEndpoint, tCtx *adk.ToolContext) (adk.InvokableToolCallEndpoint, error) {
return func(ctx context.Context, argumentsInJSON string, opts ...einotool.Option) (string, error) {
startedAt := time.Now()
result, err := endpoint(ctx, argumentsInJSON, opts...)
item := ToolTraceItem{
ResultPreview: previewToolText(result, 300),
LatencyMs: time.Since(startedAt).Milliseconds(),
Status: "ok",
}
if tCtx != nil {
item.ToolName = strings.TrimSpace(tCtx.Name)
if metadata, ok := h.toolMetadataBy[item.ToolName]; ok {
item.ToolCode = metadata.ToolCode
item.ServerCode = metadata.ServerCode
item.ToolName = metadata.ToolName
}
}
if arguments := parseToolArguments(argumentsInJSON); len(arguments) > 0 {
item.Arguments = arguments
}
if err != nil {
item.Status = "error"
item.ErrorMessage = err.Error()
}
h.collector.AddToolItem(item)
return result, err
}, nil
}
func parseToolArguments(argumentsInJSON string) map[string]any {
argumentsInJSON = strings.TrimSpace(argumentsInJSON)
if argumentsInJSON == "" {
return nil
}
ret := make(map[string]any)
if err := json.Unmarshal([]byte(argumentsInJSON), &ret); err != nil {
return nil
}
return ret
}
func previewToolText(text string, limit int) string {
if limit <= 0 {
return ""
}
text = strings.TrimSpace(text)
runes := []rune(text)
if len(runes) <= limit {
return text
}
return string(runes[:limit]) + "..."
}
@@ -0,0 +1,41 @@
package callbacks
import (
"encoding/json"
"sync"
)
type RuntimeTraceCollector struct {
mu sync.Mutex
Data RuntimeTraceData
}
func NewRuntimeTraceCollector() *RuntimeTraceCollector {
ret := &RuntimeTraceCollector{}
ret.Data.Version = "v1"
ret.Data.Status = "started"
return ret
}
func (c *RuntimeTraceCollector) Marshal() string {
if c == nil {
return ""
}
c.mu.Lock()
defer c.mu.Unlock()
buf, err := json.Marshal(c.Data)
if err != nil {
return ""
}
return string(buf)
}
func (c *RuntimeTraceCollector) AddToolItem(item ToolTraceItem) {
if c == nil {
return
}
c.mu.Lock()
defer c.mu.Unlock()
c.Data.Tools.Count++
c.Data.Tools.Items = append(c.Data.Tools.Items, item)
}
@@ -0,0 +1,63 @@
package callbacks
type ToolTraceItem struct {
ToolCode string `json:"toolCode"`
ServerCode string `json:"serverCode"`
ToolName string `json:"toolName"`
Arguments map[string]any `json:"arguments,omitempty"`
ResultPreview string `json:"resultPreview,omitempty"`
LatencyMs int64 `json:"latencyMs,omitempty"`
Status string `json:"status,omitempty"`
ErrorMessage string `json:"errorMessage,omitempty"`
}
type RetrieverTraceItem struct {
Query string `json:"query,omitempty"`
KnowledgeBaseID int64 `json:"knowledgeBaseId,omitempty"`
DocumentID int64 `json:"documentId,omitempty"`
DocumentTitle string `json:"documentTitle,omitempty"`
Score float64 `json:"score,omitempty"`
LatencyMs int64 `json:"latencyMs,omitempty"`
}
type RuntimeTraceData struct {
Version string `json:"version"`
Status string `json:"status"`
RunID string `json:"runId,omitempty"`
Interrupt struct {
CheckPointID string `json:"checkPointId,omitempty"`
Items []InterruptTraceContext `json:"items,omitempty"`
} `json:"interrupt"`
Model struct {
Provider string `json:"provider,omitempty"`
Name string `json:"name,omitempty"`
} `json:"model"`
Input struct {
HistoryMessageCount int `json:"historyMessageCount,omitempty"`
KnowledgeBaseIDs []int64 `json:"knowledgeBaseIds,omitempty"`
ToolCodes []string `json:"toolCodes,omitempty"`
CurrentUserMessagePreview string `json:"currentUserMessagePreview,omitempty"`
} `json:"input"`
Retriever struct {
Count int `json:"count,omitempty"`
Items []RetrieverTraceItem `json:"items,omitempty"`
} `json:"retriever"`
Tools struct {
Count int `json:"count,omitempty"`
Items []ToolTraceItem `json:"items,omitempty"`
} `json:"tools"`
Output struct {
ReplyText string `json:"replyText,omitempty"`
FinishReason string `json:"finishReason,omitempty"`
} `json:"output"`
Error struct {
Message string `json:"message,omitempty"`
Stage string `json:"stage,omitempty"`
} `json:"error"`
}
type InterruptTraceContext struct {
Type string `json:"type,omitempty"`
ID string `json:"id"`
InfoPreview string `json:"infoPreview,omitempty"`
}
@@ -0,0 +1,112 @@
package factory
import (
"context"
"strings"
einoadapter "cs-agent/internal/ai/runtime/internal/impl/adapter"
einoagents "cs-agent/internal/ai/runtime/internal/impl/agents"
einocallbacks "cs-agent/internal/ai/runtime/internal/impl/callbacks"
"cs-agent/internal/models"
"github.com/cloudwego/eino/adk"
einobasetool "github.com/cloudwego/eino/components/tool"
"github.com/cloudwego/eino/compose"
)
type AgentFactory struct {
chatModelFactory *ChatModelFactory
toolFactory *ToolFactory
}
func NewAgentFactory() *AgentFactory {
return &AgentFactory{
chatModelFactory: NewChatModelFactory(),
toolFactory: NewToolFactory(),
}
}
func (f *AgentFactory) BuildCustomerServiceAgent(ctx context.Context, aiAgent *models.AIAgent, aiConfig *models.AIConfig,
toolDefinitions []einoadapter.MCPToolDefinition, extraTools []einobasetool.BaseTool, extraToolCodes map[string]string,
collector *einocallbacks.RuntimeTraceCollector) (*einoagents.CustomerServiceAgent, error) {
if aiAgent == nil || aiConfig == nil {
return nil, nil
}
chatModel, err := f.chatModelFactory.Build(ctx, aiConfig)
if err != nil {
return nil, err
}
baseTools, err := f.toolFactory.BuildBaseToolsByDefinitions(ctx, toolDefinitions)
if err != nil {
return nil, err
}
allTools := make([]einobasetool.BaseTool, 0, len(baseTools)+len(extraTools))
allTools = append(allTools, extraTools...)
allTools = append(allTools, baseTools...)
handlers := make([]adk.ChatModelAgentMiddleware, 0, 1)
if collector != nil {
toolMetadataBy := make(map[string]einocallbacks.ToolMetadata, len(toolDefinitions))
for _, item := range toolDefinitions {
toolMetadataBy[item.ModelName] = einocallbacks.ToolMetadata{
ToolCode: item.ToolCode,
ServerCode: item.ServerCode,
ToolName: item.ToolName,
}
}
handlers = append(handlers, einocallbacks.NewRuntimeTraceHandler(collector, toolMetadataBy))
}
inner, err := adk.NewChatModelAgent(ctx, &adk.ChatModelAgentConfig{
Name: strings.TrimSpace(aiAgent.Name),
Description: strings.TrimSpace(aiAgent.Description),
Instruction: buildAgentInstruction(aiAgent, extraToolCodes),
Model: chatModel,
ToolsConfig: adk.ToolsConfig{
ToolsNodeConfig: compose.ToolsNodeConfig{
Tools: allTools,
},
},
Handlers: handlers,
})
if err != nil {
return nil, err
}
return &einoagents.CustomerServiceAgent{Inner: inner}, nil
}
func buildAgentInstruction(aiAgent *models.AIAgent, extraToolCodes map[string]string) string {
baseInstruction := ""
if aiAgent != nil {
baseInstruction = strings.TrimSpace(aiAgent.SystemPrompt)
}
appendixParts := make([]string, 0, 1)
if hasToolCode(extraToolCodes, "builtin/create_ticket_with_confirmation") {
appendixParts = append(appendixParts, strings.TrimSpace(`
你可以在确认信息充分后调用 create_ticket_with_confirmation 工具来创建工单,但必须遵守以下规则:
1. 只有在用户明确表达希望提交工单、投诉、报障、售后处理等诉求时,才考虑调用该工具。
2. 调用前你必须已经整理出清晰的工单标题和问题描述;如果信息不足,先继续追问,不要过早调用。
3. 一旦准备创建工单,必须调用 create_ticket_with_confirmation 工具,禁止直接口头宣称“已经创建工单”。
4. 该工具会先向用户发起确认。用户确认后才会真正创建工单;用户取消则结束本次建单流程。
5. 如果用户只是咨询、抱怨或泛泛表达不满,但没有明确要求建单,优先继续澄清,不要主动创建工单。
`))
}
if len(appendixParts) == 0 {
return baseInstruction
}
if baseInstruction == "" {
return strings.Join(appendixParts, "\n\n")
}
return baseInstruction + "\n\n" + strings.Join(appendixParts, "\n\n")
}
func hasToolCode(toolCodes map[string]string, target string) bool {
target = strings.TrimSpace(target)
if target == "" {
return false
}
for _, toolCode := range toolCodes {
if strings.TrimSpace(toolCode) == target {
return true
}
}
return false
}
@@ -0,0 +1,64 @@
package factory
import (
"context"
"strings"
"time"
"cs-agent/internal/models"
"cs-agent/internal/pkg/enums"
openai "github.com/cloudwego/eino-ext/components/model/openai"
"github.com/cloudwego/eino/components/model"
)
type ChatModelFactory struct{}
func NewChatModelFactory() *ChatModelFactory {
return &ChatModelFactory{}
}
func (f *ChatModelFactory) Build(ctx context.Context, item *models.AIConfig) (model.ToolCallingChatModel, error) {
if item == nil {
return nil, nil
}
conf := &openai.ChatModelConfig{
APIKey: strings.TrimSpace(item.APIKey),
BaseURL: strings.TrimSpace(item.BaseURL),
Model: strings.TrimSpace(item.ModelName),
}
if item.TimeoutMS > 0 {
conf.Timeout = time.Duration(item.TimeoutMS) * time.Millisecond
}
if item.MaxOutputTokens > 0 {
maxCompletionTokens := item.MaxOutputTokens
conf.MaxCompletionTokens = &maxCompletionTokens
}
if item.Provider == enums.AIProviderOpenAI && isAzureOpenAIBaseURL(item.BaseURL) {
conf.ByAzure = true
conf.APIVersion = "2024-06-01"
}
if extraFields := providerExtraFields(item); len(extraFields) > 0 {
conf.ExtraFields = extraFields
}
return openai.NewChatModel(ctx, conf)
}
func isAzureOpenAIBaseURL(baseURL string) bool {
baseURL = strings.ToLower(strings.TrimSpace(baseURL))
return strings.Contains(baseURL, ".openai.azure.com")
}
func providerExtraFields(item *models.AIConfig) map[string]any {
if item == nil {
return nil
}
baseURL := strings.ToLower(strings.TrimSpace(item.BaseURL))
modelName := strings.ToLower(strings.TrimSpace(item.ModelName))
if strings.Contains(baseURL, "dashscope.aliyuncs.com") && strings.HasPrefix(modelName, "qwen3") {
return map[string]any{
"enable_thinking": false,
}
}
return nil
}
@@ -0,0 +1,27 @@
package factory
import (
"context"
einostore "cs-agent/internal/ai/runtime/internal/impl/store"
"github.com/cloudwego/eino/adk"
)
type RunnerFactory struct{}
func NewRunnerFactory() *RunnerFactory {
return &RunnerFactory{}
}
func (f *RunnerFactory) Build(ctx context.Context, agent adk.Agent, enableStreaming bool, enableCheckpoint bool) *adk.Runner {
var checkpointStore adk.CheckPointStore
if enableCheckpoint {
checkpointStore = einostore.DefaultCheckPointStore
}
return adk.NewRunner(ctx, adk.RunnerConfig{
Agent: agent,
EnableStreaming: enableStreaming,
CheckPointStore: checkpointStore,
})
}
@@ -0,0 +1,100 @@
package factory
import (
"context"
"encoding/json"
"strings"
"cs-agent/internal/ai/mcps"
impladapter "cs-agent/internal/ai/runtime/internal/impl/adapter"
"cs-agent/internal/models"
"cs-agent/internal/pkg/dto/request"
einotool "github.com/cloudwego/eino/components/tool"
)
type ToolFactory struct{}
func NewToolFactory() *ToolFactory {
return &ToolFactory{}
}
func (f *ToolFactory) BuildMCPTools(aiAgent *models.AIAgent) ([]impladapter.MCPToolDefinition, error) {
if aiAgent == nil || strings.TrimSpace(aiAgent.AllowedMCPTools) == "" {
return nil, nil
}
var raw []request.AIAgentMCPToolRequest
if err := json.Unmarshal([]byte(aiAgent.AllowedMCPTools), &raw); err != nil {
return nil, err
}
ret := make([]impladapter.MCPToolDefinition, 0, len(raw))
for _, item := range raw {
toolCode := strings.TrimSpace(item.ServerCode) + "/" + strings.TrimSpace(item.ToolName)
definition := impladapter.MCPToolDefinition{
ToolCode: toolCode,
ServerCode: strings.TrimSpace(item.ServerCode),
ToolName: strings.TrimSpace(item.ToolName),
Title: strings.TrimSpace(item.Title),
Description: strings.TrimSpace(item.Description),
FixedArgs: cloneStringMap(item.Arguments),
}
definition.ModelName = impladapter.BuildModelToolName(definition)
ret = append(ret, definition)
}
return ret, nil
}
func (f *ToolFactory) BuildBaseTools(ctx context.Context, aiAgent *models.AIAgent) ([]einotool.BaseTool, error) {
definitions, err := f.BuildMCPTools(aiAgent)
if err != nil {
return nil, err
}
return f.BuildBaseToolsByDefinitions(ctx, definitions)
}
func (f *ToolFactory) BuildBaseToolsByDefinitions(ctx context.Context, definitions []impladapter.MCPToolDefinition) ([]einotool.BaseTool, error) {
if len(definitions) == 0 {
return nil, nil
}
metadataByCode, err := f.loadToolMetadata(ctx, definitions)
if err != nil {
return nil, err
}
ret := make([]einotool.BaseTool, 0, len(definitions))
for _, item := range definitions {
ret = append(ret, impladapter.NewMCPTool(item, metadataByCode[item.ToolCode]))
}
return ret, nil
}
func (f *ToolFactory) loadToolMetadata(ctx context.Context, definitions []impladapter.MCPToolDefinition) (map[string]*mcps.ToolInfo, error) {
toolsByCode := make(map[string]*mcps.ToolInfo, len(definitions))
serverCodes := make(map[string]struct{})
for _, item := range definitions {
serverCodes[item.ServerCode] = struct{}{}
}
for serverCode := range serverCodes {
toolInfos, err := mcps.Runtime.ListTools(ctx, serverCode)
if err != nil {
return nil, err
}
for i := range toolInfos {
toolInfo := toolInfos[i]
toolCode := strings.TrimSpace(serverCode) + "/" + strings.TrimSpace(toolInfo.Name)
toolInfoCopy := toolInfo
toolsByCode[toolCode] = &toolInfoCopy
}
}
return toolsByCode, nil
}
func cloneStringMap(input map[string]string) map[string]string {
if len(input) == 0 {
return nil
}
ret := make(map[string]string, len(input))
for key, value := range input {
ret[key] = value
}
return ret
}
@@ -0,0 +1,32 @@
package retrievers
import (
"context"
"cs-agent/internal/ai/rag"
"cs-agent/internal/models"
"cs-agent/internal/pkg/utils"
)
type KnowledgeRetriever struct {
AIAgent *models.AIAgent
}
func NewKnowledgeRetriever(aiAgent *models.AIAgent) *KnowledgeRetriever {
return &KnowledgeRetriever{AIAgent: aiAgent}
}
func (r *KnowledgeRetriever) KnowledgeBaseIDs() []int64 {
if r == nil || r.AIAgent == nil {
return nil
}
return utils.SplitInt64s(r.AIAgent.KnowledgeIDs)
}
func (r *KnowledgeRetriever) Retrieve(ctx context.Context, query string) ([]rag.RetrieveResult, *rag.RetrieveTrace, error) {
ids := r.KnowledgeBaseIDs()
return rag.Retrieve.RetrieveWithTrace(ctx, rag.RetrieveRequest{
Query: query,
KnowledgeBaseIDs: ids,
})
}
@@ -0,0 +1,69 @@
package store
import (
"context"
"encoding/base64"
"time"
"cs-agent/internal/models"
"cs-agent/internal/repositories"
"github.com/cloudwego/eino/adk"
"github.com/mlogclub/simple/sqls"
)
var DefaultCheckPointStore adk.CheckPointStore = NewDBCheckPointStore()
type DBCheckPointStore struct{}
func NewDBCheckPointStore() *DBCheckPointStore {
return &DBCheckPointStore{}
}
func (s *DBCheckPointStore) Get(_ context.Context, checkPointID string) ([]byte, bool, error) {
item := repositories.ConversationInterruptRepository.GetByCheckPointID(sqls.DB(), checkPointID)
if item == nil || item.CheckPointData == "" {
return nil, false, nil
}
return decodeCheckPointData(item.CheckPointData)
}
func (s *DBCheckPointStore) Set(_ context.Context, checkPointID string, checkPoint []byte) error {
item := repositories.ConversationInterruptRepository.GetByCheckPointID(sqls.DB(), checkPointID)
if item == nil {
item = buildEmptyInterrupt(checkPointID)
}
item.CheckPointData = encodeCheckPointData(checkPoint)
if item.ConversationID == 0 && item.AIAgentID == 0 && item.SourceMessageID == 0 && item.Status == "" {
return repositories.ConversationInterruptRepository.Create(sqls.DB(), item)
}
return repositories.ConversationInterruptRepository.UpsertByCheckPointID(sqls.DB(), item)
}
func encodeCheckPointData(data []byte) string {
if len(data) == 0 {
return ""
}
return base64.StdEncoding.EncodeToString(data)
}
func decodeCheckPointData(value string) ([]byte, bool, error) {
if value == "" {
return nil, false, nil
}
data, err := base64.StdEncoding.DecodeString(value)
if err != nil {
return nil, false, err
}
return data, true, nil
}
func buildEmptyInterrupt(checkPointID string) *models.ConversationInterrupt {
now := time.Now()
return &models.ConversationInterrupt{
CheckPointID: checkPointID,
Status: "checkpointed",
CreatedAt: now,
UpdatedAt: now,
}
}
+44
View File
@@ -0,0 +1,44 @@
package registry
import (
"strings"
einotool "github.com/cloudwego/eino/components/tool"
)
type Registry struct {
tools []Tool
}
func NewRegistry(tools ...Tool) *Registry {
return &Registry{
tools: tools,
}
}
func (r *Registry) Resolve(ctx Context) (*ToolSet, error) {
ret := &ToolSet{
Tools: make([]einotool.BaseTool, 0, len(r.tools)),
ToolCodes: make(map[string]string),
}
for _, toolDef := range r.tools {
if toolDef == nil || !toolDef.Enabled(ctx) {
continue
}
tool, err := toolDef.Build(ctx)
if err != nil {
return nil, err
}
if tool == nil {
continue
}
toolName := strings.TrimSpace(toolDef.Name())
toolCode := strings.TrimSpace(toolDef.Code())
if toolName == "" || toolCode == "" {
continue
}
ret.Tools = append(ret.Tools, tool)
ret.ToolCodes[toolName] = toolCode
}
return ret, nil
}
+26
View File
@@ -0,0 +1,26 @@
package registry
import (
"cs-agent/internal/models"
einotool "github.com/cloudwego/eino/components/tool"
)
type Context struct {
Conversation *models.Conversation
AIAgent *models.AIAgent
AIConfig *models.AIConfig
UserMessage *models.Message
}
type ToolSet struct {
Tools []einotool.BaseTool
ToolCodes map[string]string
}
type Tool interface {
Name() string
Code() string
Enabled(ctx Context) bool
Build(ctx Context) (einotool.BaseTool, error)
}
+563
View File
@@ -0,0 +1,563 @@
package runtime
import (
"context"
"encoding/json"
"fmt"
"log/slog"
"strings"
"time"
"cs-agent/internal/models"
"cs-agent/internal/pkg/dto"
"cs-agent/internal/pkg/enums"
"cs-agent/internal/repositories"
svc "cs-agent/internal/services"
"github.com/mlogclub/simple/common/strs"
"github.com/mlogclub/simple/sqls"
)
var AIReplyService = newAIReplyService()
func init() {
svc.TriggerAIReplyAsyncHook = AIReplyService.TriggerReplyAsync
}
func newAIReplyService() *aiReplyService {
return &aiReplyService{}
}
type aiReplyService struct{}
type aiReplyTraceData struct {
Status string `json:"status"`
RuntimeLatencyMs int64 `json:"runtimeLatencyMs,omitempty"`
RecheckMs int64 `json:"recheckMs,omitempty"`
CommitMs int64 `json:"commitMs,omitempty"`
FinalAction string `json:"finalAction,omitempty"`
ReplySent bool `json:"replySent,omitempty"`
ReplyMessageID int64 `json:"replyMessageId,omitempty"`
Runtime json.RawMessage `json:"runtime,omitempty"`
}
const (
defaultAIReplyAsyncTimeoutSeconds = 180
maxAIReplyAsyncTimeoutSeconds = 600
)
func (s *aiReplyService) resolveReplyTimeout(aiAgent models.AIAgent) time.Duration {
if aiAgent.ReplyTimeoutSeconds <= 0 {
return time.Duration(defaultAIReplyAsyncTimeoutSeconds) * time.Second
}
if aiAgent.ReplyTimeoutSeconds > maxAIReplyAsyncTimeoutSeconds {
return time.Duration(maxAIReplyAsyncTimeoutSeconds) * time.Second
}
return time.Duration(aiAgent.ReplyTimeoutSeconds) * time.Second
}
func (s *aiReplyService) TriggerReplyAsync(conversation models.Conversation, message models.Message) {
go func() {
aiAgent := svc.AIAgentService.Get(conversation.AIAgentID)
if aiAgent == nil || aiAgent.Status != enums.StatusOk {
return
}
startedAt := time.Now()
timeout := s.resolveReplyTimeout(*aiAgent)
ctx, cancel := context.WithTimeout(context.Background(), timeout)
defer cancel()
if err := s.TriggerReply(ctx, conversation, message, *aiAgent); err != nil {
slog.Error("failed to trigger ai reply",
"message_id", message.ID,
"timeout_ms", timeout.Milliseconds(),
"elapsed_ms", time.Since(startedAt).Milliseconds(),
"error", err)
}
}()
}
func (s *aiReplyService) TriggerReply(ctx context.Context, conversation models.Conversation, message models.Message, aiAgent models.AIAgent) (retErr error) {
startedAt := time.Now()
trace := &aiReplyTraceData{Status: "started"}
var summary *Summary
if err := ctx.Err(); err != nil {
return err
}
if message.SenderType != enums.IMSenderTypeCustomer {
return nil
}
if conversation.HandoffAt != nil || conversation.CurrentAssigneeID > 0 {
return nil
}
if aiAgent.ServiceMode == enums.IMConversationServiceModeHumanOnly {
return nil
}
if strs.IsBlank(message.Content) {
return nil
}
defer func() {
s.writeRunLog(startedAt, message, conversation, aiAgent, message.Content, retErr, trace, summary)
}()
if pendingInterrupt := svc.ConversationInterruptService.FindLatestPendingByConversationID(conversation.ID); pendingInterrupt != nil {
return s.resumePendingInterrupt(ctx, conversation, message, aiAgent, pendingInterrupt, trace, &summary)
}
if s.shouldHandoffByQuestion(message.Content, aiAgent) {
return s.handoffConversation(conversation, aiAgent, "用户主动要求人工")
}
if aiAgent.ServiceMode != enums.IMConversationServiceModeAIOnly &&
aiAgent.MaxAIReplyRounds > 0 &&
conversation.AIReplyRounds >= aiAgent.MaxAIReplyRounds {
return s.handoffConversation(conversation, aiAgent, "达到AI最大回复轮次")
}
aiConfig := svc.AIConfigService.Get(aiAgent.AIConfigID)
if aiConfig == nil {
return fmt.Errorf("ai config is nil")
}
runtimeStartedAt := time.Now()
var err error
summary, err = Service.Run(ctx, Request{
Conversation: &conversation,
UserMessage: &message,
AIAgent: &aiAgent,
AIConfig: aiConfig,
})
trace.RuntimeLatencyMs = time.Since(runtimeStartedAt).Milliseconds()
if err != nil {
trace.Status = "runtime_error"
trace.FinalAction = "error"
if summary != nil {
trace.Runtime = json.RawMessage(summary.TraceData)
}
return err
}
trace.Status = "runtime_prepared"
trace.FinalAction = toRunLogFinalAction(summary)
if summary != nil && strings.TrimSpace(summary.TraceData) != "" {
trace.Runtime = json.RawMessage(summary.TraceData)
}
if summary != nil && summary.Interrupted {
return s.handleInterruptedSummary(conversation, message, aiAgent, summary, trace)
}
if summary != nil && strings.TrimSpace(summary.ReplyText) != "" {
replyMessage, err := s.sendAIReply(conversation, message, aiAgent, summary.ReplyText, trace, "ai_reply")
if err != nil {
return err
}
if err := s.incrementAIReplyRounds(conversation.ID, conversation.AIReplyRounds+1, aiAgent.Name); err != nil {
return err
}
trace.ReplySent = replyMessage != nil
}
return nil
}
func (s *aiReplyService) resumePendingInterrupt(ctx context.Context, conversation models.Conversation, message models.Message, aiAgent models.AIAgent,
pendingInterrupt *models.ConversationInterrupt, trace *aiReplyTraceData, summaryRef **Summary) error {
if pendingInterrupt == nil {
return nil
}
aiConfig := svc.AIConfigService.Get(aiAgent.AIConfigID)
if aiConfig == nil {
return fmt.Errorf("ai config is nil")
}
runtimeStartedAt := time.Now()
summary, err := Service.Resume(ctx, ResumeRequest{
Conversation: &conversation,
AIAgent: &aiAgent,
AIConfig: aiConfig,
CheckPointID: strings.TrimSpace(pendingInterrupt.CheckPointID),
ResumeData: map[string]any{
strings.TrimSpace(pendingInterrupt.InterruptID): strings.TrimSpace(message.Content),
},
})
trace.RuntimeLatencyMs = time.Since(runtimeStartedAt).Milliseconds()
*summaryRef = summary
if err != nil {
if isCheckpointMissingError(err) {
summary = &Summary{
Status: "expired",
ReplyText: "本次确认已失效,请重新发起。",
}
*summaryRef = summary
trace.Status = "interrupt_expired"
trace.FinalAction = "expired"
replyMessage, expireErr := s.sendAIReply(conversation, message, aiAgent, summary.ReplyText, trace, "ai_interrupt_expired")
if expireErr != nil {
return expireErr
}
if err := s.incrementAIReplyRounds(conversation.ID, conversation.AIReplyRounds+1, aiAgent.Name); err != nil {
return err
}
lastResumeMessageID := int64(0)
if replyMessage != nil {
lastResumeMessageID = replyMessage.ID
}
if expireMarkErr := svc.ConversationInterruptService.MarkExpired(pendingInterrupt.ID, lastResumeMessageID); expireMarkErr != nil {
return expireMarkErr
}
return nil
}
trace.Status = "runtime_error"
trace.FinalAction = "error"
if summary != nil {
trace.Runtime = json.RawMessage(summary.TraceData)
}
return err
}
trace.Status = "runtime_prepared"
trace.FinalAction = toRunLogFinalAction(summary)
if summary != nil && strings.TrimSpace(summary.TraceData) != "" {
trace.Runtime = json.RawMessage(summary.TraceData)
}
if summary != nil && summary.Interrupted {
return s.handleInterruptedResume(conversation, message, aiAgent, pendingInterrupt, summary, trace)
}
if summary != nil && strings.TrimSpace(summary.ReplyText) != "" {
replyMessage, err := s.sendAIReply(conversation, message, aiAgent, summary.ReplyText, trace, "ai_resume")
if err != nil {
return err
}
if err := s.incrementAIReplyRounds(conversation.ID, conversation.AIReplyRounds+1, aiAgent.Name); err != nil {
return err
}
replyMessageID := int64(0)
if replyMessage != nil {
replyMessageID = replyMessage.ID
}
if isCancellationReply(summary.ReplyText) {
return svc.ConversationInterruptService.MarkCancelled(pendingInterrupt.ID, replyMessageID)
}
return svc.ConversationInterruptService.MarkResolved(pendingInterrupt.ID, replyMessageID)
}
return svc.ConversationInterruptService.MarkResolved(pendingInterrupt.ID, 0)
}
func (s *aiReplyService) handleInterruptedSummary(conversation models.Conversation, message models.Message, aiAgent models.AIAgent,
summary *Summary, trace *aiReplyTraceData) error {
pending := buildConversationInterrupt(conversation, message, aiAgent, summary)
if err := svc.ConversationInterruptService.CreateOrUpdatePending(pending); err != nil {
return err
}
pending = svc.ConversationInterruptService.GetByCheckPointID(summary.CheckPointID)
replyText := resolveInterruptPrompt(summary)
replyMessage, err := s.sendAIReply(conversation, message, aiAgent, replyText, trace, "ai_interrupt")
if err != nil {
return err
}
if err := s.incrementAIReplyRounds(conversation.ID, conversation.AIReplyRounds+1, aiAgent.Name); err != nil {
return err
}
if replyMessage != nil && pending != nil {
return svc.ConversationInterruptService.MarkPendingAgain(pending.ID, pending.InterruptID, replyText, replyMessage.ID)
}
return nil
}
func (s *aiReplyService) handleInterruptedResume(conversation models.Conversation, message models.Message, aiAgent models.AIAgent,
pendingInterrupt *models.ConversationInterrupt, summary *Summary, trace *aiReplyTraceData) error {
if pendingInterrupt == nil {
return nil
}
replyText := resolveInterruptPrompt(summary)
replyMessage, err := s.sendAIReply(conversation, message, aiAgent, replyText, trace, "ai_interrupt_resume")
if err != nil {
return err
}
if err := s.incrementAIReplyRounds(conversation.ID, conversation.AIReplyRounds+1, aiAgent.Name); err != nil {
return err
}
if replyMessage != nil {
return svc.ConversationInterruptService.MarkPendingAgain(pendingInterrupt.ID, firstInterruptID(summary), replyText, replyMessage.ID)
}
return nil
}
func (s *aiReplyService) sendAIReply(conversation models.Conversation, message models.Message, aiAgent models.AIAgent,
replyText string, trace *aiReplyTraceData, clientPrefix string) (*models.Message, error) {
replyText = strings.TrimSpace(replyText)
if replyText == "" {
return nil, nil
}
commitStartedAt := time.Now()
replyMessage, err := svc.MessageService.SendAIMessage(conversation.ID, aiAgent.ID,
fmt.Sprintf("%s_%d", strings.TrimSpace(clientPrefix), message.ID), enums.IMMessageTypeText, replyText, "", s.buildAIPrincipal(aiAgent))
if trace != nil {
trace.CommitMs = time.Since(commitStartedAt).Milliseconds()
trace.ReplySent = err == nil && replyMessage != nil
if replyMessage != nil {
trace.ReplyMessageID = replyMessage.ID
}
}
return replyMessage, err
}
func (s *aiReplyService) shouldHandoffByQuestion(question string, aiAgent models.AIAgent) bool {
if aiAgent.ServiceMode == enums.IMConversationServiceModeAIOnly {
return false
}
normalized := strings.ReplaceAll(strings.ToLower(strings.TrimSpace(question)), " ", "")
if normalized == "" {
return false
}
keywords := []string{"转人工", "人工客服"}
for _, keyword := range keywords {
if strings.Contains(normalized, keyword) {
return true
}
}
return false
}
func (s *aiReplyService) writeRunLog(startedAt time.Time, message models.Message, conversation models.Conversation, aiAgent models.AIAgent,
question string, runErr error, trace *aiReplyTraceData, summary *Summary) {
errorMessage := ""
if runErr != nil {
errorMessage = runErr.Error()
} else if summary != nil && strings.TrimSpace(summary.ErrorMessage) != "" {
errorMessage = strings.TrimSpace(summary.ErrorMessage)
}
traceData := buildAIReplyTraceData(trace)
plannedAction, plannedToolCode, planReason := buildRunLogPlan(summary)
logItem := &models.AgentRunLog{
ConversationID: conversation.ID,
MessageID: message.ID,
AIAgentID: aiAgent.ID,
AIConfigID: aiAgent.AIConfigID,
UserMessage: strings.TrimSpace(question),
PlannedAction: plannedAction,
PlannedSkillCode: strings.TrimSpace(summaryPlannedSkillCode(summary)),
PlannedToolCode: plannedToolCode,
PlanReason: planReason,
FinalAction: toRunLogFinalAction(summary),
ReplyText: buildRunLogReplyText(summary),
ErrorMessage: errorMessage,
LatencyMs: time.Since(startedAt).Milliseconds(),
TraceData: traceData,
CreatedAt: time.Now(),
}
if err := svc.AgentRunLogService.Create(logItem); err != nil {
slog.Warn("create agent run log failed",
"message_id", message.ID,
"conversation_id", logItem.ConversationID,
"ai_agent_id", aiAgent.ID,
"error", err)
}
}
func buildAIReplyTraceData(trace *aiReplyTraceData) string {
if trace == nil {
return ""
}
data, err := json.Marshal(trace)
if err != nil {
return ""
}
return string(data)
}
func buildRunLogPlan(summary *Summary) (plannedAction, plannedToolCode, planReason string) {
if summary == nil {
return "", "", ""
}
if skillCode := strings.TrimSpace(summaryPlannedSkillCode(summary)); skillCode != "" {
reason := strings.TrimSpace(summary.PlanReason)
if reason == "" {
reason = "skill_selected"
}
return "skill", "", reason
}
if strings.TrimSpace(summary.Status) == "expired" {
return "interrupt", "", "pending interrupt checkpoint expired"
}
if summary.Interrupted {
return "tool", firstInvokedToolCode(summary), "agent interrupted and is waiting for user confirmation"
}
if len(summary.InvokedToolCodes) > 0 {
return "tool", strings.TrimSpace(summary.InvokedToolCodes[0]), "agent invoked MCP tool"
}
if strings.TrimSpace(summary.ReplyText) != "" {
return "reply", "", "agent replied directly"
}
if strings.TrimSpace(summary.ErrorMessage) != "" {
return "error", "", "runtime execution failed"
}
return "fallback", "", "runtime produced empty reply"
}
func toRunLogFinalAction(summary *Summary) string {
if summary == nil {
return ""
}
if skillCode := strings.TrimSpace(summaryPlannedSkillCode(summary)); skillCode != "" && strings.TrimSpace(summary.ReplyText) != "" {
return "skill"
}
switch strings.TrimSpace(summary.Status) {
case "completed":
return "reply"
case "fallback":
return "fallback"
case "error":
return "error"
case "interrupted":
return "interrupted"
case "expired":
return "expired"
default:
return strings.TrimSpace(summary.Status)
}
}
func buildRunLogReplyText(summary *Summary) string {
if summary == nil {
return ""
}
return strings.TrimSpace(summary.ReplyText)
}
func summaryPlannedSkillCode(summary *Summary) string {
if summary == nil {
return ""
}
return strings.TrimSpace(summary.PlannedSkillCode)
}
func (s *aiReplyService) incrementAIReplyRounds(conversationID int64, nextRounds int, aiAgentName string) error {
return repositories.ConversationRepository.Updates(sqls.DB(), conversationID, map[string]any{
"ai_reply_rounds": nextRounds,
"update_user_id": 0,
"update_user_name": strings.TrimSpace(aiAgentName),
"updated_at": time.Now(),
})
}
func buildConversationInterrupt(conversation models.Conversation, message models.Message, aiAgent models.AIAgent, summary *Summary) *models.ConversationInterrupt {
if summary == nil {
return nil
}
now := time.Now()
item := svc.ConversationInterruptService.GetByCheckPointID(summary.CheckPointID)
if item == nil {
item = &models.ConversationInterrupt{
CheckPointID: summary.CheckPointID,
CreatedAt: now,
}
}
item.ConversationID = conversation.ID
item.AIAgentID = aiAgent.ID
item.SourceMessageID = message.ID
item.InterruptID = firstInterruptID(summary)
item.InterruptType = firstInterruptType(summary)
item.Status = "pending"
item.PromptText = resolveInterruptPrompt(summary)
item.UpdatedAt = now
return item
}
func resolveInterruptPrompt(summary *Summary) string {
if summary == nil || len(summary.Interrupts) == 0 {
return "请继续补充信息后再试。"
}
if prompt := extractInterruptMessage(summary.Interrupts[0].InfoPreview); prompt != "" {
return prompt
}
if prompt := strings.TrimSpace(summary.Interrupts[0].InfoPreview); prompt != "" {
return prompt
}
return "请继续补充信息后再试。"
}
func extractInterruptMessage(infoPreview string) string {
infoPreview = strings.TrimSpace(infoPreview)
if infoPreview == "" {
return ""
}
payload := make(map[string]any)
if err := json.Unmarshal([]byte(infoPreview), &payload); err != nil {
return ""
}
if message, ok := payload["message"].(string); ok {
return strings.TrimSpace(message)
}
return ""
}
func firstInterruptID(summary *Summary) string {
if summary == nil || len(summary.Interrupts) == 0 {
return ""
}
return strings.TrimSpace(summary.Interrupts[0].ID)
}
func firstInterruptType(summary *Summary) string {
if summary == nil || len(summary.Interrupts) == 0 {
return ""
}
return strings.TrimSpace(summary.Interrupts[0].Type)
}
func firstInvokedToolCode(summary *Summary) string {
if summary == nil {
return ""
}
if len(summary.InvokedToolCodes) > 0 {
return strings.TrimSpace(summary.InvokedToolCodes[0])
}
return ""
}
func isCancellationReply(replyText string) bool {
replyText = strings.TrimSpace(replyText)
return strings.Contains(replyText, "已取消本次工单创建")
}
func isCheckpointMissingError(err error) bool {
if err == nil {
return false
}
message := strings.ToLower(strings.TrimSpace(err.Error()))
return strings.Contains(message, "failed to load from checkpoint") && strings.Contains(message, "not exist")
}
func (s *aiReplyService) handoffConversation(conversation models.Conversation, aiAgent models.AIAgent, reason string) error {
now := time.Now()
if err := sqls.WithTransaction(func(ctx *sqls.TxContext) error {
if err := repositories.ConversationRepository.Updates(ctx.Tx, conversation.ID, map[string]any{
"handoff_at": now,
"handoff_reason": strings.TrimSpace(reason),
"status": enums.IMConversationStatusPending,
"current_team_id": 0,
"current_assignee_id": 0,
"update_user_id": 0,
"update_user_name": aiAgent.Name,
"updated_at": now,
}); err != nil {
return err
}
return svc.ConversationEventLogService.CreateEvent(ctx, conversation.ID, enums.IMEventTypeTransfer, enums.IMSenderTypeAI, aiAgent.ID, "AI转人工", strings.TrimSpace(reason))
}); err != nil {
return err
}
if _, err := svc.MessageService.SendAIMessage(conversation.ID, aiAgent.ID, fmt.Sprintf("ai_handoff_%d", conversation.LastMessageID), enums.IMMessageTypeText, "已为你转接人工客服,请稍候。", "", s.buildAIPrincipal(aiAgent)); err != nil {
return err
}
if _, err := svc.ConversationDispatchService.DispatchConversation(conversation.ID); err != nil {
slog.Warn("auto dispatch conversation after ai handoff failed",
"conversation_id", conversation.ID,
"ai_agent_id", aiAgent.ID,
"error", err)
}
return nil
}
func (s *aiReplyService) buildAIPrincipal(aiAgent models.AIAgent) *dto.AuthPrincipal {
username := "AI"
if strings.TrimSpace(aiAgent.Name) != "" {
username = aiAgent.Name
}
return &dto.AuthPrincipal{
UserID: 0,
Username: username,
Nickname: username,
}
}
+177
View File
@@ -0,0 +1,177 @@
package runtime
import (
"context"
"strings"
"cs-agent/internal/ai/runtime/internal/engine"
"cs-agent/internal/ai/runtime/registry"
"cs-agent/internal/ai/runtime/tools"
"cs-agent/internal/ai/skills"
)
var Service = newService()
func newService() *service {
return &service{
runtime: engine.NewService(),
registry: registry.NewRegistry(
tools.NewCreateTicketConfirmTool(),
),
}
}
type service struct {
runtime *engine.Service
registry *registry.Registry
}
func (s *service) Run(ctx context.Context, req Request) (*Summary, error) {
skillSummary, skillErr := s.tryRunSkill(ctx, req)
if skillSummary != nil && strings.TrimSpace(skillSummary.ReplyText) != "" {
return skillSummary, nil
}
if err := s.prepareToolsForRun(&req); err != nil {
return nil, err
}
summary, err := s.runtime.Run(ctx, engine.Request{
Conversation: req.Conversation,
UserMessage: req.UserMessage,
AIAgent: req.AIAgent,
AIConfig: req.AIConfig,
CheckPointID: req.CheckPointID,
ExtraTools: req.ExtraTools,
ExtraToolCodes: req.ExtraToolCodes,
})
if err != nil {
ret := toSummary(summary)
if ret != nil && skillErr != nil && strings.TrimSpace(ret.PlanReason) == "" {
ret.PlanReason = "skill_failed_fallback_runtime"
}
return ret, err
}
ret := toSummary(summary)
if ret != nil && skillErr != nil && strings.TrimSpace(ret.PlanReason) == "" {
ret.PlanReason = "skill_failed_fallback_runtime"
}
return ret, nil
}
func (s *service) Resume(ctx context.Context, req ResumeRequest) (*Summary, error) {
if err := s.prepareToolsForResume(&req); err != nil {
return nil, err
}
summary, err := s.runtime.Resume(ctx, engine.ResumeRequest{
Conversation: req.Conversation,
AIAgent: req.AIAgent,
AIConfig: req.AIConfig,
CheckPointID: req.CheckPointID,
ResumeData: req.ResumeData,
ExtraTools: req.ExtraTools,
ExtraToolCodes: req.ExtraToolCodes,
})
if err != nil {
return toSummary(summary), err
}
return toSummary(summary), nil
}
func (s *service) prepareToolsForRun(req *Request) error {
if req == nil || len(req.ExtraTools) > 0 || len(req.ExtraToolCodes) > 0 || s.registry == nil {
return nil
}
toolSet, err := s.registry.Resolve(registry.Context{
Conversation: req.Conversation,
AIAgent: req.AIAgent,
AIConfig: req.AIConfig,
UserMessage: req.UserMessage,
})
if err != nil {
return err
}
req.ExtraTools = toolSet.Tools
req.ExtraToolCodes = toolSet.ToolCodes
return nil
}
func (s *service) prepareToolsForResume(req *ResumeRequest) error {
if req == nil || len(req.ExtraTools) > 0 || len(req.ExtraToolCodes) > 0 || s.registry == nil {
return nil
}
toolSet, err := s.registry.Resolve(registry.Context{
Conversation: req.Conversation,
AIAgent: req.AIAgent,
AIConfig: req.AIConfig,
})
if err != nil {
return err
}
req.ExtraTools = toolSet.Tools
req.ExtraToolCodes = toolSet.ToolCodes
return nil
}
func toSummary(summary *engine.Summary) *Summary {
if summary == nil {
return nil
}
ret := &Summary{
RunID: summary.RunID,
Status: summary.Status,
ReplyText: summary.ReplyText,
PlannedSkillCode: "",
PlanReason: "",
ModelName: summary.ModelName,
PromptTokens: summary.PromptTokens,
CompletionTokens: summary.CompletionTokens,
HistoryMessageCount: summary.HistoryMessageCount,
RetrieverCount: summary.RetrieverCount,
ToolCallCount: summary.ToolCallCount,
ToolCodes: append([]string(nil), summary.ToolCodes...),
InvokedToolCodes: append([]string(nil), summary.InvokedToolCodes...),
CheckPointID: summary.CheckPointID,
Interrupted: summary.Interrupted,
TraceData: summary.TraceData,
ErrorMessage: summary.ErrorMessage,
}
if len(summary.Interrupts) > 0 {
ret.Interrupts = make([]InterruptContextSummary, 0, len(summary.Interrupts))
for _, item := range summary.Interrupts {
ret.Interrupts = append(ret.Interrupts, InterruptContextSummary{
Type: item.Type,
ID: item.ID,
InfoPreview: item.InfoPreview,
})
}
}
return ret
}
func (s *service) tryRunSkill(ctx context.Context, req Request) (*Summary, error) {
if req.AIAgent == nil || req.AIConfig == nil || req.UserMessage == nil || req.Conversation == nil {
return nil, nil
}
result, err := skills.Execute(ctx, skills.RuntimeContext{
AIAgentID: req.AIAgent.ID,
UserMessage: strings.TrimSpace(req.UserMessage.Content),
ConversationID: req.Conversation.ID,
})
if err != nil {
return nil, err
}
if result == nil || result.Plan == nil || result.Plan.Skill == nil {
return nil, nil
}
traceData := ""
if result.RunLog != nil {
traceData = result.RunLog.TraceData
}
return &Summary{
Status: "completed",
ReplyText: strings.TrimSpace(result.ReplyText),
PlannedSkillCode: strings.TrimSpace(result.Plan.Skill.Code),
PlanReason: strings.TrimSpace(result.Plan.MatchReason),
ModelName: req.AIConfig.ModelName,
TraceData: traceData,
}, nil
}
@@ -0,0 +1,215 @@
package tools
import (
"context"
"encoding/json"
"fmt"
"strings"
"cs-agent/internal/ai/runtime/registry"
"cs-agent/internal/models"
"cs-agent/internal/pkg/dto"
"cs-agent/internal/pkg/dto/request"
"cs-agent/internal/services"
componenttool "github.com/cloudwego/eino/components/tool"
einotool "github.com/cloudwego/eino/components/tool"
"github.com/cloudwego/eino/schema"
einojsonschema "github.com/eino-contrib/jsonschema"
orderedmap "github.com/wk8/go-ordered-map/v2"
)
const (
CreateTicketConfirmToolCode = "builtin/create_ticket_with_confirmation"
CreateTicketConfirmToolName = "create_ticket_with_confirmation"
)
type CreateTicketConfirmState struct {
Request request.CreateTicketFromConversationRequest
}
type CreateTicketConfirmInterruptInfo struct {
Type string `json:"type"`
Message string `json:"message"`
}
func init() {
schema.RegisterName[CreateTicketConfirmState]("cs_agent_create_ticket_confirm_state")
schema.RegisterName[CreateTicketConfirmInterruptInfo]("cs_agent_create_ticket_confirm_interrupt_info")
}
type CreateTicketConfirmTool struct {
conversation *models.Conversation
aiAgent *models.AIAgent
}
func NewCreateTicketConfirmTool() *CreateTicketConfirmTool {
return &CreateTicketConfirmTool{}
}
func (t *CreateTicketConfirmTool) Name() string {
return CreateTicketConfirmToolName
}
func (t *CreateTicketConfirmTool) Code() string {
return CreateTicketConfirmToolCode
}
func (t *CreateTicketConfirmTool) Enabled(ctx registry.Context) bool {
return ctx.Conversation != nil && ctx.AIAgent != nil
}
func (t *CreateTicketConfirmTool) Build(ctx registry.Context) (einotool.BaseTool, error) {
if !t.Enabled(ctx) {
return nil, nil
}
return &CreateTicketConfirmTool{
conversation: ctx.Conversation,
aiAgent: ctx.AIAgent,
}, nil
}
func (t *CreateTicketConfirmTool) Info(ctx context.Context) (*schema.ToolInfo, error) {
return &schema.ToolInfo{
Name: CreateTicketConfirmToolName,
Desc: "当用户明确希望创建工单、投诉单、报障单,且你已经整理出工单标题和描述后,调用此工具。该工具不会立即创建工单,而是会先向用户发起确认;只有用户确认后才真正创建。不要在信息不足时调用。",
ParamsOneOf: schema.NewParamsOneOfByJSONSchema(&einojsonschema.Schema{
Version: einojsonschema.Version,
Type: "object",
Required: []string{
"title",
"description",
},
Properties: orderedmap.New[string, *einojsonschema.Schema](orderedmap.WithInitialData(
orderedmap.Pair[string, *einojsonschema.Schema]{
Key: "title",
Value: &einojsonschema.Schema{
Type: "string",
Description: "工单标题,简洁概括问题。",
},
},
orderedmap.Pair[string, *einojsonschema.Schema]{
Key: "description",
Value: &einojsonschema.Schema{
Type: "string",
Description: "工单描述,清晰整理用户问题、现象和诉求。",
},
},
orderedmap.Pair[string, *einojsonschema.Schema]{
Key: "priority",
Value: &einojsonschema.Schema{
Type: "integer",
Description: "工单优先级,可选;未知时可不传。",
},
},
orderedmap.Pair[string, *einojsonschema.Schema]{
Key: "severity",
Value: &einojsonschema.Schema{
Type: "integer",
Description: "严重度,可选;1=轻微,2=严重,3=致命。",
},
},
)),
}),
Extra: map[string]any{
"toolCode": CreateTicketConfirmToolCode,
},
}, nil
}
func (t *CreateTicketConfirmTool) InvokableRun(ctx context.Context, argumentsInJSON string, opts ...einotool.Option) (string, error) {
if t == nil || t.conversation == nil || t.aiAgent == nil {
return "", fmt.Errorf("ticket confirmation tool not initialized")
}
wasInterrupted, hasState, state := componenttool.GetInterruptState[CreateTicketConfirmState](ctx)
if !wasInterrupted {
req, err := t.buildCreateRequest(argumentsInJSON)
if err != nil {
return "", err
}
info := CreateTicketConfirmInterruptInfo{
Type: "ticket_creation_confirmation",
Message: t.buildConfirmationPrompt(req),
}
return "", componenttool.StatefulInterrupt(ctx, info, CreateTicketConfirmState{Request: req})
}
if !hasState {
return "", fmt.Errorf("ticket confirmation state missing")
}
isResumeTarget, hasData, resumeText := componenttool.GetResumeContext[string](ctx)
if !isResumeTarget {
info := CreateTicketConfirmInterruptInfo{
Type: "ticket_creation_confirmation",
Message: t.buildConfirmationPrompt(state.Request),
}
return "", componenttool.StatefulInterrupt(ctx, info, state)
}
if !hasData {
info := CreateTicketConfirmInterruptInfo{
Type: "ticket_creation_confirmation",
Message: "请回复“确认”或“取消”。",
}
return "", componenttool.StatefulInterrupt(ctx, info, state)
}
decision := ParseConfirmationDecision(resumeText)
switch decision {
case DecisionConfirm:
item, err := services.TicketService.CreateFromConversation(state.Request, t.buildAIPrincipal())
if err != nil {
return "", err
}
return fmt.Sprintf("工单已创建,工单号:%s,标题:%s。", strings.TrimSpace(item.TicketNo), strings.TrimSpace(item.Title)), nil
case DecisionCancel:
return "已取消本次工单创建。", nil
default:
info := CreateTicketConfirmInterruptInfo{
Type: "ticket_creation_confirmation",
Message: "我需要你的明确确认,请直接回复“确认”或“取消”。",
}
return "", componenttool.StatefulInterrupt(ctx, info, state)
}
}
func (t *CreateTicketConfirmTool) buildCreateRequest(argumentsInJSON string) (request.CreateTicketFromConversationRequest, error) {
req := request.CreateTicketFromConversationRequest{
ConversationID: t.conversation.ID,
SyncToConversation: true,
}
raw := make(map[string]any)
if strings.TrimSpace(argumentsInJSON) != "" {
if err := json.Unmarshal([]byte(argumentsInJSON), &raw); err != nil {
return req, fmt.Errorf("invalid create ticket arguments: %w", err)
}
}
req.Title = strings.TrimSpace(getStringValue(raw, "title"))
req.Description = strings.TrimSpace(getStringValue(raw, "description"))
req.Priority = getInt64Value(raw, "priority")
req.Severity = int(getInt64Value(raw, "severity"))
if req.Title == "" {
req.Title = strings.TrimSpace(t.conversation.Subject)
}
if req.Description == "" {
req.Description = strings.TrimSpace(t.conversation.LastMessageSummary)
}
if strings.TrimSpace(req.Title) == "" {
return req, fmt.Errorf("ticket title is required")
}
return req, nil
}
func (t *CreateTicketConfirmTool) buildConfirmationPrompt(req request.CreateTicketFromConversationRequest) string {
return fmt.Sprintf("我准备为你创建工单。\n标题:%s\n描述:%s\n请直接回复“确认”或“取消”。",
strings.TrimSpace(req.Title), strings.TrimSpace(req.Description))
}
func (t *CreateTicketConfirmTool) buildAIPrincipal() *dto.AuthPrincipal {
username := "AI"
if strings.TrimSpace(t.aiAgent.Name) != "" {
username = strings.TrimSpace(t.aiAgent.Name)
}
return &dto.AuthPrincipal{
UserID: 0,
Username: username,
Nickname: username,
}
}
+63
View File
@@ -0,0 +1,63 @@
package tools
import (
"fmt"
"strings"
)
type Decision string
const (
DecisionConfirm Decision = "confirm"
DecisionCancel Decision = "cancel"
)
func ParseConfirmationDecision(value string) Decision {
value = strings.ToLower(strings.TrimSpace(value))
if value == "" {
return ""
}
confirmWords := []string{"确认", "是", "好的", "可以", "ok", "yes", "继续", "同意"}
for _, item := range confirmWords {
if strings.Contains(value, item) {
return DecisionConfirm
}
}
cancelWords := []string{"取消", "不用", "不需要", "算了", "no"}
for _, item := range cancelWords {
if strings.Contains(value, item) {
return DecisionCancel
}
}
return ""
}
func getStringValue(data map[string]any, key string) string {
value, ok := data[key]
if !ok || value == nil {
return ""
}
switch v := value.(type) {
case string:
return v
default:
return fmt.Sprintf("%v", v)
}
}
func getInt64Value(data map[string]any, key string) int64 {
value, ok := data[key]
if !ok || value == nil {
return 0
}
switch v := value.(type) {
case float64:
return int64(v)
case int64:
return v
case int:
return int64(v)
default:
return 0
}
}
+54
View File
@@ -0,0 +1,54 @@
package runtime
import (
"cs-agent/internal/models"
einotool "github.com/cloudwego/eino/components/tool"
)
type Request struct {
Conversation *models.Conversation
UserMessage *models.Message
AIAgent *models.AIAgent
AIConfig *models.AIConfig
CheckPointID string
ExtraTools []einotool.BaseTool
ExtraToolCodes map[string]string
}
type ResumeRequest struct {
Conversation *models.Conversation
AIAgent *models.AIAgent
AIConfig *models.AIConfig
CheckPointID string
ResumeData map[string]any
ExtraTools []einotool.BaseTool
ExtraToolCodes map[string]string
}
type InterruptContextSummary struct {
Type string `json:"type,omitempty"`
ID string `json:"id"`
InfoPreview string `json:"infoPreview,omitempty"`
}
type Summary struct {
RunID string
Status string
ReplyText string
PlannedSkillCode string
PlanReason string
ModelName string
PromptTokens int
CompletionTokens int
HistoryMessageCount int
RetrieverCount int
ToolCallCount int
ToolCodes []string
InvokedToolCodes []string
CheckPointID string
Interrupted bool
Interrupts []InterruptContextSummary
TraceData string
ErrorMessage string
}