0.29.1原版
This commit is contained in:
@@ -0,0 +1,284 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
)
|
||||
|
||||
func TestTranscribe(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("requires authentication", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
_, err := ts.Service.Transcribe(ctx, &v1pb.TranscribeRequest{
|
||||
Audio: &v1pb.TranscriptionAudio{
|
||||
Source: &v1pb.TranscriptionAudio_Content{Content: []byte("RIFF")},
|
||||
Filename: "voice.wav",
|
||||
ContentType: "audio/wav",
|
||||
},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "user not authenticated")
|
||||
})
|
||||
|
||||
t.Run("transcribes audio file using persisted transcription setting", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "alice")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
openAIServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
require.Equal(t, "/audio/transcriptions", r.URL.Path)
|
||||
require.Equal(t, "Bearer sk-test", r.Header.Get("Authorization"))
|
||||
require.NoError(t, r.ParseMultipartForm(10<<20))
|
||||
require.Equal(t, "whisper-1", r.FormValue("model"))
|
||||
require.Equal(t, "fr", r.FormValue("language"))
|
||||
require.Equal(t, "names: Alice", r.FormValue("prompt"))
|
||||
|
||||
file, header, err := r.FormFile("file")
|
||||
require.NoError(t, err)
|
||||
defer file.Close()
|
||||
require.Equal(t, "voice.wav", header.Filename)
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
require.NoError(t, json.NewEncoder(w).Encode(map[string]string{
|
||||
"text": "transcribed text",
|
||||
}))
|
||||
}))
|
||||
defer openAIServer.Close()
|
||||
|
||||
_, err = ts.Store.UpsertInstanceSetting(ctx, &storepb.InstanceSetting{
|
||||
Key: storepb.InstanceSettingKey_AI,
|
||||
Value: &storepb.InstanceSetting_AiSetting{
|
||||
AiSetting: &storepb.InstanceAISetting{
|
||||
Providers: []*storepb.AIProviderConfig{
|
||||
{
|
||||
Id: "openai-main",
|
||||
Title: "OpenAI",
|
||||
Type: storepb.AIProviderType_OPENAI,
|
||||
Endpoint: openAIServer.URL,
|
||||
ApiKey: "sk-test",
|
||||
},
|
||||
},
|
||||
Transcription: &storepb.TranscriptionConfig{
|
||||
ProviderId: "openai-main",
|
||||
Model: "whisper-1",
|
||||
Language: "fr",
|
||||
Prompt: "names: Alice",
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.Transcribe(userCtx, &v1pb.TranscribeRequest{
|
||||
Audio: &v1pb.TranscriptionAudio{
|
||||
Source: &v1pb.TranscriptionAudio_Content{Content: []byte("RIFF")},
|
||||
Filename: "voice.wav",
|
||||
ContentType: "audio/wav",
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "transcribed text", resp.Text)
|
||||
})
|
||||
|
||||
t.Run("returns provider error without rewriting it", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "notfound-user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
openAIServer := httptest.NewServer(http.NotFoundHandler())
|
||||
defer openAIServer.Close()
|
||||
|
||||
_, err = ts.Store.UpsertInstanceSetting(ctx, &storepb.InstanceSetting{
|
||||
Key: storepb.InstanceSettingKey_AI,
|
||||
Value: &storepb.InstanceSetting_AiSetting{
|
||||
AiSetting: &storepb.InstanceAISetting{
|
||||
Providers: []*storepb.AIProviderConfig{
|
||||
{
|
||||
Id: "openai-main",
|
||||
Title: "OpenAI",
|
||||
Type: storepb.AIProviderType_OPENAI,
|
||||
Endpoint: openAIServer.URL,
|
||||
ApiKey: "sk-test",
|
||||
},
|
||||
},
|
||||
Transcription: &storepb.TranscriptionConfig{
|
||||
ProviderId: "openai-main",
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.Transcribe(userCtx, &v1pb.TranscribeRequest{
|
||||
Audio: &v1pb.TranscriptionAudio{
|
||||
Source: &v1pb.TranscriptionAudio_Content{Content: []byte("RIFF")},
|
||||
Filename: "voice.wav",
|
||||
ContentType: "audio/wav",
|
||||
},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "failed to transcribe audio")
|
||||
})
|
||||
|
||||
t.Run("transcribes audio file with Gemini provider", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "gemini-user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
geminiServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
require.Equal(t, "/v1beta/models/gemini-2.5-flash:generateContent", r.URL.Path)
|
||||
require.Equal(t, "gemini-key", r.Header.Get("x-goog-api-key"))
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
require.NoError(t, json.NewEncoder(w).Encode(map[string]any{
|
||||
"candidates": []map[string]any{
|
||||
{
|
||||
"finishReason": "STOP",
|
||||
"content": map[string]any{
|
||||
"parts": []map[string]string{{"text": "gemini transcript"}},
|
||||
},
|
||||
},
|
||||
},
|
||||
}))
|
||||
}))
|
||||
defer geminiServer.Close()
|
||||
|
||||
_, err = ts.Store.UpsertInstanceSetting(ctx, &storepb.InstanceSetting{
|
||||
Key: storepb.InstanceSettingKey_AI,
|
||||
Value: &storepb.InstanceSetting_AiSetting{
|
||||
AiSetting: &storepb.InstanceAISetting{
|
||||
Providers: []*storepb.AIProviderConfig{
|
||||
{
|
||||
Id: "gemini-main",
|
||||
Title: "Gemini",
|
||||
Type: storepb.AIProviderType_GEMINI,
|
||||
Endpoint: geminiServer.URL + "/v1beta",
|
||||
ApiKey: "gemini-key",
|
||||
},
|
||||
},
|
||||
Transcription: &storepb.TranscriptionConfig{
|
||||
ProviderId: "gemini-main",
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.Transcribe(userCtx, &v1pb.TranscribeRequest{
|
||||
Audio: &v1pb.TranscriptionAudio{
|
||||
Source: &v1pb.TranscriptionAudio_Content{Content: []byte("mp3 bytes")},
|
||||
Filename: "voice.mp3",
|
||||
ContentType: "audio/mp3",
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "gemini transcript", resp.Text)
|
||||
})
|
||||
|
||||
t.Run("falls back to engine default model when transcription model is empty", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "bob")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
openAIServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
require.NoError(t, r.ParseMultipartForm(10<<20))
|
||||
require.Equal(t, "whisper-1", r.FormValue("model"))
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
require.NoError(t, json.NewEncoder(w).Encode(map[string]string{
|
||||
"text": "built-in model",
|
||||
}))
|
||||
}))
|
||||
defer openAIServer.Close()
|
||||
|
||||
_, err = ts.Store.UpsertInstanceSetting(ctx, &storepb.InstanceSetting{
|
||||
Key: storepb.InstanceSettingKey_AI,
|
||||
Value: &storepb.InstanceSetting_AiSetting{
|
||||
AiSetting: &storepb.InstanceAISetting{
|
||||
Providers: []*storepb.AIProviderConfig{
|
||||
{
|
||||
Id: "openai-main",
|
||||
Title: "OpenAI",
|
||||
Type: storepb.AIProviderType_OPENAI,
|
||||
Endpoint: openAIServer.URL,
|
||||
ApiKey: "sk-test",
|
||||
},
|
||||
},
|
||||
Transcription: &storepb.TranscriptionConfig{
|
||||
ProviderId: "openai-main",
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.Transcribe(userCtx, &v1pb.TranscribeRequest{
|
||||
Audio: &v1pb.TranscriptionAudio{
|
||||
Source: &v1pb.TranscriptionAudio_Content{Content: []byte("RIFF")},
|
||||
Filename: "voice.wav",
|
||||
ContentType: "audio/wav",
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "built-in model", resp.Text)
|
||||
})
|
||||
|
||||
t.Run("rejects non-audio content before provider call", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "charlie")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
_, err = ts.Service.Transcribe(userCtx, &v1pb.TranscribeRequest{
|
||||
Audio: &v1pb.TranscriptionAudio{
|
||||
Source: &v1pb.TranscriptionAudio_Content{Content: []byte("not audio")},
|
||||
Filename: "notes.txt",
|
||||
ContentType: "text/plain",
|
||||
},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "not supported")
|
||||
})
|
||||
|
||||
t.Run("returns FailedPrecondition when transcription is not configured", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "alice-empty")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
_, err = ts.Service.Transcribe(userCtx, &v1pb.TranscribeRequest{
|
||||
Audio: &v1pb.TranscriptionAudio{
|
||||
Source: &v1pb.TranscriptionAudio_Content{Content: []byte("RIFF")},
|
||||
Filename: "voice.wav",
|
||||
ContentType: "audio/wav",
|
||||
},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "transcription is not configured")
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,371 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/usememos/memos/internal/testutil"
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
apiv1 "github.com/usememos/memos/server/router/api/v1"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func TestCreateAttachment(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "test_user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
// Test case 1: Create attachment with empty type but known extension
|
||||
t.Run("EmptyType_KnownExtension", func(t *testing.T) {
|
||||
attachment, err := ts.Service.CreateAttachment(userCtx, &v1pb.CreateAttachmentRequest{
|
||||
Attachment: &v1pb.Attachment{
|
||||
Filename: "test.png",
|
||||
Content: []byte("fake png content"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "image/png", attachment.Type)
|
||||
})
|
||||
|
||||
// Test case 2: Create attachment with empty type and unknown extension, but detectable content
|
||||
t.Run("EmptyType_UnknownExtension_ContentSniffing", func(t *testing.T) {
|
||||
// PNG magic header: 89 50 4E 47 0D 0A 1A 0A
|
||||
pngContent := []byte{0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A}
|
||||
attachment, err := ts.Service.CreateAttachment(userCtx, &v1pb.CreateAttachmentRequest{
|
||||
Attachment: &v1pb.Attachment{
|
||||
Filename: "test.unknown",
|
||||
Content: pngContent,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "image/png", attachment.Type)
|
||||
})
|
||||
|
||||
// Test case 3: Empty type, unknown extension, random content -> fallback to application/octet-stream
|
||||
t.Run("EmptyType_Fallback", func(t *testing.T) {
|
||||
randomContent := []byte{0x00, 0x01, 0x02, 0x03}
|
||||
attachment, err := ts.Service.CreateAttachment(userCtx, &v1pb.CreateAttachmentRequest{
|
||||
Attachment: &v1pb.Attachment{
|
||||
Filename: "test.data",
|
||||
Content: randomContent,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "application/octet-stream", attachment.Type)
|
||||
})
|
||||
|
||||
t.Run("Type_WithParameters_NormalizedBeforeValidation", func(t *testing.T) {
|
||||
attachment, err := ts.Service.CreateAttachment(userCtx, &v1pb.CreateAttachmentRequest{
|
||||
Attachment: &v1pb.Attachment{
|
||||
Filename: "voice-note.webm",
|
||||
Type: "audio/webm;codecs=opus",
|
||||
Content: []byte("fake webm content"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "audio/webm", attachment.Type)
|
||||
})
|
||||
|
||||
t.Run("Type_InvalidFormat_Rejected", func(t *testing.T) {
|
||||
_, err := ts.Service.CreateAttachment(userCtx, &v1pb.CreateAttachmentRequest{
|
||||
Attachment: &v1pb.Attachment{
|
||||
Filename: "broken.webm",
|
||||
Type: `audio/webm;codecs="unterminated`,
|
||||
Content: []byte("fake webm content"),
|
||||
},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "invalid MIME type format")
|
||||
})
|
||||
|
||||
t.Run("LocalStorage_PathCollisionUsesUniqueReference", func(t *testing.T) {
|
||||
_, err := ts.Store.UpsertInstanceSetting(ctx, &storepb.InstanceSetting{
|
||||
Key: storepb.InstanceSettingKey_STORAGE,
|
||||
Value: &storepb.InstanceSetting_StorageSetting{
|
||||
StorageSetting: &storepb.InstanceStorageSetting{
|
||||
StorageType: storepb.InstanceStorageSetting_LOCAL,
|
||||
FilepathTemplate: "assets/{filename}",
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
first, err := ts.Service.CreateAttachment(userCtx, &v1pb.CreateAttachmentRequest{
|
||||
Attachment: &v1pb.Attachment{
|
||||
Filename: "screenshot.png",
|
||||
Type: "image/png",
|
||||
Content: []byte("first-image"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
second, err := ts.Service.CreateAttachment(userCtx, &v1pb.CreateAttachmentRequest{
|
||||
Attachment: &v1pb.Attachment{
|
||||
Filename: "screenshot.png",
|
||||
Type: "image/png",
|
||||
Content: []byte("second-image"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
firstUID, err := apiv1.ExtractAttachmentUIDFromName(first.Name)
|
||||
require.NoError(t, err)
|
||||
secondUID, err := apiv1.ExtractAttachmentUIDFromName(second.Name)
|
||||
require.NoError(t, err)
|
||||
|
||||
firstStoreAttachment, err := ts.Store.GetAttachment(ctx, &store.FindAttachment{UID: &firstUID})
|
||||
require.NoError(t, err)
|
||||
secondStoreAttachment, err := ts.Store.GetAttachment(ctx, &store.FindAttachment{UID: &secondUID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, firstStoreAttachment)
|
||||
require.NotNil(t, secondStoreAttachment)
|
||||
|
||||
require.NotEqual(t, firstStoreAttachment.Reference, secondStoreAttachment.Reference)
|
||||
|
||||
firstBlob, err := ts.Service.GetAttachmentBlob(firstStoreAttachment)
|
||||
require.NoError(t, err)
|
||||
secondBlob, err := ts.Service.GetAttachmentBlob(secondStoreAttachment)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []byte("first-image"), firstBlob)
|
||||
require.Equal(t, []byte("second-image"), secondBlob)
|
||||
})
|
||||
}
|
||||
|
||||
func TestCreateAttachmentMemoPermission(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("owner can create attachment directly linked to memo", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
owner, err := ts.CreateRegularUser(ctx, "attachment-owner")
|
||||
require.NoError(t, err)
|
||||
ownerCtx := ts.CreateUserContext(ctx, owner.ID)
|
||||
|
||||
memo, err := ts.Service.CreateMemo(ownerCtx, &v1pb.CreateMemoRequest{
|
||||
Memo: &v1pb.Memo{
|
||||
Content: "memo with direct attachment",
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
attachment, err := ts.Service.CreateAttachment(ownerCtx, &v1pb.CreateAttachmentRequest{
|
||||
Attachment: &v1pb.Attachment{
|
||||
Filename: "owner.txt",
|
||||
Type: "text/plain",
|
||||
Content: []byte("owner"),
|
||||
Memo: &memo.Name,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
attachmentUID, err := apiv1.ExtractAttachmentUIDFromName(attachment.Name)
|
||||
require.NoError(t, err)
|
||||
stored, err := ts.Store.GetAttachment(ctx, &store.FindAttachment{UID: &attachmentUID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, stored.MemoID)
|
||||
require.Equal(t, memoIDFromName(ctx, t, ts, memo.Name), *stored.MemoID)
|
||||
})
|
||||
|
||||
t.Run("admin can create attachment directly linked to memo", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
owner, err := ts.CreateRegularUser(ctx, "attachment-admin-owner")
|
||||
require.NoError(t, err)
|
||||
ownerCtx := ts.CreateUserContext(ctx, owner.ID)
|
||||
admin, err := ts.CreateHostUser(ctx, "attachment-admin")
|
||||
require.NoError(t, err)
|
||||
adminCtx := ts.CreateUserContext(ctx, admin.ID)
|
||||
|
||||
memo, err := ts.Service.CreateMemo(ownerCtx, &v1pb.CreateMemoRequest{
|
||||
Memo: &v1pb.Memo{
|
||||
Content: "memo with admin attachment",
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
attachment, err := ts.Service.CreateAttachment(adminCtx, &v1pb.CreateAttachmentRequest{
|
||||
Attachment: &v1pb.Attachment{
|
||||
Filename: "admin.txt",
|
||||
Type: "text/plain",
|
||||
Content: []byte("admin"),
|
||||
Memo: &memo.Name,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
attachmentUID, err := apiv1.ExtractAttachmentUIDFromName(attachment.Name)
|
||||
require.NoError(t, err)
|
||||
stored, err := ts.Store.GetAttachment(ctx, &store.FindAttachment{UID: &attachmentUID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, stored.MemoID)
|
||||
require.Equal(t, memoIDFromName(ctx, t, ts, memo.Name), *stored.MemoID)
|
||||
})
|
||||
|
||||
t.Run("non-owner cannot create attachment directly linked to memo", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
owner, err := ts.CreateRegularUser(ctx, "attachment-owner-denied")
|
||||
require.NoError(t, err)
|
||||
ownerCtx := ts.CreateUserContext(ctx, owner.ID)
|
||||
other, err := ts.CreateRegularUser(ctx, "attachment-other-denied")
|
||||
require.NoError(t, err)
|
||||
otherCtx := ts.CreateUserContext(ctx, other.ID)
|
||||
|
||||
memo, err := ts.Service.CreateMemo(ownerCtx, &v1pb.CreateMemoRequest{
|
||||
Memo: &v1pb.Memo{
|
||||
Content: "memo with blocked attachment",
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.CreateAttachment(otherCtx, &v1pb.CreateAttachmentRequest{
|
||||
Attachment: &v1pb.Attachment{
|
||||
Filename: "blocked.txt",
|
||||
Type: "text/plain",
|
||||
Content: []byte("blocked"),
|
||||
Memo: &memo.Name,
|
||||
},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "permission denied")
|
||||
|
||||
attachments, err := ts.Store.ListAttachments(ctx, &store.FindAttachment{
|
||||
CreatorID: &other.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, attachments)
|
||||
})
|
||||
}
|
||||
|
||||
func memoIDFromName(ctx context.Context, t *testing.T, ts *TestService, name string) int32 {
|
||||
t.Helper()
|
||||
memoUID, err := apiv1.ExtractMemoUIDFromName(name)
|
||||
require.NoError(t, err)
|
||||
memo, err := ts.Store.GetMemo(ctx, &store.FindMemo{UID: &memoUID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memo)
|
||||
return memo.ID
|
||||
}
|
||||
|
||||
func TestCreateAttachmentMotionMedia(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "motion_user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
t.Run("Apple live photo metadata roundtrip", func(t *testing.T) {
|
||||
attachment, err := ts.Service.CreateAttachment(userCtx, &v1pb.CreateAttachmentRequest{
|
||||
Attachment: &v1pb.Attachment{
|
||||
Filename: "live.heic",
|
||||
Type: "image/heic",
|
||||
Content: []byte("fake-heic-still"),
|
||||
MotionMedia: &v1pb.MotionMedia{
|
||||
Family: v1pb.MotionMediaFamily_APPLE_LIVE_PHOTO,
|
||||
Role: v1pb.MotionMediaRole_STILL,
|
||||
GroupId: "apple-group-1",
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, attachment.MotionMedia)
|
||||
require.Equal(t, v1pb.MotionMediaFamily_APPLE_LIVE_PHOTO, attachment.MotionMedia.Family)
|
||||
require.Equal(t, v1pb.MotionMediaRole_STILL, attachment.MotionMedia.Role)
|
||||
require.Equal(t, "apple-group-1", attachment.MotionMedia.GroupId)
|
||||
})
|
||||
|
||||
t.Run("Android motion photo detection", func(t *testing.T) {
|
||||
attachment, err := ts.Service.CreateAttachment(userCtx, &v1pb.CreateAttachmentRequest{
|
||||
Attachment: &v1pb.Attachment{
|
||||
Filename: "motion.jpg",
|
||||
Type: "image/jpeg",
|
||||
Content: testutil.BuildMotionPhotoJPEG(),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, attachment.MotionMedia)
|
||||
require.Equal(t, v1pb.MotionMediaFamily_ANDROID_MOTION_PHOTO, attachment.MotionMedia.Family)
|
||||
require.Equal(t, v1pb.MotionMediaRole_CONTAINER, attachment.MotionMedia.Role)
|
||||
require.True(t, attachment.MotionMedia.HasEmbeddedVideo)
|
||||
})
|
||||
}
|
||||
|
||||
func TestBatchDeleteAttachments(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "delete_user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
first, err := ts.Service.CreateAttachment(userCtx, &v1pb.CreateAttachmentRequest{
|
||||
Attachment: &v1pb.Attachment{Filename: "one.txt", Type: "text/plain", Content: []byte("one")},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
second, err := ts.Service.CreateAttachment(userCtx, &v1pb.CreateAttachmentRequest{
|
||||
Attachment: &v1pb.Attachment{Filename: "two.txt", Type: "text/plain", Content: []byte("two")},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.BatchDeleteAttachments(userCtx, &v1pb.BatchDeleteAttachmentsRequest{
|
||||
Names: []string{first.Name, second.Name},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
firstUID, err := apiv1.ExtractAttachmentUIDFromName(first.Name)
|
||||
require.NoError(t, err)
|
||||
secondUID, err := apiv1.ExtractAttachmentUIDFromName(second.Name)
|
||||
require.NoError(t, err)
|
||||
storedFirst, err := ts.Store.GetAttachment(ctx, &store.FindAttachment{UID: &firstUID})
|
||||
require.NoError(t, err)
|
||||
storedSecond, err := ts.Store.GetAttachment(ctx, &store.FindAttachment{UID: &secondUID})
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, storedFirst)
|
||||
require.Nil(t, storedSecond)
|
||||
|
||||
t.Run("deduplicates duplicate names", func(t *testing.T) {
|
||||
third, err := ts.Service.CreateAttachment(userCtx, &v1pb.CreateAttachmentRequest{
|
||||
Attachment: &v1pb.Attachment{Filename: "three.txt", Type: "text/plain", Content: []byte("three")},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.BatchDeleteAttachments(userCtx, &v1pb.BatchDeleteAttachmentsRequest{
|
||||
Names: []string{third.Name, third.Name},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
thirdUID, err := apiv1.ExtractAttachmentUIDFromName(third.Name)
|
||||
require.NoError(t, err)
|
||||
storedThird, err := ts.Store.GetAttachment(ctx, &store.FindAttachment{UID: &thirdUID})
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, storedThird)
|
||||
})
|
||||
|
||||
t.Run("rejects unauthorized deletes", func(t *testing.T) {
|
||||
ownerAttachment, err := ts.Service.CreateAttachment(userCtx, &v1pb.CreateAttachmentRequest{
|
||||
Attachment: &v1pb.Attachment{Filename: "private.txt", Type: "text/plain", Content: []byte("private")},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
otherUser, err := ts.CreateRegularUser(ctx, "other_delete_user")
|
||||
require.NoError(t, err)
|
||||
otherCtx := ts.CreateUserContext(ctx, otherUser.ID)
|
||||
|
||||
_, err = ts.Service.BatchDeleteAttachments(otherCtx, &v1pb.BatchDeleteAttachmentsRequest{
|
||||
Names: []string{ownerAttachment.Name},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "permission denied")
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,271 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
apiv1 "github.com/usememos/memos/server/router/api/v1"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func TestCreateLinkedIdentityBindsCurrentUser(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
currentUser, err := ts.CreateRegularUser(ctx, "alice")
|
||||
require.NoError(t, err)
|
||||
|
||||
mockIDP := newMockOAuthServer(t, "bind-code", "bind-access-token", map[string]any{
|
||||
"sub": "google-sub-1",
|
||||
"name": "Alice Example",
|
||||
"email": "alice@example.com",
|
||||
})
|
||||
defer mockIDP.Close()
|
||||
|
||||
idpName := createTestingOAuthIdentityProvider(ctx, t, ts, mockIDP.URL, "google-bind")
|
||||
beforeUsers, err := ts.Store.ListUsers(ctx, &store.FindUser{})
|
||||
require.NoError(t, err)
|
||||
|
||||
authCtx := ts.CreateUserContext(apiv1.WithHeaderCarrier(ctx), currentUser.ID)
|
||||
response, err := ts.Service.CreateLinkedIdentity(authCtx, &v1pb.CreateLinkedIdentityRequest{
|
||||
Parent: apiv1.BuildUserName(currentUser.Username),
|
||||
IdpName: idpName,
|
||||
Code: "bind-code",
|
||||
RedirectUri: "http://localhost:8080/auth/callback",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, response)
|
||||
require.Equal(t, apiv1.BuildUserName(currentUser.Username)+"/linkedIdentities/google-bind", response.Name)
|
||||
require.Equal(t, apiv1.IdentityProviderNamePrefix+"google-bind", response.IdpName)
|
||||
require.Equal(t, "google-sub-1", response.ExternUid)
|
||||
|
||||
afterUsers, err := ts.Store.ListUsers(ctx, &store.FindUser{})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, afterUsers, len(beforeUsers))
|
||||
|
||||
provider := "google-bind"
|
||||
externUID := "google-sub-1"
|
||||
identity, err := ts.Store.GetUserIdentity(ctx, &store.FindUserIdentity{
|
||||
Provider: &provider,
|
||||
ExternUID: &externUID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, identity)
|
||||
require.Equal(t, currentUser.ID, identity.UserID)
|
||||
}
|
||||
|
||||
func TestCreateLinkedIdentityRejectsBindingIdentityLinkedToAnotherUser(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
owner, err := ts.CreateRegularUser(ctx, "owner")
|
||||
require.NoError(t, err)
|
||||
binder, err := ts.CreateRegularUser(ctx, "binder")
|
||||
require.NoError(t, err)
|
||||
|
||||
mockIDP := newMockOAuthServer(t, "conflict-code", "conflict-access-token", map[string]any{
|
||||
"sub": "google-sub-2",
|
||||
"name": "Conflict Example",
|
||||
"email": "conflict@example.com",
|
||||
})
|
||||
defer mockIDP.Close()
|
||||
|
||||
idpName := createTestingOAuthIdentityProvider(ctx, t, ts, mockIDP.URL, "google-conflict")
|
||||
_, err = ts.Store.CreateUserIdentity(ctx, &store.UserIdentity{
|
||||
UserID: owner.ID,
|
||||
Provider: "google-conflict",
|
||||
ExternUID: "google-sub-2",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
authCtx := ts.CreateUserContext(apiv1.WithHeaderCarrier(ctx), binder.ID)
|
||||
_, err = ts.Service.CreateLinkedIdentity(authCtx, &v1pb.CreateLinkedIdentityRequest{
|
||||
Parent: apiv1.BuildUserName(binder.Username),
|
||||
IdpName: idpName,
|
||||
Code: "conflict-code",
|
||||
RedirectUri: "http://localhost:8080/auth/callback",
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Equal(t, codes.AlreadyExists, status.Code(err))
|
||||
}
|
||||
|
||||
func TestListAndDeleteLinkedIdentities(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
currentUser, err := ts.CreateRegularUser(ctx, "alice")
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Store.CreateUserIdentity(ctx, &store.UserIdentity{
|
||||
UserID: currentUser.ID,
|
||||
Provider: "google",
|
||||
ExternUID: "alice@gmail.com",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
authCtx := ts.CreateUserContext(ctx, currentUser.ID)
|
||||
listResp, err := ts.Service.ListLinkedIdentities(authCtx, &v1pb.ListLinkedIdentitiesRequest{
|
||||
Parent: apiv1.BuildUserName(currentUser.Username),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, listResp.LinkedIdentities, 1)
|
||||
linkedIdentityName := apiv1.BuildUserName(currentUser.Username) + "/linkedIdentities/google"
|
||||
require.Equal(t, linkedIdentityName, listResp.LinkedIdentities[0].Name)
|
||||
require.Equal(t, apiv1.IdentityProviderNamePrefix+"google", listResp.LinkedIdentities[0].IdpName)
|
||||
require.Equal(t, "alice@gmail.com", listResp.LinkedIdentities[0].ExternUid)
|
||||
|
||||
got, err := ts.Service.GetLinkedIdentity(authCtx, &v1pb.GetLinkedIdentityRequest{
|
||||
Name: linkedIdentityName,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, linkedIdentityName, got.Name)
|
||||
require.Equal(t, apiv1.IdentityProviderNamePrefix+"google", got.IdpName)
|
||||
require.Equal(t, "alice@gmail.com", got.ExternUid)
|
||||
|
||||
_, err = ts.Service.DeleteLinkedIdentity(authCtx, &v1pb.DeleteLinkedIdentityRequest{
|
||||
Name: linkedIdentityName,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
listResp, err = ts.Service.ListLinkedIdentities(authCtx, &v1pb.ListLinkedIdentitiesRequest{
|
||||
Parent: apiv1.BuildUserName(currentUser.Username),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, listResp.LinkedIdentities)
|
||||
}
|
||||
|
||||
func TestListLinkedIdentitiesRequiresAuthentication(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
user, err := ts.CreateRegularUser(ctx, "linked-identity-auth")
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.ListLinkedIdentities(ctx, &v1pb.ListLinkedIdentitiesRequest{
|
||||
Parent: apiv1.BuildUserName(user.Username),
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Equal(t, codes.Unauthenticated, status.Code(err))
|
||||
}
|
||||
|
||||
func TestCreateLinkedIdentityRejectsSecondIdentityForSameProvider(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
currentUser, err := ts.CreateRegularUser(ctx, "alice")
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Store.CreateUserIdentity(ctx, &store.UserIdentity{
|
||||
UserID: currentUser.ID,
|
||||
Provider: "google-provider",
|
||||
ExternUID: "google-sub-1",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
mockIDP := newMockOAuthServer(t, "second-code", "second-access-token", map[string]any{
|
||||
"sub": "google-sub-2",
|
||||
"name": "Alice Example",
|
||||
"email": "alice@example.com",
|
||||
})
|
||||
defer mockIDP.Close()
|
||||
|
||||
idpName := createTestingOAuthIdentityProvider(ctx, t, ts, mockIDP.URL, "google-provider")
|
||||
authCtx := ts.CreateUserContext(apiv1.WithHeaderCarrier(ctx), currentUser.ID)
|
||||
|
||||
_, err = ts.Service.CreateLinkedIdentity(authCtx, &v1pb.CreateLinkedIdentityRequest{
|
||||
Parent: apiv1.BuildUserName(currentUser.Username),
|
||||
IdpName: idpName,
|
||||
Code: "second-code",
|
||||
RedirectUri: "http://localhost:8080/auth/callback",
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Equal(t, codes.AlreadyExists, status.Code(err))
|
||||
}
|
||||
|
||||
func createTestingOAuthIdentityProvider(ctx context.Context, t *testing.T, ts *TestService, serverURL, uid string) string {
|
||||
t.Helper()
|
||||
|
||||
idp, err := ts.Store.CreateIdentityProvider(ctx, &storepb.IdentityProvider{
|
||||
Uid: uid,
|
||||
Name: "Google",
|
||||
Type: storepb.IdentityProvider_OAUTH2,
|
||||
Config: &storepb.IdentityProviderConfig{
|
||||
Config: &storepb.IdentityProviderConfig_Oauth2Config{
|
||||
Oauth2Config: &storepb.OAuth2Config{
|
||||
ClientId: "test-client-id",
|
||||
ClientSecret: "test-client-secret",
|
||||
AuthUrl: serverURL + "/oauth2/authorize",
|
||||
TokenUrl: serverURL + "/oauth2/token",
|
||||
UserInfoUrl: serverURL + "/oauth2/userinfo",
|
||||
FieldMapping: &storepb.FieldMapping{
|
||||
Identifier: "sub",
|
||||
DisplayName: "name",
|
||||
Email: "email",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
return apiv1.IdentityProviderNamePrefix + idp.Uid
|
||||
}
|
||||
|
||||
func newMockOAuthServer(t *testing.T, code, accessToken string, userInfo map[string]any) *httptest.Server {
|
||||
t.Helper()
|
||||
|
||||
userInfoBytes, err := json.Marshal(userInfo)
|
||||
require.NoError(t, err)
|
||||
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("/oauth2/token", func(w http.ResponseWriter, r *http.Request) {
|
||||
require.Equal(t, http.MethodPost, r.Method)
|
||||
|
||||
body, err := io.ReadAll(r.Body)
|
||||
require.NoError(t, err)
|
||||
values, err := url.ParseQuery(string(body))
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, code, values.Get("code"))
|
||||
require.Equal(t, "authorization_code", values.Get("grant_type"))
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
err = json.NewEncoder(w).Encode(map[string]any{
|
||||
"access_token": accessToken,
|
||||
"token_type": "Bearer",
|
||||
"expires_in": 3600,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
})
|
||||
mux.HandleFunc("/oauth2/userinfo", func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, err := w.Write(userInfoBytes)
|
||||
require.NoError(t, err)
|
||||
})
|
||||
|
||||
return httptest.NewServer(mux)
|
||||
}
|
||||
@@ -0,0 +1,682 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/protobuf/types/known/timestamppb"
|
||||
|
||||
"github.com/usememos/memos/internal/util"
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
"github.com/usememos/memos/server/auth"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func TestAuthenticatorAccessTokenV2(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("authenticates valid access token v2", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create a test user
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Generate access token v2
|
||||
token, _, err := auth.GenerateAccessTokenV2(
|
||||
user.ID,
|
||||
user.Username,
|
||||
string(user.Role),
|
||||
string(user.RowStatus),
|
||||
[]byte(ts.Secret),
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Authenticate
|
||||
authenticator := auth.NewAuthenticator(ts.Store, ts.Secret)
|
||||
claims, err := authenticator.AuthenticateByAccessTokenV2(token)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, claims)
|
||||
assert.Equal(t, user.ID, claims.UserID)
|
||||
assert.Equal(t, user.Username, claims.Username)
|
||||
assert.Equal(t, string(user.Role), claims.Role)
|
||||
assert.Equal(t, string(user.RowStatus), claims.Status)
|
||||
})
|
||||
|
||||
t.Run("fails with invalid token", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
authenticator := auth.NewAuthenticator(ts.Store, ts.Secret)
|
||||
_, err := authenticator.AuthenticateByAccessTokenV2("invalid-token")
|
||||
assert.Error(t, err)
|
||||
})
|
||||
|
||||
t.Run("fails with wrong secret", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Generate token with one secret
|
||||
token, _, err := auth.GenerateAccessTokenV2(
|
||||
user.ID,
|
||||
user.Username,
|
||||
string(user.Role),
|
||||
string(user.RowStatus),
|
||||
[]byte("secret-1"),
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Try to authenticate with different secret
|
||||
authenticator := auth.NewAuthenticator(ts.Store, "secret-2")
|
||||
_, err = authenticator.AuthenticateByAccessTokenV2(token)
|
||||
assert.Error(t, err)
|
||||
})
|
||||
|
||||
t.Run("request authentication rejects archived user", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "archived-access-token")
|
||||
require.NoError(t, err)
|
||||
token, _, err := auth.GenerateAccessTokenV2(
|
||||
user.ID,
|
||||
user.Username,
|
||||
string(user.Role),
|
||||
string(user.RowStatus),
|
||||
[]byte(ts.Secret),
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
archivedStatus := store.Archived
|
||||
_, err = ts.Store.UpdateUser(ctx, &store.UpdateUser{
|
||||
ID: user.ID,
|
||||
RowStatus: &archivedStatus,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
authenticator := auth.NewAuthenticator(ts.Store, ts.Secret)
|
||||
result := authenticator.Authenticate(ctx, "Bearer "+token)
|
||||
assert.Nil(t, result)
|
||||
})
|
||||
}
|
||||
|
||||
func TestAuthenticatorRefreshToken(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("authenticates valid refresh token", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create a test user
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create refresh token record in store
|
||||
tokenID := util.GenUUID()
|
||||
refreshTokenRecord := &storepb.RefreshTokensUserSetting_RefreshToken{
|
||||
TokenId: tokenID,
|
||||
ExpiresAt: timestamppb.New(time.Now().Add(auth.RefreshTokenDuration)),
|
||||
CreatedAt: timestamppb.Now(),
|
||||
}
|
||||
err = ts.Store.AddUserRefreshToken(ctx, user.ID, refreshTokenRecord)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Generate refresh token JWT
|
||||
token, _, err := auth.GenerateRefreshToken(user.ID, tokenID, []byte(ts.Secret))
|
||||
require.NoError(t, err)
|
||||
|
||||
// Authenticate
|
||||
authenticator := auth.NewAuthenticator(ts.Store, ts.Secret)
|
||||
authenticatedUser, returnedTokenID, err := authenticator.AuthenticateByRefreshToken(ctx, token)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, authenticatedUser)
|
||||
assert.Equal(t, user.ID, authenticatedUser.ID)
|
||||
assert.Equal(t, tokenID, returnedTokenID)
|
||||
})
|
||||
|
||||
t.Run("fails with revoked token", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
tokenID := util.GenUUID()
|
||||
|
||||
// Generate refresh token JWT but don't store it in database (simulates revocation)
|
||||
token, _, err := auth.GenerateRefreshToken(user.ID, tokenID, []byte(ts.Secret))
|
||||
require.NoError(t, err)
|
||||
|
||||
// Try to authenticate
|
||||
authenticator := auth.NewAuthenticator(ts.Store, ts.Secret)
|
||||
_, _, err = authenticator.AuthenticateByRefreshToken(ctx, token)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "revoked")
|
||||
})
|
||||
|
||||
t.Run("fails with expired token", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create expired refresh token record in store
|
||||
tokenID := util.GenUUID()
|
||||
expiredToken := &storepb.RefreshTokensUserSetting_RefreshToken{
|
||||
TokenId: tokenID,
|
||||
ExpiresAt: timestamppb.New(time.Now().Add(-1 * time.Hour)), // Expired
|
||||
CreatedAt: timestamppb.Now(),
|
||||
}
|
||||
err = ts.Store.AddUserRefreshToken(ctx, user.ID, expiredToken)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Generate refresh token JWT (JWT itself isn't expired yet)
|
||||
token, _, err := auth.GenerateRefreshToken(user.ID, tokenID, []byte(ts.Secret))
|
||||
require.NoError(t, err)
|
||||
|
||||
// Try to authenticate
|
||||
authenticator := auth.NewAuthenticator(ts.Store, ts.Secret)
|
||||
_, _, err = authenticator.AuthenticateByRefreshToken(ctx, token)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "expired")
|
||||
})
|
||||
|
||||
t.Run("fails with archived user", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create valid refresh token
|
||||
tokenID := util.GenUUID()
|
||||
refreshTokenRecord := &storepb.RefreshTokensUserSetting_RefreshToken{
|
||||
TokenId: tokenID,
|
||||
ExpiresAt: timestamppb.New(time.Now().Add(auth.RefreshTokenDuration)),
|
||||
CreatedAt: timestamppb.Now(),
|
||||
}
|
||||
err = ts.Store.AddUserRefreshToken(ctx, user.ID, refreshTokenRecord)
|
||||
require.NoError(t, err)
|
||||
|
||||
token, _, err := auth.GenerateRefreshToken(user.ID, tokenID, []byte(ts.Secret))
|
||||
require.NoError(t, err)
|
||||
|
||||
// Archive the user
|
||||
archivedStatus := store.Archived
|
||||
_, err = ts.Store.UpdateUser(ctx, &store.UpdateUser{
|
||||
ID: user.ID,
|
||||
RowStatus: &archivedStatus,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Try to authenticate
|
||||
authenticator := auth.NewAuthenticator(ts.Store, ts.Secret)
|
||||
_, _, err = authenticator.AuthenticateByRefreshToken(ctx, token)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "archived")
|
||||
})
|
||||
}
|
||||
|
||||
func TestAuthenticatorPAT(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("authenticates valid PAT", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create a test user
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Generate PAT
|
||||
token := auth.GeneratePersonalAccessToken()
|
||||
tokenHash := auth.HashPersonalAccessToken(token)
|
||||
tokenID := util.GenUUID()
|
||||
|
||||
// Store PAT in database
|
||||
patRecord := &storepb.PersonalAccessTokensUserSetting_PersonalAccessToken{
|
||||
TokenId: tokenID,
|
||||
TokenHash: tokenHash,
|
||||
Description: "Test PAT",
|
||||
CreatedAt: timestamppb.Now(),
|
||||
}
|
||||
err = ts.Store.AddUserPersonalAccessToken(ctx, user.ID, patRecord)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Authenticate
|
||||
authenticator := auth.NewAuthenticator(ts.Store, ts.Secret)
|
||||
authenticatedUser, pat, err := authenticator.AuthenticateByPAT(ctx, token)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, authenticatedUser)
|
||||
assert.NotNil(t, pat)
|
||||
assert.Equal(t, user.ID, authenticatedUser.ID)
|
||||
assert.Equal(t, tokenID, pat.TokenId)
|
||||
})
|
||||
|
||||
t.Run("fails with invalid PAT format", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
authenticator := auth.NewAuthenticator(ts.Store, ts.Secret)
|
||||
_, _, err := authenticator.AuthenticateByPAT(ctx, "invalid-token-without-prefix")
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "invalid PAT format")
|
||||
})
|
||||
|
||||
t.Run("fails with non-existent PAT", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Generate a PAT but don't store it
|
||||
token := auth.GeneratePersonalAccessToken()
|
||||
|
||||
authenticator := auth.NewAuthenticator(ts.Store, ts.Secret)
|
||||
_, _, err := authenticator.AuthenticateByPAT(ctx, token)
|
||||
assert.Error(t, err)
|
||||
})
|
||||
|
||||
t.Run("fails with expired PAT", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Generate and store expired PAT
|
||||
token := auth.GeneratePersonalAccessToken()
|
||||
tokenHash := auth.HashPersonalAccessToken(token)
|
||||
tokenID := util.GenUUID()
|
||||
|
||||
expiredPAT := &storepb.PersonalAccessTokensUserSetting_PersonalAccessToken{
|
||||
TokenId: tokenID,
|
||||
TokenHash: tokenHash,
|
||||
Description: "Expired PAT",
|
||||
ExpiresAt: timestamppb.New(time.Now().Add(-1 * time.Hour)), // Expired
|
||||
CreatedAt: timestamppb.Now(),
|
||||
}
|
||||
err = ts.Store.AddUserPersonalAccessToken(ctx, user.ID, expiredPAT)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Try to authenticate
|
||||
authenticator := auth.NewAuthenticator(ts.Store, ts.Secret)
|
||||
_, _, err = authenticator.AuthenticateByPAT(ctx, token)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "expired")
|
||||
})
|
||||
|
||||
t.Run("succeeds with non-expiring PAT", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Generate and store PAT without expiration
|
||||
token := auth.GeneratePersonalAccessToken()
|
||||
tokenHash := auth.HashPersonalAccessToken(token)
|
||||
tokenID := util.GenUUID()
|
||||
|
||||
patRecord := &storepb.PersonalAccessTokensUserSetting_PersonalAccessToken{
|
||||
TokenId: tokenID,
|
||||
TokenHash: tokenHash,
|
||||
Description: "Never-expiring PAT",
|
||||
ExpiresAt: nil, // No expiration
|
||||
CreatedAt: timestamppb.Now(),
|
||||
}
|
||||
err = ts.Store.AddUserPersonalAccessToken(ctx, user.ID, patRecord)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Authenticate
|
||||
authenticator := auth.NewAuthenticator(ts.Store, ts.Secret)
|
||||
authenticatedUser, pat, err := authenticator.AuthenticateByPAT(ctx, token)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, authenticatedUser)
|
||||
assert.NotNil(t, pat)
|
||||
assert.Nil(t, pat.ExpiresAt)
|
||||
})
|
||||
|
||||
t.Run("fails with archived user", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Generate and store PAT
|
||||
token := auth.GeneratePersonalAccessToken()
|
||||
tokenHash := auth.HashPersonalAccessToken(token)
|
||||
tokenID := util.GenUUID()
|
||||
|
||||
patRecord := &storepb.PersonalAccessTokensUserSetting_PersonalAccessToken{
|
||||
TokenId: tokenID,
|
||||
TokenHash: tokenHash,
|
||||
Description: "Test PAT",
|
||||
CreatedAt: timestamppb.Now(),
|
||||
}
|
||||
err = ts.Store.AddUserPersonalAccessToken(ctx, user.ID, patRecord)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Archive the user
|
||||
archivedStatus := store.Archived
|
||||
_, err = ts.Store.UpdateUser(ctx, &store.UpdateUser{
|
||||
ID: user.ID,
|
||||
RowStatus: &archivedStatus,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Try to authenticate
|
||||
authenticator := auth.NewAuthenticator(ts.Store, ts.Secret)
|
||||
_, _, err = authenticator.AuthenticateByPAT(ctx, token)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "archived")
|
||||
})
|
||||
}
|
||||
|
||||
func TestStoreRefreshTokenMethods(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("adds and retrieves refresh token", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
tokenID := util.GenUUID()
|
||||
token := &storepb.RefreshTokensUserSetting_RefreshToken{
|
||||
TokenId: tokenID,
|
||||
ExpiresAt: timestamppb.New(time.Now().Add(30 * 24 * time.Hour)),
|
||||
CreatedAt: timestamppb.Now(),
|
||||
}
|
||||
|
||||
err = ts.Store.AddUserRefreshToken(ctx, user.ID, token)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Retrieve tokens
|
||||
tokens, err := ts.Store.GetUserRefreshTokens(ctx, user.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, tokens, 1)
|
||||
assert.Equal(t, tokenID, tokens[0].TokenId)
|
||||
})
|
||||
|
||||
t.Run("retrieves specific refresh token by ID", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
tokenID := util.GenUUID()
|
||||
token := &storepb.RefreshTokensUserSetting_RefreshToken{
|
||||
TokenId: tokenID,
|
||||
ExpiresAt: timestamppb.New(time.Now().Add(30 * 24 * time.Hour)),
|
||||
CreatedAt: timestamppb.Now(),
|
||||
}
|
||||
|
||||
err = ts.Store.AddUserRefreshToken(ctx, user.ID, token)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Retrieve specific token
|
||||
retrievedToken, err := ts.Store.GetUserRefreshTokenByID(ctx, user.ID, tokenID)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, retrievedToken)
|
||||
assert.Equal(t, tokenID, retrievedToken.TokenId)
|
||||
})
|
||||
|
||||
t.Run("removes refresh token", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
tokenID := util.GenUUID()
|
||||
token := &storepb.RefreshTokensUserSetting_RefreshToken{
|
||||
TokenId: tokenID,
|
||||
ExpiresAt: timestamppb.New(time.Now().Add(30 * 24 * time.Hour)),
|
||||
CreatedAt: timestamppb.Now(),
|
||||
}
|
||||
|
||||
err = ts.Store.AddUserRefreshToken(ctx, user.ID, token)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Remove token
|
||||
err = ts.Store.RemoveUserRefreshToken(ctx, user.ID, tokenID)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Verify removal
|
||||
tokens, err := ts.Store.GetUserRefreshTokens(ctx, user.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, tokens, 0)
|
||||
})
|
||||
|
||||
t.Run("handles multiple refresh tokens", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Add multiple tokens
|
||||
tokenID1 := util.GenUUID()
|
||||
tokenID2 := util.GenUUID()
|
||||
|
||||
token1 := &storepb.RefreshTokensUserSetting_RefreshToken{
|
||||
TokenId: tokenID1,
|
||||
ExpiresAt: timestamppb.New(time.Now().Add(30 * 24 * time.Hour)),
|
||||
CreatedAt: timestamppb.Now(),
|
||||
}
|
||||
token2 := &storepb.RefreshTokensUserSetting_RefreshToken{
|
||||
TokenId: tokenID2,
|
||||
ExpiresAt: timestamppb.New(time.Now().Add(30 * 24 * time.Hour)),
|
||||
CreatedAt: timestamppb.Now(),
|
||||
}
|
||||
|
||||
err = ts.Store.AddUserRefreshToken(ctx, user.ID, token1)
|
||||
require.NoError(t, err)
|
||||
err = ts.Store.AddUserRefreshToken(ctx, user.ID, token2)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Retrieve all tokens
|
||||
tokens, err := ts.Store.GetUserRefreshTokens(ctx, user.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, tokens, 2)
|
||||
|
||||
// Remove one token
|
||||
err = ts.Store.RemoveUserRefreshToken(ctx, user.ID, tokenID1)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Verify only one token remains
|
||||
tokens, err = ts.Store.GetUserRefreshTokens(ctx, user.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, tokens, 1)
|
||||
assert.Equal(t, tokenID2, tokens[0].TokenId)
|
||||
})
|
||||
}
|
||||
|
||||
func TestStorePersonalAccessTokenMethods(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("adds and retrieves PAT", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
token := auth.GeneratePersonalAccessToken()
|
||||
tokenHash := auth.HashPersonalAccessToken(token)
|
||||
tokenID := util.GenUUID()
|
||||
|
||||
pat := &storepb.PersonalAccessTokensUserSetting_PersonalAccessToken{
|
||||
TokenId: tokenID,
|
||||
TokenHash: tokenHash,
|
||||
Description: "Test PAT",
|
||||
CreatedAt: timestamppb.Now(),
|
||||
}
|
||||
|
||||
err = ts.Store.AddUserPersonalAccessToken(ctx, user.ID, pat)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Retrieve PATs
|
||||
pats, err := ts.Store.GetUserPersonalAccessTokens(ctx, user.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, pats, 1)
|
||||
assert.Equal(t, tokenID, pats[0].TokenId)
|
||||
assert.Equal(t, tokenHash, pats[0].TokenHash)
|
||||
})
|
||||
|
||||
t.Run("removes PAT", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
token := auth.GeneratePersonalAccessToken()
|
||||
tokenHash := auth.HashPersonalAccessToken(token)
|
||||
tokenID := util.GenUUID()
|
||||
|
||||
pat := &storepb.PersonalAccessTokensUserSetting_PersonalAccessToken{
|
||||
TokenId: tokenID,
|
||||
TokenHash: tokenHash,
|
||||
Description: "Test PAT",
|
||||
CreatedAt: timestamppb.Now(),
|
||||
}
|
||||
|
||||
err = ts.Store.AddUserPersonalAccessToken(ctx, user.ID, pat)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Remove PAT
|
||||
err = ts.Store.RemoveUserPersonalAccessToken(ctx, user.ID, tokenID)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Verify removal
|
||||
pats, err := ts.Store.GetUserPersonalAccessTokens(ctx, user.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, pats, 0)
|
||||
})
|
||||
|
||||
t.Run("updates PAT last used time", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
token := auth.GeneratePersonalAccessToken()
|
||||
tokenHash := auth.HashPersonalAccessToken(token)
|
||||
tokenID := util.GenUUID()
|
||||
|
||||
pat := &storepb.PersonalAccessTokensUserSetting_PersonalAccessToken{
|
||||
TokenId: tokenID,
|
||||
TokenHash: tokenHash,
|
||||
Description: "Test PAT",
|
||||
CreatedAt: timestamppb.Now(),
|
||||
}
|
||||
|
||||
err = ts.Store.AddUserPersonalAccessToken(ctx, user.ID, pat)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Update last used time
|
||||
lastUsed := timestamppb.Now()
|
||||
err = ts.Store.UpdatePATLastUsed(ctx, user.ID, tokenID, lastUsed)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Verify update
|
||||
pats, err := ts.Store.GetUserPersonalAccessTokens(ctx, user.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, pats, 1)
|
||||
assert.NotNil(t, pats[0].LastUsedAt)
|
||||
})
|
||||
|
||||
t.Run("handles multiple PATs", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Add multiple PATs
|
||||
token1 := auth.GeneratePersonalAccessToken()
|
||||
tokenHash1 := auth.HashPersonalAccessToken(token1)
|
||||
tokenID1 := util.GenUUID()
|
||||
|
||||
token2 := auth.GeneratePersonalAccessToken()
|
||||
tokenHash2 := auth.HashPersonalAccessToken(token2)
|
||||
tokenID2 := util.GenUUID()
|
||||
|
||||
pat1 := &storepb.PersonalAccessTokensUserSetting_PersonalAccessToken{
|
||||
TokenId: tokenID1,
|
||||
TokenHash: tokenHash1,
|
||||
Description: "PAT 1",
|
||||
CreatedAt: timestamppb.Now(),
|
||||
}
|
||||
pat2 := &storepb.PersonalAccessTokensUserSetting_PersonalAccessToken{
|
||||
TokenId: tokenID2,
|
||||
TokenHash: tokenHash2,
|
||||
Description: "PAT 2",
|
||||
CreatedAt: timestamppb.Now(),
|
||||
}
|
||||
|
||||
err = ts.Store.AddUserPersonalAccessToken(ctx, user.ID, pat1)
|
||||
require.NoError(t, err)
|
||||
err = ts.Store.AddUserPersonalAccessToken(ctx, user.ID, pat2)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Retrieve all PATs
|
||||
pats, err := ts.Store.GetUserPersonalAccessTokens(ctx, user.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, pats, 2)
|
||||
|
||||
// Remove one PAT
|
||||
err = ts.Store.RemoveUserPersonalAccessToken(ctx, user.ID, tokenID1)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Verify only one PAT remains
|
||||
pats, err = ts.Store.GetUserPersonalAccessTokens(ctx, user.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, pats, 1)
|
||||
assert.Equal(t, tokenID2, pats[0].TokenId)
|
||||
})
|
||||
|
||||
t.Run("finds user by PAT hash", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
token := auth.GeneratePersonalAccessToken()
|
||||
tokenHash := auth.HashPersonalAccessToken(token)
|
||||
tokenID := util.GenUUID()
|
||||
|
||||
pat := &storepb.PersonalAccessTokensUserSetting_PersonalAccessToken{
|
||||
TokenId: tokenID,
|
||||
TokenHash: tokenHash,
|
||||
Description: "Test PAT",
|
||||
CreatedAt: timestamppb.Now(),
|
||||
}
|
||||
|
||||
err = ts.Store.AddUserPersonalAccessToken(ctx, user.ID, pat)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Find user by PAT hash
|
||||
result, err := ts.Store.GetUserByPATHash(ctx, tokenHash)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, result)
|
||||
assert.Equal(t, user.ID, result.UserID)
|
||||
assert.NotNil(t, result.User)
|
||||
assert.Equal(t, user.Username, result.User.Username)
|
||||
assert.NotNil(t, result.PAT)
|
||||
assert.Equal(t, tokenID, result.PAT.TokenId)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,613 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/protobuf/types/known/fieldmaskpb"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
apiv1 "github.com/usememos/memos/server/router/api/v1"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func TestCreateIdentityProvider(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("CreateIdentityProvider success", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create host user
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
ctx := ts.CreateUserContext(ctx, hostUser.ID)
|
||||
|
||||
// Create OAuth2 identity provider
|
||||
req := &v1pb.CreateIdentityProviderRequest{
|
||||
IdentityProvider: &v1pb.IdentityProvider{
|
||||
Title: "Test OAuth2 Provider",
|
||||
IdentifierFilter: "",
|
||||
Type: v1pb.IdentityProvider_OAUTH2,
|
||||
Config: &v1pb.IdentityProviderConfig{
|
||||
Config: &v1pb.IdentityProviderConfig_Oauth2Config{
|
||||
Oauth2Config: &v1pb.OAuth2Config{
|
||||
ClientId: "test-client-id",
|
||||
ClientSecret: "test-client-secret",
|
||||
AuthUrl: "https://example.com/oauth/authorize",
|
||||
TokenUrl: "https://example.com/oauth/token",
|
||||
UserInfoUrl: "https://example.com/oauth/userinfo",
|
||||
Scopes: []string{"openid", "profile", "email"},
|
||||
FieldMapping: &v1pb.FieldMapping{
|
||||
Identifier: "id",
|
||||
DisplayName: "name",
|
||||
Email: "email",
|
||||
AvatarUrl: "avatar_url",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
resp, err := ts.Service.CreateIdentityProvider(ctx, req)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp)
|
||||
require.Equal(t, "Test OAuth2 Provider", resp.Title)
|
||||
require.Equal(t, v1pb.IdentityProvider_OAUTH2, resp.Type)
|
||||
require.Contains(t, resp.Name, "identity-providers/")
|
||||
require.NotNil(t, resp.Config.GetOauth2Config())
|
||||
require.Equal(t, "test-client-id", resp.Config.GetOauth2Config().ClientId)
|
||||
})
|
||||
|
||||
t.Run("CreateIdentityProvider permission denied for non-host user", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create regular user
|
||||
regularUser, err := ts.CreateRegularUser(ctx, "user")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
ctx := ts.CreateUserContext(ctx, regularUser.ID)
|
||||
|
||||
req := &v1pb.CreateIdentityProviderRequest{
|
||||
IdentityProvider: &v1pb.IdentityProvider{
|
||||
Title: "Test Provider",
|
||||
Type: v1pb.IdentityProvider_OAUTH2,
|
||||
},
|
||||
}
|
||||
|
||||
_, err = ts.Service.CreateIdentityProvider(ctx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "permission denied")
|
||||
})
|
||||
|
||||
t.Run("CreateIdentityProvider unauthenticated", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
req := &v1pb.CreateIdentityProviderRequest{
|
||||
IdentityProvider: &v1pb.IdentityProvider{
|
||||
Title: "Test Provider",
|
||||
Type: v1pb.IdentityProvider_OAUTH2,
|
||||
},
|
||||
}
|
||||
|
||||
_, err := ts.Service.CreateIdentityProvider(ctx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "user not authenticated")
|
||||
})
|
||||
}
|
||||
|
||||
func TestListIdentityProviders(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("ListIdentityProviders empty", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
req := &v1pb.ListIdentityProvidersRequest{}
|
||||
resp, err := ts.Service.ListIdentityProviders(ctx, req)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp)
|
||||
require.Empty(t, resp.IdentityProviders)
|
||||
})
|
||||
|
||||
t.Run("ListIdentityProviders with providers", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create host user
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, hostUser.ID)
|
||||
|
||||
// Create a couple of identity providers
|
||||
createReq1 := &v1pb.CreateIdentityProviderRequest{
|
||||
IdentityProvider: &v1pb.IdentityProvider{
|
||||
Title: "Provider 1",
|
||||
Type: v1pb.IdentityProvider_OAUTH2,
|
||||
Config: &v1pb.IdentityProviderConfig{
|
||||
Config: &v1pb.IdentityProviderConfig_Oauth2Config{
|
||||
Oauth2Config: &v1pb.OAuth2Config{
|
||||
ClientId: "client1",
|
||||
AuthUrl: "https://example1.com/auth",
|
||||
TokenUrl: "https://example1.com/token",
|
||||
UserInfoUrl: "https://example1.com/user",
|
||||
FieldMapping: &v1pb.FieldMapping{
|
||||
Identifier: "id",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
createReq2 := &v1pb.CreateIdentityProviderRequest{
|
||||
IdentityProvider: &v1pb.IdentityProvider{
|
||||
Title: "Provider 2",
|
||||
Type: v1pb.IdentityProvider_OAUTH2,
|
||||
Config: &v1pb.IdentityProviderConfig{
|
||||
Config: &v1pb.IdentityProviderConfig_Oauth2Config{
|
||||
Oauth2Config: &v1pb.OAuth2Config{
|
||||
ClientId: "client2",
|
||||
AuthUrl: "https://example2.com/auth",
|
||||
TokenUrl: "https://example2.com/token",
|
||||
UserInfoUrl: "https://example2.com/user",
|
||||
FieldMapping: &v1pb.FieldMapping{
|
||||
Identifier: "id",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
_, err = ts.Service.CreateIdentityProvider(userCtx, createReq1)
|
||||
require.NoError(t, err)
|
||||
_, err = ts.Service.CreateIdentityProvider(userCtx, createReq2)
|
||||
require.NoError(t, err)
|
||||
|
||||
// List providers
|
||||
listReq := &v1pb.ListIdentityProvidersRequest{}
|
||||
resp, err := ts.Service.ListIdentityProviders(ctx, listReq)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp)
|
||||
require.Len(t, resp.IdentityProviders, 2)
|
||||
|
||||
// Verify response contains expected providers
|
||||
titles := []string{resp.IdentityProviders[0].Title, resp.IdentityProviders[1].Title}
|
||||
require.Contains(t, titles, "Provider 1")
|
||||
require.Contains(t, titles, "Provider 2")
|
||||
})
|
||||
}
|
||||
|
||||
func TestGetIdentityProvider(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("GetIdentityProvider success", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create host user
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, hostUser.ID)
|
||||
|
||||
// Create identity provider
|
||||
createReq := &v1pb.CreateIdentityProviderRequest{
|
||||
IdentityProvider: &v1pb.IdentityProvider{
|
||||
Title: "Test Provider",
|
||||
Type: v1pb.IdentityProvider_OAUTH2,
|
||||
Config: &v1pb.IdentityProviderConfig{
|
||||
Config: &v1pb.IdentityProviderConfig_Oauth2Config{
|
||||
Oauth2Config: &v1pb.OAuth2Config{
|
||||
ClientId: "test-client",
|
||||
ClientSecret: "test-secret",
|
||||
AuthUrl: "https://example.com/auth",
|
||||
TokenUrl: "https://example.com/token",
|
||||
UserInfoUrl: "https://example.com/user",
|
||||
Scopes: []string{"openid", "profile"},
|
||||
FieldMapping: &v1pb.FieldMapping{
|
||||
Identifier: "id",
|
||||
DisplayName: "name",
|
||||
Email: "email",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
created, err := ts.Service.CreateIdentityProvider(userCtx, createReq)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Get identity provider
|
||||
getReq := &v1pb.GetIdentityProviderRequest{
|
||||
Name: created.Name,
|
||||
}
|
||||
|
||||
// ClientSecret is write-only: never returned in responses, even to admins.
|
||||
resp, err := ts.Service.GetIdentityProvider(ctx, getReq)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp)
|
||||
require.Equal(t, created.Name, resp.Name)
|
||||
require.Equal(t, "Test Provider", resp.Title)
|
||||
require.Equal(t, v1pb.IdentityProvider_OAUTH2, resp.Type)
|
||||
require.NotNil(t, resp.Config.GetOauth2Config())
|
||||
require.Equal(t, "test-client", resp.Config.GetOauth2Config().ClientId)
|
||||
require.Empty(t, resp.Config.GetOauth2Config().ClientSecret,
|
||||
"ClientSecret must never be returned in responses")
|
||||
|
||||
// Same for admin: secret is still write-only.
|
||||
respAdmin, err := ts.Service.GetIdentityProvider(userCtx, getReq)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, respAdmin)
|
||||
require.Equal(t, "test-client", respAdmin.Config.GetOauth2Config().ClientId)
|
||||
require.Empty(t, respAdmin.Config.GetOauth2Config().ClientSecret,
|
||||
"ClientSecret must never be returned in responses, even to admins")
|
||||
})
|
||||
|
||||
t.Run("GetIdentityProvider not found", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
req := &v1pb.GetIdentityProviderRequest{
|
||||
Name: "identity-providers/999",
|
||||
}
|
||||
|
||||
_, err := ts.Service.GetIdentityProvider(ctx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "not found")
|
||||
})
|
||||
|
||||
t.Run("GetIdentityProvider invalid name", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
req := &v1pb.GetIdentityProviderRequest{
|
||||
Name: "invalid-name",
|
||||
}
|
||||
|
||||
_, err := ts.Service.GetIdentityProvider(ctx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "invalid identity provider name")
|
||||
})
|
||||
}
|
||||
|
||||
func TestUpdateIdentityProvider(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("UpdateIdentityProvider success", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create host user
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, hostUser.ID)
|
||||
|
||||
// Create identity provider
|
||||
createReq := &v1pb.CreateIdentityProviderRequest{
|
||||
IdentityProvider: &v1pb.IdentityProvider{
|
||||
Title: "Original Provider",
|
||||
IdentifierFilter: "",
|
||||
Type: v1pb.IdentityProvider_OAUTH2,
|
||||
Config: &v1pb.IdentityProviderConfig{
|
||||
Config: &v1pb.IdentityProviderConfig_Oauth2Config{
|
||||
Oauth2Config: &v1pb.OAuth2Config{
|
||||
ClientId: "original-client",
|
||||
AuthUrl: "https://original.com/auth",
|
||||
TokenUrl: "https://original.com/token",
|
||||
UserInfoUrl: "https://original.com/user",
|
||||
FieldMapping: &v1pb.FieldMapping{
|
||||
Identifier: "id",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
created, err := ts.Service.CreateIdentityProvider(userCtx, createReq)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Update identity provider
|
||||
updateReq := &v1pb.UpdateIdentityProviderRequest{
|
||||
IdentityProvider: &v1pb.IdentityProvider{
|
||||
Name: created.Name,
|
||||
Title: "Updated Provider",
|
||||
IdentifierFilter: "test@example.com",
|
||||
Type: v1pb.IdentityProvider_OAUTH2,
|
||||
Config: &v1pb.IdentityProviderConfig{
|
||||
Config: &v1pb.IdentityProviderConfig_Oauth2Config{
|
||||
Oauth2Config: &v1pb.OAuth2Config{
|
||||
ClientId: "updated-client",
|
||||
ClientSecret: "updated-secret",
|
||||
AuthUrl: "https://updated.com/auth",
|
||||
TokenUrl: "https://updated.com/token",
|
||||
UserInfoUrl: "https://updated.com/user",
|
||||
Scopes: []string{"openid", "profile", "email"},
|
||||
FieldMapping: &v1pb.FieldMapping{
|
||||
Identifier: "sub",
|
||||
DisplayName: "given_name",
|
||||
Email: "email",
|
||||
AvatarUrl: "picture",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
UpdateMask: &fieldmaskpb.FieldMask{
|
||||
Paths: []string{"title", "identifier_filter", "config"},
|
||||
},
|
||||
}
|
||||
|
||||
updated, err := ts.Service.UpdateIdentityProvider(userCtx, updateReq)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, updated)
|
||||
require.Equal(t, "Updated Provider", updated.Title)
|
||||
require.Equal(t, "test@example.com", updated.IdentifierFilter)
|
||||
require.Equal(t, "updated-client", updated.Config.GetOauth2Config().ClientId)
|
||||
})
|
||||
|
||||
t.Run("UpdateIdentityProvider empty ClientSecret preserves existing credential", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, hostUser.ID)
|
||||
|
||||
// Create IDP with a real secret.
|
||||
created, err := ts.Service.CreateIdentityProvider(userCtx, &v1pb.CreateIdentityProviderRequest{
|
||||
IdentityProvider: &v1pb.IdentityProvider{
|
||||
Title: "Preserve Secret Test",
|
||||
Type: v1pb.IdentityProvider_OAUTH2,
|
||||
Config: &v1pb.IdentityProviderConfig{
|
||||
Config: &v1pb.IdentityProviderConfig_Oauth2Config{
|
||||
Oauth2Config: &v1pb.OAuth2Config{
|
||||
ClientId: "cid",
|
||||
ClientSecret: "original-secret",
|
||||
AuthUrl: "https://ex.com/auth",
|
||||
TokenUrl: "https://ex.com/token",
|
||||
UserInfoUrl: "https://ex.com/user",
|
||||
FieldMapping: &v1pb.FieldMapping{Identifier: "id"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Update with empty ClientSecret (simulating UI that doesn't resend the secret).
|
||||
_, err = ts.Service.UpdateIdentityProvider(userCtx, &v1pb.UpdateIdentityProviderRequest{
|
||||
IdentityProvider: &v1pb.IdentityProvider{
|
||||
Name: created.Name,
|
||||
Title: "Updated Title",
|
||||
Type: v1pb.IdentityProvider_OAUTH2,
|
||||
Config: &v1pb.IdentityProviderConfig{
|
||||
Config: &v1pb.IdentityProviderConfig_Oauth2Config{
|
||||
Oauth2Config: &v1pb.OAuth2Config{
|
||||
ClientId: "cid",
|
||||
ClientSecret: "", // empty = preserve existing
|
||||
AuthUrl: "https://ex.com/auth",
|
||||
TokenUrl: "https://ex.com/token",
|
||||
UserInfoUrl: "https://ex.com/user",
|
||||
FieldMapping: &v1pb.FieldMapping{Identifier: "id"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
UpdateMask: &fieldmaskpb.FieldMask{Paths: []string{"title", "config"}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Verify the stored secret was preserved by reading from store directly.
|
||||
uid, _ := apiv1.ExtractIdentityProviderUIDFromName(created.Name)
|
||||
stored, err := ts.Store.GetIdentityProvider(ctx, &store.FindIdentityProvider{UID: &uid})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "original-secret", stored.Config.GetOauth2Config().ClientSecret,
|
||||
"existing ClientSecret must be preserved when an empty value is sent")
|
||||
require.Equal(t, "Updated Title", stored.Name)
|
||||
})
|
||||
|
||||
t.Run("UpdateIdentityProvider missing update mask", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create host user
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, hostUser.ID)
|
||||
|
||||
req := &v1pb.UpdateIdentityProviderRequest{
|
||||
IdentityProvider: &v1pb.IdentityProvider{
|
||||
Name: "identity-providers/1",
|
||||
Title: "Updated Provider",
|
||||
},
|
||||
}
|
||||
|
||||
_, err = ts.Service.UpdateIdentityProvider(userCtx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "update_mask is required")
|
||||
})
|
||||
|
||||
t.Run("UpdateIdentityProvider invalid name", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create host user
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, hostUser.ID)
|
||||
|
||||
req := &v1pb.UpdateIdentityProviderRequest{
|
||||
IdentityProvider: &v1pb.IdentityProvider{
|
||||
Name: "invalid-name",
|
||||
Title: "Updated Provider",
|
||||
},
|
||||
UpdateMask: &fieldmaskpb.FieldMask{
|
||||
Paths: []string{"title"},
|
||||
},
|
||||
}
|
||||
|
||||
_, err = ts.Service.UpdateIdentityProvider(userCtx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "invalid identity provider name")
|
||||
})
|
||||
}
|
||||
|
||||
func TestDeleteIdentityProvider(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("DeleteIdentityProvider success", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create host user
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, hostUser.ID)
|
||||
|
||||
// Create identity provider
|
||||
createReq := &v1pb.CreateIdentityProviderRequest{
|
||||
IdentityProvider: &v1pb.IdentityProvider{
|
||||
Title: "Provider to Delete",
|
||||
Type: v1pb.IdentityProvider_OAUTH2,
|
||||
Config: &v1pb.IdentityProviderConfig{
|
||||
Config: &v1pb.IdentityProviderConfig_Oauth2Config{
|
||||
Oauth2Config: &v1pb.OAuth2Config{
|
||||
ClientId: "client-to-delete",
|
||||
AuthUrl: "https://example.com/auth",
|
||||
TokenUrl: "https://example.com/token",
|
||||
UserInfoUrl: "https://example.com/user",
|
||||
FieldMapping: &v1pb.FieldMapping{
|
||||
Identifier: "id",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
created, err := ts.Service.CreateIdentityProvider(userCtx, createReq)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Delete identity provider
|
||||
deleteReq := &v1pb.DeleteIdentityProviderRequest{
|
||||
Name: created.Name,
|
||||
}
|
||||
|
||||
_, err = ts.Service.DeleteIdentityProvider(userCtx, deleteReq)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Verify deletion
|
||||
getReq := &v1pb.GetIdentityProviderRequest{
|
||||
Name: created.Name,
|
||||
}
|
||||
|
||||
_, err = ts.Service.GetIdentityProvider(ctx, getReq)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "not found")
|
||||
})
|
||||
|
||||
t.Run("DeleteIdentityProvider invalid name", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create host user
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, hostUser.ID)
|
||||
|
||||
req := &v1pb.DeleteIdentityProviderRequest{
|
||||
Name: "invalid-name",
|
||||
}
|
||||
|
||||
_, err = ts.Service.DeleteIdentityProvider(userCtx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "invalid identity provider name")
|
||||
})
|
||||
|
||||
t.Run("DeleteIdentityProvider not found", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create host user
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, hostUser.ID)
|
||||
|
||||
req := &v1pb.DeleteIdentityProviderRequest{
|
||||
Name: "identity-providers/999",
|
||||
}
|
||||
|
||||
_, err = ts.Service.DeleteIdentityProvider(userCtx, req)
|
||||
require.Error(t, err)
|
||||
// Note: Delete might succeed even if item doesn't exist, depending on store implementation
|
||||
})
|
||||
}
|
||||
|
||||
func TestIdentityProviderPermissions(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("Only host users can create identity providers", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create regular user
|
||||
regularUser, err := ts.CreateRegularUser(ctx, "regularuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, regularUser.ID)
|
||||
|
||||
req := &v1pb.CreateIdentityProviderRequest{
|
||||
IdentityProvider: &v1pb.IdentityProvider{
|
||||
Title: "Test Provider",
|
||||
Type: v1pb.IdentityProvider_OAUTH2,
|
||||
},
|
||||
}
|
||||
|
||||
_, err = ts.Service.CreateIdentityProvider(userCtx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "permission denied")
|
||||
})
|
||||
|
||||
t.Run("Authentication required", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
req := &v1pb.CreateIdentityProviderRequest{
|
||||
IdentityProvider: &v1pb.IdentityProvider{
|
||||
Title: "Test Provider",
|
||||
Type: v1pb.IdentityProvider_OAUTH2,
|
||||
},
|
||||
}
|
||||
|
||||
_, err := ts.Service.CreateIdentityProvider(ctx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "user not authenticated")
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,54 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
)
|
||||
|
||||
func TestInstanceAdminRetrieval(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("Instance becomes initialized after first admin user is created", func(t *testing.T) {
|
||||
// Create test service
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Verify instance is not initialized initially
|
||||
profile1, err := ts.Service.GetInstanceProfile(ctx, &v1pb.GetInstanceProfileRequest{})
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, profile1.Admin, "Instance should not be initialized before first admin user")
|
||||
|
||||
// Create the first admin user
|
||||
user, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, user)
|
||||
|
||||
// Verify instance is now initialized
|
||||
profile2, err := ts.Service.GetInstanceProfile(ctx, &v1pb.GetInstanceProfileRequest{})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, profile2.Admin, "Instance should be initialized after first admin user is created")
|
||||
require.Equal(t, user.Username, profile2.Admin.Username)
|
||||
})
|
||||
|
||||
t.Run("Admin retrieval is cached by Store layer", func(t *testing.T) {
|
||||
// Create test service
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create admin user
|
||||
user, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Multiple calls should return consistent admin user (from cache)
|
||||
for i := 0; i < 5; i++ {
|
||||
profile, err := ts.Service.GetInstanceProfile(ctx, &v1pb.GetInstanceProfileRequest{})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, profile.Admin)
|
||||
require.Equal(t, user.Username, profile.Admin.Username)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,945 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
colorpb "google.golang.org/genproto/googleapis/type/color"
|
||||
"google.golang.org/protobuf/types/known/fieldmaskpb"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
)
|
||||
|
||||
func TestGetInstanceProfile(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("GetInstanceProfile returns instance profile", func(t *testing.T) {
|
||||
// Create test service for this specific test
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Call GetInstanceProfile directly
|
||||
req := &v1pb.GetInstanceProfileRequest{}
|
||||
resp, err := ts.Service.GetInstanceProfile(ctx, req)
|
||||
|
||||
// Verify response
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp)
|
||||
|
||||
// Verify the response contains expected data
|
||||
require.Equal(t, "test-1.0.0", resp.Version)
|
||||
require.Equal(t, "test-commit", resp.Commit)
|
||||
require.True(t, resp.Demo)
|
||||
require.Equal(t, "http://localhost:8080", resp.InstanceUrl)
|
||||
|
||||
// Instance should not be initialized since no admin users are created
|
||||
require.Nil(t, resp.Admin)
|
||||
})
|
||||
|
||||
t.Run("GetInstanceProfile with initialized instance", func(t *testing.T) {
|
||||
// Create test service for this specific test
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create a host user in the store
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, hostUser)
|
||||
|
||||
// Call GetInstanceProfile directly
|
||||
req := &v1pb.GetInstanceProfileRequest{}
|
||||
resp, err := ts.Service.GetInstanceProfile(ctx, req)
|
||||
|
||||
// Verify response
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp)
|
||||
|
||||
// Verify the response contains expected data with initialized flag
|
||||
require.Equal(t, "test-1.0.0", resp.Version)
|
||||
require.Equal(t, "test-commit", resp.Commit)
|
||||
require.True(t, resp.Demo)
|
||||
require.Equal(t, "http://localhost:8080", resp.InstanceUrl)
|
||||
|
||||
// Instance should be initialized since an admin user exists
|
||||
require.NotNil(t, resp.Admin)
|
||||
require.Equal(t, hostUser.Username, resp.Admin.Username)
|
||||
})
|
||||
}
|
||||
|
||||
func TestGetInstanceProfile_Concurrency(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("Concurrent access to service", func(t *testing.T) {
|
||||
// Create test service for this specific test
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create a host user
|
||||
_, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Make concurrent requests
|
||||
numGoroutines := 10
|
||||
results := make(chan *v1pb.InstanceProfile, numGoroutines)
|
||||
errors := make(chan error, numGoroutines)
|
||||
|
||||
for i := 0; i < numGoroutines; i++ {
|
||||
go func() {
|
||||
req := &v1pb.GetInstanceProfileRequest{}
|
||||
resp, err := ts.Service.GetInstanceProfile(ctx, req)
|
||||
if err != nil {
|
||||
errors <- err
|
||||
return
|
||||
}
|
||||
results <- resp
|
||||
}()
|
||||
}
|
||||
|
||||
// Collect all results
|
||||
for i := 0; i < numGoroutines; i++ {
|
||||
select {
|
||||
case err := <-errors:
|
||||
t.Fatalf("Goroutine returned error: %v", err)
|
||||
case resp := <-results:
|
||||
require.NotNil(t, resp)
|
||||
require.Equal(t, "test-1.0.0", resp.Version)
|
||||
require.Equal(t, "test-commit", resp.Commit)
|
||||
require.True(t, resp.Demo)
|
||||
require.Equal(t, "http://localhost:8080", resp.InstanceUrl)
|
||||
require.NotNil(t, resp.Admin)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestGetInstanceSetting(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("GetInstanceSetting - general setting", func(t *testing.T) {
|
||||
// Create test service for this specific test
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Call GetInstanceSetting for general setting
|
||||
req := &v1pb.GetInstanceSettingRequest{
|
||||
Name: "instance/settings/GENERAL",
|
||||
}
|
||||
resp, err := ts.Service.GetInstanceSetting(ctx, req)
|
||||
|
||||
// Verify response
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp)
|
||||
require.Equal(t, "instance/settings/GENERAL", resp.Name)
|
||||
|
||||
// The general setting should have a general_setting field
|
||||
generalSetting := resp.GetGeneralSetting()
|
||||
require.NotNil(t, generalSetting)
|
||||
|
||||
// General setting should have default values
|
||||
require.False(t, generalSetting.DisallowUserRegistration)
|
||||
require.False(t, generalSetting.DisallowPasswordAuth)
|
||||
require.Empty(t, generalSetting.AdditionalScript)
|
||||
})
|
||||
|
||||
t.Run("GetInstanceSetting - storage setting", func(t *testing.T) {
|
||||
// Create test service for this specific test
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create a host user for storage setting access
|
||||
hostUser, err := ts.CreateHostUser(ctx, "testhost")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Add user to context
|
||||
userCtx := ts.CreateUserContext(ctx, hostUser.ID)
|
||||
|
||||
// Call GetInstanceSetting for storage setting
|
||||
req := &v1pb.GetInstanceSettingRequest{
|
||||
Name: "instance/settings/STORAGE",
|
||||
}
|
||||
resp, err := ts.Service.GetInstanceSetting(userCtx, req)
|
||||
|
||||
// Verify response
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp)
|
||||
require.Equal(t, "instance/settings/STORAGE", resp.Name)
|
||||
|
||||
// The storage setting should have a storage_setting field
|
||||
storageSetting := resp.GetStorageSetting()
|
||||
require.NotNil(t, storageSetting)
|
||||
})
|
||||
|
||||
t.Run("GetInstanceSetting - memo related setting", func(t *testing.T) {
|
||||
// Create test service for this specific test
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Call GetInstanceSetting for memo related setting
|
||||
req := &v1pb.GetInstanceSettingRequest{
|
||||
Name: "instance/settings/MEMO_RELATED",
|
||||
}
|
||||
resp, err := ts.Service.GetInstanceSetting(ctx, req)
|
||||
|
||||
// Verify response
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp)
|
||||
require.Equal(t, "instance/settings/MEMO_RELATED", resp.Name)
|
||||
|
||||
// The memo related setting should have a memo_related_setting field
|
||||
memoRelatedSetting := resp.GetMemoRelatedSetting()
|
||||
require.NotNil(t, memoRelatedSetting)
|
||||
})
|
||||
|
||||
t.Run("GetInstanceSetting - tags setting", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
req := &v1pb.GetInstanceSettingRequest{
|
||||
Name: "instance/settings/TAGS",
|
||||
}
|
||||
resp, err := ts.Service.GetInstanceSetting(ctx, req)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp)
|
||||
require.Equal(t, "instance/settings/TAGS", resp.Name)
|
||||
require.NotNil(t, resp.GetTagsSetting())
|
||||
require.Empty(t, resp.GetTagsSetting().GetTags())
|
||||
})
|
||||
|
||||
t.Run("GetInstanceSetting - notification setting requires admin", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
admin, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
adminCtx := ts.CreateUserContext(ctx, admin.ID)
|
||||
|
||||
regularUser, err := ts.CreateRegularUser(ctx, "user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, regularUser.ID)
|
||||
|
||||
req := &v1pb.GetInstanceSettingRequest{Name: "instance/settings/NOTIFICATION"}
|
||||
|
||||
// Unauthenticated request must be rejected.
|
||||
_, err = ts.Service.GetInstanceSetting(ctx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "not authenticated")
|
||||
|
||||
// Non-admin request must be rejected.
|
||||
_, err = ts.Service.GetInstanceSetting(userCtx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "permission denied")
|
||||
|
||||
// Admin request succeeds and does NOT expose SmtpPassword.
|
||||
resp, err := ts.Service.GetInstanceSetting(adminCtx, req)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp)
|
||||
require.Equal(t, "instance/settings/NOTIFICATION", resp.Name)
|
||||
require.NotNil(t, resp.GetNotificationSetting())
|
||||
require.Empty(t, resp.GetNotificationSetting().GetEmail().GetSmtpPassword(),
|
||||
"SmtpPassword must never be returned in responses")
|
||||
})
|
||||
|
||||
t.Run("GetInstanceSetting - AI setting requires authenticated user", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
admin, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
adminCtx := ts.CreateUserContext(ctx, admin.ID)
|
||||
|
||||
regularUser, err := ts.CreateRegularUser(ctx, "user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, regularUser.ID)
|
||||
|
||||
req := &v1pb.GetInstanceSettingRequest{Name: "instance/settings/AI"}
|
||||
|
||||
_, err = ts.Service.GetInstanceSetting(ctx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "not authenticated")
|
||||
|
||||
resp, err := ts.Service.GetInstanceSetting(userCtx, req)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp.GetAiSetting())
|
||||
require.Empty(t, resp.GetAiSetting().GetProviders())
|
||||
|
||||
resp, err = ts.Service.GetInstanceSetting(adminCtx, req)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp.GetAiSetting())
|
||||
require.Empty(t, resp.GetAiSetting().GetProviders())
|
||||
})
|
||||
|
||||
t.Run("GetInstanceSetting - invalid setting name", func(t *testing.T) {
|
||||
// Create test service for this specific test
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Call GetInstanceSetting with invalid name
|
||||
req := &v1pb.GetInstanceSettingRequest{
|
||||
Name: "invalid/setting/name",
|
||||
}
|
||||
_, err := ts.Service.GetInstanceSetting(ctx, req)
|
||||
|
||||
// Should return an error
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "invalid instance setting name")
|
||||
})
|
||||
}
|
||||
|
||||
func TestBatchGetInstanceSettings(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("BatchGetInstanceSettings - returns settings in request order", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
resp, err := ts.Service.BatchGetInstanceSettings(ctx, &v1pb.BatchGetInstanceSettingsRequest{
|
||||
Names: []string{
|
||||
"instance/settings/TAGS",
|
||||
"instance/settings/GENERAL",
|
||||
"instance/settings/MEMO_RELATED",
|
||||
},
|
||||
})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp)
|
||||
require.Len(t, resp.Settings, 3)
|
||||
require.Equal(t, "instance/settings/TAGS", resp.Settings[0].Name)
|
||||
require.NotNil(t, resp.Settings[0].GetTagsSetting())
|
||||
require.Equal(t, "instance/settings/GENERAL", resp.Settings[1].Name)
|
||||
require.NotNil(t, resp.Settings[1].GetGeneralSetting())
|
||||
require.Equal(t, "instance/settings/MEMO_RELATED", resp.Settings[2].Name)
|
||||
require.NotNil(t, resp.Settings[2].GetMemoRelatedSetting())
|
||||
})
|
||||
|
||||
t.Run("BatchGetInstanceSettings - admin-only setting requires admin", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
regularUser, err := ts.CreateRegularUser(ctx, "batch-user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, regularUser.ID)
|
||||
|
||||
_, err = ts.Service.BatchGetInstanceSettings(userCtx, &v1pb.BatchGetInstanceSettingsRequest{
|
||||
Names: []string{"instance/settings/GENERAL", "instance/settings/NOTIFICATION"},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "permission denied")
|
||||
|
||||
admin, err := ts.CreateHostUser(ctx, "batch-admin")
|
||||
require.NoError(t, err)
|
||||
adminCtx := ts.CreateUserContext(ctx, admin.ID)
|
||||
|
||||
resp, err := ts.Service.BatchGetInstanceSettings(adminCtx, &v1pb.BatchGetInstanceSettingsRequest{
|
||||
Names: []string{"instance/settings/GENERAL", "instance/settings/NOTIFICATION"},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, resp.Settings, 2)
|
||||
require.NotNil(t, resp.Settings[1].GetNotificationSetting())
|
||||
})
|
||||
|
||||
t.Run("BatchGetInstanceSettings - invalid setting name", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
_, err := ts.Service.BatchGetInstanceSettings(ctx, &v1pb.BatchGetInstanceSettingsRequest{
|
||||
Names: []string{"instance/settings/GENERAL", "invalid/setting/name"},
|
||||
})
|
||||
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "invalid instance setting name")
|
||||
})
|
||||
}
|
||||
|
||||
func TestTestInstanceEmailSettingAuthorization(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
admin, err := ts.CreateHostUser(ctx, "email-test-admin")
|
||||
require.NoError(t, err)
|
||||
adminCtx := ts.CreateUserContext(ctx, admin.ID)
|
||||
|
||||
regularUser, err := ts.CreateRegularUser(ctx, "email-test-user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, regularUser.ID)
|
||||
|
||||
req := &v1pb.TestInstanceEmailSettingRequest{}
|
||||
|
||||
_, err = ts.Service.TestInstanceEmailSetting(ctx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "not authenticated")
|
||||
|
||||
_, err = ts.Service.TestInstanceEmailSetting(userCtx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "permission denied")
|
||||
|
||||
_, err = ts.Service.TestInstanceEmailSetting(adminCtx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "invalid notification email setting")
|
||||
}
|
||||
|
||||
func TestTestInstanceEmailSettingRequiresPasswordWhenSMTPIdentityChanges(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
admin, err := ts.CreateHostUser(ctx, "email-test-identity-admin")
|
||||
require.NoError(t, err)
|
||||
adminCtx := ts.CreateUserContext(ctx, admin.ID)
|
||||
|
||||
_, err = ts.Store.UpsertInstanceSetting(ctx, &storepb.InstanceSetting{
|
||||
Key: storepb.InstanceSettingKey_NOTIFICATION,
|
||||
Value: &storepb.InstanceSetting_NotificationSetting{
|
||||
NotificationSetting: &storepb.InstanceNotificationSetting{
|
||||
Email: &storepb.InstanceNotificationSetting_EmailSetting{
|
||||
Enabled: true,
|
||||
SmtpHost: "smtp.example.com",
|
||||
SmtpPort: 587,
|
||||
SmtpUsername: "bot@example.com",
|
||||
SmtpPassword: "stored-password",
|
||||
FromEmail: "bot@example.com",
|
||||
UseTls: true,
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.TestInstanceEmailSetting(adminCtx, &v1pb.TestInstanceEmailSettingRequest{
|
||||
Email: &v1pb.InstanceSetting_NotificationSetting_EmailSetting{
|
||||
Enabled: true,
|
||||
SmtpHost: "attacker.example.com",
|
||||
SmtpPort: 587,
|
||||
SmtpUsername: "bot@example.com",
|
||||
SmtpPassword: "",
|
||||
FromEmail: "bot@example.com",
|
||||
UseTls: true,
|
||||
},
|
||||
RecipientEmail: admin.Email,
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "smtp password is required")
|
||||
}
|
||||
|
||||
func TestUpdateInstanceSetting(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("UpdateInstanceSetting - AI setting requires admin", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
regularUser, err := ts.CreateRegularUser(ctx, "user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, regularUser.ID)
|
||||
|
||||
setting := &v1pb.InstanceSetting{
|
||||
Name: "instance/settings/AI",
|
||||
Value: &v1pb.InstanceSetting_AiSetting{
|
||||
AiSetting: &v1pb.InstanceSetting_AISetting{
|
||||
Providers: []*v1pb.InstanceSetting_AIProviderConfig{
|
||||
{
|
||||
Id: "openai-main",
|
||||
Title: "OpenAI",
|
||||
Type: v1pb.InstanceSetting_OPENAI,
|
||||
ApiKey: "sk-test",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
_, err = ts.Service.UpdateInstanceSetting(ctx, &v1pb.UpdateInstanceSettingRequest{Setting: setting})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "not authenticated")
|
||||
|
||||
_, err = ts.Service.UpdateInstanceSetting(userCtx, &v1pb.UpdateInstanceSettingRequest{Setting: setting})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "permission denied")
|
||||
})
|
||||
|
||||
t.Run("UpdateInstanceSetting - tags setting", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.UpdateInstanceSetting(ts.CreateUserContext(ctx, hostUser.ID), &v1pb.UpdateInstanceSettingRequest{
|
||||
Setting: &v1pb.InstanceSetting{
|
||||
Name: "instance/settings/TAGS",
|
||||
Value: &v1pb.InstanceSetting_TagsSetting_{
|
||||
TagsSetting: &v1pb.InstanceSetting_TagsSetting{
|
||||
Tags: map[string]*v1pb.InstanceSetting_TagMetadata{
|
||||
"bug": {
|
||||
BackgroundColor: &colorpb.Color{
|
||||
Red: 0.9,
|
||||
Green: 0.1,
|
||||
Blue: 0.1,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
UpdateMask: &fieldmaskpb.FieldMask{Paths: []string{"tags"}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp.GetTagsSetting())
|
||||
require.Contains(t, resp.GetTagsSetting().GetTags(), "bug")
|
||||
})
|
||||
|
||||
t.Run("UpdateInstanceSetting - invalid tags color", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.UpdateInstanceSetting(ts.CreateUserContext(ctx, hostUser.ID), &v1pb.UpdateInstanceSettingRequest{
|
||||
Setting: &v1pb.InstanceSetting{
|
||||
Name: "instance/settings/TAGS",
|
||||
Value: &v1pb.InstanceSetting_TagsSetting_{
|
||||
TagsSetting: &v1pb.InstanceSetting_TagsSetting{
|
||||
Tags: map[string]*v1pb.InstanceSetting_TagMetadata{
|
||||
"bug": {
|
||||
BackgroundColor: &colorpb.Color{
|
||||
Red: 1.2,
|
||||
Green: 0.1,
|
||||
Blue: 0.1,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "invalid instance setting")
|
||||
})
|
||||
|
||||
t.Run("UpdateInstanceSetting - tags setting without color", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.UpdateInstanceSetting(ts.CreateUserContext(ctx, hostUser.ID), &v1pb.UpdateInstanceSettingRequest{
|
||||
Setting: &v1pb.InstanceSetting{
|
||||
Name: "instance/settings/TAGS",
|
||||
Value: &v1pb.InstanceSetting_TagsSetting_{
|
||||
TagsSetting: &v1pb.InstanceSetting_TagsSetting{
|
||||
Tags: map[string]*v1pb.InstanceSetting_TagMetadata{
|
||||
"spoiler": {
|
||||
BlurContent: true,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp.GetTagsSetting())
|
||||
require.Contains(t, resp.GetTagsSetting().GetTags(), "spoiler")
|
||||
require.Nil(t, resp.GetTagsSetting().GetTags()["spoiler"].GetBackgroundColor())
|
||||
require.True(t, resp.GetTagsSetting().GetTags()["spoiler"].GetBlurContent())
|
||||
})
|
||||
|
||||
t.Run("UpdateInstanceSetting - notification setting password is write-only", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
adminCtx := ts.CreateUserContext(ctx, hostUser.ID)
|
||||
|
||||
// Save notification setting with a password.
|
||||
resp, err := ts.Service.UpdateInstanceSetting(adminCtx, &v1pb.UpdateInstanceSettingRequest{
|
||||
Setting: &v1pb.InstanceSetting{
|
||||
Name: "instance/settings/NOTIFICATION",
|
||||
Value: &v1pb.InstanceSetting_NotificationSetting_{
|
||||
NotificationSetting: &v1pb.InstanceSetting_NotificationSetting{
|
||||
Email: &v1pb.InstanceSetting_NotificationSetting_EmailSetting{
|
||||
Enabled: true,
|
||||
SmtpHost: "smtp.example.com",
|
||||
SmtpPort: 587,
|
||||
SmtpUsername: "bot@example.com",
|
||||
SmtpPassword: "secret",
|
||||
FromEmail: "bot@example.com",
|
||||
FromName: "Memos Bot",
|
||||
ReplyTo: "support@example.com",
|
||||
UseTls: true,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
UpdateMask: &fieldmaskpb.FieldMask{Paths: []string{"notification_setting"}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.True(t, resp.GetNotificationSetting().GetEmail().GetEnabled())
|
||||
require.Equal(t, "smtp.example.com", resp.GetNotificationSetting().GetEmail().GetSmtpHost())
|
||||
// Password must not be returned even in the update response.
|
||||
require.Empty(t, resp.GetNotificationSetting().GetEmail().GetSmtpPassword(),
|
||||
"SmtpPassword must never be returned in responses")
|
||||
})
|
||||
|
||||
t.Run("UpdateInstanceSetting - empty password preserves existing credential", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
adminCtx := ts.CreateUserContext(ctx, hostUser.ID)
|
||||
|
||||
notificationSetting := &v1pb.InstanceSetting{
|
||||
Name: "instance/settings/NOTIFICATION",
|
||||
Value: &v1pb.InstanceSetting_NotificationSetting_{
|
||||
NotificationSetting: &v1pb.InstanceSetting_NotificationSetting{
|
||||
Email: &v1pb.InstanceSetting_NotificationSetting_EmailSetting{
|
||||
Enabled: true,
|
||||
SmtpHost: "smtp.example.com",
|
||||
SmtpPort: 587,
|
||||
SmtpUsername: "bot@example.com",
|
||||
SmtpPassword: "original-password",
|
||||
FromEmail: "bot@example.com",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
// First save with a real password.
|
||||
_, err = ts.Service.UpdateInstanceSetting(adminCtx, &v1pb.UpdateInstanceSettingRequest{
|
||||
Setting: notificationSetting,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Second update with an empty password (simulating a UI that doesn't re-send the secret).
|
||||
notificationSetting.GetNotificationSetting().GetEmail().SmtpPassword = ""
|
||||
notificationSetting.GetNotificationSetting().GetEmail().FromName = "Updated Bot"
|
||||
_, err = ts.Service.UpdateInstanceSetting(adminCtx, &v1pb.UpdateInstanceSettingRequest{
|
||||
Setting: notificationSetting,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// The stored setting should have preserved the original password.
|
||||
stored, err := ts.Store.GetInstanceNotificationSetting(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "original-password", stored.GetEmail().GetSmtpPassword(),
|
||||
"existing SmtpPassword must be preserved when an empty value is sent")
|
||||
require.Equal(t, "Updated Bot", stored.GetEmail().GetFromName())
|
||||
})
|
||||
|
||||
t.Run("UpdateInstanceSetting - empty password rejected when SMTP identity changes", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
adminCtx := ts.CreateUserContext(ctx, hostUser.ID)
|
||||
|
||||
notificationSetting := &v1pb.InstanceSetting{
|
||||
Name: "instance/settings/NOTIFICATION",
|
||||
Value: &v1pb.InstanceSetting_NotificationSetting_{
|
||||
NotificationSetting: &v1pb.InstanceSetting_NotificationSetting{
|
||||
Email: &v1pb.InstanceSetting_NotificationSetting_EmailSetting{
|
||||
Enabled: true,
|
||||
SmtpHost: "smtp.example.com",
|
||||
SmtpPort: 587,
|
||||
SmtpUsername: "bot@example.com",
|
||||
SmtpPassword: "original-password",
|
||||
FromEmail: "bot@example.com",
|
||||
UseTls: true,
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
_, err = ts.Service.UpdateInstanceSetting(adminCtx, &v1pb.UpdateInstanceSettingRequest{
|
||||
Setting: notificationSetting,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
notificationSetting.GetNotificationSetting().GetEmail().SmtpPassword = ""
|
||||
notificationSetting.GetNotificationSetting().GetEmail().SmtpHost = "smtp2.example.com"
|
||||
_, err = ts.Service.UpdateInstanceSetting(adminCtx, &v1pb.UpdateInstanceSettingRequest{
|
||||
Setting: notificationSetting,
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "smtp password is required")
|
||||
})
|
||||
|
||||
t.Run("UpdateInstanceSetting - S3 secret is write-only and preserved on empty", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
adminCtx := ts.CreateUserContext(ctx, hostUser.ID)
|
||||
|
||||
// Save storage setting with a real secret.
|
||||
_, err = ts.Service.UpdateInstanceSetting(adminCtx, &v1pb.UpdateInstanceSettingRequest{
|
||||
Setting: &v1pb.InstanceSetting{
|
||||
Name: "instance/settings/STORAGE",
|
||||
Value: &v1pb.InstanceSetting_StorageSetting_{
|
||||
StorageSetting: &v1pb.InstanceSetting_StorageSetting{
|
||||
S3Config: &v1pb.InstanceSetting_StorageSetting_S3Config{
|
||||
AccessKeyId: "AKID",
|
||||
AccessKeySecret: "super-secret",
|
||||
Endpoint: "s3.example.com",
|
||||
Region: "us-east-1",
|
||||
Bucket: "memos",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Read back: secret must not be returned.
|
||||
resp, err := ts.Service.GetInstanceSetting(adminCtx, &v1pb.GetInstanceSettingRequest{
|
||||
Name: "instance/settings/STORAGE",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, resp.GetStorageSetting().GetS3Config().GetAccessKeySecret(),
|
||||
"AccessKeySecret must never be returned in responses")
|
||||
|
||||
// Update with empty secret; original must be preserved in the store.
|
||||
_, err = ts.Service.UpdateInstanceSetting(adminCtx, &v1pb.UpdateInstanceSettingRequest{
|
||||
Setting: &v1pb.InstanceSetting{
|
||||
Name: "instance/settings/STORAGE",
|
||||
Value: &v1pb.InstanceSetting_StorageSetting_{
|
||||
StorageSetting: &v1pb.InstanceSetting_StorageSetting{
|
||||
S3Config: &v1pb.InstanceSetting_StorageSetting_S3Config{
|
||||
AccessKeyId: "AKID",
|
||||
AccessKeySecret: "", // omitted / not changed
|
||||
Endpoint: "s3-v2.example.com",
|
||||
Region: "us-east-1",
|
||||
Bucket: "memos",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
stored, err := ts.Store.GetInstanceStorageSetting(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "super-secret", stored.GetS3Config().GetAccessKeySecret(),
|
||||
"existing AccessKeySecret must be preserved when an empty value is sent")
|
||||
require.Equal(t, "s3-v2.example.com", stored.GetS3Config().GetEndpoint())
|
||||
})
|
||||
|
||||
t.Run("UpdateInstanceSetting - AI provider keys are write-only and preserved on empty", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
adminCtx := ts.CreateUserContext(ctx, hostUser.ID)
|
||||
|
||||
_, err = ts.Service.UpdateInstanceSetting(adminCtx, &v1pb.UpdateInstanceSettingRequest{
|
||||
Setting: &v1pb.InstanceSetting{
|
||||
Name: "instance/settings/AI",
|
||||
Value: &v1pb.InstanceSetting_AiSetting{
|
||||
AiSetting: &v1pb.InstanceSetting_AISetting{
|
||||
Providers: []*v1pb.InstanceSetting_AIProviderConfig{
|
||||
{
|
||||
Id: "openai-main",
|
||||
Title: "OpenAI",
|
||||
Type: v1pb.InstanceSetting_OPENAI,
|
||||
ApiKey: "sk-original",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.GetInstanceSetting(adminCtx, &v1pb.GetInstanceSettingRequest{
|
||||
Name: "instance/settings/AI",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, resp.GetAiSetting().GetProviders(), 1)
|
||||
provider := resp.GetAiSetting().GetProviders()[0]
|
||||
require.Empty(t, provider.GetApiKey(), "AI provider API key must never be returned in responses")
|
||||
require.True(t, provider.GetApiKeySet())
|
||||
require.Equal(t, "sk-o...inal", provider.GetApiKeyHint())
|
||||
require.Equal(t, "https://api.openai.com/v1", provider.GetEndpoint())
|
||||
|
||||
_, err = ts.Service.UpdateInstanceSetting(adminCtx, &v1pb.UpdateInstanceSettingRequest{
|
||||
Setting: &v1pb.InstanceSetting{
|
||||
Name: "instance/settings/AI",
|
||||
Value: &v1pb.InstanceSetting_AiSetting{
|
||||
AiSetting: &v1pb.InstanceSetting_AISetting{
|
||||
Providers: []*v1pb.InstanceSetting_AIProviderConfig{
|
||||
{
|
||||
Id: "openai-main",
|
||||
Title: "OpenAI primary",
|
||||
Type: v1pb.InstanceSetting_OPENAI,
|
||||
ApiKey: "",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
stored, err := ts.Store.GetInstanceAISetting(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, stored.GetProviders(), 1)
|
||||
require.Equal(t, "sk-original", stored.GetProviders()[0].GetApiKey(),
|
||||
"existing AI provider API key must be preserved when an empty value is sent")
|
||||
require.Equal(t, "OpenAI primary", stored.GetProviders()[0].GetTitle())
|
||||
})
|
||||
|
||||
t.Run("UpdateInstanceSetting - transcription provider_id must reference an existing provider", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
adminCtx := ts.CreateUserContext(ctx, hostUser.ID)
|
||||
|
||||
_, err = ts.Service.UpdateInstanceSetting(adminCtx, &v1pb.UpdateInstanceSettingRequest{
|
||||
Setting: &v1pb.InstanceSetting{
|
||||
Name: "instance/settings/AI",
|
||||
Value: &v1pb.InstanceSetting_AiSetting{
|
||||
AiSetting: &v1pb.InstanceSetting_AISetting{
|
||||
Providers: []*v1pb.InstanceSetting_AIProviderConfig{
|
||||
{
|
||||
Id: "openai-main",
|
||||
Title: "OpenAI",
|
||||
Type: v1pb.InstanceSetting_OPENAI,
|
||||
ApiKey: "sk-test",
|
||||
},
|
||||
},
|
||||
Transcription: &v1pb.InstanceSetting_TranscriptionConfig{
|
||||
ProviderId: "does-not-exist",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "transcription provider_id")
|
||||
})
|
||||
|
||||
t.Run("UpdateInstanceSetting - transcription strings are length-capped", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
adminCtx := ts.CreateUserContext(ctx, hostUser.ID)
|
||||
|
||||
base := &v1pb.InstanceSetting{
|
||||
Name: "instance/settings/AI",
|
||||
Value: &v1pb.InstanceSetting_AiSetting{
|
||||
AiSetting: &v1pb.InstanceSetting_AISetting{
|
||||
Providers: []*v1pb.InstanceSetting_AIProviderConfig{
|
||||
{
|
||||
Id: "openai-main",
|
||||
Title: "OpenAI",
|
||||
Type: v1pb.InstanceSetting_OPENAI,
|
||||
ApiKey: "sk-test",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
oversizedModel := strings.Repeat("a", 257)
|
||||
base.GetAiSetting().Transcription = &v1pb.InstanceSetting_TranscriptionConfig{
|
||||
ProviderId: "openai-main",
|
||||
Model: oversizedModel,
|
||||
}
|
||||
_, err = ts.Service.UpdateInstanceSetting(adminCtx, &v1pb.UpdateInstanceSettingRequest{Setting: base})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "transcription model")
|
||||
|
||||
oversizedLanguage := strings.Repeat("a", 33)
|
||||
base.GetAiSetting().Transcription = &v1pb.InstanceSetting_TranscriptionConfig{
|
||||
ProviderId: "openai-main",
|
||||
Language: oversizedLanguage,
|
||||
}
|
||||
_, err = ts.Service.UpdateInstanceSetting(adminCtx, &v1pb.UpdateInstanceSettingRequest{Setting: base})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "transcription language")
|
||||
|
||||
oversizedPrompt := strings.Repeat("a", 4097)
|
||||
base.GetAiSetting().Transcription = &v1pb.InstanceSetting_TranscriptionConfig{
|
||||
ProviderId: "openai-main",
|
||||
Prompt: oversizedPrompt,
|
||||
}
|
||||
_, err = ts.Service.UpdateInstanceSetting(adminCtx, &v1pb.UpdateInstanceSettingRequest{Setting: base})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "transcription prompt")
|
||||
})
|
||||
|
||||
t.Run("UpdateInstanceSetting - transcription is preserved when omitted on update", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
adminCtx := ts.CreateUserContext(ctx, hostUser.ID)
|
||||
|
||||
_, err = ts.Service.UpdateInstanceSetting(adminCtx, &v1pb.UpdateInstanceSettingRequest{
|
||||
Setting: &v1pb.InstanceSetting{
|
||||
Name: "instance/settings/AI",
|
||||
Value: &v1pb.InstanceSetting_AiSetting{
|
||||
AiSetting: &v1pb.InstanceSetting_AISetting{
|
||||
Providers: []*v1pb.InstanceSetting_AIProviderConfig{
|
||||
{
|
||||
Id: "openai-main",
|
||||
Title: "OpenAI",
|
||||
Type: v1pb.InstanceSetting_OPENAI,
|
||||
ApiKey: "sk-test",
|
||||
},
|
||||
},
|
||||
Transcription: &v1pb.InstanceSetting_TranscriptionConfig{
|
||||
ProviderId: "openai-main",
|
||||
Model: "whisper-1",
|
||||
Language: "en",
|
||||
Prompt: "names: Alice",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.UpdateInstanceSetting(adminCtx, &v1pb.UpdateInstanceSettingRequest{
|
||||
Setting: &v1pb.InstanceSetting{
|
||||
Name: "instance/settings/AI",
|
||||
Value: &v1pb.InstanceSetting_AiSetting{
|
||||
AiSetting: &v1pb.InstanceSetting_AISetting{
|
||||
Providers: []*v1pb.InstanceSetting_AIProviderConfig{
|
||||
{
|
||||
Id: "openai-main",
|
||||
Title: "OpenAI",
|
||||
Type: v1pb.InstanceSetting_OPENAI,
|
||||
ApiKey: "",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
stored, err := ts.Store.GetInstanceAISetting(ctx)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, stored.GetTranscription())
|
||||
require.Equal(t, "openai-main", stored.GetTranscription().GetProviderId())
|
||||
require.Equal(t, "whisper-1", stored.GetTranscription().GetModel())
|
||||
require.Equal(t, "en", stored.GetTranscription().GetLanguage())
|
||||
require.Equal(t, "names: Alice", stored.GetTranscription().GetPrompt())
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,72 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
)
|
||||
|
||||
func TestGetInstanceStats_HappyPath(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
admin, err := ts.CreateHostUser(ctx, "admin1")
|
||||
require.NoError(t, err)
|
||||
adminCtx := ts.CreateUserContext(ctx, admin.ID)
|
||||
|
||||
resp, err := ts.Service.GetInstanceStats(adminCtx, &v1pb.GetInstanceStatsRequest{})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp)
|
||||
|
||||
require.NotNil(t, resp.Database)
|
||||
require.Equal(t, "sqlite", resp.Database.Driver)
|
||||
require.Greater(t, resp.Database.SizeBytes, int64(0))
|
||||
|
||||
require.GreaterOrEqual(t, resp.LocalStorageBytes, int64(0))
|
||||
}
|
||||
|
||||
func TestGetInstanceStats_NonAdminDenied(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Need an admin to exist (otherwise instance is uninitialized).
|
||||
admin, err := ts.CreateHostUser(ctx, "admin1")
|
||||
require.NoError(t, err)
|
||||
_ = admin
|
||||
|
||||
regular, err := ts.CreateRegularUser(ctx, "alice")
|
||||
require.NoError(t, err)
|
||||
regularCtx := ts.CreateUserContext(ctx, regular.ID)
|
||||
|
||||
_, err = ts.Service.GetInstanceStats(regularCtx, &v1pb.GetInstanceStatsRequest{})
|
||||
require.Error(t, err)
|
||||
st, ok := status.FromError(err)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, codes.PermissionDenied, st.Code())
|
||||
}
|
||||
|
||||
func TestGetInstanceStats_Cache(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
admin, err := ts.CreateHostUser(ctx, "admin1")
|
||||
require.NoError(t, err)
|
||||
adminCtx := ts.CreateUserContext(ctx, admin.ID)
|
||||
|
||||
first, err := ts.Service.GetInstanceStats(adminCtx, &v1pb.GetInstanceStatsRequest{})
|
||||
require.NoError(t, err)
|
||||
|
||||
second, err := ts.Service.GetInstanceStats(adminCtx, &v1pb.GetInstanceStatsRequest{})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Cache hit: same pointer (the cache returns the stored *InstanceStats directly).
|
||||
require.Same(t, first, second)
|
||||
}
|
||||
@@ -0,0 +1,346 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
apiv1 "github.com/usememos/memos/proto/gen/api/v1"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func TestSetMemoAttachments(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("SetMemoAttachments success by memo owner", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create user
|
||||
user, err := ts.CreateRegularUser(ctx, "user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
// Create memo
|
||||
memo, err := ts.Service.CreateMemo(userCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "Test memo",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memo)
|
||||
|
||||
// Create attachment
|
||||
attachment, err := ts.Service.CreateAttachment(userCtx, &apiv1.CreateAttachmentRequest{
|
||||
Attachment: &apiv1.Attachment{
|
||||
Filename: "test.txt",
|
||||
Size: 5,
|
||||
Type: "text/plain",
|
||||
Content: []byte("hello"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, attachment)
|
||||
|
||||
// Set memo attachments - should succeed
|
||||
_, err = ts.Service.SetMemoAttachments(userCtx, &apiv1.SetMemoAttachmentsRequest{
|
||||
Name: memo.Name,
|
||||
Attachments: []*apiv1.Attachment{
|
||||
{Name: attachment.Name},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
})
|
||||
|
||||
t.Run("SetMemoAttachments success by host user", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create regular user
|
||||
regularUser, err := ts.CreateRegularUser(ctx, "user")
|
||||
require.NoError(t, err)
|
||||
regularUserCtx := ts.CreateUserContext(ctx, regularUser.ID)
|
||||
|
||||
// Create host user
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
hostCtx := ts.CreateUserContext(ctx, hostUser.ID)
|
||||
|
||||
// Create memo by regular user
|
||||
memo, err := ts.Service.CreateMemo(regularUserCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "Test memo",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memo)
|
||||
|
||||
// Host user can modify attachments - should succeed
|
||||
_, err = ts.Service.SetMemoAttachments(hostCtx, &apiv1.SetMemoAttachmentsRequest{
|
||||
Name: memo.Name,
|
||||
Attachments: []*apiv1.Attachment{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
})
|
||||
|
||||
t.Run("SetMemoAttachments permission denied for non-owner", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create user1
|
||||
user1, err := ts.CreateRegularUser(ctx, "user1")
|
||||
require.NoError(t, err)
|
||||
user1Ctx := ts.CreateUserContext(ctx, user1.ID)
|
||||
|
||||
// Create user2
|
||||
user2, err := ts.CreateRegularUser(ctx, "user2")
|
||||
require.NoError(t, err)
|
||||
user2Ctx := ts.CreateUserContext(ctx, user2.ID)
|
||||
|
||||
// Create memo by user1
|
||||
memo, err := ts.Service.CreateMemo(user1Ctx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "Test memo",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memo)
|
||||
|
||||
// User2 tries to modify attachments - should fail
|
||||
_, err = ts.Service.SetMemoAttachments(user2Ctx, &apiv1.SetMemoAttachmentsRequest{
|
||||
Name: memo.Name,
|
||||
Attachments: []*apiv1.Attachment{},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "permission denied")
|
||||
})
|
||||
|
||||
t.Run("SetMemoAttachments unauthenticated", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create user
|
||||
user, err := ts.CreateRegularUser(ctx, "user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
// Create memo
|
||||
memo, err := ts.Service.CreateMemo(userCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "Test memo",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memo)
|
||||
|
||||
// Unauthenticated user tries to modify attachments - should fail
|
||||
_, err = ts.Service.SetMemoAttachments(ctx, &apiv1.SetMemoAttachmentsRequest{
|
||||
Name: memo.Name,
|
||||
Attachments: []*apiv1.Attachment{},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "not authenticated")
|
||||
})
|
||||
|
||||
t.Run("SetMemoAttachments memo not found", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create user
|
||||
user, err := ts.CreateRegularUser(ctx, "user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
// Try to set attachments on non-existent memo - should fail
|
||||
_, err = ts.Service.SetMemoAttachments(userCtx, &apiv1.SetMemoAttachmentsRequest{
|
||||
Name: "memos/nonexistent-uid-12345",
|
||||
Attachments: []*apiv1.Attachment{},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "not found")
|
||||
})
|
||||
|
||||
t.Run("SetMemoAttachments removes incomplete live photo groups", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "live_group_user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
still, err := ts.Service.CreateAttachment(userCtx, &apiv1.CreateAttachmentRequest{
|
||||
Attachment: &apiv1.Attachment{
|
||||
Filename: "live.heic",
|
||||
Type: "image/heic",
|
||||
Content: []byte("still"),
|
||||
MotionMedia: &apiv1.MotionMedia{
|
||||
Family: apiv1.MotionMediaFamily_APPLE_LIVE_PHOTO,
|
||||
Role: apiv1.MotionMediaRole_STILL,
|
||||
GroupId: "memo-live-group",
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
video, err := ts.Service.CreateAttachment(userCtx, &apiv1.CreateAttachmentRequest{
|
||||
Attachment: &apiv1.Attachment{
|
||||
Filename: "live.mov",
|
||||
Type: "video/quicktime",
|
||||
Content: []byte("video"),
|
||||
MotionMedia: &apiv1.MotionMedia{
|
||||
Family: apiv1.MotionMediaFamily_APPLE_LIVE_PHOTO,
|
||||
Role: apiv1.MotionMediaRole_VIDEO,
|
||||
GroupId: "memo-live-group",
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
memo, err := ts.Service.CreateMemo(userCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "memo with live photo",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
Attachments: []*apiv1.Attachment{
|
||||
{Name: still.Name},
|
||||
{Name: video.Name},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.SetMemoAttachments(userCtx, &apiv1.SetMemoAttachmentsRequest{
|
||||
Name: memo.Name,
|
||||
Attachments: []*apiv1.Attachment{
|
||||
{Name: still.Name},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
response, err := ts.Service.ListMemoAttachments(userCtx, &apiv1.ListMemoAttachmentsRequest{Name: memo.Name})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, response.Attachments, 0)
|
||||
})
|
||||
|
||||
t.Run("SetMemoAttachments denies attaching another user's attachment", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
victim, err := ts.CreateRegularUser(ctx, "attachment_victim")
|
||||
require.NoError(t, err)
|
||||
attacker, err := ts.CreateRegularUser(ctx, "attachment_attacker")
|
||||
require.NoError(t, err)
|
||||
victimCtx := ts.CreateUserContext(ctx, victim.ID)
|
||||
attackerCtx := ts.CreateUserContext(ctx, attacker.ID)
|
||||
|
||||
victimAttachment, err := ts.Service.CreateAttachment(victimCtx, &apiv1.CreateAttachmentRequest{
|
||||
Attachment: &apiv1.Attachment{
|
||||
Filename: "secret.txt",
|
||||
Size: 6,
|
||||
Type: "text/plain",
|
||||
Content: []byte("secret"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
victimMemo, err := ts.Service.CreateMemo(victimCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "victim protected memo",
|
||||
Visibility: apiv1.Visibility_PROTECTED,
|
||||
Attachments: []*apiv1.Attachment{
|
||||
{Name: victimAttachment.Name},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
attackerMemo, err := ts.Service.CreateMemo(attackerCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "attacker public memo",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.SetMemoAttachments(attackerCtx, &apiv1.SetMemoAttachmentsRequest{
|
||||
Name: attackerMemo.Name,
|
||||
Attachments: []*apiv1.Attachment{
|
||||
{Name: victimAttachment.Name},
|
||||
},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "cannot attach another user's attachment")
|
||||
|
||||
victimAttachments, err := ts.Service.ListMemoAttachments(victimCtx, &apiv1.ListMemoAttachmentsRequest{Name: victimMemo.Name})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, victimAttachments.Attachments, 1)
|
||||
require.Equal(t, victimAttachment.Name, victimAttachments.Attachments[0].Name)
|
||||
|
||||
attackerAttachments, err := ts.Service.ListMemoAttachments(attackerCtx, &apiv1.ListMemoAttachmentsRequest{Name: attackerMemo.Name})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, attackerAttachments.Attachments)
|
||||
})
|
||||
|
||||
t.Run("SetMemoAttachments denies removing another user's attached attachment", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
victim, err := ts.CreateRegularUser(ctx, "remove_victim")
|
||||
require.NoError(t, err)
|
||||
attacker, err := ts.CreateRegularUser(ctx, "remove_attacker")
|
||||
require.NoError(t, err)
|
||||
victimCtx := ts.CreateUserContext(ctx, victim.ID)
|
||||
attackerCtx := ts.CreateUserContext(ctx, attacker.ID)
|
||||
|
||||
victimAttachment, err := ts.Service.CreateAttachment(victimCtx, &apiv1.CreateAttachmentRequest{
|
||||
Attachment: &apiv1.Attachment{
|
||||
Filename: "kept.txt",
|
||||
Size: 4,
|
||||
Type: "text/plain",
|
||||
Content: []byte("kept"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
attackerMemo, err := ts.Service.CreateMemo(attackerCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "contaminated memo",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
attachmentUID := strings.TrimPrefix(victimAttachment.Name, "attachments/")
|
||||
attachment, err := ts.Store.GetAttachment(ctx, &store.FindAttachment{UID: &attachmentUID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, attachment)
|
||||
|
||||
memoUID := strings.TrimPrefix(attackerMemo.Name, "memos/")
|
||||
memo, err := ts.Store.GetMemo(ctx, &store.FindMemo{UID: &memoUID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memo)
|
||||
|
||||
err = ts.Store.UpdateAttachment(ctx, &store.UpdateAttachment{
|
||||
ID: attachment.ID,
|
||||
MemoID: &memo.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.SetMemoAttachments(attackerCtx, &apiv1.SetMemoAttachmentsRequest{
|
||||
Name: attackerMemo.Name,
|
||||
Attachments: []*apiv1.Attachment{},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "cannot remove another user's attachment")
|
||||
|
||||
attachmentAfter, err := ts.Store.GetAttachment(ctx, &store.FindAttachment{ID: &attachment.ID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, attachmentAfter)
|
||||
require.NotNil(t, attachmentAfter.MemoID)
|
||||
require.Equal(t, memo.ID, *attachmentAfter.MemoID)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,169 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
apiv1 "github.com/usememos/memos/proto/gen/api/v1"
|
||||
)
|
||||
|
||||
func TestSetMemoRelations(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("SetMemoRelations success by memo owner", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create user
|
||||
user, err := ts.CreateRegularUser(ctx, "user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
// Create memo1
|
||||
memo1, err := ts.Service.CreateMemo(userCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "Test memo 1",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memo1)
|
||||
|
||||
// Create memo2
|
||||
memo2, err := ts.Service.CreateMemo(userCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "Test memo 2",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memo2)
|
||||
|
||||
// Set memo relations - should succeed
|
||||
_, err = ts.Service.SetMemoRelations(userCtx, &apiv1.SetMemoRelationsRequest{
|
||||
Name: memo1.Name,
|
||||
Relations: []*apiv1.MemoRelation{
|
||||
{
|
||||
RelatedMemo: &apiv1.MemoRelation_Memo{
|
||||
Name: memo2.Name,
|
||||
},
|
||||
Type: apiv1.MemoRelation_REFERENCE,
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
})
|
||||
|
||||
t.Run("SetMemoRelations success by host user", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create regular user
|
||||
regularUser, err := ts.CreateRegularUser(ctx, "user")
|
||||
require.NoError(t, err)
|
||||
regularUserCtx := ts.CreateUserContext(ctx, regularUser.ID)
|
||||
|
||||
// Create host user
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
hostCtx := ts.CreateUserContext(ctx, hostUser.ID)
|
||||
|
||||
// Create memo by regular user
|
||||
memo, err := ts.Service.CreateMemo(regularUserCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "Test memo",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memo)
|
||||
|
||||
// Host user can modify relations - should succeed
|
||||
_, err = ts.Service.SetMemoRelations(hostCtx, &apiv1.SetMemoRelationsRequest{
|
||||
Name: memo.Name,
|
||||
Relations: []*apiv1.MemoRelation{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
})
|
||||
|
||||
t.Run("SetMemoRelations permission denied for non-owner", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create user1
|
||||
user1, err := ts.CreateRegularUser(ctx, "user1")
|
||||
require.NoError(t, err)
|
||||
user1Ctx := ts.CreateUserContext(ctx, user1.ID)
|
||||
|
||||
// Create user2
|
||||
user2, err := ts.CreateRegularUser(ctx, "user2")
|
||||
require.NoError(t, err)
|
||||
user2Ctx := ts.CreateUserContext(ctx, user2.ID)
|
||||
|
||||
// Create memo by user1
|
||||
memo, err := ts.Service.CreateMemo(user1Ctx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "Test memo",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memo)
|
||||
|
||||
// User2 tries to modify relations - should fail
|
||||
_, err = ts.Service.SetMemoRelations(user2Ctx, &apiv1.SetMemoRelationsRequest{
|
||||
Name: memo.Name,
|
||||
Relations: []*apiv1.MemoRelation{},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "permission denied")
|
||||
})
|
||||
|
||||
t.Run("SetMemoRelations unauthenticated", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create user
|
||||
user, err := ts.CreateRegularUser(ctx, "user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
// Create memo
|
||||
memo, err := ts.Service.CreateMemo(userCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "Test memo",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memo)
|
||||
|
||||
// Unauthenticated user tries to modify relations - should fail
|
||||
_, err = ts.Service.SetMemoRelations(ctx, &apiv1.SetMemoRelationsRequest{
|
||||
Name: memo.Name,
|
||||
Relations: []*apiv1.MemoRelation{},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "not authenticated")
|
||||
})
|
||||
|
||||
t.Run("SetMemoRelations memo not found", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create user
|
||||
user, err := ts.CreateRegularUser(ctx, "user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
// Try to set relations on non-existent memo - should fail
|
||||
_, err = ts.Service.SetMemoRelations(userCtx, &apiv1.SetMemoRelationsRequest{
|
||||
Name: "memos/nonexistent-uid-12345",
|
||||
Relations: []*apiv1.MemoRelation{},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "not found")
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,282 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/usememos/memos/internal/profile"
|
||||
"github.com/usememos/memos/internal/version"
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
"github.com/usememos/memos/server/auth"
|
||||
apiv1 "github.com/usememos/memos/server/router/api/v1"
|
||||
"github.com/usememos/memos/store"
|
||||
"github.com/usememos/memos/store/db"
|
||||
)
|
||||
|
||||
const (
|
||||
benchmarkTopLevelMemoCount = 5000
|
||||
benchmarkPageSize = 16
|
||||
)
|
||||
|
||||
type benchmarkService struct {
|
||||
*TestService
|
||||
hostUser *store.User
|
||||
authenticatedCtx context.Context
|
||||
publicCtx context.Context
|
||||
pageTenToken string
|
||||
commentParentName string
|
||||
}
|
||||
|
||||
func newBenchmarkService(tb testing.TB) *benchmarkService {
|
||||
tb.Helper()
|
||||
|
||||
ctx := context.Background()
|
||||
testService := newTestingServiceForTB(tb)
|
||||
|
||||
hostUser, err := testService.CreateHostUser(ctx, "bench-host")
|
||||
if err != nil {
|
||||
tb.Fatalf("failed to create host user: %v", err)
|
||||
}
|
||||
|
||||
commentParentName, err := seedListMemosBenchmarkData(ctx, testService.Store, hostUser)
|
||||
if err != nil {
|
||||
tb.Fatalf("failed to seed benchmark data: %v", err)
|
||||
}
|
||||
|
||||
authenticatedCtx := context.WithValue(context.Background(), auth.UserIDContextKey, hostUser.ID)
|
||||
pageTenToken, err := getListMemosPageToken(authenticatedCtx, testService.Service, 10, benchmarkPageSize)
|
||||
if err != nil {
|
||||
tb.Fatalf("failed to build page token: %v", err)
|
||||
}
|
||||
|
||||
return &benchmarkService{
|
||||
TestService: testService,
|
||||
hostUser: hostUser,
|
||||
authenticatedCtx: authenticatedCtx,
|
||||
publicCtx: context.Background(),
|
||||
pageTenToken: pageTenToken,
|
||||
commentParentName: commentParentName,
|
||||
}
|
||||
}
|
||||
|
||||
func newTestingServiceForTB(tb testing.TB) *TestService {
|
||||
tb.Helper()
|
||||
|
||||
ctx := context.Background()
|
||||
dataDir := tb.TempDir()
|
||||
testProfile := getBenchmarkProfile(dataDir)
|
||||
dbDriver, err := db.NewDBDriver(testProfile)
|
||||
if err != nil {
|
||||
tb.Fatalf("failed to create db driver: %v", err)
|
||||
}
|
||||
|
||||
testStore := store.New(dbDriver, testProfile)
|
||||
if err := testStore.Migrate(ctx); err != nil {
|
||||
tb.Fatalf("failed to migrate db: %v", err)
|
||||
}
|
||||
tb.Cleanup(func() {
|
||||
testStore.Close()
|
||||
})
|
||||
|
||||
service := newServiceWithProfile(testProfile, testStore)
|
||||
return &TestService{
|
||||
Service: service,
|
||||
Store: testStore,
|
||||
Profile: testProfile,
|
||||
Secret: service.Secret,
|
||||
}
|
||||
}
|
||||
|
||||
func getBenchmarkProfile(dataDir string) *profile.Profile {
|
||||
return &profile.Profile{
|
||||
Demo: true,
|
||||
Version: version.GetCurrentVersion(),
|
||||
InstanceURL: "http://localhost:8080",
|
||||
Driver: "sqlite",
|
||||
DSN: filepath.Join(dataDir, "bench.db"),
|
||||
Data: dataDir,
|
||||
}
|
||||
}
|
||||
|
||||
func newServiceWithProfile(testProfile *profile.Profile, testStore *store.Store) *apiv1.APIV1Service {
|
||||
service := apiv1.NewAPIV1Service("bench-secret", testProfile, testStore)
|
||||
return service
|
||||
}
|
||||
|
||||
func seedListMemosBenchmarkData(ctx context.Context, stores *store.Store, hostUser *store.User) (string, error) {
|
||||
topLevelMemos := make([]*store.Memo, 0, benchmarkTopLevelMemoCount)
|
||||
commentParentName := ""
|
||||
|
||||
for i := 0; i < benchmarkTopLevelMemoCount; i++ {
|
||||
visibility := store.Private
|
||||
if i%4 == 0 {
|
||||
visibility = store.Public
|
||||
}
|
||||
memo, err := stores.CreateMemo(ctx, &store.Memo{
|
||||
UID: fmt.Sprintf("memo-%06d", i),
|
||||
CreatorID: hostUser.ID,
|
||||
Content: benchmarkMemoContent(i),
|
||||
Visibility: visibility,
|
||||
})
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
topLevelMemos = append(topLevelMemos, memo)
|
||||
|
||||
if i%3 == 0 {
|
||||
if _, err := stores.CreateAttachment(ctx, &store.Attachment{
|
||||
UID: fmt.Sprintf("att-%06d", i),
|
||||
CreatorID: hostUser.ID,
|
||||
Filename: fmt.Sprintf("memo-%06d.png", i),
|
||||
Type: "image/png",
|
||||
Size: 2048,
|
||||
MemoID: &memo.ID,
|
||||
}); err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
|
||||
if i%5 == 0 {
|
||||
if _, err := stores.UpsertReaction(ctx, &store.Reaction{
|
||||
CreatorID: hostUser.ID,
|
||||
ContentID: "memos/" + memo.UID,
|
||||
ReactionType: "thumbs-up",
|
||||
}); err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for i, memo := range topLevelMemos {
|
||||
if i+1 < len(topLevelMemos) && i%4 == 0 {
|
||||
if _, err := stores.UpsertMemoRelation(ctx, &store.MemoRelation{
|
||||
MemoID: memo.ID,
|
||||
RelatedMemoID: topLevelMemos[i+1].ID,
|
||||
Type: store.MemoRelationReference,
|
||||
}); err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
|
||||
if i%6 == 0 {
|
||||
commentMemo, err := stores.CreateMemo(ctx, &store.Memo{
|
||||
UID: fmt.Sprintf("comment-%06d", i),
|
||||
CreatorID: hostUser.ID,
|
||||
Content: fmt.Sprintf("Comment for memo %06d", i),
|
||||
Visibility: store.Private,
|
||||
})
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if _, err := stores.UpsertMemoRelation(ctx, &store.MemoRelation{
|
||||
MemoID: commentMemo.ID,
|
||||
RelatedMemoID: memo.ID,
|
||||
Type: store.MemoRelationComment,
|
||||
}); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if commentParentName == "" {
|
||||
commentParentName = "memos/" + memo.UID
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return commentParentName, nil
|
||||
}
|
||||
|
||||
func benchmarkMemoContent(i int) string {
|
||||
return fmt.Sprintf("# Bench Memo %06d\n\nThis is benchmark memo %06d with enough content to exercise snippet generation.\n\n- task one\n- task two\n", i, i)
|
||||
}
|
||||
|
||||
func getListMemosPageToken(ctx context.Context, service *apiv1.APIV1Service, page int, pageSize int32) (string, error) {
|
||||
pageToken := ""
|
||||
for range page - 1 {
|
||||
resp, err := service.ListMemos(ctx, &v1pb.ListMemosRequest{
|
||||
PageSize: pageSize,
|
||||
PageToken: pageToken,
|
||||
})
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
pageToken = resp.NextPageToken
|
||||
if pageToken == "" {
|
||||
break
|
||||
}
|
||||
}
|
||||
return pageToken, nil
|
||||
}
|
||||
|
||||
func BenchmarkListMemos(b *testing.B) {
|
||||
bench := newBenchmarkService(b)
|
||||
|
||||
b.Run("authenticated_first_page", func(b *testing.B) {
|
||||
req := &v1pb.ListMemosRequest{PageSize: benchmarkPageSize}
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
resp, err := bench.Service.ListMemos(bench.authenticatedCtx, req)
|
||||
if err != nil {
|
||||
b.Fatalf("ListMemos failed: %v", err)
|
||||
}
|
||||
if len(resp.Memos) == 0 {
|
||||
b.Fatal("expected memos in authenticated benchmark response")
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
b.Run("authenticated_page_ten", func(b *testing.B) {
|
||||
req := &v1pb.ListMemosRequest{PageSize: benchmarkPageSize, PageToken: bench.pageTenToken}
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
resp, err := bench.Service.ListMemos(bench.authenticatedCtx, req)
|
||||
if err != nil {
|
||||
b.Fatalf("ListMemos failed: %v", err)
|
||||
}
|
||||
if len(resp.Memos) == 0 {
|
||||
b.Fatal("expected memos in paged benchmark response")
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
b.Run("public_first_page", func(b *testing.B) {
|
||||
req := &v1pb.ListMemosRequest{PageSize: benchmarkPageSize}
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
resp, err := bench.Service.ListMemos(bench.publicCtx, req)
|
||||
if err != nil {
|
||||
b.Fatalf("ListMemos failed: %v", err)
|
||||
}
|
||||
if len(resp.Memos) == 0 {
|
||||
b.Fatal("expected memos in public benchmark response")
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func BenchmarkListMemoCommentsPreview(b *testing.B) {
|
||||
bench := newBenchmarkService(b)
|
||||
if bench.commentParentName == "" {
|
||||
b.Fatal("expected seeded memo with comments")
|
||||
}
|
||||
|
||||
req := &v1pb.ListMemoCommentsRequest{
|
||||
Name: bench.commentParentName,
|
||||
PageSize: 3,
|
||||
}
|
||||
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
resp, err := bench.Service.ListMemoComments(bench.authenticatedCtx, req)
|
||||
if err != nil {
|
||||
b.Fatalf("ListMemoComments failed: %v", err)
|
||||
}
|
||||
if len(resp.Memos) == 0 {
|
||||
b.Fatal("expected comments in benchmark response")
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,732 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"slices"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
"google.golang.org/protobuf/types/known/fieldmaskpb"
|
||||
"google.golang.org/protobuf/types/known/timestamppb"
|
||||
|
||||
apiv1 "github.com/usememos/memos/proto/gen/api/v1"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func TestListMemos(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create userOne
|
||||
userOne, err := ts.CreateRegularUser(ctx, "test-user-1")
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, userOne)
|
||||
|
||||
// Create userOne context
|
||||
userOneCtx := ts.CreateUserContext(ctx, userOne.ID)
|
||||
|
||||
// Create userTwo
|
||||
userTwo, err := ts.CreateRegularUser(ctx, "test-user-2")
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, userTwo)
|
||||
|
||||
// Create userTwo context
|
||||
userTwoCtx := ts.CreateUserContext(ctx, userTwo.ID)
|
||||
|
||||
// Create attachmentOne by userOne
|
||||
attachmentOne, err := ts.Service.CreateAttachment(userOneCtx, &apiv1.CreateAttachmentRequest{
|
||||
Attachment: &apiv1.Attachment{
|
||||
Name: "",
|
||||
Filename: "hello.txt",
|
||||
Size: 5,
|
||||
Type: "text/plain",
|
||||
Content: []byte{
|
||||
104, 101, 108, 108, 111,
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, attachmentOne)
|
||||
|
||||
// Create attachmentTwo by userOne
|
||||
attachmentTwo, err := ts.Service.CreateAttachment(userOneCtx, &apiv1.CreateAttachmentRequest{
|
||||
Attachment: &apiv1.Attachment{
|
||||
Name: "",
|
||||
Filename: "world.txt",
|
||||
Size: 5,
|
||||
Type: "text/plain",
|
||||
Content: []byte{
|
||||
119, 111, 114, 108, 100,
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, attachmentTwo)
|
||||
|
||||
// Create memoOne with two attachments by userOne
|
||||
memoOne, err := ts.Service.CreateMemo(userOneCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "Hellooo, any words after this sentence won't be in the snippet. This is the next sentence. And I also have two attachments.",
|
||||
Visibility: apiv1.Visibility_PROTECTED,
|
||||
Attachments: []*apiv1.Attachment{
|
||||
&apiv1.Attachment{
|
||||
Name: attachmentOne.Name,
|
||||
},
|
||||
&apiv1.Attachment{
|
||||
Name: attachmentTwo.Name,
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memoOne)
|
||||
|
||||
// Create memoTwo by userTwo referencing memoOne
|
||||
memoTwo, err := ts.Service.CreateMemo(userTwoCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "This is a memo reminding you to check the attachment attached to memoOne. I have referenced the memo below.⬇️",
|
||||
Visibility: apiv1.Visibility_PROTECTED,
|
||||
Relations: []*apiv1.MemoRelation{
|
||||
&apiv1.MemoRelation{
|
||||
RelatedMemo: &apiv1.MemoRelation_Memo{
|
||||
Name: memoOne.Name,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memoTwo)
|
||||
|
||||
// Create memoThree by userOne
|
||||
memoThree, err := ts.Service.CreateMemo(userOneCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "This is a very popular memo. I have 2 reactions!",
|
||||
Visibility: apiv1.Visibility_PROTECTED,
|
||||
},
|
||||
})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memoThree)
|
||||
|
||||
// Create reaction from userOne on memoThree
|
||||
reactionOne, err := ts.Service.UpsertMemoReaction(userOneCtx, &apiv1.UpsertMemoReactionRequest{
|
||||
Name: memoThree.Name,
|
||||
Reaction: &apiv1.Reaction{
|
||||
ContentId: memoThree.Name,
|
||||
ReactionType: "❤️",
|
||||
},
|
||||
})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, reactionOne)
|
||||
|
||||
// Create reaction from userTwo on memoThree
|
||||
reactionTwo, err := ts.Service.UpsertMemoReaction(userTwoCtx, &apiv1.UpsertMemoReactionRequest{
|
||||
Name: memoThree.Name,
|
||||
Reaction: &apiv1.Reaction{
|
||||
ContentId: memoThree.Name,
|
||||
ReactionType: "👍",
|
||||
},
|
||||
})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, reactionTwo)
|
||||
|
||||
memos, err := ts.Service.ListMemos(userOneCtx, &apiv1.ListMemosRequest{PageSize: 10})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memos)
|
||||
require.Equal(t, 3, len(memos.Memos))
|
||||
|
||||
// ///////////////
|
||||
// VERIFY MEMO ONE
|
||||
// ///////////////
|
||||
memoOneResIdx := slices.IndexFunc(memos.Memos, func(m *apiv1.Memo) bool { return m.GetName() == memoOne.GetName() })
|
||||
require.NotEqual(t, memoOneResIdx, -1)
|
||||
|
||||
memoOneRes := memos.Memos[memoOneResIdx]
|
||||
require.NotNil(t, memoOneRes)
|
||||
|
||||
require.Equal(t, fmt.Sprintf("users/%s", userOne.Username), memoOneRes.GetCreator())
|
||||
require.Equal(t, apiv1.Visibility_PROTECTED, memoOneRes.GetVisibility())
|
||||
require.Equal(t, memoOne.Content, memoOneRes.GetContent())
|
||||
require.Equal(t, memoOne.Content[:64]+"...", memoOneRes.GetSnippet(), "memoOne's content is snipped past the 64 char limit")
|
||||
require.Len(t, memoOneRes.Attachments, 2)
|
||||
require.Len(t, memoOneRes.Relations, 1)
|
||||
require.Empty(t, memoOneRes.Reactions)
|
||||
|
||||
// verify memoOne's attachments
|
||||
// attachment one
|
||||
attachmentOneResIdx := slices.IndexFunc(memoOneRes.Attachments, func(a *apiv1.Attachment) bool { return a.GetName() == attachmentOne.GetName() })
|
||||
require.NotEqual(t, attachmentOneResIdx, -1)
|
||||
|
||||
attachmentOneRes := memoOneRes.Attachments[attachmentOneResIdx]
|
||||
require.NotNil(t, attachmentOneRes)
|
||||
|
||||
require.Equal(t, attachmentOne.GetName(), attachmentOneRes.GetName())
|
||||
require.Equal(t, attachmentOne.GetContent(), attachmentOneRes.GetContent())
|
||||
|
||||
// attachment two
|
||||
attachmentTwoResIdx := slices.IndexFunc(memoOneRes.Attachments, func(a *apiv1.Attachment) bool { return a.GetName() == attachmentTwo.GetName() })
|
||||
require.NotEqual(t, attachmentTwoResIdx, -1)
|
||||
|
||||
attachmentTwoRes := memoOneRes.Attachments[attachmentTwoResIdx]
|
||||
require.NotNil(t, attachmentTwoRes)
|
||||
require.Equal(t, attachmentTwo.GetName(), attachmentTwoRes.GetName())
|
||||
|
||||
require.Equal(t, attachmentTwo.GetName(), attachmentTwoRes.GetName())
|
||||
require.Equal(t, attachmentTwo.GetContent(), attachmentTwoRes.GetContent())
|
||||
|
||||
// verify memoOne's relations
|
||||
require.Len(t, memoOneRes.Relations, 1)
|
||||
memoOneExpectedRelation := &apiv1.MemoRelation{
|
||||
Memo: &apiv1.MemoRelation_Memo{Name: memoTwo.GetName()},
|
||||
RelatedMemo: &apiv1.MemoRelation_Memo{Name: memoOne.GetName()},
|
||||
}
|
||||
require.Equal(t, memoOneExpectedRelation.Memo.GetName(), memoOneRes.Relations[0].Memo.GetName())
|
||||
require.Equal(t, memoOneExpectedRelation.RelatedMemo.GetName(), memoOneRes.Relations[0].RelatedMemo.GetName())
|
||||
|
||||
// ///////////////
|
||||
// VERIFY MEMO TWO
|
||||
// ///////////////
|
||||
memoTwoResIdx := slices.IndexFunc(memos.Memos, func(m *apiv1.Memo) bool { return m.GetName() == memoTwo.GetName() })
|
||||
require.NotEqual(t, memoTwoResIdx, -1)
|
||||
|
||||
memoTwoRes := memos.Memos[memoTwoResIdx]
|
||||
require.NotNil(t, memoTwoRes)
|
||||
|
||||
require.Equal(t, fmt.Sprintf("users/%s", userTwo.Username), memoTwoRes.GetCreator())
|
||||
require.Equal(t, apiv1.Visibility_PROTECTED, memoTwoRes.GetVisibility())
|
||||
require.Equal(t, memoTwo.Content, memoTwoRes.GetContent())
|
||||
require.Empty(t, memoTwoRes.Attachments)
|
||||
require.Len(t, memoTwoRes.Relations, 1)
|
||||
require.Empty(t, memoTwoRes.Reactions)
|
||||
|
||||
// verify memoTwo's relations
|
||||
require.Len(t, memoTwoRes.Relations, 1)
|
||||
memoTwoExpectedRelation := &apiv1.MemoRelation{
|
||||
Memo: &apiv1.MemoRelation_Memo{Name: memoTwo.GetName()},
|
||||
RelatedMemo: &apiv1.MemoRelation_Memo{Name: memoOne.GetName()},
|
||||
}
|
||||
require.Equal(t, memoTwoExpectedRelation.Memo.GetName(), memoTwoRes.Relations[0].Memo.GetName())
|
||||
require.Equal(t, memoTwoExpectedRelation.RelatedMemo.GetName(), memoTwoRes.Relations[0].RelatedMemo.GetName())
|
||||
|
||||
// ///////////////
|
||||
// VERIFY MEMO THREE
|
||||
// ///////////////
|
||||
memoThreeResIdx := slices.IndexFunc(memos.Memos, func(m *apiv1.Memo) bool { return m.GetName() == memoThree.GetName() })
|
||||
require.NotEqual(t, memoThreeResIdx, -1)
|
||||
|
||||
memoThreeRes := memos.Memos[memoThreeResIdx]
|
||||
require.NotNil(t, memoThreeRes)
|
||||
|
||||
require.Equal(t, fmt.Sprintf("users/%s", userOne.Username), memoThreeRes.GetCreator())
|
||||
require.Equal(t, apiv1.Visibility_PROTECTED, memoThreeRes.GetVisibility())
|
||||
require.Equal(t, memoThree.Content, memoThreeRes.GetContent())
|
||||
require.Empty(t, memoThreeRes.Attachments)
|
||||
require.Empty(t, memoThreeRes.Relations)
|
||||
require.Len(t, memoThreeRes.Reactions, 2)
|
||||
|
||||
// verify memoThree's reactions
|
||||
require.Len(t, memoThreeRes.Reactions, 2)
|
||||
// userOne's reaction
|
||||
userOneReactionIdx := slices.IndexFunc(memoThreeRes.Reactions, func(r *apiv1.Reaction) bool { return r.GetCreator() == fmt.Sprintf("users/%s", userOne.Username) })
|
||||
require.NotEqual(t, userOneReactionIdx, -1)
|
||||
|
||||
userOneReaction := memoThreeRes.Reactions[userOneReactionIdx]
|
||||
require.NotNil(t, userOneReaction)
|
||||
require.Equal(t, "❤️", userOneReaction.ReactionType)
|
||||
|
||||
// userTwo's reaction
|
||||
userTwoReactionIdx := slices.IndexFunc(memoThreeRes.Reactions, func(r *apiv1.Reaction) bool { return r.GetCreator() == fmt.Sprintf("users/%s", userTwo.Username) })
|
||||
require.NotEqual(t, userTwoReactionIdx, -1)
|
||||
|
||||
userTwoReaction := memoThreeRes.Reactions[userTwoReactionIdx]
|
||||
require.NotNil(t, userTwoReaction)
|
||||
require.Equal(t, "👍", userTwoReaction.ReactionType)
|
||||
}
|
||||
|
||||
func TestListMemosTimeOrderBy(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateHostUser(ctx, "time-order-user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
memoEarlyCreateLateUpdate, err := ts.Service.CreateMemo(userCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "early create late update",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
CreateTime: timestamppb.New(time.Date(2020, 1, 1, 0, 0, 0, 0, time.UTC)),
|
||||
UpdateTime: timestamppb.New(time.Date(2020, 1, 3, 0, 0, 0, 0, time.UTC)),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
memoMiddleCreateEarlyUpdate, err := ts.Service.CreateMemo(userCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "middle create early update",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
CreateTime: timestamppb.New(time.Date(2020, 1, 2, 0, 0, 0, 0, time.UTC)),
|
||||
UpdateTime: timestamppb.New(time.Date(2020, 1, 1, 0, 0, 0, 0, time.UTC)),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
memoLateCreateMiddleUpdate, err := ts.Service.CreateMemo(userCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "late create middle update",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
CreateTime: timestamppb.New(time.Date(2020, 1, 3, 0, 0, 0, 0, time.UTC)),
|
||||
UpdateTime: timestamppb.New(time.Date(2020, 1, 2, 0, 0, 0, 0, time.UTC)),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
orderBy string
|
||||
wantNames []string
|
||||
}{
|
||||
{
|
||||
name: "default create time",
|
||||
orderBy: "",
|
||||
wantNames: []string{
|
||||
memoLateCreateMiddleUpdate.Name,
|
||||
memoMiddleCreateEarlyUpdate.Name,
|
||||
memoEarlyCreateLateUpdate.Name,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "explicit create time",
|
||||
orderBy: "create_time desc",
|
||||
wantNames: []string{
|
||||
memoLateCreateMiddleUpdate.Name,
|
||||
memoMiddleCreateEarlyUpdate.Name,
|
||||
memoEarlyCreateLateUpdate.Name,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "explicit update time",
|
||||
orderBy: "update_time desc",
|
||||
wantNames: []string{
|
||||
memoEarlyCreateLateUpdate.Name,
|
||||
memoLateCreateMiddleUpdate.Name,
|
||||
memoMiddleCreateEarlyUpdate.Name,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "pinned with explicit create time",
|
||||
orderBy: "pinned desc, create_time desc",
|
||||
wantNames: []string{
|
||||
memoLateCreateMiddleUpdate.Name,
|
||||
memoMiddleCreateEarlyUpdate.Name,
|
||||
memoEarlyCreateLateUpdate.Name,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "explicit create time ascending",
|
||||
orderBy: "create_time asc",
|
||||
wantNames: []string{
|
||||
memoEarlyCreateLateUpdate.Name,
|
||||
memoMiddleCreateEarlyUpdate.Name,
|
||||
memoLateCreateMiddleUpdate.Name,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
resp, err := ts.Service.ListMemos(userCtx, &apiv1.ListMemosRequest{
|
||||
PageSize: 10,
|
||||
OrderBy: test.orderBy,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, resp.Memos, len(test.wantNames))
|
||||
|
||||
gotNames := make([]string, 0, len(resp.Memos))
|
||||
for _, memo := range resp.Memos {
|
||||
gotNames = append(gotNames, memo.Name)
|
||||
}
|
||||
require.Equal(t, test.wantNames, gotNames)
|
||||
})
|
||||
}
|
||||
|
||||
_, err = ts.Service.ListMemos(userCtx, &apiv1.ListMemosRequest{
|
||||
PageSize: 10,
|
||||
OrderBy: "display_time desc",
|
||||
})
|
||||
require.Error(t, err)
|
||||
}
|
||||
|
||||
func TestListMemosSkipsReactionsWithMissingCreators(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
owner, err := ts.CreateRegularUser(ctx, "memo-owner")
|
||||
require.NoError(t, err)
|
||||
ownerCtx := ts.CreateUserContext(ctx, owner.ID)
|
||||
|
||||
reactor, err := ts.CreateRegularUser(ctx, "memo-reactor")
|
||||
require.NoError(t, err)
|
||||
reactorCtx := ts.CreateUserContext(ctx, reactor.ID)
|
||||
|
||||
memo, err := ts.Service.CreateMemo(ownerCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "memo with orphan reaction",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.UpsertMemoReaction(reactorCtx, &apiv1.UpsertMemoReactionRequest{
|
||||
Name: memo.Name,
|
||||
Reaction: &apiv1.Reaction{
|
||||
ContentId: memo.Name,
|
||||
ReactionType: "👍",
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Store.DeleteUser(ctx, &store.DeleteUser{ID: reactor.ID})
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.ListMemos(ownerCtx, &apiv1.ListMemosRequest{PageSize: 10})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, resp.Memos, 1)
|
||||
require.Equal(t, memo.Name, resp.Memos[0].Name)
|
||||
require.Empty(t, resp.Memos[0].Reactions)
|
||||
}
|
||||
|
||||
func TestListMemosSkipsMemosWithMissingCreators(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
owner, err := ts.CreateRegularUser(ctx, "memo-visible-owner")
|
||||
require.NoError(t, err)
|
||||
ownerCtx := ts.CreateUserContext(ctx, owner.ID)
|
||||
|
||||
orphanCreator, err := ts.CreateRegularUser(ctx, "memo-orphan-creator")
|
||||
require.NoError(t, err)
|
||||
orphanCtx := ts.CreateUserContext(ctx, orphanCreator.ID)
|
||||
|
||||
ownerMemo, err := ts.Service.CreateMemo(ownerCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "owner memo",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.CreateMemo(orphanCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "orphan memo",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Store.DeleteUser(ctx, &store.DeleteUser{ID: orphanCreator.ID})
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.ListMemos(ownerCtx, &apiv1.ListMemosRequest{PageSize: 10})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, resp.Memos, 1)
|
||||
require.Equal(t, ownerMemo.Name, resp.Memos[0].Name)
|
||||
}
|
||||
|
||||
func TestListMemoCommentsSkipsCommentsWithMissingCreators(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
owner, err := ts.CreateRegularUser(ctx, "comment-owner")
|
||||
require.NoError(t, err)
|
||||
ownerCtx := ts.CreateUserContext(ctx, owner.ID)
|
||||
|
||||
commenter, err := ts.CreateRegularUser(ctx, "comment-orphan")
|
||||
require.NoError(t, err)
|
||||
commenterCtx := ts.CreateUserContext(ctx, commenter.ID)
|
||||
|
||||
memo, err := ts.Service.CreateMemo(ownerCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "memo with comment",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.CreateMemoComment(commenterCtx, &apiv1.CreateMemoCommentRequest{
|
||||
Name: memo.Name,
|
||||
Comment: &apiv1.Memo{
|
||||
Content: "comment to orphan",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Store.DeleteUser(ctx, &store.DeleteUser{ID: commenter.ID})
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.ListMemoComments(ownerCtx, &apiv1.ListMemoCommentsRequest{Name: memo.Name})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, resp.Memos)
|
||||
}
|
||||
|
||||
func TestListMemoCommentsPaginates(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
owner, err := ts.CreateRegularUser(ctx, "comment-page-owner")
|
||||
require.NoError(t, err)
|
||||
ownerCtx := ts.CreateUserContext(ctx, owner.ID)
|
||||
|
||||
memo, err := ts.Service.CreateMemo(ownerCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "memo with paged comments",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
for i := 0; i < 3; i++ {
|
||||
_, err = ts.Service.CreateMemoComment(ownerCtx, &apiv1.CreateMemoCommentRequest{
|
||||
Name: memo.Name,
|
||||
Comment: &apiv1.Memo{
|
||||
Content: fmt.Sprintf("comment %d", i),
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
firstPage, err := ts.Service.ListMemoComments(ownerCtx, &apiv1.ListMemoCommentsRequest{Name: memo.Name, PageSize: 2})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, firstPage.Memos, 2)
|
||||
require.NotEmpty(t, firstPage.NextPageToken)
|
||||
|
||||
secondPage, err := ts.Service.ListMemoComments(ownerCtx, &apiv1.ListMemoCommentsRequest{Name: memo.Name, PageToken: firstPage.NextPageToken})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, secondPage.Memos, 1)
|
||||
require.Empty(t, secondPage.NextPageToken)
|
||||
}
|
||||
|
||||
func TestCreateMemoCommentInheritsParentVisibility(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
owner, err := ts.CreateRegularUser(ctx, "private-comment-owner")
|
||||
require.NoError(t, err)
|
||||
ownerCtx := ts.CreateUserContext(ctx, owner.ID)
|
||||
|
||||
parent, err := ts.Service.CreateMemo(ownerCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "private parent",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
comment, err := ts.Service.CreateMemoComment(ownerCtx, &apiv1.CreateMemoCommentRequest{
|
||||
Name: parent.Name,
|
||||
Comment: &apiv1.Memo{
|
||||
Content: "client requested public comment",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, apiv1.Visibility_PRIVATE, comment.Visibility)
|
||||
|
||||
updatedComment, err := ts.Service.UpdateMemo(ownerCtx, &apiv1.UpdateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Name: comment.Name,
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
UpdateMask: &fieldmaskpb.FieldMask{Paths: []string{"visibility"}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, apiv1.Visibility_PRIVATE, updatedComment.Visibility)
|
||||
|
||||
_, err = ts.Service.GetMemo(ctx, &apiv1.GetMemoRequest{Name: comment.Name})
|
||||
require.Equal(t, codes.Unauthenticated, status.Code(err))
|
||||
}
|
||||
|
||||
func TestGetMemoCommentRequiresParentReadAccess(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
owner, err := ts.CreateRegularUser(ctx, "legacy-comment-owner")
|
||||
require.NoError(t, err)
|
||||
ownerCtx := ts.CreateUserContext(ctx, owner.ID)
|
||||
|
||||
other, err := ts.CreateRegularUser(ctx, "legacy-comment-other")
|
||||
require.NoError(t, err)
|
||||
otherCtx := ts.CreateUserContext(ctx, other.ID)
|
||||
|
||||
parent, err := ts.Service.CreateMemo(ownerCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "private parent for legacy comment",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
legacyComment, err := ts.Store.CreateMemo(ctx, &store.Memo{
|
||||
UID: "legacy-public-comment",
|
||||
CreatorID: owner.ID,
|
||||
Content: "legacy public comment under private parent",
|
||||
Visibility: store.Public,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
parentUID := parent.Name[len("memos/"):]
|
||||
parentMemo, err := ts.Store.GetMemo(ctx, &store.FindMemo{UID: &parentUID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, parentMemo)
|
||||
|
||||
_, err = ts.Store.UpsertMemoRelation(ctx, &store.MemoRelation{
|
||||
MemoID: legacyComment.ID,
|
||||
RelatedMemoID: parentMemo.ID,
|
||||
Type: store.MemoRelationComment,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
commentName := "memos/" + legacyComment.UID
|
||||
_, err = ts.Service.GetMemo(ctx, &apiv1.GetMemoRequest{Name: commentName})
|
||||
require.Equal(t, codes.Unauthenticated, status.Code(err))
|
||||
|
||||
_, err = ts.Service.GetMemo(otherCtx, &apiv1.GetMemoRequest{Name: commentName})
|
||||
require.Equal(t, codes.PermissionDenied, status.Code(err))
|
||||
|
||||
comment, err := ts.Service.GetMemo(ownerCtx, &apiv1.GetMemoRequest{Name: commentName})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, parent.Name, comment.GetParent())
|
||||
|
||||
_, err = ts.Service.ListMemoComments(ctx, &apiv1.ListMemoCommentsRequest{Name: parent.Name})
|
||||
require.Equal(t, codes.Unauthenticated, status.Code(err))
|
||||
|
||||
_, err = ts.Service.ListMemoComments(otherCtx, &apiv1.ListMemoCommentsRequest{Name: parent.Name})
|
||||
require.Equal(t, codes.PermissionDenied, status.Code(err))
|
||||
|
||||
comments, err := ts.Service.ListMemoComments(ownerCtx, &apiv1.ListMemoCommentsRequest{Name: parent.Name})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, comments.Memos, 1)
|
||||
require.Equal(t, commentName, comments.Memos[0].Name)
|
||||
}
|
||||
|
||||
// TestCreateMemoWithCustomTimestamps tests that custom timestamps can be set when creating memos and comments.
|
||||
// This addresses issue #5483: https://github.com/usememos/memos/issues/5483
|
||||
func TestCreateMemoWithCustomTimestamps(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create a test user
|
||||
user, err := ts.CreateRegularUser(ctx, "test-user-timestamps")
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, user)
|
||||
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
// Define custom timestamps (January 1, 2020)
|
||||
customCreateTime := time.Date(2020, 1, 1, 12, 0, 0, 0, time.UTC)
|
||||
customUpdateTime := time.Date(2020, 1, 2, 12, 0, 0, 0, time.UTC)
|
||||
|
||||
// Test 1: Create a memo with custom create_time
|
||||
memoWithCreateTime, err := ts.Service.CreateMemo(userCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "This memo has a custom creation time",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
CreateTime: timestamppb.New(customCreateTime),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memoWithCreateTime)
|
||||
require.Equal(t, customCreateTime.Unix(), memoWithCreateTime.CreateTime.AsTime().Unix(), "create_time should match the custom timestamp")
|
||||
|
||||
// Test 2: Create a memo with custom update_time
|
||||
memoWithUpdateTime, err := ts.Service.CreateMemo(userCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "This memo has a custom update time",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
UpdateTime: timestamppb.New(customUpdateTime),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memoWithUpdateTime)
|
||||
require.Equal(t, customUpdateTime.Unix(), memoWithUpdateTime.UpdateTime.AsTime().Unix(), "update_time should match the custom timestamp")
|
||||
|
||||
// Test 3: Create a memo with all custom timestamps
|
||||
memoWithAllTimestamps, err := ts.Service.CreateMemo(userCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "This memo has all custom timestamps",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
CreateTime: timestamppb.New(customCreateTime),
|
||||
UpdateTime: timestamppb.New(customUpdateTime),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memoWithAllTimestamps)
|
||||
require.Equal(t, customCreateTime.Unix(), memoWithAllTimestamps.CreateTime.AsTime().Unix(), "create_time should match the custom timestamp")
|
||||
require.Equal(t, customUpdateTime.Unix(), memoWithAllTimestamps.UpdateTime.AsTime().Unix(), "update_time should match the custom timestamp")
|
||||
|
||||
// Test 4: Create a comment (memo relation) with custom timestamps
|
||||
parentMemo, err := ts.Service.CreateMemo(userCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "This is the parent memo",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, parentMemo)
|
||||
|
||||
customCommentCreateTime := time.Date(2021, 6, 15, 10, 30, 0, 0, time.UTC)
|
||||
comment, err := ts.Service.CreateMemoComment(userCtx, &apiv1.CreateMemoCommentRequest{
|
||||
Name: parentMemo.Name,
|
||||
Comment: &apiv1.Memo{
|
||||
Content: "This is a comment with custom create time",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
CreateTime: timestamppb.New(customCommentCreateTime),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, comment)
|
||||
require.Equal(t, customCommentCreateTime.Unix(), comment.CreateTime.AsTime().Unix(), "comment create_time should match the custom timestamp")
|
||||
|
||||
// Test 5: Verify that memos without custom timestamps still get auto-generated ones
|
||||
memoWithoutTimestamps, err := ts.Service.CreateMemo(userCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "This memo has auto-generated timestamps",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memoWithoutTimestamps)
|
||||
require.NotNil(t, memoWithoutTimestamps.CreateTime, "create_time should be auto-generated")
|
||||
require.NotNil(t, memoWithoutTimestamps.UpdateTime, "update_time should be auto-generated")
|
||||
require.True(t, time.Now().Unix()-memoWithoutTimestamps.CreateTime.AsTime().Unix() < 5, "create_time should be recent (within 5 seconds)")
|
||||
}
|
||||
@@ -0,0 +1,263 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
apiv1 "github.com/usememos/memos/proto/gen/api/v1"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func TestDeleteMemoShare_VerifiesShareBelongsToMemo(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
userOne, err := ts.CreateRegularUser(ctx, "share-owner-one")
|
||||
require.NoError(t, err)
|
||||
userTwo, err := ts.CreateRegularUser(ctx, "share-owner-two")
|
||||
require.NoError(t, err)
|
||||
|
||||
userOneCtx := ts.CreateUserContext(ctx, userOne.ID)
|
||||
userTwoCtx := ts.CreateUserContext(ctx, userTwo.ID)
|
||||
|
||||
memoOne, err := ts.Service.CreateMemo(userOneCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "memo one",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
memoTwo, err := ts.Service.CreateMemo(userTwoCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "memo two",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
share, err := ts.Service.CreateMemoShare(userTwoCtx, &apiv1.CreateMemoShareRequest{
|
||||
Parent: memoTwo.Name,
|
||||
MemoShare: &apiv1.MemoShare{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
shareToken := share.Name[strings.LastIndex(share.Name, "/")+1:]
|
||||
forgedName := memoOne.Name + "/shares/" + shareToken
|
||||
|
||||
_, err = ts.Service.DeleteMemoShare(userOneCtx, &apiv1.DeleteMemoShareRequest{
|
||||
Name: forgedName,
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Equal(t, codes.NotFound, status.Code(err))
|
||||
|
||||
sharedMemo, err := ts.Service.GetMemoByShare(ctx, &apiv1.GetMemoByShareRequest{
|
||||
ShareId: shareToken,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, memoTwo.Name, sharedMemo.Name)
|
||||
}
|
||||
|
||||
func TestGetMemoByShare_IncludesReactions(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "share-reactions")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
memo, err := ts.Service.CreateMemo(userCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "memo with reactions",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
reaction, err := ts.Service.UpsertMemoReaction(userCtx, &apiv1.UpsertMemoReactionRequest{
|
||||
Name: memo.Name,
|
||||
Reaction: &apiv1.Reaction{
|
||||
ContentId: memo.Name,
|
||||
ReactionType: "👍",
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, reaction)
|
||||
|
||||
share, err := ts.Service.CreateMemoShare(userCtx, &apiv1.CreateMemoShareRequest{
|
||||
Parent: memo.Name,
|
||||
MemoShare: &apiv1.MemoShare{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
shareToken := share.Name[strings.LastIndex(share.Name, "/")+1:]
|
||||
sharedMemo, err := ts.Service.GetMemoByShare(ctx, &apiv1.GetMemoByShareRequest{
|
||||
ShareId: shareToken,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, sharedMemo.Reactions, 1)
|
||||
require.Equal(t, "👍", sharedMemo.Reactions[0].ReactionType)
|
||||
require.Equal(t, memo.Name, sharedMemo.Reactions[0].ContentId)
|
||||
}
|
||||
|
||||
func TestGetMemoByShare_SkipsReactionsWithMissingCreators(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
owner, err := ts.CreateRegularUser(ctx, "share-owner")
|
||||
require.NoError(t, err)
|
||||
ownerCtx := ts.CreateUserContext(ctx, owner.ID)
|
||||
|
||||
reactor, err := ts.CreateRegularUser(ctx, "share-reaction-orphan")
|
||||
require.NoError(t, err)
|
||||
reactorCtx := ts.CreateUserContext(ctx, reactor.ID)
|
||||
|
||||
memo, err := ts.Service.CreateMemo(ownerCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "memo with orphan share reaction",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.UpsertMemoReaction(reactorCtx, &apiv1.UpsertMemoReactionRequest{
|
||||
Name: memo.Name,
|
||||
Reaction: &apiv1.Reaction{
|
||||
ContentId: memo.Name,
|
||||
ReactionType: "👍",
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
share, err := ts.Service.CreateMemoShare(ownerCtx, &apiv1.CreateMemoShareRequest{
|
||||
Parent: memo.Name,
|
||||
MemoShare: &apiv1.MemoShare{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Store.DeleteUser(ctx, &store.DeleteUser{ID: reactor.ID})
|
||||
require.NoError(t, err)
|
||||
|
||||
shareToken := share.Name[strings.LastIndex(share.Name, "/")+1:]
|
||||
sharedMemo, err := ts.Service.GetMemoByShare(ctx, &apiv1.GetMemoByShareRequest{
|
||||
ShareId: shareToken,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, sharedMemo.Reactions)
|
||||
}
|
||||
|
||||
func TestGetMemoByShare_ReturnsNotFoundForUnknownShare(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
_, err := ts.Service.GetMemoByShare(ctx, &apiv1.GetMemoByShareRequest{
|
||||
ShareId: "missing-share-token",
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Equal(t, codes.NotFound, status.Code(err))
|
||||
}
|
||||
|
||||
func TestGetMemoByShare_ReturnsNotFoundForExpiredShare(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "share-expired")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
memo, err := ts.Service.CreateMemo(userCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "memo with expired share",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
expiredTs := time.Now().Add(-time.Hour).Unix()
|
||||
expiredShare, err := ts.Store.CreateMemoShare(ctx, &store.MemoShare{
|
||||
UID: "expired-share-token",
|
||||
MemoID: parseMemoIDFromNameForTest(t, ts, memo.Name),
|
||||
CreatorID: user.ID,
|
||||
ExpiresTs: &expiredTs,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.GetMemoByShare(ctx, &apiv1.GetMemoByShareRequest{
|
||||
ShareId: expiredShare.UID,
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Equal(t, codes.NotFound, status.Code(err))
|
||||
}
|
||||
|
||||
func TestGetMemoByShare_ReturnsNotFoundForArchivedMemo(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "share-archived")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
memoResp, err := ts.Service.CreateMemo(userCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "memo that will be archived",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
share, err := ts.Service.CreateMemoShare(userCtx, &apiv1.CreateMemoShareRequest{
|
||||
Parent: memoResp.Name,
|
||||
MemoShare: &apiv1.MemoShare{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
memoID := parseMemoIDFromNameForTest(t, ts, memoResp.Name)
|
||||
memo, err := ts.Store.GetMemo(ctx, &store.FindMemo{ID: &memoID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memo)
|
||||
|
||||
archived := store.Archived
|
||||
err = ts.Store.UpdateMemo(ctx, &store.UpdateMemo{
|
||||
ID: memo.ID,
|
||||
RowStatus: &archived,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
shareToken := share.Name[strings.LastIndex(share.Name, "/")+1:]
|
||||
_, err = ts.Service.GetMemoByShare(ctx, &apiv1.GetMemoByShareRequest{
|
||||
ShareId: shareToken,
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Equal(t, codes.NotFound, status.Code(err))
|
||||
}
|
||||
|
||||
func parseMemoIDFromNameForTest(t *testing.T, ts *TestService, memoName string) int32 {
|
||||
t.Helper()
|
||||
|
||||
memoUID, ok := strings.CutPrefix(memoName, "memos/")
|
||||
require.True(t, ok, "memo name must start with memos/: %s", memoName)
|
||||
|
||||
memo, err := ts.Store.GetMemo(context.Background(), &store.FindMemo{UID: &memoUID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memo)
|
||||
|
||||
return memo.ID
|
||||
}
|
||||
@@ -0,0 +1,235 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
apiv1 "github.com/usememos/memos/proto/gen/api/v1"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func TestDeleteMemoReaction(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("DeleteMemoReaction success by reaction owner", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create user
|
||||
user, err := ts.CreateRegularUser(ctx, "user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
// Create memo
|
||||
memo, err := ts.Service.CreateMemo(userCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "Test memo",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memo)
|
||||
|
||||
// Create reaction
|
||||
reaction, err := ts.Service.UpsertMemoReaction(userCtx, &apiv1.UpsertMemoReactionRequest{
|
||||
Name: memo.Name,
|
||||
Reaction: &apiv1.Reaction{
|
||||
ContentId: memo.Name,
|
||||
ReactionType: "👍",
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, reaction)
|
||||
require.Equal(t, "users/user", reaction.Creator)
|
||||
|
||||
// Delete reaction - should succeed
|
||||
_, err = ts.Service.DeleteMemoReaction(userCtx, &apiv1.DeleteMemoReactionRequest{
|
||||
Name: reaction.Name,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
})
|
||||
|
||||
t.Run("DeleteMemoReaction success by host user", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create regular user
|
||||
regularUser, err := ts.CreateRegularUser(ctx, "user")
|
||||
require.NoError(t, err)
|
||||
regularUserCtx := ts.CreateUserContext(ctx, regularUser.ID)
|
||||
|
||||
// Create host user
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
hostCtx := ts.CreateUserContext(ctx, hostUser.ID)
|
||||
|
||||
// Create memo by regular user
|
||||
memo, err := ts.Service.CreateMemo(regularUserCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "Test memo",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memo)
|
||||
|
||||
// Create reaction by regular user
|
||||
reaction, err := ts.Service.UpsertMemoReaction(regularUserCtx, &apiv1.UpsertMemoReactionRequest{
|
||||
Name: memo.Name,
|
||||
Reaction: &apiv1.Reaction{
|
||||
ContentId: memo.Name,
|
||||
ReactionType: "👍",
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, reaction)
|
||||
|
||||
// Host user can delete reaction - should succeed
|
||||
_, err = ts.Service.DeleteMemoReaction(hostCtx, &apiv1.DeleteMemoReactionRequest{
|
||||
Name: reaction.Name,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
})
|
||||
|
||||
t.Run("DeleteMemoReaction permission denied for non-owner", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create user1
|
||||
user1, err := ts.CreateRegularUser(ctx, "user1")
|
||||
require.NoError(t, err)
|
||||
user1Ctx := ts.CreateUserContext(ctx, user1.ID)
|
||||
|
||||
// Create user2
|
||||
user2, err := ts.CreateRegularUser(ctx, "user2")
|
||||
require.NoError(t, err)
|
||||
user2Ctx := ts.CreateUserContext(ctx, user2.ID)
|
||||
|
||||
// Create memo by user1
|
||||
memo, err := ts.Service.CreateMemo(user1Ctx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "Test memo",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memo)
|
||||
|
||||
// Create reaction by user1
|
||||
reaction, err := ts.Service.UpsertMemoReaction(user1Ctx, &apiv1.UpsertMemoReactionRequest{
|
||||
Name: memo.Name,
|
||||
Reaction: &apiv1.Reaction{
|
||||
ContentId: memo.Name,
|
||||
ReactionType: "👍",
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, reaction)
|
||||
|
||||
// User2 tries to delete reaction - should fail with permission denied
|
||||
_, err = ts.Service.DeleteMemoReaction(user2Ctx, &apiv1.DeleteMemoReactionRequest{
|
||||
Name: reaction.Name,
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "permission denied")
|
||||
})
|
||||
|
||||
t.Run("DeleteMemoReaction unauthenticated", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create user
|
||||
user, err := ts.CreateRegularUser(ctx, "user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
// Create memo
|
||||
memo, err := ts.Service.CreateMemo(userCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "Test memo",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memo)
|
||||
|
||||
// Create reaction
|
||||
reaction, err := ts.Service.UpsertMemoReaction(userCtx, &apiv1.UpsertMemoReactionRequest{
|
||||
Name: memo.Name,
|
||||
Reaction: &apiv1.Reaction{
|
||||
ContentId: memo.Name,
|
||||
ReactionType: "👍",
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, reaction)
|
||||
|
||||
// Unauthenticated user tries to delete reaction - should fail
|
||||
_, err = ts.Service.DeleteMemoReaction(ctx, &apiv1.DeleteMemoReactionRequest{
|
||||
Name: reaction.Name,
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "not authenticated")
|
||||
})
|
||||
|
||||
t.Run("DeleteMemoReaction not found returns permission denied", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create user
|
||||
user, err := ts.CreateRegularUser(ctx, "user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
// Try to delete non-existent reaction - should fail with permission denied
|
||||
// (not "not found" to avoid information disclosure)
|
||||
// Use new nested resource format: memos/{memo}/reactions/{reaction}
|
||||
_, err = ts.Service.DeleteMemoReaction(userCtx, &apiv1.DeleteMemoReactionRequest{
|
||||
Name: "memos/nonexistent/reactions/99999",
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "permission denied")
|
||||
require.NotContains(t, err.Error(), "not found")
|
||||
})
|
||||
}
|
||||
|
||||
func TestListMemoReactionsSkipsMissingCreators(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
owner, err := ts.CreateRegularUser(ctx, "reaction-owner")
|
||||
require.NoError(t, err)
|
||||
ownerCtx := ts.CreateUserContext(ctx, owner.ID)
|
||||
|
||||
reactor, err := ts.CreateRegularUser(ctx, "reaction-orphan")
|
||||
require.NoError(t, err)
|
||||
reactorCtx := ts.CreateUserContext(ctx, reactor.ID)
|
||||
|
||||
memo, err := ts.Service.CreateMemo(ownerCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "reaction list memo",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.UpsertMemoReaction(reactorCtx, &apiv1.UpsertMemoReactionRequest{
|
||||
Name: memo.Name,
|
||||
Reaction: &apiv1.Reaction{
|
||||
ContentId: memo.Name,
|
||||
ReactionType: "🔥",
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Store.DeleteUser(ctx, &store.DeleteUser{ID: reactor.ID})
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.ListMemoReactions(ctx, &apiv1.ListMemoReactionsRequest{Name: memo.Name})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, resp.Reactions)
|
||||
}
|
||||
@@ -0,0 +1,838 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/protobuf/types/known/fieldmaskpb"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
)
|
||||
|
||||
func TestListShortcuts(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("ListShortcuts success", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create a user
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
// List shortcuts (should be empty initially)
|
||||
req := &v1pb.ListShortcutsRequest{
|
||||
Parent: fmt.Sprintf("users/%s", user.Username),
|
||||
}
|
||||
|
||||
resp, err := ts.Service.ListShortcuts(userCtx, req)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp)
|
||||
require.Empty(t, resp.Shortcuts)
|
||||
})
|
||||
|
||||
t.Run("ListShortcuts permission denied for different user", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create two users
|
||||
user1, err := ts.CreateRegularUser(ctx, "user1")
|
||||
require.NoError(t, err)
|
||||
user2, err := ts.CreateRegularUser(ctx, "user2")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user1 context but try to list user2's shortcuts
|
||||
userCtx := ts.CreateUserContext(ctx, user1.ID)
|
||||
|
||||
req := &v1pb.ListShortcutsRequest{
|
||||
Parent: fmt.Sprintf("users/%s", user2.Username),
|
||||
}
|
||||
|
||||
_, err = ts.Service.ListShortcuts(userCtx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "permission denied")
|
||||
})
|
||||
|
||||
t.Run("ListShortcuts invalid parent format", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create user
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
req := &v1pb.ListShortcutsRequest{
|
||||
Parent: "invalid-parent-format",
|
||||
}
|
||||
|
||||
_, err = ts.Service.ListShortcuts(userCtx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "invalid user name")
|
||||
})
|
||||
|
||||
t.Run("ListShortcuts unauthenticated", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
_, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
req := &v1pb.ListShortcutsRequest{
|
||||
Parent: "users/testuser",
|
||||
}
|
||||
|
||||
_, err = ts.Service.ListShortcuts(ctx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "permission denied")
|
||||
})
|
||||
|
||||
t.Run("ListShortcuts returns not found for numeric parent", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
_, err = ts.Service.ListShortcuts(userCtx, &v1pb.ListShortcutsRequest{
|
||||
Parent: "users/1",
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "user not found")
|
||||
})
|
||||
}
|
||||
|
||||
func TestGetShortcut(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("GetShortcut success", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create a user
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
// First create a shortcut
|
||||
createReq := &v1pb.CreateShortcutRequest{
|
||||
Parent: fmt.Sprintf("users/%s", user.Username),
|
||||
Shortcut: &v1pb.Shortcut{
|
||||
Title: "Test Shortcut",
|
||||
Filter: "tag in [\"test\"]",
|
||||
},
|
||||
}
|
||||
|
||||
created, err := ts.Service.CreateShortcut(userCtx, createReq)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Now get the shortcut
|
||||
getReq := &v1pb.GetShortcutRequest{
|
||||
Name: created.Name,
|
||||
}
|
||||
|
||||
resp, err := ts.Service.GetShortcut(userCtx, getReq)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp)
|
||||
require.Equal(t, created.Name, resp.Name)
|
||||
require.Equal(t, "Test Shortcut", resp.Title)
|
||||
require.Equal(t, "tag in [\"test\"]", resp.Filter)
|
||||
})
|
||||
|
||||
t.Run("GetShortcut permission denied for different user", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create two users
|
||||
user1, err := ts.CreateRegularUser(ctx, "user1")
|
||||
require.NoError(t, err)
|
||||
user2, err := ts.CreateRegularUser(ctx, "user2")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create shortcut as user1
|
||||
user1Ctx := ts.CreateUserContext(ctx, user1.ID)
|
||||
createReq := &v1pb.CreateShortcutRequest{
|
||||
Parent: fmt.Sprintf("users/%s", user1.Username),
|
||||
Shortcut: &v1pb.Shortcut{
|
||||
Title: "User1 Shortcut",
|
||||
Filter: "tag in [\"user1\"]",
|
||||
},
|
||||
}
|
||||
|
||||
created, err := ts.Service.CreateShortcut(user1Ctx, createReq)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Try to get shortcut as user2
|
||||
user2Ctx := ts.CreateUserContext(ctx, user2.ID)
|
||||
getReq := &v1pb.GetShortcutRequest{
|
||||
Name: created.Name,
|
||||
}
|
||||
|
||||
_, err = ts.Service.GetShortcut(user2Ctx, getReq)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "permission denied")
|
||||
})
|
||||
|
||||
t.Run("GetShortcut invalid name format", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create user
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
req := &v1pb.GetShortcutRequest{
|
||||
Name: "invalid-shortcut-name",
|
||||
}
|
||||
|
||||
_, err = ts.Service.GetShortcut(userCtx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "invalid shortcut name")
|
||||
})
|
||||
|
||||
t.Run("GetShortcut not found", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create user
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
req := &v1pb.GetShortcutRequest{
|
||||
Name: fmt.Sprintf("users/%s", user.Username) + "/shortcuts/nonexistent",
|
||||
}
|
||||
|
||||
_, err = ts.Service.GetShortcut(userCtx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "not found")
|
||||
})
|
||||
}
|
||||
|
||||
func TestCreateShortcut(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("CreateShortcut success", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create a user
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
req := &v1pb.CreateShortcutRequest{
|
||||
Parent: fmt.Sprintf("users/%s", user.Username),
|
||||
Shortcut: &v1pb.Shortcut{
|
||||
Title: "My Shortcut",
|
||||
Filter: "tag in [\"important\"]",
|
||||
},
|
||||
}
|
||||
|
||||
resp, err := ts.Service.CreateShortcut(userCtx, req)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp)
|
||||
require.Equal(t, "My Shortcut", resp.Title)
|
||||
require.Equal(t, "tag in [\"important\"]", resp.Filter)
|
||||
require.Contains(t, resp.Name, fmt.Sprintf("users/%s/shortcuts/", user.Username))
|
||||
|
||||
// Verify the shortcut was created by listing
|
||||
listReq := &v1pb.ListShortcutsRequest{
|
||||
Parent: fmt.Sprintf("users/%s", user.Username),
|
||||
}
|
||||
|
||||
listResp, err := ts.Service.ListShortcuts(userCtx, listReq)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, listResp.Shortcuts, 1)
|
||||
require.Equal(t, "My Shortcut", listResp.Shortcuts[0].Title)
|
||||
})
|
||||
|
||||
t.Run("CreateShortcut permission denied for different user", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create two users
|
||||
user1, err := ts.CreateRegularUser(ctx, "user1")
|
||||
require.NoError(t, err)
|
||||
user2, err := ts.CreateRegularUser(ctx, "user2")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user1 context but try to create shortcut for user2
|
||||
userCtx := ts.CreateUserContext(ctx, user1.ID)
|
||||
|
||||
req := &v1pb.CreateShortcutRequest{
|
||||
Parent: fmt.Sprintf("users/%s", user2.Username),
|
||||
Shortcut: &v1pb.Shortcut{
|
||||
Title: "Forbidden Shortcut",
|
||||
Filter: "tag in [\"forbidden\"]",
|
||||
},
|
||||
}
|
||||
|
||||
_, err = ts.Service.CreateShortcut(userCtx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "permission denied")
|
||||
})
|
||||
|
||||
t.Run("CreateShortcut invalid parent format", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create user
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
req := &v1pb.CreateShortcutRequest{
|
||||
Parent: "invalid-parent",
|
||||
Shortcut: &v1pb.Shortcut{
|
||||
Title: "Test Shortcut",
|
||||
Filter: "tag in [\"test\"]",
|
||||
},
|
||||
}
|
||||
|
||||
_, err = ts.Service.CreateShortcut(userCtx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "invalid user name")
|
||||
})
|
||||
|
||||
t.Run("CreateShortcut invalid filter", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create user
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
req := &v1pb.CreateShortcutRequest{
|
||||
Parent: fmt.Sprintf("users/%s", user.Username),
|
||||
Shortcut: &v1pb.Shortcut{
|
||||
Title: "Invalid Filter Shortcut",
|
||||
Filter: "invalid||filter))syntax",
|
||||
},
|
||||
}
|
||||
|
||||
_, err = ts.Service.CreateShortcut(userCtx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "invalid filter")
|
||||
})
|
||||
|
||||
t.Run("CreateShortcut missing title", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create user
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
req := &v1pb.CreateShortcutRequest{
|
||||
Parent: fmt.Sprintf("users/%s", user.Username),
|
||||
Shortcut: &v1pb.Shortcut{
|
||||
Filter: "tag in [\"test\"]",
|
||||
},
|
||||
}
|
||||
|
||||
_, err = ts.Service.CreateShortcut(userCtx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "title is required")
|
||||
})
|
||||
}
|
||||
|
||||
func TestUpdateShortcut(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("UpdateShortcut success", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create a user
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
// Create a shortcut first
|
||||
createReq := &v1pb.CreateShortcutRequest{
|
||||
Parent: fmt.Sprintf("users/%s", user.Username),
|
||||
Shortcut: &v1pb.Shortcut{
|
||||
Title: "Original Title",
|
||||
Filter: "tag in [\"original\"]",
|
||||
},
|
||||
}
|
||||
|
||||
created, err := ts.Service.CreateShortcut(userCtx, createReq)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Update the shortcut
|
||||
updateReq := &v1pb.UpdateShortcutRequest{
|
||||
Shortcut: &v1pb.Shortcut{
|
||||
Name: created.Name,
|
||||
Title: "Updated Title",
|
||||
Filter: "tag in [\"updated\"]",
|
||||
},
|
||||
UpdateMask: &fieldmaskpb.FieldMask{
|
||||
Paths: []string{"title", "filter"},
|
||||
},
|
||||
}
|
||||
|
||||
updated, err := ts.Service.UpdateShortcut(userCtx, updateReq)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, updated)
|
||||
require.Equal(t, "Updated Title", updated.Title)
|
||||
require.Equal(t, "tag in [\"updated\"]", updated.Filter)
|
||||
require.Equal(t, created.Name, updated.Name)
|
||||
})
|
||||
|
||||
t.Run("UpdateShortcut permission denied for different user", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create two users
|
||||
user1, err := ts.CreateRegularUser(ctx, "user1")
|
||||
require.NoError(t, err)
|
||||
user2, err := ts.CreateRegularUser(ctx, "user2")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create shortcut as user1
|
||||
user1Ctx := ts.CreateUserContext(ctx, user1.ID)
|
||||
createReq := &v1pb.CreateShortcutRequest{
|
||||
Parent: fmt.Sprintf("users/%s", user1.Username),
|
||||
Shortcut: &v1pb.Shortcut{
|
||||
Title: "User1 Shortcut",
|
||||
Filter: "tag in [\"user1\"]",
|
||||
},
|
||||
}
|
||||
|
||||
created, err := ts.Service.CreateShortcut(user1Ctx, createReq)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Try to update shortcut as user2
|
||||
user2Ctx := ts.CreateUserContext(ctx, user2.ID)
|
||||
updateReq := &v1pb.UpdateShortcutRequest{
|
||||
Shortcut: &v1pb.Shortcut{
|
||||
Name: created.Name,
|
||||
Title: "Hacked Title",
|
||||
Filter: "tag in [\"hacked\"]",
|
||||
},
|
||||
UpdateMask: &fieldmaskpb.FieldMask{
|
||||
Paths: []string{"title", "filter"},
|
||||
},
|
||||
}
|
||||
|
||||
_, err = ts.Service.UpdateShortcut(user2Ctx, updateReq)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "permission denied")
|
||||
})
|
||||
|
||||
t.Run("UpdateShortcut missing update mask", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create a user and context for authentication
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
req := &v1pb.UpdateShortcutRequest{
|
||||
Shortcut: &v1pb.Shortcut{
|
||||
Name: fmt.Sprintf("users/%s/shortcuts/test", user.Username),
|
||||
Title: "Updated Title",
|
||||
},
|
||||
}
|
||||
|
||||
_, err = ts.Service.UpdateShortcut(userCtx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "update mask is required")
|
||||
})
|
||||
|
||||
t.Run("UpdateShortcut invalid name format", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
req := &v1pb.UpdateShortcutRequest{
|
||||
Shortcut: &v1pb.Shortcut{
|
||||
Name: "invalid-shortcut-name",
|
||||
Title: "Updated Title",
|
||||
},
|
||||
UpdateMask: &fieldmaskpb.FieldMask{
|
||||
Paths: []string{"title"},
|
||||
},
|
||||
}
|
||||
|
||||
_, err := ts.Service.UpdateShortcut(ctx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "invalid shortcut name")
|
||||
})
|
||||
|
||||
t.Run("UpdateShortcut invalid filter", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create user
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
// Create a shortcut first
|
||||
createReq := &v1pb.CreateShortcutRequest{
|
||||
Parent: fmt.Sprintf("users/%s", user.Username),
|
||||
Shortcut: &v1pb.Shortcut{
|
||||
Title: "Test Shortcut",
|
||||
Filter: "tag in [\"test\"]",
|
||||
},
|
||||
}
|
||||
|
||||
created, err := ts.Service.CreateShortcut(userCtx, createReq)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Try to update with invalid filter
|
||||
updateReq := &v1pb.UpdateShortcutRequest{
|
||||
Shortcut: &v1pb.Shortcut{
|
||||
Name: created.Name,
|
||||
Filter: "invalid||filter))syntax",
|
||||
},
|
||||
UpdateMask: &fieldmaskpb.FieldMask{
|
||||
Paths: []string{"filter"},
|
||||
},
|
||||
}
|
||||
|
||||
_, err = ts.Service.UpdateShortcut(userCtx, updateReq)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "invalid filter")
|
||||
})
|
||||
}
|
||||
|
||||
func TestDeleteShortcut(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("DeleteShortcut success", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create a user
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
// Create a shortcut first
|
||||
createReq := &v1pb.CreateShortcutRequest{
|
||||
Parent: fmt.Sprintf("users/%s", user.Username),
|
||||
Shortcut: &v1pb.Shortcut{
|
||||
Title: "Shortcut to Delete",
|
||||
Filter: "tag in [\"delete\"]",
|
||||
},
|
||||
}
|
||||
|
||||
created, err := ts.Service.CreateShortcut(userCtx, createReq)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Delete the shortcut
|
||||
deleteReq := &v1pb.DeleteShortcutRequest{
|
||||
Name: created.Name,
|
||||
}
|
||||
|
||||
_, err = ts.Service.DeleteShortcut(userCtx, deleteReq)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Verify deletion by listing shortcuts
|
||||
listReq := &v1pb.ListShortcutsRequest{
|
||||
Parent: fmt.Sprintf("users/%s", user.Username),
|
||||
}
|
||||
|
||||
listResp, err := ts.Service.ListShortcuts(userCtx, listReq)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, listResp.Shortcuts)
|
||||
|
||||
// Also verify by trying to get the deleted shortcut
|
||||
getReq := &v1pb.GetShortcutRequest{
|
||||
Name: created.Name,
|
||||
}
|
||||
|
||||
_, err = ts.Service.GetShortcut(userCtx, getReq)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "not found")
|
||||
})
|
||||
|
||||
t.Run("DeleteShortcut permission denied for different user", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create two users
|
||||
user1, err := ts.CreateRegularUser(ctx, "user1")
|
||||
require.NoError(t, err)
|
||||
user2, err := ts.CreateRegularUser(ctx, "user2")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create shortcut as user1
|
||||
user1Ctx := ts.CreateUserContext(ctx, user1.ID)
|
||||
createReq := &v1pb.CreateShortcutRequest{
|
||||
Parent: fmt.Sprintf("users/%s", user1.Username),
|
||||
Shortcut: &v1pb.Shortcut{
|
||||
Title: "User1 Shortcut",
|
||||
Filter: "tag in [\"user1\"]",
|
||||
},
|
||||
}
|
||||
|
||||
created, err := ts.Service.CreateShortcut(user1Ctx, createReq)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Try to delete shortcut as user2
|
||||
user2Ctx := ts.CreateUserContext(ctx, user2.ID)
|
||||
deleteReq := &v1pb.DeleteShortcutRequest{
|
||||
Name: created.Name,
|
||||
}
|
||||
|
||||
_, err = ts.Service.DeleteShortcut(user2Ctx, deleteReq)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "permission denied")
|
||||
})
|
||||
|
||||
t.Run("DeleteShortcut invalid name format", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
req := &v1pb.DeleteShortcutRequest{
|
||||
Name: "invalid-shortcut-name",
|
||||
}
|
||||
|
||||
_, err := ts.Service.DeleteShortcut(ctx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "invalid shortcut name")
|
||||
})
|
||||
|
||||
t.Run("DeleteShortcut not found", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create user
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
req := &v1pb.DeleteShortcutRequest{
|
||||
Name: fmt.Sprintf("users/%s", user.Username) + "/shortcuts/nonexistent",
|
||||
}
|
||||
|
||||
_, err = ts.Service.DeleteShortcut(userCtx, req)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "not found")
|
||||
})
|
||||
}
|
||||
|
||||
func TestShortcutFiltering(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("CreateShortcut with valid filters", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create user
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
// Test various valid filter formats
|
||||
validFilters := []string{
|
||||
"tag in [\"work\"]",
|
||||
"content.contains(\"meeting\")",
|
||||
"tag in [\"work\"] && content.contains(\"meeting\")",
|
||||
"tag in [\"work\"] || tag in [\"personal\"]",
|
||||
"creator_id == 1",
|
||||
"visibility == \"PUBLIC\"",
|
||||
"has_task_list == true",
|
||||
"has_task_list == false",
|
||||
}
|
||||
|
||||
for i, filter := range validFilters {
|
||||
req := &v1pb.CreateShortcutRequest{
|
||||
Parent: fmt.Sprintf("users/%s", user.Username),
|
||||
Shortcut: &v1pb.Shortcut{
|
||||
Title: "Valid Filter " + string(rune(i)),
|
||||
Filter: filter,
|
||||
},
|
||||
}
|
||||
|
||||
_, err = ts.Service.CreateShortcut(userCtx, req)
|
||||
require.NoError(t, err, "Filter should be valid: %s", filter)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("CreateShortcut with invalid filters", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create user
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
// Test various invalid filter formats
|
||||
invalidFilters := []string{
|
||||
"tag in ", // incomplete expression
|
||||
"invalid_field @in [\"value\"]", // unknown field
|
||||
"tag in [\"work\"] &&", // incomplete expression
|
||||
"tag in [\"work\"] || || tag in [\"test\"]", // double operator
|
||||
"((tag in [\"work\"]", // unmatched parentheses
|
||||
"tag in [\"work\"] && )", // mismatched parentheses
|
||||
"tag == \"work\"", // wrong operator (== not supported for tags)
|
||||
"tag in work", // missing brackets
|
||||
}
|
||||
|
||||
for _, filter := range invalidFilters {
|
||||
req := &v1pb.CreateShortcutRequest{
|
||||
Parent: fmt.Sprintf("users/%s", user.Username),
|
||||
Shortcut: &v1pb.Shortcut{
|
||||
Title: "Invalid Filter Test",
|
||||
Filter: filter,
|
||||
},
|
||||
}
|
||||
|
||||
_, err = ts.Service.CreateShortcut(userCtx, req)
|
||||
require.Error(t, err, "Filter should be invalid: %s", filter)
|
||||
require.Contains(t, err.Error(), "invalid filter", "Error should mention invalid filter for: %s", filter)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestShortcutCRUDComplete(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("Complete CRUD lifecycle", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create user
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set user context
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
// 1. Create multiple shortcuts
|
||||
shortcut1Req := &v1pb.CreateShortcutRequest{
|
||||
Parent: fmt.Sprintf("users/%s", user.Username),
|
||||
Shortcut: &v1pb.Shortcut{
|
||||
Title: "Work Notes",
|
||||
Filter: "tag in [\"work\"]",
|
||||
},
|
||||
}
|
||||
|
||||
shortcut2Req := &v1pb.CreateShortcutRequest{
|
||||
Parent: fmt.Sprintf("users/%s", user.Username),
|
||||
Shortcut: &v1pb.Shortcut{
|
||||
Title: "Personal Notes",
|
||||
Filter: "tag in [\"personal\"]",
|
||||
},
|
||||
}
|
||||
|
||||
created1, err := ts.Service.CreateShortcut(userCtx, shortcut1Req)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "Work Notes", created1.Title)
|
||||
|
||||
created2, err := ts.Service.CreateShortcut(userCtx, shortcut2Req)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "Personal Notes", created2.Title)
|
||||
|
||||
// 2. List shortcuts and verify both exist
|
||||
listReq := &v1pb.ListShortcutsRequest{
|
||||
Parent: fmt.Sprintf("users/%s", user.Username),
|
||||
}
|
||||
|
||||
listResp, err := ts.Service.ListShortcuts(userCtx, listReq)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, listResp.Shortcuts, 2)
|
||||
|
||||
// 3. Get individual shortcuts
|
||||
getReq1 := &v1pb.GetShortcutRequest{Name: created1.Name}
|
||||
getResp1, err := ts.Service.GetShortcut(userCtx, getReq1)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, created1.Name, getResp1.Name)
|
||||
require.Equal(t, "Work Notes", getResp1.Title)
|
||||
|
||||
getReq2 := &v1pb.GetShortcutRequest{Name: created2.Name}
|
||||
getResp2, err := ts.Service.GetShortcut(userCtx, getReq2)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, created2.Name, getResp2.Name)
|
||||
require.Equal(t, "Personal Notes", getResp2.Title)
|
||||
|
||||
// 4. Update one shortcut
|
||||
updateReq := &v1pb.UpdateShortcutRequest{
|
||||
Shortcut: &v1pb.Shortcut{
|
||||
Name: created1.Name,
|
||||
Title: "Work & Meeting Notes",
|
||||
Filter: "tag in [\"work\"] || tag in [\"meeting\"]",
|
||||
},
|
||||
UpdateMask: &fieldmaskpb.FieldMask{
|
||||
Paths: []string{"title", "filter"},
|
||||
},
|
||||
}
|
||||
|
||||
updated, err := ts.Service.UpdateShortcut(userCtx, updateReq)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "Work & Meeting Notes", updated.Title)
|
||||
require.Equal(t, "tag in [\"work\"] || tag in [\"meeting\"]", updated.Filter)
|
||||
|
||||
// 5. Verify update by getting it again
|
||||
getUpdatedReq := &v1pb.GetShortcutRequest{Name: created1.Name}
|
||||
getUpdatedResp, err := ts.Service.GetShortcut(userCtx, getUpdatedReq)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "Work & Meeting Notes", getUpdatedResp.Title)
|
||||
require.Equal(t, "tag in [\"work\"] || tag in [\"meeting\"]", getUpdatedResp.Filter)
|
||||
|
||||
// 6. Delete one shortcut
|
||||
deleteReq := &v1pb.DeleteShortcutRequest{
|
||||
Name: created2.Name,
|
||||
}
|
||||
|
||||
_, err = ts.Service.DeleteShortcut(userCtx, deleteReq)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 7. Verify deletion by listing (should only have 1 left)
|
||||
finalListResp, err := ts.Service.ListShortcuts(userCtx, listReq)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, finalListResp.Shortcuts, 1)
|
||||
require.Equal(t, "Work & Meeting Notes", finalListResp.Shortcuts[0].Title)
|
||||
|
||||
// 8. Verify deleted shortcut can't be accessed
|
||||
getDeletedReq := &v1pb.GetShortcutRequest{Name: created2.Name}
|
||||
_, err = ts.Service.GetShortcut(userCtx, getDeletedReq)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "not found")
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,135 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/labstack/echo/v5"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/usememos/memos/server/auth"
|
||||
apiv1 "github.com/usememos/memos/server/router/api/v1"
|
||||
)
|
||||
|
||||
func TestSSEHandler_Authentication(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "sse-user")
|
||||
require.NoError(t, err)
|
||||
|
||||
token, _, err := auth.GenerateAccessTokenV2(
|
||||
user.ID,
|
||||
user.Username,
|
||||
string(user.Role),
|
||||
string(user.RowStatus),
|
||||
[]byte(ts.Secret),
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
e := echo.New()
|
||||
apiv1.RegisterSSERoutes(e, ts.Service.SSEHub, ts.Store, ts.Secret)
|
||||
|
||||
t.Run("no token returns 401", func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/v1/sse", nil)
|
||||
rec := httptest.NewRecorder()
|
||||
e.ServeHTTP(rec, req)
|
||||
require.Equal(t, http.StatusUnauthorized, rec.Code)
|
||||
})
|
||||
|
||||
t.Run("invalid token returns 401", func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/v1/sse", nil)
|
||||
req.Header.Set("Authorization", "Bearer invalid-token")
|
||||
rec := httptest.NewRecorder()
|
||||
e.ServeHTTP(rec, req)
|
||||
require.Equal(t, http.StatusUnauthorized, rec.Code)
|
||||
})
|
||||
|
||||
t.Run("valid token returns 200 and stream", func(t *testing.T) {
|
||||
// Use a cancellable context so we can close the SSE connection after
|
||||
// confirming the headers, preventing the handler's event loop from
|
||||
// blocking the test indefinitely.
|
||||
reqCtx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/v1/sse", nil).WithContext(reqCtx)
|
||||
req.Header.Set("Authorization", "Bearer "+token)
|
||||
rec := httptest.NewRecorder()
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
defer close(done)
|
||||
e.ServeHTTP(rec, req)
|
||||
}()
|
||||
// Cancel the context to signal client disconnect, which exits the SSE loop.
|
||||
cancel()
|
||||
<-done
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
require.Equal(t, "text/event-stream", rec.Header().Get("Content-Type"))
|
||||
})
|
||||
|
||||
t.Run("token in query param returns 401", func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/v1/sse?token="+token, nil)
|
||||
rec := httptest.NewRecorder()
|
||||
e.ServeHTTP(rec, req)
|
||||
require.Equal(t, http.StatusUnauthorized, rec.Code)
|
||||
})
|
||||
|
||||
t.Run("valid token streams initial comment", func(t *testing.T) {
|
||||
server := httptest.NewServer(e)
|
||||
defer server.Close()
|
||||
|
||||
reqCtx, cancel := context.WithTimeout(ctx, time.Second)
|
||||
defer cancel()
|
||||
req, err := http.NewRequestWithContext(reqCtx, http.MethodGet, server.URL+"/api/v1/sse", nil)
|
||||
require.NoError(t, err)
|
||||
req.Header.Set("Authorization", "Bearer "+token)
|
||||
|
||||
resp, err := server.Client().Do(req)
|
||||
require.NoError(t, err)
|
||||
defer resp.Body.Close()
|
||||
require.Equal(t, http.StatusOK, resp.StatusCode)
|
||||
require.Equal(t, "text/event-stream", resp.Header.Get("Content-Type"))
|
||||
|
||||
line, err := bufio.NewReader(resp.Body).ReadString('\n')
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, ": connected\n", line)
|
||||
})
|
||||
|
||||
t.Run("hub close disconnects stream", func(t *testing.T) {
|
||||
server := httptest.NewServer(e)
|
||||
defer server.Close()
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, server.URL+"/api/v1/sse", nil)
|
||||
require.NoError(t, err)
|
||||
req.Header.Set("Authorization", "Bearer "+token)
|
||||
|
||||
resp, err := server.Client().Do(req) //nolint:bodyclose // Body is closed after verifying the SSE stream disconnects.
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
body := resp.Body
|
||||
defer body.Close()
|
||||
require.Equal(t, http.StatusOK, resp.StatusCode)
|
||||
require.Equal(t, "text/event-stream", resp.Header.Get("Content-Type"))
|
||||
|
||||
ts.Service.SSEHub.Close()
|
||||
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := io.ReadAll(body)
|
||||
done <- err
|
||||
}()
|
||||
|
||||
select {
|
||||
case err := <-done:
|
||||
require.NoError(t, err)
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("SSE stream did not close after hub close")
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/usememos/memos/internal/markdown"
|
||||
"github.com/usememos/memos/internal/profile"
|
||||
"github.com/usememos/memos/server/auth"
|
||||
apiv1 "github.com/usememos/memos/server/router/api/v1"
|
||||
"github.com/usememos/memos/store"
|
||||
teststore "github.com/usememos/memos/store/test"
|
||||
)
|
||||
|
||||
// TestService holds the test service setup for API v1 services.
|
||||
type TestService struct {
|
||||
Service *apiv1.APIV1Service
|
||||
Store *store.Store
|
||||
Profile *profile.Profile
|
||||
Secret string
|
||||
}
|
||||
|
||||
// NewTestService creates a new test service with SQLite database.
|
||||
func NewTestService(t *testing.T) *TestService {
|
||||
ctx := context.Background()
|
||||
|
||||
// Create a test store with SQLite
|
||||
testStore := teststore.NewTestingStore(ctx, t)
|
||||
|
||||
// Align the profile data directory with the test store so attachment files and
|
||||
// derived caches resolve against the same location as DeleteAttachmentStorage.
|
||||
testProfile := &profile.Profile{
|
||||
Demo: true,
|
||||
Version: "test-1.0.0",
|
||||
Commit: "test-commit",
|
||||
InstanceURL: "http://localhost:8080",
|
||||
Driver: "sqlite",
|
||||
DSN: ":memory:",
|
||||
Data: testStore.GetDataDir(),
|
||||
}
|
||||
|
||||
// Create APIV1Service with nil grpcServer since we're testing direct calls
|
||||
secret := "test-secret"
|
||||
markdownService := markdown.NewService(
|
||||
markdown.WithTagExtension(),
|
||||
markdown.WithMentionExtension(),
|
||||
)
|
||||
service := &apiv1.APIV1Service{
|
||||
Secret: secret,
|
||||
Profile: testProfile,
|
||||
Store: testStore,
|
||||
MarkdownService: markdownService,
|
||||
SSEHub: apiv1.NewSSEHub(),
|
||||
}
|
||||
|
||||
return &TestService{
|
||||
Service: service,
|
||||
Store: testStore,
|
||||
Profile: testProfile,
|
||||
Secret: secret,
|
||||
}
|
||||
}
|
||||
|
||||
// Cleanup closes resources after test.
|
||||
func (ts *TestService) Cleanup() {
|
||||
ts.Store.Close()
|
||||
}
|
||||
|
||||
// CreateHostUser creates an admin user for testing.
|
||||
func (ts *TestService) CreateHostUser(ctx context.Context, username string) (*store.User, error) {
|
||||
return ts.Store.CreateUser(ctx, &store.User{
|
||||
Username: username,
|
||||
Role: store.RoleAdmin,
|
||||
Email: username + "@example.com",
|
||||
})
|
||||
}
|
||||
|
||||
// CreateRegularUser creates a regular user for testing.
|
||||
func (ts *TestService) CreateRegularUser(ctx context.Context, username string) (*store.User, error) {
|
||||
return ts.Store.CreateUser(ctx, &store.User{
|
||||
Username: username,
|
||||
Role: store.RoleUser,
|
||||
Email: username + "@example.com",
|
||||
})
|
||||
}
|
||||
|
||||
// CreateUserContext creates a context with the given user's ID for authentication.
|
||||
func (*TestService) CreateUserContext(ctx context.Context, userID int32) context.Context {
|
||||
// Use the context key from the auth package
|
||||
return context.WithValue(ctx, auth.UserIDContextKey, userID)
|
||||
}
|
||||
@@ -0,0 +1,113 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
apiv1 "github.com/usememos/memos/proto/gen/api/v1"
|
||||
)
|
||||
|
||||
func TestUserEmailVisibility(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("GetUser redacts email for anonymous callers", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "targetuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
got, err := ts.Service.GetUser(ctx, &apiv1.GetUserRequest{
|
||||
Name: "users/targetuser",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, got)
|
||||
require.Equal(t, user.Username, got.Username)
|
||||
require.Empty(t, got.Email)
|
||||
})
|
||||
|
||||
t.Run("GetUser redacts email for other regular users", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
targetUser, err := ts.CreateRegularUser(ctx, "targetuser")
|
||||
require.NoError(t, err)
|
||||
viewer, err := ts.CreateRegularUser(ctx, "vieweruser")
|
||||
require.NoError(t, err)
|
||||
|
||||
viewerCtx := ts.CreateUserContext(ctx, viewer.ID)
|
||||
got, err := ts.Service.GetUser(viewerCtx, &apiv1.GetUserRequest{
|
||||
Name: "users/targetuser",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, got)
|
||||
require.Equal(t, targetUser.Username, got.Username)
|
||||
require.Empty(t, got.Email)
|
||||
})
|
||||
|
||||
t.Run("GetUser returns email for the same user", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "selfuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
got, err := ts.Service.GetUser(userCtx, &apiv1.GetUserRequest{
|
||||
Name: "users/selfuser",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, got)
|
||||
require.Equal(t, user.Email, got.Email)
|
||||
})
|
||||
|
||||
t.Run("GetUser returns email for admins", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
targetUser, err := ts.CreateRegularUser(ctx, "targetuser")
|
||||
require.NoError(t, err)
|
||||
admin, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
|
||||
adminCtx := ts.CreateUserContext(ctx, admin.ID)
|
||||
got, err := ts.Service.GetUser(adminCtx, &apiv1.GetUserRequest{
|
||||
Name: "users/targetuser",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, got)
|
||||
require.Equal(t, targetUser.Email, got.Email)
|
||||
})
|
||||
|
||||
t.Run("GetCurrentUser returns email for the authenticated user", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "currentuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
got, err := ts.Service.GetCurrentUser(userCtx, &apiv1.GetCurrentUserRequest{})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, got)
|
||||
require.NotNil(t, got.User)
|
||||
require.Equal(t, user.Email, got.User.Email)
|
||||
})
|
||||
|
||||
t.Run("GetInstanceProfile redacts admin email for anonymous callers", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
admin, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
|
||||
got, err := ts.Service.GetInstanceProfile(ctx, &apiv1.GetInstanceProfileRequest{})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, got)
|
||||
require.NotNil(t, got.Admin)
|
||||
require.Equal(t, admin.Username, got.Admin.Username)
|
||||
require.Empty(t, got.Admin.Email)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,526 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/protobuf/types/known/fieldmaskpb"
|
||||
|
||||
"github.com/usememos/memos/internal/email"
|
||||
apiv1 "github.com/usememos/memos/proto/gen/api/v1"
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func TestListUserNotificationsIncludesMemoCommentPayload(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
owner, err := ts.CreateRegularUser(ctx, "notification-owner")
|
||||
require.NoError(t, err)
|
||||
ownerCtx := ts.CreateUserContext(ctx, owner.ID)
|
||||
|
||||
commenter, err := ts.CreateRegularUser(ctx, "notification-commenter")
|
||||
require.NoError(t, err)
|
||||
commenterCtx := ts.CreateUserContext(ctx, commenter.ID)
|
||||
|
||||
memo, err := ts.Service.CreateMemo(ownerCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "Base memo",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
comment, err := ts.Service.CreateMemoComment(commenterCtx, &apiv1.CreateMemoCommentRequest{
|
||||
Name: memo.Name,
|
||||
Comment: &apiv1.Memo{
|
||||
Content: "Comment content",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.ListUserNotifications(ownerCtx, &apiv1.ListUserNotificationsRequest{
|
||||
Parent: fmt.Sprintf("users/%s", owner.Username),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, resp.Notifications, 1)
|
||||
|
||||
notification := resp.Notifications[0]
|
||||
require.Contains(t, notification.Name, fmt.Sprintf("users/%s/notifications/", owner.Username))
|
||||
require.Equal(t, fmt.Sprintf("users/%s", commenter.Username), notification.Sender)
|
||||
require.NotNil(t, notification.SenderUser)
|
||||
require.Equal(t, commenter.Username, notification.SenderUser.Username)
|
||||
require.Equal(t, apiv1.UserNotification_MEMO_COMMENT, notification.Type)
|
||||
require.NotNil(t, notification.GetMemoComment())
|
||||
require.Equal(t, comment.Name, notification.GetMemoComment().Memo)
|
||||
require.Equal(t, memo.Name, notification.GetMemoComment().RelatedMemo)
|
||||
require.Equal(t, "Comment content", notification.GetMemoComment().MemoSnippet)
|
||||
require.Equal(t, "Base memo", notification.GetMemoComment().RelatedMemoSnippet)
|
||||
}
|
||||
|
||||
func TestListUserNotificationsStoresMemoCommentPayloadInInbox(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
owner, err := ts.CreateRegularUser(ctx, "notification-owner")
|
||||
require.NoError(t, err)
|
||||
ownerCtx := ts.CreateUserContext(ctx, owner.ID)
|
||||
|
||||
commenter, err := ts.CreateRegularUser(ctx, "notification-commenter")
|
||||
require.NoError(t, err)
|
||||
commenterCtx := ts.CreateUserContext(ctx, commenter.ID)
|
||||
|
||||
memo, err := ts.Service.CreateMemo(ownerCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "Base memo",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.CreateMemoComment(commenterCtx, &apiv1.CreateMemoCommentRequest{
|
||||
Name: memo.Name,
|
||||
Comment: &apiv1.Memo{
|
||||
Content: "Comment content",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
messageType := storepb.InboxMessage_MEMO_COMMENT
|
||||
inboxes, err := ts.Store.ListInboxes(ctx, &store.FindInbox{
|
||||
ReceiverID: &owner.ID,
|
||||
MessageType: &messageType,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, inboxes, 1)
|
||||
require.NotNil(t, inboxes[0].Message)
|
||||
require.NotNil(t, inboxes[0].Message.GetMemoComment())
|
||||
require.NotZero(t, inboxes[0].Message.GetMemoComment().MemoId)
|
||||
require.NotZero(t, inboxes[0].Message.GetMemoComment().RelatedMemoId)
|
||||
}
|
||||
|
||||
func TestCreateMemoCommentSendsEmailNotificationWhenEnabled(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
var sentConfig *email.Config
|
||||
var sentMessage *email.Message
|
||||
ts.Service.NotificationEmailSender = func(config *email.Config, message *email.Message) {
|
||||
sentConfig = config
|
||||
sentMessage = message
|
||||
}
|
||||
|
||||
_, err := ts.Store.UpsertInstanceSetting(ctx, &storepb.InstanceSetting{
|
||||
Key: storepb.InstanceSettingKey_NOTIFICATION,
|
||||
Value: &storepb.InstanceSetting_NotificationSetting{
|
||||
NotificationSetting: &storepb.InstanceNotificationSetting{
|
||||
Email: &storepb.InstanceNotificationSetting_EmailSetting{
|
||||
Enabled: true,
|
||||
SmtpHost: "smtp.example.com",
|
||||
SmtpPort: 587,
|
||||
SmtpUsername: "bot@example.com",
|
||||
SmtpPassword: "password",
|
||||
FromEmail: "bot@example.com",
|
||||
FromName: "Memos",
|
||||
ReplyTo: "reply@example.com",
|
||||
UseTls: true,
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
owner, err := ts.CreateRegularUser(ctx, "email-comment-owner")
|
||||
require.NoError(t, err)
|
||||
ownerCtx := ts.CreateUserContext(ctx, owner.ID)
|
||||
|
||||
commenter, err := ts.CreateRegularUser(ctx, "email-commenter")
|
||||
require.NoError(t, err)
|
||||
commenterCtx := ts.CreateUserContext(ctx, commenter.ID)
|
||||
|
||||
memo, err := ts.Service.CreateMemo(ownerCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "Base memo for email",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
comment, err := ts.Service.CreateMemoComment(commenterCtx, &apiv1.CreateMemoCommentRequest{
|
||||
Name: memo.Name,
|
||||
Comment: &apiv1.Memo{
|
||||
Content: "Email comment content",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NotNil(t, sentConfig)
|
||||
require.Equal(t, "smtp.example.com", sentConfig.SMTPHost)
|
||||
require.Equal(t, 587, sentConfig.SMTPPort)
|
||||
require.Equal(t, "bot@example.com", sentConfig.FromEmail)
|
||||
require.True(t, sentConfig.UseTLS)
|
||||
|
||||
require.NotNil(t, sentMessage)
|
||||
require.Equal(t, []string{owner.Email}, sentMessage.To)
|
||||
require.Equal(t, "reply@example.com", sentMessage.ReplyTo)
|
||||
require.Contains(t, sentMessage.Subject, "commented on your memo")
|
||||
require.Contains(t, sentMessage.Body, "Hi email-comment-owner,")
|
||||
require.Contains(t, sentMessage.Body, "email-commenter commented on your memo.")
|
||||
require.Contains(t, sentMessage.Body, fmt.Sprintf("http://localhost:8080/%s#%s", memo.Name, strings.TrimPrefix(comment.Name, "memos/")))
|
||||
require.Contains(t, sentMessage.Body, "You are receiving this because you own this memo.")
|
||||
require.NotContains(t, sentMessage.Body, "Email comment content")
|
||||
require.NotContains(t, sentMessage.Body, "Base memo for email")
|
||||
}
|
||||
|
||||
func TestCreateMemoMentionSendsEmailNotificationWhenEnabled(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
var sentMessage *email.Message
|
||||
ts.Service.NotificationEmailSender = func(_ *email.Config, message *email.Message) {
|
||||
sentMessage = message
|
||||
}
|
||||
|
||||
_, err := ts.Store.UpsertInstanceSetting(ctx, &storepb.InstanceSetting{
|
||||
Key: storepb.InstanceSettingKey_NOTIFICATION,
|
||||
Value: &storepb.InstanceSetting_NotificationSetting{
|
||||
NotificationSetting: &storepb.InstanceNotificationSetting{
|
||||
Email: &storepb.InstanceNotificationSetting_EmailSetting{
|
||||
Enabled: true,
|
||||
SmtpHost: "smtp.example.com",
|
||||
SmtpPort: 587,
|
||||
FromEmail: "bot@example.com",
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
author, err := ts.CreateRegularUser(ctx, "email-mention-author")
|
||||
require.NoError(t, err)
|
||||
authorCtx := ts.CreateUserContext(ctx, author.ID)
|
||||
|
||||
target, err := ts.CreateRegularUser(ctx, "email-mention-target")
|
||||
require.NoError(t, err)
|
||||
|
||||
memo, err := ts.Service.CreateMemo(authorCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: fmt.Sprintf("Hello @%s from email", target.Username),
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NotNil(t, sentMessage)
|
||||
require.Equal(t, []string{target.Email}, sentMessage.To)
|
||||
require.Contains(t, sentMessage.Subject, "mentioned you in a memo")
|
||||
require.Contains(t, sentMessage.Body, "Hi email-mention-target,")
|
||||
require.Contains(t, sentMessage.Body, "email-mention-author mentioned you in a memo.")
|
||||
require.Contains(t, sentMessage.Body, fmt.Sprintf("http://localhost:8080/%s", memo.Name))
|
||||
require.Contains(t, sentMessage.Body, "You are receiving this because you were mentioned in this memo.")
|
||||
require.NotContains(t, sentMessage.Body, "Hello")
|
||||
require.NotContains(t, sentMessage.Body, fmt.Sprintf("Hello @%s from email", target.Username))
|
||||
}
|
||||
|
||||
func TestCreateMemoCommentSkipsEmailNotificationWithoutInstanceURL(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
ts.Profile.InstanceURL = ""
|
||||
|
||||
var sentMessage *email.Message
|
||||
ts.Service.NotificationEmailSender = func(_ *email.Config, message *email.Message) {
|
||||
sentMessage = message
|
||||
}
|
||||
|
||||
_, err := ts.Store.UpsertInstanceSetting(ctx, &storepb.InstanceSetting{
|
||||
Key: storepb.InstanceSettingKey_NOTIFICATION,
|
||||
Value: &storepb.InstanceSetting_NotificationSetting{
|
||||
NotificationSetting: &storepb.InstanceNotificationSetting{
|
||||
Email: &storepb.InstanceNotificationSetting_EmailSetting{
|
||||
Enabled: true,
|
||||
SmtpHost: "smtp.example.com",
|
||||
SmtpPort: 587,
|
||||
FromEmail: "bot@example.com",
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
owner, err := ts.CreateRegularUser(ctx, "email-comment-no-url-owner")
|
||||
require.NoError(t, err)
|
||||
ownerCtx := ts.CreateUserContext(ctx, owner.ID)
|
||||
|
||||
commenter, err := ts.CreateRegularUser(ctx, "email-comment-no-url-commenter")
|
||||
require.NoError(t, err)
|
||||
commenterCtx := ts.CreateUserContext(ctx, commenter.ID)
|
||||
|
||||
memo, err := ts.Service.CreateMemo(ownerCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "Base memo without instance URL",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.CreateMemoComment(commenterCtx, &apiv1.CreateMemoCommentRequest{
|
||||
Name: memo.Name,
|
||||
Comment: &apiv1.Memo{
|
||||
Content: "Comment without instance URL",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, sentMessage)
|
||||
}
|
||||
|
||||
func TestListUserNotificationsOmitsPayloadWhenMemosDeleted(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
owner, err := ts.CreateRegularUser(ctx, "notification-owner")
|
||||
require.NoError(t, err)
|
||||
ownerCtx := ts.CreateUserContext(ctx, owner.ID)
|
||||
|
||||
commenter, err := ts.CreateRegularUser(ctx, "notification-commenter")
|
||||
require.NoError(t, err)
|
||||
commenterCtx := ts.CreateUserContext(ctx, commenter.ID)
|
||||
|
||||
memo, err := ts.Service.CreateMemo(ownerCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "Base memo",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.CreateMemoComment(commenterCtx, &apiv1.CreateMemoCommentRequest{
|
||||
Name: memo.Name,
|
||||
Comment: &apiv1.Memo{
|
||||
Content: "Comment content",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.DeleteMemo(ownerCtx, &apiv1.DeleteMemoRequest{
|
||||
Name: memo.Name,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.ListUserNotifications(ownerCtx, &apiv1.ListUserNotificationsRequest{
|
||||
Parent: fmt.Sprintf("users/%s", owner.Username),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, resp.Notifications, 1)
|
||||
require.Equal(t, apiv1.UserNotification_MEMO_COMMENT, resp.Notifications[0].Type)
|
||||
require.Nil(t, resp.Notifications[0].GetMemoComment())
|
||||
}
|
||||
|
||||
func TestListUserNotificationsSkipsNotificationsWithMissingUsers(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
owner, err := ts.CreateRegularUser(ctx, "notification-owner")
|
||||
require.NoError(t, err)
|
||||
ownerCtx := ts.CreateUserContext(ctx, owner.ID)
|
||||
|
||||
commenter, err := ts.CreateRegularUser(ctx, "notification-orphan")
|
||||
require.NoError(t, err)
|
||||
commenterCtx := ts.CreateUserContext(ctx, commenter.ID)
|
||||
|
||||
memo, err := ts.Service.CreateMemo(ownerCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "Base memo",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.CreateMemoComment(commenterCtx, &apiv1.CreateMemoCommentRequest{
|
||||
Name: memo.Name,
|
||||
Comment: &apiv1.Memo{
|
||||
Content: "Comment content",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Store.DeleteUser(ctx, &store.DeleteUser{ID: commenter.ID})
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.ListUserNotifications(ownerCtx, &apiv1.ListUserNotificationsRequest{
|
||||
Parent: fmt.Sprintf("users/%s", owner.Username),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, resp.Notifications)
|
||||
}
|
||||
|
||||
func TestListUserNotificationsRejectsNumericParent(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
owner, err := ts.CreateRegularUser(ctx, "notification-owner")
|
||||
require.NoError(t, err)
|
||||
ownerCtx := ts.CreateUserContext(ctx, owner.ID)
|
||||
|
||||
_, err = ts.Service.ListUserNotifications(ownerCtx, &apiv1.ListUserNotificationsRequest{
|
||||
Parent: "users/1",
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "user not found")
|
||||
}
|
||||
|
||||
func TestListUserNotificationsIncludesMemoMentionPayload(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
author, err := ts.CreateRegularUser(ctx, "mention-author")
|
||||
require.NoError(t, err)
|
||||
authorCtx := ts.CreateUserContext(ctx, author.ID)
|
||||
|
||||
target, err := ts.CreateRegularUser(ctx, "mention-target")
|
||||
require.NoError(t, err)
|
||||
targetCtx := ts.CreateUserContext(ctx, target.ID)
|
||||
|
||||
memo, err := ts.Service.CreateMemo(authorCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: fmt.Sprintf("Hello @%s", target.Username),
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.ListUserNotifications(targetCtx, &apiv1.ListUserNotificationsRequest{
|
||||
Parent: fmt.Sprintf("users/%s", target.Username),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, resp.Notifications, 1)
|
||||
require.Equal(t, apiv1.UserNotification_MEMO_MENTION, resp.Notifications[0].Type)
|
||||
require.NotNil(t, resp.Notifications[0].GetMemoMention())
|
||||
require.Equal(t, memo.Name, resp.Notifications[0].GetMemoMention().Memo)
|
||||
require.Empty(t, resp.Notifications[0].GetMemoMention().RelatedMemo)
|
||||
require.Equal(t, author.Username, resp.Notifications[0].SenderUser.Username)
|
||||
require.Equal(t, "Hello", resp.Notifications[0].GetMemoMention().MemoSnippet)
|
||||
}
|
||||
|
||||
func TestCreateMemoCommentMentionDoesNotDuplicateOwnerNotification(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
owner, err := ts.CreateRegularUser(ctx, "mention-owner")
|
||||
require.NoError(t, err)
|
||||
ownerCtx := ts.CreateUserContext(ctx, owner.ID)
|
||||
|
||||
commenter, err := ts.CreateRegularUser(ctx, "mention-commenter")
|
||||
require.NoError(t, err)
|
||||
commenterCtx := ts.CreateUserContext(ctx, commenter.ID)
|
||||
|
||||
memo, err := ts.Service.CreateMemo(ownerCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "Base memo",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.CreateMemoComment(commenterCtx, &apiv1.CreateMemoCommentRequest{
|
||||
Name: memo.Name,
|
||||
Comment: &apiv1.Memo{
|
||||
Content: fmt.Sprintf("Hi @%s", owner.Username),
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.ListUserNotifications(ownerCtx, &apiv1.ListUserNotificationsRequest{
|
||||
Parent: fmt.Sprintf("users/%s", owner.Username),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, resp.Notifications, 1)
|
||||
require.Equal(t, apiv1.UserNotification_MEMO_COMMENT, resp.Notifications[0].Type)
|
||||
}
|
||||
|
||||
func TestUpdateMemoMentionOnlyNotifiesNewTargets(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
author, err := ts.CreateRegularUser(ctx, "mention-update-author")
|
||||
require.NoError(t, err)
|
||||
authorCtx := ts.CreateUserContext(ctx, author.ID)
|
||||
|
||||
firstTarget, err := ts.CreateRegularUser(ctx, "mention-update-first")
|
||||
require.NoError(t, err)
|
||||
firstTargetCtx := ts.CreateUserContext(ctx, firstTarget.ID)
|
||||
|
||||
secondTarget, err := ts.CreateRegularUser(ctx, "mention-update-second")
|
||||
require.NoError(t, err)
|
||||
secondTargetCtx := ts.CreateUserContext(ctx, secondTarget.ID)
|
||||
|
||||
memo, err := ts.Service.CreateMemo(authorCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "",
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
updatedMemo, err := ts.Service.UpdateMemo(authorCtx, &apiv1.UpdateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Name: memo.Name,
|
||||
Content: fmt.Sprintf("Hello @%s", firstTarget.Username),
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
UpdateMask: &fieldmaskpb.FieldMask{Paths: []string{"content"}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
firstResp, err := ts.Service.ListUserNotifications(firstTargetCtx, &apiv1.ListUserNotificationsRequest{
|
||||
Parent: fmt.Sprintf("users/%s", firstTarget.Username),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, firstResp.Notifications, 1)
|
||||
require.Equal(t, apiv1.UserNotification_MEMO_MENTION, firstResp.Notifications[0].Type)
|
||||
require.Equal(t, updatedMemo.Name, firstResp.Notifications[0].GetMemoMention().Memo)
|
||||
|
||||
_, err = ts.Service.UpdateMemo(authorCtx, &apiv1.UpdateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Name: memo.Name,
|
||||
Content: fmt.Sprintf("Hello again @%s and @%s", firstTarget.Username, secondTarget.Username),
|
||||
Visibility: apiv1.Visibility_PUBLIC,
|
||||
},
|
||||
UpdateMask: &fieldmaskpb.FieldMask{Paths: []string{"content"}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
firstResp, err = ts.Service.ListUserNotifications(firstTargetCtx, &apiv1.ListUserNotificationsRequest{
|
||||
Parent: fmt.Sprintf("users/%s", firstTarget.Username),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, firstResp.Notifications, 1)
|
||||
|
||||
secondResp, err := ts.Service.ListUserNotifications(secondTargetCtx, &apiv1.ListUserNotificationsRequest{
|
||||
Parent: fmt.Sprintf("users/%s", secondTarget.Username),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, secondResp.Notifications, 1)
|
||||
require.Equal(t, apiv1.UserNotification_MEMO_MENTION, secondResp.Notifications[0].Type)
|
||||
}
|
||||
@@ -0,0 +1,197 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/protobuf/types/known/fieldmaskpb"
|
||||
|
||||
apiv1 "github.com/usememos/memos/proto/gen/api/v1"
|
||||
apiv1server "github.com/usememos/memos/server/router/api/v1"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func TestUserResourceName(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("GetUser returns username-based canonical name", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
got, err := ts.Service.GetUser(ctx, &apiv1.GetUserRequest{
|
||||
Name: "users/testuser",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, got)
|
||||
require.Equal(t, "users/testuser", got.Name)
|
||||
require.Equal(t, user.Username, got.Username)
|
||||
})
|
||||
|
||||
t.Run("CreateUser returns username-based canonical name", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
created, err := ts.Service.CreateUser(ctx, &apiv1.CreateUserRequest{
|
||||
User: &apiv1.User{
|
||||
Username: "newuser",
|
||||
Email: "newuser@example.com",
|
||||
Password: "password123",
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, created)
|
||||
require.Equal(t, "users/newuser", created.Name)
|
||||
})
|
||||
|
||||
t.Run("Mixed-case username remains usable after auth", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "Gnammi")
|
||||
require.NoError(t, err)
|
||||
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
currentUser, err := ts.Service.GetCurrentUser(userCtx, &apiv1.GetCurrentUserRequest{})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, currentUser.GetUser())
|
||||
require.Equal(t, "users/Gnammi", currentUser.GetUser().Name)
|
||||
|
||||
settings, err := ts.Service.ListUserSettings(userCtx, &apiv1.ListUserSettingsRequest{
|
||||
Parent: currentUser.GetUser().Name,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, settings)
|
||||
|
||||
shortcuts, err := ts.Service.ListShortcuts(userCtx, &apiv1.ListShortcutsRequest{
|
||||
Parent: currentUser.GetUser().Name,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, shortcuts)
|
||||
})
|
||||
|
||||
t.Run("BatchGetUsers preserves mixed-case usernames", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "Gnammi")
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.BatchGetUsers(ctx, &apiv1.BatchGetUsersRequest{
|
||||
Usernames: []string{user.Username},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, resp.Users, 1)
|
||||
require.Equal(t, "users/Gnammi", resp.Users[0].Name)
|
||||
})
|
||||
|
||||
t.Run("CreateUser rejects all-numeric usernames", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
_, err := ts.Service.CreateUser(ctx, &apiv1.CreateUserRequest{
|
||||
User: &apiv1.User{
|
||||
Username: "123",
|
||||
Email: "123@example.com",
|
||||
Password: "password123",
|
||||
},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "invalid username")
|
||||
})
|
||||
|
||||
t.Run("GetUser returns not found for numeric user resource names", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
_, err := ts.CreateRegularUser(ctx, "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.GetUser(ctx, &apiv1.GetUserRequest{
|
||||
Name: "users/1",
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "user not found")
|
||||
})
|
||||
|
||||
t.Run("legacy invalid username remains addressable for get update and delete", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
legacyUser, err := ts.CreateRegularUser(ctx, "legacy_user")
|
||||
require.NoError(t, err)
|
||||
|
||||
got, err := ts.Service.GetUser(ctx, &apiv1.GetUserRequest{
|
||||
Name: "users/legacy_user",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, got)
|
||||
require.Equal(t, "users/legacy_user", got.Name)
|
||||
|
||||
authCtx := ts.CreateUserContext(apiv1server.WithHeaderCarrier(ctx), legacyUser.ID)
|
||||
updated, err := ts.Service.UpdateUser(authCtx, &apiv1.UpdateUserRequest{
|
||||
User: &apiv1.User{
|
||||
Name: apiv1server.BuildUserName(legacyUser.Username),
|
||||
DisplayName: "Legacy User",
|
||||
},
|
||||
UpdateMask: &fieldmaskpb.FieldMask{Paths: []string{"display_name"}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "Legacy User", updated.DisplayName)
|
||||
|
||||
_, err = ts.Service.DeleteUser(authCtx, &apiv1.DeleteUserRequest{
|
||||
Name: apiv1server.BuildUserName(legacyUser.Username),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
deleted, err := ts.Store.GetUser(ctx, &store.FindUser{ID: &legacyUser.ID})
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, deleted)
|
||||
})
|
||||
|
||||
t.Run("email-like legacy username can be renamed to a valid username", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
legacyUser, err := ts.CreateRegularUser(ctx, "alice@example.com")
|
||||
require.NoError(t, err)
|
||||
|
||||
authCtx := ts.CreateUserContext(apiv1server.WithHeaderCarrier(ctx), legacyUser.ID)
|
||||
updated, err := ts.Service.UpdateUser(authCtx, &apiv1.UpdateUserRequest{
|
||||
User: &apiv1.User{
|
||||
Name: apiv1server.BuildUserName(legacyUser.Username),
|
||||
Username: "alice",
|
||||
},
|
||||
UpdateMask: &fieldmaskpb.FieldMask{Paths: []string{"username"}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "users/alice", updated.Name)
|
||||
require.Equal(t, "alice", updated.Username)
|
||||
|
||||
renamed, err := ts.Store.GetUser(ctx, &store.FindUser{ID: &legacyUser.ID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, renamed)
|
||||
require.Equal(t, "alice", renamed.Username)
|
||||
})
|
||||
|
||||
t.Run("email-like legacy username can be deleted", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
legacyUser, err := ts.CreateRegularUser(ctx, "bob@example.com")
|
||||
require.NoError(t, err)
|
||||
|
||||
authCtx := ts.CreateUserContext(apiv1server.WithHeaderCarrier(ctx), legacyUser.ID)
|
||||
_, err = ts.Service.DeleteUser(authCtx, &apiv1.DeleteUserRequest{
|
||||
Name: apiv1server.BuildUserName(legacyUser.Username),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
deleted, err := ts.Store.GetUser(ctx, &store.FindUser{ID: &legacyUser.ID})
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, deleted)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,71 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
apiv1 "github.com/usememos/memos/proto/gen/api/v1"
|
||||
)
|
||||
|
||||
func TestBatchGetUsersReturnsExactUsernamesWithoutAuthentication(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
_, err := ts.CreateRegularUser(ctx, "batch-alpha")
|
||||
require.NoError(t, err)
|
||||
_, err = ts.CreateRegularUser(ctx, "batch-beta")
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.BatchGetUsers(ctx, &apiv1.BatchGetUsersRequest{
|
||||
Usernames: []string{"batch-alpha", "batch-beta", "missing-user", "batch-alpha"},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, resp.Users, 2)
|
||||
|
||||
got := map[string]struct{}{}
|
||||
for _, user := range resp.Users {
|
||||
got[user.Username] = struct{}{}
|
||||
}
|
||||
_, ok := got["batch-alpha"]
|
||||
require.True(t, ok)
|
||||
_, ok = got["batch-beta"]
|
||||
require.True(t, ok)
|
||||
}
|
||||
|
||||
func TestBatchGetUsersRejectsTooManyUsernames(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
usernames := make([]string, 0, 101)
|
||||
for i := range 101 {
|
||||
usernames = append(usernames, fmt.Sprintf("user-%d", i))
|
||||
}
|
||||
|
||||
_, err := ts.Service.BatchGetUsers(ctx, &apiv1.BatchGetUsersRequest{
|
||||
Usernames: usernames,
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "too many usernames")
|
||||
}
|
||||
|
||||
func TestBatchGetUsersRejectsTooManyNonEmptyUsernamesBeforeDedupe(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
usernames := make([]string, 0, 101)
|
||||
for range 101 {
|
||||
usernames = append(usernames, "legacy@example.com")
|
||||
}
|
||||
|
||||
_, err := ts.Service.BatchGetUsers(ctx, &apiv1.BatchGetUsersRequest{
|
||||
Usernames: usernames,
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "too many usernames")
|
||||
}
|
||||
@@ -0,0 +1,489 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/protobuf/types/known/timestamppb"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
apiv1 "github.com/usememos/memos/server/router/api/v1"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func TestDeleteUserSelfDeleteCleansAccountDataAndAuthCookies(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
user, err := ts.CreateRegularUser(ctx, "alice")
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Store.CreateUserIdentity(ctx, &store.UserIdentity{
|
||||
UserID: user.ID,
|
||||
Provider: "google",
|
||||
ExternUID: "alice-google-sub",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
err = ts.Store.AddUserRefreshToken(ctx, user.ID, &storepb.RefreshTokensUserSetting_RefreshToken{
|
||||
TokenId: "refresh-token-id",
|
||||
ExpiresAt: timestamppb.New(time.Now().Add(time.Hour)),
|
||||
CreatedAt: timestamppb.Now(),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
headerCtx := apiv1.WithHeaderCarrier(ctx)
|
||||
authCtx := ts.CreateUserContext(headerCtx, user.ID)
|
||||
_, err = ts.Service.DeleteUser(authCtx, &v1pb.DeleteUserRequest{
|
||||
Name: apiv1.BuildUserName(user.Username),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
deletedUser, err := ts.Store.GetUser(ctx, &store.FindUser{ID: &user.ID})
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, deletedUser)
|
||||
|
||||
identities, err := ts.Store.ListUserIdentities(ctx, &store.FindUserIdentity{UserID: &user.ID})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, identities)
|
||||
|
||||
refreshSetting, err := ts.Store.GetUserSetting(ctx, &store.FindUserSetting{
|
||||
UserID: &user.ID,
|
||||
Key: storepb.UserSetting_REFRESH_TOKENS,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, refreshSetting)
|
||||
|
||||
carrier := apiv1.GetHeaderCarrier(authCtx)
|
||||
require.NotNil(t, carrier)
|
||||
require.Contains(t, strings.ToLower(carrier.Get("Set-Cookie")), "memos_refresh=")
|
||||
}
|
||||
|
||||
func TestDeleteUserSelfDeleteRemovesOwnedResourcesAndMemoSubtrees(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
user, err := ts.CreateRegularUser(ctx, "resource-owner")
|
||||
require.NoError(t, err)
|
||||
peer, err := ts.CreateRegularUser(ctx, "resource-peer")
|
||||
require.NoError(t, err)
|
||||
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
peerCtx := ts.CreateUserContext(ctx, peer.ID)
|
||||
|
||||
ownMemo, err := ts.Service.CreateMemo(userCtx, &v1pb.CreateMemoRequest{
|
||||
Memo: &v1pb.Memo{
|
||||
Content: "owner memo",
|
||||
Visibility: v1pb.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
foreignMemo, err := ts.Service.CreateMemo(peerCtx, &v1pb.CreateMemoRequest{
|
||||
Memo: &v1pb.Memo{
|
||||
Content: "peer memo",
|
||||
Visibility: v1pb.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
peerCommentOnOwnMemo, err := ts.Service.CreateMemoComment(peerCtx, &v1pb.CreateMemoCommentRequest{
|
||||
Name: ownMemo.Name,
|
||||
Comment: &v1pb.Memo{
|
||||
Content: "peer comment on owner memo",
|
||||
Visibility: v1pb.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
peerNestedCommentOnOwnMemo, err := ts.Service.CreateMemoComment(peerCtx, &v1pb.CreateMemoCommentRequest{
|
||||
Name: peerCommentOnOwnMemo.Name,
|
||||
Comment: &v1pb.Memo{
|
||||
Content: "peer nested comment on owner memo",
|
||||
Visibility: v1pb.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
userCommentOnForeignMemo, err := ts.Service.CreateMemoComment(userCtx, &v1pb.CreateMemoCommentRequest{
|
||||
Name: foreignMemo.Name,
|
||||
Comment: &v1pb.Memo{
|
||||
Content: "owner comment on peer memo",
|
||||
Visibility: v1pb.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
peerReplyToUserComment, err := ts.Service.CreateMemoComment(peerCtx, &v1pb.CreateMemoCommentRequest{
|
||||
Name: userCommentOnForeignMemo.Name,
|
||||
Comment: &v1pb.Memo{
|
||||
Content: "peer reply to owner comment",
|
||||
Visibility: v1pb.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
ownMemoUID, err := apiv1.ExtractMemoUIDFromName(ownMemo.Name)
|
||||
require.NoError(t, err)
|
||||
ownMemoStore, err := ts.Store.GetMemo(ctx, &store.FindMemo{UID: &ownMemoUID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, ownMemoStore)
|
||||
|
||||
foreignMemoUID, err := apiv1.ExtractMemoUIDFromName(foreignMemo.Name)
|
||||
require.NoError(t, err)
|
||||
foreignMemoStore, err := ts.Store.GetMemo(ctx, &store.FindMemo{UID: &foreignMemoUID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, foreignMemoStore)
|
||||
|
||||
attachedAttachment, err := ts.Store.CreateAttachment(ctx, &store.Attachment{
|
||||
UID: "attach-owner-memo",
|
||||
CreatorID: user.ID,
|
||||
Filename: "owner.txt",
|
||||
Type: "text/plain",
|
||||
Size: 4,
|
||||
Blob: []byte("memo"),
|
||||
MemoID: &ownMemoStore.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
thumbnailCachePath := filepath.Join(ts.Profile.Data, ".thumbnail_cache", attachedAttachment.UID+".jpeg")
|
||||
motionCachePath := filepath.Join(ts.Profile.Data, ".motion_cache", attachedAttachment.UID+".mp4")
|
||||
require.NoError(t, os.MkdirAll(filepath.Dir(thumbnailCachePath), 0o755))
|
||||
require.NoError(t, os.WriteFile(thumbnailCachePath, []byte("thumb"), 0o644))
|
||||
require.NoError(t, os.MkdirAll(filepath.Dir(motionCachePath), 0o755))
|
||||
require.NoError(t, os.WriteFile(motionCachePath, []byte("motion"), 0o644))
|
||||
|
||||
unattachedAttachment, err := ts.Store.CreateAttachment(ctx, &store.Attachment{
|
||||
UID: "attach-owner-loose",
|
||||
CreatorID: user.ID,
|
||||
Filename: "loose.txt",
|
||||
Type: "text/plain",
|
||||
Size: 5,
|
||||
Blob: []byte("loose"),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
peerAttachment, err := ts.Store.CreateAttachment(ctx, &store.Attachment{
|
||||
UID: "attach-peer-keep",
|
||||
CreatorID: peer.ID,
|
||||
Filename: "peer.txt",
|
||||
Type: "text/plain",
|
||||
Size: 4,
|
||||
Blob: []byte("peer"),
|
||||
MemoID: &foreignMemoStore.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Store.UpsertReaction(ctx, &store.Reaction{
|
||||
CreatorID: peer.ID,
|
||||
ContentID: ownMemo.Name,
|
||||
ReactionType: "👍",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
_, err = ts.Store.UpsertReaction(ctx, &store.Reaction{
|
||||
CreatorID: peer.ID,
|
||||
ContentID: userCommentOnForeignMemo.Name,
|
||||
ReactionType: "🔥",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
_, err = ts.Store.UpsertReaction(ctx, &store.Reaction{
|
||||
CreatorID: user.ID,
|
||||
ContentID: foreignMemo.Name,
|
||||
ReactionType: "👋",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
peerReactionOnForeignMemo, err := ts.Store.UpsertReaction(ctx, &store.Reaction{
|
||||
CreatorID: peer.ID,
|
||||
ContentID: foreignMemo.Name,
|
||||
ReactionType: "✅",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Store.CreateMemoShare(ctx, &store.MemoShare{
|
||||
UID: "share-owner-ownmemo",
|
||||
MemoID: ownMemoStore.ID,
|
||||
CreatorID: user.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
_, err = ts.Store.CreateMemoShare(ctx, &store.MemoShare{
|
||||
UID: "share-owner-foreignmemo",
|
||||
MemoID: foreignMemoStore.ID,
|
||||
CreatorID: user.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
peerShare, err := ts.Store.CreateMemoShare(ctx, &store.MemoShare{
|
||||
UID: "share-peer-foreignmemo",
|
||||
MemoID: foreignMemoStore.ID,
|
||||
CreatorID: peer.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Store.CreateUserIdentity(ctx, &store.UserIdentity{
|
||||
UserID: user.ID,
|
||||
Provider: "google",
|
||||
ExternUID: "resource-owner-google-sub",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
err = ts.Store.AddUserPersonalAccessToken(ctx, user.ID, &storepb.PersonalAccessTokensUserSetting_PersonalAccessToken{
|
||||
TokenId: "pat-owner",
|
||||
TokenHash: "pat-owner-hash",
|
||||
Description: "owner pat",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
headerCtx := apiv1.WithHeaderCarrier(ctx)
|
||||
authCtx := ts.CreateUserContext(headerCtx, user.ID)
|
||||
_, err = ts.Service.DeleteUser(authCtx, &v1pb.DeleteUserRequest{
|
||||
Name: apiv1.BuildUserName(user.Username),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
deletedUser, err := ts.Store.GetUser(ctx, &store.FindUser{ID: &user.ID})
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, deletedUser)
|
||||
|
||||
for _, memoName := range []string{
|
||||
ownMemo.Name,
|
||||
peerCommentOnOwnMemo.Name,
|
||||
peerNestedCommentOnOwnMemo.Name,
|
||||
userCommentOnForeignMemo.Name,
|
||||
peerReplyToUserComment.Name,
|
||||
} {
|
||||
memoUID, extractErr := apiv1.ExtractMemoUIDFromName(memoName)
|
||||
require.NoError(t, extractErr)
|
||||
memo, getErr := ts.Store.GetMemo(ctx, &store.FindMemo{UID: &memoUID})
|
||||
require.NoError(t, getErr)
|
||||
require.Nil(t, memo, memoName)
|
||||
}
|
||||
|
||||
foreignMemoAfterDelete, err := ts.Store.GetMemo(ctx, &store.FindMemo{UID: &foreignMemoUID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, foreignMemoAfterDelete)
|
||||
|
||||
for _, attachmentID := range []int32{attachedAttachment.ID, unattachedAttachment.ID} {
|
||||
attachment, getErr := ts.Store.GetAttachment(ctx, &store.FindAttachment{ID: &attachmentID})
|
||||
require.NoError(t, getErr)
|
||||
require.Nil(t, attachment)
|
||||
}
|
||||
peerAttachmentAfterDelete, err := ts.Store.GetAttachment(ctx, &store.FindAttachment{ID: &peerAttachment.ID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, peerAttachmentAfterDelete)
|
||||
_, err = os.Stat(thumbnailCachePath)
|
||||
require.ErrorIs(t, err, os.ErrNotExist)
|
||||
_, err = os.Stat(motionCachePath)
|
||||
require.ErrorIs(t, err, os.ErrNotExist)
|
||||
|
||||
ownMemoReactions, err := ts.Store.ListReactions(ctx, &store.FindReaction{ContentID: &ownMemo.Name})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, ownMemoReactions)
|
||||
|
||||
userCommentReactions, err := ts.Store.ListReactions(ctx, &store.FindReaction{ContentID: &userCommentOnForeignMemo.Name})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, userCommentReactions)
|
||||
|
||||
foreignMemoReactions, err := ts.Store.ListReactions(ctx, &store.FindReaction{ContentID: &foreignMemo.Name})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, foreignMemoReactions, 1)
|
||||
require.Equal(t, peerReactionOnForeignMemo.ID, foreignMemoReactions[0].ID)
|
||||
|
||||
ownerShares, err := ts.Store.ListMemoShares(ctx, &store.FindMemoShare{CreatorID: &user.ID})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, ownerShares)
|
||||
|
||||
peerShares, err := ts.Store.ListMemoShares(ctx, &store.FindMemoShare{CreatorID: &peer.ID})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, peerShares, 1)
|
||||
require.Equal(t, peerShare.ID, peerShares[0].ID)
|
||||
|
||||
sentInboxes, err := ts.Store.ListInboxes(ctx, &store.FindInbox{SenderID: &user.ID})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, sentInboxes)
|
||||
receivedInboxes, err := ts.Store.ListInboxes(ctx, &store.FindInbox{ReceiverID: &user.ID})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, receivedInboxes)
|
||||
|
||||
identities, err := ts.Store.ListUserIdentities(ctx, &store.FindUserIdentity{UserID: &user.ID})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, identities)
|
||||
|
||||
patSetting, err := ts.Store.GetUserSetting(ctx, &store.FindUserSetting{
|
||||
UserID: &user.ID,
|
||||
Key: storepb.UserSetting_PERSONAL_ACCESS_TOKENS,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, patSetting)
|
||||
}
|
||||
|
||||
func TestDeleteUserRollbackPreservesAllResources(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
user, err := ts.CreateRegularUser(ctx, "rollback-owner")
|
||||
require.NoError(t, err)
|
||||
peer, err := ts.CreateRegularUser(ctx, "rollback-peer")
|
||||
require.NoError(t, err)
|
||||
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
ownMemo, err := ts.Service.CreateMemo(userCtx, &v1pb.CreateMemoRequest{
|
||||
Memo: &v1pb.Memo{
|
||||
Content: "rollback owner memo",
|
||||
Visibility: v1pb.Visibility_PUBLIC,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
ownMemoUID, err := apiv1.ExtractMemoUIDFromName(ownMemo.Name)
|
||||
require.NoError(t, err)
|
||||
ownMemoStore, err := ts.Store.GetMemo(ctx, &store.FindMemo{UID: &ownMemoUID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, ownMemoStore)
|
||||
|
||||
attachment, err := ts.Store.CreateAttachment(ctx, &store.Attachment{
|
||||
UID: "attach-rollback-owner",
|
||||
CreatorID: user.ID,
|
||||
Filename: "rollback.txt",
|
||||
Type: "text/plain",
|
||||
Size: 8,
|
||||
Blob: []byte("rollback"),
|
||||
MemoID: &ownMemoStore.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
reaction, err := ts.Store.UpsertReaction(ctx, &store.Reaction{
|
||||
CreatorID: user.ID,
|
||||
ContentID: ownMemo.Name,
|
||||
ReactionType: "💥",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
share, err := ts.Store.CreateMemoShare(ctx, &store.MemoShare{
|
||||
UID: "share-rollback-owner",
|
||||
MemoID: ownMemoStore.ID,
|
||||
CreatorID: user.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
inbox, err := ts.Store.CreateInbox(ctx, &store.Inbox{
|
||||
SenderID: peer.ID,
|
||||
ReceiverID: user.ID,
|
||||
Status: store.UNREAD,
|
||||
Message: &storepb.InboxMessage{
|
||||
Type: storepb.InboxMessage_MEMO_COMMENT,
|
||||
Payload: &storepb.InboxMessage_MemoComment{
|
||||
MemoComment: &storepb.InboxMessage_MemoCommentPayload{
|
||||
MemoId: ownMemoStore.ID,
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Store.CreateUserIdentity(ctx, &store.UserIdentity{
|
||||
UserID: user.ID,
|
||||
Provider: "google",
|
||||
ExternUID: "rollback-owner-google-sub",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
err = ts.Store.AddUserPersonalAccessToken(ctx, user.ID, &storepb.PersonalAccessTokensUserSetting_PersonalAccessToken{
|
||||
TokenId: "pat-rollback-owner",
|
||||
TokenHash: "pat-rollback-owner-hash",
|
||||
Description: "rollback pat",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
headerCtx := apiv1.WithHeaderCarrier(ctx)
|
||||
failCtx := store.WithDeleteUserFailpoint(headerCtx, store.DeleteUserFailpointBeforeCommit)
|
||||
authCtx := ts.CreateUserContext(failCtx, user.ID)
|
||||
_, err = ts.Service.DeleteUser(authCtx, &v1pb.DeleteUserRequest{
|
||||
Name: apiv1.BuildUserName(user.Username),
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.ErrorContains(t, err, "delete user failpoint before commit")
|
||||
|
||||
userAfterRollback, err := ts.Store.GetUser(ctx, &store.FindUser{ID: &user.ID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, userAfterRollback)
|
||||
|
||||
memoAfterRollback, err := ts.Store.GetMemo(ctx, &store.FindMemo{UID: &ownMemoUID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memoAfterRollback)
|
||||
|
||||
attachmentAfterRollback, err := ts.Store.GetAttachment(ctx, &store.FindAttachment{ID: &attachment.ID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, attachmentAfterRollback)
|
||||
|
||||
reactionAfterRollback, err := ts.Store.GetReaction(ctx, &store.FindReaction{ID: &reaction.ID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, reactionAfterRollback)
|
||||
|
||||
shareAfterRollback, err := ts.Store.GetMemoShare(ctx, &store.FindMemoShare{ID: &share.ID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, shareAfterRollback)
|
||||
|
||||
inboxesAfterRollback, err := ts.Store.ListInboxes(ctx, &store.FindInbox{ID: &inbox.ID})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, inboxesAfterRollback, 1)
|
||||
|
||||
identitiesAfterRollback, err := ts.Store.ListUserIdentities(ctx, &store.FindUserIdentity{UserID: &user.ID})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, identitiesAfterRollback, 1)
|
||||
|
||||
patSetting, err := ts.Store.GetUserSetting(ctx, &store.FindUserSetting{
|
||||
UserID: &user.ID,
|
||||
Key: storepb.UserSetting_PERSONAL_ACCESS_TOKENS,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, patSetting)
|
||||
}
|
||||
|
||||
func TestDeleteUserReturnsErrorWhenAttachmentStorageCleanupFails(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
user, err := ts.CreateRegularUser(ctx, "cleanup-failure-owner")
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Store.CreateAttachment(ctx, &store.Attachment{
|
||||
UID: "attach-cleanup-failure",
|
||||
CreatorID: user.ID,
|
||||
Filename: "failure.txt",
|
||||
Type: "text/plain",
|
||||
Size: 7,
|
||||
Blob: []byte("failure"),
|
||||
StorageType: storepb.AttachmentStorageType_LOCAL,
|
||||
Reference: "cleanup-failure.txt",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
headerCtx := apiv1.WithHeaderCarrier(ctx)
|
||||
failCtx := store.WithDeleteAttachmentStorageFailpoint(headerCtx)
|
||||
authCtx := ts.CreateUserContext(failCtx, user.ID)
|
||||
_, err = ts.Service.DeleteUser(authCtx, &v1pb.DeleteUserRequest{
|
||||
Name: apiv1.BuildUserName(user.Username),
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.ErrorContains(t, err, "attachment storage cleanup failed")
|
||||
require.ErrorContains(t, err, "attachment_id=")
|
||||
require.ErrorContains(t, err, store.ErrDeleteAttachmentStorageFailpoint.Error())
|
||||
|
||||
deletedUser, err := ts.Store.GetUser(ctx, &store.FindUser{ID: &user.ID})
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, deletedUser)
|
||||
}
|
||||
@@ -0,0 +1,259 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
"google.golang.org/protobuf/types/known/fieldmaskpb"
|
||||
|
||||
apiv1 "github.com/usememos/memos/proto/gen/api/v1"
|
||||
apiv1server "github.com/usememos/memos/server/router/api/v1"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func TestUserServiceWithEmailLikeUsername(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("SignIn accepts email-like legacy username", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user := createLegacyPasswordUser(ctx, t, ts, "signin@example.com", "password123")
|
||||
|
||||
signInCtx := apiv1server.WithHeaderCarrier(ctx)
|
||||
resp, err := ts.Service.SignIn(signInCtx, &apiv1.SignInRequest{
|
||||
Credentials: &apiv1.SignInRequest_PasswordCredentials_{
|
||||
PasswordCredentials: &apiv1.SignInRequest_PasswordCredentials{
|
||||
Username: user.Username,
|
||||
Password: "password123",
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, user.Username, resp.User.Username)
|
||||
require.NotEmpty(t, resp.AccessToken)
|
||||
})
|
||||
|
||||
t.Run("GetUser accepts email-like username in resource name", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "alice@example.com")
|
||||
require.NoError(t, err)
|
||||
|
||||
got, err := ts.Service.GetUser(ctx, &apiv1.GetUserRequest{
|
||||
Name: "users/alice@example.com",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, got)
|
||||
require.Equal(t, user.Username, got.Username)
|
||||
require.Equal(t, "users/alice@example.com", got.Name)
|
||||
})
|
||||
|
||||
t.Run("BatchGetUsers accepts email-like legacy username", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "batch@example.com")
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.BatchGetUsers(ctx, &apiv1.BatchGetUsersRequest{
|
||||
Usernames: []string{" batch@example.com ", "missing@example.com", "batch@example.com"},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, resp.Users, 1)
|
||||
require.Equal(t, user.Username, resp.Users[0].Username)
|
||||
require.Equal(t, "users/batch@example.com", resp.Users[0].Name)
|
||||
})
|
||||
|
||||
t.Run("BatchGetUsers accepts underscore legacy username", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "legacy_batch")
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.BatchGetUsers(ctx, &apiv1.BatchGetUsersRequest{
|
||||
Usernames: []string{"legacy_batch"},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, resp.Users, 1)
|
||||
require.Equal(t, user.Username, resp.Users[0].Username)
|
||||
require.Equal(t, "users/legacy_batch", resp.Users[0].Name)
|
||||
})
|
||||
|
||||
t.Run("ListUserSettings accepts email-like username in parent", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "alice@example.com")
|
||||
require.NoError(t, err)
|
||||
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
resp, err := ts.Service.ListUserSettings(userCtx, &apiv1.ListUserSettingsRequest{
|
||||
Parent: "users/alice@example.com",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp)
|
||||
require.NotEmpty(t, resp.Settings)
|
||||
})
|
||||
|
||||
t.Run("UpdateUser can change non-username fields for email-like username", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "alice@example.com")
|
||||
require.NoError(t, err)
|
||||
|
||||
authCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
updated, err := ts.Service.UpdateUser(authCtx, &apiv1.UpdateUserRequest{
|
||||
User: &apiv1.User{
|
||||
Name: "users/alice@example.com",
|
||||
DisplayName: "Alice Example",
|
||||
},
|
||||
UpdateMask: &fieldmaskpb.FieldMask{Paths: []string{"display_name"}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "Alice Example", updated.DisplayName)
|
||||
require.Equal(t, "users/alice@example.com", updated.Name)
|
||||
})
|
||||
|
||||
t.Run("UpdateUser can rename email-like username to valid username", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "bob@example.com")
|
||||
require.NoError(t, err)
|
||||
|
||||
authCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
updated, err := ts.Service.UpdateUser(authCtx, &apiv1.UpdateUserRequest{
|
||||
User: &apiv1.User{
|
||||
Name: "users/bob@example.com",
|
||||
Username: "bob",
|
||||
},
|
||||
UpdateMask: &fieldmaskpb.FieldMask{Paths: []string{"username"}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "bob", updated.Username)
|
||||
require.Equal(t, apiv1server.BuildUserName("bob"), updated.Name)
|
||||
|
||||
stored, err := ts.Store.GetUser(ctx, &store.FindUser{ID: &user.ID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, stored)
|
||||
require.Equal(t, "bob", stored.Username)
|
||||
})
|
||||
|
||||
t.Run("UpdateUser rejects writing invalid username values", func(t *testing.T) {
|
||||
for _, username := range []string{"alice@example.com", "legacy_user"} {
|
||||
t.Run(username, func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "rename@example.com")
|
||||
require.NoError(t, err)
|
||||
|
||||
authCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
_, err = ts.Service.UpdateUser(authCtx, &apiv1.UpdateUserRequest{
|
||||
User: &apiv1.User{
|
||||
Name: "users/rename@example.com",
|
||||
Username: username,
|
||||
},
|
||||
UpdateMask: &fieldmaskpb.FieldMask{Paths: []string{"username"}},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "invalid username")
|
||||
|
||||
stored, err := ts.Store.GetUser(ctx, &store.FindUser{ID: &user.ID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, stored)
|
||||
require.Equal(t, "rename@example.com", stored.Username)
|
||||
})
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("admin cannot rename user to invalid username", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "admin-rename-target")
|
||||
require.NoError(t, err)
|
||||
admin, err := ts.CreateHostUser(ctx, "rename-admin")
|
||||
require.NoError(t, err)
|
||||
|
||||
adminCtx := ts.CreateUserContext(ctx, admin.ID)
|
||||
_, err = ts.Service.UpdateUser(adminCtx, &apiv1.UpdateUserRequest{
|
||||
User: &apiv1.User{
|
||||
Name: apiv1server.BuildUserName(user.Username),
|
||||
Username: "admin@example.com",
|
||||
},
|
||||
UpdateMask: &fieldmaskpb.FieldMask{Paths: []string{"username"}},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "invalid username")
|
||||
|
||||
stored, err := ts.Store.GetUser(ctx, &store.FindUser{ID: &user.ID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, stored)
|
||||
require.Equal(t, "admin-rename-target", stored.Username)
|
||||
})
|
||||
|
||||
t.Run("UpdateUser can archive email-like username account", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "dave@example.com")
|
||||
require.NoError(t, err)
|
||||
admin, err := ts.CreateHostUser(ctx, "email-admin")
|
||||
require.NoError(t, err)
|
||||
|
||||
authCtx := ts.CreateUserContext(ctx, admin.ID)
|
||||
updated, err := ts.Service.UpdateUser(authCtx, &apiv1.UpdateUserRequest{
|
||||
User: &apiv1.User{
|
||||
Name: "users/dave@example.com",
|
||||
State: apiv1.State_ARCHIVED,
|
||||
},
|
||||
UpdateMask: &fieldmaskpb.FieldMask{Paths: []string{"state"}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, apiv1.State_ARCHIVED, updated.State)
|
||||
|
||||
stored, err := ts.Store.GetUser(ctx, &store.FindUser{ID: &user.ID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, stored)
|
||||
require.Equal(t, store.Archived, stored.RowStatus)
|
||||
})
|
||||
|
||||
t.Run("DeleteUser can remove email-like username account", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "carol@example.com")
|
||||
require.NoError(t, err)
|
||||
|
||||
authCtx := ts.CreateUserContext(apiv1server.WithHeaderCarrier(ctx), user.ID)
|
||||
_, err = ts.Service.DeleteUser(authCtx, &apiv1.DeleteUserRequest{
|
||||
Name: "users/carol@example.com",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
deleted, err := ts.Store.GetUser(ctx, &store.FindUser{ID: &user.ID})
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, deleted)
|
||||
})
|
||||
}
|
||||
|
||||
func createLegacyPasswordUser(ctx context.Context, t *testing.T, ts *TestService, username, password string) *store.User {
|
||||
passwordHash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
|
||||
require.NoError(t, err)
|
||||
|
||||
user, err := ts.Store.CreateUser(ctx, &store.User{
|
||||
Username: username,
|
||||
Role: store.RoleUser,
|
||||
Email: username,
|
||||
PasswordHash: string(passwordHash),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
return user
|
||||
}
|
||||
@@ -0,0 +1,387 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/protobuf/types/known/fieldmaskpb"
|
||||
|
||||
apiv1 "github.com/usememos/memos/proto/gen/api/v1"
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
apiv1server "github.com/usememos/memos/server/router/api/v1"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func TestCreateUserRegistration(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("CreateUser success when registration enabled", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// User registration is enabled by default, no need to set it explicitly
|
||||
|
||||
// Create user without authentication - should succeed
|
||||
_, err := ts.Service.CreateUser(ctx, &apiv1.CreateUserRequest{
|
||||
User: &apiv1.User{
|
||||
Username: "newuser",
|
||||
Email: "newuser@example.com",
|
||||
Password: "password123",
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
})
|
||||
|
||||
t.Run("CreateUser blocked when registration disabled", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create a host user first so we're not in first-user setup mode
|
||||
_, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Disable user registration
|
||||
_, err = ts.Store.UpsertInstanceSetting(ctx, &storepb.InstanceSetting{
|
||||
Key: storepb.InstanceSettingKey_GENERAL,
|
||||
Value: &storepb.InstanceSetting_GeneralSetting{
|
||||
GeneralSetting: &storepb.InstanceGeneralSetting{
|
||||
DisallowUserRegistration: true,
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Try to create user without authentication - should fail
|
||||
_, err = ts.Service.CreateUser(ctx, &apiv1.CreateUserRequest{
|
||||
User: &apiv1.User{
|
||||
Username: "newuser",
|
||||
Email: "newuser@example.com",
|
||||
Password: "password123",
|
||||
},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "not allowed")
|
||||
})
|
||||
|
||||
t.Run("CreateUser succeeds for superuser even when registration disabled", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create host user
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
hostCtx := ts.CreateUserContext(ctx, hostUser.ID)
|
||||
|
||||
// Disable user registration
|
||||
_, err = ts.Store.UpsertInstanceSetting(ctx, &storepb.InstanceSetting{
|
||||
Key: storepb.InstanceSettingKey_GENERAL,
|
||||
Value: &storepb.InstanceSetting_GeneralSetting{
|
||||
GeneralSetting: &storepb.InstanceGeneralSetting{
|
||||
DisallowUserRegistration: true,
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Host user can create users even when registration is disabled - should succeed
|
||||
_, err = ts.Service.CreateUser(hostCtx, &apiv1.CreateUserRequest{
|
||||
User: &apiv1.User{
|
||||
Username: "newuser",
|
||||
Email: "newuser@example.com",
|
||||
Password: "password123",
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
})
|
||||
|
||||
t.Run("CreateUser regular user cannot create users when registration disabled", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create regular user
|
||||
regularUser, err := ts.CreateRegularUser(ctx, "regularuser")
|
||||
require.NoError(t, err)
|
||||
regularUserCtx := ts.CreateUserContext(ctx, regularUser.ID)
|
||||
|
||||
// Disable user registration
|
||||
_, err = ts.Store.UpsertInstanceSetting(ctx, &storepb.InstanceSetting{
|
||||
Key: storepb.InstanceSettingKey_GENERAL,
|
||||
Value: &storepb.InstanceSetting_GeneralSetting{
|
||||
GeneralSetting: &storepb.InstanceGeneralSetting{
|
||||
DisallowUserRegistration: true,
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Regular user tries to create user when registration is disabled - should fail
|
||||
_, err = ts.Service.CreateUser(regularUserCtx, &apiv1.CreateUserRequest{
|
||||
User: &apiv1.User{
|
||||
Username: "newuser",
|
||||
Email: "newuser@example.com",
|
||||
Password: "password123",
|
||||
},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "not allowed")
|
||||
})
|
||||
|
||||
t.Run("CreateUser host can assign roles", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create host user
|
||||
hostUser, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
hostCtx := ts.CreateUserContext(ctx, hostUser.ID)
|
||||
|
||||
// Host user can create user with specific role - should succeed
|
||||
createdUser, err := ts.Service.CreateUser(hostCtx, &apiv1.CreateUserRequest{
|
||||
User: &apiv1.User{
|
||||
Username: "newadmin",
|
||||
Email: "newadmin@example.com",
|
||||
Password: "password123",
|
||||
Role: apiv1.User_ADMIN,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "users/newadmin", createdUser.Name)
|
||||
require.NotNil(t, createdUser)
|
||||
require.Equal(t, apiv1.User_ADMIN, createdUser.Role)
|
||||
})
|
||||
|
||||
t.Run("CreateUser unauthenticated user can only create regular user", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create a host user first so we're not in first-user setup mode
|
||||
_, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
|
||||
// User registration is enabled by default
|
||||
|
||||
// Unauthenticated user tries to create admin user - role should be ignored
|
||||
createdUser, err := ts.Service.CreateUser(ctx, &apiv1.CreateUserRequest{
|
||||
User: &apiv1.User{
|
||||
Username: "wannabeadmin",
|
||||
Email: "wannabeadmin@example.com",
|
||||
Password: "password123",
|
||||
Role: apiv1.User_ADMIN, // This should be ignored
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, createdUser)
|
||||
require.Equal(t, "users/wannabeadmin", createdUser.Name)
|
||||
require.Equal(t, apiv1.User_USER, createdUser.Role, "Unauthenticated users can only create USER role")
|
||||
})
|
||||
|
||||
t.Run("CreateUser blocked when password auth disabled for self signup", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
_, err := ts.CreateHostUser(ctx, "admin")
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Store.UpsertInstanceSetting(ctx, &storepb.InstanceSetting{
|
||||
Key: storepb.InstanceSettingKey_GENERAL,
|
||||
Value: &storepb.InstanceSetting_GeneralSetting{
|
||||
GeneralSetting: &storepb.InstanceGeneralSetting{
|
||||
DisallowPasswordAuth: true,
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.CreateUser(ctx, &apiv1.CreateUserRequest{
|
||||
User: &apiv1.User{
|
||||
Username: "newuser",
|
||||
Email: "newuser@example.com",
|
||||
Password: "password123",
|
||||
},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "password signup is not allowed")
|
||||
})
|
||||
|
||||
t.Run("CreateUser rejects empty password", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
_, err := ts.Service.CreateUser(ctx, &apiv1.CreateUserRequest{
|
||||
User: &apiv1.User{
|
||||
Username: "newuser",
|
||||
Email: "newuser@example.com",
|
||||
Password: "",
|
||||
},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "password must not be empty")
|
||||
})
|
||||
|
||||
t.Run("CreateUser rejects invalid writable usernames", func(t *testing.T) {
|
||||
for _, username := range []string{"alice@example.com", "legacy_user", "123"} {
|
||||
t.Run(username, func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
_, err := ts.Service.CreateUser(ctx, &apiv1.CreateUserRequest{
|
||||
User: &apiv1.User{
|
||||
Username: username,
|
||||
Email: "newuser@example.com",
|
||||
Password: "password123",
|
||||
},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "invalid username")
|
||||
})
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("CreateUser validate only rejects invalid writable username", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
_, err := ts.Service.CreateUser(ctx, &apiv1.CreateUserRequest{
|
||||
User: &apiv1.User{
|
||||
Username: "alice@example.com",
|
||||
Email: "newuser@example.com",
|
||||
Password: "password123",
|
||||
},
|
||||
ValidateOnly: true,
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "invalid username")
|
||||
})
|
||||
|
||||
t.Run("UpdateUser rejects empty password", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "alice")
|
||||
require.NoError(t, err)
|
||||
|
||||
authCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
_, err = ts.Service.UpdateUser(authCtx, &apiv1.UpdateUserRequest{
|
||||
User: &apiv1.User{
|
||||
Name: apiv1server.BuildUserName(user.Username),
|
||||
Password: "",
|
||||
},
|
||||
UpdateMask: &fieldmaskpb.FieldMask{Paths: []string{"password"}},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "password must not be empty")
|
||||
})
|
||||
|
||||
t.Run("UpdateUser rejects missing user message", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "missing-message")
|
||||
require.NoError(t, err)
|
||||
|
||||
authCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
_, err = ts.Service.UpdateUser(authCtx, &apiv1.UpdateUserRequest{
|
||||
UpdateMask: &fieldmaskpb.FieldMask{Paths: []string{"display_name"}},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "user is required")
|
||||
})
|
||||
|
||||
t.Run("CreateUser concurrent first setup creates one admin", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
const workers = 12
|
||||
var wg sync.WaitGroup
|
||||
for i := range workers {
|
||||
wg.Go(func() {
|
||||
_, _ = ts.Service.CreateUser(ctx, &apiv1.CreateUserRequest{
|
||||
User: &apiv1.User{
|
||||
Username: fmt.Sprintf("setup-user-%d", i),
|
||||
Email: "setup-user@example.com",
|
||||
Password: "password123",
|
||||
},
|
||||
})
|
||||
})
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
users, err := ts.Store.ListUsers(ctx, &store.FindUser{})
|
||||
require.NoError(t, err)
|
||||
adminCount := 0
|
||||
for _, user := range users {
|
||||
if user.Role == store.RoleAdmin {
|
||||
adminCount++
|
||||
}
|
||||
}
|
||||
require.Equal(t, 1, adminCount)
|
||||
})
|
||||
|
||||
t.Run("UpdateUser state requires admin", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "state-user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
_, err = ts.Service.UpdateUser(userCtx, &apiv1.UpdateUserRequest{
|
||||
User: &apiv1.User{
|
||||
Name: apiv1server.BuildUserName(user.Username),
|
||||
State: apiv1.State_ARCHIVED,
|
||||
},
|
||||
UpdateMask: &fieldmaskpb.FieldMask{Paths: []string{"state"}},
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "permission denied")
|
||||
|
||||
admin, err := ts.CreateHostUser(ctx, "state-admin")
|
||||
require.NoError(t, err)
|
||||
adminCtx := ts.CreateUserContext(ctx, admin.ID)
|
||||
updated, err := ts.Service.UpdateUser(adminCtx, &apiv1.UpdateUserRequest{
|
||||
User: &apiv1.User{
|
||||
Name: apiv1server.BuildUserName(user.Username),
|
||||
State: apiv1.State_ARCHIVED,
|
||||
},
|
||||
UpdateMask: &fieldmaskpb.FieldMask{Paths: []string{"state"}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, apiv1.State_ARCHIVED, updated.State)
|
||||
})
|
||||
|
||||
t.Run("archived user context is rejected", func(t *testing.T) {
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "archived-access-user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
archived := store.Archived
|
||||
_, err = ts.Store.UpdateUser(ctx, &store.UpdateUser{
|
||||
ID: user.ID,
|
||||
RowStatus: &archived,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Service.GetCurrentUser(userCtx, &apiv1.GetCurrentUserRequest{})
|
||||
require.Error(t, err)
|
||||
|
||||
_, err = ts.Service.CreateMemo(userCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "should not be created",
|
||||
},
|
||||
})
|
||||
require.Error(t, err)
|
||||
|
||||
_, err = ts.Service.UpdateUser(userCtx, &apiv1.UpdateUserRequest{
|
||||
User: &apiv1.User{
|
||||
Name: apiv1server.BuildUserName(user.Username),
|
||||
State: apiv1.State_NORMAL,
|
||||
},
|
||||
UpdateMask: &fieldmaskpb.FieldMask{Paths: []string{"state"}},
|
||||
})
|
||||
require.Error(t, err)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,202 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func TestGetUserStats_TagCount(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
// Create test service
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
// Create a test host user
|
||||
user, err := ts.CreateHostUser(ctx, "test-user")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create user context for authentication
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
// Create a memo with a single tag
|
||||
memo, err := ts.Store.CreateMemo(ctx, &store.Memo{
|
||||
UID: "test-memo-1",
|
||||
CreatorID: user.ID,
|
||||
Content: "This is a test memo with #test tag",
|
||||
Visibility: store.Public,
|
||||
Payload: &storepb.MemoPayload{
|
||||
Tags: []string{"test"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memo)
|
||||
|
||||
// Test GetUserStats
|
||||
userName := fmt.Sprintf("users/%s", user.Username)
|
||||
response, err := ts.Service.GetUserStats(userCtx, &v1pb.GetUserStatsRequest{
|
||||
Name: userName,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, response)
|
||||
require.Equal(t, fmt.Sprintf("users/%s/stats", user.Username), response.Name)
|
||||
|
||||
// Check that the tag count is exactly 1, not 2
|
||||
require.Contains(t, response.TagCount, "test")
|
||||
require.Equal(t, int32(1), response.TagCount["test"], "Tag count should be 1 for a single occurrence")
|
||||
|
||||
// Create another memo with the same tag
|
||||
memo2, err := ts.Store.CreateMemo(ctx, &store.Memo{
|
||||
UID: "test-memo-2",
|
||||
CreatorID: user.ID,
|
||||
Content: "Another memo with #test tag",
|
||||
Visibility: store.Public,
|
||||
Payload: &storepb.MemoPayload{
|
||||
Tags: []string{"test"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memo2)
|
||||
|
||||
// Test GetUserStats again
|
||||
response2, err := ts.Service.GetUserStats(userCtx, &v1pb.GetUserStatsRequest{
|
||||
Name: userName,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, response2)
|
||||
|
||||
// Check that the tag count is exactly 2, not 3
|
||||
require.Contains(t, response2.TagCount, "test")
|
||||
require.Equal(t, int32(2), response2.TagCount["test"], "Tag count should be 2 for two occurrences")
|
||||
|
||||
// Test with a new unique tag
|
||||
memo3, err := ts.Store.CreateMemo(ctx, &store.Memo{
|
||||
UID: "test-memo-3",
|
||||
CreatorID: user.ID,
|
||||
Content: "Memo with #unique tag",
|
||||
Visibility: store.Public,
|
||||
Payload: &storepb.MemoPayload{
|
||||
Tags: []string{"unique"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memo3)
|
||||
|
||||
// Test GetUserStats for the new tag
|
||||
response3, err := ts.Service.GetUserStats(userCtx, &v1pb.GetUserStatsRequest{
|
||||
Name: userName,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, response3)
|
||||
|
||||
// Check that the unique tag count is exactly 1
|
||||
require.Contains(t, response3.TagCount, "unique")
|
||||
require.Equal(t, int32(1), response3.TagCount["unique"], "New tag count should be 1 for first occurrence")
|
||||
|
||||
// The original test tag should still be 2
|
||||
require.Contains(t, response3.TagCount, "test")
|
||||
require.Equal(t, int32(2), response3.TagCount["test"], "Original tag count should remain 2")
|
||||
|
||||
_, err = ts.Service.GetUserStats(userCtx, &v1pb.GetUserStatsRequest{
|
||||
Name: "users/1",
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "user not found")
|
||||
}
|
||||
|
||||
func TestGetUserStats_MemoUpdatedTimestamps(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateHostUser(ctx, "ts-user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
memo, err := ts.Store.CreateMemo(ctx, &store.Memo{
|
||||
UID: "ts-memo-1",
|
||||
CreatorID: user.ID,
|
||||
Content: "first content",
|
||||
Visibility: store.Public,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, memo)
|
||||
|
||||
// SQLite UpdateMemo only sets fields explicitly passed (created_ts default
|
||||
// fires on INSERT only). So bump updated_ts explicitly to simulate an edit
|
||||
// happening after creation.
|
||||
newContent := "second content"
|
||||
newUpdatedTs := memo.UpdatedTs + 100
|
||||
require.NoError(t, ts.Store.UpdateMemo(ctx, &store.UpdateMemo{
|
||||
ID: memo.ID,
|
||||
Content: &newContent,
|
||||
UpdatedTs: &newUpdatedTs,
|
||||
}))
|
||||
|
||||
userName := fmt.Sprintf("users/%s", user.Username)
|
||||
resp, err := ts.Service.GetUserStats(userCtx, &v1pb.GetUserStatsRequest{Name: userName})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp)
|
||||
|
||||
require.Len(t, resp.MemoCreatedTimestamps, 1, "should have one created timestamp")
|
||||
require.Len(t, resp.MemoUpdatedTimestamps, 1, "should have one updated timestamp")
|
||||
|
||||
require.Equal(t, memo.CreatedTs, resp.MemoCreatedTimestamps[0].AsTime().Unix())
|
||||
require.Equal(t, newUpdatedTs, resp.MemoUpdatedTimestamps[0].AsTime().Unix())
|
||||
require.Greater(
|
||||
t,
|
||||
resp.MemoUpdatedTimestamps[0].AsTime().Unix(),
|
||||
resp.MemoCreatedTimestamps[0].AsTime().Unix(),
|
||||
"updated_ts should be after created_ts after an edit",
|
||||
)
|
||||
}
|
||||
|
||||
func TestListAllUserStats_FilterExcludesPrivateMemos(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateHostUser(ctx, "stats-filter-user")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
_, err = ts.Store.CreateMemo(ctx, &store.Memo{
|
||||
UID: "stats-filter-public",
|
||||
CreatorID: user.ID,
|
||||
Content: "public memo",
|
||||
Visibility: store.Public,
|
||||
Payload: &storepb.MemoPayload{Tags: []string{"public"}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
_, err = ts.Store.CreateMemo(ctx, &store.Memo{
|
||||
UID: "stats-filter-private",
|
||||
CreatorID: user.ID,
|
||||
Content: "private memo",
|
||||
Visibility: store.Private,
|
||||
Payload: &storepb.MemoPayload{Tags: []string{"private"}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
unfilteredResp, err := ts.Service.ListAllUserStats(userCtx, &v1pb.ListAllUserStatsRequest{})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, unfilteredResp.Stats, 1)
|
||||
require.Equal(t, int32(1), unfilteredResp.Stats[0].TagCount["public"])
|
||||
require.Equal(t, int32(1), unfilteredResp.Stats[0].TagCount["private"])
|
||||
|
||||
filteredResp, err := ts.Service.ListAllUserStats(userCtx, &v1pb.ListAllUserStatsRequest{
|
||||
Filter: `visibility in ["PUBLIC", "PROTECTED"]`,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, filteredResp.Stats, 1)
|
||||
require.Equal(t, int32(1), filteredResp.Stats[0].TagCount["public"])
|
||||
require.NotContains(t, filteredResp.Stats[0].TagCount, "private")
|
||||
}
|
||||
@@ -0,0 +1,58 @@
|
||||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
apiv1 "github.com/usememos/memos/proto/gen/api/v1"
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
apiv1server "github.com/usememos/memos/server/router/api/v1"
|
||||
)
|
||||
|
||||
func TestListUserSettingsOmitsInternalStoreSettings(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "locale-user")
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Store.UpsertUserSetting(ctx, &storepb.UserSetting{
|
||||
UserId: user.ID,
|
||||
Key: storepb.UserSetting_REFRESH_TOKENS,
|
||||
Value: &storepb.UserSetting_RefreshTokens{
|
||||
RefreshTokens: &storepb.RefreshTokensUserSetting{},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Store.UpsertUserSetting(ctx, &storepb.UserSetting{
|
||||
UserId: user.ID,
|
||||
Key: storepb.UserSetting_SHORTCUTS,
|
||||
Value: &storepb.UserSetting_Shortcuts{
|
||||
Shortcuts: &storepb.ShortcutsUserSetting{},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.Store.UpsertUserSetting(ctx, &storepb.UserSetting{
|
||||
UserId: user.ID,
|
||||
Key: storepb.UserSetting_GENERAL,
|
||||
Value: &storepb.UserSetting_General{
|
||||
General: &storepb.GeneralUserSetting{
|
||||
Locale: "ja",
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.ListUserSettings(ts.CreateUserContext(ctx, user.ID), &apiv1.ListUserSettingsRequest{
|
||||
Parent: apiv1server.BuildUserName(user.Username),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, resp.Settings, 1)
|
||||
require.Equal(t, "users/locale-user/settings/GENERAL", resp.Settings[0].Name)
|
||||
require.Equal(t, "ja", resp.Settings[0].GetGeneralSetting().Locale)
|
||||
}
|
||||
Reference in New Issue
Block a user