Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
48 changes: 45 additions & 3 deletions internal/agent/execute.go
Original file line number Diff line number Diff line change
Expand Up @@ -66,9 +66,10 @@ func (e *Engine) runAgent(ctx context.Context, req Request, chatModel model.Tool
logger.Infof("[Agent] 工具: name=%s, desc=%s", info.Name, truncateStr(info.Desc, 80))
}

// 5. 构建 system prompt
systemPrompt := buildReActSystemPrompt(ctx, userTools)
logger.Infof("[Agent] SystemPrompt (前200字符): %s", truncateStr(systemPrompt, 200))
// 5. 构建 system prompt(ReAct 规则 + 摘要/记忆/用户上下文增强)
baseSystemPrompt := buildReActSystemPrompt(ctx, userTools)
systemPrompt := buildEnhancedSystemPromptForAgent(baseSystemPrompt, req.Summary, req.Memories, req.UserCtx)
logger.Infof("[Agent] SystemPrompt (前400字符): %s", truncateStr(systemPrompt, 400))

// 6. 构建输入消息(历史 + 当前问题)
inputMessages := buildInputMessages(req.Query, req.History)
Expand Down Expand Up @@ -276,3 +277,44 @@ func isToolChoiceUnsupportedError(errMsg string) bool {
}
return false
}

// buildEnhancedSystemPromptForAgent 在 ReAct 系统提示词上注入:时间/用户信息 + 对话摘要 + 用户记忆
// 与快速模式的 buildEnhancedSystemPrompt 逻辑保持一致,确保双模式行为统一
func buildEnhancedSystemPromptForAgent(base string, summary *entity.ChatSummary, memories []entity.UserMemory, userCtx PromptUserContext) string {
var extras []string

userInfo := "## 当前信息\n"
if userCtx.TimeStr != "" {
userInfo += "- 当前时间:" + userCtx.TimeStr + "\n"
}
if userCtx.Username != "" {
userInfo += "- 用户:" + userCtx.Username + "\n"
}
if userCtx.Role != "" {
userInfo += "- 角色:" + userCtx.Role + "\n"
}
if userInfo != "## 当前信息\n" {
extras = append(extras, userInfo)
}

if summary != nil && summary.Summary != "" {
extras = append(extras, "## 本次对话摘要\n"+summary.Summary)
}

if len(memories) > 0 {
var memoryText strings.Builder
memoryText.WriteString("## 关于用户的已知信息\n")
for _, m := range memories {
memoryText.WriteString("- ")
memoryText.WriteString(m.Content)
memoryText.WriteString("\n")
}
extras = append(extras, memoryText.String())
}

if len(extras) == 0 {
return base
}

return base + "\n\n" + strings.Join(extras, "\n\n")
}
10 changes: 10 additions & 0 deletions internal/agent/types.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,13 @@ import (
"solvify-agent/internal/model/entity"
)

type PromptUserContext struct {
ID string
Username string
Role string
TimeStr string
}

// Request 描述 Agent 执行请求
type Request struct {
UserID string // 用户 ID(用于知识库检索权限)
Expand All @@ -13,6 +20,9 @@ type Request struct {
KnowledgeBaseIDs []string // 知识库 ID 列表
ModelID string // 模型 ID
ModelType string // 模型类型(user/system)
Summary *entity.ChatSummary // 会话摘要(长对话压缩内容)
Memories []entity.UserMemory // 用户长期记忆(偏好/事实/约束/决策)
UserCtx PromptUserContext // 用户基本信息 + 当前时间
}

// Event 描述 Agent SSE 事件
Expand Down
4 changes: 4 additions & 0 deletions internal/repository/chat_message_interface.go
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,11 @@ type ChatMessageSearchRow struct {
type ChatMessageRepo interface {
Create(ctx context.Context, message *entity.ChatMessage) error
FindBySessionID(ctx context.Context, sessionID string) ([]entity.ChatMessage, error)
// FindBySessionIDForContext 摘要/记忆抽取场景专用:全量消息但只取 5 个必要字段,sources/metadata 不传
FindBySessionIDForContext(ctx context.Context, sessionID string) ([]entity.ChatMessage, error)
FindRecent(ctx context.Context, sessionID string, limit int) ([]entity.ChatMessage, error)
// FindRecentForContext 上下文构建专用:只取构建 Prompt 需要的 5 个字段,避免 sources/metadata 大字段浪费 IO
FindRecentForContext(ctx context.Context, sessionID string, limit int) ([]entity.ChatMessage, error)
DeleteBySessionID(ctx context.Context, sessionID string) error
// SearchByKeyword 按关键字搜索用户历史消息
SearchByKeyword(ctx context.Context, userID, query string, topK int) ([]ChatMessageSearchRow, error)
Expand Down
33 changes: 33 additions & 0 deletions internal/repository/chat_message_repository.go
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,19 @@ func (r *chatMessageRepository) FindBySessionID(ctx context.Context, sessionID s
return messages, err
}

// FindBySessionIDForContext 会话摘要/记忆抽取场景专用:
// 全量消息,但只 SELECT 构建 Prompt 需要的 5 个字段,sources/metadata 不传
// 对 50 轮以上长会话可减少 90%+ 的数据传输
func (r *chatMessageRepository) FindBySessionIDForContext(ctx context.Context, sessionID string) ([]entity.ChatMessage, error) {
var messages []entity.ChatMessage
err := r.db.WithContext(ctx).
Select("id, session_id, role, content, created_at").
Where("session_id = ?", sessionID).
Order("created_at ASC").
Find(&messages).Error
return messages, err
}

// FindRecent 获取会话的最近 N 条消息
func (r *chatMessageRepository) FindRecent(ctx context.Context, sessionID string, limit int) ([]entity.ChatMessage, error) {
var messages []entity.ChatMessage
Expand All @@ -50,6 +63,26 @@ func (r *chatMessageRepository) FindRecent(ctx context.Context, sessionID string
return messages, err
}

// FindRecentForContext 上下文构建专用:只取构建 Prompt 必需的 5 个字段
// 避免 SELECT * 把 sources/metadata 两个 JSON 大字段也传回来(可能 > 100KB/条),
// 构建上下文只看 role+content,单条消息体积从 100KB 降到 ~100Byte
func (r *chatMessageRepository) FindRecentForContext(ctx context.Context, sessionID string, limit int) ([]entity.ChatMessage, error) {
var messages []entity.ChatMessage
err := r.db.WithContext(ctx).
Select("id, session_id, role, content, created_at").
Where("session_id = ?", sessionID).
Order("created_at DESC").
Limit(limit).
Find(&messages).Error

// 反转顺序,使其按时间正序
for i, j := 0, len(messages)-1; i < j; i, j = i+1, j-1 {
messages[i], messages[j] = messages[j], messages[i]
}

return messages, err
}

// DeleteBySessionID 删除会话的所有消息
func (r *chatMessageRepository) DeleteBySessionID(ctx context.Context, sessionID string) error {
return r.db.WithContext(ctx).Where("session_id = ?", sessionID).Delete(&entity.ChatMessage{}).Error
Expand Down
Loading
Loading