Init
This commit is contained in:
@@ -0,0 +1,101 @@
|
||||
package ai
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
openai "github.com/openai/openai-go/v3"
|
||||
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"cs-agent/internal/pkg/errorsx"
|
||||
)
|
||||
|
||||
type EmbeddingResult struct {
|
||||
Vector []float32
|
||||
TokensUsed int
|
||||
ModelName string
|
||||
Dimension int
|
||||
}
|
||||
|
||||
type embedding struct{}
|
||||
|
||||
var Embedding = &embedding{}
|
||||
|
||||
func (s *embedding) GetModel(ctx context.Context) (*models.AIConfig, error) {
|
||||
config, err := GetEnabledAIConfig(enums.AIModelTypeEmbedding)
|
||||
if err != nil {
|
||||
return nil, errorsx.BusinessError(2001, "未配置可用的 Embedding 模型")
|
||||
}
|
||||
return config, nil
|
||||
}
|
||||
|
||||
func (s *embedding) GenerateEmbedding(ctx context.Context, text string) (*EmbeddingResult, error) {
|
||||
if text == "" {
|
||||
return nil, errorsx.InvalidParam("文本内容不能为空")
|
||||
}
|
||||
|
||||
result, err := s.callEmbeddingAPI(ctx, text)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (s *embedding) GenerateBatchEmbeddings(ctx context.Context, texts []string) ([]EmbeddingResult, error) {
|
||||
if len(texts) == 0 {
|
||||
return nil, errorsx.InvalidParam("文本列表不能为空")
|
||||
}
|
||||
|
||||
results := make([]EmbeddingResult, 0, len(texts))
|
||||
for _, text := range texts {
|
||||
result, err := s.callEmbeddingAPI(ctx, text)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to generate embedding for text: %w", err)
|
||||
}
|
||||
results = append(results, *result)
|
||||
}
|
||||
|
||||
return results, nil
|
||||
}
|
||||
|
||||
func (s *embedding) callEmbeddingAPI(ctx context.Context, text string) (*EmbeddingResult, error) {
|
||||
config, err := GetEnabledAIConfig(enums.AIModelTypeEmbedding)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
client := newOpenAIClient(config)
|
||||
embeddingResp, err := client.Embeddings.New(ctx, openai.EmbeddingNewParams{
|
||||
Input: openai.EmbeddingNewParamsInputUnion{
|
||||
OfString: openai.String(text),
|
||||
},
|
||||
Model: openai.EmbeddingModel(config.ModelName),
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to call embedding api: %w", err)
|
||||
}
|
||||
|
||||
if len(embeddingResp.Data) == 0 {
|
||||
return nil, fmt.Errorf("no embedding data in response")
|
||||
}
|
||||
vector := make([]float32, 0, len(embeddingResp.Data[0].Embedding))
|
||||
for _, item := range embeddingResp.Data[0].Embedding {
|
||||
vector = append(vector, float32(item))
|
||||
}
|
||||
|
||||
return &EmbeddingResult{
|
||||
Vector: vector,
|
||||
TokensUsed: int(embeddingResp.Usage.TotalTokens),
|
||||
ModelName: embeddingResp.Model,
|
||||
Dimension: len(vector),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *embedding) GetDimension(ctx context.Context) (int, error) {
|
||||
model, err := s.GetModel(ctx)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return model.Dimension, nil
|
||||
}
|
||||
@@ -0,0 +1,105 @@
|
||||
package ai
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"unicode/utf8"
|
||||
|
||||
"github.com/mlogclub/simple/common/strs"
|
||||
openai "github.com/openai/openai-go/v3"
|
||||
"github.com/openai/openai-go/v3/shared"
|
||||
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
)
|
||||
|
||||
type ChatCompletionResult struct {
|
||||
Content string
|
||||
ModelName string
|
||||
PromptTokens int
|
||||
CompletionTokens int
|
||||
}
|
||||
|
||||
type llm struct{}
|
||||
|
||||
var LLM = &llm{}
|
||||
|
||||
func (s *llm) Chat(ctx context.Context, systemPrompt string, userPrompt string) (*ChatCompletionResult, error) {
|
||||
config, err := GetEnabledAIConfig(enums.AIModelTypeLLM)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return s.ChatWithConfig(ctx, config, systemPrompt, userPrompt)
|
||||
}
|
||||
|
||||
func (s *llm) ChatWithConfig(ctx context.Context, config *models.AIConfig, systemPrompt string, userPrompt string) (*ChatCompletionResult, error) {
|
||||
if config == nil {
|
||||
return nil, fmt.Errorf("ai config is nil")
|
||||
}
|
||||
|
||||
messages := make([]openai.ChatCompletionMessageParamUnion, 0, 2)
|
||||
if strs.IsNotBlank(systemPrompt) {
|
||||
messages = append(messages, openai.ChatCompletionMessageParamUnion{
|
||||
OfSystem: &openai.ChatCompletionSystemMessageParam{
|
||||
Content: openai.ChatCompletionSystemMessageParamContentUnion{
|
||||
OfString: openai.String(systemPrompt),
|
||||
},
|
||||
},
|
||||
})
|
||||
}
|
||||
messages = append(messages, openai.ChatCompletionMessageParamUnion{
|
||||
OfUser: &openai.ChatCompletionUserMessageParam{
|
||||
Content: openai.ChatCompletionUserMessageParamContentUnion{
|
||||
OfString: openai.String(userPrompt),
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
params := openai.ChatCompletionNewParams{
|
||||
Messages: messages,
|
||||
Model: shared.ChatModel(config.ModelName),
|
||||
}
|
||||
if config.MaxOutputTokens > 0 {
|
||||
params.MaxCompletionTokens = openai.Int(int64(config.MaxOutputTokens))
|
||||
}
|
||||
applyProviderSpecificChatParams(¶ms, config)
|
||||
|
||||
client := newOpenAIClient(config)
|
||||
chatResp, err := client.Chat.Completions.New(ctx, params)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to call llm api (model=%s provider=%s system_chars=%d user_chars=%d max_output_tokens=%d): %w",
|
||||
config.ModelName, config.Provider, utf8.RuneCountInString(systemPrompt), utf8.RuneCountInString(userPrompt), config.MaxOutputTokens, err)
|
||||
}
|
||||
if len(chatResp.Choices) == 0 {
|
||||
return nil, fmt.Errorf("no llm choices in response")
|
||||
}
|
||||
|
||||
content := strings.TrimSpace(chatResp.Choices[0].Message.Content)
|
||||
return &ChatCompletionResult{
|
||||
Content: content,
|
||||
ModelName: config.ModelName,
|
||||
PromptTokens: int(chatResp.Usage.PromptTokens),
|
||||
CompletionTokens: int(chatResp.Usage.CompletionTokens),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func applyProviderSpecificChatParams(params *openai.ChatCompletionNewParams, config *models.AIConfig) {
|
||||
if params == nil || config == nil {
|
||||
return
|
||||
}
|
||||
if isDashScopeQwenThinkingModel(config) {
|
||||
params.SetExtraFields(map[string]any{
|
||||
"enable_thinking": false,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func isDashScopeQwenThinkingModel(config *models.AIConfig) bool {
|
||||
if config == nil {
|
||||
return false
|
||||
}
|
||||
baseURL := strings.ToLower(strings.TrimSpace(config.BaseURL))
|
||||
modelName := strings.ToLower(strings.TrimSpace(config.ModelName))
|
||||
return strings.Contains(baseURL, "dashscope.aliyuncs.com") && strings.HasPrefix(modelName, "qwen3")
|
||||
}
|
||||
@@ -0,0 +1,198 @@
|
||||
package mcps
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"cs-agent/internal/pkg/errorsx"
|
||||
|
||||
"github.com/modelcontextprotocol/go-sdk/mcp"
|
||||
)
|
||||
|
||||
type Client struct{}
|
||||
|
||||
func NewClient() *Client {
|
||||
return &Client{}
|
||||
}
|
||||
|
||||
func (c *Client) TestConnection(ctx context.Context, cfg ServerConfig) (*ConnectionResult, error) {
|
||||
session, closeFn, err := c.connect(ctx, cfg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer closeFn()
|
||||
|
||||
initResult := session.InitializeResult()
|
||||
serverName := ""
|
||||
version := ""
|
||||
protocol := ""
|
||||
if initResult != nil {
|
||||
serverName = initResult.ServerInfo.Name
|
||||
version = initResult.ServerInfo.Version
|
||||
protocol = initResult.ProtocolVersion
|
||||
}
|
||||
return &ConnectionResult{
|
||||
ServerCode: cfg.Code,
|
||||
Endpoint: cfg.Endpoint,
|
||||
Protocol: protocol,
|
||||
ServerName: serverName,
|
||||
Version: version,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (c *Client) ListTools(ctx context.Context, cfg ServerConfig) ([]ToolInfo, error) {
|
||||
session, closeFn, err := c.connect(ctx, cfg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer closeFn()
|
||||
|
||||
result, err := session.ListTools(ctx, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("列出 MCP 工具失败: %w", err)
|
||||
}
|
||||
ret := make([]ToolInfo, 0, len(result.Tools))
|
||||
for _, tool := range result.Tools {
|
||||
ret = append(ret, ToolInfo{
|
||||
Name: tool.Name,
|
||||
Title: tool.Title,
|
||||
Description: tool.Description,
|
||||
InputSchema: tool.InputSchema,
|
||||
OutputSchema: tool.OutputSchema,
|
||||
})
|
||||
}
|
||||
return ret, nil
|
||||
}
|
||||
|
||||
func (c *Client) CallTool(ctx context.Context, cfg ServerConfig, toolName string, arguments map[string]any) (*ToolCallResult, error) {
|
||||
toolName = strings.TrimSpace(toolName)
|
||||
if toolName == "" {
|
||||
return nil, errorsx.InvalidParam("toolName不能为空")
|
||||
}
|
||||
|
||||
session, closeFn, err := c.connect(ctx, cfg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer closeFn()
|
||||
|
||||
result, err := session.CallTool(ctx, &mcp.CallToolParams{
|
||||
Name: toolName,
|
||||
Arguments: arguments,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("调用 MCP 工具失败: %w", err)
|
||||
}
|
||||
return &ToolCallResult{
|
||||
ServerCode: cfg.Code,
|
||||
ToolName: toolName,
|
||||
IsError: result.IsError,
|
||||
Content: convertContents(result.Content),
|
||||
StructuredContent: result.StructuredContent,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (c *Client) connect(ctx context.Context, cfg ServerConfig) (*mcp.ClientSession, func(), error) {
|
||||
if strings.TrimSpace(cfg.Code) == "" {
|
||||
return nil, nil, errorsx.InvalidParam("serverCode不能为空")
|
||||
}
|
||||
if strings.TrimSpace(cfg.Endpoint) == "" {
|
||||
return nil, nil, errorsx.InvalidParam("MCP endpoint不能为空")
|
||||
}
|
||||
|
||||
timeout := time.Duration(cfg.TimeoutMS) * time.Millisecond
|
||||
if timeout <= 0 {
|
||||
timeout = 15 * time.Second
|
||||
}
|
||||
connCtx, cancel := context.WithTimeout(ctx, timeout)
|
||||
|
||||
httpClient := &http.Client{
|
||||
Transport: &headerRoundTripper{
|
||||
next: http.DefaultTransport,
|
||||
headers: cfg.Headers,
|
||||
},
|
||||
}
|
||||
client := mcp.NewClient(&mcp.Implementation{
|
||||
Name: "cs-agent-mcp-client",
|
||||
Version: "v1",
|
||||
}, nil)
|
||||
transport := &mcp.StreamableClientTransport{
|
||||
Endpoint: cfg.Endpoint,
|
||||
HTTPClient: httpClient,
|
||||
MaxRetries: 0,
|
||||
DisableStandaloneSSE: true,
|
||||
}
|
||||
session, err := client.Connect(connCtx, transport, nil)
|
||||
if err != nil {
|
||||
cancel()
|
||||
return nil, nil, fmt.Errorf("连接 MCP Server 失败: %w", err)
|
||||
}
|
||||
return session, func() {
|
||||
_ = session.Close()
|
||||
cancel()
|
||||
}, nil
|
||||
}
|
||||
|
||||
func convertContents(contents []mcp.Content) []ToolResultContent {
|
||||
ret := make([]ToolResultContent, 0, len(contents))
|
||||
for _, item := range contents {
|
||||
switch v := item.(type) {
|
||||
case *mcp.TextContent:
|
||||
ret = append(ret, ToolResultContent{
|
||||
Type: "text",
|
||||
Text: v.Text,
|
||||
})
|
||||
case *mcp.ImageContent:
|
||||
ret = append(ret, ToolResultContent{
|
||||
Type: "image",
|
||||
Data: map[string]any{
|
||||
"mimeType": v.MIMEType,
|
||||
"data": v.Data,
|
||||
},
|
||||
})
|
||||
case *mcp.AudioContent:
|
||||
ret = append(ret, ToolResultContent{
|
||||
Type: "audio",
|
||||
Data: map[string]any{
|
||||
"mimeType": v.MIMEType,
|
||||
"data": v.Data,
|
||||
},
|
||||
})
|
||||
case *mcp.EmbeddedResource:
|
||||
ret = append(ret, ToolResultContent{
|
||||
Type: "resource",
|
||||
Data: v.Resource,
|
||||
})
|
||||
default:
|
||||
ret = append(ret, ToolResultContent{
|
||||
Type: "unknown",
|
||||
Data: v,
|
||||
})
|
||||
}
|
||||
}
|
||||
return ret
|
||||
}
|
||||
|
||||
type headerRoundTripper struct {
|
||||
next http.RoundTripper
|
||||
headers map[string]string
|
||||
}
|
||||
|
||||
func (r *headerRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
next := r.next
|
||||
if next == nil {
|
||||
next = http.DefaultTransport
|
||||
}
|
||||
clone := req.Clone(req.Context())
|
||||
for key, value := range r.headers {
|
||||
key = strings.TrimSpace(key)
|
||||
if key == "" {
|
||||
continue
|
||||
}
|
||||
clone.Header.Set(key, value)
|
||||
}
|
||||
return next.RoundTrip(clone)
|
||||
}
|
||||
@@ -0,0 +1,69 @@
|
||||
package providers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"cs-agent/internal/pkg/config"
|
||||
"time"
|
||||
|
||||
"github.com/modelcontextprotocol/go-sdk/mcp"
|
||||
)
|
||||
|
||||
type systemToolProvider struct{}
|
||||
|
||||
func NewSystemToolProvider() ToolProvider {
|
||||
return &systemToolProvider{}
|
||||
}
|
||||
|
||||
func (p *systemToolProvider) Name() string {
|
||||
return "system"
|
||||
}
|
||||
|
||||
func (p *systemToolProvider) Register(server *mcp.Server) error {
|
||||
mcp.AddTool(
|
||||
server,
|
||||
&mcp.Tool{
|
||||
Name: "server_time",
|
||||
Description: "获取当前服务端时间,可选传入时区。",
|
||||
},
|
||||
func(_ context.Context, _ *mcp.CallToolRequest, args serverTimeArgs) (*mcp.CallToolResult, map[string]any, error) {
|
||||
loc := time.Local
|
||||
timezone := args.Timezone
|
||||
if timezone == "" {
|
||||
timezone = "Local"
|
||||
} else if loaded, err := time.LoadLocation(timezone); err == nil {
|
||||
loc = loaded
|
||||
}
|
||||
now := time.Now().In(loc)
|
||||
return nil, map[string]any{
|
||||
"timezone": timezone,
|
||||
"timestamp": now.Format("2006-01-02 15:04:05"),
|
||||
"unix": now.Unix(),
|
||||
}, nil
|
||||
},
|
||||
)
|
||||
|
||||
mcp.AddTool(
|
||||
server,
|
||||
&mcp.Tool{
|
||||
Name: "service_info",
|
||||
Description: "查看当前 cs-agent 服务的基础运行信息。",
|
||||
},
|
||||
func(ctx context.Context, req *mcp.CallToolRequest, _ struct{}) (*mcp.CallToolResult, map[string]any, error) {
|
||||
cfg := config.Current()
|
||||
return nil, map[string]any{
|
||||
"name": "cs-agent",
|
||||
"version": "v1",
|
||||
"mcpPath": "/api/mcp",
|
||||
"port": cfg.Server.Port,
|
||||
"mcpEnabled": cfg.MCP.Enabled,
|
||||
"vectorDb": cfg.VectorDB.Type,
|
||||
"storageType": cfg.Storage.Default,
|
||||
}, nil
|
||||
},
|
||||
)
|
||||
return nil
|
||||
}
|
||||
|
||||
type serverTimeArgs struct {
|
||||
Timezone string `json:"timezone,omitempty" jsonschema:"可选时区名称,例如 Asia/Shanghai 或 UTC"`
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
package providers
|
||||
|
||||
import (
|
||||
"github.com/modelcontextprotocol/go-sdk/mcp"
|
||||
)
|
||||
|
||||
type ToolProvider interface {
|
||||
Name() string
|
||||
Register(server *mcp.Server) error
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
package mcps
|
||||
|
||||
import (
|
||||
"cs-agent/internal/ai/mcps/providers"
|
||||
|
||||
"github.com/modelcontextprotocol/go-sdk/mcp"
|
||||
)
|
||||
|
||||
func defaultProviders() []providers.ToolProvider {
|
||||
return []providers.ToolProvider{
|
||||
providers.NewSystemToolProvider(),
|
||||
// 在这里注册其他的 ToolProvider
|
||||
}
|
||||
}
|
||||
|
||||
func registerProviders(server *mcp.Server) error {
|
||||
for _, provider := range defaultProviders() {
|
||||
if err := provider.Register(server); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,72 @@
|
||||
package mcps
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
|
||||
"cs-agent/internal/pkg/config"
|
||||
"cs-agent/internal/pkg/errorsx"
|
||||
)
|
||||
|
||||
type RuntimeService struct {
|
||||
client *Client
|
||||
}
|
||||
|
||||
var Runtime = NewRuntimeService()
|
||||
|
||||
func NewRuntimeService() *RuntimeService {
|
||||
return &RuntimeService{
|
||||
client: NewClient(),
|
||||
}
|
||||
}
|
||||
|
||||
func (s *RuntimeService) CallTool(ctx context.Context, serverCode string, toolName string, arguments map[string]any) (*ToolCallResult, error) {
|
||||
server, err := s.resolveServer(serverCode)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return s.client.CallTool(ctx, server, toolName, arguments)
|
||||
}
|
||||
|
||||
func (s *RuntimeService) ListTools(ctx context.Context, serverCode string) ([]ToolInfo, error) {
|
||||
server, err := s.resolveServer(serverCode)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return s.client.ListTools(ctx, server)
|
||||
}
|
||||
|
||||
func (s *RuntimeService) resolveServer(serverCode string) (ServerConfig, error) {
|
||||
cfg := config.Current()
|
||||
if !cfg.MCP.Enabled {
|
||||
return ServerConfig{}, errorsx.InvalidParam("MCP未启用")
|
||||
}
|
||||
serverCode = strings.TrimSpace(serverCode)
|
||||
if serverCode == "" {
|
||||
return ServerConfig{}, errorsx.InvalidParam("serverCode不能为空")
|
||||
}
|
||||
server, ok := cfg.MCP.Servers[serverCode]
|
||||
if !ok {
|
||||
return ServerConfig{}, errorsx.InvalidParam("MCP服务配置不存在")
|
||||
}
|
||||
if !server.Enabled {
|
||||
return ServerConfig{}, errorsx.InvalidParam("MCP服务未启用")
|
||||
}
|
||||
return ServerConfig{
|
||||
Code: serverCode,
|
||||
Endpoint: strings.TrimSpace(server.Endpoint),
|
||||
TimeoutMS: server.TimeoutMS,
|
||||
Headers: cloneRuntimeHeaders(server.Headers),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func cloneRuntimeHeaders(headers map[string]string) map[string]string {
|
||||
if len(headers) == 0 {
|
||||
return nil
|
||||
}
|
||||
ret := make(map[string]string, len(headers))
|
||||
for key, value := range headers {
|
||||
ret[key] = value
|
||||
}
|
||||
return ret
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
package mcps
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/modelcontextprotocol/go-sdk/mcp"
|
||||
)
|
||||
|
||||
func NewHTTPHandler() http.Handler {
|
||||
server := newServer()
|
||||
return mcp.NewStreamableHTTPHandler(func(_ *http.Request) *mcp.Server {
|
||||
return server
|
||||
}, &mcp.StreamableHTTPOptions{
|
||||
JSONResponse: true,
|
||||
SessionTimeout: 2 * time.Minute,
|
||||
})
|
||||
}
|
||||
|
||||
func newServer() *mcp.Server {
|
||||
server := mcp.NewServer(&mcp.Implementation{
|
||||
Name: "cs-agent-mcp-server",
|
||||
Title: "CS Agent MCP Server",
|
||||
Version: "v1",
|
||||
WebsiteURL: "https://github.com/modelcontextprotocol",
|
||||
}, nil)
|
||||
if err := registerProviders(server); err != nil {
|
||||
panic(fmt.Sprintf("register mcp providers failed: %v", err))
|
||||
}
|
||||
return server
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
package mcps
|
||||
|
||||
type ServerConfig struct {
|
||||
Code string
|
||||
Endpoint string
|
||||
TimeoutMS int
|
||||
Headers map[string]string
|
||||
}
|
||||
|
||||
type ServerInfo struct {
|
||||
Code string `json:"code"`
|
||||
Enabled bool `json:"enabled"`
|
||||
Endpoint string `json:"endpoint"`
|
||||
TimeoutMS int `json:"timeoutMs"`
|
||||
}
|
||||
|
||||
type ConnectionResult struct {
|
||||
ServerCode string `json:"serverCode"`
|
||||
Endpoint string `json:"endpoint"`
|
||||
Protocol string `json:"protocol"`
|
||||
ServerName string `json:"serverName"`
|
||||
Version string `json:"version"`
|
||||
}
|
||||
|
||||
type ToolInfo struct {
|
||||
Name string `json:"name"`
|
||||
Title string `json:"title"`
|
||||
Description string `json:"description"`
|
||||
InputSchema any `json:"inputSchema"`
|
||||
OutputSchema any `json:"outputSchema,omitempty"`
|
||||
}
|
||||
|
||||
type ToolResultContent struct {
|
||||
Type string `json:"type"`
|
||||
Text string `json:"text,omitempty"`
|
||||
Data any `json:"data,omitempty"`
|
||||
}
|
||||
|
||||
type ToolCallResult struct {
|
||||
ServerCode string `json:"serverCode"`
|
||||
ToolName string `json:"toolName"`
|
||||
IsError bool `json:"isError"`
|
||||
Content []ToolResultContent `json:"content"`
|
||||
StructuredContent any `json:"structuredContent,omitempty"`
|
||||
}
|
||||
@@ -0,0 +1,37 @@
|
||||
package ai
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/mlogclub/simple/sqls"
|
||||
openai "github.com/openai/openai-go/v3"
|
||||
"github.com/openai/openai-go/v3/option"
|
||||
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"cs-agent/internal/pkg/errorsx"
|
||||
"cs-agent/internal/repositories"
|
||||
)
|
||||
|
||||
func newOpenAIClient(config *models.AIConfig) openai.Client {
|
||||
opts := []option.RequestOption{
|
||||
option.WithAPIKey(config.APIKey),
|
||||
option.WithBaseURL(config.BaseURL),
|
||||
}
|
||||
if config.TimeoutMS > 0 {
|
||||
opts = append(opts, option.WithRequestTimeout(time.Duration(config.TimeoutMS)*time.Millisecond))
|
||||
}
|
||||
if config.MaxRetryCount >= 0 {
|
||||
opts = append(opts, option.WithMaxRetries(config.MaxRetryCount))
|
||||
}
|
||||
|
||||
return openai.NewClient(opts...)
|
||||
}
|
||||
|
||||
func GetEnabledAIConfig(modelType enums.AIModelType) (*models.AIConfig, error) {
|
||||
item := repositories.AIConfigRepository.GetEnabled(sqls.DB(), modelType)
|
||||
if item == nil {
|
||||
return nil, errorsx.BusinessError(2005, "未配置可用的 AI 配置")
|
||||
}
|
||||
return item, nil
|
||||
}
|
||||
@@ -0,0 +1,432 @@
|
||||
package rag
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/mlogclub/simple/sqls"
|
||||
|
||||
"cs-agent/internal/ai"
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/dto"
|
||||
"cs-agent/internal/pkg/dto/request"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"cs-agent/internal/pkg/errorsx"
|
||||
"cs-agent/internal/repositories"
|
||||
)
|
||||
|
||||
type answer struct {
|
||||
}
|
||||
|
||||
var Answer = &answer{}
|
||||
|
||||
func (s *answer) DebugSearch(ctx context.Context, req request.KnowledgeSearchRequest) (*response.KnowledgeSearchResponse, error) {
|
||||
if strings.TrimSpace(req.Question) == "" {
|
||||
return nil, errorsx.InvalidParam("问题不能为空")
|
||||
}
|
||||
startedAt := time.Now()
|
||||
results, err := s.retrieve(req, ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
respResults := make([]response.KnowledgeSearchResult, 0, len(results))
|
||||
for _, item := range results {
|
||||
respResults = append(respResults, response.KnowledgeSearchResult{
|
||||
KnowledgeBaseID: item.KnowledgeBaseID,
|
||||
ChunkID: item.ChunkID,
|
||||
DocumentID: item.DocumentID,
|
||||
DocumentTitle: item.DocumentTitle,
|
||||
FaqID: item.FaqID,
|
||||
FaqQuestion: item.FaqQuestion,
|
||||
ChunkNo: item.ChunkNo,
|
||||
Title: item.Title,
|
||||
SectionPath: item.SectionPath,
|
||||
Content: item.Content,
|
||||
Score: float64(item.Score),
|
||||
})
|
||||
}
|
||||
|
||||
return &response.KnowledgeSearchResponse{
|
||||
Question: req.Question,
|
||||
Results: respResults,
|
||||
HitCount: len(respResults),
|
||||
LatencyMs: time.Since(startedAt).Milliseconds(),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *answer) DebugAnswer(ctx context.Context, req request.KnowledgeAnswerRequest, operator *dto.AuthPrincipal) (*response.KnowledgeAnswerResponse, error) {
|
||||
if strings.TrimSpace(req.Question) == "" {
|
||||
return nil, errorsx.InvalidParam("问题不能为空")
|
||||
}
|
||||
startedAt := time.Now()
|
||||
|
||||
retrieveStartedAt := time.Now()
|
||||
results, err := s.retrieve(request.KnowledgeSearchRequest{
|
||||
KnowledgeBaseIDs: req.KnowledgeBaseIDs,
|
||||
Question: req.Question,
|
||||
TopK: req.TopK,
|
||||
ScoreThreshold: req.ScoreThreshold,
|
||||
RerankLimit: req.RerankLimit,
|
||||
}, ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
retrieveMs := time.Since(retrieveStartedAt).Milliseconds()
|
||||
knowledgeBase := s.resolveAnswerKnowledgeBase(req.KnowledgeBaseIDs, results)
|
||||
contextResults := buildContextHits(Retrieve.SelectContextResults(results, 4000))
|
||||
|
||||
hits := make([]response.KnowledgeSearchResult, 0, len(results))
|
||||
topScore := 0.0
|
||||
for i, item := range results {
|
||||
score := float64(item.Score)
|
||||
if i == 0 {
|
||||
topScore = score
|
||||
}
|
||||
hits = append(hits, response.KnowledgeSearchResult{
|
||||
KnowledgeBaseID: item.KnowledgeBaseID,
|
||||
ChunkID: item.ChunkID,
|
||||
DocumentID: item.DocumentID,
|
||||
DocumentTitle: item.DocumentTitle,
|
||||
FaqID: item.FaqID,
|
||||
FaqQuestion: item.FaqQuestion,
|
||||
ChunkNo: item.ChunkNo,
|
||||
Title: item.Title,
|
||||
SectionPath: item.SectionPath,
|
||||
Content: item.Content,
|
||||
Score: score,
|
||||
})
|
||||
}
|
||||
citations := buildKnowledgeCitations(hits, 3)
|
||||
|
||||
answerMode := enums.KnowledgeAnswerMode(req.AnswerMode)
|
||||
if answerMode == 0 {
|
||||
if knowledgeBase != nil {
|
||||
answerMode = enums.KnowledgeAnswerMode(knowledgeBase.AnswerMode)
|
||||
}
|
||||
if answerMode == 0 {
|
||||
answerMode = enums.KnowledgeAnswerModeStrict
|
||||
}
|
||||
}
|
||||
|
||||
fallbackMode := enums.KnowledgeFallbackMode(req.FallbackMode)
|
||||
if fallbackMode == 0 {
|
||||
if knowledgeBase != nil {
|
||||
fallbackMode = enums.KnowledgeFallbackMode(knowledgeBase.FallbackMode)
|
||||
}
|
||||
if fallbackMode == 0 {
|
||||
fallbackMode = enums.KnowledgeFallbackModeNoAnswer
|
||||
}
|
||||
}
|
||||
|
||||
answerStatus := enums.KnowledgeAnswerStatusNormal
|
||||
answer := ""
|
||||
modelName := ""
|
||||
promptTokens := 0
|
||||
completionTokens := 0
|
||||
generateStartedAt := time.Now()
|
||||
|
||||
if len(hits) == 0 {
|
||||
answerStatus = enums.KnowledgeAnswerStatusNoAnswer
|
||||
answer = buildFallbackAnswer(fallbackMode)
|
||||
} else {
|
||||
contextText := Retrieve.BuildContext(ctx, results, 4000)
|
||||
systemPrompt := buildAnswerSystemPrompt(answerMode)
|
||||
userPrompt := fmt.Sprintf("用户问题:%s\n\n参考资料:\n%s", req.Question, contextText)
|
||||
llmResult, llmErr := ai.LLM.Chat(ctx, systemPrompt, userPrompt)
|
||||
if llmErr != nil {
|
||||
answerStatus = enums.KnowledgeAnswerStatusFallback
|
||||
answer = buildFallbackAnswer(fallbackMode)
|
||||
} else {
|
||||
answer = llmResult.Content
|
||||
modelName = llmResult.ModelName
|
||||
promptTokens = llmResult.PromptTokens
|
||||
completionTokens = llmResult.CompletionTokens
|
||||
if strings.TrimSpace(answer) == "" {
|
||||
answerStatus = enums.KnowledgeAnswerStatusFallback
|
||||
answer = buildFallbackAnswer(fallbackMode)
|
||||
}
|
||||
}
|
||||
}
|
||||
generateMs := time.Since(generateStartedAt).Milliseconds()
|
||||
rerankLimit := 0
|
||||
chunkProvider := ""
|
||||
chunkTargetTokens := 0
|
||||
chunkMaxTokens := 0
|
||||
chunkOverlapTokens := 0
|
||||
if knowledgeBase != nil {
|
||||
rerankLimit = resolveRerankLimit(req.RerankLimit, knowledgeBase.DefaultRerankLimit)
|
||||
chunkProvider = knowledgeBase.ChunkProvider
|
||||
chunkTargetTokens = knowledgeBase.ChunkTargetTokens
|
||||
chunkMaxTokens = knowledgeBase.ChunkMaxTokens
|
||||
chunkOverlapTokens = knowledgeBase.ChunkOverlapTokens
|
||||
}
|
||||
|
||||
logItem, err := RetrieveLog.CreateRetrieveLog(&CreateRetrieveLogRequest{
|
||||
KnowledgeBaseID: firstKnowledgeBaseID(req.KnowledgeBaseIDs),
|
||||
Channel: defaultRetrieveChannel(req.Channel),
|
||||
Scene: defaultRetrieveScene(req.Scene),
|
||||
SessionID: req.SessionID,
|
||||
ConversationID: req.ConversationID,
|
||||
Question: req.Question,
|
||||
RewriteQuestion: "",
|
||||
Answer: answer,
|
||||
AnswerStatus: int(answerStatus),
|
||||
ChunkProvider: chunkProvider,
|
||||
ChunkTargetTokens: chunkTargetTokens,
|
||||
ChunkMaxTokens: chunkMaxTokens,
|
||||
ChunkOverlapTokens: chunkOverlapTokens,
|
||||
RerankEnabled: rerankLimit > 0,
|
||||
RerankLimit: rerankLimit,
|
||||
Hits: hits,
|
||||
UsedHits: contextResults,
|
||||
Citations: citations,
|
||||
LatencyMs: time.Since(startedAt).Milliseconds(),
|
||||
RetrieveMs: retrieveMs,
|
||||
GenerateMs: generateMs,
|
||||
PromptTokens: promptTokens,
|
||||
CompletionTokens: completionTokens,
|
||||
ModelName: modelName,
|
||||
}, operator)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &response.KnowledgeAnswerResponse{
|
||||
Question: req.Question,
|
||||
Answer: answer,
|
||||
AnswerStatus: int(answerStatus),
|
||||
AnswerStatusName: getAnswerStatusName(answerStatus),
|
||||
Citations: citations,
|
||||
Hits: hits,
|
||||
HitCount: len(hits),
|
||||
TopScore: topScore,
|
||||
LatencyMs: time.Since(startedAt).Milliseconds(),
|
||||
RetrieveMs: retrieveMs,
|
||||
GenerateMs: generateMs,
|
||||
PromptTokens: promptTokens,
|
||||
CompletionTokens: completionTokens,
|
||||
ModelName: modelName,
|
||||
RetrieveLogID: logItem.ID,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func buildContextHits(results []RetrieveResult) []response.KnowledgeSearchResult {
|
||||
if len(results) == 0 {
|
||||
return nil
|
||||
}
|
||||
hits := make([]response.KnowledgeSearchResult, 0, len(results))
|
||||
for _, item := range results {
|
||||
hits = append(hits, response.KnowledgeSearchResult{
|
||||
KnowledgeBaseID: item.KnowledgeBaseID,
|
||||
ChunkID: item.ChunkID,
|
||||
DocumentID: item.DocumentID,
|
||||
DocumentTitle: item.DocumentTitle,
|
||||
FaqID: item.FaqID,
|
||||
FaqQuestion: item.FaqQuestion,
|
||||
ChunkNo: item.ChunkNo,
|
||||
Title: item.Title,
|
||||
SectionPath: item.SectionPath,
|
||||
Content: item.Content,
|
||||
Score: float64(item.Score),
|
||||
})
|
||||
}
|
||||
return hits
|
||||
}
|
||||
|
||||
func buildKnowledgeCitations(hits []response.KnowledgeSearchResult, limit int) []response.KnowledgeCitation {
|
||||
if len(hits) == 0 || limit <= 0 {
|
||||
return nil
|
||||
}
|
||||
citations := make([]response.KnowledgeCitation, 0, limit)
|
||||
seen := make(map[string]struct{})
|
||||
for _, item := range hits {
|
||||
key := fmt.Sprintf("%d|%d|%s|%d", item.DocumentID, item.FaqID, item.SectionPath, item.ChunkNo)
|
||||
if _, ok := seen[key]; ok {
|
||||
continue
|
||||
}
|
||||
seen[key] = struct{}{}
|
||||
citations = append(citations, response.KnowledgeCitation{
|
||||
DocumentID: item.DocumentID,
|
||||
DocumentTitle: item.DocumentTitle,
|
||||
FaqID: item.FaqID,
|
||||
FaqQuestion: item.FaqQuestion,
|
||||
ChunkNo: item.ChunkNo,
|
||||
Title: item.Title,
|
||||
SectionPath: item.SectionPath,
|
||||
Snippet: truncateCitationSnippet(item.Content, 120),
|
||||
Score: item.Score,
|
||||
})
|
||||
if len(citations) >= limit {
|
||||
break
|
||||
}
|
||||
}
|
||||
return citations
|
||||
}
|
||||
|
||||
func truncateCitationSnippet(text string, limit int) string {
|
||||
if limit <= 0 {
|
||||
return ""
|
||||
}
|
||||
runes := []rune(strings.TrimSpace(text))
|
||||
if len(runes) <= limit {
|
||||
return string(runes)
|
||||
}
|
||||
return string(runes[:limit]) + "..."
|
||||
}
|
||||
|
||||
func (s *answer) BuildDocumentIndex(ctx context.Context, documentID int64) error {
|
||||
return Index.IndexDocumentByID(ctx, documentID)
|
||||
}
|
||||
|
||||
func (s *answer) retrieve(req request.KnowledgeSearchRequest, ctx context.Context) ([]RetrieveResult, error) {
|
||||
if len(normalizeKnowledgeBaseIDs(req.KnowledgeBaseIDs)) == 0 {
|
||||
return nil, errorsx.InvalidParam("知识库不能为空")
|
||||
}
|
||||
knowledgeBases := s.loadKnowledgeBases(req.KnowledgeBaseIDs)
|
||||
|
||||
results, err := Retrieve.Retrieve(ctx, RetrieveRequest{
|
||||
KnowledgeBaseIDs: req.KnowledgeBaseIDs,
|
||||
Query: req.Question,
|
||||
TopK: req.TopK,
|
||||
ScoreThreshold: req.ScoreThreshold,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
defaultRerankLimit := resolveDefaultRerankLimit(knowledgeBases)
|
||||
rerankLimit := resolveRerankLimit(req.RerankLimit, defaultRerankLimit)
|
||||
if rerankLimit > 0 && len(results) > rerankLimit {
|
||||
return Retrieve.RetrieveWithRerank(ctx, RetrieveRequest{
|
||||
KnowledgeBaseIDs: req.KnowledgeBaseIDs,
|
||||
Query: req.Question,
|
||||
TopK: req.TopK,
|
||||
ScoreThreshold: req.ScoreThreshold,
|
||||
}, rerankLimit)
|
||||
}
|
||||
return results, nil
|
||||
}
|
||||
|
||||
func (s *answer) loadKnowledgeBases(knowledgeBaseIDs []int64) []models.KnowledgeBase {
|
||||
normalized := normalizeKnowledgeBaseIDs(knowledgeBaseIDs)
|
||||
if len(normalized) == 0 {
|
||||
return nil
|
||||
}
|
||||
items := repositories.KnowledgeBaseRepository.Find(sqls.DB(), sqls.NewCnd().In("id", normalized))
|
||||
if len(items) == 0 {
|
||||
return nil
|
||||
}
|
||||
itemMap := make(map[int64]models.KnowledgeBase, len(items))
|
||||
for _, item := range items {
|
||||
itemMap[item.ID] = item
|
||||
}
|
||||
results := make([]models.KnowledgeBase, 0, len(normalized))
|
||||
for _, id := range normalized {
|
||||
if item, ok := itemMap[id]; ok {
|
||||
results = append(results, item)
|
||||
}
|
||||
}
|
||||
return results
|
||||
}
|
||||
|
||||
func (s *answer) resolvePrimaryKnowledgeBase(knowledgeBaseIDs []int64) *models.KnowledgeBase {
|
||||
items := s.loadKnowledgeBases(knowledgeBaseIDs)
|
||||
for _, item := range items {
|
||||
return &item
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *answer) resolveAnswerKnowledgeBase(knowledgeBaseIDs []int64, results []RetrieveResult) *models.KnowledgeBase {
|
||||
items := s.loadKnowledgeBases(knowledgeBaseIDs)
|
||||
if len(items) == 0 {
|
||||
return nil
|
||||
}
|
||||
if len(results) > 0 {
|
||||
for _, item := range items {
|
||||
if item.ID == results[0].KnowledgeBaseID {
|
||||
return &item
|
||||
}
|
||||
}
|
||||
}
|
||||
return &items[0]
|
||||
}
|
||||
|
||||
func firstKnowledgeBaseID(ids []int64) int64 {
|
||||
normalized := normalizeKnowledgeBaseIDs(ids)
|
||||
if len(normalized) == 0 {
|
||||
return 0
|
||||
}
|
||||
return normalized[0]
|
||||
}
|
||||
|
||||
func resolveRerankLimit(requestLimit int, defaultLimit int) int {
|
||||
if requestLimit > 0 {
|
||||
return requestLimit
|
||||
}
|
||||
if defaultLimit > 0 {
|
||||
return defaultLimit
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func resolveDefaultRerankLimit(items []models.KnowledgeBase) int {
|
||||
limit := 0
|
||||
for _, item := range items {
|
||||
if item.DefaultRerankLimit > limit {
|
||||
limit = item.DefaultRerankLimit
|
||||
}
|
||||
}
|
||||
return limit
|
||||
}
|
||||
|
||||
func buildAnswerSystemPrompt(answerMode enums.KnowledgeAnswerMode) string {
|
||||
if answerMode == enums.KnowledgeAnswerModeAssist {
|
||||
return "你是客服知识库助手。请优先依据提供的知识片段回答,可以做轻度归纳,但不要编造未提供的事实。"
|
||||
}
|
||||
return "你是严格的客服知识库助手。只能依据提供的知识片段回答;如果资料不足,请明确说明知识库暂无明确信息。"
|
||||
}
|
||||
|
||||
func buildFallbackAnswer(fallbackMode enums.KnowledgeFallbackMode) string {
|
||||
switch fallbackMode {
|
||||
case enums.KnowledgeFallbackModeSuggestRetry:
|
||||
return "当前知识库里没有找到足够明确的信息,你可以换个更具体的问法再试一次。"
|
||||
case enums.KnowledgeFallbackModeTransferHuman:
|
||||
return "当前知识库里没有找到足够明确的信息,建议转人工进一步处理。"
|
||||
default:
|
||||
return "当前知识库暂无明确信息。"
|
||||
}
|
||||
}
|
||||
|
||||
func defaultRetrieveChannel(channel string) string {
|
||||
if strings.TrimSpace(channel) == "" {
|
||||
return string(enums.KnowledgeRetrieveChannelDebug)
|
||||
}
|
||||
return channel
|
||||
}
|
||||
|
||||
func defaultRetrieveScene(scene string) string {
|
||||
if strings.TrimSpace(scene) == "" {
|
||||
return string(enums.KnowledgeRetrieveSceneQA)
|
||||
}
|
||||
return scene
|
||||
}
|
||||
|
||||
func getAnswerStatusName(status enums.KnowledgeAnswerStatus) string {
|
||||
switch status {
|
||||
case enums.KnowledgeAnswerStatusNormal:
|
||||
return "正常"
|
||||
case enums.KnowledgeAnswerStatusNoAnswer:
|
||||
return "无答案"
|
||||
case enums.KnowledgeAnswerStatusFallback:
|
||||
return "兜底"
|
||||
case enums.KnowledgeAnswerStatusBlocked:
|
||||
return "风控拦截"
|
||||
default:
|
||||
return "未知"
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,128 @@
|
||||
package rag
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
)
|
||||
|
||||
func TestBuildFallbackAnswer(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
mode enums.KnowledgeFallbackMode
|
||||
expected string
|
||||
}{
|
||||
{
|
||||
name: "no answer",
|
||||
mode: enums.KnowledgeFallbackModeNoAnswer,
|
||||
expected: "当前知识库暂无明确信息。",
|
||||
},
|
||||
{
|
||||
name: "suggest retry",
|
||||
mode: enums.KnowledgeFallbackModeSuggestRetry,
|
||||
expected: "当前知识库里没有找到足够明确的信息,你可以换个更具体的问法再试一次。",
|
||||
},
|
||||
{
|
||||
name: "transfer human",
|
||||
mode: enums.KnowledgeFallbackModeTransferHuman,
|
||||
expected: "当前知识库里没有找到足够明确的信息,建议转人工进一步处理。",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
if got := buildFallbackAnswer(tt.mode); got != tt.expected {
|
||||
t.Fatalf("%s: expected %q, got %q", tt.name, tt.expected, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetAnswerStatusName(t *testing.T) {
|
||||
if got := getAnswerStatusName(enums.KnowledgeAnswerStatusNoAnswer); got != "无答案" {
|
||||
t.Fatalf("expected no-answer label, got %q", got)
|
||||
}
|
||||
if got := getAnswerStatusName(enums.KnowledgeAnswerStatusFallback); got != "兜底" {
|
||||
t.Fatalf("expected fallback label, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveRerankLimit(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
requestLimit int
|
||||
defaultLimit int
|
||||
expected int
|
||||
}{
|
||||
{
|
||||
name: "request overrides default",
|
||||
requestLimit: 3,
|
||||
defaultLimit: 5,
|
||||
expected: 3,
|
||||
},
|
||||
{
|
||||
name: "default used when request missing",
|
||||
requestLimit: 0,
|
||||
defaultLimit: 5,
|
||||
expected: 5,
|
||||
},
|
||||
{
|
||||
name: "zero when both missing",
|
||||
requestLimit: 0,
|
||||
defaultLimit: 0,
|
||||
expected: 0,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
if got := resolveRerankLimit(tt.requestLimit, tt.defaultLimit); got != tt.expected {
|
||||
t.Fatalf("%s: expected %d, got %d", tt.name, tt.expected, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveDefaultRerankLimit(t *testing.T) {
|
||||
items := []models.KnowledgeBase{
|
||||
{ID: 11, DefaultRerankLimit: 3},
|
||||
{ID: 22, DefaultRerankLimit: 7},
|
||||
{ID: 33, DefaultRerankLimit: 5},
|
||||
}
|
||||
|
||||
if got := resolveDefaultRerankLimit(items); got != 7 {
|
||||
t.Fatalf("expected max rerank limit 7, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildKnowledgeCitations(t *testing.T) {
|
||||
hits := []response.KnowledgeSearchResult{
|
||||
{
|
||||
DocumentID: 11,
|
||||
DocumentTitle: "退款手册",
|
||||
ChunkNo: 0,
|
||||
Title: "退款说明",
|
||||
SectionPath: "售后 > 退款说明",
|
||||
Content: "退款申请提交后,预计1-3个工作日到账。",
|
||||
Score: 0.91,
|
||||
},
|
||||
{
|
||||
DocumentID: 11,
|
||||
DocumentTitle: "退款手册",
|
||||
ChunkNo: 0,
|
||||
Title: "退款说明",
|
||||
SectionPath: "售后 > 退款说明",
|
||||
Content: "重复内容",
|
||||
Score: 0.89,
|
||||
},
|
||||
}
|
||||
|
||||
citations := buildKnowledgeCitations(hits, 3)
|
||||
if len(citations) != 1 {
|
||||
t.Fatalf("expected 1 citation, got %d", len(citations))
|
||||
}
|
||||
if citations[0].DocumentID != 11 {
|
||||
t.Fatalf("expected document id 11, got %d", citations[0].DocumentID)
|
||||
}
|
||||
if citations[0].SectionPath != "售后 > 退款说明" {
|
||||
t.Fatalf("unexpected section path: %q", citations[0].SectionPath)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,44 @@
|
||||
package chunk
|
||||
|
||||
import (
|
||||
"context"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
)
|
||||
|
||||
type fixedProvider struct{}
|
||||
|
||||
func NewFixedProvider() Provider {
|
||||
return &fixedProvider{}
|
||||
}
|
||||
|
||||
func (p *fixedProvider) Name() string {
|
||||
return string(enums.KnowledgeChunkProviderFixed)
|
||||
}
|
||||
|
||||
func (p *fixedProvider) Supports(contentType enums.KnowledgeDocumentContentType) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
func (p *fixedProvider) Chunk(ctx context.Context, req *ChunkRequest) ([]ChunkResult, error) {
|
||||
text := req.PlainText
|
||||
if text == "" {
|
||||
text = req.Content
|
||||
}
|
||||
parts := splitPlainText(text, req.Options)
|
||||
results := make([]ChunkResult, 0, len(parts))
|
||||
for i, part := range parts {
|
||||
results = append(results, ChunkResult{
|
||||
ChunkNo: i,
|
||||
Title: req.DocumentTitle,
|
||||
Content: part,
|
||||
ChunkType: enums.KnowledgeChunkTypeText,
|
||||
SectionPath: req.DocumentTitle,
|
||||
CharCount: len([]rune(part)),
|
||||
TokenCount: estimateTokenCount(part),
|
||||
Metadata: map[string]any{
|
||||
"provider": enums.KnowledgeChunkProviderFixed,
|
||||
},
|
||||
})
|
||||
}
|
||||
return results, nil
|
||||
}
|
||||
@@ -0,0 +1,12 @@
|
||||
package chunk
|
||||
|
||||
import (
|
||||
"context"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
)
|
||||
|
||||
type Provider interface {
|
||||
Name() string
|
||||
Supports(contentType enums.KnowledgeDocumentContentType) bool
|
||||
Chunk(ctx context.Context, req *ChunkRequest) ([]ChunkResult, error)
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
package chunk
|
||||
|
||||
import (
|
||||
"context"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
type Registry struct {
|
||||
providers map[string]Provider
|
||||
}
|
||||
|
||||
func NewRegistry() *Registry {
|
||||
return &Registry{
|
||||
providers: make(map[string]Provider),
|
||||
}
|
||||
}
|
||||
|
||||
func NewDefaultRegistry() *Registry {
|
||||
r := NewRegistry()
|
||||
r.Register(NewFixedProvider())
|
||||
r.Register(NewStructuredProvider())
|
||||
return r
|
||||
}
|
||||
|
||||
func (r *Registry) Register(p Provider) {
|
||||
if p == nil {
|
||||
return
|
||||
}
|
||||
r.providers[p.Name()] = p
|
||||
}
|
||||
|
||||
func (r *Registry) Get(name string) Provider {
|
||||
if name == "" {
|
||||
return nil
|
||||
}
|
||||
return r.providers[name]
|
||||
}
|
||||
|
||||
func (r *Registry) Resolve(name string, contentType enums.KnowledgeDocumentContentType) Provider {
|
||||
if p := r.Get(name); p != nil && p.Supports(contentType) {
|
||||
return p
|
||||
}
|
||||
if p := r.Get(string(enums.KnowledgeChunkProviderStructured)); p != nil && p.Supports(contentType) {
|
||||
return p
|
||||
}
|
||||
return r.Get(string(enums.KnowledgeChunkProviderFixed))
|
||||
}
|
||||
|
||||
func (r *Registry) Chunk(ctx context.Context, req *ChunkRequest) ([]ChunkResult, error) {
|
||||
if req == nil {
|
||||
return nil, fmt.Errorf("chunk request is nil")
|
||||
}
|
||||
provider := r.Resolve(req.Options.Provider, req.ContentType)
|
||||
if provider == nil {
|
||||
return nil, fmt.Errorf("chunk provider not found")
|
||||
}
|
||||
return provider.Chunk(ctx, req)
|
||||
}
|
||||
@@ -0,0 +1,261 @@
|
||||
package chunk
|
||||
|
||||
import (
|
||||
"context"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"strings"
|
||||
|
||||
"github.com/gomarkdown/markdown"
|
||||
"golang.org/x/net/html"
|
||||
)
|
||||
|
||||
type structuredProvider struct{}
|
||||
|
||||
type contentBlock struct {
|
||||
Type string
|
||||
Level int
|
||||
Text string
|
||||
Title string
|
||||
SectionPath string
|
||||
}
|
||||
|
||||
func NewStructuredProvider() Provider {
|
||||
return &structuredProvider{}
|
||||
}
|
||||
|
||||
func (p *structuredProvider) Name() string {
|
||||
return string(enums.KnowledgeChunkProviderStructured)
|
||||
}
|
||||
|
||||
func (p *structuredProvider) Supports(contentType enums.KnowledgeDocumentContentType) bool {
|
||||
switch contentType {
|
||||
case enums.KnowledgeDocumentContentTypeHTML, enums.KnowledgeDocumentContentTypeMarkdown:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func (p *structuredProvider) Chunk(ctx context.Context, req *ChunkRequest) ([]ChunkResult, error) {
|
||||
content := req.Content
|
||||
if req.ContentType == enums.KnowledgeDocumentContentTypeMarkdown {
|
||||
content = string(markdown.ToHTML([]byte(content), nil, nil))
|
||||
}
|
||||
|
||||
blocks := parseStructuredBlocks(content, req.DocumentTitle)
|
||||
if len(blocks) == 0 {
|
||||
return NewFixedProvider().Chunk(ctx, req)
|
||||
}
|
||||
|
||||
results := make([]ChunkResult, 0)
|
||||
chunkNo := 0
|
||||
for _, block := range blocks {
|
||||
parts := splitPlainText(block.Text, req.Options)
|
||||
for _, part := range parts {
|
||||
if part == "" {
|
||||
continue
|
||||
}
|
||||
results = append(results, ChunkResult{
|
||||
ChunkNo: chunkNo,
|
||||
Title: block.Title,
|
||||
Content: part,
|
||||
ChunkType: mapBlockType(block.Type),
|
||||
SectionPath: block.SectionPath,
|
||||
CharCount: len([]rune(part)),
|
||||
TokenCount: estimateTokenCount(part),
|
||||
Metadata: map[string]any{
|
||||
"provider": enums.KnowledgeChunkProviderStructured,
|
||||
"blockType": block.Type,
|
||||
"sectionPath": block.SectionPath,
|
||||
"sectionTitle": block.Title,
|
||||
},
|
||||
})
|
||||
chunkNo++
|
||||
}
|
||||
}
|
||||
if len(results) == 0 {
|
||||
return NewFixedProvider().Chunk(ctx, req)
|
||||
}
|
||||
return results, nil
|
||||
}
|
||||
|
||||
func parseStructuredBlocks(content string, documentTitle string) []contentBlock {
|
||||
content = strings.TrimSpace(content)
|
||||
if content == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
parent := &html.Node{Type: html.ElementNode, Data: "div"}
|
||||
nodes, err := html.ParseFragment(strings.NewReader(content), parent)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
var blocks []contentBlock
|
||||
headings := make([]string, 0)
|
||||
var walk func(node *html.Node)
|
||||
walk = func(node *html.Node) {
|
||||
if node == nil {
|
||||
return
|
||||
}
|
||||
if node.Type == html.ElementNode {
|
||||
switch node.Data {
|
||||
case "h1", "h2", "h3", "h4", "h5", "h6":
|
||||
title := normalizeText(nodeText(node))
|
||||
if title != "" {
|
||||
level := int(node.Data[1] - '0')
|
||||
if level <= 0 {
|
||||
level = 1
|
||||
}
|
||||
headings = updateHeadingPath(headings, level, title)
|
||||
}
|
||||
return
|
||||
case "p":
|
||||
appendBlock(&blocks, "paragraph", normalizeText(nodeText(node)), currentTitle(headings, documentTitle), strings.Join(headings, " > "))
|
||||
return
|
||||
case "ul", "ol":
|
||||
appendBlock(&blocks, "list", normalizeText(listText(node)), currentTitle(headings, documentTitle), strings.Join(headings, " > "))
|
||||
return
|
||||
case "table":
|
||||
appendBlock(&blocks, "table", normalizeText(tableText(node)), currentTitle(headings, documentTitle), strings.Join(headings, " > "))
|
||||
return
|
||||
case "pre", "code":
|
||||
appendBlock(&blocks, "code", normalizeText(nodeText(node)), currentTitle(headings, documentTitle), strings.Join(headings, " > "))
|
||||
return
|
||||
}
|
||||
}
|
||||
for child := node.FirstChild; child != nil; child = child.NextSibling {
|
||||
walk(child)
|
||||
}
|
||||
}
|
||||
for _, node := range nodes {
|
||||
walk(node)
|
||||
}
|
||||
return blocks
|
||||
}
|
||||
|
||||
func appendBlock(blocks *[]contentBlock, blockType string, text string, title string, sectionPath string) {
|
||||
text = normalizeText(text)
|
||||
if text == "" {
|
||||
return
|
||||
}
|
||||
if sectionPath == "" {
|
||||
sectionPath = title
|
||||
}
|
||||
*blocks = append(*blocks, contentBlock{
|
||||
Type: blockType,
|
||||
Text: text,
|
||||
Title: title,
|
||||
SectionPath: sectionPath,
|
||||
})
|
||||
}
|
||||
|
||||
func updateHeadingPath(headings []string, level int, title string) []string {
|
||||
if level <= 0 {
|
||||
level = 1
|
||||
}
|
||||
if len(headings) >= level {
|
||||
headings = headings[:level-1]
|
||||
}
|
||||
headings = append(headings, title)
|
||||
return headings
|
||||
}
|
||||
|
||||
func currentTitle(headings []string, documentTitle string) string {
|
||||
if len(headings) == 0 {
|
||||
return documentTitle
|
||||
}
|
||||
return headings[len(headings)-1]
|
||||
}
|
||||
|
||||
func mapBlockType(blockType string) enums.KnowledgeChunkType {
|
||||
switch blockType {
|
||||
case "table":
|
||||
return enums.KnowledgeChunkTypeTable
|
||||
case "code":
|
||||
return enums.KnowledgeChunkTypeCode
|
||||
default:
|
||||
return enums.KnowledgeChunkTypeText
|
||||
}
|
||||
}
|
||||
|
||||
func nodeText(node *html.Node) string {
|
||||
if node == nil {
|
||||
return ""
|
||||
}
|
||||
var builder strings.Builder
|
||||
writeNodeText(&builder, node)
|
||||
return builder.String()
|
||||
}
|
||||
|
||||
func writeNodeText(builder *strings.Builder, node *html.Node) {
|
||||
if node == nil {
|
||||
return
|
||||
}
|
||||
switch node.Type {
|
||||
case html.TextNode:
|
||||
builder.WriteString(node.Data)
|
||||
case html.ElementNode:
|
||||
if shouldSeparate(node.Data) {
|
||||
builder.WriteByte(' ')
|
||||
}
|
||||
}
|
||||
for child := node.FirstChild; child != nil; child = child.NextSibling {
|
||||
writeNodeText(builder, child)
|
||||
}
|
||||
if node.Type == html.ElementNode && shouldSeparate(node.Data) {
|
||||
builder.WriteByte(' ')
|
||||
}
|
||||
}
|
||||
|
||||
func shouldSeparate(tag string) bool {
|
||||
switch tag {
|
||||
case "p", "div", "br", "li", "ul", "ol", "blockquote", "pre", "table", "tr", "td", "th", "h1", "h2", "h3", "h4", "h5", "h6":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func listText(node *html.Node) string {
|
||||
items := make([]string, 0)
|
||||
for child := node.FirstChild; child != nil; child = child.NextSibling {
|
||||
if child.Type == html.ElementNode && child.Data == "li" {
|
||||
item := normalizeText(nodeText(child))
|
||||
if item != "" {
|
||||
items = append(items, item)
|
||||
}
|
||||
}
|
||||
}
|
||||
return strings.Join(items, " ")
|
||||
}
|
||||
|
||||
func tableText(node *html.Node) string {
|
||||
rows := make([]string, 0)
|
||||
var walk func(*html.Node)
|
||||
walk = func(n *html.Node) {
|
||||
if n == nil {
|
||||
return
|
||||
}
|
||||
if n.Type == html.ElementNode && n.Data == "tr" {
|
||||
cells := make([]string, 0)
|
||||
for child := n.FirstChild; child != nil; child = child.NextSibling {
|
||||
if child.Type == html.ElementNode && (child.Data == "td" || child.Data == "th") {
|
||||
cell := normalizeText(nodeText(child))
|
||||
if cell != "" {
|
||||
cells = append(cells, cell)
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(cells) > 0 {
|
||||
rows = append(rows, strings.Join(cells, " | "))
|
||||
}
|
||||
return
|
||||
}
|
||||
for child := n.FirstChild; child != nil; child = child.NextSibling {
|
||||
walk(child)
|
||||
}
|
||||
}
|
||||
walk(node)
|
||||
return strings.Join(rows, " ")
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
package chunk
|
||||
|
||||
import "cs-agent/internal/pkg/enums"
|
||||
|
||||
type ChunkRequest struct {
|
||||
KnowledgeBaseID int64
|
||||
DocumentID int64
|
||||
DocumentTitle string
|
||||
ContentType enums.KnowledgeDocumentContentType
|
||||
Content string
|
||||
PlainText string
|
||||
Options ChunkOptions
|
||||
}
|
||||
|
||||
type ChunkOptions struct {
|
||||
Provider string
|
||||
TargetTokens int
|
||||
MaxTokens int
|
||||
OverlapTokens int
|
||||
EnableFallback bool
|
||||
}
|
||||
|
||||
type ChunkResult struct {
|
||||
ChunkNo int
|
||||
Title string
|
||||
Content string
|
||||
ChunkType enums.KnowledgeChunkType
|
||||
SectionPath string
|
||||
CharCount int
|
||||
TokenCount int
|
||||
Metadata map[string]any
|
||||
}
|
||||
@@ -0,0 +1,218 @@
|
||||
package chunk
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"encoding/hex"
|
||||
"strings"
|
||||
"unicode"
|
||||
"unicode/utf8"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultTargetTokens = 300
|
||||
defaultMaxTokens = 400
|
||||
defaultOverlapTokens = 40
|
||||
)
|
||||
|
||||
func normalizeOptions(opts ChunkOptions) ChunkOptions {
|
||||
if opts.TargetTokens <= 0 {
|
||||
opts.TargetTokens = defaultTargetTokens
|
||||
}
|
||||
if opts.MaxTokens <= 0 {
|
||||
opts.MaxTokens = defaultMaxTokens
|
||||
}
|
||||
if opts.MaxTokens < opts.TargetTokens {
|
||||
opts.MaxTokens = opts.TargetTokens
|
||||
}
|
||||
if opts.OverlapTokens < 0 {
|
||||
opts.OverlapTokens = 0
|
||||
}
|
||||
if opts.OverlapTokens == 0 {
|
||||
opts.OverlapTokens = defaultOverlapTokens
|
||||
}
|
||||
if opts.Provider == "" {
|
||||
opts.Provider = string(enums.KnowledgeChunkProviderStructured)
|
||||
}
|
||||
return opts
|
||||
}
|
||||
|
||||
func normalizeText(text string) string {
|
||||
return strings.Join(strings.Fields(strings.TrimSpace(text)), " ")
|
||||
}
|
||||
|
||||
func estimateTokenCount(text string) int {
|
||||
text = strings.TrimSpace(text)
|
||||
if text == "" {
|
||||
return 0
|
||||
}
|
||||
count := 0
|
||||
inWord := false
|
||||
for _, r := range text {
|
||||
switch {
|
||||
case unicode.IsSpace(r):
|
||||
inWord = false
|
||||
case unicode.Is(unicode.Han, r):
|
||||
count++
|
||||
inWord = false
|
||||
case unicode.IsLetter(r) || unicode.IsDigit(r):
|
||||
if !inWord {
|
||||
count++
|
||||
inWord = true
|
||||
}
|
||||
default:
|
||||
count++
|
||||
inWord = false
|
||||
}
|
||||
}
|
||||
if count == 0 {
|
||||
return utf8.RuneCountInString(text)
|
||||
}
|
||||
return count
|
||||
}
|
||||
|
||||
func contentHash(text string) string {
|
||||
sum := sha256.Sum256([]byte(text))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
func splitSentences(text string) []string {
|
||||
text = strings.TrimSpace(text)
|
||||
if text == "" {
|
||||
return nil
|
||||
}
|
||||
var sentences []string
|
||||
var builder strings.Builder
|
||||
for _, r := range text {
|
||||
builder.WriteRune(r)
|
||||
switch r {
|
||||
case '\n', '。', '!', '?', '!', '?', ';', ';':
|
||||
sentence := normalizeText(builder.String())
|
||||
if sentence != "" {
|
||||
sentences = append(sentences, sentence)
|
||||
}
|
||||
builder.Reset()
|
||||
}
|
||||
}
|
||||
if builder.Len() > 0 {
|
||||
sentence := normalizeText(builder.String())
|
||||
if sentence != "" {
|
||||
sentences = append(sentences, sentence)
|
||||
}
|
||||
}
|
||||
if len(sentences) == 0 {
|
||||
return []string{normalizeText(text)}
|
||||
}
|
||||
return sentences
|
||||
}
|
||||
|
||||
func tailTextByTokens(text string, tokenLimit int) string {
|
||||
if tokenLimit <= 0 {
|
||||
return ""
|
||||
}
|
||||
sentences := splitSentences(text)
|
||||
if len(sentences) == 0 {
|
||||
return ""
|
||||
}
|
||||
var selected []string
|
||||
total := 0
|
||||
for i := len(sentences) - 1; i >= 0; i-- {
|
||||
sentence := sentences[i]
|
||||
tokens := estimateTokenCount(sentence)
|
||||
if total > 0 && total+tokens > tokenLimit {
|
||||
break
|
||||
}
|
||||
selected = append([]string{sentence}, selected...)
|
||||
total += tokens
|
||||
}
|
||||
return strings.TrimSpace(strings.Join(selected, " "))
|
||||
}
|
||||
|
||||
func splitPlainText(text string, opts ChunkOptions) []string {
|
||||
text = normalizeText(text)
|
||||
if text == "" {
|
||||
return nil
|
||||
}
|
||||
opts = normalizeOptions(opts)
|
||||
sentences := splitSentences(text)
|
||||
if len(sentences) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
chunks := make([]string, 0)
|
||||
current := make([]string, 0)
|
||||
currentTokens := 0
|
||||
|
||||
flush := func() {
|
||||
if len(current) == 0 {
|
||||
return
|
||||
}
|
||||
chunks = append(chunks, strings.Join(current, " "))
|
||||
}
|
||||
|
||||
for _, sentence := range sentences {
|
||||
sentenceTokens := estimateTokenCount(sentence)
|
||||
if sentenceTokens > opts.MaxTokens {
|
||||
if len(current) > 0 {
|
||||
flush()
|
||||
overlap := tailTextByTokens(strings.Join(current, " "), opts.OverlapTokens)
|
||||
current = nil
|
||||
currentTokens = 0
|
||||
if overlap != "" {
|
||||
current = append(current, overlap)
|
||||
currentTokens = estimateTokenCount(overlap)
|
||||
}
|
||||
}
|
||||
for _, piece := range splitLongSentence(sentence, opts.MaxTokens) {
|
||||
piece = normalizeText(piece)
|
||||
if piece != "" {
|
||||
chunks = append(chunks, piece)
|
||||
}
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
if currentTokens > 0 && currentTokens+sentenceTokens > opts.MaxTokens {
|
||||
flush()
|
||||
overlap := tailTextByTokens(strings.Join(current, " "), opts.OverlapTokens)
|
||||
current = nil
|
||||
currentTokens = 0
|
||||
if overlap != "" {
|
||||
current = append(current, overlap)
|
||||
currentTokens = estimateTokenCount(overlap)
|
||||
}
|
||||
}
|
||||
|
||||
current = append(current, sentence)
|
||||
currentTokens += sentenceTokens
|
||||
}
|
||||
|
||||
flush()
|
||||
return chunks
|
||||
}
|
||||
|
||||
func splitLongSentence(text string, maxTokens int) []string {
|
||||
runes := []rune(strings.TrimSpace(text))
|
||||
if len(runes) == 0 {
|
||||
return nil
|
||||
}
|
||||
if maxTokens <= 0 {
|
||||
return []string{text}
|
||||
}
|
||||
window := maxTokens * 2
|
||||
if window < 50 {
|
||||
window = 50
|
||||
}
|
||||
var result []string
|
||||
for start := 0; start < len(runes); start += window {
|
||||
end := start + window
|
||||
if end > len(runes) {
|
||||
end = len(runes)
|
||||
}
|
||||
part := normalizeText(string(runes[start:end]))
|
||||
if part != "" {
|
||||
result = append(result, part)
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
@@ -0,0 +1,702 @@
|
||||
package rag
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"cs-agent/internal/ai"
|
||||
ragchunk "cs-agent/internal/ai/rag/chunk"
|
||||
"cs-agent/internal/ai/rag/vectordb"
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"cs-agent/internal/repositories"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/mlogclub/simple/common/strs"
|
||||
"github.com/mlogclub/simple/sqls"
|
||||
)
|
||||
|
||||
type ChunkingConfig struct {
|
||||
Provider string
|
||||
TargetTokens int
|
||||
MaxTokens int
|
||||
OverlapTokens int
|
||||
EnableFallback bool
|
||||
}
|
||||
|
||||
type index struct {
|
||||
chunkConfig ChunkingConfig
|
||||
registry *ragchunk.Registry
|
||||
}
|
||||
|
||||
const knowledgeCollectionName = "knowledge_chunks"
|
||||
|
||||
var Index = &index{
|
||||
chunkConfig: ChunkingConfig{
|
||||
Provider: string(enums.KnowledgeChunkProviderStructured),
|
||||
TargetTokens: 300,
|
||||
MaxTokens: 400,
|
||||
OverlapTokens: 40,
|
||||
EnableFallback: true,
|
||||
},
|
||||
registry: ragchunk.NewDefaultRegistry(),
|
||||
}
|
||||
|
||||
func (s *index) IndexDocumentByID(ctx context.Context, documentID int64) error {
|
||||
document := repositories.KnowledgeDocumentRepository.Get(sqls.DB(), documentID)
|
||||
if document == nil {
|
||||
return fmt.Errorf("document not found: %d", documentID)
|
||||
}
|
||||
return s.IndexDocument(ctx, document)
|
||||
}
|
||||
|
||||
func (s *index) IndexDocument(ctx context.Context, document *models.KnowledgeDocument) error {
|
||||
start := time.Now()
|
||||
if err := s.markDocumentIndexPending(document.ID); err != nil {
|
||||
slog.Error("Failed to mark knowledge document index as pending", "document_id", document.ID, "error", err)
|
||||
}
|
||||
|
||||
fail := func(err error) error {
|
||||
if updateErr := s.markDocumentIndexFailed(document.ID, err); updateErr != nil {
|
||||
slog.Error("Failed to mark knowledge document index as failed", "document_id", document.ID, "error", updateErr)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// TODO 这里每次都查询下知识库不太友好
|
||||
knowledgeBase := repositories.KnowledgeBaseRepository.Get(sqls.DB(), document.KnowledgeBaseID)
|
||||
if knowledgeBase == nil {
|
||||
return fail(fmt.Errorf("knowledge base not found: %d", document.KnowledgeBaseID))
|
||||
}
|
||||
|
||||
existingChunks := repositories.KnowledgeChunkRepository.FindByDocumentID(sqls.DB(), document.ID)
|
||||
|
||||
chunks, err := s.registry.Chunk(ctx, &ragchunk.ChunkRequest{
|
||||
KnowledgeBaseID: document.KnowledgeBaseID,
|
||||
DocumentID: document.ID,
|
||||
DocumentTitle: document.Title,
|
||||
ContentType: document.ContentType,
|
||||
Content: document.Content,
|
||||
PlainText: ExtractPlainText(document.Content, document.ContentType),
|
||||
Options: ragchunk.ChunkOptions{
|
||||
Provider: firstNonEmptyString(knowledgeBase.ChunkProvider, s.chunkConfig.Provider),
|
||||
TargetTokens: firstPositiveInt(knowledgeBase.ChunkTargetTokens, s.chunkConfig.TargetTokens),
|
||||
MaxTokens: firstPositiveInt(knowledgeBase.ChunkMaxTokens, s.chunkConfig.MaxTokens),
|
||||
OverlapTokens: firstPositiveInt(knowledgeBase.ChunkOverlapTokens, s.chunkConfig.OverlapTokens),
|
||||
EnableFallback: s.chunkConfig.EnableFallback,
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return fail(fmt.Errorf("failed to chunk document: %w", err))
|
||||
}
|
||||
if len(chunks) == 0 {
|
||||
return fail(fmt.Errorf("no chunks generated from document"))
|
||||
}
|
||||
|
||||
collectionName := s.getCollectionName()
|
||||
provider := vectordb.GetProvider()
|
||||
if provider == nil {
|
||||
return fail(fmt.Errorf("vectordb provider not initialized"))
|
||||
}
|
||||
|
||||
if _, err := ai.Embedding.GetModel(ctx); err != nil {
|
||||
return fail(fmt.Errorf("failed to get embedding model: %w", err))
|
||||
}
|
||||
|
||||
existingVectorIDs := make([]string, 0, len(existingChunks))
|
||||
for _, chunk := range existingChunks {
|
||||
if strs.IsNotBlank(chunk.VectorID) {
|
||||
existingVectorIDs = append(existingVectorIDs, chunk.VectorID)
|
||||
}
|
||||
}
|
||||
|
||||
vectors := make([]vectordb.Vector, 0, len(chunks))
|
||||
chunkModels := make([]models.KnowledgeChunk, 0, len(chunks))
|
||||
dimension := 0
|
||||
|
||||
for i, chunk := range chunks {
|
||||
embeddingResult, err := ai.Embedding.GenerateEmbedding(ctx, chunk.Content)
|
||||
if err != nil {
|
||||
slog.Error("Failed to generate embedding for chunk", "document_id", document.ID, "chunk_index", i, "error", err)
|
||||
return fail(fmt.Errorf("failed to generate embedding for chunk %d: %w", i, err))
|
||||
}
|
||||
if dimension == 0 {
|
||||
dimension = embeddingResult.Dimension
|
||||
}
|
||||
|
||||
chunkID := buildKnowledgeChunkVectorID(knowledgeBase.ID, document.ID, chunk.ChunkNo)
|
||||
providerName := ""
|
||||
if chunk.Metadata != nil {
|
||||
if value, ok := chunk.Metadata["provider"].(string); ok {
|
||||
providerName = value
|
||||
}
|
||||
}
|
||||
chunkModel := models.KnowledgeChunk{
|
||||
KnowledgeBaseID: knowledgeBase.ID,
|
||||
DocumentID: document.ID,
|
||||
ChunkNo: chunk.ChunkNo,
|
||||
Title: chunk.Title,
|
||||
Content: chunk.Content,
|
||||
ContentHash: buildChunkContentHash(chunk.Content),
|
||||
CharCount: chunk.CharCount,
|
||||
TokenCount: chunk.TokenCount,
|
||||
ChunkType: string(chunk.ChunkType),
|
||||
SectionPath: chunk.SectionPath,
|
||||
Provider: providerName,
|
||||
VectorID: chunkID,
|
||||
Status: enums.StatusOk,
|
||||
CreatedAt: time.Now(),
|
||||
UpdatedAt: time.Now(),
|
||||
}
|
||||
chunkModels = append(chunkModels, chunkModel)
|
||||
|
||||
vectors = append(vectors, vectordb.Vector{
|
||||
ID: chunkID,
|
||||
Vector: embeddingResult.Vector,
|
||||
Payload: vectordb.ChunkPayload{
|
||||
KnowledgeBaseID: knowledgeBase.ID,
|
||||
DocumentID: document.ID,
|
||||
DocumentTitle: document.Title,
|
||||
ChunkNo: chunk.ChunkNo,
|
||||
ChunkType: string(chunk.ChunkType),
|
||||
SectionPath: chunk.SectionPath,
|
||||
Content: chunk.Content,
|
||||
Title: chunk.Title,
|
||||
Provider: providerName,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
if len(vectors) == 0 {
|
||||
return fail(fmt.Errorf("no vectors generated"))
|
||||
}
|
||||
|
||||
collectionInfo, err := provider.GetCollection(ctx, collectionName)
|
||||
if err != nil || collectionInfo == nil {
|
||||
if dimension <= 0 {
|
||||
return fail(fmt.Errorf("invalid embedding dimension: %d", dimension))
|
||||
}
|
||||
if err := provider.CreateCollection(ctx, collectionName, dimension); err != nil {
|
||||
return fail(fmt.Errorf("failed to create collection: %w", err))
|
||||
}
|
||||
slog.Info("Created collection for knowledge base", "collection", collectionName, "dimension", dimension)
|
||||
}
|
||||
|
||||
if len(existingVectorIDs) > 0 {
|
||||
if err := provider.DeleteVectors(ctx, collectionName, existingVectorIDs); err != nil {
|
||||
return fail(fmt.Errorf("failed to delete old vectors: %w", err))
|
||||
}
|
||||
}
|
||||
|
||||
if err := provider.UpsertVectors(ctx, collectionName, vectors); err != nil {
|
||||
return fail(fmt.Errorf("failed to upsert vectors: %w", err))
|
||||
}
|
||||
|
||||
if err := sqls.WithTransaction(func(ctx *sqls.TxContext) error {
|
||||
if err := ctx.Tx.Where("document_id = ?", document.ID).Delete(&models.KnowledgeChunk{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
for _, chunk := range chunkModels {
|
||||
if err := ctx.Tx.Create(&chunk).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}); err != nil {
|
||||
return fail(fmt.Errorf("failed to save chunks: %w", err))
|
||||
}
|
||||
|
||||
if err := s.markDocumentIndexIndexed(document.ID); err != nil {
|
||||
slog.Error("Failed to mark knowledge document index as indexed", "document_id", document.ID, "error", err)
|
||||
}
|
||||
|
||||
slog.Info("Document indexed successfully",
|
||||
slog.Any("document_id", document.ID),
|
||||
slog.Any("chunks_count", len(chunks)),
|
||||
slog.Any("vectors_count", len(vectors)),
|
||||
slog.Any("time_taken", time.Since(start).String()),
|
||||
)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *index) IndexFAQByID(ctx context.Context, faqID int64) error {
|
||||
faq := repositories.KnowledgeFAQRepository.Get(sqls.DB(), faqID)
|
||||
if faq == nil {
|
||||
return fmt.Errorf("faq not found: %d", faqID)
|
||||
}
|
||||
if err := s.markFAQIndexPending(faq.ID); err != nil {
|
||||
slog.Error("Failed to mark knowledge faq index as pending", "faq_id", faq.ID, "error", err)
|
||||
}
|
||||
fail := func(err error) error {
|
||||
if updateErr := s.markFAQIndexFailed(faq.ID, err); updateErr != nil {
|
||||
slog.Error("Failed to mark knowledge faq index as failed", "faq_id", faq.ID, "error", updateErr)
|
||||
}
|
||||
return err
|
||||
}
|
||||
knowledgeBase := repositories.KnowledgeBaseRepository.Get(sqls.DB(), faq.KnowledgeBaseID)
|
||||
if knowledgeBase == nil {
|
||||
return fail(fmt.Errorf("knowledge base not found: %d", faq.KnowledgeBaseID))
|
||||
}
|
||||
if knowledgeBase.KnowledgeType != string(enums.KnowledgeBaseTypeFAQ) {
|
||||
return fail(fmt.Errorf("knowledge base %d is not faq type", knowledgeBase.ID))
|
||||
}
|
||||
existingChunks := repositories.KnowledgeChunkRepository.FindByFaqID(sqls.DB(), faq.ID)
|
||||
content := buildFAQChunkContent(faq)
|
||||
if content == "" {
|
||||
return fail(fmt.Errorf("faq content is empty"))
|
||||
}
|
||||
|
||||
provider := vectordb.GetProvider()
|
||||
if provider == nil {
|
||||
return fail(fmt.Errorf("vectordb provider not initialized"))
|
||||
}
|
||||
if _, err := ai.Embedding.GetModel(ctx); err != nil {
|
||||
return fail(fmt.Errorf("failed to get embedding model: %w", err))
|
||||
}
|
||||
embeddingResult, err := ai.Embedding.GenerateEmbedding(ctx, content)
|
||||
if err != nil {
|
||||
return fail(fmt.Errorf("failed to generate embedding for faq %d: %w", faq.ID, err))
|
||||
}
|
||||
|
||||
chunkID := buildKnowledgeFAQChunkVectorID(knowledgeBase.ID, faq.ID, 0)
|
||||
chunkModel := models.KnowledgeChunk{
|
||||
KnowledgeBaseID: knowledgeBase.ID,
|
||||
FaqID: faq.ID,
|
||||
ChunkNo: 0,
|
||||
Title: faq.Question,
|
||||
Content: content,
|
||||
ContentHash: buildChunkContentHash(content),
|
||||
CharCount: len([]rune(content)),
|
||||
TokenCount: len([]rune(content)) / 2,
|
||||
ChunkType: string(enums.KnowledgeChunkTypeFAQ),
|
||||
Provider: string(enums.KnowledgeChunkProviderFAQ),
|
||||
VectorID: chunkID,
|
||||
Status: enums.StatusOk,
|
||||
CreatedAt: time.Now(),
|
||||
UpdatedAt: time.Now(),
|
||||
}
|
||||
|
||||
collectionName := s.getCollectionName()
|
||||
collectionInfo, err := provider.GetCollection(ctx, collectionName)
|
||||
if err != nil || collectionInfo == nil {
|
||||
if err := provider.CreateCollection(ctx, collectionName, embeddingResult.Dimension); err != nil {
|
||||
return fail(fmt.Errorf("failed to create collection: %w", err))
|
||||
}
|
||||
}
|
||||
|
||||
existingVectorIDs := make([]string, 0, len(existingChunks))
|
||||
for _, chunk := range existingChunks {
|
||||
if strs.IsNotBlank(chunk.VectorID) {
|
||||
existingVectorIDs = append(existingVectorIDs, chunk.VectorID)
|
||||
}
|
||||
}
|
||||
if len(existingVectorIDs) > 0 {
|
||||
if err := provider.DeleteVectors(ctx, collectionName, existingVectorIDs); err != nil {
|
||||
return fail(fmt.Errorf("failed to delete old vectors: %w", err))
|
||||
}
|
||||
}
|
||||
|
||||
if err := provider.UpsertVectors(ctx, collectionName, []vectordb.Vector{{
|
||||
ID: chunkID,
|
||||
Vector: embeddingResult.Vector,
|
||||
Payload: vectordb.ChunkPayload{
|
||||
KnowledgeBaseID: knowledgeBase.ID,
|
||||
FaqID: faq.ID,
|
||||
FaqQuestion: faq.Question,
|
||||
ChunkNo: 0,
|
||||
ChunkType: string(enums.KnowledgeChunkTypeFAQ),
|
||||
Content: content,
|
||||
Title: faq.Question,
|
||||
Provider: string(enums.KnowledgeChunkProviderFAQ),
|
||||
},
|
||||
}}); err != nil {
|
||||
return fail(fmt.Errorf("failed to upsert vectors: %w", err))
|
||||
}
|
||||
|
||||
if err := sqls.WithTransaction(func(ctx *sqls.TxContext) error {
|
||||
if err := ctx.Tx.Where("faq_id = ?", faq.ID).Delete(&models.KnowledgeChunk{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return ctx.Tx.Create(&chunkModel).Error
|
||||
}); err != nil {
|
||||
return fail(fmt.Errorf("failed to save faq chunk: %w", err))
|
||||
}
|
||||
if err := s.markFAQIndexIndexed(faq.ID); err != nil {
|
||||
slog.Error("Failed to mark knowledge faq index as indexed", "faq_id", faq.ID, "error", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *index) RemoveDocumentIndex(ctx context.Context, documentID int64) error {
|
||||
document := repositories.KnowledgeDocumentRepository.Get(sqls.DB(), documentID)
|
||||
if document == nil {
|
||||
return nil
|
||||
}
|
||||
chunks := repositories.KnowledgeChunkRepository.Find(sqls.DB(), sqls.NewCnd().Eq("document_id", documentID))
|
||||
return s.removeDocumentIndexByChunks(ctx, document.KnowledgeBaseID, documentID, chunks)
|
||||
}
|
||||
|
||||
func (s *index) RemoveDocumentIndexFromKnowledgeBase(ctx context.Context, knowledgeBaseID int64, documentID int64) error {
|
||||
chunks := repositories.KnowledgeChunkRepository.Find(sqls.DB(), sqls.NewCnd().Eq("document_id", documentID))
|
||||
return s.removeDocumentIndexByChunks(ctx, knowledgeBaseID, documentID, chunks)
|
||||
}
|
||||
|
||||
func (s *index) RemoveDocumentIndexByChunkModels(ctx context.Context, knowledgeBaseID int64, documentID int64, chunks []models.KnowledgeChunk) error {
|
||||
return s.removeDocumentIndexByChunks(ctx, knowledgeBaseID, documentID, chunks)
|
||||
}
|
||||
|
||||
func (s *index) removeDocumentIndexByChunks(ctx context.Context, knowledgeBaseID int64, documentID int64, chunks []models.KnowledgeChunk) error {
|
||||
if len(chunks) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
collectionName := s.getCollectionName()
|
||||
provider := vectordb.GetProvider()
|
||||
if provider == nil {
|
||||
return fmt.Errorf("vectordb provider not initialized")
|
||||
}
|
||||
|
||||
vectorIDs := make([]string, 0, len(chunks))
|
||||
for _, chunk := range chunks {
|
||||
if chunk.VectorID != "" {
|
||||
vectorIDs = append(vectorIDs, chunk.VectorID)
|
||||
}
|
||||
}
|
||||
|
||||
if len(vectorIDs) > 0 {
|
||||
if err := provider.DeleteVectors(ctx, collectionName, vectorIDs); err != nil {
|
||||
slog.Error("Failed to delete vectors", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
if err := sqls.WithTransaction(func(ctx *sqls.TxContext) error {
|
||||
return ctx.Tx.Where("document_id = ?", documentID).Delete(&models.KnowledgeChunk{}).Error
|
||||
}); err != nil {
|
||||
return fmt.Errorf("failed to delete chunks: %w", err)
|
||||
}
|
||||
|
||||
slog.Info("Document index removed", "document_id", documentID, "chunks_removed", len(chunks))
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *index) RemoveFAQIndex(ctx context.Context, faqID int64) error {
|
||||
faq := repositories.KnowledgeFAQRepository.Get(sqls.DB(), faqID)
|
||||
if faq == nil {
|
||||
return nil
|
||||
}
|
||||
chunks := repositories.KnowledgeChunkRepository.FindByFaqID(sqls.DB(), faqID)
|
||||
return s.removeFAQIndexByChunks(ctx, faq.KnowledgeBaseID, faqID, chunks)
|
||||
}
|
||||
|
||||
func (s *index) RemoveFAQIndexByChunkModels(ctx context.Context, knowledgeBaseID int64, faqID int64, chunks []models.KnowledgeChunk) error {
|
||||
return s.removeFAQIndexByChunks(ctx, knowledgeBaseID, faqID, chunks)
|
||||
}
|
||||
|
||||
func (s *index) removeFAQIndexByChunks(ctx context.Context, knowledgeBaseID int64, faqID int64, chunks []models.KnowledgeChunk) error {
|
||||
if len(chunks) == 0 {
|
||||
return nil
|
||||
}
|
||||
collectionName := s.getCollectionName()
|
||||
provider := vectordb.GetProvider()
|
||||
if provider == nil {
|
||||
return fmt.Errorf("vectordb provider not initialized")
|
||||
}
|
||||
vectorIDs := make([]string, 0, len(chunks))
|
||||
for _, chunk := range chunks {
|
||||
if chunk.VectorID != "" {
|
||||
vectorIDs = append(vectorIDs, chunk.VectorID)
|
||||
}
|
||||
}
|
||||
if len(vectorIDs) > 0 {
|
||||
if err := provider.DeleteVectors(ctx, collectionName, vectorIDs); err != nil {
|
||||
slog.Error("Failed to delete faq vectors", "error", err)
|
||||
}
|
||||
}
|
||||
if err := sqls.WithTransaction(func(ctx *sqls.TxContext) error {
|
||||
return ctx.Tx.Where("faq_id = ?", faqID).Delete(&models.KnowledgeChunk{}).Error
|
||||
}); err != nil {
|
||||
return fmt.Errorf("failed to delete faq chunks: %w", err)
|
||||
}
|
||||
slog.Info("FAQ index removed", "faq_id", faqID, "chunks_removed", len(chunks))
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *index) getCollectionName() string {
|
||||
return knowledgeCollectionName
|
||||
}
|
||||
|
||||
func buildKnowledgeChunkVectorID(knowledgeBaseID int64, documentID int64, chunkNo int) string {
|
||||
raw := fmt.Sprintf("kb:%d:doc:%d:chunk:%d", knowledgeBaseID, documentID, chunkNo)
|
||||
return uuid.NewSHA1(uuid.NameSpaceOID, []byte(raw)).String()
|
||||
}
|
||||
|
||||
func buildKnowledgeFAQChunkVectorID(knowledgeBaseID int64, faqID int64, chunkNo int) string {
|
||||
raw := fmt.Sprintf("kb:%d:faq:%d:chunk:%d", knowledgeBaseID, faqID, chunkNo)
|
||||
return uuid.NewSHA1(uuid.NameSpaceOID, []byte(raw)).String()
|
||||
}
|
||||
|
||||
func buildChunkContentHash(content string) string {
|
||||
sum := sha256.Sum256([]byte(content))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
func firstPositiveInt(values ...int) int {
|
||||
for _, value := range values {
|
||||
if value > 0 {
|
||||
return value
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func firstNonEmptyString(values ...string) string {
|
||||
for _, value := range values {
|
||||
if value != "" {
|
||||
return value
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (s *index) EnsureCollection(ctx context.Context) error {
|
||||
dimension, err := ai.Embedding.GetDimension(ctx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get embedding dimension: %w", err)
|
||||
}
|
||||
|
||||
collectionName := s.getCollectionName()
|
||||
provider := vectordb.GetProvider()
|
||||
if provider == nil {
|
||||
return fmt.Errorf("vectordb provider not initialized")
|
||||
}
|
||||
|
||||
existing, err := provider.GetCollection(ctx, collectionName)
|
||||
if err == nil && existing != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
return provider.CreateCollection(ctx, collectionName, dimension)
|
||||
}
|
||||
|
||||
func (s *index) RebuildKnowledgeBaseIndex(ctx context.Context, knowledgeBaseID int64) error {
|
||||
knowledgeBase := repositories.KnowledgeBaseRepository.Get(sqls.DB(), knowledgeBaseID)
|
||||
if knowledgeBase == nil {
|
||||
return fmt.Errorf("knowledge base not found: %d", knowledgeBaseID)
|
||||
}
|
||||
|
||||
if err := s.resetKnowledgeBaseIndexStorage(ctx, knowledgeBaseID); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
successCount := 0
|
||||
failedCount := 0
|
||||
if knowledgeBase.KnowledgeType == string(enums.KnowledgeBaseTypeFAQ) {
|
||||
faqs := repositories.KnowledgeFAQRepository.Find(sqls.DB(), sqls.NewCnd().
|
||||
Eq("knowledge_base_id", knowledgeBaseID).
|
||||
Where("status != ?", enums.StatusDeleted))
|
||||
if len(faqs) == 0 {
|
||||
slog.Info("No faqs found in knowledge base, nothing to rebuild", "knowledge_base_id", knowledgeBaseID)
|
||||
return nil
|
||||
}
|
||||
slog.Info("Rebuilding faq knowledge base index", "knowledge_base_id", knowledgeBaseID, "faq_count", len(faqs))
|
||||
for _, faq := range faqs {
|
||||
if err := s.IndexFAQByID(ctx, faq.ID); err != nil {
|
||||
slog.Error("Failed to index faq", "faq_id", faq.ID, "error", err)
|
||||
failedCount++
|
||||
} else {
|
||||
successCount++
|
||||
}
|
||||
}
|
||||
} else {
|
||||
documents := repositories.KnowledgeDocumentRepository.Find(sqls.DB(), sqls.NewCnd().
|
||||
Eq("knowledge_base_id", knowledgeBaseID).
|
||||
Where("status != ?", enums.StatusDeleted))
|
||||
if len(documents) == 0 {
|
||||
slog.Info("No documents found in knowledge base, nothing to rebuild", "knowledge_base_id", knowledgeBaseID)
|
||||
return nil
|
||||
}
|
||||
|
||||
documentIDs := make([]int64, 0, len(documents))
|
||||
for _, doc := range documents {
|
||||
documentIDs = append(documentIDs, doc.ID)
|
||||
}
|
||||
if err := s.markKnowledgeBaseDocumentsIndexPending(knowledgeBaseID, documentIDs); err != nil {
|
||||
slog.Error("Failed to mark knowledge base documents index as pending", "knowledge_base_id", knowledgeBaseID, "error", err)
|
||||
}
|
||||
|
||||
slog.Info("Rebuilding knowledge base index", "knowledge_base_id", knowledgeBaseID, "document_count", len(documents))
|
||||
for _, doc := range documents {
|
||||
if err := s.IndexDocumentByID(ctx, doc.ID); err != nil {
|
||||
slog.Error("Failed to index document", "document_id", doc.ID, "error", err)
|
||||
failedCount++
|
||||
} else {
|
||||
successCount++
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
slog.Info("Knowledge base index rebuild completed",
|
||||
"knowledge_base_id", knowledgeBaseID,
|
||||
"success_count", successCount,
|
||||
"failed_count", failedCount)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func buildFAQChunkContent(faq *models.KnowledgeFAQ) string {
|
||||
if faq == nil {
|
||||
return ""
|
||||
}
|
||||
parts := []string{fmt.Sprintf("问题:%s", faq.Question)}
|
||||
var similarQuestions []string
|
||||
if faq.SimilarQuestions != "" {
|
||||
_ = json.Unmarshal([]byte(faq.SimilarQuestions), &similarQuestions)
|
||||
}
|
||||
if len(similarQuestions) > 0 {
|
||||
parts = append(parts, fmt.Sprintf("相似问:%s", joinSimilarQuestions(similarQuestions)))
|
||||
}
|
||||
parts = append(parts, fmt.Sprintf("回答:%s", faq.Answer))
|
||||
content := ""
|
||||
for _, part := range parts {
|
||||
if part == "" {
|
||||
continue
|
||||
}
|
||||
if content != "" {
|
||||
content += "\n"
|
||||
}
|
||||
content += part
|
||||
}
|
||||
return content
|
||||
}
|
||||
|
||||
func joinSimilarQuestions(items []string) string {
|
||||
result := ""
|
||||
for _, item := range items {
|
||||
if item == "" {
|
||||
continue
|
||||
}
|
||||
if result != "" {
|
||||
result += ";"
|
||||
}
|
||||
result += item
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func (s *index) markDocumentIndexPending(documentID int64) error {
|
||||
return repositories.KnowledgeDocumentRepository.Updates(sqls.DB(), documentID, map[string]any{
|
||||
"index_status": enums.KnowledgeDocumentIndexStatusPending,
|
||||
"indexed_at": nil,
|
||||
"index_error": "",
|
||||
"updated_at": time.Now(),
|
||||
})
|
||||
}
|
||||
|
||||
func (s *index) markDocumentIndexIndexed(documentID int64) error {
|
||||
now := time.Now()
|
||||
return repositories.KnowledgeDocumentRepository.Updates(sqls.DB(), documentID, map[string]any{
|
||||
"index_status": enums.KnowledgeDocumentIndexStatusIndexed,
|
||||
"indexed_at": &now,
|
||||
"index_error": "",
|
||||
"updated_at": now,
|
||||
})
|
||||
}
|
||||
|
||||
func (s *index) markDocumentIndexFailed(documentID int64, err error) error {
|
||||
return repositories.KnowledgeDocumentRepository.Updates(sqls.DB(), documentID, map[string]any{
|
||||
"index_status": enums.KnowledgeDocumentIndexStatusFailed,
|
||||
"index_error": truncateIndexError(err),
|
||||
"updated_at": time.Now(),
|
||||
})
|
||||
}
|
||||
|
||||
func (s *index) markKnowledgeBaseDocumentsIndexPending(knowledgeBaseID int64, documentIDs []int64) error {
|
||||
if len(documentIDs) == 0 {
|
||||
return nil
|
||||
}
|
||||
return sqls.DB().Model(&models.KnowledgeDocument{}).
|
||||
Where("knowledge_base_id = ?", knowledgeBaseID).
|
||||
Where("id IN ?", documentIDs).
|
||||
Updates(map[string]any{
|
||||
"index_status": enums.KnowledgeDocumentIndexStatusPending,
|
||||
"indexed_at": nil,
|
||||
"index_error": "",
|
||||
"updated_at": time.Now(),
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (s *index) markFAQIndexPending(faqID int64) error {
|
||||
return repositories.KnowledgeFAQRepository.Updates(sqls.DB(), faqID, map[string]any{
|
||||
"index_status": enums.KnowledgeDocumentIndexStatusPending,
|
||||
"indexed_at": nil,
|
||||
"index_error": "",
|
||||
"updated_at": time.Now(),
|
||||
})
|
||||
}
|
||||
|
||||
func (s *index) markFAQIndexIndexed(faqID int64) error {
|
||||
now := time.Now()
|
||||
return repositories.KnowledgeFAQRepository.Updates(sqls.DB(), faqID, map[string]any{
|
||||
"index_status": enums.KnowledgeDocumentIndexStatusIndexed,
|
||||
"indexed_at": &now,
|
||||
"index_error": "",
|
||||
"updated_at": now,
|
||||
})
|
||||
}
|
||||
|
||||
func (s *index) markFAQIndexFailed(faqID int64, err error) error {
|
||||
return repositories.KnowledgeFAQRepository.Updates(sqls.DB(), faqID, map[string]any{
|
||||
"index_status": enums.KnowledgeDocumentIndexStatusFailed,
|
||||
"index_error": truncateIndexError(err),
|
||||
"updated_at": time.Now(),
|
||||
})
|
||||
}
|
||||
|
||||
func truncateIndexError(err error) string {
|
||||
if err == nil {
|
||||
return ""
|
||||
}
|
||||
message := err.Error()
|
||||
if len(message) <= 1000 {
|
||||
return message
|
||||
}
|
||||
return message[:1000]
|
||||
}
|
||||
|
||||
func (s *index) resetKnowledgeBaseIndexStorage(ctx context.Context, knowledgeBaseID int64) error {
|
||||
collectionName := s.getCollectionName()
|
||||
provider := vectordb.GetProvider()
|
||||
if provider == nil {
|
||||
return fmt.Errorf("vectordb provider not initialized")
|
||||
}
|
||||
|
||||
chunks := repositories.KnowledgeChunkRepository.Find(sqls.DB(), sqls.NewCnd().Eq("knowledge_base_id", knowledgeBaseID))
|
||||
vectorIDs := make([]string, 0, len(chunks))
|
||||
for _, chunk := range chunks {
|
||||
if strs.IsNotBlank(chunk.VectorID) {
|
||||
vectorIDs = append(vectorIDs, chunk.VectorID)
|
||||
}
|
||||
}
|
||||
if len(vectorIDs) > 0 {
|
||||
if err := provider.DeleteVectors(ctx, collectionName, vectorIDs); err != nil {
|
||||
return fmt.Errorf("failed to delete vectors for knowledge base %d before rebuild: %w", knowledgeBaseID, err)
|
||||
}
|
||||
}
|
||||
|
||||
if err := sqls.WithTransaction(func(ctx *sqls.TxContext) error {
|
||||
return ctx.Tx.Where("knowledge_base_id = ?", knowledgeBaseID).Delete(&models.KnowledgeChunk{}).Error
|
||||
}); err != nil {
|
||||
return fmt.Errorf("failed to clear chunks before rebuild: %w", err)
|
||||
}
|
||||
|
||||
slog.Info("Knowledge base index storage reset",
|
||||
"knowledge_base_id", knowledgeBaseID,
|
||||
"collection", collectionName)
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,140 @@
|
||||
package rag
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"sort"
|
||||
"time"
|
||||
|
||||
"cs-agent/internal/ai"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
)
|
||||
|
||||
type rerank struct{}
|
||||
|
||||
var Rerank = &rerank{}
|
||||
|
||||
func (s *rerank) Rerank(ctx context.Context, query string, documents []string, topN int) ([]RerankResult, error) {
|
||||
if len(documents) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
if topN <= 0 {
|
||||
topN = len(documents)
|
||||
}
|
||||
|
||||
results, err := s.callRerankAPI(ctx, query, documents, topN)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return results, nil
|
||||
}
|
||||
|
||||
func (s *rerank) callRerankAPI(ctx context.Context, query string, documents []string, topN int) ([]RerankResult, error) {
|
||||
config, err := ai.GetEnabledAIConfig(enums.AIModelTypeRerank)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
reqBody := RerankRequest{
|
||||
Model: config.ModelName,
|
||||
Query: query,
|
||||
Documents: documents,
|
||||
TopN: topN,
|
||||
}
|
||||
|
||||
jsonBody, err := json.Marshal(reqBody)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal request: %w", err)
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", config.BaseURL+"/v1/rerank", bytes.NewBuffer(jsonBody))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create request: %w", err)
|
||||
}
|
||||
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+config.APIKey)
|
||||
|
||||
client := &http.Client{
|
||||
Timeout: time.Duration(config.TimeoutMS) * time.Millisecond,
|
||||
}
|
||||
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to call rerank API: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read response: %w", err)
|
||||
}
|
||||
|
||||
var rerankResp RerankResponse
|
||||
if err := json.Unmarshal(body, &rerankResp); err != nil {
|
||||
return nil, fmt.Errorf("failed to unmarshal response: %w", err)
|
||||
}
|
||||
|
||||
results := make([]RerankResult, 0, len(rerankResp.Results))
|
||||
for _, r := range rerankResp.Results {
|
||||
results = append(results, RerankResult{
|
||||
Index: r.Index,
|
||||
RelevanceScore: r.RelevanceScore,
|
||||
})
|
||||
}
|
||||
|
||||
return results, nil
|
||||
}
|
||||
|
||||
func (s *rerank) RerankResults(ctx context.Context, query string, results []RetrieveResult, topN int) ([]RetrieveResult, error) {
|
||||
if len(results) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
if topN <= 0 {
|
||||
topN = len(results)
|
||||
}
|
||||
|
||||
documents := make([]string, 0, len(results))
|
||||
for _, r := range results {
|
||||
documents = append(documents, r.Content)
|
||||
}
|
||||
|
||||
rerankResults, err := s.Rerank(ctx, query, documents, topN)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
rerankedResults := make([]RetrieveResult, 0, len(rerankResults))
|
||||
for _, rr := range rerankResults {
|
||||
if rr.Index < len(results) {
|
||||
result := results[rr.Index]
|
||||
result.Score = float32(rr.RelevanceScore)
|
||||
rerankedResults = append(rerankedResults, result)
|
||||
}
|
||||
}
|
||||
|
||||
return rerankedResults, nil
|
||||
}
|
||||
|
||||
func (s *rerank) SimpleRerank(query string, results []RetrieveResult, topN int) []RetrieveResult {
|
||||
if len(results) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
sort.Slice(results, func(i, j int) bool {
|
||||
return results[i].Score > results[j].Score
|
||||
})
|
||||
|
||||
if topN > 0 && len(results) > topN {
|
||||
return results[:topN]
|
||||
}
|
||||
|
||||
return results
|
||||
}
|
||||
@@ -0,0 +1,507 @@
|
||||
package rag
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"cs-agent/internal/models"
|
||||
"github.com/mlogclub/simple/sqls"
|
||||
|
||||
"cs-agent/internal/ai"
|
||||
"cs-agent/internal/ai/rag/vectordb"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"cs-agent/internal/repositories"
|
||||
)
|
||||
|
||||
type retrieve struct {
|
||||
}
|
||||
|
||||
var Retrieve = &retrieve{}
|
||||
|
||||
func (s *retrieve) Retrieve(ctx context.Context, req RetrieveRequest) ([]RetrieveResult, error) {
|
||||
results, _, err := s.RetrieveWithTrace(ctx, req)
|
||||
return results, err
|
||||
}
|
||||
|
||||
type RetrieveTrace struct {
|
||||
EmbeddingMs int64
|
||||
VectorSearchMs int64
|
||||
HydrateMs int64
|
||||
}
|
||||
|
||||
func (s *retrieve) RetrieveWithTrace(ctx context.Context, req RetrieveRequest) ([]RetrieveResult, *RetrieveTrace, error) {
|
||||
trace := &RetrieveTrace{}
|
||||
if req.Query == "" {
|
||||
return nil, trace, nil
|
||||
}
|
||||
knowledgeBaseIDs := normalizeKnowledgeBaseIDs(req.KnowledgeBaseIDs)
|
||||
if len(knowledgeBaseIDs) == 0 {
|
||||
return nil, trace, nil
|
||||
}
|
||||
|
||||
retrievableKnowledgeBases := s.loadRetrievableKnowledgeBases(knowledgeBaseIDs)
|
||||
if len(retrievableKnowledgeBases) == 0 {
|
||||
slog.Info("Skip retrieve for non-enabled knowledge bases",
|
||||
"knowledge_base_ids", fmt.Sprint(knowledgeBaseIDs))
|
||||
return nil, trace, nil
|
||||
}
|
||||
|
||||
embeddingStartedAt := time.Now()
|
||||
embeddingResult, err := ai.Embedding.GenerateEmbedding(ctx, req.Query)
|
||||
trace.EmbeddingMs = time.Since(embeddingStartedAt).Milliseconds()
|
||||
if err != nil {
|
||||
return nil, trace, fmt.Errorf("failed to generate query embedding: %w", err)
|
||||
}
|
||||
|
||||
collectionName := knowledgeCollectionName
|
||||
provider := vectordb.GetProvider()
|
||||
if provider == nil {
|
||||
return nil, trace, fmt.Errorf("vectordb provider not initialized")
|
||||
}
|
||||
|
||||
searchResults := make([]vectordb.SearchResult, 0)
|
||||
vectorSearchStartedAt := time.Now()
|
||||
for _, knowledgeBase := range retrievableKnowledgeBases {
|
||||
topK, scoreThreshold := resolveKnowledgeBaseSearchOptions(req, &knowledgeBase)
|
||||
kbResults, searchErr := provider.Search(ctx, &vectordb.SearchRequest{
|
||||
CollectionName: collectionName,
|
||||
Vector: embeddingResult.Vector,
|
||||
TopK: topK,
|
||||
ScoreThreshold: scoreThreshold,
|
||||
Filter: &vectordb.SearchFilter{
|
||||
KnowledgeBaseIDs: []int64{knowledgeBase.ID},
|
||||
},
|
||||
})
|
||||
if searchErr != nil {
|
||||
slog.Error("Failed to search vectors",
|
||||
"knowledge_base_id", knowledgeBase.ID,
|
||||
"error", searchErr)
|
||||
trace.VectorSearchMs = time.Since(vectorSearchStartedAt).Milliseconds()
|
||||
return nil, trace, fmt.Errorf("failed to search vectors: %w", searchErr)
|
||||
}
|
||||
if len(kbResults) == 0 && scoreThreshold > 0 {
|
||||
s.logEmptySearchDiagnostics(ctx, provider, collectionName, embeddingResult.Vector, topK, scoreThreshold, []int64{knowledgeBase.ID}, req)
|
||||
}
|
||||
searchResults = append(searchResults, kbResults...)
|
||||
}
|
||||
trace.VectorSearchMs = time.Since(vectorSearchStartedAt).Milliseconds()
|
||||
|
||||
if len(searchResults) == 0 {
|
||||
return nil, trace, nil
|
||||
}
|
||||
sort.SliceStable(searchResults, func(i, j int) bool {
|
||||
if searchResults[i].Score == searchResults[j].Score {
|
||||
return searchResults[i].ID < searchResults[j].ID
|
||||
}
|
||||
return searchResults[i].Score > searchResults[j].Score
|
||||
})
|
||||
|
||||
results := make([]RetrieveResult, 0, len(searchResults))
|
||||
hydrateStartedAt := time.Now()
|
||||
vectorIDs := make([]string, 0, len(searchResults))
|
||||
for _, sr := range searchResults {
|
||||
if strings.TrimSpace(sr.ID) == "" {
|
||||
continue
|
||||
}
|
||||
vectorIDs = append(vectorIDs, sr.ID)
|
||||
}
|
||||
chunks := repositories.KnowledgeChunkRepository.FindByVectorIDs(sqls.DB(), vectorIDs)
|
||||
chunkByVectorID := make(map[string]*models.KnowledgeChunk, len(chunks))
|
||||
documentIDs := make([]int64, 0)
|
||||
faqIDs := make([]int64, 0)
|
||||
documentSeen := make(map[int64]struct{})
|
||||
faqSeen := make(map[int64]struct{})
|
||||
for i := range chunks {
|
||||
chunk := &chunks[i]
|
||||
chunkByVectorID[chunk.VectorID] = chunk
|
||||
if chunk.DocumentID > 0 {
|
||||
if _, ok := documentSeen[chunk.DocumentID]; !ok {
|
||||
documentSeen[chunk.DocumentID] = struct{}{}
|
||||
documentIDs = append(documentIDs, chunk.DocumentID)
|
||||
}
|
||||
}
|
||||
if chunk.FaqID > 0 {
|
||||
if _, ok := faqSeen[chunk.FaqID]; !ok {
|
||||
faqSeen[chunk.FaqID] = struct{}{}
|
||||
faqIDs = append(faqIDs, chunk.FaqID)
|
||||
}
|
||||
}
|
||||
}
|
||||
documents := repositories.KnowledgeDocumentRepository.FindByIDs(sqls.DB(), documentIDs)
|
||||
documentByID := make(map[int64]*models.KnowledgeDocument, len(documents))
|
||||
for i := range documents {
|
||||
document := &documents[i]
|
||||
documentByID[document.ID] = document
|
||||
}
|
||||
faqs := repositories.KnowledgeFAQRepository.FindByIDs(sqls.DB(), faqIDs)
|
||||
faqByID := make(map[int64]*models.KnowledgeFAQ, len(faqs))
|
||||
for i := range faqs {
|
||||
faq := &faqs[i]
|
||||
faqByID[faq.ID] = faq
|
||||
}
|
||||
for _, sr := range searchResults {
|
||||
chunk := chunkByVectorID[sr.ID]
|
||||
if chunk == nil || chunk.Status != enums.StatusOk {
|
||||
continue
|
||||
}
|
||||
|
||||
documentTitle := ""
|
||||
faqQuestion := ""
|
||||
if chunk.DocumentID > 0 {
|
||||
document := documentByID[chunk.DocumentID]
|
||||
if document == nil || document.Status != enums.StatusOk {
|
||||
continue
|
||||
}
|
||||
documentTitle = document.Title
|
||||
}
|
||||
if chunk.FaqID > 0 {
|
||||
faq := faqByID[chunk.FaqID]
|
||||
if faq == nil || faq.Status != enums.StatusOk {
|
||||
continue
|
||||
}
|
||||
faqQuestion = faq.Question
|
||||
}
|
||||
|
||||
results = append(results, RetrieveResult{
|
||||
KnowledgeBaseID: chunk.KnowledgeBaseID,
|
||||
ChunkID: chunk.ID,
|
||||
DocumentID: chunk.DocumentID,
|
||||
DocumentTitle: documentTitle,
|
||||
FaqID: chunk.FaqID,
|
||||
FaqQuestion: faqQuestion,
|
||||
ChunkNo: chunk.ChunkNo,
|
||||
Title: chunk.Title,
|
||||
SectionPath: chunk.SectionPath,
|
||||
Content: chunk.Content,
|
||||
Score: sr.Score,
|
||||
ChunkType: extractChunkType(sr.Payload),
|
||||
})
|
||||
}
|
||||
trace.HydrateMs = time.Since(hydrateStartedAt).Milliseconds()
|
||||
|
||||
return results, trace, nil
|
||||
}
|
||||
|
||||
func extractChunkType(payload vectordb.ChunkPayload) string {
|
||||
if payload.ChunkType != "" {
|
||||
return payload.ChunkType
|
||||
}
|
||||
return string(enums.KnowledgeChunkTypeText)
|
||||
}
|
||||
|
||||
func (s *retrieve) logEmptySearchDiagnostics(ctx context.Context, provider vectordb.Provider, collectionName string, vector []float32, topK int, scoreThreshold float32, knowledgeBaseIDs []int64, req RetrieveRequest) {
|
||||
rawResults, err := provider.Search(ctx, &vectordb.SearchRequest{
|
||||
CollectionName: collectionName,
|
||||
Vector: vector,
|
||||
TopK: topK,
|
||||
ScoreThreshold: 0,
|
||||
Filter: &vectordb.SearchFilter{
|
||||
KnowledgeBaseIDs: knowledgeBaseIDs,
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
slog.Warn("Knowledge retrieve diagnostics failed",
|
||||
"knowledge_base_ids", fmt.Sprint(knowledgeBaseIDs),
|
||||
"collection", collectionName,
|
||||
"query", truncateForLog(req.Query, 80),
|
||||
"score_threshold", scoreThreshold,
|
||||
"error", err)
|
||||
return
|
||||
}
|
||||
if len(rawResults) == 0 {
|
||||
slog.Info("Knowledge retrieve returned no candidates even without threshold",
|
||||
"knowledge_base_ids", fmt.Sprint(knowledgeBaseIDs),
|
||||
"collection", collectionName,
|
||||
"query", truncateForLog(req.Query, 80),
|
||||
"score_threshold", scoreThreshold)
|
||||
return
|
||||
}
|
||||
|
||||
candidates := make([]string, 0, len(rawResults))
|
||||
for _, item := range rawResults {
|
||||
candidates = append(candidates, fmt.Sprintf("%s:%.4f", item.ID, item.Score))
|
||||
}
|
||||
|
||||
slog.Info("Knowledge retrieve filtered all candidates by score threshold",
|
||||
"knowledge_base_ids", fmt.Sprint(knowledgeBaseIDs),
|
||||
"collection", collectionName,
|
||||
"query", truncateForLog(req.Query, 80),
|
||||
"score_threshold", scoreThreshold,
|
||||
"top_candidates", strings.Join(candidates, ","))
|
||||
}
|
||||
|
||||
func truncateForLog(text string, limit int) string {
|
||||
if limit <= 0 {
|
||||
return ""
|
||||
}
|
||||
runes := []rune(strings.TrimSpace(text))
|
||||
if len(runes) <= limit {
|
||||
return string(runes)
|
||||
}
|
||||
return string(runes[:limit]) + "..."
|
||||
}
|
||||
|
||||
func (s *retrieve) RetrieveWithRerank(ctx context.Context, req RetrieveRequest, rerankLimit int) ([]RetrieveResult, error) {
|
||||
results, err := s.Retrieve(ctx, req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if len(results) <= rerankLimit {
|
||||
return results, nil
|
||||
}
|
||||
|
||||
rerankedResults, err := s.rerank(ctx, req.Query, results, rerankLimit)
|
||||
if err != nil {
|
||||
slog.Warn("Rerank failed, returning original results", "error", err)
|
||||
if len(results) > rerankLimit {
|
||||
return results[:rerankLimit], nil
|
||||
}
|
||||
return results, nil
|
||||
}
|
||||
|
||||
return rerankedResults, nil
|
||||
}
|
||||
|
||||
func (s *retrieve) rerank(ctx context.Context, query string, results []RetrieveResult, limit int) ([]RetrieveResult, error) {
|
||||
return Rerank.RerankResults(ctx, query, results, limit)
|
||||
}
|
||||
|
||||
func (s *retrieve) SelectContextResults(results []RetrieveResult, maxTokens int) []RetrieveResult {
|
||||
if len(results) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
normalizedResults := normalizeContextResults(results)
|
||||
selected := make([]RetrieveResult, 0, len(normalizedResults))
|
||||
totalTokens := 0
|
||||
documentUsage := make(map[int64]int)
|
||||
|
||||
for _, item := range normalizedResults {
|
||||
if documentUsage[item.DocumentID] >= 2 {
|
||||
continue
|
||||
}
|
||||
chunkText := buildContextChunkText(item)
|
||||
estimatedTokens := len(chunkText) / 2
|
||||
if totalTokens+estimatedTokens > maxTokens {
|
||||
break
|
||||
}
|
||||
selected = append(selected, item)
|
||||
totalTokens += estimatedTokens
|
||||
documentUsage[item.DocumentID]++
|
||||
}
|
||||
return selected
|
||||
}
|
||||
|
||||
func (s *retrieve) BuildContext(ctx context.Context, results []RetrieveResult, maxTokens int) string {
|
||||
if len(results) == 0 {
|
||||
return ""
|
||||
}
|
||||
|
||||
normalizedResults := s.SelectContextResults(results, maxTokens)
|
||||
context := ""
|
||||
for _, r := range normalizedResults {
|
||||
chunkText := buildContextChunkText(r)
|
||||
context += chunkText
|
||||
}
|
||||
|
||||
return context
|
||||
}
|
||||
|
||||
func normalizeContextResults(results []RetrieveResult) []RetrieveResult {
|
||||
if len(results) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
merged := mergeAdjacentResults(results)
|
||||
return dedupeSectionResults(merged)
|
||||
}
|
||||
|
||||
func dedupeSectionResults(results []RetrieveResult) []RetrieveResult {
|
||||
seen := make(map[string]struct{})
|
||||
deduped := make([]RetrieveResult, 0, len(results))
|
||||
for _, item := range results {
|
||||
key := buildSectionKey(item)
|
||||
if _, ok := seen[key]; ok {
|
||||
continue
|
||||
}
|
||||
seen[key] = struct{}{}
|
||||
deduped = append(deduped, item)
|
||||
}
|
||||
return deduped
|
||||
}
|
||||
|
||||
func mergeAdjacentResults(results []RetrieveResult) []RetrieveResult {
|
||||
if len(results) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
merged := make([]RetrieveResult, 0, len(results))
|
||||
for _, item := range results {
|
||||
if len(merged) == 0 {
|
||||
merged = append(merged, item)
|
||||
continue
|
||||
}
|
||||
|
||||
last := &merged[len(merged)-1]
|
||||
if canMergeContextResult(*last, item) {
|
||||
last.Content = strings.TrimSpace(last.Content + "\n" + item.Content)
|
||||
if item.Score > last.Score {
|
||||
last.Score = item.Score
|
||||
}
|
||||
continue
|
||||
}
|
||||
merged = append(merged, item)
|
||||
}
|
||||
return merged
|
||||
}
|
||||
|
||||
func canMergeContextResult(left, right RetrieveResult) bool {
|
||||
if left.FaqID > 0 || right.FaqID > 0 {
|
||||
return false
|
||||
}
|
||||
if left.DocumentID != right.DocumentID {
|
||||
return false
|
||||
}
|
||||
if left.SectionPath == "" || right.SectionPath == "" {
|
||||
return false
|
||||
}
|
||||
if left.SectionPath != right.SectionPath {
|
||||
return false
|
||||
}
|
||||
return right.ChunkNo == left.ChunkNo+1
|
||||
}
|
||||
|
||||
func buildSectionKey(item RetrieveResult) string {
|
||||
if item.FaqID > 0 {
|
||||
return fmt.Sprintf("faq:%d", item.FaqID)
|
||||
}
|
||||
sectionPath := strings.TrimSpace(item.SectionPath)
|
||||
if sectionPath != "" {
|
||||
return fmt.Sprintf("%d|%s", item.DocumentID, sectionPath)
|
||||
}
|
||||
title := strings.TrimSpace(item.Title)
|
||||
if title != "" {
|
||||
return fmt.Sprintf("%d|%s", item.DocumentID, title)
|
||||
}
|
||||
return fmt.Sprintf("%d|chunk:%d", item.DocumentID, item.ChunkNo)
|
||||
}
|
||||
|
||||
func buildContextChunkText(item RetrieveResult) string {
|
||||
if item.FaqID > 0 {
|
||||
title := strings.TrimSpace(item.FaqQuestion)
|
||||
if title == "" {
|
||||
title = strings.TrimSpace(item.Title)
|
||||
}
|
||||
if title == "" {
|
||||
title = fmt.Sprintf("FAQ#%d", item.FaqID)
|
||||
}
|
||||
return fmt.Sprintf("【FAQ:%s】\n%s\n\n", title, item.Content)
|
||||
}
|
||||
title := strings.TrimSpace(item.DocumentTitle)
|
||||
if title == "" {
|
||||
title = fmt.Sprintf("文档#%d", item.DocumentID)
|
||||
}
|
||||
if item.SectionPath != "" {
|
||||
return fmt.Sprintf("【文档:%s|章节:%s】\n%s\n\n", title, item.SectionPath, item.Content)
|
||||
}
|
||||
if item.Title != "" {
|
||||
return fmt.Sprintf("【文档:%s|标题:%s】\n%s\n\n", title, item.Title, item.Content)
|
||||
}
|
||||
return fmt.Sprintf("【文档:%s】\n%s\n\n", title, item.Content)
|
||||
}
|
||||
|
||||
func (s *retrieve) GetKnowledgeBaseStats(ctx context.Context, knowledgeBaseID int64) (*KnowledgeBaseStats, error) {
|
||||
knowledgeBase := repositories.KnowledgeBaseRepository.Get(sqls.DB(), knowledgeBaseID)
|
||||
if knowledgeBase == nil {
|
||||
return nil, fmt.Errorf("knowledge base not found")
|
||||
}
|
||||
|
||||
documentCount := repositories.KnowledgeDocumentRepository.CountByKnowledgeBaseID(sqls.DB(), knowledgeBaseID)
|
||||
chunkCount := repositories.KnowledgeChunkRepository.CountByKnowledgeBaseID(sqls.DB(), knowledgeBaseID)
|
||||
|
||||
publishedCount := repositories.KnowledgeDocumentRepository.Count(sqls.DB(), sqls.NewCnd().
|
||||
Eq("knowledge_base_id", knowledgeBaseID).
|
||||
Eq("status", enums.StatusOk))
|
||||
|
||||
return &KnowledgeBaseStats{
|
||||
KnowledgeBaseID: knowledgeBaseID,
|
||||
DocumentCount: documentCount,
|
||||
PublishedCount: publishedCount,
|
||||
ChunkCount: chunkCount,
|
||||
VectorCount: int(chunkCount),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func normalizeKnowledgeBaseIDs(ids []int64) []int64 {
|
||||
if len(ids) == 0 {
|
||||
return nil
|
||||
}
|
||||
seen := make(map[int64]struct{}, len(ids))
|
||||
normalized := make([]int64, 0, len(ids))
|
||||
for _, id := range ids {
|
||||
if id <= 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[id]; ok {
|
||||
continue
|
||||
}
|
||||
seen[id] = struct{}{}
|
||||
normalized = append(normalized, id)
|
||||
}
|
||||
return normalized
|
||||
}
|
||||
|
||||
func resolveKnowledgeBaseSearchOptions(req RetrieveRequest, knowledgeBase *models.KnowledgeBase) (int, float32) {
|
||||
topK := req.TopK
|
||||
if topK <= 0 && knowledgeBase != nil && knowledgeBase.DefaultTopK > 0 {
|
||||
topK = knowledgeBase.DefaultTopK
|
||||
}
|
||||
if topK <= 0 {
|
||||
topK = 8
|
||||
}
|
||||
|
||||
scoreThreshold := float32(req.ScoreThreshold)
|
||||
if scoreThreshold <= 0 && knowledgeBase != nil && knowledgeBase.DefaultScoreThreshold > 0 {
|
||||
scoreThreshold = float32(knowledgeBase.DefaultScoreThreshold)
|
||||
}
|
||||
if scoreThreshold <= 0 {
|
||||
scoreThreshold = 0.3
|
||||
}
|
||||
return topK, scoreThreshold
|
||||
}
|
||||
|
||||
func (s *retrieve) loadRetrievableKnowledgeBases(ids []int64) []models.KnowledgeBase {
|
||||
if len(ids) == 0 {
|
||||
return nil
|
||||
}
|
||||
items := repositories.KnowledgeBaseRepository.Find(sqls.DB(), sqls.NewCnd().In("id", ids))
|
||||
if len(items) == 0 {
|
||||
return nil
|
||||
}
|
||||
allowed := make(map[int64]models.KnowledgeBase, len(items))
|
||||
for _, item := range items {
|
||||
if item.Status == enums.StatusOk {
|
||||
allowed[item.ID] = item
|
||||
}
|
||||
}
|
||||
filtered := make([]models.KnowledgeBase, 0, len(ids))
|
||||
for _, id := range ids {
|
||||
if item, ok := allowed[id]; ok {
|
||||
filtered = append(filtered, item)
|
||||
}
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
type KnowledgeBaseStats struct {
|
||||
KnowledgeBaseID int64 `json:"knowledgeBaseId"`
|
||||
DocumentCount int64 `json:"documentCount"`
|
||||
PublishedCount int64 `json:"publishedCount"`
|
||||
ChunkCount int64 `json:"chunkCount"`
|
||||
VectorCount int `json:"vectorCount"`
|
||||
}
|
||||
@@ -0,0 +1,300 @@
|
||||
package rag
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/dto"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/mlogclub/simple/sqls"
|
||||
)
|
||||
|
||||
var RetrieveLog = &retrieveLog{}
|
||||
|
||||
type retrieveLog struct {
|
||||
}
|
||||
|
||||
type CreateRetrieveLogRequest struct {
|
||||
KnowledgeBaseID int64
|
||||
Channel string
|
||||
Scene string
|
||||
SessionID string
|
||||
ConversationID int64
|
||||
Question string
|
||||
RewriteQuestion string
|
||||
Answer string
|
||||
AnswerStatus int
|
||||
ChunkProvider string
|
||||
ChunkTargetTokens int
|
||||
ChunkMaxTokens int
|
||||
ChunkOverlapTokens int
|
||||
RerankEnabled bool
|
||||
RerankLimit int
|
||||
Hits []response.KnowledgeSearchResult
|
||||
UsedHits []response.KnowledgeSearchResult
|
||||
Citations []response.KnowledgeCitation
|
||||
LatencyMs int64
|
||||
RetrieveMs int64
|
||||
GenerateMs int64
|
||||
PromptTokens int
|
||||
CompletionTokens int
|
||||
ModelName string
|
||||
}
|
||||
|
||||
type retrieveTraceData struct {
|
||||
Retrieve retrieveTraceRetrieve `json:"retrieve"`
|
||||
ChunkConfig retrieveTraceChunkConfig `json:"chunkConfig"`
|
||||
Context retrieveTraceContext `json:"context"`
|
||||
Citations []retrieveTraceCitation `json:"citations"`
|
||||
}
|
||||
|
||||
type retrieveTraceRetrieve struct {
|
||||
Provider string `json:"provider"`
|
||||
RerankEnabled bool `json:"rerankEnabled"`
|
||||
RerankLimit int `json:"rerankLimit"`
|
||||
RawHitCount int `json:"rawHitCount"`
|
||||
ContextHitCount int `json:"contextHitCount"`
|
||||
CitationCount int `json:"citationCount"`
|
||||
}
|
||||
|
||||
type retrieveTraceChunkConfig struct {
|
||||
Provider string `json:"provider"`
|
||||
TargetTokens int `json:"targetTokens"`
|
||||
MaxTokens int `json:"maxTokens"`
|
||||
OverlapTokens int `json:"overlapTokens"`
|
||||
}
|
||||
|
||||
type retrieveTraceContext struct {
|
||||
KnowledgeBaseIDs []int64 `json:"knowledgeBaseIds"`
|
||||
DocumentIDs []int64 `json:"documentIds"`
|
||||
SectionPaths []string `json:"sectionPaths"`
|
||||
UsedChunkKeys []string `json:"usedChunkKeys"`
|
||||
}
|
||||
|
||||
type retrieveTraceCitation struct {
|
||||
DocumentID int64 `json:"documentId"`
|
||||
ChunkNo int `json:"chunkNo"`
|
||||
SectionPath string `json:"sectionPath"`
|
||||
}
|
||||
|
||||
func (s *retrieveLog) FindHitsByRetrieveLogID(retrieveLogID int64) []models.KnowledgeRetrieveHit {
|
||||
if retrieveLogID <= 0 {
|
||||
return nil
|
||||
}
|
||||
var list []models.KnowledgeRetrieveHit
|
||||
sqls.DB().Where("retrieve_log_id = ?", retrieveLogID).Order("rank_no asc, id asc").Find(&list)
|
||||
return list
|
||||
}
|
||||
|
||||
func (s *retrieveLog) CreateRetrieveLog(req *CreateRetrieveLogRequest, _ *dto.AuthPrincipal) (*models.KnowledgeRetrieveLog, error) {
|
||||
if req == nil {
|
||||
return nil, fmt.Errorf("retrieve log request is nil")
|
||||
}
|
||||
now := time.Now()
|
||||
topScore := 0.0
|
||||
if len(req.Hits) > 0 {
|
||||
topScore = req.Hits[0].Score
|
||||
}
|
||||
traceData := buildRetrieveTraceData(req)
|
||||
|
||||
log := &models.KnowledgeRetrieveLog{
|
||||
KnowledgeBaseID: req.KnowledgeBaseID,
|
||||
Channel: req.Channel,
|
||||
Scene: req.Scene,
|
||||
SessionID: req.SessionID,
|
||||
ConversationID: req.ConversationID,
|
||||
RequestID: uuid.New().String(),
|
||||
Question: req.Question,
|
||||
RewriteQuestion: req.RewriteQuestion,
|
||||
Answer: req.Answer,
|
||||
AnswerStatus: req.AnswerStatus,
|
||||
HitCount: len(req.Hits),
|
||||
TopScore: topScore,
|
||||
ChunkProvider: req.ChunkProvider,
|
||||
ChunkTargetTokens: req.ChunkTargetTokens,
|
||||
ChunkMaxTokens: req.ChunkMaxTokens,
|
||||
ChunkOverlapTokens: req.ChunkOverlapTokens,
|
||||
RerankEnabled: req.RerankEnabled,
|
||||
RerankLimit: req.RerankLimit,
|
||||
CitationCount: len(req.Citations),
|
||||
UsedChunkCount: len(req.UsedHits),
|
||||
LatencyMs: req.LatencyMs,
|
||||
RetrieveMs: req.RetrieveMs,
|
||||
GenerateMs: req.GenerateMs,
|
||||
PromptTokens: req.PromptTokens,
|
||||
CompletionTokens: req.CompletionTokens,
|
||||
ModelName: req.ModelName,
|
||||
TraceData: traceData,
|
||||
CreatedAt: now,
|
||||
}
|
||||
|
||||
usedHitKeys := make(map[string]struct{}, len(req.UsedHits))
|
||||
for _, item := range req.UsedHits {
|
||||
usedHitKeys[buildKnowledgeSearchResultKey(item)] = struct{}{}
|
||||
}
|
||||
citationKeys := make(map[string]struct{}, len(req.Citations))
|
||||
for _, item := range req.Citations {
|
||||
citationKeys[buildKnowledgeCitationKey(item)] = struct{}{}
|
||||
}
|
||||
|
||||
if err := sqls.WithTransaction(func(ctx *sqls.TxContext) error {
|
||||
if err := ctx.Tx.Create(log).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
for i, hit := range req.Hits {
|
||||
hitKey := buildKnowledgeSearchResultKey(hit)
|
||||
hitRecord := &models.KnowledgeRetrieveHit{
|
||||
RetrieveLogID: log.ID,
|
||||
KnowledgeBaseID: hit.KnowledgeBaseID,
|
||||
ChunkID: hit.ChunkID,
|
||||
DocumentID: hit.DocumentID,
|
||||
DocumentTitle: hit.DocumentTitle,
|
||||
FaqID: hit.FaqID,
|
||||
FaqQuestion: hit.FaqQuestion,
|
||||
ChunkNo: hit.ChunkNo,
|
||||
Title: hit.Title,
|
||||
SectionPath: hit.SectionPath,
|
||||
ChunkType: "",
|
||||
Provider: req.ChunkProvider,
|
||||
RankNo: i + 1,
|
||||
Score: hit.Score,
|
||||
RerankScore: hit.RerankScore,
|
||||
UsedInAnswer: hasHitKey(usedHitKeys, hitKey),
|
||||
IsCitation: hasHitKey(citationKeys, hitKey),
|
||||
Snippet: hit.Content,
|
||||
CreatedAt: now,
|
||||
}
|
||||
if err := ctx.Tx.Create(hitRecord).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return log, nil
|
||||
}
|
||||
|
||||
func buildRetrieveTraceData(req *CreateRetrieveLogRequest) string {
|
||||
trace := retrieveTraceData{
|
||||
Retrieve: retrieveTraceRetrieve{
|
||||
Provider: req.ChunkProvider,
|
||||
RerankEnabled: req.RerankEnabled,
|
||||
RerankLimit: req.RerankLimit,
|
||||
RawHitCount: len(req.Hits),
|
||||
ContextHitCount: len(req.UsedHits),
|
||||
CitationCount: len(req.Citations),
|
||||
},
|
||||
ChunkConfig: retrieveTraceChunkConfig{
|
||||
Provider: req.ChunkProvider,
|
||||
TargetTokens: req.ChunkTargetTokens,
|
||||
MaxTokens: req.ChunkMaxTokens,
|
||||
OverlapTokens: req.ChunkOverlapTokens,
|
||||
},
|
||||
Context: retrieveTraceContext{
|
||||
KnowledgeBaseIDs: distinctKnowledgeBaseIDs(req.UsedHits),
|
||||
DocumentIDs: distinctDocumentIDs(req.UsedHits),
|
||||
SectionPaths: distinctSectionPaths(req.UsedHits),
|
||||
UsedChunkKeys: buildUsedChunkKeys(req.UsedHits),
|
||||
},
|
||||
Citations: buildTraceCitations(req.Citations),
|
||||
}
|
||||
data, err := json.Marshal(trace)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return string(data)
|
||||
}
|
||||
|
||||
func buildTraceCitations(citations []response.KnowledgeCitation) []retrieveTraceCitation {
|
||||
items := make([]retrieveTraceCitation, 0, len(citations))
|
||||
for _, item := range citations {
|
||||
items = append(items, retrieveTraceCitation{
|
||||
DocumentID: item.DocumentID,
|
||||
ChunkNo: item.ChunkNo,
|
||||
SectionPath: item.SectionPath,
|
||||
})
|
||||
}
|
||||
return items
|
||||
}
|
||||
|
||||
func buildUsedChunkKeys(hits []response.KnowledgeSearchResult) []string {
|
||||
keys := make([]string, 0, len(hits))
|
||||
for _, item := range hits {
|
||||
keys = append(keys, buildKnowledgeSearchResultKey(item))
|
||||
}
|
||||
return keys
|
||||
}
|
||||
|
||||
func distinctKnowledgeBaseIDs(hits []response.KnowledgeSearchResult) []int64 {
|
||||
ids := make([]int64, 0)
|
||||
seen := make(map[int64]struct{})
|
||||
for _, item := range hits {
|
||||
if item.KnowledgeBaseID <= 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[item.KnowledgeBaseID]; ok {
|
||||
continue
|
||||
}
|
||||
seen[item.KnowledgeBaseID] = struct{}{}
|
||||
ids = append(ids, item.KnowledgeBaseID)
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
func distinctDocumentIDs(hits []response.KnowledgeSearchResult) []int64 {
|
||||
seen := make(map[int64]struct{})
|
||||
items := make([]int64, 0)
|
||||
for _, item := range hits {
|
||||
if item.DocumentID <= 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[item.DocumentID]; ok {
|
||||
continue
|
||||
}
|
||||
seen[item.DocumentID] = struct{}{}
|
||||
items = append(items, item.DocumentID)
|
||||
}
|
||||
return items
|
||||
}
|
||||
|
||||
func distinctSectionPaths(hits []response.KnowledgeSearchResult) []string {
|
||||
seen := make(map[string]struct{})
|
||||
items := make([]string, 0)
|
||||
for _, item := range hits {
|
||||
sectionPath := item.SectionPath
|
||||
if sectionPath == "" {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[sectionPath]; ok {
|
||||
continue
|
||||
}
|
||||
seen[sectionPath] = struct{}{}
|
||||
items = append(items, sectionPath)
|
||||
}
|
||||
return items
|
||||
}
|
||||
|
||||
func buildKnowledgeSearchResultKey(item response.KnowledgeSearchResult) string {
|
||||
if item.FaqID > 0 {
|
||||
return fmt.Sprintf("faq:%d|%d", item.FaqID, item.ChunkNo)
|
||||
}
|
||||
return fmt.Sprintf("%d|%s|%d", item.DocumentID, item.SectionPath, item.ChunkNo)
|
||||
}
|
||||
|
||||
func buildKnowledgeCitationKey(item response.KnowledgeCitation) string {
|
||||
if item.FaqID > 0 {
|
||||
return fmt.Sprintf("faq:%d|%d", item.FaqID, item.ChunkNo)
|
||||
}
|
||||
return fmt.Sprintf("%d|%s|%d", item.DocumentID, item.SectionPath, item.ChunkNo)
|
||||
}
|
||||
|
||||
func hasHitKey(items map[string]struct{}, key string) bool {
|
||||
_, ok := items[key]
|
||||
return ok
|
||||
}
|
||||
@@ -0,0 +1,49 @@
|
||||
package rag
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"cs-agent/internal/models"
|
||||
)
|
||||
|
||||
func TestResolveKnowledgeBaseSearchOptionsUsesKnowledgeBaseDefaults(t *testing.T) {
|
||||
topK, scoreThreshold := resolveKnowledgeBaseSearchOptions(RetrieveRequest{}, &models.KnowledgeBase{
|
||||
DefaultTopK: 6,
|
||||
DefaultScoreThreshold: 0.42,
|
||||
})
|
||||
|
||||
if topK != 6 {
|
||||
t.Fatalf("expected topK 6, got %d", topK)
|
||||
}
|
||||
if scoreThreshold != float32(0.42) {
|
||||
t.Fatalf("expected score threshold 0.42, got %v", scoreThreshold)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveKnowledgeBaseSearchOptionsRequestOverridesKnowledgeBaseDefaults(t *testing.T) {
|
||||
topK, scoreThreshold := resolveKnowledgeBaseSearchOptions(RetrieveRequest{
|
||||
TopK: 9,
|
||||
ScoreThreshold: 0.55,
|
||||
}, &models.KnowledgeBase{
|
||||
DefaultTopK: 6,
|
||||
DefaultScoreThreshold: 0.42,
|
||||
})
|
||||
|
||||
if topK != 9 {
|
||||
t.Fatalf("expected request topK 9, got %d", topK)
|
||||
}
|
||||
if scoreThreshold != float32(0.55) {
|
||||
t.Fatalf("expected request score threshold 0.55, got %v", scoreThreshold)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveKnowledgeBaseSearchOptionsUsesSystemDefaults(t *testing.T) {
|
||||
topK, scoreThreshold := resolveKnowledgeBaseSearchOptions(RetrieveRequest{}, nil)
|
||||
|
||||
if topK != 8 {
|
||||
t.Fatalf("expected fallback topK 8, got %d", topK)
|
||||
}
|
||||
if scoreThreshold != float32(0.3) {
|
||||
t.Fatalf("expected fallback score threshold 0.3, got %v", scoreThreshold)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,48 @@
|
||||
package rag
|
||||
|
||||
type RetrieveRequest struct {
|
||||
KnowledgeBaseIDs []int64
|
||||
Query string
|
||||
TopK int
|
||||
ScoreThreshold float64
|
||||
}
|
||||
|
||||
type RetrieveResult struct {
|
||||
KnowledgeBaseID int64 `json:"knowledgeBaseId"`
|
||||
ChunkID int64 `json:"chunkId"`
|
||||
DocumentID int64 `json:"documentId"`
|
||||
DocumentTitle string `json:"documentTitle"`
|
||||
FaqID int64 `json:"faqId"`
|
||||
FaqQuestion string `json:"faqQuestion"`
|
||||
ChunkNo int `json:"chunkNo"`
|
||||
Title string `json:"title"`
|
||||
SectionPath string `json:"sectionPath"`
|
||||
Content string `json:"content"`
|
||||
Score float32 `json:"score"`
|
||||
ChunkType string `json:"chunkType"`
|
||||
}
|
||||
|
||||
type RerankRequest struct {
|
||||
Model string `json:"model"`
|
||||
Query string `json:"query"`
|
||||
Documents []string `json:"documents"`
|
||||
TopN int `json:"top_n"`
|
||||
}
|
||||
|
||||
type RerankResponse struct {
|
||||
Results []struct {
|
||||
Document string `json:"document"`
|
||||
Index int `json:"index"`
|
||||
RelevanceScore float64 `json:"relevance_score"`
|
||||
} `json:"results"`
|
||||
Meta struct {
|
||||
APIVersion struct {
|
||||
Version string `json:"version"`
|
||||
} `json:"api_version"`
|
||||
} `json:"meta"`
|
||||
}
|
||||
|
||||
type RerankResult struct {
|
||||
Index int `json:"index"`
|
||||
RelevanceScore float64 `json:"relevanceScore"`
|
||||
}
|
||||
@@ -0,0 +1,112 @@
|
||||
package rag
|
||||
|
||||
import (
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"strings"
|
||||
|
||||
"github.com/yuin/goldmark"
|
||||
"github.com/yuin/goldmark/extension"
|
||||
"github.com/yuin/goldmark/parser"
|
||||
"github.com/yuin/goldmark/renderer/html"
|
||||
htmlparser "golang.org/x/net/html"
|
||||
)
|
||||
|
||||
var plainTextMarkdown = goldmark.New(
|
||||
goldmark.WithExtensions(extension.GFM),
|
||||
goldmark.WithParserOptions(
|
||||
parser.WithAutoHeadingID(),
|
||||
),
|
||||
goldmark.WithRendererOptions(
|
||||
html.WithHardWraps(),
|
||||
html.WithXHTML(),
|
||||
),
|
||||
)
|
||||
|
||||
func ExtractPlainText(content string, contentType enums.KnowledgeDocumentContentType) string {
|
||||
switch contentType {
|
||||
case enums.KnowledgeDocumentContentTypeMarkdown:
|
||||
return ExtractPlainTextFromMarkdown(content)
|
||||
case enums.KnowledgeDocumentContentTypeHTML:
|
||||
return ExtractPlainTextFromHTML(content)
|
||||
default:
|
||||
return normalizeWhitespace(content)
|
||||
}
|
||||
}
|
||||
|
||||
func ExtractPlainTextFromMarkdown(content string) string {
|
||||
content = strings.TrimSpace(content)
|
||||
if content == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
var buf strings.Builder
|
||||
if err := plainTextMarkdown.Convert([]byte(content), &buf); err != nil {
|
||||
return normalizeWhitespace(content)
|
||||
}
|
||||
return ExtractPlainTextFromHTML(buf.String())
|
||||
}
|
||||
|
||||
func ExtractPlainTextFromHTML(content string) string {
|
||||
content = strings.TrimSpace(content)
|
||||
if content == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
var builder strings.Builder
|
||||
parent := &htmlparser.Node{
|
||||
Type: htmlparser.ElementNode,
|
||||
Data: "div",
|
||||
}
|
||||
nodes, err := htmlparser.ParseFragment(strings.NewReader(content), parent)
|
||||
if err == nil {
|
||||
for _, node := range nodes {
|
||||
writeHTMLNodeText(&builder, node)
|
||||
}
|
||||
return normalizeWhitespace(builder.String())
|
||||
}
|
||||
|
||||
// 兜底:部分输入在 ParseFragment 下会失败(例如不符合 fragment 规则或上下文不匹配)。
|
||||
// 这里用完整 HTML 解析保证可用性。
|
||||
doc, err := htmlparser.Parse(strings.NewReader("<div>" + content + "</div>"))
|
||||
if err != nil {
|
||||
return normalizeWhitespace(content)
|
||||
}
|
||||
writeHTMLNodeText(&builder, doc)
|
||||
return normalizeWhitespace(builder.String())
|
||||
}
|
||||
|
||||
func writeHTMLNodeText(builder *strings.Builder, node *htmlparser.Node) {
|
||||
if node == nil {
|
||||
return
|
||||
}
|
||||
|
||||
switch node.Type {
|
||||
case htmlparser.TextNode:
|
||||
builder.WriteString(node.Data)
|
||||
case htmlparser.ElementNode:
|
||||
if shouldSeparateHTMLText(node.Data) {
|
||||
builder.WriteByte(' ')
|
||||
}
|
||||
}
|
||||
|
||||
for child := node.FirstChild; child != nil; child = child.NextSibling {
|
||||
writeHTMLNodeText(builder, child)
|
||||
}
|
||||
|
||||
if node.Type == htmlparser.ElementNode && shouldSeparateHTMLText(node.Data) {
|
||||
builder.WriteByte(' ')
|
||||
}
|
||||
}
|
||||
|
||||
func shouldSeparateHTMLText(tag string) bool {
|
||||
switch tag {
|
||||
case "p", "div", "br", "li", "ul", "ol", "blockquote", "pre", "table", "tr", "td", "th", "h1", "h2", "h3", "h4", "h5", "h6":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeWhitespace(content string) string {
|
||||
return strings.Join(strings.Fields(strings.TrimSpace(content)), " ")
|
||||
}
|
||||
@@ -0,0 +1,43 @@
|
||||
package rag
|
||||
|
||||
import (
|
||||
"cs-agent/internal/ai/rag/vectordb"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestExtractPlainTextFromHTMLSeparatesBlockContent(t *testing.T) {
|
||||
got := ExtractPlainTextFromHTML("<div>Hello</div><div>World</div><p>Again</p>")
|
||||
want := "Hello World Again"
|
||||
if got != want {
|
||||
t.Fatalf("expected %q, got %q", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractPlainTextMarkdownUsesGoldmark(t *testing.T) {
|
||||
got := ExtractPlainText("# Title\n\n- one\n- two", enums.KnowledgeDocumentContentTypeMarkdown)
|
||||
want := "Title one two"
|
||||
if got != want {
|
||||
t.Fatalf("expected %q, got %q", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChunkPayloadFromMapSupportsTypedConversion(t *testing.T) {
|
||||
got := vectordb.ChunkPayloadFromMap(map[string]any{
|
||||
"knowledge_base_id": "1",
|
||||
"document_id": "123",
|
||||
"document_title": "Doc",
|
||||
"chunk_no": "2",
|
||||
"chunk_type": "text",
|
||||
"section_path": "A > B",
|
||||
"title": "hello",
|
||||
"content": "world",
|
||||
"provider": "structured",
|
||||
})
|
||||
if got.KnowledgeBaseID != 1 || got.DocumentID != 123 || got.ChunkNo != 2 {
|
||||
t.Fatalf("unexpected numeric conversion result: %+v", got)
|
||||
}
|
||||
if got.DocumentTitle != "Doc" || got.SectionPath != "A > B" || got.Provider != "structured" {
|
||||
t.Fatalf("unexpected string conversion result: %+v", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,43 @@
|
||||
package vectordb
|
||||
|
||||
import (
|
||||
"github.com/mlogclub/simple/common/structs"
|
||||
"github.com/spf13/cast"
|
||||
)
|
||||
|
||||
type ChunkPayload struct {
|
||||
KnowledgeBaseID int64 `json:"knowledge_base_id"`
|
||||
DocumentID int64 `json:"document_id"`
|
||||
DocumentTitle string `json:"document_title"`
|
||||
FaqID int64 `json:"faq_id"`
|
||||
FaqQuestion string `json:"faq_question"`
|
||||
ChunkNo int `json:"chunk_no"`
|
||||
ChunkType string `json:"chunk_type"`
|
||||
SectionPath string `json:"section_path"`
|
||||
Title string `json:"title"`
|
||||
Content string `json:"content"`
|
||||
Provider string `json:"provider"`
|
||||
}
|
||||
|
||||
func (p ChunkPayload) ToMap() map[string]any {
|
||||
return structs.StructToMap(p)
|
||||
}
|
||||
|
||||
func ChunkPayloadFromMap(data map[string]any) ChunkPayload {
|
||||
if data == nil {
|
||||
return ChunkPayload{}
|
||||
}
|
||||
return ChunkPayload{
|
||||
KnowledgeBaseID: cast.ToInt64(data["knowledge_base_id"]),
|
||||
DocumentID: cast.ToInt64(data["document_id"]),
|
||||
DocumentTitle: cast.ToString(data["document_title"]),
|
||||
FaqID: cast.ToInt64(data["faq_id"]),
|
||||
FaqQuestion: cast.ToString(data["faq_question"]),
|
||||
ChunkNo: cast.ToInt(data["chunk_no"]),
|
||||
ChunkType: cast.ToString(data["chunk_type"]),
|
||||
SectionPath: cast.ToString(data["section_path"]),
|
||||
Title: cast.ToString(data["title"]),
|
||||
Content: cast.ToString(data["content"]),
|
||||
Provider: cast.ToString(data["provider"]),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,86 @@
|
||||
package vectordb
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"cs-agent/internal/pkg/config"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
)
|
||||
|
||||
var defaultProvider Provider
|
||||
|
||||
func Init(cfg *config.VectorDBConfig) error {
|
||||
if cfg == nil || cfg.Type == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
var err error
|
||||
switch enums.VectorDBType(cfg.Type) {
|
||||
case enums.VectorDBTypeQdrant:
|
||||
defaultProvider, err = NewQdrantProvider(cfg)
|
||||
default:
|
||||
return fmt.Errorf("unsupported vectordb type: %s", cfg.Type)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func GetProvider() Provider {
|
||||
return defaultProvider
|
||||
}
|
||||
|
||||
func Close() error {
|
||||
if defaultProvider != nil {
|
||||
return defaultProvider.Close()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func CreateCollection(ctx context.Context, name string, dimension int) error {
|
||||
if defaultProvider == nil {
|
||||
return fmt.Errorf("vectordb provider not initialized")
|
||||
}
|
||||
return defaultProvider.CreateCollection(ctx, name, dimension)
|
||||
}
|
||||
|
||||
func DeleteCollection(ctx context.Context, name string) error {
|
||||
if defaultProvider == nil {
|
||||
return fmt.Errorf("vectordb provider not initialized")
|
||||
}
|
||||
return defaultProvider.DeleteCollection(ctx, name)
|
||||
}
|
||||
|
||||
func GetCollection(ctx context.Context, name string) (*CollectionInfo, error) {
|
||||
if defaultProvider == nil {
|
||||
return nil, fmt.Errorf("vectordb provider not initialized")
|
||||
}
|
||||
return defaultProvider.GetCollection(ctx, name)
|
||||
}
|
||||
|
||||
func ListCollections(ctx context.Context) ([]string, error) {
|
||||
if defaultProvider == nil {
|
||||
return nil, fmt.Errorf("vectordb provider not initialized")
|
||||
}
|
||||
return defaultProvider.ListCollections(ctx)
|
||||
}
|
||||
|
||||
func UpsertVectors(ctx context.Context, collectionName string, vectors []Vector) error {
|
||||
if defaultProvider == nil {
|
||||
return fmt.Errorf("vectordb provider not initialized")
|
||||
}
|
||||
return defaultProvider.UpsertVectors(ctx, collectionName, vectors)
|
||||
}
|
||||
|
||||
func DeleteVectors(ctx context.Context, collectionName string, ids []string) error {
|
||||
if defaultProvider == nil {
|
||||
return fmt.Errorf("vectordb provider not initialized")
|
||||
}
|
||||
return defaultProvider.DeleteVectors(ctx, collectionName, ids)
|
||||
}
|
||||
|
||||
func Search(ctx context.Context, req *SearchRequest) ([]SearchResult, error) {
|
||||
if defaultProvider == nil {
|
||||
return nil, fmt.Errorf("vectordb provider not initialized")
|
||||
}
|
||||
return defaultProvider.Search(ctx, req)
|
||||
}
|
||||
@@ -0,0 +1,292 @@
|
||||
package vectordb
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/qdrant/go-client/qdrant"
|
||||
|
||||
"cs-agent/internal/pkg/config"
|
||||
)
|
||||
|
||||
type Vector struct {
|
||||
ID string `json:"id"`
|
||||
Vector []float32 `json:"vector"`
|
||||
Payload ChunkPayload `json:"payload"`
|
||||
}
|
||||
|
||||
type SearchRequest struct {
|
||||
CollectionName string `json:"collectionName"`
|
||||
Vector []float32 `json:"vector"`
|
||||
TopK int `json:"topK"`
|
||||
ScoreThreshold float32 `json:"scoreThreshold"`
|
||||
Filter *SearchFilter `json:"filter,omitempty"`
|
||||
}
|
||||
|
||||
type SearchFilter struct {
|
||||
KnowledgeBaseIDs []int64 `json:"knowledgeBaseIds,omitempty"`
|
||||
DocumentIDs []int64 `json:"documentIds,omitempty"`
|
||||
}
|
||||
|
||||
type SearchResult struct {
|
||||
ID string `json:"id"`
|
||||
Score float32 `json:"score"`
|
||||
Payload ChunkPayload `json:"payload"`
|
||||
}
|
||||
|
||||
type CollectionInfo struct {
|
||||
Name string `json:"name"`
|
||||
Dimension int `json:"dimension"`
|
||||
PointCount int `json:"pointCount"`
|
||||
Status string `json:"status"`
|
||||
}
|
||||
|
||||
type Provider interface {
|
||||
CreateCollection(ctx context.Context, name string, dimension int) error
|
||||
DeleteCollection(ctx context.Context, name string) error
|
||||
GetCollection(ctx context.Context, name string) (*CollectionInfo, error)
|
||||
ListCollections(ctx context.Context) ([]string, error)
|
||||
|
||||
UpsertVectors(ctx context.Context, collectionName string, vectors []Vector) error
|
||||
DeleteVectors(ctx context.Context, collectionName string, ids []string) error
|
||||
|
||||
Search(ctx context.Context, req *SearchRequest) ([]SearchResult, error)
|
||||
Close() error
|
||||
}
|
||||
|
||||
type QdrantProvider struct {
|
||||
client *qdrant.Client
|
||||
}
|
||||
|
||||
func NewQdrantProvider(cfg *config.VectorDBConfig) (*QdrantProvider, error) {
|
||||
if cfg == nil {
|
||||
return nil, fmt.Errorf("vectordb config is nil")
|
||||
}
|
||||
|
||||
host := cfg.Host
|
||||
if host == "" {
|
||||
host = "localhost"
|
||||
}
|
||||
|
||||
port := cfg.GrpcPort
|
||||
if port <= 0 {
|
||||
port = 6334
|
||||
}
|
||||
|
||||
client, err := qdrant.NewClient(&qdrant.Config{
|
||||
Host: host,
|
||||
Port: port,
|
||||
APIKey: cfg.APIKey,
|
||||
UseTLS: cfg.UseTLS,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create qdrant client: %w", err)
|
||||
}
|
||||
|
||||
return &QdrantProvider{client: client}, nil
|
||||
}
|
||||
|
||||
func (p *QdrantProvider) Close() error {
|
||||
if p.client != nil {
|
||||
return p.client.Close()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *QdrantProvider) CreateCollection(ctx context.Context, name string, dimension int) error {
|
||||
err := p.client.CreateCollection(ctx, &qdrant.CreateCollection{
|
||||
CollectionName: name,
|
||||
VectorsConfig: qdrant.NewVectorsConfig(&qdrant.VectorParams{
|
||||
Size: uint64(dimension),
|
||||
Distance: qdrant.Distance_Cosine,
|
||||
}),
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create collection %s: %w", name, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *QdrantProvider) DeleteCollection(ctx context.Context, name string) error {
|
||||
err := p.client.DeleteCollection(ctx, name)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to delete collection %s: %w", name, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *QdrantProvider) GetCollection(ctx context.Context, name string) (*CollectionInfo, error) {
|
||||
info, err := p.client.GetCollectionInfo(ctx, name)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get collection %s: %w", name, err)
|
||||
}
|
||||
|
||||
status := info.GetStatus().String()
|
||||
pointCount := int(info.GetPointsCount())
|
||||
|
||||
dimension := 0
|
||||
if info.Config != nil && info.Config.Params != nil {
|
||||
vectorsConfig := info.Config.Params.VectorsConfig
|
||||
if vectorsConfig != nil {
|
||||
params := vectorsConfig.GetParams()
|
||||
if params != nil {
|
||||
dimension = int(params.Size)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return &CollectionInfo{
|
||||
Name: name,
|
||||
Dimension: dimension,
|
||||
PointCount: pointCount,
|
||||
Status: status,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (p *QdrantProvider) ListCollections(ctx context.Context) ([]string, error) {
|
||||
collections, err := p.client.ListCollections(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to list collections: %w", err)
|
||||
}
|
||||
|
||||
return collections, nil
|
||||
}
|
||||
|
||||
func (p *QdrantProvider) UpsertVectors(ctx context.Context, collectionName string, vectors []Vector) error {
|
||||
if len(vectors) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
points := make([]*qdrant.PointStruct, 0, len(vectors))
|
||||
for _, v := range vectors {
|
||||
points = append(points, &qdrant.PointStruct{
|
||||
Id: qdrant.NewID(v.ID),
|
||||
Vectors: qdrant.NewVectors(v.Vector...),
|
||||
Payload: qdrant.NewValueMap(v.Payload.ToMap()),
|
||||
})
|
||||
}
|
||||
|
||||
_, err := p.client.Upsert(ctx, &qdrant.UpsertPoints{
|
||||
CollectionName: collectionName,
|
||||
Points: points,
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to upsert vectors to collection %s: %w", collectionName, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *QdrantProvider) DeleteVectors(ctx context.Context, collectionName string, ids []string) error {
|
||||
if len(ids) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
pointIDs := make([]*qdrant.PointId, 0, len(ids))
|
||||
for _, id := range ids {
|
||||
pointIDs = append(pointIDs, qdrant.NewID(id))
|
||||
}
|
||||
|
||||
_, err := p.client.Delete(ctx, &qdrant.DeletePoints{
|
||||
CollectionName: collectionName,
|
||||
Points: &qdrant.PointsSelector{
|
||||
PointsSelectorOneOf: &qdrant.PointsSelector_Points{
|
||||
Points: &qdrant.PointsIdsList{
|
||||
Ids: pointIDs,
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to delete vectors from collection %s: %w", collectionName, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *QdrantProvider) Search(ctx context.Context, req *SearchRequest) ([]SearchResult, error) {
|
||||
filter := p.buildFilter(req.Filter)
|
||||
|
||||
results, err := p.client.Query(ctx, &qdrant.QueryPoints{
|
||||
CollectionName: req.CollectionName,
|
||||
Query: qdrant.NewQuery(req.Vector...),
|
||||
Limit: qdrant.PtrOf(uint64(req.TopK)),
|
||||
ScoreThreshold: &req.ScoreThreshold,
|
||||
Filter: filter,
|
||||
WithPayload: qdrant.NewWithPayload(true),
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to search collection %s: %w", req.CollectionName, err)
|
||||
}
|
||||
|
||||
searchResults := make([]SearchResult, 0, len(results))
|
||||
for _, r := range results {
|
||||
payload := make(map[string]any)
|
||||
if r.Payload != nil {
|
||||
for k, v := range r.Payload {
|
||||
payload[k] = p.extractPayloadValue(v)
|
||||
}
|
||||
}
|
||||
|
||||
id := ""
|
||||
if r.Id != nil {
|
||||
id = r.Id.GetUuid()
|
||||
}
|
||||
|
||||
searchResults = append(searchResults, SearchResult{
|
||||
ID: id,
|
||||
Score: r.Score,
|
||||
Payload: ChunkPayloadFromMap(payload),
|
||||
})
|
||||
}
|
||||
|
||||
return searchResults, nil
|
||||
}
|
||||
|
||||
func (p *QdrantProvider) buildFilter(filter *SearchFilter) *qdrant.Filter {
|
||||
if filter == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
must := make([]*qdrant.Condition, 0, 2)
|
||||
if len(filter.KnowledgeBaseIDs) > 0 {
|
||||
must = append(must, qdrant.NewMatchInts("knowledge_base_id", filter.KnowledgeBaseIDs...))
|
||||
}
|
||||
if len(filter.DocumentIDs) > 0 {
|
||||
must = append(must, qdrant.NewMatchInts("document_id", filter.DocumentIDs...))
|
||||
}
|
||||
if len(must) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
return &qdrant.Filter{Must: must}
|
||||
}
|
||||
|
||||
func (p *QdrantProvider) extractPayloadValue(v *qdrant.Value) interface{} {
|
||||
if v == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
switch val := v.Kind.(type) {
|
||||
case *qdrant.Value_StringValue:
|
||||
return val.StringValue
|
||||
case *qdrant.Value_IntegerValue:
|
||||
return val.IntegerValue
|
||||
case *qdrant.Value_DoubleValue:
|
||||
return val.DoubleValue
|
||||
case *qdrant.Value_BoolValue:
|
||||
return val.BoolValue
|
||||
case *qdrant.Value_ListValue:
|
||||
list := make([]interface{}, 0, len(val.ListValue.Values))
|
||||
for _, item := range val.ListValue.Values {
|
||||
list = append(list, p.extractPayloadValue(item))
|
||||
}
|
||||
return list
|
||||
case *qdrant.Value_StructValue:
|
||||
m := make(map[string]interface{})
|
||||
for k, v := range val.StructValue.Fields {
|
||||
m[k] = p.extractPayloadValue(v)
|
||||
}
|
||||
return m
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -0,0 +1,248 @@
|
||||
package skills
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
"text/template"
|
||||
"time"
|
||||
|
||||
"cs-agent/internal/ai"
|
||||
"cs-agent/internal/ai/mcps"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"cs-agent/internal/pkg/errorsx"
|
||||
)
|
||||
|
||||
type mcpToolExecutionConfig struct {
|
||||
ServerCode string `json:"serverCode"`
|
||||
ToolName string `json:"toolName"`
|
||||
Arguments map[string]string `json:"arguments"`
|
||||
}
|
||||
|
||||
func executeByPlan(ctx context.Context, plan *ExecutionPlan, runtimeCtx RuntimeContext) (string, *ExecutionTrace, error) {
|
||||
if plan == nil || plan.Skill == nil {
|
||||
return "", nil, nil
|
||||
}
|
||||
trace := &ExecutionTrace{
|
||||
Status: "started",
|
||||
ExecutionMode: string(plan.Skill.ExecutionMode),
|
||||
}
|
||||
switch plan.Skill.ExecutionMode {
|
||||
case "", enums.SkillExecutionModePromptOnly:
|
||||
replyText, err := executePromptOnly(ctx, plan, runtimeCtx, trace)
|
||||
return replyText, trace, err
|
||||
case enums.SkillExecutionModeMCPTool:
|
||||
replyText, err := executeMCPTool(ctx, plan, runtimeCtx, trace)
|
||||
return replyText, trace, err
|
||||
default:
|
||||
trace.Status = "invalid_execution_mode"
|
||||
return "", trace, errorsx.InvalidParam("Skill执行模式不支持")
|
||||
}
|
||||
}
|
||||
|
||||
func executePromptOnly(ctx context.Context, plan *ExecutionPlan, runtimeCtx RuntimeContext, trace *ExecutionTrace) (string, error) {
|
||||
if plan == nil || plan.Skill == nil {
|
||||
return "", nil
|
||||
}
|
||||
if plan.AIConfig == nil {
|
||||
return "", errorsx.InvalidParam("Skill 关联的 AI 配置不可用")
|
||||
}
|
||||
systemPrompt := strings.TrimSpace(plan.Skill.Prompt)
|
||||
if systemPrompt == "" {
|
||||
return "", errorsx.InvalidParam("Skill Prompt 不能为空")
|
||||
}
|
||||
userPrompt := strings.TrimSpace(runtimeCtx.UserMessage)
|
||||
if userPrompt == "" {
|
||||
return "", errorsx.InvalidParam("用户消息不能为空")
|
||||
}
|
||||
promptTrace := &PromptTrace{Status: "started"}
|
||||
if trace != nil {
|
||||
trace.Prompt = promptTrace
|
||||
}
|
||||
startedAt := time.Now()
|
||||
result, err := ai.LLM.ChatWithConfig(ctx, plan.AIConfig, systemPrompt, userPrompt)
|
||||
promptTrace.LatencyMs = time.Since(startedAt).Milliseconds()
|
||||
if err != nil {
|
||||
promptTrace.Status = "error"
|
||||
promptTrace.Error = err.Error()
|
||||
if trace != nil {
|
||||
trace.Status = "error"
|
||||
}
|
||||
return "", err
|
||||
}
|
||||
promptTrace.Status = "ok"
|
||||
promptTrace.ModelName = result.ModelName
|
||||
promptTrace.PromptTokens = result.PromptTokens
|
||||
promptTrace.CompletionTokens = result.CompletionTokens
|
||||
if trace != nil {
|
||||
trace.Status = "ok"
|
||||
}
|
||||
return strings.TrimSpace(result.Content), nil
|
||||
}
|
||||
|
||||
func executeMCPTool(ctx context.Context, plan *ExecutionPlan, runtimeCtx RuntimeContext, trace *ExecutionTrace) (string, error) {
|
||||
cfg, err := parseMCPToolExecutionConfig(plan.Skill.ExecutionConfig)
|
||||
if err != nil {
|
||||
if trace != nil {
|
||||
trace.Status = "config_error"
|
||||
}
|
||||
return "", err
|
||||
}
|
||||
arguments, err := buildToolArguments(cfg.Arguments, runtimeCtx)
|
||||
if err != nil {
|
||||
if trace != nil {
|
||||
trace.Status = "argument_error"
|
||||
}
|
||||
return "", err
|
||||
}
|
||||
mcpTrace := &MCPExecutionTrace{
|
||||
Status: "started",
|
||||
ServerCode: cfg.ServerCode,
|
||||
ToolName: cfg.ToolName,
|
||||
Arguments: arguments,
|
||||
}
|
||||
if trace != nil {
|
||||
trace.MCP = mcpTrace
|
||||
}
|
||||
toolStartedAt := time.Now()
|
||||
toolResult, err := mcps.Runtime.CallTool(ctx, cfg.ServerCode, cfg.ToolName, arguments)
|
||||
mcpTrace.LatencyMs = time.Since(toolStartedAt).Milliseconds()
|
||||
if err != nil {
|
||||
mcpTrace.Status = "error"
|
||||
mcpTrace.Error = err.Error()
|
||||
if trace != nil {
|
||||
trace.Status = "error"
|
||||
}
|
||||
return "", err
|
||||
}
|
||||
mcpTrace.Status = "ok"
|
||||
mcpTrace.IsError = toolResult.IsError
|
||||
mcpTrace.ContentItemCount = len(toolResult.Content)
|
||||
mcpTrace.HasStructuredContent = toolResult.StructuredContent != nil
|
||||
toolSummary := buildToolSummary(toolResult)
|
||||
mcpTrace.ResultPreview = truncateTraceText(toolSummary, 500)
|
||||
if strings.TrimSpace(toolSummary) == "" {
|
||||
if trace != nil {
|
||||
trace.Status = "empty_tool_result"
|
||||
}
|
||||
return "", errorsx.InvalidParam("MCP工具未返回有效结果")
|
||||
}
|
||||
systemPrompt := strings.TrimSpace(plan.Skill.Prompt)
|
||||
if systemPrompt == "" {
|
||||
systemPrompt = "你是客服技能助手。请依据工具结果准确回答用户问题,不要编造工具结果中不存在的事实。"
|
||||
}
|
||||
userPrompt := fmt.Sprintf("用户问题:%s\n\n工具结果:\n%s", strings.TrimSpace(runtimeCtx.UserMessage), toolSummary)
|
||||
summaryTrace := &PromptTrace{Status: "started"}
|
||||
mcpTrace.SummaryPrompt = summaryTrace
|
||||
summaryStartedAt := time.Now()
|
||||
result, err := ai.LLM.ChatWithConfig(ctx, plan.AIConfig, systemPrompt, userPrompt)
|
||||
summaryTrace.LatencyMs = time.Since(summaryStartedAt).Milliseconds()
|
||||
if err != nil {
|
||||
summaryTrace.Status = "error"
|
||||
summaryTrace.Error = err.Error()
|
||||
if trace != nil {
|
||||
trace.Status = "error"
|
||||
}
|
||||
return "", err
|
||||
}
|
||||
summaryTrace.Status = "ok"
|
||||
summaryTrace.ModelName = result.ModelName
|
||||
summaryTrace.PromptTokens = result.PromptTokens
|
||||
summaryTrace.CompletionTokens = result.CompletionTokens
|
||||
if trace != nil {
|
||||
trace.Status = "ok"
|
||||
}
|
||||
return strings.TrimSpace(result.Content), nil
|
||||
}
|
||||
|
||||
func parseMCPToolExecutionConfig(raw string) (*mcpToolExecutionConfig, error) {
|
||||
raw = strings.TrimSpace(raw)
|
||||
if raw == "" {
|
||||
return nil, errorsx.InvalidParam("ExecutionConfig不能为空")
|
||||
}
|
||||
cfg := &mcpToolExecutionConfig{}
|
||||
if err := json.Unmarshal([]byte(raw), cfg); err != nil {
|
||||
return nil, errorsx.InvalidParam("ExecutionConfig格式不合法")
|
||||
}
|
||||
if strings.TrimSpace(cfg.ServerCode) == "" {
|
||||
return nil, errorsx.InvalidParam("ExecutionConfig.serverCode不能为空")
|
||||
}
|
||||
if strings.TrimSpace(cfg.ToolName) == "" {
|
||||
return nil, errorsx.InvalidParam("ExecutionConfig.toolName不能为空")
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
func buildToolArguments(templateArgs map[string]string, runtimeCtx RuntimeContext) (map[string]any, error) {
|
||||
if len(templateArgs) == 0 {
|
||||
return map[string]any{
|
||||
"query": strings.TrimSpace(runtimeCtx.UserMessage),
|
||||
}, nil
|
||||
}
|
||||
data := map[string]any{
|
||||
"userMessage": strings.TrimSpace(runtimeCtx.UserMessage),
|
||||
"conversationId": runtimeCtx.ConversationID,
|
||||
"aiAgentId": runtimeCtx.AIAgentID,
|
||||
"manualSkillCode": strings.TrimSpace(runtimeCtx.ManualSkillCode),
|
||||
"intentCode": strings.TrimSpace(runtimeCtx.IntentCode),
|
||||
}
|
||||
ret := make(map[string]any, len(templateArgs))
|
||||
for key, value := range templateArgs {
|
||||
rendered, err := renderTemplate(value, data)
|
||||
if err != nil {
|
||||
return nil, errorsx.InvalidParam("ExecutionConfig.arguments模板不合法")
|
||||
}
|
||||
ret[key] = rendered
|
||||
}
|
||||
return ret, nil
|
||||
}
|
||||
|
||||
func renderTemplate(raw string, data map[string]any) (string, error) {
|
||||
tpl, err := template.New("skill_arg").Option("missingkey=zero").Parse(raw)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
if err := tpl.Execute(&buf, data); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return strings.TrimSpace(buf.String()), nil
|
||||
}
|
||||
|
||||
func buildToolSummary(result *mcps.ToolCallResult) string {
|
||||
if result == nil {
|
||||
return ""
|
||||
}
|
||||
lines := make([]string, 0, len(result.Content)+2)
|
||||
if result.StructuredContent != nil {
|
||||
if data, err := json.Marshal(result.StructuredContent); err == nil {
|
||||
lines = append(lines, string(data))
|
||||
}
|
||||
}
|
||||
for _, item := range result.Content {
|
||||
if strings.TrimSpace(item.Text) != "" {
|
||||
lines = append(lines, strings.TrimSpace(item.Text))
|
||||
continue
|
||||
}
|
||||
if item.Data != nil {
|
||||
if data, err := json.Marshal(item.Data); err == nil {
|
||||
lines = append(lines, string(data))
|
||||
}
|
||||
}
|
||||
}
|
||||
return strings.TrimSpace(strings.Join(lines, "\n"))
|
||||
}
|
||||
|
||||
func truncateTraceText(raw string, limit int) string {
|
||||
raw = strings.TrimSpace(raw)
|
||||
if raw == "" || limit <= 0 {
|
||||
return raw
|
||||
}
|
||||
runes := []rune(raw)
|
||||
if len(runes) <= limit {
|
||||
return raw
|
||||
}
|
||||
return strings.TrimSpace(string(runes[:limit])) + "..."
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
package skills
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"time"
|
||||
|
||||
"cs-agent/internal/models"
|
||||
)
|
||||
|
||||
// BuildRunLog 根据执行计划与运行结果构建 Skill 运行日志。
|
||||
func BuildRunLog(ctx RuntimeContext, plan *ExecutionPlan, trace *ExecutionTrace, err error) *models.SkillRunLog {
|
||||
log := &models.SkillRunLog{
|
||||
ConversationID: ctx.ConversationID,
|
||||
AIAgentID: ctx.AIAgentID,
|
||||
ManualSkillCode: ctx.ManualSkillCode,
|
||||
IntentCode: ctx.IntentCode,
|
||||
UserMessage: ctx.UserMessage,
|
||||
TraceData: buildTraceData(trace),
|
||||
CreatedAt: time.Now(),
|
||||
}
|
||||
if plan != nil {
|
||||
if plan.AIConfig != nil {
|
||||
log.AIConfigID = plan.AIConfig.ID
|
||||
log.UsedModel = plan.AIConfig.ModelName
|
||||
log.UsedProvider = plan.AIConfig.Provider
|
||||
}
|
||||
if plan.Skill != nil {
|
||||
log.SkillDefinitionID = plan.Skill.ID
|
||||
log.SkillCode = plan.Skill.Code
|
||||
log.Matched = true
|
||||
log.FinalSelected = true
|
||||
log.MatchReason = plan.MatchReason
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
log.ErrorMessage = err.Error()
|
||||
} else if !log.Matched {
|
||||
if plan != nil && plan.MatchReason != "" {
|
||||
log.MatchReason = plan.MatchReason
|
||||
} else {
|
||||
log.MatchReason = "not_matched"
|
||||
}
|
||||
}
|
||||
return log
|
||||
}
|
||||
|
||||
func buildTraceData(trace *ExecutionTrace) string {
|
||||
if trace == nil {
|
||||
return ""
|
||||
}
|
||||
data, err := json.Marshal(trace)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return string(data)
|
||||
}
|
||||
@@ -0,0 +1,183 @@
|
||||
package skills
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"cs-agent/internal/ai"
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"cs-agent/internal/pkg/errorsx"
|
||||
"cs-agent/internal/pkg/utils"
|
||||
"cs-agent/internal/repositories"
|
||||
|
||||
"github.com/mlogclub/simple/common/strs"
|
||||
"github.com/mlogclub/simple/sqls"
|
||||
)
|
||||
|
||||
type intentTriggerConfig struct {
|
||||
Intents []string `json:"intents"`
|
||||
}
|
||||
|
||||
// MatchSkill 对单个 SkillDefinition 执行命中判断。
|
||||
func MatchSkill(execCtx context.Context, ctx RuntimeContext, aiAgent *models.AIAgent, aiConfig *models.AIConfig) (*models.SkillDefinition, string, *RouteTrace, error) {
|
||||
if strs.IsNotBlank(ctx.ManualSkillCode) {
|
||||
skill := repositories.SkillDefinitionRepository.GetByCode(sqls.DB(), ctx.ManualSkillCode)
|
||||
if skill == nil || skill.Status != enums.StatusOk {
|
||||
return nil, "", nil, errorsx.InvalidParam("Skill 不存在或未启用")
|
||||
}
|
||||
return skill, "manual_skill_code", &RouteTrace{
|
||||
Status: "manual_selected",
|
||||
SelectedSkillCode: skill.Code,
|
||||
}, nil
|
||||
}
|
||||
|
||||
candidates := loadCandidateSkills(aiAgent)
|
||||
trace := &RouteTrace{
|
||||
Status: "started",
|
||||
CandidateSkillCodes: make([]string, 0, len(candidates)),
|
||||
}
|
||||
for _, item := range candidates {
|
||||
trace.CandidateSkillCodes = append(trace.CandidateSkillCodes, item.Code)
|
||||
}
|
||||
if len(candidates) == 0 {
|
||||
trace.Status = "no_candidate"
|
||||
return nil, "no_enabled_skill_bound", trace, nil
|
||||
}
|
||||
|
||||
intentCode := strings.TrimSpace(ctx.IntentCode)
|
||||
if intentCode != "" {
|
||||
for _, item := range candidates {
|
||||
if strings.EqualFold(strings.TrimSpace(item.Code), intentCode) {
|
||||
trace.Status = "intent_selected"
|
||||
trace.SelectedSkillCode = item.Code
|
||||
return &item, "intent_code", trace, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if len(candidates) == 1 {
|
||||
trace.Status = "single_candidate"
|
||||
trace.SelectedSkillCode = candidates[0].Code
|
||||
return &candidates[0], "single_candidate", trace, nil
|
||||
}
|
||||
|
||||
selected, routeTrace, err := routeSkillWithLLM(execCtx, aiConfig, ctx.UserMessage, candidates)
|
||||
if routeTrace != nil {
|
||||
trace.Status = routeTrace.Status
|
||||
trace.SelectedSkillCode = routeTrace.SelectedSkillCode
|
||||
trace.RawDecision = routeTrace.RawDecision
|
||||
trace.LatencyMs = routeTrace.LatencyMs
|
||||
trace.Error = routeTrace.Error
|
||||
}
|
||||
if err != nil {
|
||||
if trace.Error == "" {
|
||||
trace.Error = err.Error()
|
||||
}
|
||||
return nil, "route_error", trace, err
|
||||
}
|
||||
if selected == nil {
|
||||
if trace.Status == "started" {
|
||||
trace.Status = "not_matched"
|
||||
}
|
||||
return nil, "route_none", trace, nil
|
||||
}
|
||||
return selected, "llm_route", trace, nil
|
||||
}
|
||||
|
||||
func loadCandidateSkills(aiAgent *models.AIAgent) []models.SkillDefinition {
|
||||
if aiAgent == nil {
|
||||
return nil
|
||||
}
|
||||
skillIDs := utils.SplitInt64s(aiAgent.SkillIDs)
|
||||
if len(skillIDs) == 0 {
|
||||
return nil
|
||||
}
|
||||
ret := make([]models.SkillDefinition, 0, len(skillIDs))
|
||||
for _, id := range skillIDs {
|
||||
skill := repositories.SkillDefinitionRepository.Get(sqls.DB(), id)
|
||||
if skill == nil || skill.Status != enums.StatusOk {
|
||||
continue
|
||||
}
|
||||
ret = append(ret, *skill)
|
||||
}
|
||||
return ret
|
||||
}
|
||||
|
||||
func routeSkillWithLLM(ctx context.Context, aiConfig *models.AIConfig, userMessage string, candidates []models.SkillDefinition) (*models.SkillDefinition, *RouteTrace, error) {
|
||||
trace := &RouteTrace{Status: "started"}
|
||||
if aiConfig == nil {
|
||||
trace.Status = "config_error"
|
||||
trace.Error = "ai config is nil"
|
||||
return nil, trace, errorsx.InvalidParam("Skill 路由依赖的 AI 配置不可用")
|
||||
}
|
||||
if len(candidates) == 0 {
|
||||
trace.Status = "no_candidate"
|
||||
return nil, trace, nil
|
||||
}
|
||||
userMessage = strings.TrimSpace(userMessage)
|
||||
if userMessage == "" {
|
||||
trace.Status = "empty_user_message"
|
||||
return nil, trace, nil
|
||||
}
|
||||
systemPrompt := "你是客服技能路由器。你只能在候选 Skill 中选择一个最合适的 skillCode,或者返回 NONE。只有当用户问题与 Skill 的职责边界明确匹配时才选择;如果不明确、信息不足、多个 Skill 都不够确定,就返回 NONE。输出只能是 skillCode 或 NONE,不能输出其他内容。"
|
||||
userPrompt := buildSkillRoutePrompt(userMessage, candidates)
|
||||
startedAt := time.Now()
|
||||
result, err := ai.LLM.ChatWithConfig(ctx, aiConfig, systemPrompt, userPrompt)
|
||||
trace.LatencyMs = time.Since(startedAt).Milliseconds()
|
||||
if err != nil {
|
||||
trace.Status = "route_error"
|
||||
trace.Error = err.Error()
|
||||
return nil, trace, err
|
||||
}
|
||||
decision := normalizeRouteDecision(result.Content)
|
||||
trace.RawDecision = strings.TrimSpace(result.Content)
|
||||
if decision == "" || decision == "NONE" {
|
||||
trace.Status = "not_matched"
|
||||
return nil, trace, nil
|
||||
}
|
||||
for _, item := range candidates {
|
||||
if strings.EqualFold(item.Code, decision) {
|
||||
trace.Status = "llm_selected"
|
||||
trace.SelectedSkillCode = item.Code
|
||||
return &item, trace, nil
|
||||
}
|
||||
}
|
||||
trace.Status = "invalid_decision"
|
||||
trace.Error = fmt.Sprintf("invalid route decision: %s", decision)
|
||||
return nil, trace, nil
|
||||
}
|
||||
|
||||
func buildSkillRoutePrompt(userMessage string, candidates []models.SkillDefinition) string {
|
||||
lines := make([]string, 0, len(candidates)+4)
|
||||
lines = append(lines, "用户问题:")
|
||||
lines = append(lines, strings.TrimSpace(userMessage))
|
||||
lines = append(lines, "")
|
||||
lines = append(lines, "候选 Skills:")
|
||||
for _, item := range candidates {
|
||||
lines = append(lines, fmt.Sprintf("- skillCode=%s; name=%s; description=%s", strings.TrimSpace(item.Code), strings.TrimSpace(item.Name), strings.TrimSpace(item.Description)))
|
||||
}
|
||||
lines = append(lines, "")
|
||||
lines = append(lines, "请只输出一个 skillCode 或 NONE。")
|
||||
return strings.Join(lines, "\n")
|
||||
}
|
||||
|
||||
func normalizeRouteDecision(raw string) string {
|
||||
raw = strings.TrimSpace(raw)
|
||||
if raw == "" {
|
||||
return ""
|
||||
}
|
||||
raw = strings.Trim(raw, "`")
|
||||
raw = strings.TrimSpace(raw)
|
||||
if idx := strings.Index(raw, "\n"); idx >= 0 {
|
||||
raw = raw[:idx]
|
||||
}
|
||||
raw = strings.TrimSpace(raw)
|
||||
raw = strings.Trim(raw, "\"'")
|
||||
if strings.EqualFold(raw, "NONE") {
|
||||
return "NONE"
|
||||
}
|
||||
return raw
|
||||
}
|
||||
@@ -0,0 +1,94 @@
|
||||
package skills
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/errorsx"
|
||||
"cs-agent/internal/repositories"
|
||||
|
||||
"github.com/mlogclub/simple/sqls"
|
||||
)
|
||||
|
||||
// BuildExecutionPlan 构建当前请求的 Skill 执行计划。
|
||||
func BuildExecutionPlan(execCtx context.Context, ctx RuntimeContext) (*ExecutionPlan, error) {
|
||||
if ctx.AIAgentID <= 0 {
|
||||
return nil, errorsx.InvalidParam("AIAgentID不能为空")
|
||||
}
|
||||
|
||||
aiAgent := repositories.AIAgentRepository.Get(sqls.DB(), ctx.AIAgentID)
|
||||
if aiAgent == nil {
|
||||
return nil, errorsx.InvalidParam("AI Agent不存在")
|
||||
}
|
||||
aiConfig := repositories.AIConfigRepository.Get(sqls.DB(), aiAgent.AIConfigID)
|
||||
if aiConfig == nil {
|
||||
return nil, errorsx.InvalidParam("AI Agent关联的AI配置不存在")
|
||||
}
|
||||
|
||||
skill, matchReason, routeTrace, err := MatchSkill(execCtx, ctx, aiAgent, aiConfig)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &ExecutionPlan{
|
||||
AIAgent: aiAgent,
|
||||
AIConfig: aiConfig,
|
||||
Skill: skill,
|
||||
MatchReason: strings.TrimSpace(matchReason),
|
||||
RouteTrace: routeTrace,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// WriteRunLog 写入 Skill 运行日志。
|
||||
func WriteRunLog(log *models.SkillRunLog) error {
|
||||
if log == nil {
|
||||
return nil
|
||||
}
|
||||
return repositories.SkillRunLogRepository.Create(sqls.DB(), log)
|
||||
}
|
||||
|
||||
// Execute 执行一次 Skill 运行,当前阶段仅支持 prompt_only 风格的手动 Skill。
|
||||
func Execute(ctx context.Context, runtimeCtx RuntimeContext) (*ExecutionResult, error) {
|
||||
plan, err := BuildExecutionPlan(ctx, runtimeCtx)
|
||||
if err != nil {
|
||||
trace := &ExecutionTrace{Status: "plan_error"}
|
||||
log := BuildRunLog(runtimeCtx, nil, trace, err)
|
||||
_ = WriteRunLog(log)
|
||||
return nil, err
|
||||
}
|
||||
if plan == nil || plan.Skill == nil {
|
||||
trace := &ExecutionTrace{Status: "noop"}
|
||||
if plan != nil {
|
||||
trace.MatchReason = strings.TrimSpace(plan.MatchReason)
|
||||
trace.Route = plan.RouteTrace
|
||||
}
|
||||
log := BuildRunLog(runtimeCtx, plan, trace, nil)
|
||||
_ = WriteRunLog(log)
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
replyText, trace, err := executeByPlan(ctx, plan, runtimeCtx)
|
||||
if trace != nil {
|
||||
trace.MatchReason = strings.TrimSpace(plan.MatchReason)
|
||||
if trace.Route == nil {
|
||||
trace.Route = plan.RouteTrace
|
||||
}
|
||||
}
|
||||
log := BuildRunLog(runtimeCtx, plan, trace, err)
|
||||
if strings.TrimSpace(replyText) != "" && strings.TrimSpace(log.MatchReason) == "" {
|
||||
log.MatchReason = string(plan.Skill.ExecutionMode)
|
||||
}
|
||||
if writeErr := WriteRunLog(log); writeErr != nil && err == nil {
|
||||
err = writeErr
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &ExecutionResult{
|
||||
Plan: plan,
|
||||
ReplyText: strings.TrimSpace(replyText),
|
||||
RunLog: log,
|
||||
Trace: trace,
|
||||
}, nil
|
||||
}
|
||||
@@ -0,0 +1,70 @@
|
||||
package skills
|
||||
|
||||
import "cs-agent/internal/models"
|
||||
|
||||
// RuntimeContext 表示一次 Skill 运行的输入上下文。
|
||||
type RuntimeContext struct {
|
||||
AIAgentID int64 // AIAgentID 为当前请求所属的 AI Agent ID,必填。
|
||||
UserMessage string // UserMessage 为当前用户输入。
|
||||
ConversationID int64 // ConversationID 为当前会话 ID,无会话上下文时为 0。
|
||||
ManualSkillCode string // ManualSkillCode 为显式指定的 Skill 编码。
|
||||
IntentCode string // IntentCode 为上游识别出的意图编码。
|
||||
}
|
||||
|
||||
// ExecutionPlan 表示 Skill Runtime 计算出的最终执行计划。
|
||||
type ExecutionPlan struct {
|
||||
AIAgent *models.AIAgent // AIAgent 为本次请求所属的 AI Agent。
|
||||
AIConfig *models.AIConfig // AIConfig 为本次请求实际使用的模型配置。
|
||||
Skill *models.SkillDefinition // Skill 为最终命中的 Skill,未命中时为空。
|
||||
MatchReason string // MatchReason 为命中原因。
|
||||
RouteTrace *RouteTrace // RouteTrace 为匹配阶段的路由追踪。
|
||||
}
|
||||
|
||||
// ExecutionResult 表示一次 Skill 执行的最终结果。
|
||||
type ExecutionResult struct {
|
||||
Plan *ExecutionPlan
|
||||
ReplyText string
|
||||
RunLog *models.SkillRunLog
|
||||
Trace *ExecutionTrace
|
||||
}
|
||||
|
||||
type ExecutionTrace struct {
|
||||
Status string `json:"status"`
|
||||
MatchReason string `json:"matchReason,omitempty"`
|
||||
Route *RouteTrace `json:"route,omitempty"`
|
||||
ExecutionMode string `json:"executionMode,omitempty"`
|
||||
Prompt *PromptTrace `json:"prompt,omitempty"`
|
||||
MCP *MCPExecutionTrace `json:"mcp,omitempty"`
|
||||
}
|
||||
|
||||
type RouteTrace struct {
|
||||
Status string `json:"status"`
|
||||
CandidateSkillCodes []string `json:"candidateSkillCodes,omitempty"`
|
||||
SelectedSkillCode string `json:"selectedSkillCode,omitempty"`
|
||||
RawDecision string `json:"rawDecision,omitempty"`
|
||||
LatencyMs int64 `json:"latencyMs,omitempty"`
|
||||
Error string `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
type PromptTrace struct {
|
||||
Status string `json:"status"`
|
||||
LatencyMs int64 `json:"latencyMs,omitempty"`
|
||||
ModelName string `json:"modelName,omitempty"`
|
||||
PromptTokens int `json:"promptTokens,omitempty"`
|
||||
CompletionTokens int `json:"completionTokens,omitempty"`
|
||||
Error string `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
type MCPExecutionTrace struct {
|
||||
Status string `json:"status"`
|
||||
ServerCode string `json:"serverCode,omitempty"`
|
||||
ToolName string `json:"toolName,omitempty"`
|
||||
Arguments map[string]any `json:"arguments,omitempty"`
|
||||
LatencyMs int64 `json:"latencyMs,omitempty"`
|
||||
IsError bool `json:"isError,omitempty"`
|
||||
ContentItemCount int `json:"contentItemCount,omitempty"`
|
||||
HasStructuredContent bool `json:"hasStructuredContent,omitempty"`
|
||||
ResultPreview string `json:"resultPreview,omitempty"`
|
||||
Error string `json:"error,omitempty"`
|
||||
SummaryPrompt *PromptTrace `json:"summaryPrompt,omitempty"`
|
||||
}
|
||||
@@ -0,0 +1,107 @@
|
||||
package bootstrap
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"cs-agent/internal/pkg/config"
|
||||
|
||||
"github.com/mlogclub/simple/sqls"
|
||||
"gorm.io/driver/mysql"
|
||||
|
||||
// "gorm.io/driver/sqlite" // Sqlite driver based on CGO
|
||||
"github.com/glebarez/sqlite" // Pure go SQLite driver, checkout https://github.com/glebarez/sqlite for details
|
||||
"gorm.io/gorm"
|
||||
|
||||
// "gorm.io/gorm/logger"
|
||||
"gorm.io/gorm/schema"
|
||||
)
|
||||
|
||||
func InitDB(cfg config.DBConfig) (*gorm.DB, error) {
|
||||
var dialector gorm.Dialector
|
||||
switch cfg.Type {
|
||||
case "sqlite":
|
||||
if err := ensureSQLiteDir(cfg.DSN); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
dialector = sqlite.Open(cfg.DSN)
|
||||
case "mysql":
|
||||
dialector = mysql.Open(cfg.DSN)
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported db type: %s", cfg.Type)
|
||||
}
|
||||
|
||||
db, err := gorm.Open(dialector, &gorm.Config{
|
||||
// Logger: logger.New(
|
||||
// log.New(os.Stdout, "\r\n", log.LstdFlags),
|
||||
// logger.Config{
|
||||
// SlowThreshold: time.Second,
|
||||
// LogLevel: logger.Info,
|
||||
// IgnoreRecordNotFoundError: true,
|
||||
// Colorful: true,
|
||||
// },
|
||||
// ),
|
||||
NamingStrategy: schema.NamingStrategy{
|
||||
TablePrefix: "t_",
|
||||
SingularTable: true,
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
sqlDB, err := db.DB()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if cfg.MaxIdleConns > 0 {
|
||||
sqlDB.SetMaxIdleConns(cfg.MaxIdleConns)
|
||||
}
|
||||
if cfg.MaxOpenConns > 0 {
|
||||
sqlDB.SetMaxOpenConns(cfg.MaxOpenConns)
|
||||
}
|
||||
if cfg.ConnMaxIdleTimeSeconds > 0 {
|
||||
sqlDB.SetConnMaxIdleTime(time.Duration(cfg.ConnMaxIdleTimeSeconds) * time.Second)
|
||||
}
|
||||
if cfg.ConnMaxLifetimeSeconds > 0 {
|
||||
sqlDB.SetConnMaxLifetime(time.Duration(cfg.ConnMaxLifetimeSeconds) * time.Second)
|
||||
}
|
||||
|
||||
sqls.SetDB(db)
|
||||
return db, nil
|
||||
}
|
||||
|
||||
func ensureSQLiteDir(dsn string) error {
|
||||
dbPath := sqliteFilePath(dsn)
|
||||
if dbPath == "" {
|
||||
return nil
|
||||
}
|
||||
dir := filepath.Dir(dbPath)
|
||||
if dir == "." || dir == "" {
|
||||
return nil
|
||||
}
|
||||
return os.MkdirAll(dir, 0o755)
|
||||
}
|
||||
|
||||
func sqliteFilePath(dsn string) string {
|
||||
if dsn == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
path := dsn
|
||||
if after, ok := strings.CutPrefix(path, "file:"); ok {
|
||||
path = after
|
||||
}
|
||||
if idx := strings.Index(path, "?"); idx >= 0 {
|
||||
path = path[:idx]
|
||||
}
|
||||
|
||||
normalized := strings.TrimSpace(path)
|
||||
if normalized == "" || normalized == ":memory:" || strings.Contains(normalized, "mode=memory") {
|
||||
return ""
|
||||
}
|
||||
return normalized
|
||||
}
|
||||
@@ -0,0 +1,66 @@
|
||||
package bootstrap
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestSQLiteFilePath(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
dsn string
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "plain relative path",
|
||||
dsn: "./data/app.db",
|
||||
want: "./data/app.db",
|
||||
},
|
||||
{
|
||||
name: "file uri with query",
|
||||
dsn: "file:./data/app.db?_busy_timeout=5000",
|
||||
want: "./data/app.db",
|
||||
},
|
||||
{
|
||||
name: "memory dsn",
|
||||
dsn: "file::memory:?cache=shared",
|
||||
want: "",
|
||||
},
|
||||
{
|
||||
name: "memory alias",
|
||||
dsn: ":memory:",
|
||||
want: "",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range cases {
|
||||
tt := tt
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
if got := sqliteFilePath(tt.dsn); got != tt.want {
|
||||
t.Fatalf("sqliteFilePath(%q) = %q, want %q", tt.dsn, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureSQLiteDir(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
baseDir := t.TempDir()
|
||||
dbPath := filepath.Join(baseDir, "nested", "app.db")
|
||||
dsn := "file:" + dbPath + "?_busy_timeout=5000"
|
||||
|
||||
if err := ensureSQLiteDir(dsn); err != nil {
|
||||
t.Fatalf("ensureSQLiteDir() error = %v", err)
|
||||
}
|
||||
|
||||
if info, err := os.Stat(filepath.Dir(dbPath)); err != nil {
|
||||
t.Fatalf("os.Stat() error = %v", err)
|
||||
} else if !info.IsDir() {
|
||||
t.Fatalf("expected %q to be a directory", filepath.Dir(dbPath))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,44 @@
|
||||
package bootstrap
|
||||
|
||||
import (
|
||||
"cs-agent/internal/ai/rag/vectordb"
|
||||
"cs-agent/internal/pkg/config"
|
||||
"cs-agent/internal/pkg/logx"
|
||||
"cs-agent/internal/services/cronx"
|
||||
"cs-agent/internal/wxwork"
|
||||
"log/slog"
|
||||
)
|
||||
|
||||
func Init(configPath string) error {
|
||||
cfg, err := config.Load(configPath)
|
||||
if err != nil {
|
||||
slog.Error("init config failed", "error", err)
|
||||
return err
|
||||
}
|
||||
config.SetCurrent(cfg)
|
||||
|
||||
logx.Init(logx.Config{
|
||||
Level: cfg.Logger.Level,
|
||||
Format: cfg.Logger.Format,
|
||||
AddSource: cfg.Logger.AddSource,
|
||||
})
|
||||
|
||||
if _, err := InitDB(cfg.DB); err != nil {
|
||||
slog.Error("init db failed", "error", err)
|
||||
return err
|
||||
}
|
||||
if err := InitMigrations(); err != nil {
|
||||
slog.Error("init migrations failed", "error", err)
|
||||
return err
|
||||
}
|
||||
if err := vectordb.Init(&cfg.VectorDB); err != nil {
|
||||
slog.Error("init vector db failed", "error", err)
|
||||
return err
|
||||
}
|
||||
|
||||
// 启动任务调度器
|
||||
cronx.Init()
|
||||
|
||||
wxwork.Init()
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
package bootstrap
|
||||
|
||||
import (
|
||||
"cs-agent/internal/migration"
|
||||
"cs-agent/internal/models"
|
||||
|
||||
"github.com/mlogclub/simple/sqls"
|
||||
)
|
||||
|
||||
func InitMigrations() error {
|
||||
if err := sqls.DB().AutoMigrate(models.Models...); err != nil {
|
||||
return err
|
||||
}
|
||||
return migration.Migrate()
|
||||
}
|
||||
@@ -0,0 +1,150 @@
|
||||
package bootstrap
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"cs-agent/internal/ai/mcps"
|
||||
_ "cs-agent/internal/ai/runtime"
|
||||
"cs-agent/internal/controllers/api"
|
||||
"cs-agent/internal/controllers/console"
|
||||
"cs-agent/internal/controllers/open"
|
||||
"cs-agent/internal/controllers/third"
|
||||
"cs-agent/internal/middleware"
|
||||
"cs-agent/internal/pkg/config"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/kataras/iris/v12/middleware/cors"
|
||||
"github.com/kataras/iris/v12/middleware/recover"
|
||||
"github.com/kataras/iris/v12/mvc"
|
||||
|
||||
_ "cs-agent/internal/services/wx_callback_handlers"
|
||||
)
|
||||
|
||||
func NewServer() (*iris.Application, error) {
|
||||
cfg := config.Current()
|
||||
app := iris.New()
|
||||
corsHandler := cors.New().
|
||||
AllowOrigin("*").
|
||||
AllowHeaders("Origin", "Content-Type", "Accept", "Authorization", "X-Requested-With", "X-Visitor-Id", "X-Channel-Id", "X-External-Source", "X-External-Id", "X-External-Name").
|
||||
MaxAge(600).
|
||||
ExposeHeaders("Content-Length", "Content-Type", "Authorization", "X-Visitor-Id", "X-Channel-Id", "X-External-Source", "X-External-Id", "X-External-Name").
|
||||
Handler()
|
||||
app.UseRouter(func(ctx iris.Context) {
|
||||
// WebSocket upgrade is validated by the upgrader's origin policy.
|
||||
if isWebsocketUpgrade(ctx) {
|
||||
ctx.Next()
|
||||
return
|
||||
}
|
||||
corsHandler(ctx)
|
||||
})
|
||||
app.UseRouter(recover.New())
|
||||
app.UseRouter(func(ctx iris.Context) {
|
||||
start := time.Now()
|
||||
path := ctx.Path()
|
||||
method := ctx.Method()
|
||||
ctx.Next()
|
||||
|
||||
slog.Info("http request",
|
||||
"method", method,
|
||||
"path", path,
|
||||
"status", ctx.GetStatusCode(),
|
||||
"elapsed", time.Since(start).Milliseconds(),
|
||||
"clientIp", ctx.RemoteAddr(),
|
||||
)
|
||||
})
|
||||
app.UseRouter(func(ctx iris.Context) {
|
||||
ctx.SetMaxRequestBodySize(cfg.Storage.MaxRequestBodySizeBytes())
|
||||
ctx.Next()
|
||||
})
|
||||
|
||||
// 注册路由
|
||||
addRouter(app)
|
||||
|
||||
// 注册本地存储静态资源服务
|
||||
app.HandleDir(cfg.Storage.Local.BaseURL, iris.Dir(cfg.Storage.Local.Root), iris.DirOptions{
|
||||
ShowList: false,
|
||||
})
|
||||
|
||||
// 注册web静态资源服务
|
||||
app.HandleDir("/", iris.Dir("web/out"), iris.DirOptions{
|
||||
IndexName: "index.html",
|
||||
Compress: true,
|
||||
ShowList: false,
|
||||
})
|
||||
|
||||
// 注册widget静态资源服务
|
||||
app.HandleDir("/widget", iris.Dir("widget/out"), iris.DirOptions{
|
||||
IndexName: "index.html",
|
||||
Compress: true,
|
||||
ShowList: false,
|
||||
})
|
||||
|
||||
return app, nil
|
||||
}
|
||||
|
||||
func isWebsocketUpgrade(ctx iris.Context) bool {
|
||||
if !strings.EqualFold(ctx.GetHeader("Upgrade"), "websocket") {
|
||||
return false
|
||||
}
|
||||
return strings.Contains(strings.ToLower(ctx.GetHeader("Connection")), "upgrade")
|
||||
}
|
||||
|
||||
func addRouter(app *iris.Application) {
|
||||
mcpHandler := mcps.NewHTTPHandler()
|
||||
|
||||
app.Any("/api/mcp", iris.FromStd(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
mcpHandler.ServeHTTP(w, r)
|
||||
})))
|
||||
|
||||
app.Get("/api/console/ws", middleware.AuthMiddleware, middleware.ConsoleWsMiddleware)
|
||||
app.Get("/api/open/im/ws", middleware.OpenImWsMiddleware)
|
||||
|
||||
mvc.Configure(app.Party("/api/auth"), func(m *mvc.Application) {
|
||||
m.Handle(new(api.AuthController))
|
||||
})
|
||||
|
||||
mvc.Configure(app.Party("/api/console", middleware.AuthMiddleware), func(m *mvc.Application) {
|
||||
m.Party("/dashboard").Handle(new(console.DashboardController))
|
||||
m.Party("/user").Handle(new(console.UserController))
|
||||
m.Party("/company").Handle(new(console.CompanyController))
|
||||
m.Party("/customer").Handle(new(console.CustomerController))
|
||||
m.Party("/customer-contact").Handle(new(console.CustomerContactController))
|
||||
m.Party("/role").Handle(new(console.RoleController))
|
||||
m.Party("/permission").Handle(new(console.PermissionController))
|
||||
m.Party("/session").Handle(new(console.SessionController))
|
||||
m.Party("/tag").Handle(new(console.TagController))
|
||||
m.Party("/conversation").Handle(new(console.ConversationController))
|
||||
m.Party("/ticket").Handle(new(console.TicketController))
|
||||
m.Party("/ticket-resolution-code").Handle(new(console.TicketResolutionCodeController))
|
||||
m.Party("/ticket-priority-config").Handle(new(console.TicketPriorityConfigController))
|
||||
m.Party("/quick-reply").Handle(new(console.QuickReplyController))
|
||||
m.Party("/channel").Handle(new(console.ChannelController))
|
||||
m.Party("/agent").Handle(new(console.AgentController))
|
||||
m.Party("/agent-team").Handle(new(console.AgentTeamController))
|
||||
m.Party("/agent-team-schedule").Handle(new(console.AgentTeamScheduleController))
|
||||
m.Party("/ai-agent").Handle(new(console.AIAgentController))
|
||||
m.Party("/ai-config").Handle(new(console.AIConfigController))
|
||||
m.Party("/asset").Handle(new(console.AssetController))
|
||||
m.Party("/knowledge-base").Handle(new(console.KnowledgeBaseController))
|
||||
m.Party("/knowledge-document").Handle(new(console.KnowledgeDocumentController))
|
||||
m.Party("/knowledge-faq").Handle(new(console.KnowledgeFAQController))
|
||||
m.Party("/knowledge-retrieve").Handle(new(console.KnowledgeRetrieveController))
|
||||
m.Party("/knowledge-retrieve-log").Handle(new(console.KnowledgeRetrieveLogController))
|
||||
m.Party("/agent-run-log").Handle(new(console.AgentRunLogController))
|
||||
m.Party("/skill-definition").Handle(new(console.SkillDefinitionController))
|
||||
m.Party("/mcp").Handle(new(console.MCPController))
|
||||
})
|
||||
|
||||
mvc.Configure(app.Party("/api/open/im", middleware.OpenImContextMiddleware), func(m *mvc.Application) {
|
||||
m.Party("/widget").Handle(new(open.ImWidgetController))
|
||||
m.Party("/conversation").Handle(new(open.ImConversationController))
|
||||
m.Party("/message").Handle(new(open.ImMessageController))
|
||||
})
|
||||
|
||||
mvc.Configure(app.Party("/api/third"), func(m *mvc.Application) {
|
||||
m.Party("/wechat").Handle(new(third.WechatController))
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,82 @@
|
||||
package builders
|
||||
|
||||
import (
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/pkg/utils"
|
||||
"cs-agent/internal/services"
|
||||
)
|
||||
|
||||
func BuildAgentProfileList(items []models.AgentProfile) []response.AgentProfileResponse {
|
||||
if len(items) == 0 {
|
||||
return []response.AgentProfileResponse{}
|
||||
}
|
||||
|
||||
userIDs := make([]int64, 0, len(items))
|
||||
teamIDs := make([]int64, 0, len(items))
|
||||
for _, item := range items {
|
||||
if item.UserID > 0 {
|
||||
userIDs = append(userIDs, item.UserID)
|
||||
}
|
||||
if item.TeamID > 0 {
|
||||
teamIDs = append(teamIDs, item.TeamID)
|
||||
}
|
||||
}
|
||||
|
||||
users := services.UserService.FindByIds(userIDs)
|
||||
teams := services.AgentTeamService.FindByIds(teamIDs)
|
||||
|
||||
userMap := make(map[int64]*models.User, len(users))
|
||||
for i := range users {
|
||||
userMap[users[i].ID] = &users[i]
|
||||
}
|
||||
|
||||
teamMap := make(map[int64]*models.AgentTeam, len(teams))
|
||||
for i := range teams {
|
||||
teamMap[teams[i].ID] = &teams[i]
|
||||
}
|
||||
|
||||
results := make([]response.AgentProfileResponse, 0, len(items))
|
||||
for _, item := range items {
|
||||
if result := doBuildAgentProfileResponse(&item, userMap[item.UserID], teamMap[item.TeamID]); result != nil {
|
||||
results = append(results, *result)
|
||||
}
|
||||
}
|
||||
return results
|
||||
}
|
||||
|
||||
func BuildAgentProfileResponse(item *models.AgentProfile) *response.AgentProfileResponse {
|
||||
user := services.UserService.Get(item.UserID)
|
||||
team := services.AgentTeamService.Get(item.TeamID)
|
||||
return doBuildAgentProfileResponse(item, user, team)
|
||||
}
|
||||
|
||||
func doBuildAgentProfileResponse(item *models.AgentProfile, user *models.User, team *models.AgentTeam) *response.AgentProfileResponse {
|
||||
if item == nil {
|
||||
return nil
|
||||
}
|
||||
ret := &response.AgentProfileResponse{
|
||||
ID: item.ID,
|
||||
UserID: item.UserID,
|
||||
TeamID: item.TeamID,
|
||||
AgentCode: item.AgentCode,
|
||||
DisplayName: item.DisplayName,
|
||||
Avatar: item.Avatar,
|
||||
ServiceStatus: item.ServiceStatus,
|
||||
MaxConcurrentCount: item.MaxConcurrentCount,
|
||||
PriorityLevel: item.PriorityLevel,
|
||||
AutoAssignEnabled: item.AutoAssignEnabled,
|
||||
ReceiveOfflineMessage: item.ReceiveOfflineMessage,
|
||||
LastOnlineAt: utils.FormatTimePtr(item.LastOnlineAt),
|
||||
LastStatusAt: utils.FormatTimePtr(item.LastStatusAt),
|
||||
Remark: item.Remark,
|
||||
}
|
||||
if user != nil {
|
||||
ret.Username = user.Username
|
||||
ret.Nickname = user.Nickname
|
||||
}
|
||||
if team != nil {
|
||||
ret.TeamName = team.Name
|
||||
}
|
||||
return ret
|
||||
}
|
||||
@@ -0,0 +1,30 @@
|
||||
package builders
|
||||
|
||||
import (
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
)
|
||||
|
||||
func BuildAgentRunLog(item *models.AgentRunLog) response.AgentRunLogResponse {
|
||||
if item == nil {
|
||||
return response.AgentRunLogResponse{}
|
||||
}
|
||||
return response.AgentRunLogResponse{
|
||||
ID: item.ID,
|
||||
ConversationID: item.ConversationID,
|
||||
MessageID: item.MessageID,
|
||||
AIAgentID: item.AIAgentID,
|
||||
AIConfigID: item.AIConfigID,
|
||||
UserMessage: item.UserMessage,
|
||||
PlannedAction: item.PlannedAction,
|
||||
PlannedSkillCode: item.PlannedSkillCode,
|
||||
PlannedToolCode: item.PlannedToolCode,
|
||||
PlanReason: item.PlanReason,
|
||||
FinalAction: item.FinalAction,
|
||||
ReplyText: item.ReplyText,
|
||||
ErrorMessage: item.ErrorMessage,
|
||||
LatencyMs: item.LatencyMs,
|
||||
TraceData: item.TraceData,
|
||||
CreatedAt: item.CreatedAt.Format("2006-01-02 15:04:05"),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
package builders
|
||||
|
||||
import (
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
)
|
||||
|
||||
func BuildAsset(item *models.Asset) response.AssetResponse {
|
||||
ret := response.AssetResponse{
|
||||
ID: item.ID,
|
||||
AssetID: item.AssetID,
|
||||
Provider: item.Provider,
|
||||
Filename: item.Filename,
|
||||
FileSize: item.FileSize,
|
||||
MimeType: item.MimeType,
|
||||
URL: item.URL,
|
||||
StorageKey: item.StorageKey,
|
||||
Status: item.Status,
|
||||
CreatedAt: item.CreatedAt.Format("2006-01-02 15:04:05"),
|
||||
UpdatedAt: item.UpdatedAt.Format("2006-01-02 15:04:05"),
|
||||
CreateUserID: item.CreateUserID,
|
||||
CreateUserName: item.CreateUserName,
|
||||
UpdateUserID: item.UpdateUserID,
|
||||
UpdateUserName: item.UpdateUserName,
|
||||
}
|
||||
return ret
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
package builders
|
||||
|
||||
import (
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"time"
|
||||
)
|
||||
|
||||
func BuildCompany(item *models.Company) *response.CompanyResponse {
|
||||
if item == nil {
|
||||
return nil
|
||||
}
|
||||
return &response.CompanyResponse{
|
||||
ID: item.ID,
|
||||
Name: item.Name,
|
||||
Code: item.Code,
|
||||
Status: item.Status,
|
||||
Remark: item.Remark,
|
||||
CreatedAt: item.CreatedAt.Format(time.DateTime),
|
||||
UpdatedAt: item.UpdatedAt.Format(time.DateTime),
|
||||
}
|
||||
}
|
||||
|
||||
func BuildCompanyList(list []models.Company) []response.CompanyResponse {
|
||||
results := make([]response.CompanyResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
if company := BuildCompany(&item); company != nil {
|
||||
results = append(results, *company)
|
||||
}
|
||||
}
|
||||
return results
|
||||
}
|
||||
@@ -0,0 +1,267 @@
|
||||
package builders
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"cs-agent/internal/pkg/utils"
|
||||
"cs-agent/internal/services"
|
||||
|
||||
"github.com/mlogclub/simple/sqls"
|
||||
)
|
||||
|
||||
func BuildConversation(item *models.Conversation) response.ConversationResponse {
|
||||
agentReadState, customerReadState := services.ConversationReadStateService.GetConversationReadStates(item.ID)
|
||||
ret := response.ConversationResponse{
|
||||
ID: item.ID,
|
||||
AIAgentID: item.AIAgentID,
|
||||
CustomerID: item.CustomerID,
|
||||
ExternalSource: item.ExternalSource,
|
||||
ExternalID: item.ExternalID,
|
||||
Subject: item.Subject,
|
||||
Status: item.Status,
|
||||
ServiceMode: item.ServiceMode,
|
||||
Priority: item.Priority,
|
||||
CurrentAssigneeID: item.CurrentAssigneeID,
|
||||
LastMessageID: item.LastMessageID,
|
||||
LastMessageAt: utils.FormatTime(item.LastMessageAt),
|
||||
LastActiveAt: utils.FormatTime(item.LastActiveAt),
|
||||
LastMessageSummary: item.LastMessageSummary,
|
||||
CustomerUnreadCount: item.CustomerUnreadCount,
|
||||
AgentUnreadCount: item.AgentUnreadCount,
|
||||
CustomerLastReadMessageID: readStateMessageID(customerReadState),
|
||||
CustomerLastReadSeqNo: readStateSeqNo(customerReadState),
|
||||
CustomerLastReadAt: readStateAt(customerReadState),
|
||||
AgentLastReadMessageID: readStateMessageID(agentReadState),
|
||||
AgentLastReadSeqNo: readStateSeqNo(agentReadState),
|
||||
AgentLastReadAt: readStateAt(agentReadState),
|
||||
ClosedAt: utils.FormatTimePtr(item.ClosedAt),
|
||||
ClosedBy: item.ClosedBy,
|
||||
CloseReason: item.CloseReason,
|
||||
}
|
||||
if item.CurrentAssigneeID > 0 {
|
||||
if user := services.UserService.Get(item.CurrentAssigneeID); user != nil {
|
||||
ret.CurrentAssigneeName = user.Nickname
|
||||
if ret.CurrentAssigneeName == "" {
|
||||
ret.CurrentAssigneeName = user.Username
|
||||
}
|
||||
}
|
||||
}
|
||||
if item.ClosedBy > 0 {
|
||||
if user := services.UserService.Get(item.ClosedBy); user != nil {
|
||||
ret.ClosedByName = user.Nickname
|
||||
if ret.ClosedByName == "" {
|
||||
ret.ClosedByName = user.Username
|
||||
}
|
||||
}
|
||||
}
|
||||
return ret
|
||||
}
|
||||
|
||||
func BuildParticipantResponses(conversationID int64) []response.ConversationParticipantResponse {
|
||||
list := services.ConversationParticipantService.Find(sqls.NewCnd().Eq("conversation_id", conversationID).Asc("id"))
|
||||
if len(list) == 0 {
|
||||
return nil
|
||||
}
|
||||
ret := make([]response.ConversationParticipantResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
ret = append(ret, response.ConversationParticipantResponse{
|
||||
ID: item.ID,
|
||||
ParticipantType: item.ParticipantType,
|
||||
ParticipantID: item.ParticipantID,
|
||||
ExternalParticipantID: item.ExternalParticipantID,
|
||||
JoinedAt: utils.FormatTimePtr(item.JoinedAt),
|
||||
LeftAt: utils.FormatTimePtr(item.LeftAt),
|
||||
Status: item.Status,
|
||||
})
|
||||
}
|
||||
return ret
|
||||
}
|
||||
|
||||
func BuildMessages(list []models.Message) []response.MessageResponse {
|
||||
if len(list) == 0 {
|
||||
return nil
|
||||
}
|
||||
agentReadState, customerReadState := services.ConversationReadStateService.GetConversationReadStates(list[0].ConversationID)
|
||||
aiSenderNames, userSenderNames := collectMessageSenderNameMaps(list)
|
||||
agentProfiles := collectAgentProfilesByMessages(list)
|
||||
ret := make([]response.MessageResponse, 0, len(list))
|
||||
for i := range list {
|
||||
ret = append(ret, BuildMessageWithReadStates(&list[i], agentReadState, customerReadState, aiSenderNames, userSenderNames, agentProfiles))
|
||||
}
|
||||
return ret
|
||||
}
|
||||
|
||||
func BuildMessage(item *models.Message) response.MessageResponse {
|
||||
agentReadState, customerReadState := services.ConversationReadStateService.GetConversationReadStates(item.ConversationID)
|
||||
return BuildMessageWithReadStates(item, agentReadState, customerReadState, nil, nil, nil)
|
||||
}
|
||||
|
||||
func BuildMessageWithReadStates(item *models.Message, agentReadState, customerReadState *models.ConversationReadState, aiSenderNames, userSenderNames map[int64]string, agentProfiles map[int64]*models.AgentProfile) response.MessageResponse {
|
||||
content, payload := buildMessageResponseContent(item)
|
||||
ret := response.MessageResponse{
|
||||
ID: item.ID,
|
||||
ConversationID: item.ConversationID,
|
||||
ClientMsgID: item.ClientMsgID,
|
||||
SenderType: item.SenderType,
|
||||
SenderID: item.SenderID,
|
||||
MessageType: item.MessageType,
|
||||
Content: content,
|
||||
Payload: payload,
|
||||
SeqNo: item.SeqNo,
|
||||
SendStatus: item.SendStatus,
|
||||
SentAt: utils.FormatTimePtr(item.SentAt),
|
||||
DeliveredAt: utils.FormatTimePtr(item.DeliveredAt),
|
||||
ReadAt: utils.FormatTimePtr(item.ReadAt),
|
||||
CustomerRead: isMessageRead(item, customerReadState),
|
||||
CustomerReadAt: readMessageAt(item, customerReadState),
|
||||
AgentRead: isMessageRead(item, agentReadState),
|
||||
AgentReadAt: readMessageAt(item, agentReadState),
|
||||
RecalledAt: utils.FormatTimePtr(item.RecalledAt),
|
||||
QuotedMessageID: item.QuotedMessageID,
|
||||
}
|
||||
if item.SenderID > 0 {
|
||||
if item.SenderType == enums.IMSenderTypeAI {
|
||||
if aiSenderNames != nil {
|
||||
ret.SenderName = aiSenderNames[item.SenderID]
|
||||
} else if aiAgent := services.AIAgentService.Get(item.SenderID); aiAgent != nil {
|
||||
ret.SenderName = aiAgent.Name
|
||||
}
|
||||
} else if item.SenderType == enums.IMSenderTypeAgent {
|
||||
profile := agentProfiles[item.SenderID]
|
||||
if profile != nil {
|
||||
if dn := strings.TrimSpace(profile.DisplayName); dn != "" {
|
||||
ret.SenderName = dn
|
||||
}
|
||||
if av := strings.TrimSpace(profile.Avatar); av != "" {
|
||||
ret.SenderAvatar = av
|
||||
}
|
||||
}
|
||||
if ret.SenderName == "" {
|
||||
if userSenderNames != nil {
|
||||
ret.SenderName = userSenderNames[item.SenderID]
|
||||
} else if user := services.UserService.Get(item.SenderID); user != nil {
|
||||
ret.SenderName = user.Nickname
|
||||
if ret.SenderName == "" {
|
||||
ret.SenderName = user.Username
|
||||
}
|
||||
}
|
||||
}
|
||||
} else if userSenderNames != nil {
|
||||
ret.SenderName = userSenderNames[item.SenderID]
|
||||
} else if user := services.UserService.Get(item.SenderID); user != nil {
|
||||
ret.SenderName = user.Nickname
|
||||
if ret.SenderName == "" {
|
||||
ret.SenderName = user.Username
|
||||
}
|
||||
}
|
||||
}
|
||||
return ret
|
||||
}
|
||||
|
||||
func buildMessageResponseContent(item *models.Message) (content, payload string) {
|
||||
if item == nil {
|
||||
return "", ""
|
||||
}
|
||||
if item.RecalledAt != nil || item.SendStatus == int(enums.IMMessageStatusRecalled) {
|
||||
return "该消息已撤回", ""
|
||||
}
|
||||
return item.Content, item.Payload
|
||||
}
|
||||
|
||||
func collectAgentProfilesByMessages(list []models.Message) map[int64]*models.AgentProfile {
|
||||
var agentUserIDs []int64
|
||||
seen := make(map[int64]struct{})
|
||||
for i := range list {
|
||||
m := &list[i]
|
||||
if m.SenderType != enums.IMSenderTypeAgent || m.SenderID <= 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[m.SenderID]; ok {
|
||||
continue
|
||||
}
|
||||
seen[m.SenderID] = struct{}{}
|
||||
agentUserIDs = append(agentUserIDs, m.SenderID)
|
||||
}
|
||||
if len(agentUserIDs) == 0 {
|
||||
return nil
|
||||
}
|
||||
profiles := services.AgentProfileService.Find(sqls.NewCnd().In("user_id", agentUserIDs))
|
||||
out := make(map[int64]*models.AgentProfile, len(profiles))
|
||||
for i := range profiles {
|
||||
out[profiles[i].UserID] = &profiles[i]
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func collectMessageSenderNameMaps(list []models.Message) (aiNames map[int64]string, userNames map[int64]string) {
|
||||
aiNames = make(map[int64]string)
|
||||
userNames = make(map[int64]string)
|
||||
var aiIDs, userIDs []int64
|
||||
seenAI := make(map[int64]struct{})
|
||||
seenUser := make(map[int64]struct{})
|
||||
for i := range list {
|
||||
m := &list[i]
|
||||
if m.SenderID <= 0 {
|
||||
continue
|
||||
}
|
||||
if m.SenderType == enums.IMSenderTypeAI {
|
||||
if _, ok := seenAI[m.SenderID]; ok {
|
||||
continue
|
||||
}
|
||||
seenAI[m.SenderID] = struct{}{}
|
||||
aiIDs = append(aiIDs, m.SenderID)
|
||||
continue
|
||||
}
|
||||
if _, ok := seenUser[m.SenderID]; ok {
|
||||
continue
|
||||
}
|
||||
seenUser[m.SenderID] = struct{}{}
|
||||
userIDs = append(userIDs, m.SenderID)
|
||||
}
|
||||
for _, a := range services.AIAgentService.FindByIds(aiIDs) {
|
||||
aiNames[a.ID] = a.Name
|
||||
}
|
||||
for _, u := range services.UserService.FindByIds(userIDs) {
|
||||
name := u.Nickname
|
||||
if name == "" {
|
||||
name = u.Username
|
||||
}
|
||||
userNames[u.ID] = name
|
||||
}
|
||||
return aiNames, userNames
|
||||
}
|
||||
|
||||
func isMessageRead(item *models.Message, state *models.ConversationReadState) bool {
|
||||
return item != nil && state != nil && state.LastReadSeqNo >= item.SeqNo
|
||||
}
|
||||
|
||||
func readMessageAt(item *models.Message, state *models.ConversationReadState) string {
|
||||
if !isMessageRead(item, state) {
|
||||
return ""
|
||||
}
|
||||
return utils.FormatTimePtr(state.LastReadAt)
|
||||
}
|
||||
|
||||
func readStateMessageID(state *models.ConversationReadState) int64 {
|
||||
if state == nil {
|
||||
return 0
|
||||
}
|
||||
return state.LastReadMessageID
|
||||
}
|
||||
|
||||
func readStateSeqNo(state *models.ConversationReadState) int64 {
|
||||
if state == nil {
|
||||
return 0
|
||||
}
|
||||
return state.LastReadSeqNo
|
||||
}
|
||||
|
||||
func readStateAt(state *models.ConversationReadState) string {
|
||||
if state == nil {
|
||||
return ""
|
||||
}
|
||||
return utils.FormatTimePtr(state.LastReadAt)
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
package builders
|
||||
|
||||
import (
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/pkg/utils"
|
||||
"cs-agent/internal/services"
|
||||
"time"
|
||||
)
|
||||
|
||||
func BuildCustomer(item *models.Customer) *response.CustomerResponse {
|
||||
if item == nil {
|
||||
return nil
|
||||
}
|
||||
return &response.CustomerResponse{
|
||||
ID: item.ID,
|
||||
Name: item.Name,
|
||||
Gender: item.Gender,
|
||||
CompanyID: item.CompanyID,
|
||||
Company: BuildCompany(services.CompanyService.Get(item.CompanyID)),
|
||||
LastActiveAt: utils.FormatTimePtr(item.LastActiveAt),
|
||||
PrimaryMobile: item.PrimaryMobile,
|
||||
PrimaryEmail: item.PrimaryEmail,
|
||||
Status: item.Status,
|
||||
Remark: item.Remark,
|
||||
CreatedAt: item.CreatedAt.Format(time.DateTime),
|
||||
UpdatedAt: item.UpdatedAt.Format(time.DateTime),
|
||||
}
|
||||
}
|
||||
|
||||
func BuildCustomerList(list []models.Customer) []response.CustomerResponse {
|
||||
results := make([]response.CustomerResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
if customer := BuildCustomer(&item); customer != nil {
|
||||
results = append(results, *customer)
|
||||
}
|
||||
}
|
||||
return results
|
||||
}
|
||||
@@ -0,0 +1,36 @@
|
||||
package builders
|
||||
|
||||
import (
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/pkg/utils"
|
||||
"time"
|
||||
)
|
||||
|
||||
func BuildCustomerContactResponse(item *models.CustomerContact) response.CustomerContactResponse {
|
||||
if item == nil {
|
||||
return response.CustomerContactResponse{}
|
||||
}
|
||||
return response.CustomerContactResponse{
|
||||
ID: item.ID,
|
||||
CustomerID: item.CustomerID,
|
||||
ContactType: item.ContactType,
|
||||
ContactValue: item.ContactValue,
|
||||
IsPrimary: item.IsPrimary,
|
||||
IsVerified: item.IsVerified,
|
||||
VerifiedAt: utils.FormatTimePtr(item.VerifiedAt),
|
||||
Source: item.Source,
|
||||
Status: item.Status,
|
||||
Remark: item.Remark,
|
||||
CreatedAt: item.CreatedAt.Format(time.DateTime),
|
||||
UpdatedAt: item.UpdatedAt.Format(time.DateTime),
|
||||
}
|
||||
}
|
||||
|
||||
func BuildCustomerContactList(list []models.CustomerContact) []response.CustomerContactResponse {
|
||||
results := make([]response.CustomerContactResponse, 0, len(list))
|
||||
for i := range list {
|
||||
results = append(results, BuildCustomerContactResponse(&list[i]))
|
||||
}
|
||||
return results
|
||||
}
|
||||
@@ -0,0 +1,152 @@
|
||||
package builders
|
||||
|
||||
import (
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"encoding/json"
|
||||
)
|
||||
|
||||
func BuildKnowledgeBase(item *models.KnowledgeBase) response.KnowledgeBaseResponse {
|
||||
return response.KnowledgeBaseResponse{
|
||||
ID: item.ID,
|
||||
Name: item.Name,
|
||||
Description: item.Description,
|
||||
KnowledgeType: item.KnowledgeType,
|
||||
KnowledgeTypeName: enums.GetKnowledgeBaseTypeLabel(enums.KnowledgeBaseType(item.KnowledgeType)),
|
||||
Status: item.Status,
|
||||
StatusName: enums.GetStatusLabel(item.Status),
|
||||
DefaultTopK: item.DefaultTopK,
|
||||
DefaultScoreThreshold: item.DefaultScoreThreshold,
|
||||
DefaultRerankLimit: item.DefaultRerankLimit,
|
||||
ChunkProvider: item.ChunkProvider,
|
||||
ChunkTargetTokens: item.ChunkTargetTokens,
|
||||
ChunkMaxTokens: item.ChunkMaxTokens,
|
||||
ChunkOverlapTokens: item.ChunkOverlapTokens,
|
||||
AnswerMode: item.AnswerMode,
|
||||
AnswerModeName: enums.GetKnowledgeAnswerModeLabel(enums.KnowledgeAnswerMode(item.AnswerMode)),
|
||||
FallbackMode: item.FallbackMode,
|
||||
FallbackModeName: enums.GetKnowledgeFallbackModeLabel(enums.KnowledgeFallbackMode(item.FallbackMode)),
|
||||
Remark: item.Remark,
|
||||
CreatedAt: item.CreatedAt,
|
||||
UpdatedAt: item.UpdatedAt,
|
||||
CreateUserName: item.CreateUserName,
|
||||
UpdateUserName: item.UpdateUserName,
|
||||
}
|
||||
}
|
||||
|
||||
func BuildKnowledgeDocument(item *models.KnowledgeDocument) response.KnowledgeDocumentResponse {
|
||||
return response.KnowledgeDocumentResponse{
|
||||
ID: item.ID,
|
||||
KnowledgeBaseID: item.KnowledgeBaseID,
|
||||
Title: item.Title,
|
||||
Status: item.Status,
|
||||
StatusName: enums.GetStatusLabel(item.Status),
|
||||
IndexStatus: item.IndexStatus,
|
||||
IndexStatusName: enums.GetKnowledgeDocumentIndexStatusLabel(item.IndexStatus),
|
||||
IndexedAt: item.IndexedAt,
|
||||
IndexError: item.IndexError,
|
||||
ContentHash: item.ContentHash,
|
||||
ContentType: item.ContentType,
|
||||
Content: item.Content,
|
||||
CreatedAt: item.CreatedAt,
|
||||
UpdatedAt: item.UpdatedAt,
|
||||
CreateUserName: item.CreateUserName,
|
||||
UpdateUserName: item.UpdateUserName,
|
||||
}
|
||||
}
|
||||
|
||||
func BuildKnowledgeFAQ(item *models.KnowledgeFAQ) response.KnowledgeFAQResponse {
|
||||
return response.KnowledgeFAQResponse{
|
||||
ID: item.ID,
|
||||
KnowledgeBaseID: item.KnowledgeBaseID,
|
||||
Question: item.Question,
|
||||
Answer: item.Answer,
|
||||
SimilarQuestions: parseSimilarQuestions(item.SimilarQuestions),
|
||||
Status: item.Status,
|
||||
StatusName: enums.GetStatusLabel(item.Status),
|
||||
IndexStatus: item.IndexStatus,
|
||||
IndexStatusName: enums.GetKnowledgeDocumentIndexStatusLabel(item.IndexStatus),
|
||||
IndexedAt: item.IndexedAt,
|
||||
IndexError: item.IndexError,
|
||||
Remark: item.Remark,
|
||||
CreatedAt: item.CreatedAt,
|
||||
UpdatedAt: item.UpdatedAt,
|
||||
CreateUserName: item.CreateUserName,
|
||||
UpdateUserName: item.UpdateUserName,
|
||||
}
|
||||
}
|
||||
|
||||
func BuildKnowledgeRetrieveLog(item *models.KnowledgeRetrieveLog) response.KnowledgeRetrieveLogResponse {
|
||||
return response.KnowledgeRetrieveLogResponse{
|
||||
ID: item.ID,
|
||||
KnowledgeBaseID: item.KnowledgeBaseID,
|
||||
Channel: item.Channel,
|
||||
ChannelName: enums.GetKnowledgeRetrieveChannelLabel(enums.KnowledgeRetrieveChannel(item.Channel)),
|
||||
Scene: item.Scene,
|
||||
SceneName: enums.GetKnowledgeRetrieveSceneLabel(enums.KnowledgeRetrieveScene(item.Scene)),
|
||||
SessionID: item.SessionID,
|
||||
ConversationID: item.ConversationID,
|
||||
RequestID: item.RequestID,
|
||||
Question: item.Question,
|
||||
RewriteQuestion: item.RewriteQuestion,
|
||||
Answer: item.Answer,
|
||||
AnswerStatus: item.AnswerStatus,
|
||||
AnswerStatusName: enums.GetKnowledgeAnswerStatusLabel(enums.KnowledgeAnswerStatus(item.AnswerStatus)),
|
||||
HitCount: item.HitCount,
|
||||
TopScore: item.TopScore,
|
||||
ChunkProvider: item.ChunkProvider,
|
||||
ChunkTargetTokens: item.ChunkTargetTokens,
|
||||
ChunkMaxTokens: item.ChunkMaxTokens,
|
||||
ChunkOverlapTokens: item.ChunkOverlapTokens,
|
||||
RerankEnabled: item.RerankEnabled,
|
||||
RerankLimit: item.RerankLimit,
|
||||
CitationCount: item.CitationCount,
|
||||
UsedChunkCount: item.UsedChunkCount,
|
||||
LatencyMs: item.LatencyMs,
|
||||
RetrieveMs: item.RetrieveMs,
|
||||
GenerateMs: item.GenerateMs,
|
||||
PromptTokens: item.PromptTokens,
|
||||
CompletionTokens: item.CompletionTokens,
|
||||
ModelName: item.ModelName,
|
||||
TraceData: item.TraceData,
|
||||
CreatedAt: item.CreatedAt,
|
||||
}
|
||||
}
|
||||
|
||||
func BuildKnowledgeRetrieveHitResponse(item *models.KnowledgeRetrieveHit) response.KnowledgeRetrieveHitResponse {
|
||||
return response.KnowledgeRetrieveHitResponse{
|
||||
ID: item.ID,
|
||||
RetrieveLogID: item.RetrieveLogID,
|
||||
KnowledgeBaseID: item.KnowledgeBaseID,
|
||||
ChunkID: item.ChunkID,
|
||||
DocumentID: item.DocumentID,
|
||||
DocumentTitle: item.DocumentTitle,
|
||||
FaqID: item.FaqID,
|
||||
FaqQuestion: item.FaqQuestion,
|
||||
ChunkNo: item.ChunkNo,
|
||||
Title: item.Title,
|
||||
SectionPath: item.SectionPath,
|
||||
ChunkType: item.ChunkType,
|
||||
ChunkTypeName: enums.GetKnowledgeChunkTypeLabel(enums.KnowledgeChunkType(item.ChunkType)),
|
||||
Provider: item.Provider,
|
||||
RankNo: item.RankNo,
|
||||
Score: item.Score,
|
||||
RerankScore: item.RerankScore,
|
||||
UsedInAnswer: item.UsedInAnswer,
|
||||
IsCitation: item.IsCitation,
|
||||
Snippet: item.Snippet,
|
||||
CreatedAt: item.CreatedAt,
|
||||
}
|
||||
}
|
||||
|
||||
func parseSimilarQuestions(raw string) []string {
|
||||
if raw == "" {
|
||||
return []string{}
|
||||
}
|
||||
var items []string
|
||||
if err := json.Unmarshal([]byte(raw), &items); err != nil {
|
||||
return []string{}
|
||||
}
|
||||
return items
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
package builders
|
||||
|
||||
import (
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
)
|
||||
|
||||
func BuildSkillDefinitionResponse(item *models.SkillDefinition) response.SkillDefinitionResponse {
|
||||
return response.SkillDefinitionResponse{
|
||||
ID: item.ID,
|
||||
Code: item.Code,
|
||||
Name: item.Name,
|
||||
Description: item.Description,
|
||||
Prompt: item.Prompt,
|
||||
ExecutionMode: string(item.ExecutionMode),
|
||||
ExecutionModeName: enums.GetSkillExecutionModeLabel(item.ExecutionMode),
|
||||
ExecutionConfig: item.ExecutionConfig,
|
||||
Priority: item.Priority,
|
||||
Status: int(item.Status),
|
||||
StatusName: getSkillStatusName(item.Status),
|
||||
Remark: item.Remark,
|
||||
CreatedAt: item.CreatedAt,
|
||||
UpdatedAt: item.UpdatedAt,
|
||||
CreateUserName: item.CreateUserName,
|
||||
UpdateUserName: item.UpdateUserName,
|
||||
}
|
||||
}
|
||||
|
||||
func getSkillStatusName(status enums.Status) string {
|
||||
if label := enums.GetStatusLabel(status); label != "" {
|
||||
return label
|
||||
}
|
||||
return "未知"
|
||||
}
|
||||
@@ -0,0 +1,75 @@
|
||||
package builders
|
||||
|
||||
import (
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"time"
|
||||
)
|
||||
|
||||
func BuildTagResponse(item *models.Tag) response.TagResponse {
|
||||
if item == nil {
|
||||
return response.TagResponse{}
|
||||
}
|
||||
return response.TagResponse{
|
||||
ID: item.ID,
|
||||
ParentID: item.ParentID,
|
||||
Name: item.Name,
|
||||
Remark: item.Remark,
|
||||
SortNo: item.SortNo,
|
||||
Status: item.Status,
|
||||
CreatedAt: item.CreatedAt.Format(time.DateTime),
|
||||
UpdatedAt: item.UpdatedAt.Format(time.DateTime),
|
||||
}
|
||||
}
|
||||
|
||||
func BuildTagResponses(list []models.Tag) []response.TagResponse {
|
||||
if len(list) == 0 {
|
||||
return nil
|
||||
}
|
||||
results := make([]response.TagResponse, 0, len(list))
|
||||
for i := range list {
|
||||
results = append(results, BuildTagResponse(&list[i]))
|
||||
}
|
||||
return results
|
||||
}
|
||||
|
||||
func BuildTagTreeResponses(list []models.Tag) []*response.TagTreeResponse {
|
||||
if len(list) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
nodeMap := make(map[int64]*response.TagTreeResponse, len(list))
|
||||
roots := make([]*response.TagTreeResponse, 0)
|
||||
|
||||
for i := range list {
|
||||
item := &list[i]
|
||||
nodeMap[item.ID] = &response.TagTreeResponse{
|
||||
ID: item.ID,
|
||||
ParentID: item.ParentID,
|
||||
Name: item.Name,
|
||||
Remark: item.Remark,
|
||||
SortNo: item.SortNo,
|
||||
Status: item.Status,
|
||||
CreatedAt: item.CreatedAt.Format(time.DateTime),
|
||||
UpdatedAt: item.UpdatedAt.Format(time.DateTime),
|
||||
Children: make([]*response.TagTreeResponse, 0),
|
||||
}
|
||||
}
|
||||
|
||||
for i := range list {
|
||||
item := &list[i]
|
||||
node := nodeMap[item.ID]
|
||||
if item.ParentID == 0 {
|
||||
roots = append(roots, node)
|
||||
continue
|
||||
}
|
||||
parent, ok := nodeMap[item.ParentID]
|
||||
if !ok {
|
||||
roots = append(roots, node)
|
||||
continue
|
||||
}
|
||||
parent.Children = append(parent.Children, node)
|
||||
}
|
||||
|
||||
return roots
|
||||
}
|
||||
@@ -0,0 +1,567 @@
|
||||
package builders
|
||||
|
||||
import (
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"cs-agent/internal/pkg/utils"
|
||||
"cs-agent/internal/services"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
|
||||
"github.com/mlogclub/simple/sqls"
|
||||
)
|
||||
|
||||
type TicketBuildContext struct {
|
||||
TagsByTicketID map[int64][]models.Tag
|
||||
Priorities map[int64]*models.TicketPriorityConfig
|
||||
ResolutionCodes map[string]*models.TicketResolutionCode
|
||||
Users map[int64]*models.User
|
||||
Teams map[int64]*models.AgentTeam
|
||||
Customers map[int64]*models.Customer
|
||||
SLAByTicketID map[int64][]models.TicketSLARecord
|
||||
WatchedTicketIDs map[int64]struct{}
|
||||
}
|
||||
|
||||
type TicketDetailBuildContext struct {
|
||||
Users map[int64]*models.User
|
||||
Teams map[int64]*models.AgentTeam
|
||||
AgentProfiles map[int64]*models.AgentProfile
|
||||
RelatedTickets map[int64]*models.Ticket
|
||||
Customers map[int64]*models.Customer
|
||||
AIAgents map[int64]*models.AIAgent
|
||||
}
|
||||
|
||||
func BuildTicket(item *models.Ticket) *response.TicketResponse {
|
||||
return BuildTicketWithContext(item, nil)
|
||||
}
|
||||
|
||||
func BuildTicketWithContext(item *models.Ticket, ctx *TicketBuildContext) *response.TicketResponse {
|
||||
if item == nil {
|
||||
return nil
|
||||
}
|
||||
ret := &response.TicketResponse{
|
||||
ID: item.ID,
|
||||
TicketNo: item.TicketNo,
|
||||
Title: item.Title,
|
||||
Description: item.Description,
|
||||
Source: item.Source,
|
||||
Channel: item.Channel,
|
||||
CustomerID: item.CustomerID,
|
||||
ConversationID: item.ConversationID,
|
||||
Type: item.Type,
|
||||
Priority: item.Priority,
|
||||
Severity: item.Severity,
|
||||
Status: item.Status,
|
||||
CurrentTeamID: item.CurrentTeamID,
|
||||
CurrentAssigneeID: item.CurrentAssigneeID,
|
||||
PendingReason: item.PendingReason,
|
||||
CloseReason: item.CloseReason,
|
||||
ResolutionCode: item.ResolutionCode,
|
||||
ResolutionSummary: item.ResolutionSummary,
|
||||
FirstResponseAt: utils.FormatTimePtr(item.FirstResponseAt),
|
||||
ResolvedAt: utils.FormatTimePtr(item.ResolvedAt),
|
||||
ClosedAt: utils.FormatTimePtr(item.ClosedAt),
|
||||
DueAt: utils.FormatTimePtr(item.DueAt),
|
||||
NextReplyDeadlineAt: utils.FormatTimePtr(item.NextReplyDeadlineAt),
|
||||
ResolveDeadlineAt: utils.FormatTimePtr(item.ResolveDeadlineAt),
|
||||
ReopenedCount: item.ReopenedCount,
|
||||
CreatedAt: utils.FormatTime(item.CreatedAt),
|
||||
UpdatedAt: utils.FormatTime(item.UpdatedAt),
|
||||
}
|
||||
if ctx != nil {
|
||||
if _, ok := ctx.WatchedTicketIDs[item.ID]; ok {
|
||||
ret.WatchedByMe = true
|
||||
}
|
||||
}
|
||||
if ctx != nil && ctx.TagsByTicketID != nil {
|
||||
ret.Tags = BuildTagResponses(ctx.TagsByTicketID[item.ID])
|
||||
} else {
|
||||
ret.Tags = BuildTagResponses(services.TicketService.GetTags(item.ID))
|
||||
}
|
||||
if item.Priority > 0 && ctx != nil && ctx.Priorities != nil {
|
||||
if priority := ctx.Priorities[item.Priority]; priority != nil {
|
||||
ret.PriorityName = priority.Name
|
||||
}
|
||||
} else if item.Priority > 0 {
|
||||
if priority := services.TicketPriorityConfigService.Get(item.Priority); priority != nil {
|
||||
ret.PriorityName = priority.Name
|
||||
}
|
||||
}
|
||||
if item.ResolutionCode != "" && ctx != nil && ctx.ResolutionCodes != nil {
|
||||
if code := ctx.ResolutionCodes[item.ResolutionCode]; code != nil {
|
||||
ret.ResolutionCodeName = code.Name
|
||||
}
|
||||
} else if item.ResolutionCode != "" {
|
||||
if code := services.TicketResolutionCodeService.Take("code = ? AND status <> ?", item.ResolutionCode, enums.StatusDeleted); code != nil {
|
||||
ret.ResolutionCodeName = code.Name
|
||||
}
|
||||
}
|
||||
if item.CurrentAssigneeID > 0 && ctx != nil && ctx.Users != nil {
|
||||
if user := ctx.Users[item.CurrentAssigneeID]; user != nil {
|
||||
ret.CurrentAssigneeName = user.Nickname
|
||||
if ret.CurrentAssigneeName == "" {
|
||||
ret.CurrentAssigneeName = user.Username
|
||||
}
|
||||
}
|
||||
} else if item.CurrentAssigneeID > 0 {
|
||||
if user := services.UserService.Get(item.CurrentAssigneeID); user != nil {
|
||||
ret.CurrentAssigneeName = user.Nickname
|
||||
if ret.CurrentAssigneeName == "" {
|
||||
ret.CurrentAssigneeName = user.Username
|
||||
}
|
||||
}
|
||||
}
|
||||
if item.CurrentTeamID > 0 && ctx != nil && ctx.Teams != nil {
|
||||
if team := ctx.Teams[item.CurrentTeamID]; team != nil {
|
||||
ret.CurrentTeamName = team.Name
|
||||
}
|
||||
} else if item.CurrentTeamID > 0 {
|
||||
if team := services.AgentTeamService.Get(item.CurrentTeamID); team != nil {
|
||||
ret.CurrentTeamName = team.Name
|
||||
}
|
||||
}
|
||||
if item.CustomerID > 0 && ctx != nil && ctx.Customers != nil {
|
||||
ret.Customer = BuildCustomer(ctx.Customers[item.CustomerID])
|
||||
} else if item.CustomerID > 0 {
|
||||
ret.Customer = BuildCustomer(services.CustomerService.Get(item.CustomerID))
|
||||
}
|
||||
if ctx != nil && ctx.SLAByTicketID != nil {
|
||||
ret.SLA = BuildTicketSLAList(ctx.SLAByTicketID[item.ID])
|
||||
} else {
|
||||
ret.SLA = BuildTicketSLAList(
|
||||
services.TicketSLARecordService.Find(
|
||||
sqls.NewCnd().Eq("ticket_id", item.ID).Asc("id"),
|
||||
),
|
||||
)
|
||||
}
|
||||
return ret
|
||||
}
|
||||
|
||||
func BuildTicketList(list []models.Ticket) []response.TicketResponse {
|
||||
return BuildTicketListWithContext(list, nil)
|
||||
}
|
||||
|
||||
func BuildTicketListWithContext(list []models.Ticket, ctx *TicketBuildContext) []response.TicketResponse {
|
||||
if len(list) == 0 {
|
||||
return nil
|
||||
}
|
||||
results := make([]response.TicketResponse, 0, len(list))
|
||||
for i := range list {
|
||||
if item := BuildTicketWithContext(&list[i], ctx); item != nil {
|
||||
results = append(results, *item)
|
||||
}
|
||||
}
|
||||
return results
|
||||
}
|
||||
|
||||
func BuildTicketComment(item *models.TicketComment) *response.TicketCommentResponse {
|
||||
return BuildTicketCommentWithContext(item, nil)
|
||||
}
|
||||
|
||||
func BuildTicketCommentWithContext(item *models.TicketComment, ctx *TicketDetailBuildContext) *response.TicketCommentResponse {
|
||||
if item == nil {
|
||||
return nil
|
||||
}
|
||||
ret := &response.TicketCommentResponse{
|
||||
ID: item.ID,
|
||||
TicketID: item.TicketID,
|
||||
CommentType: item.CommentType,
|
||||
AuthorType: item.AuthorType,
|
||||
AuthorID: item.AuthorID,
|
||||
ContentType: item.ContentType,
|
||||
Content: item.Content,
|
||||
Payload: item.Payload,
|
||||
CreatedAt: utils.FormatTime(item.CreatedAt),
|
||||
}
|
||||
ret.AuthorName = buildTicketOperatorName(item.AuthorType, item.AuthorID, ctx)
|
||||
return ret
|
||||
}
|
||||
|
||||
func BuildTicketCommentList(list []models.TicketComment) []response.TicketCommentResponse {
|
||||
return BuildTicketCommentListWithContext(list, nil)
|
||||
}
|
||||
|
||||
func BuildTicketCommentListWithContext(list []models.TicketComment, ctx *TicketDetailBuildContext) []response.TicketCommentResponse {
|
||||
if len(list) == 0 {
|
||||
return nil
|
||||
}
|
||||
results := make([]response.TicketCommentResponse, 0, len(list))
|
||||
for i := range list {
|
||||
if item := BuildTicketCommentWithContext(&list[i], ctx); item != nil {
|
||||
results = append(results, *item)
|
||||
}
|
||||
}
|
||||
return results
|
||||
}
|
||||
|
||||
func BuildTicketEventLog(item *models.TicketEventLog) *response.TicketEventLogResponse {
|
||||
return BuildTicketEventLogWithContext(item, nil)
|
||||
}
|
||||
|
||||
func BuildTicketEventLogWithContext(item *models.TicketEventLog, ctx *TicketDetailBuildContext) *response.TicketEventLogResponse {
|
||||
if item == nil {
|
||||
return nil
|
||||
}
|
||||
return &response.TicketEventLogResponse{
|
||||
ID: item.ID,
|
||||
TicketID: item.TicketID,
|
||||
EventType: item.EventType,
|
||||
OperatorType: item.OperatorType,
|
||||
OperatorID: item.OperatorID,
|
||||
OperatorName: buildTicketOperatorName(item.OperatorType, item.OperatorID, ctx),
|
||||
OldValue: item.OldValue,
|
||||
NewValue: item.NewValue,
|
||||
Content: item.Content,
|
||||
Payload: item.Payload,
|
||||
CreatedAt: utils.FormatTime(item.CreatedAt),
|
||||
}
|
||||
}
|
||||
|
||||
func BuildTicketEventLogList(list []models.TicketEventLog) []response.TicketEventLogResponse {
|
||||
return BuildTicketEventLogListWithContext(list, nil)
|
||||
}
|
||||
|
||||
func BuildTicketEventLogListWithContext(list []models.TicketEventLog, ctx *TicketDetailBuildContext) []response.TicketEventLogResponse {
|
||||
if len(list) == 0 {
|
||||
return nil
|
||||
}
|
||||
results := make([]response.TicketEventLogResponse, 0, len(list))
|
||||
for i := range list {
|
||||
if item := BuildTicketEventLogWithContext(&list[i], ctx); item != nil {
|
||||
results = append(results, *item)
|
||||
}
|
||||
}
|
||||
return results
|
||||
}
|
||||
|
||||
func BuildTicketSLAList(list []models.TicketSLARecord) []response.TicketSLAResponse {
|
||||
if len(list) == 0 {
|
||||
return nil
|
||||
}
|
||||
results := make([]response.TicketSLAResponse, 0, len(list))
|
||||
for i := range list {
|
||||
item := &list[i]
|
||||
results = append(results, response.TicketSLAResponse{
|
||||
SLAType: item.SLAType,
|
||||
TargetMinutes: item.TargetMinutes,
|
||||
Status: item.Status,
|
||||
StartedAt: utils.FormatTimePtr(item.StartedAt),
|
||||
PausedAt: utils.FormatTimePtr(item.PausedAt),
|
||||
StoppedAt: utils.FormatTimePtr(item.StoppedAt),
|
||||
BreachedAt: utils.FormatTimePtr(item.BreachedAt),
|
||||
ElapsedMin: item.ElapsedMin,
|
||||
})
|
||||
}
|
||||
return results
|
||||
}
|
||||
|
||||
func BuildTicketDetail(aggregate *services.TicketDetailAggregate) *response.TicketDetailResponse {
|
||||
return BuildTicketDetailWithContext(aggregate, nil)
|
||||
}
|
||||
|
||||
func BuildTicketDetailWithContext(aggregate *services.TicketDetailAggregate, ctx *TicketDetailBuildContext) *response.TicketDetailResponse {
|
||||
if aggregate == nil || aggregate.Ticket == nil {
|
||||
return nil
|
||||
}
|
||||
if ctx == nil {
|
||||
users := make(map[int64]*models.User, len(aggregate.Users)+len(aggregate.OperatorUsers))
|
||||
for id, item := range aggregate.Users {
|
||||
users[id] = item
|
||||
}
|
||||
for id, item := range aggregate.OperatorUsers {
|
||||
users[id] = item
|
||||
}
|
||||
ctx = &TicketDetailBuildContext{
|
||||
Users: users,
|
||||
Teams: aggregate.Teams,
|
||||
AgentProfiles: aggregate.AgentProfiles,
|
||||
RelatedTickets: aggregate.RelatedMap,
|
||||
Customers: aggregate.OperatorCustomers,
|
||||
AIAgents: aggregate.OperatorAIAgents,
|
||||
}
|
||||
}
|
||||
ret := &response.TicketDetailResponse{
|
||||
Ticket: *BuildTicket(aggregate.Ticket),
|
||||
Watchers: BuildTicketWatcherListWithContext(aggregate.Watchers, ctx),
|
||||
Collaborators: BuildTicketCollaboratorListWithContext(aggregate.Collaborators, ctx),
|
||||
RelatedTickets: BuildTicketRelationListWithContext(aggregate.RelatedTickets, ctx),
|
||||
}
|
||||
if len(aggregate.Comments) > 0 {
|
||||
ret.Comments = make([]response.TicketCommentResponse, 0, len(aggregate.Comments))
|
||||
for i := range aggregate.Comments {
|
||||
if item := BuildTicketCommentWithContext(&aggregate.Comments[i], ctx); item != nil {
|
||||
ret.Comments = append(ret.Comments, *item)
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(aggregate.Events) > 0 {
|
||||
ret.Events = make([]response.TicketEventLogResponse, 0, len(aggregate.Events))
|
||||
for i := range aggregate.Events {
|
||||
if item := BuildTicketEventLogWithContext(&aggregate.Events[i], ctx); item != nil {
|
||||
ret.Events = append(ret.Events, *item)
|
||||
}
|
||||
}
|
||||
}
|
||||
if aggregate.Customer != nil {
|
||||
ret.Ticket.Customer = BuildCustomer(aggregate.Customer)
|
||||
}
|
||||
ret.Ticket.Tags = BuildTagResponses(aggregate.Tags)
|
||||
ret.Ticket.SLA = BuildTicketSLAList(aggregate.SLAs)
|
||||
for i := range ret.Comments {
|
||||
if ret.Comments[i].AuthorName == "" {
|
||||
ret.Comments[i].AuthorName = buildTicketOperatorName(ret.Comments[i].AuthorType, ret.Comments[i].AuthorID, ctx)
|
||||
}
|
||||
}
|
||||
for i := range ret.Events {
|
||||
if ret.Events[i].OperatorName == "" {
|
||||
ret.Events[i].OperatorName = buildTicketOperatorName(ret.Events[i].OperatorType, ret.Events[i].OperatorID, ctx)
|
||||
}
|
||||
}
|
||||
return ret
|
||||
}
|
||||
|
||||
func BuildTicketRelation(item *models.TicketRelation) *response.TicketRelationResponse {
|
||||
return BuildTicketRelationWithContext(item, nil)
|
||||
}
|
||||
|
||||
func BuildTicketRelationWithContext(item *models.TicketRelation, ctx *TicketDetailBuildContext) *response.TicketRelationResponse {
|
||||
if item == nil {
|
||||
return nil
|
||||
}
|
||||
ret := &response.TicketRelationResponse{
|
||||
ID: item.ID,
|
||||
TicketID: item.TicketID,
|
||||
RelatedTicketID: item.RelatedTicketID,
|
||||
RelationType: item.RelationType,
|
||||
}
|
||||
related := (*models.Ticket)(nil)
|
||||
if ctx != nil && ctx.RelatedTickets != nil {
|
||||
related = ctx.RelatedTickets[item.RelatedTicketID]
|
||||
}
|
||||
if related == nil {
|
||||
related = services.TicketService.Get(item.RelatedTicketID)
|
||||
}
|
||||
if related != nil {
|
||||
ret.RelatedTicketNo = related.TicketNo
|
||||
ret.RelatedTicketTitle = related.Title
|
||||
ret.RelatedTicketStatus = related.Status
|
||||
ret.UpdatedAt = utils.FormatTime(related.UpdatedAt)
|
||||
if related.CurrentTeamID > 0 && ctx != nil && ctx.Teams != nil {
|
||||
if team := ctx.Teams[related.CurrentTeamID]; team != nil {
|
||||
ret.CurrentTeamName = team.Name
|
||||
}
|
||||
} else if related.CurrentTeamID > 0 {
|
||||
if team := services.AgentTeamService.Get(related.CurrentTeamID); team != nil {
|
||||
ret.CurrentTeamName = team.Name
|
||||
}
|
||||
}
|
||||
if related.CurrentAssigneeID > 0 && ctx != nil && ctx.Users != nil {
|
||||
ret.CurrentAssigneeName = buildTicketUserDisplayName(ctx.Users[related.CurrentAssigneeID])
|
||||
} else if related.CurrentAssigneeID > 0 {
|
||||
if user := services.UserService.Get(related.CurrentAssigneeID); user != nil {
|
||||
ret.CurrentAssigneeName = user.Nickname
|
||||
if ret.CurrentAssigneeName == "" {
|
||||
ret.CurrentAssigneeName = user.Username
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return ret
|
||||
}
|
||||
|
||||
func BuildTicketRelationList(list []models.TicketRelation) []response.TicketRelationResponse {
|
||||
return BuildTicketRelationListWithContext(list, nil)
|
||||
}
|
||||
|
||||
func BuildTicketRelationListWithContext(list []models.TicketRelation, ctx *TicketDetailBuildContext) []response.TicketRelationResponse {
|
||||
if len(list) == 0 {
|
||||
return nil
|
||||
}
|
||||
results := make([]response.TicketRelationResponse, 0, len(list))
|
||||
for i := range list {
|
||||
if item := BuildTicketRelationWithContext(&list[i], ctx); item != nil {
|
||||
results = append(results, *item)
|
||||
}
|
||||
}
|
||||
return results
|
||||
}
|
||||
|
||||
func BuildTicketSummary(summary *services.TicketSummaryAggregate) *response.TicketSummaryResponse {
|
||||
if summary == nil {
|
||||
return nil
|
||||
}
|
||||
return &response.TicketSummaryResponse{
|
||||
All: summary.All,
|
||||
Mine: summary.Mine,
|
||||
Watching: summary.Watching,
|
||||
Collaboration: summary.Collaboration,
|
||||
Participating: summary.Participating,
|
||||
Mentioned: summary.Mentioned,
|
||||
Unassigned: summary.Unassigned,
|
||||
PendingCustomer: summary.PendingCustomer,
|
||||
PendingInternal: summary.PendingInternal,
|
||||
Overdue: summary.Overdue,
|
||||
}
|
||||
}
|
||||
|
||||
func BuildTicketRiskOverview(overview *services.TicketRiskOverviewAggregate) *response.TicketRiskOverviewResponse {
|
||||
if overview == nil {
|
||||
return nil
|
||||
}
|
||||
ret := &response.TicketRiskOverviewResponse{
|
||||
Overdue: overview.Overdue,
|
||||
HighRisk: overview.HighRisk,
|
||||
Unassigned: overview.Unassigned,
|
||||
PendingInternal: overview.PendingInternal,
|
||||
PendingCustomer: overview.PendingCustomer,
|
||||
RiskWindowMins: overview.RiskWindowMins,
|
||||
}
|
||||
if len(overview.Reasons) > 0 {
|
||||
ret.Reasons = make([]response.TicketRiskReasonResponse, 0, len(overview.Reasons))
|
||||
for _, item := range overview.Reasons {
|
||||
ret.Reasons = append(ret.Reasons, response.TicketRiskReasonResponse{
|
||||
Code: item.Code,
|
||||
Title: item.Title,
|
||||
Description: item.Description,
|
||||
Count: item.Count,
|
||||
})
|
||||
}
|
||||
}
|
||||
return ret
|
||||
}
|
||||
|
||||
func BuildTicketWatcherList(list []models.TicketWatcher) []response.TicketWatcherResponse {
|
||||
return BuildTicketWatcherListWithContext(list, nil)
|
||||
}
|
||||
|
||||
func BuildTicketWatcherListWithContext(list []models.TicketWatcher, ctx *TicketDetailBuildContext) []response.TicketWatcherResponse {
|
||||
if len(list) == 0 {
|
||||
return nil
|
||||
}
|
||||
results := make([]response.TicketWatcherResponse, 0, len(list))
|
||||
for i := range list {
|
||||
item := &list[i]
|
||||
out := response.TicketWatcherResponse{
|
||||
ID: item.ID,
|
||||
UserID: item.UserID,
|
||||
}
|
||||
if ctx != nil && ctx.Users != nil {
|
||||
out.UserName = buildTicketUserDisplayName(ctx.Users[item.UserID])
|
||||
} else if user := services.UserService.Get(item.UserID); user != nil {
|
||||
out.UserName = buildTicketUserDisplayName(user)
|
||||
}
|
||||
results = append(results, out)
|
||||
}
|
||||
return results
|
||||
}
|
||||
|
||||
func BuildTicketCollaboratorList(list []models.TicketCollaborator) []response.TicketCollaboratorResponse {
|
||||
return BuildTicketCollaboratorListWithContext(list, nil)
|
||||
}
|
||||
|
||||
func BuildTicketCollaboratorListWithContext(list []models.TicketCollaborator, ctx *TicketDetailBuildContext) []response.TicketCollaboratorResponse {
|
||||
if len(list) == 0 {
|
||||
return nil
|
||||
}
|
||||
results := make([]response.TicketCollaboratorResponse, 0, len(list))
|
||||
for i := range list {
|
||||
item := &list[i]
|
||||
out := response.TicketCollaboratorResponse{
|
||||
ID: item.ID,
|
||||
UserID: item.UserID,
|
||||
}
|
||||
if ctx != nil && ctx.Users != nil {
|
||||
out.UserName = buildTicketUserDisplayName(ctx.Users[item.UserID])
|
||||
} else if user := services.UserService.Get(item.UserID); user != nil {
|
||||
out.UserName = buildTicketUserDisplayName(user)
|
||||
}
|
||||
if ctx != nil && ctx.AgentProfiles != nil {
|
||||
if profile := ctx.AgentProfiles[item.UserID]; profile != nil && profile.TeamID > 0 {
|
||||
if team := ctx.Teams[profile.TeamID]; team != nil {
|
||||
out.TeamName = team.Name
|
||||
}
|
||||
}
|
||||
} else if profile := services.AgentProfileService.GetByUserID(item.UserID); profile != nil && profile.TeamID > 0 {
|
||||
if team := services.AgentTeamService.Get(profile.TeamID); team != nil {
|
||||
out.TeamName = team.Name
|
||||
}
|
||||
}
|
||||
results = append(results, out)
|
||||
}
|
||||
return results
|
||||
}
|
||||
|
||||
func BuildTicketView(item *models.TicketView) *response.TicketViewResponse {
|
||||
if item == nil {
|
||||
return nil
|
||||
}
|
||||
ret := &response.TicketViewResponse{
|
||||
ID: item.ID,
|
||||
Name: item.Name,
|
||||
SortNo: item.SortNo,
|
||||
}
|
||||
if strings.TrimSpace(item.FiltersJSON) != "" {
|
||||
_ = json.Unmarshal([]byte(item.FiltersJSON), &ret.Filters)
|
||||
}
|
||||
return ret
|
||||
}
|
||||
|
||||
func BuildTicketViewList(list []models.TicketView) []response.TicketViewResponse {
|
||||
if len(list) == 0 {
|
||||
return nil
|
||||
}
|
||||
results := make([]response.TicketViewResponse, 0, len(list))
|
||||
for i := range list {
|
||||
if item := BuildTicketView(&list[i]); item != nil {
|
||||
results = append(results, *item)
|
||||
}
|
||||
}
|
||||
return results
|
||||
}
|
||||
|
||||
func buildTicketUserDisplayName(user *models.User) string {
|
||||
if user == nil {
|
||||
return ""
|
||||
}
|
||||
if user.Nickname != "" {
|
||||
return user.Nickname
|
||||
}
|
||||
return user.Username
|
||||
}
|
||||
|
||||
func buildTicketOperatorName(senderType enums.IMSenderType, senderID int64, ctx *TicketDetailBuildContext) string {
|
||||
if senderID <= 0 {
|
||||
return ""
|
||||
}
|
||||
switch senderType {
|
||||
case enums.IMSenderTypeAgent:
|
||||
if ctx != nil && ctx.Users != nil {
|
||||
if user := ctx.Users[senderID]; user != nil {
|
||||
return buildTicketUserDisplayName(user)
|
||||
}
|
||||
}
|
||||
if user := services.UserService.Get(senderID); user != nil {
|
||||
return buildTicketUserDisplayName(user)
|
||||
}
|
||||
case enums.IMSenderTypeCustomer:
|
||||
if ctx != nil && ctx.Customers != nil {
|
||||
if customer := ctx.Customers[senderID]; customer != nil {
|
||||
return customer.Name
|
||||
}
|
||||
}
|
||||
if customer := services.CustomerService.Get(senderID); customer != nil {
|
||||
return customer.Name
|
||||
}
|
||||
case enums.IMSenderTypeAI:
|
||||
if ctx != nil && ctx.AIAgents != nil {
|
||||
if ai := ctx.AIAgents[senderID]; ai != nil {
|
||||
return ai.Name
|
||||
}
|
||||
}
|
||||
if ai := services.AIAgentService.Get(senderID); ai != nil {
|
||||
return ai.Name
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
@@ -0,0 +1,61 @@
|
||||
package builders
|
||||
|
||||
import (
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
)
|
||||
|
||||
func BuildTicketResolutionCode(item *models.TicketResolutionCode) *response.TicketResolutionCodeResponse {
|
||||
if item == nil {
|
||||
return nil
|
||||
}
|
||||
return &response.TicketResolutionCodeResponse{
|
||||
ID: item.ID,
|
||||
Name: item.Name,
|
||||
Code: item.Code,
|
||||
SortNo: item.SortNo,
|
||||
Status: item.Status,
|
||||
Remark: item.Remark,
|
||||
}
|
||||
}
|
||||
|
||||
func BuildTicketResolutionCodeList(list []models.TicketResolutionCode) []response.TicketResolutionCodeResponse {
|
||||
if len(list) == 0 {
|
||||
return make([]response.TicketResolutionCodeResponse, 0)
|
||||
}
|
||||
ret := make([]response.TicketResolutionCodeResponse, 0, len(list))
|
||||
for i := range list {
|
||||
if item := BuildTicketResolutionCode(&list[i]); item != nil {
|
||||
ret = append(ret, *item)
|
||||
}
|
||||
}
|
||||
return ret
|
||||
}
|
||||
|
||||
func BuildTicketPriorityConfig(item *models.TicketPriorityConfig) *response.TicketPriorityConfigResponse {
|
||||
if item == nil {
|
||||
return nil
|
||||
}
|
||||
return &response.TicketPriorityConfigResponse{
|
||||
ID: item.ID,
|
||||
Name: item.Name,
|
||||
SortNo: item.SortNo,
|
||||
FirstResponseMinutes: item.FirstResponseMinutes,
|
||||
ResolutionMinutes: item.ResolutionMinutes,
|
||||
Status: item.Status,
|
||||
Remark: item.Remark,
|
||||
}
|
||||
}
|
||||
|
||||
func BuildTicketPriorityConfigList(list []models.TicketPriorityConfig) []response.TicketPriorityConfigResponse {
|
||||
if len(list) == 0 {
|
||||
return make([]response.TicketPriorityConfigResponse, 0)
|
||||
}
|
||||
ret := make([]response.TicketPriorityConfigResponse, 0, len(list))
|
||||
for i := range list {
|
||||
if item := BuildTicketPriorityConfig(&list[i]); item != nil {
|
||||
ret = append(ret, *item)
|
||||
}
|
||||
}
|
||||
return ret
|
||||
}
|
||||
@@ -0,0 +1,68 @@
|
||||
package builders
|
||||
|
||||
import (
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/pkg/utils"
|
||||
"cs-agent/internal/services"
|
||||
)
|
||||
|
||||
type UserBuildOptions struct {
|
||||
Roles bool
|
||||
Permissions bool
|
||||
}
|
||||
|
||||
func BuildUserList(items []models.User, options UserBuildOptions) []response.UserResponse {
|
||||
results := make([]response.UserResponse, 0, len(items))
|
||||
for _, item := range items {
|
||||
results = append(results, *BuildUserResponse(&item, options))
|
||||
}
|
||||
return results
|
||||
}
|
||||
|
||||
func BuildUserResponse(item *models.User, options UserBuildOptions) *response.UserResponse {
|
||||
if item == nil {
|
||||
return nil
|
||||
}
|
||||
ret := &response.UserResponse{
|
||||
ID: item.ID,
|
||||
Username: item.Username,
|
||||
Nickname: item.Nickname,
|
||||
Avatar: item.Avatar,
|
||||
Status: item.Status,
|
||||
LastLoginAt: utils.FormatTimePtr(item.LastLoginAt),
|
||||
LastLoginIP: item.LastLoginIP,
|
||||
}
|
||||
|
||||
if item.Mobile != nil {
|
||||
ret.Mobile = *item.Mobile
|
||||
}
|
||||
if item.Email != nil {
|
||||
ret.Email = *item.Email
|
||||
}
|
||||
|
||||
if options.Roles {
|
||||
ret.Roles = buildAssignedRoles(item.ID)
|
||||
}
|
||||
if options.Permissions {
|
||||
permissionCodes, _ := services.AuthService.GetUserPermissions(item.ID)
|
||||
ret.Permissions = permissionCodes
|
||||
}
|
||||
return ret
|
||||
}
|
||||
|
||||
func buildAssignedRoles(userID int64) []response.RoleResponse {
|
||||
roles, _ := services.AuthService.GetUserRoles(userID)
|
||||
results := make([]response.RoleResponse, 0, len(roles))
|
||||
for _, role := range roles {
|
||||
results = append(results, response.RoleResponse{
|
||||
ID: role.ID,
|
||||
Name: role.Name,
|
||||
Code: role.Code,
|
||||
Status: role.Status,
|
||||
IsSystem: role.IsSystem,
|
||||
SortNo: role.SortNo,
|
||||
})
|
||||
}
|
||||
return results
|
||||
}
|
||||
@@ -0,0 +1,120 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"cs-agent/internal/pkg/config"
|
||||
"cs-agent/internal/pkg/dto/request"
|
||||
"cs-agent/internal/services"
|
||||
"net/url"
|
||||
"strings"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
)
|
||||
|
||||
type AuthController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
func (c *AuthController) PostLogin() *web.JsonResult {
|
||||
cfg := config.Current()
|
||||
req := request.LoginRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
ret, err := services.AuthService.Login(req, cfg.Auth, c.Ctx.RemoteAddr(), c.Ctx.GetHeader("User-Agent"))
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(ret)
|
||||
}
|
||||
|
||||
func (c *AuthController) PostRefresh_token() *web.JsonResult {
|
||||
cfg := config.Current()
|
||||
req := request.RefreshTokenRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
ret, err := services.AuthService.RefreshToken(req.RefreshToken, cfg.Auth, c.Ctx.RemoteAddr(), c.Ctx.GetHeader("User-Agent"))
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(ret)
|
||||
}
|
||||
|
||||
func (c *AuthController) GetWxwork_login() {
|
||||
loginURL, err := services.WxWorkLoginService.BuildWxWorkLoginURL(c.Ctx.URLParam("next"))
|
||||
if err != nil {
|
||||
c.redirectWxWorkError(err.Error())
|
||||
return
|
||||
}
|
||||
c.Ctx.Redirect(loginURL, iris.StatusFound)
|
||||
}
|
||||
|
||||
func (c *AuthController) GetWxwork_qr_login() {
|
||||
loginURL, err := services.WxWorkLoginService.BuildWxWorkQRCodeLoginURL(c.Ctx.URLParam("next"))
|
||||
if err != nil {
|
||||
c.redirectWxWorkError(err.Error())
|
||||
return
|
||||
}
|
||||
c.Ctx.Redirect(loginURL, iris.StatusFound)
|
||||
}
|
||||
|
||||
func (c *AuthController) GetWxwork_callback() {
|
||||
cfg := config.Current()
|
||||
ticket, next, err := services.WxWorkLoginService.LoginByWxWork(
|
||||
c.Ctx.URLParam("code"),
|
||||
c.Ctx.URLParam("state"),
|
||||
cfg.Auth,
|
||||
c.Ctx.RemoteAddr(),
|
||||
c.Ctx.GetHeader("User-Agent"),
|
||||
)
|
||||
if err != nil {
|
||||
c.redirectWxWorkError(err.Error())
|
||||
return
|
||||
}
|
||||
c.Ctx.Redirect("/login/wxwork/callback?ticket="+url.QueryEscape(ticket)+"&next="+url.QueryEscape(next), iris.StatusFound)
|
||||
}
|
||||
|
||||
func (c *AuthController) PostWxwork_exchange() *web.JsonResult {
|
||||
req := request.WxWorkExchangeRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
ret, err := services.WxWorkLoginService.ExchangeWxWorkLoginTicket(req.Ticket)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(ret)
|
||||
}
|
||||
|
||||
func (c *AuthController) PostLogout() *web.JsonResult {
|
||||
req := request.LogoutRequest{}
|
||||
if c.Ctx.GetContentLength() > 0 {
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
}
|
||||
|
||||
if err := services.AuthService.Logout(c.Ctx.GetHeader("Authorization"), req.RefreshToken); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *AuthController) GetProfile() *web.JsonResult {
|
||||
ret, err := services.AuthService.CurrentProfile(c.Ctx)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(ret)
|
||||
}
|
||||
|
||||
func (c *AuthController) redirectWxWorkError(message string) {
|
||||
if idx := strings.Index(message, ": "); idx >= 0 {
|
||||
message = message[idx+2:]
|
||||
}
|
||||
c.Ctx.Redirect("/login?wxworkError="+url.QueryEscape(message), iris.StatusFound)
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
package console
|
||||
|
||||
import (
|
||||
"cs-agent/internal/builders"
|
||||
"cs-agent/internal/pkg/constants"
|
||||
"cs-agent/internal/pkg/dto/request"
|
||||
"cs-agent/internal/services"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
)
|
||||
|
||||
type AgentController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
func (c *AgentController) AnyList() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAgentView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
list, paging := services.AgentProfileService.FindPageByCnd(params.NewPagedSqlCnd(c.Ctx,
|
||||
params.QueryFilter{ParamName: "userId"},
|
||||
params.QueryFilter{ParamName: "teamId"},
|
||||
params.QueryFilter{ParamName: "serviceStatus"},
|
||||
params.QueryFilter{ParamName: "agentCode", Op: params.Like},
|
||||
params.QueryFilter{ParamName: "displayName", Op: params.Like},
|
||||
).Desc("id"))
|
||||
results := builders.BuildAgentProfileList(list)
|
||||
return web.JsonData(&web.PageResult{Results: results, Page: paging})
|
||||
}
|
||||
|
||||
func (c *AgentController) GetList_all() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAgentView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
list := services.AgentProfileService.Find(params.NewPagedSqlCnd(c.Ctx,
|
||||
params.QueryFilter{ParamName: "userId"},
|
||||
params.QueryFilter{ParamName: "teamId"},
|
||||
params.QueryFilter{ParamName: "serviceStatus"},
|
||||
params.QueryFilter{ParamName: "agentCode", Op: params.Like},
|
||||
).Desc("id"))
|
||||
|
||||
return web.JsonData(builders.BuildAgentProfileList(list))
|
||||
}
|
||||
|
||||
func (c *AgentController) GetBy(id int64) *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAgentView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
item := services.AgentProfileService.Get(id)
|
||||
if item == nil {
|
||||
return web.JsonErrorMsg("客服档案不存在")
|
||||
}
|
||||
return web.JsonData(builders.BuildAgentProfileResponse(item))
|
||||
}
|
||||
|
||||
func (c *AgentController) PostCreate() *web.JsonResult {
|
||||
user, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAgentCreate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.CreateAgentProfileRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
item, err := services.AgentProfileService.CreateAgentProfile(req, user)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(builders.BuildAgentProfileResponse(item))
|
||||
}
|
||||
|
||||
func (c *AgentController) PostUpdate() *web.JsonResult {
|
||||
user, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAgentUpdate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.UpdateAgentProfileRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.AgentProfileService.UpdateAgentProfile(req, user); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *AgentController) PostDelete() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAgentDelete); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.DeleteAgentProfileRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.AgentProfileService.DeleteAgentProfile(req.ID); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
@@ -0,0 +1,52 @@
|
||||
package console
|
||||
|
||||
import (
|
||||
"cs-agent/internal/builders"
|
||||
"cs-agent/internal/pkg/constants"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/services"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
)
|
||||
|
||||
type AgentRunLogController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
func (c *AgentRunLogController) AnyList() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionConversationView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
cnd := params.NewPagedSqlCnd(c.Ctx,
|
||||
params.QueryFilter{ParamName: "conversationId"},
|
||||
params.QueryFilter{ParamName: "messageId"},
|
||||
params.QueryFilter{ParamName: "aiAgentId"},
|
||||
params.QueryFilter{ParamName: "plannedAction"},
|
||||
params.QueryFilter{ParamName: "plannedSkillCode", Op: params.Like},
|
||||
params.QueryFilter{ParamName: "finalAction"},
|
||||
params.QueryFilter{ParamName: "userMessage", Op: params.Like},
|
||||
).Desc("id")
|
||||
queryParams := params.NewQueryParams(c.Ctx)
|
||||
queryParams.Cnd = *cnd
|
||||
list, paging := services.AgentRunLogService.FindPageByParams(queryParams)
|
||||
results := make([]response.AgentRunLogResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
results = append(results, builders.BuildAgentRunLog(&item))
|
||||
}
|
||||
return web.JsonData(&web.PageResult{Results: results, Page: paging})
|
||||
}
|
||||
|
||||
func (c *AgentRunLogController) GetBy(id int64) *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionConversationView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
item := services.AgentRunLogService.Get(id)
|
||||
if item == nil {
|
||||
return web.JsonErrorMsg("Agent 运行日志不存在")
|
||||
}
|
||||
return web.JsonData(builders.BuildAgentRunLog(item))
|
||||
}
|
||||
@@ -0,0 +1,124 @@
|
||||
package console
|
||||
|
||||
import (
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/constants"
|
||||
"cs-agent/internal/pkg/dto/request"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"cs-agent/internal/services"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/sqls"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
)
|
||||
|
||||
type AgentTeamController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
func (c *AgentTeamController) AnyList() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAgentTeamView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
cnd := params.NewSqlCnd(c.Ctx,
|
||||
params.QueryFilter{ParamName: "status"},
|
||||
params.QueryFilter{ParamName: "leaderUserId"},
|
||||
params.QueryFilter{ParamName: "name", Op: params.Like},
|
||||
).Desc("id")
|
||||
if _, ok := params.Get(c.Ctx, "status"); !ok {
|
||||
cnd.Where("status <> ?", enums.StatusDeleted)
|
||||
}
|
||||
list := services.AgentTeamService.Find(cnd)
|
||||
results := make([]response.AgentTeamResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
results = append(results, buildAgentTeamResponse(&item))
|
||||
}
|
||||
return web.JsonData(results)
|
||||
}
|
||||
|
||||
func (c *AgentTeamController) GetList_all() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAgentTeamView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
list := services.AgentTeamService.Find(sqls.NewCnd().Eq("status", enums.StatusOk))
|
||||
results := make([]response.AgentTeamResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
results = append(results, buildAgentTeamResponse(&item))
|
||||
}
|
||||
return web.JsonData(results)
|
||||
}
|
||||
|
||||
func (c *AgentTeamController) GetBy(id int64) *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAgentTeamView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
item := services.AgentTeamService.Get(id)
|
||||
if item == nil || item.Status == enums.StatusDeleted {
|
||||
return web.JsonErrorMsg("客服组不存在")
|
||||
}
|
||||
return web.JsonData(buildAgentTeamResponse(item))
|
||||
}
|
||||
|
||||
func (c *AgentTeamController) PostCreate() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAgentTeamCreate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.CreateAgentTeamRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
item, err := services.AgentTeamService.CreateAgentTeam(req, operator)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(buildAgentTeamResponse(item))
|
||||
}
|
||||
|
||||
func (c *AgentTeamController) PostUpdate() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAgentTeamUpdate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.UpdateAgentTeamRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.AgentTeamService.UpdateAgentTeam(req, operator); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *AgentTeamController) PostDelete() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAgentTeamDelete)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.DeleteAgentTeamRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.AgentTeamService.DeleteAgentTeam(req.ID, operator); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func buildAgentTeamResponse(item *models.AgentTeam) response.AgentTeamResponse {
|
||||
ret := response.AgentTeamResponse{
|
||||
ID: item.ID,
|
||||
Name: item.Name,
|
||||
LeaderUserID: item.LeaderUserID,
|
||||
Status: item.Status,
|
||||
Description: item.Description,
|
||||
Remark: item.Remark,
|
||||
}
|
||||
if user := services.UserService.Get(item.LeaderUserID); user != nil {
|
||||
ret.LeaderUsername = user.Username
|
||||
ret.LeaderNickname = user.Nickname
|
||||
}
|
||||
return ret
|
||||
}
|
||||
@@ -0,0 +1,103 @@
|
||||
package console
|
||||
|
||||
import (
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/constants"
|
||||
"cs-agent/internal/pkg/dto/request"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/services"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
)
|
||||
|
||||
type AgentTeamScheduleController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
func (c *AgentTeamScheduleController) AnyList() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAgentTeamScheduleView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
cnd := params.NewPagedSqlCnd(c.Ctx,
|
||||
params.QueryFilter{ParamName: "teamId"},
|
||||
).Desc("start_at").Desc("id")
|
||||
list, paging := services.AgentTeamScheduleService.FindPageByCnd(cnd)
|
||||
results := make([]response.AgentTeamScheduleResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
results = append(results, buildAgentTeamScheduleResponse(&item))
|
||||
}
|
||||
return web.JsonData(&web.PageResult{Results: results, Page: paging})
|
||||
}
|
||||
|
||||
func (c *AgentTeamScheduleController) GetBy(id int64) *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAgentTeamScheduleView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
item := services.AgentTeamScheduleService.Get(id)
|
||||
if item == nil {
|
||||
return web.JsonErrorMsg("客服组排班不存在")
|
||||
}
|
||||
return web.JsonData(buildAgentTeamScheduleResponse(item))
|
||||
}
|
||||
|
||||
func (c *AgentTeamScheduleController) PostCreate() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAgentTeamScheduleCreate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.CreateAgentTeamScheduleRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
item, err := services.AgentTeamScheduleService.CreateAgentTeamSchedule(req, operator)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(buildAgentTeamScheduleResponse(item))
|
||||
}
|
||||
|
||||
func (c *AgentTeamScheduleController) PostUpdate() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAgentTeamScheduleUpdate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.UpdateAgentTeamScheduleRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.AgentTeamScheduleService.UpdateAgentTeamSchedule(req, operator); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *AgentTeamScheduleController) PostDelete() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAgentTeamScheduleDelete); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.DeleteAgentTeamScheduleRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.AgentTeamScheduleService.DeleteAgentTeamSchedule(req.ID); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func buildAgentTeamScheduleResponse(item *models.AgentTeamSchedule) response.AgentTeamScheduleResponse {
|
||||
ret := response.AgentTeamScheduleResponse{
|
||||
ID: item.ID,
|
||||
TeamID: item.TeamID,
|
||||
StartAt: item.StartAt.Format("2006-01-02 15:04:05"),
|
||||
EndAt: item.EndAt.Format("2006-01-02 15:04:05"),
|
||||
SourceType: item.SourceType,
|
||||
Remark: item.Remark,
|
||||
}
|
||||
if team := services.AgentTeamService.Get(item.TeamID); team != nil {
|
||||
ret.TeamName = team.Name
|
||||
}
|
||||
return ret
|
||||
}
|
||||
@@ -0,0 +1,183 @@
|
||||
package console
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"strings"
|
||||
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/dto/request"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"cs-agent/internal/pkg/utils"
|
||||
"cs-agent/internal/services"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/sqls"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
)
|
||||
|
||||
type AIAgentController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
func (c *AIAgentController) AnyList() *web.JsonResult {
|
||||
cnd := params.NewPagedSqlCnd(c.Ctx,
|
||||
params.QueryFilter{ParamName: "status"},
|
||||
params.QueryFilter{ParamName: "name", Op: params.Like},
|
||||
params.QueryFilter{ParamName: "code", Op: params.Like},
|
||||
).Desc("sort_no").Desc("id")
|
||||
list, paging := services.AIAgentService.FindPageByCnd(cnd)
|
||||
results := make([]response.AIAgentResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
results = append(results, buildAIAgentResponse(&item))
|
||||
}
|
||||
return web.JsonData(&web.PageResult{Results: results, Page: paging})
|
||||
}
|
||||
|
||||
func (c *AIAgentController) GetList_all() *web.JsonResult {
|
||||
list := services.AIAgentService.Find(sqls.NewCnd().Where("status = ?", enums.StatusOk).Desc("sort_no").Desc("id"))
|
||||
results := make([]response.AIAgentResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
results = append(results, buildAIAgentResponse(&item))
|
||||
}
|
||||
return web.JsonData(results)
|
||||
}
|
||||
|
||||
func (c *AIAgentController) GetBy(id int64) *web.JsonResult {
|
||||
item := services.AIAgentService.Get(id)
|
||||
if item == nil {
|
||||
return web.JsonErrorMsg("AI Agent 不存在")
|
||||
}
|
||||
return web.JsonData(buildAIAgentResponse(item))
|
||||
}
|
||||
|
||||
func (c *AIAgentController) PostCreate() *web.JsonResult {
|
||||
req := request.CreateAIAgentRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
item, err := services.AIAgentService.CreateAIAgent(req, services.AuthService.GetAuthPrincipal(c.Ctx))
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(buildAIAgentResponse(item))
|
||||
}
|
||||
|
||||
func (c *AIAgentController) PostUpdate() *web.JsonResult {
|
||||
req := request.UpdateAIAgentRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.AIAgentService.UpdateAIAgent(req, services.AuthService.GetAuthPrincipal(c.Ctx)); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *AIAgentController) PostDelete() *web.JsonResult {
|
||||
req := request.DeleteAIAgentRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.AIAgentService.DeleteAIAgent(req.ID, services.AuthService.GetAuthPrincipal(c.Ctx)); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *AIAgentController) PostUpdate_sort() *web.JsonResult {
|
||||
var ids []int64
|
||||
if err := c.Ctx.ReadJSON(&ids); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.AIAgentService.UpdateSort(ids); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *AIAgentController) PostUpdate_status() *web.JsonResult {
|
||||
req := request.UpdateAIAgentStatusRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.AIAgentService.UpdateStatus(req.ID, req.Status, services.AuthService.GetAuthPrincipal(c.Ctx)); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func buildAIAgentResponse(item *models.AIAgent) response.AIAgentResponse {
|
||||
ret := response.AIAgentResponse{
|
||||
ID: item.ID,
|
||||
Name: item.Name,
|
||||
Description: item.Description,
|
||||
Status: item.Status,
|
||||
StatusName: enums.GetStatusLabel(item.Status),
|
||||
AIConfigID: item.AIConfigID,
|
||||
ServiceMode: item.ServiceMode,
|
||||
ServiceModeName: enums.GetIMConversationServiceModeLabel(enums.IMConversationServiceMode(item.ServiceMode)),
|
||||
SystemPrompt: item.SystemPrompt,
|
||||
WelcomeMessage: item.WelcomeMessage,
|
||||
ReplyTimeoutSeconds: item.ReplyTimeoutSeconds,
|
||||
HandoffMode: item.HandoffMode,
|
||||
HandoffModeName: enums.GetAIAgentHandoffModeLabel(enums.AIAgentHandoffMode(item.HandoffMode)),
|
||||
MaxAIReplyRounds: item.MaxAIReplyRounds,
|
||||
FallbackMode: item.FallbackMode,
|
||||
FallbackModeName: enums.GetAIAgentFallbackModeLabel(enums.AIAgentFallbackMode(item.FallbackMode)),
|
||||
FallbackMessage: item.FallbackMessage,
|
||||
KnowledgeIDs: utils.SplitInt64s(item.KnowledgeIDs),
|
||||
SkillIDs: utils.SplitInt64s(item.SkillIDs),
|
||||
KnowledgeBaseNames: make([]string, 0),
|
||||
Skills: make([]response.AIAgentSkillResponse, 0),
|
||||
Teams: make([]response.AIAgentTeamResponse, 0),
|
||||
DirectTools: make([]response.AIAgentMCPToolResponse, 0),
|
||||
SortNo: item.SortNo,
|
||||
Remark: item.Remark,
|
||||
CreatedAt: item.CreatedAt.Format("2006-01-02 15:04:05"),
|
||||
UpdatedAt: item.UpdatedAt.Format("2006-01-02 15:04:05"),
|
||||
CreateUserName: item.CreateUserName,
|
||||
UpdateUserName: item.UpdateUserName,
|
||||
}
|
||||
if aiConfig := services.AIConfigService.Get(item.AIConfigID); aiConfig != nil {
|
||||
ret.AIConfigName = aiConfig.Name
|
||||
}
|
||||
for _, id := range utils.SplitInt64s(item.TeamIDs) {
|
||||
if team := services.AgentTeamService.Get(id); team != nil {
|
||||
ret.Teams = append(ret.Teams, response.AIAgentTeamResponse{
|
||||
ID: team.ID,
|
||||
Name: team.Name,
|
||||
})
|
||||
}
|
||||
}
|
||||
for _, id := range ret.KnowledgeIDs {
|
||||
if knowledgeBase := services.KnowledgeBaseService.Get(id); knowledgeBase != nil {
|
||||
ret.KnowledgeBaseNames = append(ret.KnowledgeBaseNames, knowledgeBase.Name)
|
||||
}
|
||||
}
|
||||
for _, id := range ret.SkillIDs {
|
||||
if skill := services.SkillDefinitionService.Get(id); skill != nil {
|
||||
ret.Skills = append(ret.Skills, response.AIAgentSkillResponse{
|
||||
ID: skill.ID,
|
||||
Code: skill.Code,
|
||||
Name: skill.Name,
|
||||
})
|
||||
}
|
||||
}
|
||||
if raw := strings.TrimSpace(item.AllowedMCPTools); raw != "" {
|
||||
var directTools []request.AIAgentMCPToolRequest
|
||||
if err := json.Unmarshal([]byte(raw), &directTools); err == nil {
|
||||
for _, tool := range directTools {
|
||||
ret.DirectTools = append(ret.DirectTools, response.AIAgentMCPToolResponse{
|
||||
ServerCode: strings.TrimSpace(tool.ServerCode),
|
||||
ToolName: strings.TrimSpace(tool.ToolName),
|
||||
Title: strings.TrimSpace(tool.Title),
|
||||
Description: strings.TrimSpace(tool.Description),
|
||||
Arguments: tool.Arguments,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
return ret
|
||||
}
|
||||
@@ -0,0 +1,143 @@
|
||||
package console
|
||||
|
||||
import (
|
||||
"cs-agent/internal/pkg/constants"
|
||||
"cs-agent/internal/pkg/dto/request"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"cs-agent/internal/services"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
)
|
||||
|
||||
type AIConfigController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
func (c *AIConfigController) AnyList() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAIConfigView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
list, paging := services.AIConfigService.FindPageByCnd(params.NewPagedSqlCnd(c.Ctx,
|
||||
params.QueryFilter{ParamName: "status"},
|
||||
params.QueryFilter{ParamName: "provider"},
|
||||
params.QueryFilter{ParamName: "modelType"},
|
||||
params.QueryFilter{ParamName: "name", Op: params.Like},
|
||||
params.QueryFilter{ParamName: "modelName", Op: params.Like},
|
||||
).Desc("sort_no").Desc("id"))
|
||||
results := make([]response.AIConfigResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
results = append(results, response.BuildAIConfigResponse(&item))
|
||||
}
|
||||
return web.JsonData(&web.PageResult{Results: results, Page: paging})
|
||||
}
|
||||
|
||||
func (c *AIConfigController) AnyList_all() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAIConfigView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
list := services.AIConfigService.Find(params.NewSqlCnd(c.Ctx,
|
||||
params.QueryFilter{ParamName: "modelType"},
|
||||
).Eq("status", enums.StatusOk).Desc("sort_no").Desc("id"))
|
||||
|
||||
results := make([]response.AIConfigResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
results = append(results, response.BuildAIConfigResponse(&item))
|
||||
}
|
||||
return web.JsonData(results)
|
||||
}
|
||||
|
||||
func (c *AIConfigController) GetBy(id int64) *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAIConfigView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
item := services.AIConfigService.Get(id)
|
||||
if item == nil || item.Status == enums.StatusDeleted {
|
||||
return web.JsonErrorMsg("AI配置不存在")
|
||||
}
|
||||
return web.JsonData(response.BuildAIConfigResponse(item))
|
||||
}
|
||||
|
||||
func (c *AIConfigController) PostCreate() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAIConfigCreate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.CreateAIConfigRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
item, err := services.AIConfigService.CreateAIConfig(req, operator)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(response.BuildAIConfigResponse(item))
|
||||
}
|
||||
|
||||
func (c *AIConfigController) PostUpdate() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAIConfigUpdate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.UpdateAIConfigRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.AIConfigService.UpdateAIConfig(req, operator); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *AIConfigController) PostDelete() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAIConfigDelete)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.DeleteAIConfigRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.AIConfigService.DeleteAIConfig(req.ID, operator); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *AIConfigController) PostUpdate_status() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAIConfigUpdate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.UpdateAIConfigStatusRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.AIConfigService.UpdateStatus(req.ID, req.Status, operator); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *AIConfigController) PostUpdate_sort() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAIConfigUpdate); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
var ids []int64
|
||||
if err := c.Ctx.ReadJSON(&ids); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.AIConfigService.UpdateSort(ids); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
@@ -0,0 +1,92 @@
|
||||
package console
|
||||
|
||||
import (
|
||||
"cs-agent/internal/builders"
|
||||
"cs-agent/internal/pkg/constants"
|
||||
"cs-agent/internal/pkg/dto/request"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"cs-agent/internal/services"
|
||||
"strings"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
)
|
||||
|
||||
type AssetController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
func (c *AssetController) AnyList() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAssetView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
cnd := params.NewPagedSqlCnd(c.Ctx,
|
||||
params.QueryFilter{ParamName: "provider"},
|
||||
params.QueryFilter{ParamName: "status"},
|
||||
params.QueryFilter{ParamName: "createUserId"},
|
||||
params.QueryFilter{ParamName: "filename", Op: params.Like},
|
||||
).Desc("id")
|
||||
if strings.TrimSpace(c.Ctx.URLParam("status")) == "" {
|
||||
cnd = cnd.Eq("status", enums.AssetStatusSuccess)
|
||||
}
|
||||
|
||||
list, paging := services.AssetService.FindPageByCnd(cnd)
|
||||
results := make([]response.AssetResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
results = append(results, builders.BuildAsset(&item))
|
||||
}
|
||||
return web.JsonData(&web.PageResult{Results: results, Page: paging})
|
||||
}
|
||||
|
||||
func (c *AssetController) GetBy(id int64) *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAssetView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
item := services.AssetService.Get(id)
|
||||
if item == nil {
|
||||
return web.JsonErrorMsg("文件不存在")
|
||||
}
|
||||
return web.JsonData(builders.BuildAsset(item))
|
||||
}
|
||||
|
||||
func (c *AssetController) PostCreate() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAssetCreate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.CreateAssetRequest{}
|
||||
if err := params.ReadForm(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
f, header, err := c.Ctx.FormFile("file")
|
||||
if err != nil {
|
||||
return web.JsonErrorMsg("请选择上传文件")
|
||||
}
|
||||
_ = f.Close()
|
||||
|
||||
item, err := services.AssetService.UploadFile(header, req.Prefix, operator)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(builders.BuildAsset(item))
|
||||
}
|
||||
|
||||
func (c *AssetController) PostDelete() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionAssetDelete)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.DeleteAssetRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.AssetService.DeleteAsset(req.ID, operator); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
@@ -0,0 +1,118 @@
|
||||
package console
|
||||
|
||||
import (
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/constants"
|
||||
"cs-agent/internal/pkg/dto/request"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"cs-agent/internal/services"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
)
|
||||
|
||||
type ChannelController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
func (c *ChannelController) AnyList() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionChannelView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
list, paging := services.ChannelService.FindPageByCnd(params.NewPagedSqlCnd(c.Ctx,
|
||||
params.QueryFilter{ParamName: "status"},
|
||||
params.QueryFilter{ParamName: "name", Op: params.Like},
|
||||
params.QueryFilter{ParamName: "channelType"},
|
||||
params.QueryFilter{ParamName: "channelId", Op: params.Like},
|
||||
).Where("status <> ?", enums.StatusDeleted).Desc("id"))
|
||||
results := make([]response.ChannelResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
results = append(results, buildChannelResponse(&item))
|
||||
}
|
||||
return web.JsonData(&web.PageResult{Results: results, Page: paging})
|
||||
}
|
||||
|
||||
func (c *ChannelController) GetBy(id int64) *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionChannelView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
item := services.ChannelService.Get(id)
|
||||
if item == nil || item.Status == enums.StatusDeleted {
|
||||
return web.JsonErrorMsg("channel not found")
|
||||
}
|
||||
return web.JsonData(buildChannelResponse(item))
|
||||
}
|
||||
|
||||
func (c *ChannelController) PostCreate() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionChannelCreate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.CreateChannelRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
item, err := services.ChannelService.CreateChannel(req, operator)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(buildChannelResponse(item))
|
||||
}
|
||||
|
||||
func (c *ChannelController) PostUpdate() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionChannelUpdate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.UpdateChannelRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.ChannelService.UpdateChannel(req, operator); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *ChannelController) PostUpdate_status() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionChannelUpdate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.UpdateChannelStatusRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.ChannelService.UpdateStatus(req.ID, req.Status, operator); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *ChannelController) PostDelete() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionChannelDelete)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.DeleteChannelRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.ChannelService.DeleteChannel(req.ID, operator); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func buildChannelResponse(item *models.Channel) response.ChannelResponse {
|
||||
ret := response.BuildChannelResponse(item)
|
||||
if item == nil {
|
||||
return ret
|
||||
}
|
||||
if aiAgent := services.AIAgentService.Get(item.AIAgentID); aiAgent != nil {
|
||||
ret.AIAgentName = aiAgent.Name
|
||||
}
|
||||
return ret
|
||||
}
|
||||
@@ -0,0 +1,113 @@
|
||||
package console
|
||||
|
||||
import (
|
||||
"cs-agent/internal/builders"
|
||||
"cs-agent/internal/pkg/constants"
|
||||
"cs-agent/internal/pkg/dto/request"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"cs-agent/internal/services"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
)
|
||||
|
||||
type CompanyController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
func (c *CompanyController) AnyList() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionCompanyView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
list, paging := services.CompanyService.FindPageByCnd(params.NewPagedSqlCnd(c.Ctx,
|
||||
params.QueryFilter{ParamName: "status"},
|
||||
params.QueryFilter{ParamName: "name", Op: params.Like},
|
||||
params.QueryFilter{ParamName: "code", Op: params.Like},
|
||||
).Where("status <> ?", enums.StatusDeleted).Desc("id"))
|
||||
|
||||
results := builders.BuildCompanyList(list)
|
||||
companyIDs := make([]int64, 0, len(results))
|
||||
for _, item := range results {
|
||||
companyIDs = append(companyIDs, item.ID)
|
||||
}
|
||||
countMap := services.CustomerService.CountByCompanyIDs(companyIDs)
|
||||
for i := range results {
|
||||
results[i].CustomerCount = countMap[results[i].ID]
|
||||
}
|
||||
return web.JsonData(&web.PageResult{Results: results, Page: paging})
|
||||
}
|
||||
|
||||
func (c *CompanyController) GetBy(id int64) *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionCompanyView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
item := services.CompanyService.Get(id)
|
||||
if item == nil || item.Status == enums.StatusDeleted {
|
||||
return web.JsonData(nil)
|
||||
}
|
||||
ret := builders.BuildCompany(item)
|
||||
return web.JsonData(&ret)
|
||||
}
|
||||
|
||||
func (c *CompanyController) PostCreate() *web.JsonResult {
|
||||
user, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionCompanyCreate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.CreateCompanyRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
item, err := services.CompanyService.CreateCompany(req, user)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
ret := builders.BuildCompany(item)
|
||||
return web.JsonData(&ret)
|
||||
}
|
||||
|
||||
func (c *CompanyController) PostUpdate() *web.JsonResult {
|
||||
user, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionCompanyUpdate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.UpdateCompanyRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.CompanyService.UpdateCompany(req, user); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *CompanyController) PostDelete() *web.JsonResult {
|
||||
user, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionCompanyDelete)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.DeleteCompanyRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.CompanyService.DeleteCompany(req.ID, *user); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *CompanyController) PostUpdate_status() *web.JsonResult {
|
||||
user, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionCompanyUpdate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.UpdateCompanyStatusRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.CompanyService.UpdateStatus(req.ID, req.Status, user); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
@@ -0,0 +1,365 @@
|
||||
package console
|
||||
|
||||
import (
|
||||
"cs-agent/internal/builders"
|
||||
"cs-agent/internal/pkg/constants"
|
||||
"cs-agent/internal/pkg/dto/request"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"cs-agent/internal/services"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/common/strs"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
"github.com/spf13/cast"
|
||||
)
|
||||
|
||||
type ConversationController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
func (c *ConversationController) AnyList() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionConversationView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
cnd := params.NewPagedSqlCnd(c.Ctx,
|
||||
params.QueryFilter{ParamName: "status"},
|
||||
params.QueryFilter{ParamName: "externalSource"},
|
||||
params.QueryFilter{ParamName: "serviceMode"},
|
||||
params.QueryFilter{ParamName: "currentAssigneeId"},
|
||||
).Desc("last_message_at").Desc("id")
|
||||
|
||||
paging := params.GetPaging(c.Ctx)
|
||||
|
||||
if keyword, _ := params.Get(c.Ctx, "keyword"); strs.IsNotBlank(keyword) {
|
||||
keywordLike := "%" + strings.TrimSpace(keyword) + "%"
|
||||
cnd.Where("subject LIKE ? OR external_id LIKE ? OR last_message_summary LIKE ?", keywordLike, keywordLike, keywordLike)
|
||||
}
|
||||
|
||||
// 标签搜索
|
||||
if tagID, _ := params.GetInt64(c.Ctx, "tagId"); tagID > 0 {
|
||||
tagIDs := services.TagService.GetSelfAndDescendantIDs(tagID)
|
||||
if len(tagIDs) == 0 {
|
||||
return web.JsonData(&web.PageResult{
|
||||
Results: []response.ConversationResponse{},
|
||||
Page: paging,
|
||||
})
|
||||
}
|
||||
cnd.Where("id IN (SELECT conversation_id FROM conversation_tag_rels WHERE tag_id IN (?))", tagIDs)
|
||||
}
|
||||
if agentTeamID, _ := params.GetInt64(c.Ctx, "agentTeamId"); agentTeamID > 0 {
|
||||
userIDs := services.AgentProfileService.GetUserIDsByTeamID(agentTeamID)
|
||||
if len(userIDs) == 0 {
|
||||
return web.JsonData(&web.PageResult{
|
||||
Results: []response.ConversationResponse{},
|
||||
Page: paging,
|
||||
})
|
||||
}
|
||||
cnd.In("current_assignee_id", userIDs)
|
||||
}
|
||||
|
||||
list, paging := services.ConversationService.FindPageByCnd(cnd)
|
||||
results := make([]response.ConversationResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
results = append(results, builders.BuildConversation(&item))
|
||||
}
|
||||
return web.JsonData(&web.PageResult{Results: results, Page: paging})
|
||||
}
|
||||
|
||||
func (c *ConversationController) AnyConversations() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionConversationView)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
filterValue, _ := params.Get(c.Ctx, "filter")
|
||||
keyword, _ := params.Get(c.Ctx, "keyword")
|
||||
paging := params.GetPaging(c.Ctx)
|
||||
|
||||
list, paging, err := services.ConversationService.ListConversations(
|
||||
operator.UserID,
|
||||
request.AgentConversationFilter(strings.TrimSpace(filterValue)),
|
||||
keyword,
|
||||
paging,
|
||||
)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
results := make([]response.ConversationResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
results = append(results, builders.BuildConversation(&item))
|
||||
}
|
||||
return web.JsonData(&web.PageResult{Results: results, Page: paging})
|
||||
}
|
||||
|
||||
func (c *ConversationController) GetBy(id int64) *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionConversationView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
item := services.ConversationService.Get(id)
|
||||
if item == nil {
|
||||
return web.JsonErrorMsg("会话不存在")
|
||||
}
|
||||
|
||||
detail := response.ConversationDetailResponse{
|
||||
ConversationResponse: builders.BuildConversation(item),
|
||||
Participants: builders.BuildParticipantResponses(id),
|
||||
}
|
||||
return web.JsonData(detail)
|
||||
}
|
||||
|
||||
func (c *ConversationController) AnyMessage_list() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionConversationView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
var (
|
||||
conversationID, _ = params.GetInt64(c.Ctx, "conversationId")
|
||||
senderType, _ = params.Get(c.Ctx, "senderType")
|
||||
messageType, _ = params.Get(c.Ctx, "messageType")
|
||||
cursor, _ = params.GetInt64(c.Ctx, "cursor")
|
||||
limit, _ = params.GetInt(c.Ctx, "limit")
|
||||
)
|
||||
if conversation := services.ConversationService.Get(conversationID); conversation == nil {
|
||||
return web.JsonErrorMsg("会话不存在")
|
||||
}
|
||||
|
||||
list, nextCursor, hasMore := services.MessageService.FindByConversationIDCursor(
|
||||
conversationID, cursor, limit, senderType, messageType,
|
||||
)
|
||||
results := builders.BuildMessages(list)
|
||||
|
||||
return web.JsonCursorData(results, cast.ToString(nextCursor), hasMore)
|
||||
}
|
||||
|
||||
func (c *ConversationController) PostAssign() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionConversationAssign)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.AssignConversationRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.ConversationService.AssignConversation(req, operator); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *ConversationController) PostDispatch() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionConversationAssign)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.DispatchConversationRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.ConversationService.AutoAssignConversation(req.ConversationID, operator); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *ConversationController) PostTransfer() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionConversationTransfer)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.TransferConversationRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.ConversationService.TransferConversation(req.ConversationID, req.ToUserID, req.Reason, operator); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *ConversationController) PostClose() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionConversationClose)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.CloseConversationRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.ConversationService.CloseConversation(req.ConversationID, req.CloseReason, operator); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *ConversationController) PostLink_customer() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionConversationLinkCustomer)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.LinkConversationCustomerRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.ConversationService.LinkConversationCustomer(req.ConversationID, req.CustomerID, operator); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *ConversationController) PostSend_message() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionConversationSend)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.SendConversationMessageRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
item, err := services.MessageService.SendAgentMessage(req.ConversationID, 0, req.ClientMsgID, req.MessageType, req.Content, req.Payload, operator)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(builders.BuildMessage(item))
|
||||
}
|
||||
|
||||
func (c *ConversationController) PostRecall_message() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionConversationSend)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.RecallConversationMessageRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
item, err := services.MessageService.RecallAgentMessage(req.MessageID, operator)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(builders.BuildMessage(item))
|
||||
}
|
||||
|
||||
func (c *ConversationController) PostRead() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionConversationView)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.ReadConversationRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.ConversationService.MarkAgentConversationReadToMessage(req.ConversationID, req.MessageID, operator); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *ConversationController) PostUpload_image() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionConversationSend)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
rawConv := strings.TrimSpace(c.Ctx.FormValue("conversationId"))
|
||||
if rawConv == "" {
|
||||
return web.JsonErrorMsg("conversationId不能为空")
|
||||
}
|
||||
conversationID, err := strconv.ParseInt(rawConv, 10, 64)
|
||||
if err != nil || conversationID <= 0 {
|
||||
return web.JsonErrorMsg("conversationId不能为空")
|
||||
}
|
||||
if _, err := services.MessageService.ValidateConversationSender(conversationID, enums.IMSenderTypeAgent, operator, nil); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
f, header, err := c.Ctx.FormFile("file")
|
||||
if err != nil {
|
||||
return web.JsonErrorMsg("请选择上传图片")
|
||||
}
|
||||
_ = f.Close()
|
||||
|
||||
if !strings.HasPrefix(strings.ToLower(header.Header.Get("Content-Type")), "image/") {
|
||||
return web.JsonErrorMsg("仅支持上传图片文件")
|
||||
}
|
||||
|
||||
item, err := services.AssetService.UploadFile(header, "images", operator)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(builders.BuildAsset(item))
|
||||
}
|
||||
|
||||
func (c *ConversationController) PostUpload_attachment() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionConversationSend)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
rawConv := strings.TrimSpace(c.Ctx.FormValue("conversationId"))
|
||||
if rawConv == "" {
|
||||
return web.JsonErrorMsg("conversationId不能为空")
|
||||
}
|
||||
conversationID, err := strconv.ParseInt(rawConv, 10, 64)
|
||||
if err != nil || conversationID <= 0 {
|
||||
return web.JsonErrorMsg("conversationId不能为空")
|
||||
}
|
||||
if _, err := services.MessageService.ValidateConversationSender(conversationID, enums.IMSenderTypeAgent, operator, nil); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
f, header, err := c.Ctx.FormFile("file")
|
||||
if err != nil {
|
||||
return web.JsonErrorMsg("请选择上传附件")
|
||||
}
|
||||
_ = f.Close()
|
||||
|
||||
item, err := services.AssetService.UploadFile(header, "attachments", operator)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(builders.BuildAsset(item))
|
||||
}
|
||||
|
||||
func (c *ConversationController) PostAdd_tag() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionConversationTag)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.AddConversationTagRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.ConversationTagService.AddTag(req, operator); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *ConversationController) PostRemove_tag() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionConversationTag); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.RemoveConversationTagRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.ConversationTagService.RemoveTag(req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
@@ -0,0 +1,76 @@
|
||||
package console
|
||||
|
||||
import (
|
||||
"cs-agent/internal/builders"
|
||||
"cs-agent/internal/pkg/constants"
|
||||
"cs-agent/internal/pkg/dto/request"
|
||||
"cs-agent/internal/services"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
)
|
||||
|
||||
type CustomerContactController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
// AnyList GET/POST /customer-contact/list?customerId=
|
||||
func (c *CustomerContactController) AnyList() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionCustomerView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
customerID, _ := params.GetInt64(c.Ctx, "customerId")
|
||||
if customerID <= 0 {
|
||||
return web.JsonErrorMsg("customerId 必填")
|
||||
}
|
||||
list := services.CustomerContactService.FindActiveByCustomerID(customerID)
|
||||
return web.JsonData(builders.BuildCustomerContactList(list))
|
||||
}
|
||||
|
||||
func (c *CustomerContactController) PostCreate() *web.JsonResult {
|
||||
user, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionCustomerUpdate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.CreateCustomerContactRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
item, err := services.CustomerContactService.CreateCustomerContact(req, user)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
ret := builders.BuildCustomerContactResponse(item)
|
||||
return web.JsonData(&ret)
|
||||
}
|
||||
|
||||
func (c *CustomerContactController) PostUpdate() *web.JsonResult {
|
||||
user, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionCustomerUpdate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.UpdateCustomerContactRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.CustomerContactService.UpdateCustomerContact(req, user); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *CustomerContactController) PostDelete() *web.JsonResult {
|
||||
user, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionCustomerUpdate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.DeleteCustomerContactRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.CustomerContactService.DeleteCustomerContact(req.ID, user); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
@@ -0,0 +1,129 @@
|
||||
package console
|
||||
|
||||
import (
|
||||
"cs-agent/internal/builders"
|
||||
"cs-agent/internal/pkg/constants"
|
||||
"cs-agent/internal/pkg/dto"
|
||||
"cs-agent/internal/pkg/dto/request"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"cs-agent/internal/services"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
)
|
||||
|
||||
type CustomerController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
func (c *CustomerController) PostList() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionCustomerView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
var req request.CustomerListRequest
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
list, paging := services.CustomerService.ListCustomers(req)
|
||||
return web.JsonData(&web.PageResult{Results: builders.BuildCustomerList(list), Page: paging})
|
||||
}
|
||||
|
||||
func (c *CustomerController) GetBy(id int64) *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionCustomerView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
item := services.CustomerService.Get(id)
|
||||
if item == nil || item.Status == enums.StatusDeleted {
|
||||
return web.JsonData(nil)
|
||||
}
|
||||
ret := builders.BuildCustomer(item)
|
||||
return web.JsonData(&ret)
|
||||
}
|
||||
|
||||
// PostSave_profile POST /save_profile — 客户主信息与联系方式在同一事务中保存。
|
||||
func (c *CustomerController) PostSave_profile() *web.JsonResult {
|
||||
req := request.SaveCustomerProfileRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
createMode := req.ID == nil || *req.ID <= 0
|
||||
var user *dto.AuthPrincipal
|
||||
var err error
|
||||
if createMode {
|
||||
user, err = services.AuthService.RequirePermission(c.Ctx, constants.PermissionCustomerCreate)
|
||||
} else {
|
||||
user, err = services.AuthService.RequirePermission(c.Ctx, constants.PermissionCustomerUpdate)
|
||||
}
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
item, err := services.CustomerService.SaveCustomerProfile(req, user)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
ret := builders.BuildCustomer(item)
|
||||
return web.JsonData(&ret)
|
||||
}
|
||||
|
||||
func (c *CustomerController) PostCreate() *web.JsonResult {
|
||||
user, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionCustomerCreate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.CreateCustomerRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
item, err := services.CustomerService.CreateCustomer(req, user)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
ret := builders.BuildCustomer(item)
|
||||
return web.JsonData(&ret)
|
||||
}
|
||||
|
||||
func (c *CustomerController) PostUpdate() *web.JsonResult {
|
||||
user, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionCustomerUpdate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.UpdateCustomerRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.CustomerService.UpdateCustomer(req, user); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *CustomerController) PostDelete() *web.JsonResult {
|
||||
user, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionCustomerDelete)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.DeleteCustomerRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.CustomerService.DeleteCustomer(req.ID, *user); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *CustomerController) PostUpdate_status() *web.JsonResult {
|
||||
user, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionCustomerUpdate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.UpdateCustomerStatusRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.CustomerService.UpdateStatus(req.ID, req.Status, user); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
package console
|
||||
|
||||
import (
|
||||
"cs-agent/internal/services"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
)
|
||||
|
||||
type DashboardController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
func (c *DashboardController) GetOverview() *web.JsonResult {
|
||||
rangeValue, _ := params.Get(c.Ctx, "range")
|
||||
return web.JsonData(services.DashboardService.GetOverview(rangeValue))
|
||||
}
|
||||
@@ -0,0 +1,164 @@
|
||||
package console
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
|
||||
"cs-agent/internal/ai/rag"
|
||||
"cs-agent/internal/builders"
|
||||
"cs-agent/internal/pkg/constants"
|
||||
"cs-agent/internal/pkg/dto/request"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/repositories"
|
||||
"cs-agent/internal/services"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/sqls"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
)
|
||||
|
||||
type KnowledgeBaseController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
func (c *KnowledgeBaseController) AnyList() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionKnowledgeBaseView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
cnd := params.NewPagedSqlCnd(c.Ctx,
|
||||
params.QueryFilter{ParamName: "status"},
|
||||
params.QueryFilter{ParamName: "name", Op: params.Like},
|
||||
).Asc("sort_no").Desc("id")
|
||||
list, paging := services.KnowledgeBaseService.FindPageByCnd(cnd)
|
||||
results := make([]response.KnowledgeBaseResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
docCount := repositories.KnowledgeDocumentRepository.CountByKnowledgeBaseID(sqls.DB(), item.ID)
|
||||
faqCount := repositories.KnowledgeFAQRepository.CountByKnowledgeBaseID(sqls.DB(), item.ID)
|
||||
resp := builders.BuildKnowledgeBase(&item)
|
||||
resp.DocumentCount = docCount
|
||||
resp.FAQCount = faqCount
|
||||
results = append(results, resp)
|
||||
}
|
||||
return web.JsonData(&web.PageResult{Results: results, Page: paging})
|
||||
}
|
||||
|
||||
func (c *KnowledgeBaseController) AnyList_all() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionKnowledgeBaseView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
list := services.KnowledgeBaseService.Find(params.NewSqlCnd(c.Ctx,
|
||||
params.QueryFilter{ParamName: "status"},
|
||||
).Asc("sort_no").Desc("id"))
|
||||
results := make([]response.KnowledgeBaseResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
resp := builders.BuildKnowledgeBase(&item)
|
||||
results = append(results, resp)
|
||||
}
|
||||
return web.JsonData(results)
|
||||
}
|
||||
|
||||
func (c *KnowledgeBaseController) GetBy(id int64) *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionKnowledgeBaseView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
item := services.KnowledgeBaseService.Get(id)
|
||||
if item == nil {
|
||||
return web.JsonErrorMsg("知识库不存在")
|
||||
}
|
||||
resp := builders.BuildKnowledgeBase(item)
|
||||
resp.DocumentCount = repositories.KnowledgeDocumentRepository.CountByKnowledgeBaseID(sqls.DB(), item.ID)
|
||||
resp.FAQCount = repositories.KnowledgeFAQRepository.CountByKnowledgeBaseID(sqls.DB(), item.ID)
|
||||
return web.JsonData(resp)
|
||||
}
|
||||
|
||||
func (c *KnowledgeBaseController) PostCreate() *web.JsonResult {
|
||||
user, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionKnowledgeBaseCreate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.CreateKnowledgeBaseRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
item, err := services.KnowledgeBaseService.CreateKnowledgeBase(req, user)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(builders.BuildKnowledgeBase(item))
|
||||
}
|
||||
|
||||
func (c *KnowledgeBaseController) PostUpdate() *web.JsonResult {
|
||||
user, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionKnowledgeBaseUpdate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.UpdateKnowledgeBaseRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.KnowledgeBaseService.UpdateKnowledgeBase(req, user); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *KnowledgeBaseController) PostDelete() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionKnowledgeBaseDelete); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
var req struct {
|
||||
ID int64 `json:"id"`
|
||||
}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.KnowledgeBaseService.DeleteKnowledgeBase(req.ID); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *KnowledgeBaseController) PostUpdate_sort() *web.JsonResult {
|
||||
var ids []int64
|
||||
if err := c.Ctx.ReadJSON(&ids); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.KnowledgeBaseService.UpdateSort(ids); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *KnowledgeBaseController) PostRebuild_index() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionKnowledgeBaseUpdate); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
var req struct {
|
||||
ID int64 `json:"id"`
|
||||
}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
knowledgeBase := services.KnowledgeBaseService.Get(req.ID)
|
||||
if knowledgeBase == nil {
|
||||
return web.JsonErrorMsg("知识库不存在")
|
||||
}
|
||||
|
||||
go func() {
|
||||
ctx := context.Background()
|
||||
if err := rag.Index.RebuildKnowledgeBaseIndex(ctx, req.ID); err != nil {
|
||||
slog.Error("Failed to rebuild knowledge base index", "knowledge_base_id", req.ID, "error", err)
|
||||
}
|
||||
}()
|
||||
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
@@ -0,0 +1,111 @@
|
||||
package console
|
||||
|
||||
import (
|
||||
"cs-agent/internal/builders"
|
||||
"cs-agent/internal/pkg/constants"
|
||||
"cs-agent/internal/pkg/dto/request"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"cs-agent/internal/services"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
)
|
||||
|
||||
type KnowledgeDocumentController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
func (c *KnowledgeDocumentController) AnyList() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionKnowledgeDocumentView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
cnd := params.NewPagedSqlCnd(c.Ctx,
|
||||
params.QueryFilter{ParamName: "knowledgeBaseId"},
|
||||
params.QueryFilter{ParamName: "title", Op: params.Like},
|
||||
).Desc("id")
|
||||
|
||||
if status, ok := params.GetInt64(c.Ctx, "status"); ok {
|
||||
cnd.Where("status = ?", status)
|
||||
} else {
|
||||
cnd.Where("status != ?", enums.StatusDeleted)
|
||||
}
|
||||
if indexStatus, ok := params.Get(c.Ctx, "indexStatus"); ok {
|
||||
if !enums.IsValidKnowledgeDocumentIndexStatus(indexStatus) {
|
||||
return web.JsonErrorMsg("indexStatus参数不合法")
|
||||
}
|
||||
cnd.Where("index_status = ?", indexStatus)
|
||||
}
|
||||
|
||||
list, paging := services.KnowledgeDocumentService.FindPageByCnd(cnd)
|
||||
results := make([]response.KnowledgeDocumentResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
results = append(results, builders.BuildKnowledgeDocument(&item))
|
||||
}
|
||||
return web.JsonData(&web.PageResult{Results: results, Page: paging})
|
||||
}
|
||||
|
||||
func (c *KnowledgeDocumentController) GetBy(id int64) *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionKnowledgeDocumentView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
item := services.KnowledgeDocumentService.Get(id)
|
||||
if item == nil {
|
||||
return web.JsonErrorMsg("文档不存在")
|
||||
}
|
||||
return web.JsonData(builders.BuildKnowledgeDocument(item))
|
||||
}
|
||||
|
||||
func (c *KnowledgeDocumentController) PostCreate() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionKnowledgeDocumentCreate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.CreateKnowledgeDocumentRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
item, err := services.KnowledgeDocumentService.CreateKnowledgeDocument(req, operator)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(builders.BuildKnowledgeDocument(item))
|
||||
}
|
||||
|
||||
func (c *KnowledgeDocumentController) PostUpdate() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionKnowledgeDocumentUpdate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.UpdateKnowledgeDocumentRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.KnowledgeDocumentService.UpdateKnowledgeDocument(req, operator); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *KnowledgeDocumentController) PostDelete() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionKnowledgeDocumentDelete); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
var req struct {
|
||||
ID int64 `json:"id"`
|
||||
}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
if err := services.KnowledgeDocumentService.DeleteKnowledgeDocument(req.ID); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
@@ -0,0 +1,94 @@
|
||||
package console
|
||||
|
||||
import (
|
||||
"cs-agent/internal/builders"
|
||||
"cs-agent/internal/pkg/constants"
|
||||
"cs-agent/internal/pkg/dto/request"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/services"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
)
|
||||
|
||||
type KnowledgeFAQController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
func (c *KnowledgeFAQController) AnyList() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionKnowledgeFAQView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
cnd := params.NewPagedSqlCnd(c.Ctx,
|
||||
params.QueryFilter{ParamName: "knowledgeBaseId"},
|
||||
params.QueryFilter{ParamName: "question", Op: params.Like},
|
||||
params.QueryFilter{ParamName: "indexStatus"},
|
||||
).Desc("id")
|
||||
list, paging := services.KnowledgeFAQService.FindPageByCnd(cnd)
|
||||
results := make([]response.KnowledgeFAQResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
results = append(results, builders.BuildKnowledgeFAQ(&item))
|
||||
}
|
||||
return web.JsonData(&web.PageResult{Results: results, Page: paging})
|
||||
}
|
||||
|
||||
func (c *KnowledgeFAQController) GetBy(id int64) *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionKnowledgeFAQView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
item := services.KnowledgeFAQService.Get(id)
|
||||
if item == nil {
|
||||
return web.JsonErrorMsg("FAQ不存在")
|
||||
}
|
||||
return web.JsonData(builders.BuildKnowledgeFAQ(item))
|
||||
}
|
||||
|
||||
func (c *KnowledgeFAQController) PostCreate() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionKnowledgeFAQCreate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.CreateKnowledgeFAQRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
item, err := services.KnowledgeFAQService.CreateKnowledgeFAQ(req, operator)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(builders.BuildKnowledgeFAQ(item))
|
||||
}
|
||||
|
||||
func (c *KnowledgeFAQController) PostUpdate() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionKnowledgeFAQUpdate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.UpdateKnowledgeFAQRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.KnowledgeFAQService.UpdateKnowledgeFAQ(req, operator); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *KnowledgeFAQController) PostDelete() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionKnowledgeFAQDelete); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
var req struct {
|
||||
ID int64 `json:"id"`
|
||||
}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.KnowledgeFAQService.DeleteKnowledgeFAQ(req.ID); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
@@ -0,0 +1,85 @@
|
||||
package console
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"cs-agent/internal/ai/rag"
|
||||
"cs-agent/internal/pkg/constants"
|
||||
"cs-agent/internal/pkg/dto/request"
|
||||
"cs-agent/internal/services"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
)
|
||||
|
||||
type KnowledgeRetrieveController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
func (c *KnowledgeRetrieveController) PostDebugSearch() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionKnowledgeDocumentView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.KnowledgeSearchRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
resp, err := rag.Answer.DebugSearch(context.Background(), req)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(resp)
|
||||
}
|
||||
|
||||
func (c *KnowledgeRetrieveController) PostDebugAnswer() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionKnowledgeDocumentView)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.KnowledgeAnswerRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
resp, err := rag.Answer.DebugAnswer(context.Background(), req, operator)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(resp)
|
||||
}
|
||||
|
||||
func (c *KnowledgeRetrieveController) PostBuild() *web.JsonResult {
|
||||
req := struct {
|
||||
DocumentID int64 `json:"documentId"`
|
||||
FAQID int64 `json:"faqId"`
|
||||
}{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
if req.DocumentID > 0 {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionKnowledgeDocumentUpdate); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := rag.Answer.BuildDocumentIndex(context.Background(), req.DocumentID); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
if req.FAQID > 0 {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionKnowledgeFAQUpdate); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := rag.Index.IndexFAQByID(context.Background(), req.FAQID); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
return web.JsonErrorMsg("documentId或faqId不能为空")
|
||||
}
|
||||
@@ -0,0 +1,77 @@
|
||||
package console
|
||||
|
||||
import (
|
||||
"cs-agent/internal/builders"
|
||||
"cs-agent/internal/pkg/constants"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/services"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
)
|
||||
|
||||
type KnowledgeRetrieveLogController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
func (c *KnowledgeRetrieveLogController) AnyList() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionKnowledgeDocumentView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
cnd := params.NewPagedSqlCnd(c.Ctx,
|
||||
params.QueryFilter{ParamName: "knowledgeBaseId"},
|
||||
params.QueryFilter{ParamName: "question", Op: params.Like},
|
||||
params.QueryFilter{ParamName: "channel"},
|
||||
params.QueryFilter{ParamName: "scene"},
|
||||
params.QueryFilter{ParamName: "chunkProvider"},
|
||||
).Desc("id")
|
||||
|
||||
if answerStatus, ok := params.GetInt64(c.Ctx, "answerStatus"); ok && answerStatus > 0 {
|
||||
cnd.Where("answer_status = ?", answerStatus)
|
||||
}
|
||||
if rerankEnabled, ok := params.GetInt64(c.Ctx, "rerankEnabled"); ok {
|
||||
cnd.Where("rerank_enabled = ?", rerankEnabled > 0)
|
||||
}
|
||||
|
||||
queryParams := params.NewQueryParams(c.Ctx)
|
||||
queryParams.Cnd = *cnd
|
||||
list, paging := services.KnowledgeRetrieveLogService.FindPageByParams(queryParams)
|
||||
results := make([]response.KnowledgeRetrieveLogResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
resp := builders.BuildKnowledgeRetrieveLog(&item)
|
||||
if knowledgeBase := services.KnowledgeBaseService.Get(item.KnowledgeBaseID); knowledgeBase != nil {
|
||||
resp.KnowledgeBaseName = knowledgeBase.Name
|
||||
}
|
||||
results = append(results, resp)
|
||||
}
|
||||
return web.JsonData(&web.PageResult{Results: results, Page: paging})
|
||||
}
|
||||
|
||||
func (c *KnowledgeRetrieveLogController) GetBy(id int64) *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionKnowledgeDocumentView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
logItem := services.KnowledgeRetrieveLogService.Get(id)
|
||||
if logItem == nil {
|
||||
return web.JsonErrorMsg("检索日志不存在")
|
||||
}
|
||||
|
||||
logResp := builders.BuildKnowledgeRetrieveLog(logItem)
|
||||
if knowledgeBase := services.KnowledgeBaseService.Get(logItem.KnowledgeBaseID); knowledgeBase != nil {
|
||||
logResp.KnowledgeBaseName = knowledgeBase.Name
|
||||
}
|
||||
|
||||
hits := services.KnowledgeRetrieveLogService.FindHitsByRetrieveLogID(id)
|
||||
hitResults := make([]response.KnowledgeRetrieveHitResponse, 0, len(hits))
|
||||
for _, item := range hits {
|
||||
hitResults = append(hitResults, builders.BuildKnowledgeRetrieveHitResponse(&item))
|
||||
}
|
||||
|
||||
return web.JsonData(response.KnowledgeRetrieveLogDetailResponse{
|
||||
Log: logResp,
|
||||
Hits: hitResults,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,70 @@
|
||||
package console
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"cs-agent/internal/pkg/constants"
|
||||
"cs-agent/internal/pkg/dto/request"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/services"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
)
|
||||
|
||||
type MCPController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
func (c *MCPController) AnyList_servers() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionMCPView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(response.BuildMCPServerInfoResponses(services.MCPDebugService.ListServers()))
|
||||
}
|
||||
|
||||
func (c *MCPController) PostTest_connection() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionMCPView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.MCPServerDebugRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
result, err := services.MCPDebugService.TestConnection(context.Background(), req.ServerCode)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(response.BuildMCPConnectionResponse(result))
|
||||
}
|
||||
|
||||
func (c *MCPController) PostList_tools() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionMCPView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.MCPServerDebugRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
result, err := services.MCPDebugService.ListTools(context.Background(), req.ServerCode)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(response.BuildMCPToolInfoResponses(result))
|
||||
}
|
||||
|
||||
func (c *MCPController) PostCall_tool() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionMCPCall); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
req := request.MCPCallToolRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
result, err := services.MCPDebugService.CallTool(context.Background(), req.ServerCode, req.ToolName, req.Arguments)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(response.BuildMCPCallToolResponse(result))
|
||||
}
|
||||
@@ -0,0 +1,71 @@
|
||||
package console
|
||||
|
||||
import (
|
||||
"cs-agent/internal/pkg/constants"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/services"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/common/strs"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
)
|
||||
|
||||
type PermissionController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
func (c *PermissionController) AnyList() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionPermissionView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
cnd := params.NewPagedSqlCnd(c.Ctx,
|
||||
params.QueryFilter{ParamName: "groupName"},
|
||||
params.QueryFilter{ParamName: "type"},
|
||||
params.QueryFilter{ParamName: "status"},
|
||||
).Desc("id")
|
||||
|
||||
if keyword, _ := params.Get(c.Ctx, "keyword"); strs.IsNotBlank(keyword) {
|
||||
cnd.Where("(name LIKE ? OR code LIKE ?)", "%"+keyword+"%", "%"+keyword+"%")
|
||||
}
|
||||
|
||||
list, paging := services.PermissionService.FindPageByCnd(cnd)
|
||||
results := make([]response.PermissionResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
results = append(results, response.PermissionResponse{
|
||||
ID: item.ID,
|
||||
Name: item.Name,
|
||||
Code: item.Code,
|
||||
Type: item.Type,
|
||||
GroupName: item.GroupName,
|
||||
Method: item.Method,
|
||||
ApiPath: item.APIPath,
|
||||
Status: item.Status,
|
||||
SortNo: item.SortNo,
|
||||
})
|
||||
}
|
||||
return web.JsonData(&web.PageResult{Results: results, Page: paging})
|
||||
}
|
||||
|
||||
func (c *PermissionController) GetBy(id int64) *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionPermissionView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
item := services.PermissionService.Get(id)
|
||||
if item == nil {
|
||||
return web.JsonErrorMsg("权限不存在")
|
||||
}
|
||||
return web.JsonData(&response.PermissionResponse{
|
||||
ID: item.ID,
|
||||
Name: item.Name,
|
||||
Code: item.Code,
|
||||
Type: item.Type,
|
||||
GroupName: item.GroupName,
|
||||
Method: item.Method,
|
||||
ApiPath: item.APIPath,
|
||||
Status: item.Status,
|
||||
SortNo: item.SortNo,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,118 @@
|
||||
package console
|
||||
|
||||
import (
|
||||
"cs-agent/internal/pkg/constants"
|
||||
"cs-agent/internal/pkg/dto/request"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"cs-agent/internal/services"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/sqls"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
)
|
||||
|
||||
type QuickReplyController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
func (c *QuickReplyController) AnyList() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionQuickReplyView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
cnd := params.NewPagedSqlCnd(c.Ctx,
|
||||
params.QueryFilter{ParamName: "status"},
|
||||
params.QueryFilter{ParamName: "groupName"},
|
||||
params.QueryFilter{ParamName: "title", Op: params.Like},
|
||||
).Asc("sort_no").Desc("id")
|
||||
|
||||
list, paging := services.QuickReplyService.FindPageByCnd(cnd)
|
||||
results := make([]response.QuickReplyResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
results = append(results, response.QuickReplyResponse{
|
||||
ID: item.ID,
|
||||
GroupName: item.GroupName,
|
||||
Title: item.Title,
|
||||
Content: item.Content,
|
||||
Status: item.Status,
|
||||
SortNo: item.SortNo,
|
||||
})
|
||||
}
|
||||
return web.JsonData(&web.PageResult{Results: results, Page: paging})
|
||||
}
|
||||
|
||||
func (c *QuickReplyController) GetList_all() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionQuickReplyView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
list := services.QuickReplyService.Find(sqls.NewCnd().Eq("status", enums.StatusOk).Asc("sort_no").Desc("id"))
|
||||
results := make([]response.QuickReplyResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
results = append(results, response.QuickReplyResponse{
|
||||
ID: item.ID,
|
||||
GroupName: item.GroupName,
|
||||
Title: item.Title,
|
||||
Content: item.Content,
|
||||
Status: item.Status,
|
||||
SortNo: item.SortNo,
|
||||
})
|
||||
}
|
||||
return web.JsonData(results)
|
||||
}
|
||||
|
||||
func (c *QuickReplyController) PostCreate() *web.JsonResult {
|
||||
user, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionQuickReplyCreate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.CreateQuickReplyRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
item, err := services.QuickReplyService.CreateQuickReply(req, user)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(&response.QuickReplyResponse{
|
||||
ID: item.ID,
|
||||
GroupName: item.GroupName,
|
||||
Title: item.Title,
|
||||
Content: item.Content,
|
||||
Status: item.Status,
|
||||
SortNo: item.SortNo,
|
||||
})
|
||||
}
|
||||
|
||||
func (c *QuickReplyController) PostUpdate() *web.JsonResult {
|
||||
user, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionQuickReplyUpdate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.UpdateQuickReplyRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.QuickReplyService.UpdateQuickReply(req, user); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *QuickReplyController) PostDelete() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionQuickReplyDelete); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.DeleteQuickReplyRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.QuickReplyService.DeleteQuickReply(req.ID); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
@@ -0,0 +1,187 @@
|
||||
package console
|
||||
|
||||
import (
|
||||
"cs-agent/internal/pkg/constants"
|
||||
"cs-agent/internal/pkg/dto/request"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/services"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/sqls"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
)
|
||||
|
||||
type RoleController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
func (c *RoleController) AnyList() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionRoleView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
cnd := params.NewPagedSqlCnd(c.Ctx,
|
||||
params.QueryFilter{ParamName: "status"},
|
||||
params.QueryFilter{ParamName: "code", Op: params.Like},
|
||||
).Asc("sort_no").Desc("id")
|
||||
list, paging := services.RoleService.FindPageByCnd(cnd)
|
||||
results := make([]response.RoleResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
results = append(results, response.RoleResponse{
|
||||
ID: item.ID,
|
||||
Name: item.Name,
|
||||
Code: item.Code,
|
||||
Status: item.Status,
|
||||
IsSystem: item.IsSystem,
|
||||
SortNo: item.SortNo,
|
||||
})
|
||||
}
|
||||
return web.JsonData(&web.PageResult{Results: results, Page: paging})
|
||||
}
|
||||
|
||||
func (c *RoleController) GetList_all() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionRoleView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
list := services.RoleService.Find(sqls.NewCnd().Asc("sort_no").Desc("id"))
|
||||
results := make([]response.RoleResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
results = append(results, response.RoleResponse{
|
||||
ID: item.ID,
|
||||
Name: item.Name,
|
||||
Code: item.Code,
|
||||
Status: item.Status,
|
||||
IsSystem: item.IsSystem,
|
||||
SortNo: item.SortNo,
|
||||
})
|
||||
}
|
||||
return web.JsonData(results)
|
||||
}
|
||||
|
||||
func (c *RoleController) GetBy(id int64) *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionRoleView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
item := services.RoleService.Get(id)
|
||||
if item == nil {
|
||||
return web.JsonErrorMsg("角色不存在")
|
||||
}
|
||||
|
||||
permissionCodes := make([]string, 0)
|
||||
list := services.RolePermissionService.Find(sqls.NewCnd().Eq("role_id", item.ID))
|
||||
for _, relation := range list {
|
||||
permission := services.PermissionService.Get(relation.PermissionID)
|
||||
if permission != nil {
|
||||
permissionCodes = append(permissionCodes, permission.Code)
|
||||
}
|
||||
}
|
||||
return web.JsonData(&response.RoleResponse{
|
||||
ID: item.ID,
|
||||
Name: item.Name,
|
||||
Code: item.Code,
|
||||
Status: item.Status,
|
||||
IsSystem: item.IsSystem,
|
||||
SortNo: item.SortNo,
|
||||
Permissions: permissionCodes,
|
||||
})
|
||||
}
|
||||
|
||||
func (c *RoleController) PostCreate() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionRoleCreate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.CreateRoleRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
role, err := services.RoleService.CreateRole(req, operator)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(&response.RoleResponse{
|
||||
ID: role.ID,
|
||||
Name: role.Name,
|
||||
Code: role.Code,
|
||||
Status: role.Status,
|
||||
IsSystem: role.IsSystem,
|
||||
SortNo: role.SortNo,
|
||||
})
|
||||
}
|
||||
|
||||
func (c *RoleController) PostUpdate() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionRoleUpdate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.UpdateRoleRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.RoleService.UpdateRole(req, operator); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *RoleController) PostDelete() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionRoleDelete); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.DeleteRoleRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.RoleService.DeleteRole(req.ID); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *RoleController) PostUpdate_status() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionRoleUpdate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.UpdateRoleStatusRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.RoleService.UpdateStatus(req.ID, req.Status, operator); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *RoleController) PostAssign_permission() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionRoleAssignPermission)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.AssignPermissionRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.RoleService.AssignPermissions(req.RoleID, req.PermissionIDs, operator); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *RoleController) PostUpdate_sort() *web.JsonResult {
|
||||
var ids []int64
|
||||
if err := c.Ctx.ReadJSON(&ids); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.RoleService.UpdateSort(ids); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
@@ -0,0 +1,83 @@
|
||||
package console
|
||||
|
||||
import (
|
||||
"cs-agent/internal/pkg/constants"
|
||||
"cs-agent/internal/pkg/dto/request"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/pkg/utils"
|
||||
"cs-agent/internal/services"
|
||||
"time"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
)
|
||||
|
||||
type SessionController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
func (c *SessionController) AnyList() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionSessionView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
cnd := params.NewPagedSqlCnd(c.Ctx,
|
||||
params.QueryFilter{ParamName: "userId"},
|
||||
params.QueryFilter{ParamName: "tokenType"},
|
||||
params.QueryFilter{ParamName: "clientType"},
|
||||
).Desc("id")
|
||||
list, paging := services.LoginSessionService.FindPageByCnd(cnd)
|
||||
results := make([]response.SessionResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
username := ""
|
||||
if user := services.UserService.Get(item.UserID); user != nil {
|
||||
username = user.Username
|
||||
}
|
||||
results = append(results, response.SessionResponse{
|
||||
ID: item.ID,
|
||||
UserID: item.UserID,
|
||||
Username: username,
|
||||
TokenType: item.TokenType,
|
||||
ClientType: item.ClientType,
|
||||
ClientIP: item.ClientIP,
|
||||
UserAgent: item.UserAgent,
|
||||
ExpiredAt: item.ExpiredAt.Format(time.DateTime),
|
||||
RevokedAt: utils.FormatTimePtr(item.RevokedAt),
|
||||
LastSeenAt: utils.FormatTimePtr(item.LastSeenAt),
|
||||
})
|
||||
}
|
||||
return web.JsonData(&web.PageResult{Results: results, Page: paging})
|
||||
}
|
||||
|
||||
func (c *SessionController) PostRevoke() *web.JsonResult {
|
||||
user, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionSessionRevoke)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.RevokeSessionRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.LoginSessionService.Revoke(req.ID, user.UserID, user.Username); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *SessionController) PostRevokeByUser() *web.JsonResult {
|
||||
user, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionSessionRevoke)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.RevokeUserSessionsRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.LoginSessionService.RevokeByUser(req.UserID, user.UserID, user.Username); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
@@ -0,0 +1,280 @@
|
||||
package console
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"cs-agent/internal/builders"
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/constants"
|
||||
"cs-agent/internal/pkg/dto/request"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"cs-agent/internal/pkg/errorsx"
|
||||
"cs-agent/internal/pkg/utils"
|
||||
"cs-agent/internal/services"
|
||||
|
||||
"github.com/kataras/iris/v12"
|
||||
"github.com/mlogclub/simple/web"
|
||||
"github.com/mlogclub/simple/web/params"
|
||||
)
|
||||
|
||||
type SkillDefinitionController struct {
|
||||
Ctx iris.Context
|
||||
}
|
||||
|
||||
func (c *SkillDefinitionController) AnyList() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionSkillDefinitionView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
cnd := params.NewPagedSqlCnd(c.Ctx,
|
||||
params.QueryFilter{ParamName: "status"},
|
||||
params.QueryFilter{ParamName: "name", Op: params.Like},
|
||||
params.QueryFilter{ParamName: "code", Op: params.Like},
|
||||
).Asc("priority").Desc("id")
|
||||
list, paging := services.SkillDefinitionService.FindPageByCnd(cnd)
|
||||
results := make([]response.SkillDefinitionResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
results = append(results, builders.BuildSkillDefinitionResponse(&item))
|
||||
}
|
||||
return web.JsonData(&web.PageResult{Results: results, Page: paging})
|
||||
}
|
||||
|
||||
func (c *SkillDefinitionController) GetList_all() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionSkillDefinitionView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
list := services.SkillDefinitionService.Find(params.NewSqlCnd(c.Ctx,
|
||||
params.QueryFilter{ParamName: "status"},
|
||||
).Asc("priority").Desc("id"))
|
||||
results := make([]response.SkillDefinitionResponse, 0, len(list))
|
||||
for _, item := range list {
|
||||
results = append(results, builders.BuildSkillDefinitionResponse(&item))
|
||||
}
|
||||
return web.JsonData(results)
|
||||
}
|
||||
|
||||
func (c *SkillDefinitionController) GetBy(id int64) *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionSkillDefinitionView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
item := services.SkillDefinitionService.Get(id)
|
||||
if item == nil {
|
||||
return web.JsonErrorMsg("Skill 不存在")
|
||||
}
|
||||
return web.JsonData(builders.BuildSkillDefinitionResponse(item))
|
||||
}
|
||||
|
||||
func (c *SkillDefinitionController) PostCreate() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionSkillDefinitionCreate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.CreateSkillDefinitionRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := validateSkillDefinitionRequest(req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if services.SkillDefinitionService.Take("code = ?", strings.TrimSpace(req.Code)) != nil {
|
||||
return web.JsonErrorMsg("Skill 编码已存在")
|
||||
}
|
||||
|
||||
item := &models.SkillDefinition{
|
||||
Code: strings.TrimSpace(req.Code),
|
||||
Name: strings.TrimSpace(req.Name),
|
||||
Description: strings.TrimSpace(req.Description),
|
||||
Prompt: strings.TrimSpace(req.Prompt),
|
||||
ExecutionMode: normalizeExecutionMode(req.ExecutionMode),
|
||||
ExecutionConfig: normalizeExecutionConfig(req.ExecutionConfig),
|
||||
Priority: 0,
|
||||
Status: enums.StatusOk,
|
||||
Remark: strings.TrimSpace(req.Remark),
|
||||
AuditFields: utils.BuildAuditFields(operator),
|
||||
}
|
||||
if err := services.SkillDefinitionService.Create(item); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(builders.BuildSkillDefinitionResponse(item))
|
||||
}
|
||||
|
||||
func (c *SkillDefinitionController) PostUpdate() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionSkillDefinitionUpdate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.UpdateSkillDefinitionRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if req.ID <= 0 {
|
||||
return web.JsonErrorMsg("Skill ID 不合法")
|
||||
}
|
||||
if err := validateSkillDefinitionRequest(req.CreateSkillDefinitionRequest); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
item := services.SkillDefinitionService.Get(req.ID)
|
||||
if item == nil {
|
||||
return web.JsonErrorMsg("Skill 不存在")
|
||||
}
|
||||
exists := services.SkillDefinitionService.Take("code = ? AND id <> ?", strings.TrimSpace(req.Code), req.ID)
|
||||
if exists != nil {
|
||||
return web.JsonErrorMsg("Skill 编码已存在")
|
||||
}
|
||||
|
||||
if err := services.SkillDefinitionService.Updates(req.ID, map[string]any{
|
||||
"code": strings.TrimSpace(req.Code),
|
||||
"name": strings.TrimSpace(req.Name),
|
||||
"description": strings.TrimSpace(req.Description),
|
||||
"prompt": strings.TrimSpace(req.Prompt),
|
||||
"execution_mode": normalizeExecutionMode(req.ExecutionMode),
|
||||
"execution_config": normalizeExecutionConfig(req.ExecutionConfig),
|
||||
"remark": strings.TrimSpace(req.Remark),
|
||||
"update_user_id": operator.UserID,
|
||||
"update_user_name": operator.Username,
|
||||
"updated_at": time.Now(),
|
||||
}); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *SkillDefinitionController) PostUpdate_status() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionSkillDefinitionUpdate)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.UpdateSkillDefinitionStatusRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if req.ID <= 0 {
|
||||
return web.JsonErrorMsg("Skill ID 不合法")
|
||||
}
|
||||
if !enums.IsValidStatus(req.Status) {
|
||||
return web.JsonErrorMsg("状态值不合法")
|
||||
}
|
||||
if services.SkillDefinitionService.Get(req.ID) == nil {
|
||||
return web.JsonErrorMsg("Skill 不存在")
|
||||
}
|
||||
|
||||
if err := services.SkillDefinitionService.Updates(req.ID, map[string]any{
|
||||
"status": req.Status,
|
||||
"update_user_id": operator.UserID,
|
||||
"update_user_name": operator.Username,
|
||||
"updated_at": time.Now(),
|
||||
}); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *SkillDefinitionController) PostDelete() *web.JsonResult {
|
||||
operator, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionSkillDefinitionDelete)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.DeleteSkillDefinitionRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if req.ID <= 0 {
|
||||
return web.JsonErrorMsg("Skill ID 不合法")
|
||||
}
|
||||
if services.SkillDefinitionService.Get(req.ID) == nil {
|
||||
return web.JsonErrorMsg("Skill 不存在")
|
||||
}
|
||||
if err := services.SkillDefinitionService.Updates(req.ID, map[string]any{
|
||||
"status": enums.StatusDeleted,
|
||||
"update_user_id": operator.UserID,
|
||||
"update_user_name": operator.Username,
|
||||
"updated_at": time.Now(),
|
||||
}); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *SkillDefinitionController) PostUpdate_priority() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionSkillDefinitionUpdate); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
var ids []int64
|
||||
if err := c.Ctx.ReadJSON(&ids); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
if err := services.SkillDefinitionService.UpdatePriority(ids); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonSuccess()
|
||||
}
|
||||
|
||||
func (c *SkillDefinitionController) PostDebug_run() *web.JsonResult {
|
||||
if _, err := services.AuthService.RequirePermission(c.Ctx, constants.PermissionSkillDefinitionView); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
|
||||
req := request.SkillDebugRunRequest{}
|
||||
if err := params.ReadJSON(c.Ctx, &req); err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
resp, err := services.SkillRuntimeService.DebugRun(context.Background(), req)
|
||||
if err != nil {
|
||||
return web.JsonError(err)
|
||||
}
|
||||
return web.JsonData(resp)
|
||||
}
|
||||
|
||||
func validateSkillDefinitionRequest(req request.CreateSkillDefinitionRequest) error {
|
||||
code := strings.TrimSpace(req.Code)
|
||||
name := strings.TrimSpace(req.Name)
|
||||
prompt := strings.TrimSpace(req.Prompt)
|
||||
if code == "" {
|
||||
return errorsx.InvalidParam("Skill 编码不能为空")
|
||||
}
|
||||
if name == "" {
|
||||
return errorsx.InvalidParam("Skill 名称不能为空")
|
||||
}
|
||||
mode := normalizeExecutionMode(req.ExecutionMode)
|
||||
if prompt == "" {
|
||||
return errorsx.InvalidParam("Prompt 不能为空")
|
||||
}
|
||||
switch mode {
|
||||
case enums.SkillExecutionModePromptOnly:
|
||||
case enums.SkillExecutionModeMCPTool:
|
||||
configText := strings.TrimSpace(req.ExecutionConfig)
|
||||
if configText == "" {
|
||||
return errorsx.InvalidParam("MCP工具模式必须填写ExecutionConfig")
|
||||
}
|
||||
var payload map[string]any
|
||||
if err := json.Unmarshal([]byte(configText), &payload); err != nil {
|
||||
return errorsx.InvalidParam("ExecutionConfig 必须是合法JSON")
|
||||
}
|
||||
default:
|
||||
return errorsx.InvalidParam("Skill 执行模式不合法")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func normalizeExecutionMode(mode enums.SkillExecutionMode) enums.SkillExecutionMode {
|
||||
if mode == "" {
|
||||
return enums.SkillExecutionModePromptOnly
|
||||
}
|
||||
return mode
|
||||
}
|
||||
|
||||
func normalizeExecutionConfig(raw string) string {
|
||||
return strings.TrimSpace(raw)
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user