mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-08-02 05:47:31 +08:00
354 lines
12 KiB
Go
354 lines
12 KiB
Go
|
|
//
|
||
|
|
// Copyright 2026 The InfiniFlow Authors. All Rights Reserved.
|
||
|
|
//
|
||
|
|
// Licensed under the Apache License, Version 2.0 (the "License");
|
||
|
|
// you may not use this file except in compliance with the License.
|
||
|
|
// You may obtain a copy of the License at
|
||
|
|
//
|
||
|
|
// http://www.apache.org/licenses/LICENSE-2.0
|
||
|
|
//
|
||
|
|
// Unless required by applicable law or agreed to in writing, software
|
||
|
|
// distributed under the License is distributed on an "AS IS" BASIS,
|
||
|
|
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||
|
|
// See the License for the specific language governing permissions and
|
||
|
|
// limitations under the License.
|
||
|
|
//
|
||
|
|
|
||
|
|
// memory_extractor.go — async memory extraction worker.
|
||
|
|
//
|
||
|
|
// Port of the Python task-executor memory path:
|
||
|
|
//
|
||
|
|
// rag/svr/task_executor.py:handle_task (task_type == "memory")
|
||
|
|
// api/db/joint_services/memory_message_service.py:
|
||
|
|
// handle_save_to_memory_task / save_extracted_to_memory_only / extract_by_llm
|
||
|
|
//
|
||
|
|
// QueueSaveToMemoryTask persists the raw message and enqueues a
|
||
|
|
// task_type="memory" message on the Redis stream te.<priority>.common.
|
||
|
|
// StartTaskConsumer drains that stream (consumer group
|
||
|
|
// "rag_flow_svr_task_broker", same as the Python executor), runs LLM
|
||
|
|
// extraction for the non-raw memory types configured on the memory, and
|
||
|
|
// persists the extracted messages with source_id pointing at the raw
|
||
|
|
// message so listMemoryMessages can aggregate them under `extract`.
|
||
|
|
package service
|
||
|
|
|
||
|
|
import (
|
||
|
|
"context"
|
||
|
|
"encoding/json"
|
||
|
|
"errors"
|
||
|
|
"fmt"
|
||
|
|
"strings"
|
||
|
|
"time"
|
||
|
|
|
||
|
|
"ragflow/internal/common"
|
||
|
|
"ragflow/internal/dao"
|
||
|
|
redisengine "ragflow/internal/engine/redis"
|
||
|
|
"ragflow/internal/entity"
|
||
|
|
models "ragflow/internal/entity/models"
|
||
|
|
"ragflow/internal/utility"
|
||
|
|
|
||
|
|
"go.uber.org/zap"
|
||
|
|
)
|
||
|
|
|
||
|
|
// memoryTaskConsumerGroup matches Python common.constants.SVR_CONSUMER_GROUP_NAME
|
||
|
|
// so Go and Python executors never double-process the same stream entry.
|
||
|
|
const memoryTaskConsumerGroup = "rag_flow_svr_task_broker"
|
||
|
|
|
||
|
|
// memoryTimeLayout is the storage format for valid_at / invalid_at,
|
||
|
|
// matching timestamp_to_date on the Python side.
|
||
|
|
const memoryTimeLayout = "2006-01-02 15:04:05"
|
||
|
|
|
||
|
|
// extractedMemory is one LLM-extracted memory item ready for persistence.
|
||
|
|
type extractedMemory struct {
|
||
|
|
MessageType string
|
||
|
|
Content string
|
||
|
|
ValidAt string
|
||
|
|
InvalidAt string // empty means still valid
|
||
|
|
}
|
||
|
|
|
||
|
|
// StartTaskConsumer is the long-running loop that drains memory
|
||
|
|
// extraction tasks from the Redis stream. It returns when ctx is
|
||
|
|
// cancelled. Per-message failures are logged and acked so one bad
|
||
|
|
// message cannot stall the stream.
|
||
|
|
func (s *MemoryMessageService) StartTaskConsumer(ctx context.Context) {
|
||
|
|
redisClient := redisengine.Get()
|
||
|
|
if redisClient == nil {
|
||
|
|
common.Error("memory task consumer: Redis is not available", nil)
|
||
|
|
return
|
||
|
|
}
|
||
|
|
consumerName := fmt.Sprintf("go_memory_extractor_%s", utility.GenerateUUID())
|
||
|
|
queueName := memoryTaskQueueName(0)
|
||
|
|
common.Info(fmt.Sprintf("Memory task consumer %s started on queue %s", consumerName, queueName))
|
||
|
|
|
||
|
|
for {
|
||
|
|
if ctx.Err() != nil {
|
||
|
|
return
|
||
|
|
}
|
||
|
|
msg, err := redisClient.QueueConsumer(queueName, memoryTaskConsumerGroup, consumerName, ">")
|
||
|
|
if err != nil {
|
||
|
|
common.Error("memory task consumer: consume error", err)
|
||
|
|
select {
|
||
|
|
case <-time.After(time.Second):
|
||
|
|
case <-ctx.Done():
|
||
|
|
return
|
||
|
|
}
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
if msg == nil {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
payload := msg.GetMessage()
|
||
|
|
if taskType, _ := payload["task_type"].(string); taskType != "memory" {
|
||
|
|
common.Warn(fmt.Sprintf("memory task consumer: skip task_type %q", taskType))
|
||
|
|
msg.Ack()
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
if err := s.HandleSaveToMemoryTask(ctx, payload); err != nil {
|
||
|
|
common.Error("memory task consumer: handle task failed", err)
|
||
|
|
}
|
||
|
|
msg.Ack()
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// HandleSaveToMemoryTask processes one queued memory task. Mirrors
|
||
|
|
// Python handle_save_to_memory_task: validate the task row, then
|
||
|
|
// extract + persist, settling task progress on the way out.
|
||
|
|
func (s *MemoryMessageService) HandleSaveToMemoryTask(ctx context.Context, payload map[string]any) error {
|
||
|
|
taskID, _ := payload["id"].(string)
|
||
|
|
if taskID == "" {
|
||
|
|
taskID, _ = payload["task_id"].(string)
|
||
|
|
}
|
||
|
|
memoryID, _ := payload["memory_id"].(string)
|
||
|
|
sourceID := payloadInt64(payload["source_id"])
|
||
|
|
msgDict, _ := payload["message_dict"].(map[string]any)
|
||
|
|
msg := MemoryMessage{
|
||
|
|
UserID: payloadString(msgDict["user_id"]),
|
||
|
|
AgentID: payloadString(msgDict["agent_id"]),
|
||
|
|
SessionID: payloadString(msgDict["session_id"]),
|
||
|
|
UserInput: payloadString(msgDict["user_input"]),
|
||
|
|
AgentResponse: payloadString(msgDict["agent_response"]),
|
||
|
|
}
|
||
|
|
|
||
|
|
task, err := s.taskDAO.GetByID(ctx, dao.DB, taskID)
|
||
|
|
if err != nil {
|
||
|
|
return fmt.Errorf("memory: task %s is not found", taskID)
|
||
|
|
}
|
||
|
|
if task.Progress == -1 {
|
||
|
|
return fmt.Errorf("memory: task %s is already failed", taskID)
|
||
|
|
}
|
||
|
|
|
||
|
|
if err := s.saveExtractedToMemory(ctx, memoryID, msg, sourceID, taskID); err != nil {
|
||
|
|
s.updateTaskProgress(taskID, -1, err.Error())
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
return nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// saveExtractedToMemory mirrors Python save_extracted_to_memory_only:
|
||
|
|
// skip raw-only memories, run LLM extraction, embed and persist the
|
||
|
|
// extracted messages under the raw message's id.
|
||
|
|
func (s *MemoryMessageService) saveExtractedToMemory(ctx context.Context, memoryID string, msg MemoryMessage, sourceID int64, taskID string) error {
|
||
|
|
mem, err := s.memories.getMemoryConfig(ctx, memoryID)
|
||
|
|
if err != nil {
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
memoryTypes := mem.MemoryType
|
||
|
|
if len(memoryTypes) == 0 {
|
||
|
|
memoryTypes = dao.GetMemoryTypeHuman(mem.Memory.MemoryType)
|
||
|
|
}
|
||
|
|
extractTypes := getTypesToExtract(memoryTypes)
|
||
|
|
if len(extractTypes) == 0 {
|
||
|
|
s.updateTaskProgress(taskID, 1.0, fmt.Sprintf("Memory '%s' don't need to extract.", memoryID))
|
||
|
|
return nil
|
||
|
|
}
|
||
|
|
|
||
|
|
extracted, err := s.extractByLLM(ctx, mem, extractTypes, msg, taskID)
|
||
|
|
if err != nil {
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
if len(extracted) == 0 {
|
||
|
|
s.updateTaskProgress(taskID, 1.0, "No memory extracted from raw message.")
|
||
|
|
return nil
|
||
|
|
}
|
||
|
|
s.updateTaskProgress(taskID, 0.5, fmt.Sprintf("Extracted %d messages from raw dialogue.", len(extracted)))
|
||
|
|
|
||
|
|
now := time.Now().UTC()
|
||
|
|
messages := make([]map[string]any, 0, len(extracted))
|
||
|
|
for _, item := range extracted {
|
||
|
|
messages = append(messages, buildExtractedMessage(generateRawMessageID(), sourceID, memoryID, msg, item, now))
|
||
|
|
}
|
||
|
|
if err := s.embedAndSaveMessages(ctx, mem, messages); err != nil {
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
s.updateTaskProgress(taskID, 1.0, "Message saved successfully.")
|
||
|
|
return nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// extractByLLM mirrors Python extract_by_llm: build the system/user
|
||
|
|
// prompts from the memory config, chat with the configured model, and
|
||
|
|
// parse the JSON result into per-type extracted items.
|
||
|
|
func (s *MemoryMessageService) extractByLLM(ctx context.Context, mem *CreateMemoryResponse, extractTypes []string, msg MemoryMessage, taskID string) ([]extractedMemory, error) {
|
||
|
|
systemPrompt := ""
|
||
|
|
if mem.SystemPrompt != nil {
|
||
|
|
systemPrompt = *mem.SystemPrompt
|
||
|
|
}
|
||
|
|
if strings.TrimSpace(systemPrompt) == "" {
|
||
|
|
systemPrompt = PromptAssembler{}.AssembleSystemPrompt(extractTypes)
|
||
|
|
}
|
||
|
|
|
||
|
|
conversation := fmt.Sprintf("User Input: %s\nAgent Response: %s", msg.UserInput, msg.AgentResponse)
|
||
|
|
now := time.Now().UTC().Format(memoryTimeLayout)
|
||
|
|
messages := []models.Message{{Role: "system", Content: systemPrompt}}
|
||
|
|
if mem.UserPrompt != nil && strings.TrimSpace(*mem.UserPrompt) != "" {
|
||
|
|
messages = append(messages,
|
||
|
|
models.Message{Role: "user", Content: *mem.UserPrompt},
|
||
|
|
models.Message{Role: "user", Content: fmt.Sprintf("Conversation: %s\nConversation Time: %s\nCurrent Time: %s", conversation, now, now)},
|
||
|
|
)
|
||
|
|
} else {
|
||
|
|
messages = append(messages, models.Message{Role: "user", Content: PromptAssembler{}.AssembleUserPrompt(conversation, now, now)})
|
||
|
|
}
|
||
|
|
|
||
|
|
// Python prefers tenant_llm_id and falls back to llm_id;
|
||
|
|
// ResolveModelConfig accepts both tenant-model ids and model names.
|
||
|
|
llmRef := mem.LLMID
|
||
|
|
if mem.TenantLLMID != nil && *mem.TenantLLMID != "" {
|
||
|
|
llmRef = *mem.TenantLLMID
|
||
|
|
}
|
||
|
|
driver, modelName, apiConfig, _, err := NewModelProviderService().ResolveModelConfig(ctx, mem.TenantID, entity.ModelTypeChat, llmRef)
|
||
|
|
if err != nil {
|
||
|
|
return nil, fmt.Errorf("resolve chat model: %w", err)
|
||
|
|
}
|
||
|
|
chatModel := models.NewChatModel(driver, &modelName, apiConfig)
|
||
|
|
|
||
|
|
s.updateTaskProgress(taskID, 0.15, "Prepared prompts and LLM.")
|
||
|
|
temperature := mem.Temperature
|
||
|
|
resp, err := chatModel.ModelDriver.ChatWithMessages(ctx, modelName, messages, apiConfig, &models.ChatConfig{Temperature: &temperature}, nil)
|
||
|
|
if err != nil {
|
||
|
|
return nil, fmt.Errorf("chat model: %w", err)
|
||
|
|
}
|
||
|
|
if resp == nil || resp.Answer == nil {
|
||
|
|
return nil, errors.New("empty response from chat model")
|
||
|
|
}
|
||
|
|
s.updateTaskProgress(taskID, 0.35, "Get extracted result from LLM.")
|
||
|
|
|
||
|
|
return parseMemoryExtraction(*resp.Answer, extractTypes), nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// buildExtractedMessage builds the persisted envelope for one extracted
|
||
|
|
// memory item. Field set matches buildRawMessage except message_type and
|
||
|
|
// source_id, which listMemoryMessages uses to aggregate extracts. Only
|
||
|
|
// logical message fields are set here; the doc engine maps them to
|
||
|
|
// storage fields (including tokenization) at insert time.
|
||
|
|
func buildExtractedMessage(messageID, sourceID int64, memoryID string, msg MemoryMessage, item extractedMemory, now time.Time) map[string]any {
|
||
|
|
var invalidAt any
|
||
|
|
if strings.TrimSpace(item.InvalidAt) != "" {
|
||
|
|
invalidAt = formatMemoryTime(item.InvalidAt, now)
|
||
|
|
}
|
||
|
|
return map[string]any{
|
||
|
|
"message_id": messageID,
|
||
|
|
"message_type": item.MessageType,
|
||
|
|
"source_id": sourceID,
|
||
|
|
"memory_id": memoryID,
|
||
|
|
"user_id": msg.UserID,
|
||
|
|
"agent_id": msg.AgentID,
|
||
|
|
"session_id": msg.SessionID,
|
||
|
|
"content": item.Content,
|
||
|
|
"valid_at": formatMemoryTime(item.ValidAt, now),
|
||
|
|
"invalid_at": invalidAt,
|
||
|
|
"forget_at": nil,
|
||
|
|
"status": true,
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// parseMemoryExtraction ports memory.utils.msg_util.get_json_result_from_llm_response
|
||
|
|
// plus the per-type flattening in extract_by_llm. Only the configured
|
||
|
|
// extract types are collected; unparseable responses yield an empty list.
|
||
|
|
func parseMemoryExtraction(answer string, extractTypes []string) []extractedMemory {
|
||
|
|
clean := strings.TrimSpace(answer)
|
||
|
|
clean = strings.TrimPrefix(clean, "```json")
|
||
|
|
clean = strings.TrimPrefix(clean, "```")
|
||
|
|
clean = strings.TrimSuffix(clean, "```")
|
||
|
|
if start := strings.Index(clean, "{"); start >= 0 {
|
||
|
|
if end := strings.LastIndex(clean, "}"); end > start {
|
||
|
|
clean = clean[start : end+1]
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
var parsed map[string][]struct {
|
||
|
|
Content string `json:"content"`
|
||
|
|
ValidAt string `json:"valid_at"`
|
||
|
|
InvalidAt string `json:"invalid_at"`
|
||
|
|
}
|
||
|
|
if err := json.Unmarshal([]byte(strings.TrimSpace(clean)), &parsed); err != nil {
|
||
|
|
common.Warn("memory: failed to parse LLM extraction result", zap.Error(err))
|
||
|
|
return nil
|
||
|
|
}
|
||
|
|
|
||
|
|
var out []extractedMemory
|
||
|
|
for _, memoryType := range extractTypes {
|
||
|
|
for _, item := range parsed[memoryType] {
|
||
|
|
if strings.TrimSpace(item.Content) == "" {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
out = append(out, extractedMemory{
|
||
|
|
MessageType: memoryType,
|
||
|
|
Content: item.Content,
|
||
|
|
ValidAt: item.ValidAt,
|
||
|
|
InvalidAt: item.InvalidAt,
|
||
|
|
})
|
||
|
|
}
|
||
|
|
}
|
||
|
|
return out
|
||
|
|
}
|
||
|
|
|
||
|
|
// formatMemoryTime normalizes an LLM-supplied timestamp (ISO 8601 or
|
||
|
|
// already-formatted) into memoryTimeLayout. Unparseable or empty input
|
||
|
|
// falls back to the supplied time.
|
||
|
|
func formatMemoryTime(value string, fallback time.Time) string {
|
||
|
|
v := strings.TrimSpace(value)
|
||
|
|
if v != "" {
|
||
|
|
for _, layout := range []string{time.RFC3339, "2006-01-02T15:04:05", "2006-01-02 15:04:05", "2006-01-02"} {
|
||
|
|
if t, err := time.Parse(layout, v); err == nil {
|
||
|
|
return t.UTC().Format(memoryTimeLayout)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
return fallback.UTC().Format(memoryTimeLayout)
|
||
|
|
}
|
||
|
|
|
||
|
|
// updateTaskProgress stamps and persists task progress, mirroring
|
||
|
|
// Python TaskService.update_progress call sites. Failures are logged
|
||
|
|
// and swallowed so progress reporting never breaks extraction.
|
||
|
|
func (s *MemoryMessageService) updateTaskProgress(taskID string, progress float64, msg string) {
|
||
|
|
if s == nil || s.taskDAO == nil || taskID == "" {
|
||
|
|
return
|
||
|
|
}
|
||
|
|
stamped := time.Now().Format(memoryTimeLayout) + " " + msg
|
||
|
|
if err := s.taskDAO.UpdateProgress(taskID, progress, stamped); err != nil {
|
||
|
|
common.Warn("memory: update task progress failed", zap.Error(err))
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func payloadInt64(v any) int64 {
|
||
|
|
switch n := v.(type) {
|
||
|
|
case float64:
|
||
|
|
return int64(n)
|
||
|
|
case int64:
|
||
|
|
return n
|
||
|
|
case int:
|
||
|
|
return int64(n)
|
||
|
|
case json.Number:
|
||
|
|
i, _ := n.Int64()
|
||
|
|
return i
|
||
|
|
case string:
|
||
|
|
var i int64
|
||
|
|
fmt.Sscanf(n, "%d", &i)
|
||
|
|
return i
|
||
|
|
}
|
||
|
|
return 0
|
||
|
|
}
|
||
|
|
|
||
|
|
func payloadString(v any) string {
|
||
|
|
s, _ := v.(string)
|
||
|
|
return s
|
||
|
|
}
|