diff --git a/internal/dao/memory.go b/internal/dao/memory.go index a69be44a0b..cad03e4bcb 100644 --- a/internal/dao/memory.go +++ b/internal/dao/memory.go @@ -24,6 +24,8 @@ import ( "fmt" "ragflow/internal/entity" "strings" + + "gorm.io/gorm" ) // Memory type bit flag constants, consistent with Python MemoryType enum @@ -111,8 +113,8 @@ func NewMemoryDAO() *MemoryDAO { // // Returns: // - error: Database operation error -func (dao *MemoryDAO) Create(memory *entity.Memory) error { - return DB.Create(memory).Error +func (dao *MemoryDAO) Create(ctx context.Context, db *gorm.DB, memory *entity.Memory) error { + return db.WithContext(ctx).Create(memory).Error } // GetByID retrieves a memory record by ID from database @@ -123,14 +125,14 @@ func (dao *MemoryDAO) Create(memory *entity.Memory) error { // Returns: // - *model.Memory: Memory model pointer // - error: Database operation error -func (dao *MemoryDAO) GetByID(id string) (*entity.Memory, error) { - return dao.GetByIDWithContext(context.Background(), id) +func (dao *MemoryDAO) GetByID(ctx context.Context, db *gorm.DB, id string) (*entity.Memory, error) { + return dao.GetByIDWithContext(ctx, db, id) } // GetByIDWithContext retrieves a memory record by ID from database with context. -func (dao *MemoryDAO) GetByIDWithContext(ctx context.Context, id string) (*entity.Memory, error) { +func (dao *MemoryDAO) GetByIDWithContext(ctx context.Context, db *gorm.DB, id string) (*entity.Memory, error) { var memory entity.Memory - err := DB.WithContext(ctx).Where("id = ?", id).First(&memory).Error + err := db.WithContext(ctx).WithContext(ctx).Where("id = ?", id).First(&memory).Error if err != nil { return nil, err } @@ -145,9 +147,9 @@ func (dao *MemoryDAO) GetByIDWithContext(ctx context.Context, id string) (*entit // Returns: // - []*model.Memory: Memory model pointer array // - error: Database operation error -func (dao *MemoryDAO) GetByTenantID(tenantID string) ([]*entity.Memory, error) { +func (dao *MemoryDAO) GetByTenantID(ctx context.Context, db *gorm.DB, tenantID string) ([]*entity.Memory, error) { var memories []*entity.Memory - err := DB.Where("tenant_id = ?", tenantID).Find(&memories).Error + err := db.WithContext(ctx).Where("tenant_id = ?", tenantID).Find(&memories).Error return memories, err } @@ -161,9 +163,9 @@ func (dao *MemoryDAO) GetByTenantID(tenantID string) ([]*entity.Memory, error) { // Returns: // - []*model.Memory: Matching memory list (for existence check) // - error: Database operation error -func (dao *MemoryDAO) GetByNameAndTenant(name string, tenantID string) ([]*entity.Memory, error) { +func (dao *MemoryDAO) GetByNameAndTenant(ctx context.Context, db *gorm.DB, name string, tenantID string) ([]*entity.Memory, error) { var memories []*entity.Memory - err := DB.Where("name = ? AND tenant_id = ?", name, tenantID).Find(&memories).Error + err := db.WithContext(ctx).Where("name = ? AND tenant_id = ?", name, tenantID).Find(&memories).Error return memories, err } @@ -175,9 +177,9 @@ func (dao *MemoryDAO) GetByNameAndTenant(name string, tenantID string) ([]*entit // Returns: // - []*model.Memory: Memory model pointer array // - error: Database operation error -func (dao *MemoryDAO) GetByIDs(ids []string) ([]*entity.Memory, error) { +func (dao *MemoryDAO) GetByIDs(ctx context.Context, db *gorm.DB, ids []string) ([]*entity.Memory, error) { var memories []*entity.Memory - err := DB.Where("id IN ?", ids).Find(&memories).Error + err := db.WithContext(ctx).Where("id IN ?", ids).Find(&memories).Error return memories, err } @@ -202,7 +204,7 @@ func (dao *MemoryDAO) GetByIDs(ids []string) ([]*entity.Memory, error) { // // updates := map[string]interface{}{"name": "NewName", "memory_type": []string{"semantic"}} // err := dao.UpdateByID("memory123", updates) -func (dao *MemoryDAO) UpdateByID(id string, updates map[string]interface{}) error { +func (dao *MemoryDAO) UpdateByID(ctx context.Context, db *gorm.DB, id string, updates map[string]interface{}) error { if updates == nil || len(updates) == 0 { return nil } @@ -222,7 +224,7 @@ func (dao *MemoryDAO) UpdateByID(id string, updates map[string]interface{}) erro } } - return DB.Model(&entity.Memory{}).Where("id = ?", id).Updates(updates).Error + return db.WithContext(ctx).Model(&entity.Memory{}).Where("id = ?", id).Updates(updates).Error } // DeleteByID deletes a memory by ID @@ -236,8 +238,8 @@ func (dao *MemoryDAO) UpdateByID(id string, updates map[string]interface{}) erro // Example: // // err := dao.DeleteByID("memory123") -func (dao *MemoryDAO) DeleteByID(id string) error { - return DB.Where("id = ?", id).Delete(&entity.Memory{}).Error +func (dao *MemoryDAO) DeleteByID(ctx context.Context, db *gorm.DB, id string) error { + return db.WithContext(ctx).Where("id = ?", id).Delete(&entity.Memory{}).Error } // GetWithOwnerNameByID retrieves a memory with owner name by ID @@ -253,7 +255,7 @@ func (dao *MemoryDAO) DeleteByID(id string) error { // Example: // // memory, err := dao.GetWithOwnerNameByID("memory123") -func (dao *MemoryDAO) GetWithOwnerNameByID(id string) (*entity.MemoryListItem, error) { +func (dao *MemoryDAO) GetWithOwnerNameByID(ctx context.Context, db *gorm.DB, id string) (*entity.MemoryListItem, error) { querySQL := ` SELECT m.id, m.name, m.avatar, m.tenant_id, m.memory_type, m.storage_type, m.embd_id, m.tenant_embd_id, m.llm_id, m.tenant_llm_id, @@ -271,7 +273,7 @@ func (dao *MemoryDAO) GetWithOwnerNameByID(id string) (*entity.MemoryListItem, e OwnerName *string `gorm:"column:owner_name"` } - if err := DB.Raw(querySQL, id).Scan(&rawResult).Error; err != nil { + if err := db.WithContext(ctx).Raw(querySQL, id).Scan(&rawResult).Error; err != nil { return nil, err } @@ -301,7 +303,7 @@ func (dao *MemoryDAO) GetWithOwnerNameByID(id string) (*entity.MemoryListItem, e // Example: // // memories, total, err := dao.GetByFilter([]string{"tenant1"}, []string{"semantic"}, "table", "test", 1, 10) -func (dao *MemoryDAO) GetByFilter(userID string, tenantIDs []string, memoryTypes []string, storageType string, keywords string, page int, pageSize int) ([]*entity.MemoryListItem, int64, error) { +func (dao *MemoryDAO) GetByFilter(ctx context.Context, db *gorm.DB, userID string, tenantIDs []string, memoryTypes []string, storageType string, keywords string, page int, pageSize int) ([]*entity.MemoryListItem, int64, error) { var conditions []string var args []interface{} @@ -338,7 +340,7 @@ func (dao *MemoryDAO) GetByFilter(userID string, tenantIDs []string, memoryTypes countSQL := fmt.Sprintf("SELECT COUNT(*) FROM memory m %s", whereClause) var total int64 - if err := DB.Raw(countSQL, args...).Scan(&total).Error; err != nil { + if err := db.WithContext(ctx).Raw(countSQL, args...).Scan(&total).Error; err != nil { return nil, 0, err } @@ -364,7 +366,7 @@ func (dao *MemoryDAO) GetByFilter(userID string, tenantIDs []string, memoryTypes OwnerName *string `gorm:"column:owner_name"` } - if err := DB.Raw(querySQL, queryArgs...).Scan(&rawResults).Error; err != nil { + if err := db.WithContext(ctx).Raw(querySQL, queryArgs...).Scan(&rawResults).Error; err != nil { return nil, 0, err } @@ -380,8 +382,8 @@ func (dao *MemoryDAO) GetByFilter(userID string, tenantIDs []string, memoryTypes } // Accessible check if it is possible for user to access the memory -func (dao *MemoryDAO) Accessible(userID, memoryID string) (bool, error) { - memory, err := dao.GetByID(memoryID) +func (dao *MemoryDAO) Accessible(ctx context.Context, db *gorm.DB, userID, memoryID string) (bool, error) { + memory, err := dao.GetByID(ctx, db, memoryID) if err != nil { return false, err } @@ -395,7 +397,7 @@ func (dao *MemoryDAO) Accessible(userID, memoryID string) (bool, error) { } var count int64 - err = DB.Table("user_tenant"). + err = db.WithContext(ctx).Table("user_tenant"). Where("tenant_id = ? AND user_id = ? AND status = ?", memory.TenantID, userID, "1"). Count(&count).Error if err != nil { diff --git a/internal/handler/memory.go b/internal/handler/memory.go index 49bccc7c3b..6e54bec5a3 100644 --- a/internal/handler/memory.go +++ b/internal/handler/memory.go @@ -130,9 +130,10 @@ func (h *MemoryHandler) CreateMemory(c *gin.Context) { // Record request parsing completion time (for timing) tParsed := time.Now() + ctx := c.Request.Context() // Call service layer to create memory - result, err := h.memoryService.CreateMemory(userID, &req) + result, err := h.memoryService.CreateMemory(ctx, userID, &req) if err != nil { // Log error if timing is enabled if timingEnabled != "" { @@ -218,8 +219,10 @@ func (h *MemoryHandler) UpdateMemory(c *gin.Context) { return } + ctx := c.Request.Context() + // Call service layer to update memory - result, err := h.memoryService.UpdateMemory(userID, memoryID, &req) + result, err := h.memoryService.UpdateMemory(ctx, userID, memoryID, &req) if err != nil { errMsg := err.Error() // Check if it's a "not found" error @@ -261,9 +264,10 @@ func (h *MemoryHandler) DeleteMemory(c *gin.Context) { common.ResponseWithCodeData(c, common.CodeArgumentError, nil, "memory ID is required") return } + ctx := c.Request.Context() // Call service layer to delete memory - err := h.memoryService.DeleteMemory(memoryID) + err := h.memoryService.DeleteMemory(ctx, memoryID) if err != nil { errMsg := err.Error() // Check if it's a "not found" error @@ -351,8 +355,10 @@ func (h *MemoryHandler) ListMemories(c *gin.Context) { } } + ctx := c.Request.Context() + // Call service layer to get memory list - result, err := h.memoryService.ListMemories(user.ID, tenantIDs, memoryTypes, storageType, keywords, page, pageSize) + result, err := h.memoryService.ListMemories(ctx, user.ID, tenantIDs, memoryTypes, storageType, keywords, page, pageSize) if err != nil { common.ResponseWithCodeData(c, common.CodeServerError, nil, err.Error()) return @@ -380,9 +386,10 @@ func (h *MemoryHandler) GetMemoryConfig(c *gin.Context) { common.ResponseWithCodeData(c, common.CodeArgumentError, nil, "memory ID is required") return } + ctx := c.Request.Context() // Call service layer to get memory configuration - result, err := h.memoryService.GetMemoryConfig(memoryID) + result, err := h.memoryService.GetMemoryConfig(ctx, memoryID) if err != nil { errMsg := err.Error() // Check if it's a "not found" error diff --git a/internal/service/memory.go b/internal/service/memory.go index 1d8e67090e..cd11b1d3e8 100644 --- a/internal/service/memory.go +++ b/internal/service/memory.go @@ -343,7 +343,7 @@ type ListMemoryResponse struct { // // req := &CreateMemoryRequest{Name: "MyMemory", MemoryType: []string{"semantic"}, EmbdID: "embd1", LLMID: "llm1"} // resp, err := service.CreateMemory("tenant123", req) -func (s *MemoryService) CreateMemory(tenantID string, req *CreateMemoryRequest) (*CreateMemoryResponse, error) { +func (s *MemoryService) CreateMemory(ctx context.Context, tenantID string, req *CreateMemoryRequest) (*CreateMemoryResponse, error) { // Resolve tenant model IDs, mirroring Python's ensure_tenant_model_ids_for_params. // Resolution failure is non-fatal (e.g. Builtin models that have no // tenant_model row) — we leave the tenant_*_id fields nil and proceed. @@ -389,7 +389,7 @@ func (s *MemoryService) CreateMemory(tenantID string, req *CreateMemoryRequest) } memoryName, err := common.DuplicateName(func(name string, tid string) bool { - existing, _ := s.memoryDAO.GetByNameAndTenant(name, tid) + existing, _ := s.memoryDAO.GetByNameAndTenant(ctx, dao.DB, name, tid) return len(existing) > 0 }, memoryName, tenantID) if err != nil { @@ -423,11 +423,11 @@ func (s *MemoryService) CreateMemory(tenantID string, req *CreateMemoryRequest) if req.TenantLLMID != nil { memory.TenantLLMID = req.TenantLLMID } - if err := s.memoryDAO.Create(memory); err != nil { + if err = s.memoryDAO.Create(ctx, dao.DB, memory); err != nil { return nil, errors.New("could not create new memory") } - createdMemory, err := s.memoryDAO.GetByID(newID) + createdMemory, err := s.memoryDAO.GetByID(ctx, dao.DB, newID) if err != nil { return nil, errors.New("could not create new memory") } @@ -451,25 +451,25 @@ func (s *MemoryService) CreateMemory(tenantID string, req *CreateMemoryRequest) // // req := &UpdateMemoryRequest{Name: ptr("NewName"), MemorySize: ptr(int64(1000000))} // resp, err := service.UpdateMemory("tenant123", "memory456", req) -func (s *MemoryService) UpdateMemory(tenantID string, memoryID string, req *UpdateMemoryRequest) (*CreateMemoryResponse, error) { +func (s *MemoryService) UpdateMemory(ctx context.Context, tenantID string, memoryID string, req *UpdateMemoryRequest) (*CreateMemoryResponse, error) { updateDict := make(map[string]interface{}) - if ok, err := s.memoryDAO.Accessible(tenantID, memoryID); !ok || err != nil { + if ok, err := s.memoryDAO.Accessible(ctx, dao.DB, tenantID, memoryID); !ok || err != nil { return nil, err } - currentMemory, err := s.memoryDAO.GetByID(memoryID) + currentMemory, err := s.memoryDAO.GetByID(ctx, dao.DB, memoryID) if err != nil { return nil, fmt.Errorf("memory '%s' not found", memoryID) } if req.Name != nil { memoryName := strings.TrimSpace(*req.Name) - if err := common.ValidateName(memoryName); err != nil { + if err = common.ValidateName(memoryName); err != nil { return nil, err } if memoryName != strings.TrimSpace(currentMemory.Name) { - memoryName, err := common.DuplicateName(func(name string, tid string) bool { - existing, _ := s.memoryDAO.GetByNameAndTenant(name, tid) + memoryName, err = common.DuplicateName(func(name string, tid string) bool { + existing, _ := s.memoryDAO.GetByNameAndTenant(ctx, dao.DB, name, tid) return len(existing) > 0 }, memoryName, tenantID) if err != nil { @@ -700,7 +700,7 @@ func (s *MemoryService) UpdateMemory(tenantID string, memoryID string, req *Upda } } if len(notAllowedUpdate) > 0 { - messages, err := s.listMemoryMessages(context.Background(), currentMemory, []string{}, "", 1, 1) + messages, err := s.listMemoryMessages(ctx, currentMemory, []string{}, "", 1, 1) if err != nil { return nil, fmt.Errorf("failed to check memory messages: %w", err) } @@ -723,11 +723,11 @@ func (s *MemoryService) UpdateMemory(tenantID string, memoryID string, req *Upda } } - if err := s.memoryDAO.UpdateByID(memoryID, updateDict); err != nil { + if err = s.memoryDAO.UpdateByID(ctx, dao.DB, memoryID, updateDict); err != nil { return nil, errors.New("failed to update memory") } - updatedMemory, err := s.memoryDAO.GetByID(memoryID) + updatedMemory, err := s.memoryDAO.GetByID(ctx, dao.DB, memoryID) if err != nil { return nil, errors.New("failed to get updated memory") } @@ -786,8 +786,8 @@ func sameStringSet(a, b []string) bool { // Example: // // err := service.DeleteMemory("memory456") -func (s *MemoryService) DeleteMemory(memoryID string) error { - _, err := s.memoryDAO.GetByID(memoryID) +func (s *MemoryService) DeleteMemory(ctx context.Context, memoryID string) error { + _, err := s.memoryDAO.GetByID(ctx, dao.DB, memoryID) if err != nil { return fmt.Errorf("memory '%s' not found", memoryID) } @@ -800,7 +800,7 @@ func (s *MemoryService) DeleteMemory(memoryID string) error { // } // Delete memory record - if err := s.memoryDAO.DeleteByID(memoryID); err != nil { + if err = s.memoryDAO.DeleteByID(ctx, dao.DB, memoryID); err != nil { return errors.New("failed to delete memory") } @@ -832,7 +832,7 @@ func (s *MemoryService) ForgetMessage(ctx context.Context, userID string, memory } indexName := memoryIndexName(memory.TenantID) - if err := s.docEngine.UpdateChunks(ctx, condition, updates, indexName, memoryID); err != nil { + if err = s.docEngine.UpdateChunks(ctx, condition, updates, indexName, memoryID); err != nil { if isMessageDocumentNotFound(err) { // Match Python delete-by-query behavior: forgetting an already-missing // message document is idempotent and still considered successful. @@ -952,7 +952,7 @@ func (s *MemoryService) UpdateMessageStatus(ctx context.Context, userID, memoryI "id": messageDocID, } indexName := memoryIndexName(memory.TenantID) - if err := s.docEngine.UpdateChunks(ctx, condition, updates, indexName, memoryID); err != nil { + if err = s.docEngine.UpdateChunks(ctx, condition, updates, indexName, memoryID); err != nil { if isMessageDocumentNotFound(err) { return false, &ResourceNotFoundError{Resource: "Message", ID: messageDocID} } @@ -1096,7 +1096,7 @@ func (s *MemoryService) filterAccessibleMemories(ctx context.Context, userID str return []*entity.Memory{}, nil } - memories, err := s.memoryDAO.GetByIDs(memoryIDs) + memories, err := s.memoryDAO.GetByIDs(ctx, dao.DB, memoryIDs) if err != nil { return nil, err } @@ -1389,7 +1389,7 @@ func (s *MemoryService) requireMemoryAccess(ctx context.Context, userID string, if err := ctx.Err(); err != nil { return nil, err } - memory, err := s.memoryDAO.GetByIDWithContext(ctx, memoryID) + memory, err := s.memoryDAO.GetByIDWithContext(ctx, dao.DB, memoryID) if err != nil { if dao.IsNotFoundErr(err) { return nil, &ResourceNotFoundError{Resource: "Memory", ID: memoryID} @@ -1439,7 +1439,7 @@ func (s *MemoryService) requireMemoryAccess(ctx context.Context, userID string, // Example: // // resp, err := service.ListMemories("user123", []string{}, []string{"semantic"}, "table", "test", 1, 10) -func (s *MemoryService) ListMemories(userID string, tenantIDs []string, memoryTypes []string, storageType string, keywords string, page int, pageSize int) (*ListMemoryResponse, error) { +func (s *MemoryService) ListMemories(ctx context.Context, userID string, tenantIDs []string, memoryTypes []string, storageType string, keywords string, page int, pageSize int) (*ListMemoryResponse, error) { // If tenantIDs is empty, get all tenants associated with the user if len(tenantIDs) == 0 { userTenantService := NewUserTenantService() @@ -1454,7 +1454,7 @@ func (s *MemoryService) ListMemories(userID string, tenantIDs []string, memoryTy } } - memories, total, err := s.memoryDAO.GetByFilter(userID, tenantIDs, memoryTypes, storageType, keywords, page, pageSize) + memories, total, err := s.memoryDAO.GetByFilter(ctx, dao.DB, userID, tenantIDs, memoryTypes, storageType, keywords, page, pageSize) if err != nil { return nil, err } @@ -1539,8 +1539,8 @@ func resolveTenantModelDisplayName(tenantModelID, rawModelID string, cache map[s // Example: // // resp, err := service.GetMemoryConfig("memory456") -func (s *MemoryService) GetMemoryConfig(memoryID string) (*CreateMemoryResponse, error) { - memory, err := s.memoryDAO.GetWithOwnerNameByID(memoryID) +func (s *MemoryService) GetMemoryConfig(ctx context.Context, memoryID string) (*CreateMemoryResponse, error) { + memory, err := s.memoryDAO.GetWithOwnerNameByID(ctx, dao.DB, memoryID) if err != nil { return nil, fmt.Errorf("memory '%s' not found", memoryID) } diff --git a/internal/service/memory_message_service.go b/internal/service/memory_message_service.go index 78a727a997..f7c9548377 100644 --- a/internal/service/memory_message_service.go +++ b/internal/service/memory_message_service.go @@ -154,7 +154,7 @@ func (s *MemoryMessageService) QueueSaveToMemoryTask( res := &QueueSaveResult{} for _, memoryID := range memoryIDs { // (1) Look up the memory. - mem, err := s.memories.GetMemoryConfig(memoryID) + mem, err := s.memories.GetMemoryConfig(ctx, memoryID) if err != nil { res.NotFound = append(res.NotFound, memoryID) continue diff --git a/internal/service/memory_message_test.go b/internal/service/memory_message_test.go index 685c35dcff..058613fc18 100644 --- a/internal/service/memory_message_test.go +++ b/internal/service/memory_message_test.go @@ -153,7 +153,8 @@ func TestListMemoriesUsesTenantModelIDForDisplayName(t *testing.T) { t.Fatalf("seed memory: %v", err) } - resp, err := NewMemoryService().ListMemories("user-1", []string{"user-1"}, nil, "", "", 1, 10) + ctx := t.Context() + resp, err := NewMemoryService().ListMemories(ctx, "user-1", []string{"user-1"}, nil, "", "", 1, 10) if err != nil { t.Fatalf("ListMemories: %v", err) } @@ -189,7 +190,8 @@ func TestListMemoriesFallsBackToRawModelIDWithoutTenantModelID(t *testing.T) { t.Fatalf("seed memory: %v", err) } - resp, err := NewMemoryService().ListMemories("user-1", []string{"user-1"}, nil, "", "", 1, 10) + ctx := t.Context() + resp, err := NewMemoryService().ListMemories(ctx, "user-1", []string{"user-1"}, nil, "", "", 1, 10) if err != nil { t.Fatalf("ListMemories: %v", err) }