fix: failed to set right status in memory (#17472)

As title

Co-authored-by: coderabbitai[bot] <136622811+coderabbitai[bot]@users.noreply.github.com>
This commit is contained in:
Haruko386
2026-07-28 19:13:22 +08:00
committed by GitHub
parent 7f21a7ba18
commit 4885dda32a
8 changed files with 161 additions and 33 deletions

View File

@@ -1492,7 +1492,12 @@ func (s *AgentService) buildRunFunc(canvasID string, versionRow *entity.UserCanv
if uid, ok := root["user_id"].(string); ok && uid != "" {
state.Sys["user_id"] = uid
}
state.Sys["agent_id"] = canvasID
if canvasID != "" {
state.Sys["canvas_id"] = canvasID
}
if sessionID != "" {
state.Sys["session_id"] = sessionID
}
if tid, ok := root["tenant_id"].(string); ok && tid != "" {
state.Sys["tenant_id"] = tid
}

View File

@@ -886,6 +886,18 @@ func (s *MemoryService) AddMessage(ctx context.Context, currentUserID string, me
return true, "All add to task.", nil
}
func (s *MemoryService) saveAgentMessage(ctx context.Context, memoryIDs []string, msg MemoryMessage) (bool, string, error) {
res, err := NewMemoryMessageService(s).QueueSaveToMemoryTask(ctx, splitFilterValues(memoryIDs), msg)
if err != nil {
return false, err.Error(), err
}
errorMsg := memorySaveErrorMessage(res)
if errorMsg != "" {
return false, errorMsg, nil
}
return true, "All add to task.", nil
}
func missingRequestedMemoryIDs(requestedMemoryIDs, accessibleMemoryIDs []string) []string {
if len(requestedMemoryIDs) == 0 {
return []string{}

View File

@@ -136,11 +136,7 @@ func NewMemoryMessageService(memories *MemoryService) *MemoryMessageService {
// the per-memory outcomes. The outer error is reserved for
// call-level failures (e.g. invalid input); per-memory failures
// go into Failed, mirroring the Python tuple shape.
func (s *MemoryMessageService) QueueSaveToMemoryTask(
ctx context.Context,
memoryIDs []string,
msg MemoryMessage,
) (*QueueSaveResult, error) {
func (s *MemoryMessageService) QueueSaveToMemoryTask(ctx context.Context, memoryIDs []string, msg MemoryMessage) (*QueueSaveResult, error) {
if len(memoryIDs) == 0 {
return &QueueSaveResult{}, nil
}
@@ -206,12 +202,7 @@ func generateRawMessageID() int64 {
// buildRawMessage constructs the raw_message envelope that gets
// passed to embed_and_save (and persisted in the message table
// for the async extractor to read).
func buildRawMessage(
rawMessageID int64,
memoryID string,
mem *CreateMemoryResponse, // from MemoryService.GetMemoryConfig
msg MemoryMessage,
) map[string]any {
func buildRawMessage(rawMessageID int64, memoryID string, mem *CreateMemoryResponse, msg MemoryMessage) map[string]any {
content := fmt.Sprintf("User Input: %s\nAgent Response: %s",
msg.UserInput, msg.AgentResponse)
out := map[string]any{
@@ -244,12 +235,13 @@ func buildRawMessage(
// buildTaskRow constructs the Task row the async extractor polls.
func buildTaskRow(rawMessageID int64, memoryID string) map[string]any {
return map[string]any{
"id": newUUIDString(),
"doc_id": memoryID,
"task_type": "memory",
"progress": 0.0,
"begin_at": time.Now(),
"digest": fmt.Sprintf("%d", rawMessageID),
"id": newUUIDString(),
"doc_id": memoryID,
"task_type": "memory",
"progress": 0.0,
"progress_msg": "",
"begin_at": time.Now(),
"digest": fmt.Sprintf("%d", rawMessageID),
}
}
@@ -321,18 +313,20 @@ func newUUIDString() string {
func taskFromRow(row map[string]any) *entity.Task {
digest := fmt.Sprint(row["digest"])
progressMsg := fmt.Sprint(row["progress_msg"])
beginAt, _ := row["begin_at"].(time.Time)
if beginAt.IsZero() {
now := time.Now()
beginAt = now
}
return &entity.Task{
ID: fmt.Sprint(row["id"]),
DocID: fmt.Sprint(row["doc_id"]),
TaskType: fmt.Sprint(row["task_type"]),
Progress: 0,
BeginAt: &beginAt,
Digest: &digest,
ID: fmt.Sprint(row["id"]),
DocID: fmt.Sprint(row["doc_id"]),
TaskType: fmt.Sprint(row["task_type"]),
Progress: 0,
ProgressMsg: &progressMsg,
BeginAt: &beginAt,
Digest: &digest,
}
}

View File

@@ -139,6 +139,9 @@ func TestBuildTaskRow_Shape(t *testing.T) {
if row["progress"] != 0.0 {
t.Errorf("progress = %v, want 0.0", row["progress"])
}
if row["progress_msg"] != "" {
t.Errorf("progress_msg = %v, want empty string", row["progress_msg"])
}
if row["digest"] != "99" {
t.Errorf("digest = %v, want \"99\"", row["digest"])
}
@@ -147,6 +150,16 @@ func TestBuildTaskRow_Shape(t *testing.T) {
}
}
func TestTaskFromRow_InitializesProgressMessage(t *testing.T) {
task := taskFromRow(buildTaskRow(99, "mem-1"))
if task.ProgressMsg == nil {
t.Fatal("ProgressMsg is nil")
}
if *task.ProgressMsg != "" {
t.Fatalf("ProgressMsg = %q, want empty string", *task.ProgressMsg)
}
}
// TestEmbedAndSave_LoudFails: until the embedder port lands,
// the call must return ErrEmbedderNotWired — not panic, not
// silently succeed. A successful embed_and_save would mask the

View File

@@ -5,6 +5,7 @@ import (
"errors"
"fmt"
"reflect"
"strings"
"testing"
"github.com/glebarez/sqlite"
@@ -241,6 +242,54 @@ func seedMemoryMessages(t *testing.T) {
}
}
func TestSaveAgentMessageBypassesRequestAccessFilter(t *testing.T) {
setupMemoryMessageTestDB(t)
if err := dao.DB.Create(&entity.Memory{
ID: "mem-owned",
Name: "Owned",
TenantID: "user-1",
MemoryType: dao.MemoryTypeRaw,
StorageType: "table",
EmbdID: "embd-1",
LLMID: "llm-1",
Permissions: string(TenantPermissionMe),
ForgettingPolicy: string(ForgettingPolicyFIFO),
}).Error; err != nil {
t.Fatalf("seed memory: %v", err)
}
svc := &MemoryService{memoryDAO: dao.NewMemoryDAO()}
msg := MemoryMessage{
AgentID: "agent-1",
SessionID: "session-1",
UserInput: "hi",
AgentResponse: "hello",
}
ok, detail, err := svc.AddMessage(context.Background(), "", []string{"mem-owned"}, msg)
if err != nil {
t.Fatalf("AddMessage: %v", err)
}
if ok || detail != "Memory not found." {
t.Fatalf("AddMessage with empty current user = (%v, %q), want permission-filtered not found", ok, detail)
}
ok, detail, err = svc.saveAgentMessage(context.Background(), []string{"mem-owned"}, msg)
if err != nil {
t.Fatalf("saveAgentMessage: %v", err)
}
if ok {
t.Fatal("saveAgentMessage unexpectedly succeeded without a message store")
}
if strings.Contains(detail, "Memory not found") {
t.Fatalf("saveAgentMessage was filtered by request user: %q", detail)
}
if !strings.Contains(detail, "message store is not initialized") {
t.Fatalf("saveAgentMessage detail = %q, want message-store failure after memory lookup", detail)
}
}
func TestGetMessagesFiltersAccessibleMemoryAndBuildsRecentSearch(t *testing.T) {
setupMemoryMessageTestDB(t)
seedMemoryMessages(t)

View File

@@ -23,8 +23,8 @@ import (
"ragflow/internal/agent/component"
)
// memorySaverAdapter bridges the component.MemorySaver interface to
// MemoryService.AddMessage. It is installed via component.SetMemorySaver
// memorySaverAdapter bridges the component.MemorySaver interface to the
// agent-internal memory queue. It is installed via component.SetMemorySaver
// at boot time so that the Message component can persist conversation
// turns to memory stores declared in the canvas DSL.
type memorySaverAdapter struct {
@@ -37,10 +37,9 @@ func NewMemorySaverAdapter(svc *MemoryService) component.MemorySaver {
return &memorySaverAdapter{svc: svc}
}
// Save implements component.MemorySaver. It delegates to
// MemoryService.AddMessage which handles access filtering, message
// construction, embedding, and async task queueing — the same pipeline
// used by the REST API add_message endpoint.
// Save implements component.MemorySaver. Agent message components already run
// inside an authorized canvas, so this path queues the selected memories
// directly instead of applying REST request access filtering.
func (a *memorySaverAdapter) Save(ctx context.Context, req component.MemorySaveRequest) error {
if a == nil || a.svc == nil {
return fmt.Errorf("memory: saver adapter not initialised")
@@ -52,7 +51,7 @@ func (a *memorySaverAdapter) Save(ctx context.Context, req component.MemorySaveR
UserInput: req.UserInput,
AgentResponse: req.AgentResponse,
}
ok, detail, err := a.svc.AddMessage(ctx, req.UserID, req.MemoryIDs, msg)
ok, detail, err := a.svc.saveAgentMessage(ctx, req.MemoryIDs, msg)
if err != nil {
return err
}