mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-08-01 13:33:48 +08:00
@@ -482,6 +482,7 @@ func (s *MemoryService) UpdateMemory(ctx context.Context, tenantID string, memor
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("memory '%s' not found", memoryID)
|
||||
}
|
||||
ownerTenantID := currentMemory.TenantID
|
||||
|
||||
if req.Name != nil {
|
||||
memoryName := strings.TrimSpace(*req.Name)
|
||||
@@ -492,7 +493,7 @@ func (s *MemoryService) UpdateMemory(ctx context.Context, tenantID string, memor
|
||||
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)
|
||||
}, memoryName, ownerTenantID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -501,7 +502,10 @@ func (s *MemoryService) UpdateMemory(ctx context.Context, tenantID string, memor
|
||||
}
|
||||
|
||||
if req.Permissions != nil {
|
||||
perm := TenantPermission(strings.ToLower(*req.Permissions))
|
||||
perm := TenantPermission(strings.ToLower(strings.TrimSpace(*req.Permissions)))
|
||||
if currentMemory.TenantID != tenantID && strings.ToLower(strings.TrimSpace(currentMemory.Permissions)) != string(perm) {
|
||||
return nil, fmt.Errorf("tenant '%s' is not allowed to modify the memory's permission", tenantID)
|
||||
}
|
||||
if !validPermissions[perm] {
|
||||
return nil, fmt.Errorf("unknown permission '%s'", *req.Permissions)
|
||||
}
|
||||
@@ -518,9 +522,9 @@ func (s *MemoryService) UpdateMemory(ctx context.Context, tenantID string, memor
|
||||
if req.LLMID != nil {
|
||||
updateDict["llm_id"] = *req.LLMID
|
||||
if req.TenantLLMID == nil && *req.LLMID != "" {
|
||||
resolved, err := modelProvider.ResolveModelID(ctx, tenantID, entity.ModelTypeChat, *req.LLMID)
|
||||
resolved, err := modelProvider.ResolveModelID(ctx, ownerTenantID, entity.ModelTypeChat, *req.LLMID)
|
||||
if err != nil {
|
||||
slog.Warn("UpdateMemory: failed to resolve tenant LLM id", "tenant_id", tenantID, "llm_id", *req.LLMID, "err", err)
|
||||
slog.Warn("UpdateMemory: failed to resolve tenant LLM id", "tenant_id", ownerTenantID, "llm_id", *req.LLMID, "err", err)
|
||||
} else if resolved != "" {
|
||||
updateDict["tenant_llm_id"] = resolved
|
||||
}
|
||||
@@ -530,9 +534,9 @@ func (s *MemoryService) UpdateMemory(ctx context.Context, tenantID string, memor
|
||||
if req.EmbdID != nil {
|
||||
updateDict["embd_id"] = *req.EmbdID
|
||||
if req.TenantEmbdID == nil && *req.EmbdID != "" {
|
||||
resolved, err := modelProvider.ResolveModelID(ctx, tenantID, entity.ModelTypeEmbedding, *req.EmbdID)
|
||||
resolved, err := modelProvider.ResolveModelID(ctx, ownerTenantID, entity.ModelTypeEmbedding, *req.EmbdID)
|
||||
if err != nil {
|
||||
slog.Warn("UpdateMemory: failed to resolve tenant embedding id", "tenant_id", tenantID, "embd_id", *req.EmbdID, "err", err)
|
||||
slog.Warn("UpdateMemory: failed to resolve tenant embedding id", "tenant_id", ownerTenantID, "embd_id", *req.EmbdID, "err", err)
|
||||
} else if resolved != "" {
|
||||
updateDict["tenant_embd_id"] = resolved
|
||||
}
|
||||
|
||||
@@ -94,6 +94,132 @@ func setupMemoryMessageTestDB(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestUpdateMemoryTeamMemberCannotChangePermissions(t *testing.T) {
|
||||
setupMemoryMessageTestDB(t)
|
||||
|
||||
status := "1"
|
||||
if err := dao.DB.Create(&entity.Memory{
|
||||
ID: "mem-team",
|
||||
Name: "Shared memory",
|
||||
TenantID: "owner-1",
|
||||
MemoryType: dao.MemoryTypeRaw,
|
||||
StorageType: "table",
|
||||
EmbdID: "embd",
|
||||
LLMID: "llm",
|
||||
Permissions: string(entity.TenantPermissionTeam),
|
||||
MemorySize: MemorySizeLimit,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("seed memory: %v", err)
|
||||
}
|
||||
if err := dao.DB.Create(&entity.UserTenant{
|
||||
ID: "ut-team",
|
||||
UserID: "member-1",
|
||||
TenantID: "owner-1",
|
||||
Role: "normal",
|
||||
InvitedBy: "owner-1",
|
||||
Status: &status,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("seed user tenant: %v", err)
|
||||
}
|
||||
|
||||
svc := NewMemoryService()
|
||||
samePermission := " TEAM "
|
||||
if _, err := svc.UpdateMemory(context.Background(), "member-1", "mem-team", &UpdateMemoryRequest{
|
||||
Description: sptr("member edit"),
|
||||
Permissions: &samePermission,
|
||||
}); err != nil {
|
||||
t.Fatalf("UpdateMemory same permission error = %v", err)
|
||||
}
|
||||
|
||||
nextPermission := "me"
|
||||
if _, err := svc.UpdateMemory(context.Background(), "member-1", "mem-team", &UpdateMemoryRequest{
|
||||
Permissions: &nextPermission,
|
||||
}); err == nil {
|
||||
t.Fatal("UpdateMemory permission change error = nil, want error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateMemoryTeamMemberResolvesModelsAgainstOwnerTenant(t *testing.T) {
|
||||
setupMemoryMessageTestDB(t)
|
||||
|
||||
status := "1"
|
||||
if err := dao.DB.Create(&entity.Memory{
|
||||
ID: "mem-model",
|
||||
Name: "Shared model memory",
|
||||
TenantID: "owner-1",
|
||||
MemoryType: dao.MemoryTypeRaw,
|
||||
StorageType: "table",
|
||||
EmbdID: "old-embd",
|
||||
LLMID: "old-llm",
|
||||
Permissions: string(entity.TenantPermissionTeam),
|
||||
MemorySize: MemorySizeLimit,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("seed memory: %v", err)
|
||||
}
|
||||
if err := dao.DB.Create(&entity.UserTenant{
|
||||
ID: "ut-model",
|
||||
UserID: "member-1",
|
||||
TenantID: "owner-1",
|
||||
Role: "normal",
|
||||
InvitedBy: "owner-1",
|
||||
Status: &status,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("seed user tenant: %v", err)
|
||||
}
|
||||
|
||||
for _, row := range []struct {
|
||||
providerID string
|
||||
tenantID string
|
||||
modelID string
|
||||
}{
|
||||
{providerID: "provider-owner", tenantID: "owner-1", modelID: "tenant-llm-owner"},
|
||||
{providerID: "provider-member", tenantID: "member-1", modelID: "tenant-llm-member"},
|
||||
} {
|
||||
if err := dao.DB.Create(&entity.TenantModelProvider{
|
||||
ID: row.providerID,
|
||||
ProviderName: "OpenAI",
|
||||
TenantID: row.tenantID,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("seed provider %s: %v", row.providerID, err)
|
||||
}
|
||||
instanceID := row.providerID + "-default"
|
||||
if err := dao.DB.Create(&entity.TenantModelInstance{
|
||||
ID: instanceID,
|
||||
InstanceName: "default",
|
||||
ProviderID: row.providerID,
|
||||
APIKey: "test-key",
|
||||
Status: "active",
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("seed instance %s: %v", instanceID, err)
|
||||
}
|
||||
if err := dao.DB.Create(&entity.TenantModel{
|
||||
ID: row.modelID,
|
||||
ModelName: "gpt-4o",
|
||||
ProviderID: row.providerID,
|
||||
InstanceID: instanceID,
|
||||
ModelType: int(entity.ModelTypeChat),
|
||||
Status: "active",
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("seed model %s: %v", row.modelID, err)
|
||||
}
|
||||
}
|
||||
|
||||
llmID := "gpt-4o@default@OpenAI"
|
||||
if _, err := NewMemoryService().UpdateMemory(context.Background(), "member-1", "mem-model", &UpdateMemoryRequest{
|
||||
LLMID: &llmID,
|
||||
}); err != nil {
|
||||
t.Fatalf("UpdateMemory model error = %v", err)
|
||||
}
|
||||
|
||||
updated, err := dao.NewMemoryDAO().GetByID(context.Background(), dao.DB, "mem-model")
|
||||
if err != nil {
|
||||
t.Fatalf("get updated memory: %v", err)
|
||||
}
|
||||
if updated.TenantLLMID == nil || *updated.TenantLLMID != "tenant-llm-owner" {
|
||||
t.Fatalf("tenant_llm_id = %v, want tenant-llm-owner", updated.TenantLLMID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListMemoriesUsesTenantModelIDForDisplayName(t *testing.T) {
|
||||
setupMemoryMessageTestDB(t)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user