diff --git a/internal/service/memory.go b/internal/service/memory.go index 9a5892e350..12b28d727e 100644 --- a/internal/service/memory.go +++ b/internal/service/memory.go @@ -1460,6 +1460,7 @@ func (s *MemoryService) ListMemories(userID string, tenantIDs []string, memoryTy } memoryList := make([]map[string]interface{}, 0, len(memories)) + modelNameCache := make(map[string]string) for _, m := range memories { resp := formatRetDataFromMemoryListItem(m) var createDateStr *string @@ -1468,6 +1469,8 @@ func (s *MemoryService) ListMemories(userID string, tenantIDs []string, memoryTy } memoryMap := map[string]interface{}{ "id": resp.ID, + "llm_id": resolveTenantModelDisplayName(ptrStringValue(resp.TenantLLMID), resp.LLMID, modelNameCache), + "embd_id": resolveTenantModelDisplayName(ptrStringValue(resp.TenantEmbdID), resp.EmbdID, modelNameCache), "name": resp.Name, "avatar": resp.Avatar, "tenant_id": resp.TenantID, @@ -1488,6 +1491,42 @@ func (s *MemoryService) ListMemories(userID string, tenantIDs []string, memoryTy }, nil } +// resolveTenantModelDisplayName turns a tenant_model ID into +// modelName@instance@provider. rawModelID is the API-facing fallback +// stored on memory.llm_id / memory.embd_id. +func resolveTenantModelDisplayName(tenantModelID, rawModelID string, cache map[string]string) string { + tenantModelID = strings.TrimSpace(tenantModelID) + rawModelID = strings.TrimSpace(rawModelID) + if tenantModelID == "" || strings.Contains(tenantModelID, "@") { + return rawModelID + } + + if displayName, ok := cache[tenantModelID]; ok { + return displayName + } + + displayName := rawModelID + defer func() { + cache[tenantModelID] = displayName + }() + + model, err := dao.NewTenantModelDAO().GetByID(tenantModelID) + if err != nil { + return displayName + } + instance, err := dao.NewTenantModelInstanceDAO().GetByID(model.InstanceID) + if err != nil { + return displayName + } + provider, err := dao.NewTenantModelProviderDAO().GetByID(model.ProviderID) + if err != nil { + return displayName + } + + displayName = fmt.Sprintf("%s@%s@%s", model.ModelName, instance.InstanceName, provider.ProviderName) + return displayName +} + // GetMemoryConfig retrieves the full configuration of a memory by ID // // Parameters: diff --git a/internal/service/memory_message_test.go b/internal/service/memory_message_test.go index 9a5b40c751..6ff427f3c5 100644 --- a/internal/service/memory_message_test.go +++ b/internal/service/memory_message_test.go @@ -75,7 +75,14 @@ func setupMemoryMessageTestDB(t *testing.T) { if err != nil { t.Fatalf("failed to open sqlite: %v", err) } - if err := db.AutoMigrate(&entity.Memory{}, &entity.UserTenant{}); err != nil { + if err := db.AutoMigrate( + &entity.Memory{}, + &entity.User{}, + &entity.UserTenant{}, + &entity.TenantModelProvider{}, + &entity.TenantModelInstance{}, + &entity.TenantModel{}, + ); err != nil { t.Fatalf("failed to migrate memory test tables: %v", err) } @@ -86,6 +93,118 @@ func setupMemoryMessageTestDB(t *testing.T) { }) } +func TestListMemoriesUsesTenantModelIDForDisplayName(t *testing.T) { + setupMemoryMessageTestDB(t) + + if err := dao.DB.Create(&entity.User{ID: "user-1", Nickname: "Owner"}).Error; err != nil { + t.Fatalf("seed user: %v", err) + } + if err := dao.DB.Create(&entity.TenantModelProvider{ + ID: "provider-1", + ProviderName: "OpenAI", + TenantID: "user-1", + }).Error; err != nil { + t.Fatalf("seed provider: %v", err) + } + if err := dao.DB.Create(&entity.TenantModelInstance{ + ID: "instance-1", + InstanceName: "default", + ProviderID: "provider-1", + APIKey: "test-key", + }).Error; err != nil { + t.Fatalf("seed instance: %v", err) + } + if err := dao.DB.Create(&entity.TenantModel{ + ID: "tenant-llm-1", + ModelName: "gpt-4o", + ProviderID: "provider-1", + InstanceID: "instance-1", + ModelType: int(entity.ModelTypeChat), + Status: "active", + }).Error; err != nil { + t.Fatalf("seed chat model: %v", err) + } + if err := dao.DB.Create(&entity.TenantModel{ + ID: "tenant-embd-1", + ModelName: "text-embedding-3-small", + ProviderID: "provider-1", + InstanceID: "instance-1", + ModelType: int(entity.ModelTypeEmbedding), + Status: "active", + }).Error; err != nil { + t.Fatalf("seed embedding model: %v", err) + } + + tenantLLMID := "tenant-llm-1" + tenantEmbdID := "tenant-embd-1" + if err := dao.DB.Create(&entity.Memory{ + ID: "mem-with-tenant-models", + Name: "With tenant models", + TenantID: "user-1", + MemoryType: dao.MemoryTypeRaw, + StorageType: "table", + LLMID: "gpt-4o@OpenAI", + TenantLLMID: &tenantLLMID, + EmbdID: "text-embedding-3-small@OpenAI", + TenantEmbdID: &tenantEmbdID, + Permissions: string(TenantPermissionMe), + ForgettingPolicy: string(ForgettingPolicyFIFO), + }).Error; err != nil { + t.Fatalf("seed memory: %v", err) + } + + resp, err := NewMemoryService().ListMemories("user-1", []string{"user-1"}, nil, "", "", 1, 10) + if err != nil { + t.Fatalf("ListMemories: %v", err) + } + if resp.TotalCount != 1 || len(resp.MemoryList) != 1 { + t.Fatalf("ListMemories returned total=%d len=%d, want 1", resp.TotalCount, len(resp.MemoryList)) + } + memory := resp.MemoryList[0] + if got, want := memory["llm_id"], "gpt-4o@default@OpenAI"; got != want { + t.Fatalf("llm_id = %v, want %v", got, want) + } + if got, want := memory["embd_id"], "text-embedding-3-small@default@OpenAI"; got != want { + t.Fatalf("embd_id = %v, want %v", got, want) + } +} + +func TestListMemoriesFallsBackToRawModelIDWithoutTenantModelID(t *testing.T) { + setupMemoryMessageTestDB(t) + + if err := dao.DB.Create(&entity.User{ID: "user-1", Nickname: "Owner"}).Error; err != nil { + t.Fatalf("seed user: %v", err) + } + if err := dao.DB.Create(&entity.Memory{ + ID: "mem-without-tenant-models", + Name: "Without tenant models", + TenantID: "user-1", + MemoryType: dao.MemoryTypeRaw, + StorageType: "table", + LLMID: "raw-llm", + EmbdID: "raw-embd", + Permissions: string(TenantPermissionMe), + ForgettingPolicy: string(ForgettingPolicyFIFO), + }).Error; err != nil { + t.Fatalf("seed memory: %v", err) + } + + resp, err := NewMemoryService().ListMemories("user-1", []string{"user-1"}, nil, "", "", 1, 10) + if err != nil { + t.Fatalf("ListMemories: %v", err) + } + if resp.TotalCount != 1 || len(resp.MemoryList) != 1 { + t.Fatalf("ListMemories returned total=%d len=%d, want 1", resp.TotalCount, len(resp.MemoryList)) + } + memory := resp.MemoryList[0] + if got, want := memory["llm_id"], "raw-llm"; got != want { + t.Fatalf("llm_id = %v, want %v", got, want) + } + if got, want := memory["embd_id"], "raw-embd"; got != want { + t.Fatalf("embd_id = %v, want %v", got, want) + } +} + func seedMemoryMessages(t *testing.T) { t.Helper()