Go: context, part8 (#17405)

Signed-off-by: Jin Hai <haijin.chn@gmail.com>
This commit is contained in:
Jin Hai
2026-07-27 13:38:15 +08:00
committed by GitHub
parent c48eb70e67
commit f5aa5f7d94
5 changed files with 67 additions and 56 deletions

View File

@@ -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 {

View File

@@ -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

View File

@@ -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)
}

View File

@@ -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

View File

@@ -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)
}